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


##########
datafusion/physical-plan/src/joins/piecewise_merge_join/classic_join.rs:
##########
@@ -631,18 +588,121 @@ fn build_matched_indices_and_mark_buffered(
     )?)
 }
 
-// Creates a record batch from the unmatched indices on the streamed side
-fn create_unmatched_batch(
-    streamed_indices: &mut PrimitiveBuilder<UInt32Type>,
-    stream_batch: &SortedStreamBatch,
+// The last key of the sorted buffered side, or `None` when it is empty or 
every key is NULL:
+// NULLs sort first, so the last key is NULL only when every buffered key is.
+fn buffered_extreme(values: &ArrayRef) -> Result<Option<ColumnarValue>> {
+    Ok(match values.len().checked_sub(1) {
+        // `apply_cmp` normalizes `-0.0` only in flat float scalars, not 
inside a
+        // `ScalarValue::Dictionary`, so normalize the key before taking it.
+        Some(last) if values.is_valid(last) => {
+            let extreme = normalize_float_zero(&values.slice(last, 1));
+            Some(ColumnarValue::Scalar(ScalarValue::try_from_array(
+                &extreme, 0,
+            )?))
+        }
+        _ => None,
+    })
+}
+
+// Which rows of `stream_values` can match at least one buffered row.
+//
+// Every match set is a suffix `[k, buffered_len)` of the sorted buffered 
side, so a streamed
+// row matches anything at all iff it matches the last buffered row, 
`buffered_extreme`
+// (`None` when no buffered key is non-null, so nothing can match). A NULL on 
either side
+// compares to NULL, which is no match.
+//
+// This must agree exactly with the scan's `JoinKeyComparator`, or it would 
change which rows
+// match. For flat keys one vectorized `apply_cmp` does: it normalizes `-0.0` 
to `+0.0` as the
+// comparator does. For nested keys it does not -- `apply_cmp` orders NULL 
elements inside a
+// key ascending, while the comparator applies the sort options (descending 
for `<`/`<=`) at
+// every level -- so those are decided with the scan's own comparator.
+fn matchable_rows(
+    stream_values: &ArrayRef,
+    buffered_values: &ArrayRef,
+    buffered_extreme: Option<&ColumnarValue>,
+    operator: Operator,
+    sort_options: SortOptions,
+) -> Result<BooleanArray> {
+    let num_rows = stream_values.len();
+    let Some(extreme) = buffered_extreme else {
+        return Ok(BooleanArray::new(BooleanBuffer::new_unset(num_rows), None));
+    };
+
+    // The comparator's nested ordering is itself wrong for `<`/`<=` (#25957). 
Once that is
+    // fixed, nested keys can go through `apply_cmp` too and this branch can 
be removed.
+    //
+    // Run-end encoded keys are decided here too: arrow's run-end comparison 
kernel overflows
+    // on an empty batch sliced past its first run, which the scan never 
reaches.
+    if stream_values.data_type().is_nested()
+        || matches!(stream_values.data_type(), DataType::RunEndEncoded(_, _))
+    {
+        let match_on_equal = matches_on_equal(operator)?;
+        let cmp = JoinKeyComparator::new(
+            &[Arc::clone(stream_values)],
+            &[Arc::clone(buffered_values)],
+            &[sort_options],
+            NullEquality::NullEqualsNothing,
+        )?;
+        let last = buffered_values.len() - 1;
+        let matchable = BooleanBuffer::collect_bool(num_rows, |row| {
+            stream_values.is_valid(row)

Review Comment:
   `is_valid` only checks the physical null buffer, and a `RunArray` has none. 
So a run-end encoded NULL key passes as matchable, gets past the `null_count() 
> 0` guard at :492, and the scan joins it with every buffered row 
(`nulls_first` → `Less`). `main` gives the same rows, so this is fine as a 
follow-up. But `join_right_run_end_encoded_keys_empty_sliced_batch` already 
feeds such a NULL (a2 = 0) and asserts nothing.
   
   ```diff
            let last = buffered_values.len() - 1;
   +        // A run-end encoded NULL has no physical null buffer.
   +        let nulls = stream_values.logical_nulls();
            let matchable = BooleanBuffer::collect_bool(num_rows, |row| {
   -            stream_values.is_valid(row)
   +            nulls.as_ref().is_none_or(|n| n.is_valid(row))
                    && is_match(cmp.compare(row, last), match_on_equal)
            });
   ```
   ```diff
   -    if stream_values[0].null_count() > 0 {
   +    if stream_values[0].logical_null_count() > 0 {
   ```
   
   In the test, after importing `batches_to_sort_string`. This fails on this 
head and passes with the fix:
   ```rs
   let (_, batches, _) =
       join_collect(left, right, on, Operator::Lt, JoinType::Right).await?;
   // The NULL key (a2 = 0) and 0 (a2 = 3) match nothing.
   assert_snapshot!(batches_to_sort_string(&batches), @r"
   +----+----+----+----+
   | a1 | b1 | a2 | b2 |
   +----+----+----+----+
   |    |    | 0  |    |
   |    |    | 3  | 0  |
   | 1  | 5  | 4  | 6  |
   | 2  | 5  | 4  | 6  |
   | 3  | 1  | 1  | 2  |
   | 3  | 1  | 2  | 2  |
   | 3  | 1  | 4  | 6  |
   | 4  | 1  | 1  | 2  |
   | 4  | 1  | 2  | 2  |
   | 4  | 1  | 4  | 6  |
   +----+----+----+----+
   ");
   ```



##########
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:
   This PR also fixes dictionary keys whose *values* are NULL (a valid key 
pointing at a NULL value). `apply_cmp` sees the logical NULL and sets the row 
aside, while `main` joins it with every buffered row. Nothing pins this, and 
the oracle here uses physical `is_valid`, so it would call such a row 
matchable. Fine as a follow-up; this passes on the current head:
   
   ```diff
   +                    let nulls = streamed.logical_nulls();
                        for row in 0..streamed.len() {
   -                        let expected = streamed.is_valid(row)
   +                        let expected = nulls.as_ref().is_none_or(|n| 
n.is_valid(row))
   ```
   ```rs
   (
       // A valid key pointing at a NULL value: logically NULL, physically 
valid.
       "dictionary_null_values",
       Arc::new(DictionaryArray::<Int32Type>::new(
           Int32Array::from(vec![0, 1, 2]),
           Arc::new(Int32Array::from(vec![Some(5), None, Some(1)])),
       )),
       Arc::new(DictionaryArray::<Int32Type>::new(
           Int32Array::from(vec![0, 1, 2, 3]),
           Arc::new(Int32Array::from(vec![Some(2), None, Some(6), Some(0)])),
       )),
   ),
   ```



-- 
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