Skip to content

Commit aa5a037

Browse files
Treat signed zeros as equal in IN LIST filters
1 parent 224cc56 commit aa5a037

5 files changed

Lines changed: 566 additions & 53 deletions

File tree

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

Lines changed: 102 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -1407,9 +1407,9 @@ fn fsl_values_row_number(list_size: i32, array_len: usize) -> Result<Int32Array>
14071407
Ok(PrimitiveArray::new(rows_number.into(), None))
14081408
}
14091409

1410-
/// Replace `-0.0` with `+0.0` in any `Float16`, `Float32`, or `Float64` array.
1411-
/// For non-float arrays returns the input unchanged. NaN payloads are
1412-
/// preserved.
1410+
/// Replace `-0.0` with `+0.0` in any `Float16`, `Float32`, or `Float64` array,
1411+
/// including dictionary-wrapped floats. For other arrays, returns the input
1412+
/// unchanged. NaN payloads are preserved.
14131413
///
14141414
/// Arrow's comparison kernels (`arrow::compute::kernels::cmp::eq` etc.) and
14151415
/// row-encoding (`arrow::row::RowConverter`) use IEEE 754 totalOrder
@@ -1430,6 +1430,17 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14301430
const NEG_ZERO_F32_BITS: u32 = (-0.0_f32).to_bits();
14311431
const NEG_ZERO_F64_BITS: u64 = (-0.0_f64).to_bits();
14321432
match array.data_type() {
1433+
DataType::Dictionary(_, value_type)
1434+
if is_float_or_dictionary_float(value_type) =>
1435+
{
1436+
let dictionary = array.as_any_dictionary();
1437+
let values = normalize_float_zero(dictionary.values());
1438+
if Arc::ptr_eq(&values, dictionary.values()) {
1439+
Arc::clone(array)
1440+
} else {
1441+
dictionary.with_values(values)
1442+
}
1443+
}
14331444
DataType::Float32 => {
14341445
let arr: &Float32Array = array.as_primitive::<Float32Type>();
14351446
if !arr
@@ -1439,8 +1450,13 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14391450
{
14401451
return Arc::clone(array);
14411452
}
1442-
let normalized: Float32Array =
1443-
arr.unary(|v| if v.to_bits() << 1 == 0 { 0.0_f32 } else { v });
1453+
let normalized: Float32Array = arr.unary(|v| {
1454+
if v.to_bits() == NEG_ZERO_F32_BITS {
1455+
0.0_f32
1456+
} else {
1457+
v
1458+
}
1459+
});
14441460
Arc::new(normalized)
14451461
}
14461462
DataType::Float64 => {
@@ -1452,8 +1468,13 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14521468
{
14531469
return Arc::clone(array);
14541470
}
1455-
let normalized: Float64Array =
1456-
arr.unary(|v| if v.to_bits() << 1 == 0 { 0.0_f64 } else { v });
1471+
let normalized: Float64Array = arr.unary(|v| {
1472+
if v.to_bits() == NEG_ZERO_F64_BITS {
1473+
0.0_f64
1474+
} else {
1475+
v
1476+
}
1477+
});
14571478
Arc::new(normalized)
14581479
}
14591480
DataType::Float16 => {
@@ -1466,8 +1487,8 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14661487
return Arc::clone(array);
14671488
}
14681489
let normalized: Float16Array = arr.unary(|v| {
1469-
if v.to_bits() << 1 == 0 {
1470-
half::f16::from_bits(0)
1490+
if v.to_bits() == NEG_ZERO_F16_BITS {
1491+
half::f16::ZERO
14711492
} else {
14721493
v
14731494
}
@@ -1478,22 +1499,38 @@ pub fn normalize_float_zero(array: &ArrayRef) -> ArrayRef {
14781499
}
14791500
}
14801501

1502+
fn is_float_or_dictionary_float(mut data_type: &DataType) -> bool {
1503+
while let DataType::Dictionary(_, value_type) = data_type {
1504+
data_type = value_type;
1505+
}
1506+
data_type.is_floating()
1507+
}
1508+
14811509
/// Replace `-0.0` with `+0.0` in `Float16`, `Float32`, or `Float64` scalar
1482-
/// values. Other variants are returned unchanged. See [`normalize_float_zero`]
1483-
/// for context.
1484-
pub fn normalize_float_zero_scalar(scalar: ScalarValue) -> ScalarValue {
1485-
match scalar {
1486-
ScalarValue::Float32(Some(v)) if v.to_bits() << 1 == 0 => {
1487-
ScalarValue::Float32(Some(0.0))
1510+
/// values, including dictionary-wrapped floats. Other variants are returned
1511+
/// unchanged. See [`normalize_float_zero`] for context.
1512+
pub fn normalize_float_zero_scalar(mut scalar: ScalarValue) -> ScalarValue {
1513+
let mut value = &mut scalar;
1514+
while let ScalarValue::Dictionary(_, dictionary_value) = value {
1515+
value = dictionary_value.as_mut();
1516+
}
1517+
1518+
match value {
1519+
ScalarValue::Float32(Some(value)) if value.to_bits() == (-0.0_f32).to_bits() => {
1520+
*value = 0.0
14881521
}
1489-
ScalarValue::Float64(Some(v)) if v.to_bits() << 1 == 0 => {
1490-
ScalarValue::Float64(Some(0.0))
1522+
ScalarValue::Float64(Some(value)) if value.to_bits() == (-0.0_f64).to_bits() => {
1523+
*value = 0.0
14911524
}
1492-
ScalarValue::Float16(Some(v)) if v.to_bits() << 1 == 0 => {
1493-
ScalarValue::Float16(Some(half::f16::from_bits(0)))
1525+
ScalarValue::Float16(Some(value))
1526+
if value.to_bits() == half::f16::NEG_ZERO.to_bits() =>
1527+
{
1528+
*value = half::f16::ZERO;
14941529
}
1495-
other => other,
1530+
_ => {}
14961531
}
1532+
1533+
scalar
14971534
}
14981535

14991536
#[cfg(test)]
@@ -1503,9 +1540,9 @@ mod tests {
15031540
use super::*;
15041541
use crate::ScalarValue::Null;
15051542
use arrow::{
1506-
array::{Float64Array, Int32Array},
1543+
array::{DictionaryArray, Float64Array, Int8Array, Int32Array},
15071544
buffer::NullBuffer,
1508-
datatypes::Int32Type,
1545+
datatypes::{Float64Type, Int8Type, Int32Type},
15091546
};
15101547
#[cfg(feature = "sql")]
15111548
use sqlparser::ast::Ident;
@@ -1534,6 +1571,49 @@ mod tests {
15341571
}
15351572
}
15361573

1574+
#[test]
1575+
fn normalize_float_zero_in_dictionary_values() -> Result<()> {
1576+
let nan = f64::from_bits(0x7ff8_0000_0000_0001);
1577+
let keys = Int8Array::from(vec![Some(0), Some(1), None, Some(2)]);
1578+
let array: ArrayRef = Arc::new(DictionaryArray::try_new(
1579+
keys.clone(),
1580+
Arc::new(Float64Array::from(vec![-0.0, nan, 1.0])),
1581+
)?);
1582+
1583+
let normalized = normalize_float_zero(&array);
1584+
assert!(!Arc::ptr_eq(&normalized, &array));
1585+
let dictionary = normalized.as_dictionary::<Int8Type>();
1586+
assert_eq!(dictionary.keys(), &keys);
1587+
let values = dictionary.values().as_primitive::<Float64Type>();
1588+
assert_eq!(values.value(0).to_bits(), 0.0_f64.to_bits());
1589+
assert_eq!(values.value(1).to_bits(), nan.to_bits());
1590+
assert_eq!(values.value(2), 1.0);
1591+
1592+
let without_negative_zero: ArrayRef = Arc::new(DictionaryArray::try_new(
1593+
Int8Array::from(vec![0, 1]),
1594+
Arc::new(Float64Array::from(vec![0.0, nan])),
1595+
)?);
1596+
assert!(Arc::ptr_eq(
1597+
&normalize_float_zero(&without_negative_zero),
1598+
&without_negative_zero
1599+
));
1600+
1601+
let scalar = ScalarValue::Dictionary(
1602+
Box::new(DataType::Int8),
1603+
Box::new(ScalarValue::Float64(Some(-0.0))),
1604+
);
1605+
let ScalarValue::Dictionary(_, value) = normalize_float_zero_scalar(scalar)
1606+
else {
1607+
unreachable!()
1608+
};
1609+
let ScalarValue::Float64(Some(value)) = *value else {
1610+
unreachable!()
1611+
};
1612+
assert_eq!(value.to_bits(), 0.0_f64.to_bits());
1613+
1614+
Ok(())
1615+
}
1616+
15371617
#[test]
15381618
fn test_bisect_linear_left_and_right() -> Result<()> {
15391619
let arrays: Vec<ArrayRef> = vec![

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

Lines changed: 155 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ use arrow::compute::kernels::boolean::{not, or_kleene};
3131
use arrow::compute::kernels::cmp::eq as arrow_eq;
3232
use arrow::datatypes::*;
3333

34+
use datafusion_common::utils::{normalize_float_zero, normalize_float_zero_scalar};
3435
use datafusion_common::{
3536
DFSchema, Result, ScalarValue, assert_or_internal_err, exec_err,
3637
};
@@ -81,6 +82,27 @@ fn supports_arrow_eq(dt: &DataType) -> bool {
8182
}
8283
}
8384

85+
fn normalize_in_list_float_zero_value(value: ColumnarValue) -> ColumnarValue {
86+
match value {
87+
ColumnarValue::Array(array)
88+
if is_float_or_dictionary_float(array.data_type()) =>
89+
{
90+
ColumnarValue::Array(normalize_float_zero(&array))
91+
}
92+
ColumnarValue::Scalar(scalar) => {
93+
ColumnarValue::Scalar(normalize_float_zero_scalar(scalar))
94+
}
95+
value => value,
96+
}
97+
}
98+
99+
fn is_float_or_dictionary_float(mut data_type: &DataType) -> bool {
100+
while let DataType::Dictionary(_, value_type) = data_type {
101+
data_type = value_type;
102+
}
103+
data_type.is_floating()
104+
}
105+
84106
/// Evaluates the list of expressions into an array, flattening any dictionaries
85107
fn evaluate_list(
86108
list: &[Arc<dyn PhysicalExpr>],
@@ -369,12 +391,15 @@ impl PhysicalExpr for InListExpr {
369391
// Use Arrow's vectorized eq kernel for types it supports (primitive,
370392
// boolean, string, binary, dictionary), falling back to row-by-row
371393
// comparator for unsupported types (nested, RunEndEncoded, etc.).
372-
let value = value.into_array(num_rows)?;
394+
// Normalize the left side once for the whole list. Doing this
395+
// outside `compare_one` avoids rescanning it for every item.
396+
let value =
397+
normalize_in_list_float_zero_value(value).into_array(num_rows)?;
373398
let lhs_supports_arrow_eq = supports_arrow_eq(value.data_type());
374399

375400
// Helper: compare value against a single list expression
376401
let compare_one = |expr: &Arc<dyn PhysicalExpr>| -> Result<BooleanArray> {
377-
match expr.evaluate(batch)? {
402+
match normalize_in_list_float_zero_value(expr.evaluate(batch)?) {
378403
ColumnarValue::Array(array) => {
379404
if lhs_supports_arrow_eq
380405
&& supports_arrow_eq(array.data_type())
@@ -3363,6 +3388,111 @@ mod tests {
33633388
Ok(())
33643389
}
33653390

3391+
#[test]
3392+
fn test_in_list_with_columns_float_signed_zero() -> Result<()> {
3393+
let schema = Schema::new(vec![
3394+
Field::new("a", DataType::Float64, false),
3395+
Field::new("b", DataType::Float64, false),
3396+
]);
3397+
let batch = RecordBatch::try_new(
3398+
Arc::new(schema.clone()),
3399+
vec![
3400+
Arc::new(Float64Array::from(vec![0.0, -0.0, 1.0])),
3401+
Arc::new(Float64Array::from(vec![-0.0, 0.0, 2.0])),
3402+
],
3403+
)?;
3404+
3405+
let expr = make_in_list_with_columns(
3406+
col("a", &schema)?,
3407+
vec![col("b", &schema)?],
3408+
false,
3409+
);
3410+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3411+
assert_eq!(
3412+
as_boolean_array(&result),
3413+
&BooleanArray::from(vec![true, true, false])
3414+
);
3415+
Ok(())
3416+
}
3417+
3418+
#[test]
3419+
fn test_in_list_with_columns_float_scalar_signed_zero() -> Result<()> {
3420+
let schema = Schema::new(vec![Field::new("a", DataType::Float32, false)]);
3421+
let batch = RecordBatch::try_new(
3422+
Arc::new(schema.clone()),
3423+
vec![Arc::new(Float32Array::from(vec![0.0, -0.0, 1.0]))],
3424+
)?;
3425+
let list = vec![lit(ScalarValue::Float32(Some(-0.0)))];
3426+
3427+
for (negated, expected) in [
3428+
(false, BooleanArray::from(vec![true, true, false])),
3429+
(true, BooleanArray::from(vec![false, false, true])),
3430+
] {
3431+
let expr =
3432+
make_in_list_with_columns(col("a", &schema)?, list.clone(), negated);
3433+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3434+
assert_eq!(as_boolean_array(&result), &expected);
3435+
}
3436+
3437+
// A scalar left-hand side is normalized before it is broadcast.
3438+
let expr = make_in_list_with_columns(
3439+
lit(ScalarValue::Float32(Some(-0.0))),
3440+
vec![col("a", &schema)?],
3441+
false,
3442+
);
3443+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3444+
assert_eq!(
3445+
as_boolean_array(&result),
3446+
&BooleanArray::from(vec![true, true, false])
3447+
);
3448+
3449+
Ok(())
3450+
}
3451+
3452+
#[test]
3453+
fn test_in_list_with_columns_dictionary_float_signed_zero() -> Result<()> {
3454+
let left: ArrayRef = Arc::new(DictionaryArray::try_new(
3455+
Int8Array::from(vec![0, 1, 2]),
3456+
Arc::new(Float64Array::from(vec![0.0, -0.0, 1.0])),
3457+
)?);
3458+
let right: ArrayRef = Arc::new(DictionaryArray::try_new(
3459+
Int8Array::from(vec![0, 1, 2]),
3460+
Arc::new(Float64Array::from(vec![-0.0, 0.0, 2.0])),
3461+
)?);
3462+
let data_type = left.data_type().clone();
3463+
let schema = Schema::new(vec![
3464+
Field::new("a", data_type.clone(), false),
3465+
Field::new("b", data_type, false),
3466+
]);
3467+
let batch = RecordBatch::try_new(Arc::new(schema.clone()), vec![left, right])?;
3468+
3469+
for (negated, expected) in [
3470+
(false, BooleanArray::from(vec![true, true, false])),
3471+
(true, BooleanArray::from(vec![false, false, true])),
3472+
] {
3473+
let expr = make_in_list_with_columns(
3474+
col("a", &schema)?,
3475+
vec![col("b", &schema)?],
3476+
negated,
3477+
);
3478+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3479+
assert_eq!(as_boolean_array(&result), &expected);
3480+
}
3481+
3482+
let scalar = lit(ScalarValue::Dictionary(
3483+
Box::new(DataType::Int8),
3484+
Box::new(ScalarValue::Float64(Some(-0.0))),
3485+
));
3486+
let expr = make_in_list_with_columns(col("a", &schema)?, vec![scalar], false);
3487+
let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?;
3488+
assert_eq!(
3489+
as_boolean_array(&result),
3490+
&BooleanArray::from(vec![true, true, false])
3491+
);
3492+
3493+
Ok(())
3494+
}
3495+
33663496
/// Tests that short-circuit evaluation produces correct results.
33673497
/// When all rows match after the first list item, remaining items
33683498
/// should be skipped without affecting correctness.
@@ -3871,6 +4001,29 @@ mod tests {
38714001
Ok(())
38724002
}
38734003

4004+
#[test]
4005+
fn test_try_new_from_array_dict_haystack_float64_signed_zero() -> Result<()> {
4006+
// One value beyond the branchless limit selects the hash-set strategy.
4007+
let list_len =
4008+
<Float64Type as branchless_filter::BranchlessFilterType>::MAX_LIST_LEN + 1;
4009+
let mut list_values = vec![Some(-0.0)];
4010+
list_values.extend((1..list_len).map(|value| Some(value as f64)));
4011+
let haystack = make_f64_dict_array(list_values);
4012+
let needles: ArrayRef = Arc::new(Float64Array::from(vec![0.0, -0.0, -1.0]));
4013+
let expected = BooleanArray::from(vec![true, true, false]);
4014+
4015+
assert_eq!(
4016+
eval_in_list_from_array(Arc::clone(&needles), Arc::clone(&haystack))?,
4017+
expected
4018+
);
4019+
assert_eq!(
4020+
eval_in_list_from_array(wrap_in_dict(needles), haystack)?,
4021+
expected
4022+
);
4023+
4024+
Ok(())
4025+
}
4026+
38744027
#[test]
38754028
fn test_try_new_from_array_type_mismatch_rejects() -> Result<()> {
38764029
let schema = Schema::new(vec![Field::new("a", DataType::Int32, false)]);

0 commit comments

Comments
 (0)