Skip to content

Commit cc384a5

Browse files
fix: normalize signed zero in nested IN-list values
1 parent ba584cd commit cc384a5

5 files changed

Lines changed: 190 additions & 19 deletions

File tree

‎datafusion/common/src/utils/mod.rs‎

Lines changed: 107 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1544,9 +1544,24 @@ pub fn has_float_leaf(data_type: &DataType) -> bool {
15441544
}
15451545

15461546
/// Replace `-0.0` with `+0.0` in `Float16`, `Float32`, or `Float64` scalar
1547-
/// values, including dictionary-wrapped floats. Other variants are returned
1548-
/// unchanged. See [`normalize_float_zero`] for context.
1547+
/// values, including floats inside nested and encoded values. Other variants
1548+
/// are returned unchanged. See [`normalize_float_zero`] for context.
15491549
pub fn normalize_float_zero_scalar(scalar: ScalarValue) -> ScalarValue {
1550+
fn normalize_array<A: Array + Clone + 'static>(array: Arc<A>) -> Arc<A> {
1551+
let original = Arc::clone(&array) as ArrayRef;
1552+
let normalized = normalize_float_zero(&original);
1553+
if Arc::ptr_eq(&original, &normalized) {
1554+
// Preserve the original scalar array when no values changed.
1555+
return array;
1556+
}
1557+
1558+
let normalized_array = normalized
1559+
.as_any()
1560+
.downcast_ref::<A>()
1561+
.expect("float-zero normalization preserves the array type");
1562+
Arc::new(normalized_array.clone())
1563+
}
1564+
15501565
match scalar {
15511566
ScalarValue::Float32(Some(v)) if v.to_bits() << 1 == 0 => {
15521567
ScalarValue::Float32(Some(0.0))
@@ -1557,10 +1572,29 @@ pub fn normalize_float_zero_scalar(scalar: ScalarValue) -> ScalarValue {
15571572
ScalarValue::Float16(Some(v)) if v.to_bits() << 1 == 0 => {
15581573
ScalarValue::Float16(Some(half::f16::from_bits(0)))
15591574
}
1575+
ScalarValue::FixedSizeList(array) => {
1576+
ScalarValue::FixedSizeList(normalize_array(array))
1577+
}
1578+
ScalarValue::List(array) => ScalarValue::List(normalize_array(array)),
1579+
ScalarValue::LargeList(array) => ScalarValue::LargeList(normalize_array(array)),
1580+
ScalarValue::ListView(array) => ScalarValue::ListView(normalize_array(array)),
1581+
ScalarValue::LargeListView(array) => {
1582+
ScalarValue::LargeListView(normalize_array(array))
1583+
}
1584+
ScalarValue::Struct(array) => ScalarValue::Struct(normalize_array(array)),
1585+
ScalarValue::Map(array) => ScalarValue::Map(normalize_array(array)),
1586+
ScalarValue::Union(Some((type_id, mut value)), fields, mode) => {
1587+
*value = normalize_float_zero_scalar(*value);
1588+
ScalarValue::Union(Some((type_id, value)), fields, mode)
1589+
}
15601590
ScalarValue::Dictionary(key, mut value) => {
15611591
*value = normalize_float_zero_scalar(*value);
15621592
ScalarValue::Dictionary(key, value)
15631593
}
1594+
ScalarValue::RunEndEncoded(run_ends, values, mut value) => {
1595+
*value = normalize_float_zero_scalar(*value);
1596+
ScalarValue::RunEndEncoded(run_ends, values, value)
1597+
}
15641598
other => other,
15651599
}
15661600
}
@@ -1723,6 +1757,77 @@ mod tests {
17231757
Ok(())
17241758
}
17251759

1760+
#[test]
1761+
fn normalize_float_zero_scalar_preserves_union_and_run_metadata() -> Result<()> {
1762+
use arrow::datatypes::{UnionFields, UnionMode};
1763+
1764+
let run_ends = Arc::new(Field::new("ends", DataType::Int32, false));
1765+
let values = Arc::new(Field::new("samples", DataType::Float64, true));
1766+
let fields = UnionFields::try_new(
1767+
[7, 42],
1768+
[
1769+
Field::new(
1770+
"runs",
1771+
DataType::RunEndEncoded(Arc::clone(&run_ends), Arc::clone(&values)),
1772+
true,
1773+
),
1774+
Field::new("other", DataType::Int32, true),
1775+
],
1776+
)?;
1777+
let nan = f64::from_bits(0x7ff8_0000_0000_0001);
1778+
for mode in [UnionMode::Sparse, UnionMode::Dense] {
1779+
let wrap = |value| {
1780+
ScalarValue::Union(
1781+
Some((
1782+
7,
1783+
Box::new(ScalarValue::RunEndEncoded(
1784+
Arc::clone(&run_ends),
1785+
Arc::clone(&values),
1786+
Box::new(ScalarValue::Float64(value)),
1787+
)),
1788+
)),
1789+
fields.clone(),
1790+
mode,
1791+
)
1792+
};
1793+
for (value, expected) in [
1794+
(Some(-0.0), Some(0.0)),
1795+
(Some(nan), Some(nan)),
1796+
(None, None),
1797+
] {
1798+
assert_eq!(normalize_float_zero_scalar(wrap(value)), wrap(expected));
1799+
}
1800+
let null = ScalarValue::Union(None, fields.clone(), mode);
1801+
assert_eq!(normalize_float_zero_scalar(null.clone()), null);
1802+
}
1803+
Ok(())
1804+
}
1805+
1806+
#[test]
1807+
fn normalize_float_zero_scalar_preserves_sliced_list() {
1808+
let nan = f64::from_bits(0x7ff8_0000_0000_0001);
1809+
let list = Arc::new(
1810+
ListArray::from_iter_primitive::<Float64Type, _, _>([
1811+
Some(vec![Some(99.0)]),
1812+
Some(vec![Some(-0.0), Some(nan), None]),
1813+
Some(vec![Some(42.0)]),
1814+
])
1815+
.slice(1, 1),
1816+
);
1817+
let ScalarValue::List(normalized) =
1818+
normalize_float_zero_scalar(ScalarValue::List(list))
1819+
else {
1820+
panic!("normalization must preserve the scalar variant");
1821+
};
1822+
assert_eq!(normalized.len(), 1);
1823+
let child = normalized.value(0);
1824+
let child = child.as_primitive::<Float64Type>();
1825+
assert_eq!(child.len(), 3);
1826+
assert_eq!(child.value(0).to_bits(), 0.0_f64.to_bits());
1827+
assert_eq!(child.value(1).to_bits(), nan.to_bits());
1828+
assert!(child.is_null(2));
1829+
}
1830+
17261831
#[test]
17271832
fn test_bisect_linear_left_and_right() -> Result<()> {
17281833
let arrays: Vec<ArrayRef> = vec![

‎datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs‎

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ use std::sync::LazyLock;
3030

3131
use datafusion_common::config::ConfigOptions;
3232
use datafusion_common::nested_struct::has_one_of_more_common_fields;
33+
use datafusion_common::utils::has_float_leaf;
3334
use datafusion_common::{
3435
DFSchema, DataFusionError, Result, ScalarValue, exec_datafusion_err, internal_err,
3536
};
@@ -2318,8 +2319,9 @@ fn are_inlist_and_eq_and_match_neg(
23182319
fn inlists_have_set_comparable_literals(left: &Expr, right: &Expr) -> bool {
23192320
match (left, right) {
23202321
(Expr::InList(l), Expr::InList(r)) => l.list.iter().chain(&r.list).all(|item| {
2321-
item.as_literal()
2322-
.is_some_and(|value| !value.is_null() && !value.data_type().is_floating())
2322+
item.as_literal().is_some_and(|value| {
2323+
!value.is_null() && !has_float_leaf(&value.data_type())
2324+
})
23232325
}),
23242326
_ => false,
23252327
}

‎datafusion/physical-expr/src/expressions/in_list.rs‎

Lines changed: 2 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -48,7 +48,7 @@ mod static_filter;
4848
mod strategy;
4949

5050
use static_filter::StaticFilterRef;
51-
use strategy::{dictionary_value_type, instantiate_static_filter};
51+
use strategy::instantiate_static_filter;
5252

5353
/// InList
5454
pub struct InListExpr {
@@ -85,15 +85,10 @@ fn supports_arrow_eq(dt: &DataType) -> bool {
8585

8686
fn normalize_in_list_float_zero_value(value: ColumnarValue) -> ColumnarValue {
8787
match value {
88-
ColumnarValue::Array(array)
89-
if dictionary_value_type(array.data_type()).is_floating() =>
90-
{
91-
ColumnarValue::Array(normalize_float_zero(&array))
92-
}
88+
ColumnarValue::Array(array) => ColumnarValue::Array(normalize_float_zero(&array)),
9389
ColumnarValue::Scalar(scalar) => {
9490
ColumnarValue::Scalar(normalize_float_zero_scalar(scalar))
9591
}
96-
value => value,
9792
}
9893
}
9994

‎datafusion/physical-expr/src/expressions/in_list/array_static_filter.rs‎

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -15,13 +15,14 @@
1515
// specific language governing permissions and limitations
1616
// under the License.
1717

18-
use arrow::array::{Array, ArrayRef, BooleanArray, make_comparator};
18+
use arrow::array::{Array, ArrayRef, BooleanArray, make_array, make_comparator};
1919
use arrow::buffer::{BooleanBuffer, NullBuffer};
2020
use arrow::compute::SortOptions;
2121
use arrow::datatypes::DataType;
2222
use arrow::util::bit_iterator::BitIndexIterator;
2323
use datafusion_common::Result;
2424
use datafusion_common::hash_utils::{RandomState, with_hashes};
25+
use datafusion_common::utils::{has_float_leaf, normalize_float_zero};
2526
use hashbrown::HashTable;
2627

2728
use super::result::build_in_list_result;
@@ -53,6 +54,8 @@ impl ArrayStaticFilter {
5354
});
5455
}
5556

57+
// Hashing treats both signed zeros alike; the comparator must do so too.
58+
let in_array = normalize_float_zero(&in_array);
5659
let state = RandomState::default();
5760
let table = Self::build_haystack_table(&in_array, &state)?;
5861

@@ -138,6 +141,11 @@ impl StaticFilter for ArrayStaticFilter {
138141
));
139142
}
140143

141-
self.find_needles_in_haystack(v, negated)
144+
if has_float_leaf(v.data_type()) {
145+
let normalized = normalize_float_zero(&make_array(v.to_data()));
146+
self.find_needles_in_haystack(normalized.as_ref(), negated)
147+
} else {
148+
self.find_needles_in_haystack(v, negated)
149+
}
142150
}
143151
}

‎datafusion/sqllogictest/test_files/negative_zero.slt‎

Lines changed: 67 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -290,13 +290,74 @@ true false true
290290
statement ok
291291
CREATE TABLE nested_zeros(id INT, a DOUBLE[]) AS VALUES (1, [0.0]), (2, [-0.0]);
292292

293-
query II
294-
SELECT l.id, r.id FROM nested_zeros l JOIN nested_zeros r ON l.a = r.a ORDER BY l.id, r.id;
293+
# Column-valued list items exercise dynamic array comparisons in both directions.
294+
query IIBB
295+
SELECT l.id, r.id,
296+
l.a IN (r.a, [1.0], [2.0], [3.0]),
297+
l.a NOT IN (r.a, [1.0], [2.0], [3.0])
298+
FROM nested_zeros l JOIN nested_zeros r ON l.a = r.a
299+
ORDER BY l.id, r.id;
295300
----
296-
1 1
297-
1 2
298-
2 1
299-
2 2
301+
1 1 true false
302+
1 2 true false
303+
2 1 true false
304+
2 2 true false
305+
306+
# Short lists become equality comparisons; four-item constant lists use a static filter.
307+
query IBBBB
308+
SELECT id,
309+
a IN ([0.0]),
310+
a IN ([0.0], [1.0], [2.0], [3.0]),
311+
a IN ([-0.0], [1.0], [2.0], [3.0]),
312+
a NOT IN ([-0.0], [1.0], [2.0], [3.0])
313+
FROM nested_zeros ORDER BY id;
314+
----
315+
1 true true true false
316+
2 true true true false
317+
318+
# The nonmatching column-derived item keeps scalar RHS comparisons dynamic.
319+
# Also check scalar LHS values against array-valued list items.
320+
query IBBBB
321+
SELECT id,
322+
a IN ([-0.0], [1.0], [2.0], [CAST(id AS DOUBLE)]),
323+
a NOT IN ([0.0], [1.0], [2.0], [CAST(id AS DOUBLE)]),
324+
[-0.0] IN (a, [1.0], [2.0], [3.0]),
325+
[0.0] NOT IN (a, [1.0], [2.0], [3.0])
326+
FROM nested_zeros ORDER BY id;
327+
----
328+
1 true false true false
329+
2 true false true false
330+
331+
# Structural literal equality must not drive IN intersection or difference for nested floats.
332+
query IBB
333+
SELECT id,
334+
a IN ([0.0], [1.0], [2.0], [3.0])
335+
AND a IN ([-0.0], [4.0], [5.0], [6.0]),
336+
a IN ([0.0], [1.0], [2.0], [3.0])
337+
AND a NOT IN ([-0.0], [4.0], [5.0], [6.0])
338+
FROM nested_zeros ORDER BY id;
339+
----
340+
1 true false
341+
2 true false
342+
343+
statement ok
344+
INSERT INTO nested_zeros VALUES (3, NULL), (4, [NULL]), (5, [9.0]);
345+
346+
# Matches win over an outer NULL list item; nested NULL elements remain comparable.
347+
# A missing match or a NULL input produces NULL, for both static and dynamic lists.
348+
query IBBBB
349+
SELECT id,
350+
a IN ([-0.0], [NULL], [1.0], NULL),
351+
a NOT IN ([-0.0], [NULL], [1.0], NULL),
352+
a IN ([-0.0], [NULL], [CAST(id AS DOUBLE)], NULL),
353+
a NOT IN ([-0.0], [NULL], [CAST(id AS DOUBLE)], NULL)
354+
FROM nested_zeros ORDER BY id;
355+
----
356+
1 true false true false
357+
2 true false true false
358+
3 NULL NULL NULL NULL
359+
4 true false true false
360+
5 NULL NULL NULL NULL
300361

301362
statement ok
302363
DROP TABLE nested_zeros;

0 commit comments

Comments
 (0)