Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 49 additions & 4 deletions native/spark-expr/src/array_funcs/array_position.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<ColumnarValue, DataFusionError> {
Expand Down Expand Up @@ -273,7 +275,15 @@ fn position_fallback<O: OffsetSizeTrait>(
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)) {
Expand Down Expand Up @@ -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::<Float64Type, _, _>([
Some(vec![Some(1.0)]),
Some(vec![Some(f64::NAN)]),
Expand All @@ -324,7 +332,44 @@ mod tests {

let result = array_position_inner(&[Arc::new(array), Arc::new(element)])?;
let result = result.as_any().downcast_ref::<Int64Array>().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::<Float64Builder>(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::<Float64Builder>(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::<Int64Array>().unwrap();
// {-0.0} is the first element and now matches {0.0}, matching Spark.
assert_eq!(result, &Int64Array::from(vec![1]));
Ok(())
}
}
Expand Down
72 changes: 69 additions & 3 deletions native/spark-expr/src/array_funcs/arrays_overlap.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -388,11 +390,34 @@ where
}
}

fn normalize_list_element_floats<OffsetSize: OffsetSizeTrait>(
list: &GenericListArray<OffsetSize>,
) -> GenericListArray<OffsetSize> {
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<OffsetSize: OffsetSizeTrait>(
left: &GenericListArray<OffsetSize>,
right: &GenericListArray<OffsetSize>,
) -> Result<ArrayRef> {
let left_owned =
has_float_leaf(left.values().data_type()).then(|| normalize_list_element_floats(left));
let left: &GenericListArray<OffsetSize> = 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<OffsetSize> = right_owned.as_ref().unwrap_or(right);

let len = left.len();
let mut builder = BooleanArray::builder(len);

Expand Down Expand Up @@ -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::<i32>(&left, &right)?;
Expand All @@ -718,7 +742,18 @@ mod tests {
let right = make_nested_float_list(&[&[-0.0]]);
let result = arrays_overlap_list::<i32>(&left, &right)?;
let result = result.as_any().downcast_ref::<BooleanArray>().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::<i32>(&left, &right)?;
let result = result.as_any().downcast_ref::<BooleanArray>().unwrap();
assert!(result.value(0));
Ok(())
}

Expand Down Expand Up @@ -905,6 +940,37 @@ mod tests {
Ok(())
}

/// Build a single-row ListArray of structs: List<Struct<a: Float64>>
fn make_struct_float_list(elements: Vec<Option<f64>>) -> 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::<Float64Builder>(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::<i32>(&left, &right)?;
let result = result.as_any().downcast_ref::<BooleanArray>().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)
Expand Down
1 change: 1 addition & 0 deletions native/spark-expr/src/array_funcs/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
Loading