SubhamSinghal commented on code in PR #25840:
URL: https://github.com/apache/datafusion/pull/25840#discussion_r4178689803


##########
datafusion/physical-plan/src/joins/piecewise_merge_join/classic_join.rs:
##########
@@ -1185,6 +1245,347 @@ mod tests {
         Ok(())
     }
 
+    /// An empty streamed batch of run-end encoded keys, sliced past the first 
run, must not
+    /// reach arrow's run-end comparison kernel, which overflows on it.
+    #[tokio::test]
+    async fn join_right_run_end_encoded_keys_empty_sliced_batch() -> 
Result<()> {
+        use arrow::array::{Int32Array, RunArray};
+        use arrow::datatypes::Int32Type;
+
+        let ree_exec = |ids: &str,
+                        keys: &str,
+                        run_ends: Vec<i32>,
+                        values: Vec<Option<i32>>,
+                        slices: &[(usize, usize)]|
+         -> Result<Arc<dyn ExecutionPlan>> {
+            let keys_array = RunArray::<Int32Type>::try_new(
+                &Int32Array::from(run_ends),
+                &Int32Array::from(values),
+            )?;
+            let schema = Arc::new(Schema::new(vec![
+                Field::new(ids, DataType::Int32, false),
+                Field::new(keys, keys_array.data_type().clone(), true),
+            ]));
+            let len = keys_array.len() as i32;
+            let batch = RecordBatch::try_new(
+                Arc::clone(&schema),
+                vec![
+                    Arc::new(Int32Array::from((0..len).collect::<Vec<_>>())),
+                    Arc::new(keys_array),
+                ],
+            )?;
+            let batches = slices
+                .iter()
+                .map(|&(offset, len)| batch.slice(offset, len))
+                .collect::<Vec<_>>();
+            Ok(TestMemoryExec::try_new_exec(&[batches], schema, None)?)
+        };
+        // Buffered keys sorted for `<` (descending, NULLs first): {NULL, 5, 
5, 1, 1}.
+        let left = ree_exec(
+            "a1",
+            "b1",
+            vec![1, 3, 5],
+            vec![None, Some(5), Some(1)],
+            &[(0, 5)],
+        )?;
+        // Streamed keys {NULL, 2, 2, 0, 6}, as batches of 3, 0 and 2 rows. 
The empty one
+        // starts at offset 3, past the end of the first run.
+        let right = ree_exec(
+            "a2",
+            "b2",
+            vec![1, 3, 4, 5],
+            vec![None, Some(2), Some(0), Some(6)],
+            &[(0, 3), (3, 0), (3, 2)],
+        )?;
+        let on = (
+            Arc::new(Column::new_with_schema("b1", &left.schema())?) as _,
+            Arc::new(Column::new_with_schema("b2", &right.schema())?) as _,
+        );
+
+        join_collect(left, right, on, Operator::Lt, JoinType::Right).await?;
+        Ok(())
+    }
+
+    /// `matchable_rows` against the scan's own definition of a match -- some 
non-NULL
+    /// buffered row `idx` with `is_match(cmp.compare(row, idx))` under the 
scan's
+    /// `JoinKeyComparator` -- for every operator and every kind of key it 
dispatches on:
+    /// flat keys through `apply_cmp` (incl. signed zeros, NaN, strings, 
dictionaries) and
+    /// nested keys with NULL elements through the comparator. Each is also 
checked against
+    /// an all-NULL and an empty buffered side.
+    #[test]
+    fn matchable_rows_agrees_with_scan() -> Result<()> {
+        use arrow::array::{
+            BinaryArray, BooleanArray, Decimal128Array, DictionaryArray, 
Float64Array,
+            Int32Array, ListArray, StringArray, StringViewArray,
+            TimestampMicrosecondArray,
+        };
+        use arrow::datatypes::Int32Type;
+
+        let strings =
+            |v: &[Option<&str>]| Arc::new(StringArray::from(v.to_vec())) as 
ArrayRef;
+        let cases: Vec<(&str, ArrayRef, ArrayRef)> = vec![
+            (
+                "int32",
+                Arc::new(Int32Array::from(vec![
+                    Some(5),
+                    None,
+                    Some(1),
+                    Some(9),
+                    Some(3),
+                ])),
+                Arc::new(Int32Array::from(vec![
+                    Some(0),
+                    Some(1),
+                    Some(2),
+                    Some(3),
+                    Some(5),
+                    Some(9),
+                    Some(10),
+                    None,
+                ])),
+            ),
+            (
+                "float64",
+                Arc::new(Float64Array::from(vec![
+                    Some(0.0),
+                    Some(-0.0),
+                    Some(1.5),
+                    Some(f64::NAN),
+                    None,
+                ])),
+                Arc::new(Float64Array::from(vec![
+                    Some(-0.0),
+                    Some(0.0),
+                    Some(1.5),
+                    Some(f64::NAN),
+                    Some(-1.0),
+                    Some(2.0),
+                    Some(f64::INFINITY),
+                    Some(f64::NEG_INFINITY),
+                    None,
+                ])),
+            ),
+            (
+                "utf8",
+                strings(&[Some("b"), Some(""), Some("d"), None]),
+                strings(&[
+                    Some(""),
+                    Some("a"),
+                    Some("b"),
+                    Some("c"),
+                    Some("d"),
+                    Some("e"),
+                    None,
+                ]),
+            ),
+            (
+                "utf8_view",
+                Arc::new(StringViewArray::from(vec![
+                    Some("b"),
+                    Some(""),
+                    Some("d"),
+                    None,
+                ])),
+                Arc::new(StringViewArray::from(vec![
+                    Some(""),
+                    Some("a"),
+                    Some("b"),
+                    Some("d"),
+                    Some("e"),
+                    None,
+                ])),
+            ),
+            (
+                "binary",
+                Arc::new(BinaryArray::from(vec![Some(&b"b"[..]), Some(b""), 
None])),
+                Arc::new(BinaryArray::from(vec![
+                    Some(&b""[..]),
+                    Some(b"a"),
+                    Some(b"b"),
+                    Some(b"c"),
+                    None,
+                ])),
+            ),
+            (
+                "dictionary",
+                Arc::new(
+                    vec![Some("b"), None, Some("d"), Some("b")]
+                        .into_iter()
+                        .collect::<DictionaryArray<Int32Type>>(),
+                ),
+                Arc::new(
+                    vec![Some("a"), Some("b"), Some("c"), Some("d"), 
Some("e"), None]
+                        .into_iter()
+                        .collect::<DictionaryArray<Int32Type>>(),
+                ),
+            ),
+            (
+                // `-0.0` as the buffered extreme inside a dictionary: the 
smallest
+                // key for `<`/`<=`, the largest for `>`/`>=`.
+                "dictionary_float64_neg_zero_min",
+                Arc::new(DictionaryArray::<Int32Type>::new(
+                    Int32Array::from(vec![0, 1]),
+                    Arc::new(Float64Array::from(vec![-0.0, 0.5])),
+                )),
+                Arc::new(DictionaryArray::<Int32Type>::new(
+                    Int32Array::from(vec![0, 1, 2]),
+                    Arc::new(Float64Array::from(vec![0.0, -0.0, 0.25])),
+                )),
+            ),
+            (
+                "dictionary_float64_neg_zero_max",
+                Arc::new(DictionaryArray::<Int32Type>::new(
+                    Int32Array::from(vec![0, 1]),
+                    Arc::new(Float64Array::from(vec![-1.0, -0.0])),
+                )),
+                Arc::new(DictionaryArray::<Int32Type>::new(
+                    Int32Array::from(vec![0, 1, 2]),
+                    Arc::new(Float64Array::from(vec![0.0, -0.0, -0.5])),
+                )),
+            ),
+            (
+                "decimal128",
+                Arc::new(
+                    Decimal128Array::from(vec![Some(150), None, Some(-25)])
+                        .with_precision_and_scale(10, 2)?,
+                ),
+                Arc::new(
+                    Decimal128Array::from(vec![
+                        Some(-26),
+                        Some(-25),
+                        Some(0),
+                        Some(150),
+                        Some(151),
+                        None,
+                    ])
+                    .with_precision_and_scale(10, 2)?,
+                ),
+            ),
+            (
+                "date32",
+                Arc::new(Date32Array::from(vec![Some(10), None, Some(20)])),
+                Arc::new(Date32Array::from(vec![
+                    Some(9),
+                    Some(10),
+                    Some(15),
+                    Some(20),
+                    Some(21),
+                    None,
+                ])),
+            ),
+            (
+                "timestamp_us",
+                Arc::new(TimestampMicrosecondArray::from(vec![
+                    Some(-5),
+                    Some(7),
+                    None,
+                ])),
+                Arc::new(TimestampMicrosecondArray::from(vec![
+                    Some(-6),
+                    Some(-5),
+                    Some(7),
+                    Some(8),
+                    None,
+                ])),
+            ),
+            (
+                "boolean",
+                Arc::new(BooleanArray::from(vec![Some(true), None])),
+                Arc::new(BooleanArray::from(vec![Some(false), Some(true), 
None])),
+            ),
+            (
+                "list_with_null_elements",
+                Arc::new(ListArray::from_iter_primitive::<Int32Type, _, 
_>(vec![
+                    Some(vec![Some(5)]),
+                    Some(vec![None]),
+                    Some(vec![Some(1), Some(2)]),
+                    Some(vec![]),
+                    None,
+                ])),
+                Arc::new(ListArray::from_iter_primitive::<Int32Type, _, 
_>(vec![
+                    Some(vec![None]),
+                    Some(vec![Some(7)]),
+                    Some(vec![Some(1)]),
+                    Some(vec![Some(5)]),
+                    Some(vec![]),
+                    Some(vec![None, Some(1)]),
+                    Some(vec![Some(5), None]),
+                    None,
+                ])),
+            ),
+            (
+                // No empty list, so the extreme key has a first element to 
compare
+                // against the streamed NULL elements.
+                "list_null_elements_vs_extreme",
+                Arc::new(ListArray::from_iter_primitive::<Int32Type, _, 
_>(vec![
+                    Some(vec![Some(5)]),
+                    Some(vec![Some(3)]),
+                    None,
+                ])),
+                Arc::new(ListArray::from_iter_primitive::<Int32Type, _, 
_>(vec![
+                    Some(vec![None]),
+                    Some(vec![Some(3), None]),
+                    Some(vec![Some(4)]),
+                    Some(vec![Some(7)]),
+                    Some(vec![Some(1)]),
+                    None,
+                ])),
+            ),
+        ];
+
+        for (name, buffered, streamed) in cases {
+            let all_null = new_null_array(buffered.data_type(), 3);
+            let empty = buffered.slice(0, 0);
+            for (side, buffered) in
+                [("full", buffered), ("all_null", all_null), ("empty", empty)]
+            {
+                for operator in
+                    [Operator::Lt, Operator::LtEq, Operator::Gt, 
Operator::GtEq]
+                {
+                    // As `PiecewiseMergeJoinExec::try_new` sorts both sides.
+                    let sort_options = match operator {
+                        Operator::Lt | Operator::LtEq => 
SortOptions::new(true, true),
+                        _ => SortOptions::new(false, true),
+                    };
+                    let sorted = take(
+                        buffered.as_ref(),
+                        &sort_to_indices(buffered.as_ref(), 
Some(sort_options), None)?,
+                        None,
+                    )?;
+                    let extreme = buffered_extreme(&sorted)?;
+
+                    let actual = matchable_rows(
+                        &streamed,
+                        &sorted,
+                        extreme.as_ref(),
+                        operator,
+                        sort_options,
+                    )?;
+
+                    let match_on_equal = matches_on_equal(operator)?;
+                    let cmp = JoinKeyComparator::new(
+                        &[Arc::clone(&streamed)],
+                        &[Arc::clone(&sorted)],
+                        &[sort_options],
+                        NullEquality::NullEqualsNothing,
+                    )?;
+                    for row in 0..streamed.len() {

Review Comment:
   Addressed in 2627267b525c9eece069b89c55afc8099be349dc



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to