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]