diff --git a/native/spark-expr/src/array_funcs/array_position.rs b/native/spark-expr/src/array_funcs/array_position.rs index 191091aabf..b2b54c40e7 100644 --- a/native/spark-expr/src/array_funcs/array_position.rs +++ b/native/spark-expr/src/array_funcs/array_position.rs @@ -33,6 +33,8 @@ use num::Float; use std::cmp::Ordering; use std::sync::Arc; +use super::nested_float_normalize::{has_float_leaf, normalize_nested_floats}; + /// Spark array_position() function that returns the 1-based position of an element in an array. /// Returns 0 if the element is not found (Spark behavior differs from DataFusion which returns null). fn spark_array_position(args: &[ColumnarValue]) -> Result { @@ -273,7 +275,15 @@ fn position_fallback( let num_rows = list_array.len(); let nulls = combined_nulls(list_array.nulls(), element.nulls()); let mut result = vec![0i64; num_rows]; - let comparator = make_comparator(values.as_ref(), element.as_ref(), SortOptions::default())?; + let values_normalized = + has_float_leaf(values.data_type()).then(|| normalize_nested_floats(values)); + let element_normalized = + has_float_leaf(element.data_type()).then(|| normalize_nested_floats(element)); + let comparator = make_comparator( + values_normalized.as_ref().unwrap_or(values).as_ref(), + element_normalized.as_ref().unwrap_or(element).as_ref(), + SortOptions::default(), + )?; for (row_index, w) in offsets.windows(2).enumerate() { if nulls.as_ref().is_some_and(|n| n.is_null(row_index)) { @@ -301,8 +311,6 @@ mod tests { #[test] fn test_nested_float_and_null_position() -> DataFusionResult<()> { - // Arrow and the previous ScalarValue fallback distinguish signed zeros, so the second - // row matches at position 2 rather than position 1. let values = ListArray::from_iter_primitive::([ Some(vec![Some(1.0)]), Some(vec![Some(f64::NAN)]), @@ -324,7 +332,44 @@ mod tests { let result = array_position_inner(&[Arc::new(array), Arc::new(element)])?; let result = result.as_any().downcast_ref::().unwrap(); - assert_eq!(result, &Int64Array::from(vec![2, 2, 1])); + assert_eq!(result, &Int64Array::from(vec![2, 1, 1])); + Ok(()) + } + + #[test] + fn test_struct_float_field_signed_zero_position() -> DataFusionResult<()> { + use arrow::array::{Float64Builder, StructBuilder}; + + let fields = vec![Arc::new(Field::new("a", DataType::Float64, true))]; + let mut values_builder = + StructBuilder::new(fields.clone(), vec![Box::new(Float64Builder::new())]); + for v in [-0.0, 1.0] { + values_builder + .field_builder::(0) + .unwrap() + .append_value(v); + values_builder.append(true); + } + let values = Arc::new(values_builder.finish()); + let array = ListArray::new( + Arc::new(Field::new("item", values.data_type().clone(), true)), + OffsetBuffer::new(vec![0, 2].into()), + values, + None, + ); + + let mut element_builder = StructBuilder::new(fields, vec![Box::new(Float64Builder::new())]); + element_builder + .field_builder::(0) + .unwrap() + .append_value(0.0); + element_builder.append(true); + let element = element_builder.finish(); + + let result = array_position_inner(&[Arc::new(array), Arc::new(element)])?; + let result = result.as_any().downcast_ref::().unwrap(); + // {-0.0} is the first element and now matches {0.0}, matching Spark. + assert_eq!(result, &Int64Array::from(vec![1])); Ok(()) } } diff --git a/native/spark-expr/src/array_funcs/arrays_overlap.rs b/native/spark-expr/src/array_funcs/arrays_overlap.rs index bd75a6ddcc..62fe0dc697 100644 --- a/native/spark-expr/src/array_funcs/arrays_overlap.rs +++ b/native/spark-expr/src/array_funcs/arrays_overlap.rs @@ -51,6 +51,8 @@ use std::hash::Hash; use std::ops::Range; use std::sync::Arc; +use super::nested_float_normalize::{has_float_leaf, normalize_nested_floats}; + #[derive(Debug, PartialEq, Eq, Hash)] pub struct SparkArraysOverlap { signature: Signature, @@ -388,11 +390,34 @@ where } } +fn normalize_list_element_floats( + list: &GenericListArray, +) -> GenericListArray { + let field = match list.data_type() { + DataType::List(f) | DataType::LargeList(f) => Arc::clone(f), + _ => unreachable!("GenericListArray always has List or LargeList data type"), + }; + let normalized_values = normalize_nested_floats(list.values()); + GenericListArray::new( + field, + list.offsets().clone(), + normalized_values, + list.nulls().cloned(), + ) +} + /// Fallback for nested and otherwise unhandled element types. fn arrays_overlap_list_generic( left: &GenericListArray, right: &GenericListArray, ) -> Result { + let left_owned = + has_float_leaf(left.values().data_type()).then(|| normalize_list_element_floats(left)); + let left: &GenericListArray = left_owned.as_ref().unwrap_or(left); + let right_owned = + has_float_leaf(right.values().data_type()).then(|| normalize_list_element_floats(right)); + let right: &GenericListArray = right_owned.as_ref().unwrap_or(right); + let len = left.len(); let mut builder = BooleanArray::builder(len); @@ -706,8 +731,7 @@ mod tests { #[test] fn test_nested_float_total_order() -> Result<()> { - // Preserve the existing Arrow total-order behavior: NaN matches itself, while signed - // zeros are distinct. + // NaN matches itself, and signed zeros are equal, matching Spark. let left = make_nested_float_list(&[&[f64::NAN]]); let right = make_nested_float_list(&[&[f64::NAN]]); let result = arrays_overlap_list::(&left, &right)?; @@ -718,7 +742,18 @@ mod tests { let right = make_nested_float_list(&[&[-0.0]]); let result = arrays_overlap_list::(&left, &right)?; let result = result.as_any().downcast_ref::().unwrap(); - assert!(!result.value(0)); + assert!(result.value(0)); + Ok(()) + } + + #[test] + fn test_nested_float_signed_nan_total_order() -> Result<()> { + // [[-NaN]] vs [[NaN]] => true + let left = make_nested_float_list(&[&[-f64::NAN]]); + let right = make_nested_float_list(&[&[f64::NAN]]); + let result = arrays_overlap_list::(&left, &right)?; + let result = result.as_any().downcast_ref::().unwrap(); + assert!(result.value(0)); Ok(()) } @@ -905,6 +940,37 @@ mod tests { Ok(()) } + /// Build a single-row ListArray of structs: List> + fn make_struct_float_list(elements: Vec>) -> ListArray { + let fields = vec![Arc::new(Field::new("a", DataType::Float64, true))]; + let struct_builder = + StructBuilder::new(fields.clone(), vec![Box::new(Float64Builder::new())]); + let mut list_builder = ListBuilder::new(struct_builder); + + for elem in &elements { + let sb = list_builder.values(); + sb.field_builder::(0) + .unwrap() + .append_option(*elem); + sb.append(true); + } + list_builder.append(true); + list_builder.finish() + } + + #[test] + fn test_struct_float_field_signed_zero_overlap() -> Result<()> { + // [{-0.0}] vs [{0.0}] => true, matching Spark + let left = make_struct_float_list(vec![Some(-0.0)]); + let right = make_struct_float_list(vec![Some(0.0)]); + + let result = arrays_overlap_list::(&left, &right)?; + let result = result.as_any().downcast_ref::().unwrap(); + assert!(result.is_valid(0)); + assert!(result.value(0)); + Ok(()) + } + #[test] fn test_struct_null_element() -> Result<()> { // [NULL] vs [{1,2}] => null (null outer element) diff --git a/native/spark-expr/src/array_funcs/mod.rs b/native/spark-expr/src/array_funcs/mod.rs index 0c2c68dc6d..cd7c126f56 100644 --- a/native/spark-expr/src/array_funcs/mod.rs +++ b/native/spark-expr/src/array_funcs/mod.rs @@ -23,6 +23,7 @@ mod arrays_zip; mod flatten; mod get_array_struct_fields; mod list_extract; +mod nested_float_normalize; mod size; pub use array_insert::ArrayInsert; diff --git a/native/spark-expr/src/array_funcs/nested_float_normalize.rs b/native/spark-expr/src/array_funcs/nested_float_normalize.rs new file mode 100644 index 0000000000..aa348acba3 --- /dev/null +++ b/native/spark-expr/src/array_funcs/nested_float_normalize.rs @@ -0,0 +1,208 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use crate::math_funcs::internal::normalize_float; +use arrow::array::{ + Array, ArrayRef, AsArray, FixedSizeListArray, Float32Array, Float64Array, LargeListArray, + ListArray, StructArray, +}; +use arrow::datatypes::{DataType, Float32Type, Float64Type}; +use std::sync::Arc; + +pub(super) fn has_float_leaf(dt: &DataType) -> bool { + match dt { + DataType::Float32 | DataType::Float64 => true, + DataType::List(field) | DataType::LargeList(field) | DataType::FixedSizeList(field, _) => { + has_float_leaf(field.data_type()) + } + DataType::Struct(fields) => fields.iter().any(|f| has_float_leaf(f.data_type())), + _ => false, + } +} + +/// Recursively rebuilds nested arrays with `-0.0` normalized to `0.0` and NaN canonicalized +/// in any Float32/Float64 leaves. +pub(super) fn normalize_nested_floats(array: &ArrayRef) -> ArrayRef { + match array.data_type() { + DataType::Float32 => { + let normalized: Float32Array = + array.as_primitive::().unary(normalize_float); + Arc::new(normalized) + } + DataType::Float64 => { + let normalized: Float64Array = + array.as_primitive::().unary(normalize_float); + Arc::new(normalized) + } + DataType::List(field) => { + let list = array.as_list::(); + let normalized_values = normalize_nested_floats(list.values()); + Arc::new(ListArray::new( + Arc::clone(field), + list.offsets().clone(), + normalized_values, + list.nulls().cloned(), + )) + } + DataType::LargeList(field) => { + let list = array.as_list::(); + let normalized_values = normalize_nested_floats(list.values()); + Arc::new(LargeListArray::new( + Arc::clone(field), + list.offsets().clone(), + normalized_values, + list.nulls().cloned(), + )) + } + DataType::FixedSizeList(field, size) => { + let list = array.as_fixed_size_list(); + let normalized_values = normalize_nested_floats(list.values()); + Arc::new(FixedSizeListArray::new( + Arc::clone(field), + *size, + normalized_values, + list.nulls().cloned(), + )) + } + DataType::Struct(_) => { + let s = array.as_struct(); + let normalized_columns: Vec = + s.columns().iter().map(normalize_nested_floats).collect(); + Arc::new(StructArray::new( + s.fields().clone(), + normalized_columns, + s.nulls().cloned(), + )) + } + _ => Arc::clone(array), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::Float64Builder; + use arrow::array::ListBuilder; + use arrow::datatypes::Field; + + #[test] + fn test_has_float_leaf() { + assert!(has_float_leaf(&DataType::Float64)); + assert!(has_float_leaf(&DataType::List(Arc::new(Field::new( + "item", + DataType::Float32, + true + ))))); + assert!(has_float_leaf(&DataType::Struct( + vec![ + Arc::new(Field::new("a", DataType::Int32, true)), + Arc::new(Field::new("b", DataType::Float64, true)), + ] + .into() + ))); + assert!(!has_float_leaf(&DataType::Int32)); + assert!(!has_float_leaf(&DataType::List(Arc::new(Field::new( + "item", + DataType::Int32, + true + ))))); + } + + #[test] + fn test_normalize_flat_floats() { + let arr: ArrayRef = Arc::new(Float64Array::from(vec![ + Some(-0.0), + Some(0.0), + Some(f64::NAN), + Some(-f64::NAN), + None, + Some(1.5), + ])); + let normalized = normalize_nested_floats(&arr); + let normalized = normalized.as_primitive::(); + + assert_eq!(normalized.value(0).to_bits(), 0.0f64.to_bits()); + assert_eq!(normalized.value(1).to_bits(), 0.0f64.to_bits()); + assert_eq!(normalized.value(2).to_bits(), f64::NAN.to_bits()); + assert_eq!(normalized.value(3).to_bits(), f64::NAN.to_bits()); + assert!(normalized.is_null(4)); + assert_eq!(normalized.value(5), 1.5); + } + + #[test] + fn test_normalize_nested_list_floats() { + let mut builder = ListBuilder::new(Float64Builder::new()); + builder.values().append_value(-0.0); + builder.values().append_value(-f64::NAN); + builder.append(true); + let arr: ArrayRef = Arc::new(builder.finish()); + + let normalized = normalize_nested_floats(&arr); + let normalized = normalized.as_list::(); + let inner = normalized.value(0); + let inner = inner.as_primitive::(); + + assert_eq!(inner.value(0).to_bits(), 0.0f64.to_bits()); + assert_eq!(inner.value(1).to_bits(), f64::NAN.to_bits()); + } + + #[test] + fn test_normalize_struct_floats() { + let a = Float64Array::from(vec![Some(-0.0), Some(1.0)]); + let b = Float64Array::from(vec![Some(-f64::NAN), Some(-0.0)]); + let fields = vec![ + Arc::new(Field::new("a", DataType::Float64, true)), + Arc::new(Field::new("b", DataType::Float64, true)), + ]; + let arr: ArrayRef = Arc::new(StructArray::new( + fields.into(), + vec![Arc::new(a), Arc::new(b)], + None, + )); + + let normalized = normalize_nested_floats(&arr); + let normalized = normalized.as_struct(); + let col_a = normalized.column(0).as_primitive::(); + let col_b = normalized.column(1).as_primitive::(); + + assert_eq!(col_a.value(0).to_bits(), 0.0f64.to_bits()); + assert_eq!(col_a.value(1), 1.0); + assert_eq!(col_b.value(0).to_bits(), f64::NAN.to_bits()); + assert_eq!(col_b.value(1).to_bits(), 0.0f64.to_bits()); + } + + #[test] + fn test_normalize_fixed_size_list_floats() { + let values = Float64Array::from(vec![Some(-0.0), Some(-f64::NAN), Some(1.0), Some(-0.0)]); + let field = Arc::new(Field::new("item", DataType::Float64, true)); + let arr: ArrayRef = Arc::new(FixedSizeListArray::new( + Arc::clone(&field), + 2, + Arc::new(values), + None, + )); + + let normalized = normalize_nested_floats(&arr); + let normalized = normalized.as_fixed_size_list(); + let flat = normalized.values().as_primitive::(); + + assert_eq!(flat.value(0).to_bits(), 0.0f64.to_bits()); + assert_eq!(flat.value(1).to_bits(), f64::NAN.to_bits()); + assert_eq!(flat.value(2), 1.0); + assert_eq!(flat.value(3).to_bits(), 0.0f64.to_bits()); + } +}