diff --git a/native/spark-expr/benches/arrays_overlap.rs b/native/spark-expr/benches/arrays_overlap.rs index c20efdd8b67..49a19795e5f 100644 --- a/native/spark-expr/benches/arrays_overlap.rs +++ b/native/spark-expr/benches/arrays_overlap.rs @@ -15,7 +15,7 @@ // specific language governing permissions and limitations // under the License. -use arrow::array::{ArrayRef, Int32Array, ListArray, StringArray}; +use arrow::array::{ArrayRef, Int32Array, ListArray, StringArray, StructArray}; use arrow::buffer::{NullBuffer, OffsetBuffer}; use arrow::datatypes::{DataType, Field}; use criterion::{criterion_group, criterion_main, Criterion}; @@ -23,6 +23,7 @@ use datafusion::common::config::ConfigOptions; use datafusion::logical_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl}; use datafusion_comet_spark_expr::SparkArraysOverlap; use std::hint::black_box; +use std::ops::Range; use std::sync::Arc; fn list_of(values: ArrayRef, rows: usize, elems_per_row: usize) -> ArrayRef { @@ -71,6 +72,65 @@ fn string_lists(rows: usize, elems_per_row: usize, offset: usize) -> (ArrayRef, ) } +fn nested_int_lists(rows: usize, elems_per_row: usize, offset: i32) -> (ArrayRef, ArrayRef) { + let total = rows * elems_per_row; + let build = |value_offset: i32| { + let values: ArrayRef = Arc::new(Int32Array::from_iter_values( + (0..total).flat_map(|i| [0, 1, 2, i as i32 + value_offset]), + )); + list_of(values, total, 4) + }; + ( + list_of(build(0), rows, elems_per_row), + list_of(build(offset), rows, elems_per_row), + ) +} + +/// Nested int32 lists of one-value lists, none of them null. Each left row holds `[v]` for each +/// `v` in `left`, and each right row likewise for `right`. +fn nested_singleton_lists( + rows: usize, + left: Range, + right: Range, +) -> (ArrayRef, ArrayRef) { + let build = |values: Range| { + let len = values.len(); + let inner: ArrayRef = Arc::new(Int32Array::from_iter_values( + (0..rows).flat_map(|_| values.clone()), + )); + let offsets: Vec = (0..=rows * len).map(|i| i as i32).collect(); + let singletons: ArrayRef = Arc::new(ListArray::new( + Arc::new(Field::new("item", DataType::Int32, true)), + OffsetBuffer::new(offsets.into()), + inner, + None, + )); + list_of(singletons, rows, len) + }; + (build(left), build(right)) +} + +fn struct_lists(rows: usize, elems_per_row: usize) -> (ArrayRef, ArrayRef) { + let total = rows * elems_per_row; + let build = |offset: i32| -> ArrayRef { + let first: ArrayRef = Arc::new(Int32Array::from_value(0, total)); + let second: ArrayRef = Arc::new(Int32Array::from_iter_values( + (0..total).map(|i| i as i32 + offset), + )); + Arc::new(StructArray::from(vec![ + (Arc::new(Field::new("first", DataType::Int32, false)), first), + ( + Arc::new(Field::new("second", DataType::Int32, false)), + second, + ), + ])) + }; + ( + list_of(build(0), rows, elems_per_row), + list_of(build(total as i32), rows, elems_per_row), + ) +} + fn invoke(udf: &SparkArraysOverlap, left: &ArrayRef, right: &ArrayRef) -> ColumnarValue { udf.invoke_with_args(ScalarFunctionArgs { args: vec![ @@ -113,6 +173,45 @@ fn criterion_benchmark(c: &mut Criterion) { c.bench_function("spark_arrays_overlap: utf8 long lists", |b| { b.iter(|| black_box(invoke(&udf, black_box(&left), black_box(&right)))) }); + + let (left, right) = nested_int_lists(rows, 8, (rows * 8) as i32); + c.bench_function("spark_arrays_overlap: nested int32 short lists", |b| { + b.iter(|| black_box(invoke(&udf, black_box(&left), black_box(&right)))) + }); + + let (left, right) = nested_int_lists(64, 64, 64 * 64); + c.bench_function("spark_arrays_overlap: nested int32 long lists", |b| { + b.iter(|| black_box(invoke(&udf, black_box(&left), black_box(&right)))) + }); + + let (left, right) = nested_int_lists(rows, 8, 4); + c.bench_function("spark_arrays_overlap: nested int32 early match", |b| { + b.iter(|| black_box(invoke(&udf, black_box(&left), black_box(&right)))) + }); + + // Rows hold 128 and 64 elements with one match. The first bench puts it last in the longer + // side, the second last in the shorter side. + let (left, right) = nested_singleton_lists(256, 0..128, 127..191); + c.bench_function( + "spark_arrays_overlap: nested int32 unequal lengths, late match in longer side", + |b| b.iter(|| black_box(invoke(&udf, black_box(&left), black_box(&right)))), + ); + + let (left, right) = nested_singleton_lists(256, 0..128, -63..1); + c.bench_function( + "spark_arrays_overlap: nested int32 unequal lengths, late match in shorter side", + |b| b.iter(|| black_box(invoke(&udf, black_box(&left), black_box(&right)))), + ); + + let (left, right) = struct_lists(rows, 8); + c.bench_function("spark_arrays_overlap: nested struct short lists", |b| { + b.iter(|| black_box(invoke(&udf, black_box(&left), black_box(&right)))) + }); + + let (left, right) = struct_lists(64, 64); + c.bench_function("spark_arrays_overlap: nested struct long lists", |b| { + b.iter(|| black_box(invoke(&udf, black_box(&left), black_box(&right)))) + }); } criterion_group!(benches, criterion_benchmark); diff --git a/native/spark-expr/src/array_funcs/arrays_overlap.rs b/native/spark-expr/src/array_funcs/arrays_overlap.rs index 542687c7d16..fa5410fb98e 100644 --- a/native/spark-expr/src/array_funcs/arrays_overlap.rs +++ b/native/spark-expr/src/array_funcs/arrays_overlap.rs @@ -170,7 +170,13 @@ fn arrays_overlap_list( let left_values = left.values(); let right_values = right.values(); - if left_values.data_type() != right_values.data_type() { + // Nested element types can differ in field nullability, such as a struct built by + // `array_repeat` beside one built by `array(...)`. `make_comparator` compares those, so keep + // them on the shared comparator below; only flat types need identical types for the fast + // paths and otherwise take the generic fallback. + let both_nested = + needs_comparator(left_values.data_type()) && needs_comparator(right_values.data_type()); + if left_values.data_type() != right_values.data_type() && !both_nested { return arrays_overlap_list_generic(left, right); } @@ -223,6 +229,29 @@ fn arrays_overlap_list( left_values.as_string::(), right_values.as_string::() ), + dt if needs_comparator(dt) => { + // Spark's nested path compares with ordering.equiv, where -0.0 == 0.0 and every NaN + // is equal, but Arrow's comparator uses total order. Normalize float leaves once per + // column so the comparator built over the full child arrays matches Spark. + let (left_values, right_values) = if has_float_leaf(dt) { + ( + normalize_nested_floats(left_values), + normalize_nested_floats(right_values), + ) + } else { + (Arc::clone(left_values), Arc::clone(right_values)) + }; + let comparator = make_comparator( + left_values.as_ref(), + right_values.as_ref(), + SortOptions::default(), + )?; + Ok(overlap_rows( + left, + right, + nested_row_overlap(&left_values, &right_values, comparator.as_ref()), + )) + } _ => arrays_overlap_list_generic(left, right), } } @@ -397,6 +426,50 @@ where } } +/// Row overlap for nested element types using one comparator for the full child arrays. +fn nested_row_overlap<'a>( + left: &'a ArrayRef, + right: &'a ArrayRef, + comparator: &'a dyn Fn(usize, usize) -> Ordering, +) -> impl FnMut(Range, Range) -> bool + 'a { + move |left_range, right_range| { + // Probe from the shorter side, as the per-row comparator did, so its early exits stay. + // The comparator takes the left index first, also when the right side is the outer loop. + if left_range.len() <= right_range.len() { + any_equal(left, left_range, right, right_range, comparator) + } else { + any_equal(right, right_range, left, left_range, |ri, li| { + comparator(li, ri) + }) + } + } +} + +/// True when a non-null `outer` element equals a non-null `inner` element. `compare` takes an +/// `outer` index, then an `inner` index. +fn any_equal( + outer: &ArrayRef, + outer_range: Range, + inner: &ArrayRef, + inner_range: Range, + compare: impl Fn(usize, usize) -> Ordering, +) -> bool { + for o in outer_range { + if outer.is_null(o) { + continue; + } + for i in inner_range.clone() { + if inner.is_null(i) { + continue; + } + if compare(o, i) == Ordering::Equal { + return true; + } + } + } + false +} + fn normalize_list_element_floats( list: &GenericListArray, ) -> GenericListArray { @@ -413,10 +486,11 @@ fn normalize_list_element_floats( ) } -/// Fallback for nested and otherwise unhandled element types. +/// Fallback for otherwise unhandled element types, including nested elements whose two sides +/// have different data types. /// /// note: Spark's flat arrays_overlap (HashSet) treats -0.0 and 0.0 as different, -/// only the nested path here treats them as equal. this normalization can't move into the +/// only the nested path treats them as equal. this normalization can't move into the /// flat fast path in arrays_overlap_list without breaking that difference. fn arrays_overlap_list_generic( left: &GenericListArray, @@ -464,26 +538,12 @@ fn arrays_overlap_list_generic( (&right_values, &left_values) }; - let comparator = if needs_comparator(probe.data_type()) { - Some(make_comparator( - probe.as_ref(), - search.as_ref(), - SortOptions::default(), - )?) - } else { - None - }; - for pi in 0..probe.len() { if probe.is_null(pi) { has_null = true; continue; } - let (found, null_eq) = if let Some(comparator) = &comparator { - find_in_array_nested(pi, search, comparator.as_ref()) - } else { - find_in_array_flat(probe, pi, search)? - }; + let (found, null_eq) = find_in_array_flat(probe, pi, search)?; if null_eq { has_null = true; } @@ -513,25 +573,6 @@ fn find_in_array_flat(probe: &ArrayRef, pi: usize, search: &ArrayRef) -> Result< Ok((eq_result.true_count() > 0, eq_result.null_count() > 0)) } -/// Element-by-element search using Arrow's nested comparator. -fn find_in_array_nested( - pi: usize, - search: &ArrayRef, - comparator: &dyn Fn(usize, usize) -> Ordering, -) -> (bool, bool) { - let mut has_null = false; - for si in 0..search.len() { - if search.is_null(si) { - has_null = true; - continue; - } - if comparator(pi, si) == Ordering::Equal { - return (true, has_null); - } - } - (false, has_null) -} - fn needs_comparator(dt: &DataType) -> bool { matches!( dt, @@ -895,6 +936,69 @@ mod tests { Ok(()) } + #[test] + fn test_nested_scan_probes_from_the_shorter_side() { + // The match is the last element of the longer side and the first of the shorter side. + // A reversed argument order reads past the shorter side and panics. + let long: ArrayRef = Arc::new(Int32Array::from_iter_values(0..128)); + let short: ArrayRef = Arc::new(Int32Array::from_iter_values(127..191)); + for (left, right) in [(&long, &short), (&short, &long)] { + let calls = std::cell::Cell::new(0); + let (left_values, right_values) = ( + left.as_primitive::(), + right.as_primitive::(), + ); + let comparator = |li: usize, ri: usize| { + calls.set(calls.get() + 1); + left_values.value(li).cmp(&right_values.value(ri)) + }; + let mut row_overlap = nested_row_overlap(left, right, &comparator); + assert!(row_overlap(0..left.len(), 0..right.len())); + assert_eq!(calls.get(), 128); + } + } + + #[test] + fn test_nested_array_sliced_offsets_and_nulls() -> Result<()> { + let make_rows = |rows: &[&[Option<&[i32]>]]| { + let mut builder = ListBuilder::new(ListBuilder::new(Int32Builder::new())); + for row in rows { + for element in *row { + if let Some(values) = element { + builder.values().values().append_slice(values); + builder.values().append(true); + } else { + builder.values().append(false); + } + } + builder.append(true); + } + builder.finish() + }; + let left = make_rows(&[ + &[Some(&[999])], + &[Some(&[10])], + &[Some(&[10]), None], + &[Some(&[50]), Some(&[60]), Some(&[70])], + ]) + .slice(1, 3); + let right = make_rows(&[ + &[Some(&[999])], + &[Some(&[20]), Some(&[30]), Some(&[40])], + &[Some(&[20])], + &[Some(&[60])], + ]) + .slice(1, 3); + + let result = arrays_overlap_list::(&left, &right)?; + let result = result.as_any().downcast_ref::().unwrap(); + assert_eq!( + result, + &BooleanArray::from(vec![Some(false), None, Some(true)]) + ); + Ok(()) + } + #[test] fn test_nested_array_basic_overlap() -> Result<()> { // [[1,2], [3,4]] vs [[3,4], [5,6]] => true @@ -1026,6 +1130,37 @@ mod tests { list_builder.finish() } + #[test] + fn test_struct_overlap_with_different_field_nullability() -> Result<()> { + // The same struct values where one side declares `a` non-nullable, as `array_repeat` + // produces while `array(...)` widens it to nullable: [{1,2}] vs [{1,2}] => true + fn single_struct_list(a_nullable: bool) -> ListArray { + let fields = vec![ + Arc::new(Field::new("a", DataType::Int32, a_nullable)), + Arc::new(Field::new("b", DataType::Int32, true)), + ]; + let mut list_builder = ListBuilder::new(StructBuilder::new( + fields, + vec![Box::new(Int32Builder::new()), Box::new(Int32Builder::new())], + )); + let sb = list_builder.values(); + sb.field_builder::(0).unwrap().append_value(1); + sb.field_builder::(1).unwrap().append_value(2); + sb.append(true); + list_builder.append(true); + list_builder.finish() + } + let left = single_struct_list(false); + let right = single_struct_list(true); + assert_ne!(left.values().data_type(), right.values().data_type()); + + 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_basic_overlap() -> Result<()> { // [{1,2}, {3,4}] vs [{3,4}, {5,6}] => true diff --git a/spark/src/test/resources/sql-tests/expressions/array/arrays_overlap.sql b/spark/src/test/resources/sql-tests/expressions/array/arrays_overlap.sql index 1ddfe14dbcd..5ca5db839b9 100644 --- a/spark/src/test/resources/sql-tests/expressions/array/arrays_overlap.sql +++ b/spark/src/test/resources/sql-tests/expressions/array/arrays_overlap.sql @@ -275,6 +275,18 @@ INSERT INTO test_overlap_struct VALUES (array(named_struct('x', 1, 'y', 2)), arr query SELECT a, b, arrays_overlap(a, b) FROM test_overlap_struct +-- array_repeat keeps the non-nullable x field while array() widens it to nullable, so the two +-- sides' element types differ only in nested nullability +statement +CREATE TABLE test_overlap_mixed_ctor(i int) USING parquet + +statement +INSERT INTO test_overlap_mixed_ctor VALUES (1), (2), (NULL) + +query +SELECT i, arrays_overlap(array_repeat(named_struct('x', 1, 'y', i), 1), array(named_struct('x', 1, 'y', i))) +FROM test_overlap_mixed_ctor + -- mixed column and literal with NULL elements query SELECT arrays_overlap(a, array(99, NULL)) FROM test_arrays_overlap