Skip to content
Merged
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
4 changes: 2 additions & 2 deletions datafusion/core/tests/fuzz_cases/join_fuzz.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1604,14 +1604,14 @@ fn pwmj_plan(
join_type: JoinType,
) -> Arc<dyn ExecutionPlan> {
// Matches `PiecewiseMergeJoinExec::required_input_ordering`: descending for `<`/`<=`,
// ascending for `>`/`>=`, NULLs first either way. Right existence joins require no
// ascending for `>`/`>=`, reversing NULL placement too. Right existence joins require no
// ordering at all -- they only read the buffered side's min/max -- so they are fed the
// left side unsorted, which is the input shape they will see in a real plan.
let buffered = match join_type {
JoinType::RightSemi | JoinType::RightAnti | JoinType::RightMark => left,
_ => {
let sort_options = match op {
Operator::Lt | Operator::LtEq => SortOptions::new(true, true),
Operator::Lt | Operator::LtEq => SortOptions::new(true, false),
Operator::Gt | Operator::GtEq => SortOptions::new(false, true),
other => panic!("not a range operator: {other:?}"),
};
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,12 +17,12 @@

//! Stream Implementation for PiecewiseMergeJoin's Classic Join (Left, Right, Full, Inner)

use arrow::array::{Array, PrimitiveBuilder, new_null_array};
use arrow::array::{PrimitiveBuilder, new_null_array};
use arrow::compute::{BatchCoalescer, take};
use arrow::datatypes::UInt32Type;
use arrow::{
array::{ArrayRef, RecordBatch, UInt32Array},
compute::{sort_to_indices, take_record_batch},
compute::{SortColumn, lexsort_to_indices, take_record_batch},
};
use arrow_schema::{Schema, SchemaRef, SortOptions};
use datafusion_common::NullEquality;
Expand All @@ -37,7 +37,9 @@ use std::{sync::Arc, task::Poll};

use crate::handle_state;
use crate::joins::piecewise_merge_join::exec::{BufferedSide, BufferedSideReadyState};
use crate::joins::piecewise_merge_join::utils::need_produce_result_in_final;
use crate::joins::piecewise_merge_join::utils::{
need_produce_result_in_final, non_null_range, unmatched_buffered_batch,
};
use crate::joins::utils::JoinKeyComparator;
use crate::joins::utils::{BuildProbeJoinMetrics, StatefulStreamResult};
use crate::stream::EmptyRecordBatchStream;
Expand Down Expand Up @@ -244,9 +246,11 @@ impl ClassicPWMJStream {
self.join_metrics.input_rows.add(batch.num_rows());

// Sort stream values and change the streamed record batch accordingly
let indices = sort_to_indices(
stream_values.as_ref(),
Some(self.sort_option),
let indices = lexsort_to_indices(
&[SortColumn {
values: Arc::clone(&stream_values),
options: Some(self.sort_option),
}],
None,
)?;
let stream_batch = take_record_batch(&batch, &indices)?;
Expand Down Expand Up @@ -336,16 +340,11 @@ impl ClassicPWMJStream {
let buffered_data = Arc::clone(&self.buffered_side.try_as_ready()?.buffered_data);
let buffered_batch = buffered_data.batch();

// Every match marks the suffix `[k, buffered_len)`, so the buffered rows that were
// never matched are exactly the complementary prefix `[0, min_marked)` -- which
// includes the null-keyed rows, since nulls sort first and the scan starts past
// them. That makes the final pass a zero-copy slice instead of building an index
// array and running `take` over it.
let min_marked = buffered_data
.min_marked
.load(AtomicOrdering::SeqCst)
.min(buffered_batch.num_rows());
let new_buffered_batch = buffered_batch.slice(0, min_marked);
let (_, non_null_end) =
non_null_range(buffered_data.values().as_ref(), self.sort_option);
let min_marked = buffered_data.min_marked.load(AtomicOrdering::SeqCst);
let new_buffered_batch =
unmatched_buffered_batch(buffered_batch, min_marked, non_null_end)?;
let mut buffered_columns = new_buffered_batch.columns().to_vec();

let streamed_columns: Vec<ArrayRef> = self
Expand Down Expand Up @@ -444,7 +443,8 @@ fn resolve_classic_join(
join_type: JoinType,
batch_process_state: &mut BatchProcessState,
) -> Result<RecordBatch> {
let buffered_len = buffered_side.buffered_data.values().len();
let (buffered_start, buffered_len) =
non_null_range(buffered_side.buffered_data.values().as_ref(), sort_options);
let stream_values = stream_batch.compare_key_values();

// Build comparator once for the batch pair
Expand All @@ -459,26 +459,28 @@ fn resolve_classic_join(
let mut stream_idx = batch_process_state.start_stream_idx;

if !batch_process_state.processed_null_count {
let buffered_null_idx = buffered_side.buffered_data.values().null_count();
let stream_null_idx = stream_values[0].null_count();
buffer_idx = buffered_null_idx;
stream_idx = stream_null_idx;
let (stream_start, stream_end) =
non_null_range(stream_values[0].as_ref(), sort_options);
buffer_idx = buffered_start;
stream_idx = stream_start;
batch_process_state.processed_null_count = true;

// The scan below starts past the streamed side's NULL-keyed rows, which
// sit at the front (`nulls_first`). A NULL join key never matches under
// `NullEqualsNothing`, so for `Right`/`Full` those rows are unmatched and
// must still be emitted; record them here since the scan will skip them.
// NULL keys never match, but outer joins must emit them on either end.
if matches!(join_type, JoinType::Right | JoinType::Full) {
for row_idx in 0..stream_null_idx as u32 {
batch_process_state.unmatched_indices.append_value(row_idx);
for row_idx in
(0..stream_start).chain(stream_end..stream_batch.batch.num_rows())
{
batch_process_state
.unmatched_indices
.append_value(row_idx as u32);
}
}
}

// Our buffer_idx variable allows us to start probing on the buffered side where we last matched
// in the previous stream row.
for row_idx in stream_idx..stream_batch.batch.num_rows() {
let (_, stream_end) = non_null_range(stream_values[0].as_ref(), sort_options);
for row_idx in stream_idx..stream_end {
while buffer_idx < buffered_len {
let compare = cmp.compare(row_idx, buffer_idx);

Expand Down
12 changes: 6 additions & 6 deletions datafusion/physical-plan/src/joins/piecewise_merge_join/exec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -330,7 +330,7 @@ impl PiecewiseMergeJoinExec {
// Take the operator and enforce a sort order on the streamed + buffered side based on
// the operator type.
let sort_options = match operator {
Operator::Lt | Operator::LtEq => SortOptions::new(true, true),
Operator::Lt | Operator::LtEq => SortOptions::new(true, false),
Operator::Gt | Operator::GtEq => SortOptions::new(false, true),
_ => {
return internal_err!(
Expand Down Expand Up @@ -1034,14 +1034,14 @@ pub(super) struct BufferedSideData {
pub(super) batch: RecordBatch,
values: ArrayRef,
pub(super) remaining_partitions: AtomicUsize,
/// The start of the matched suffix of the buffered side, or `usize::MAX` before the
/// first match. `[min_marked, len)` *is* the matched set and `[0, min_marked)` the
/// unmatched one -- no bitmap is allocated.
/// The start of the matched non-null suffix, or `usize::MAX` before the
/// first match. `[min_marked, non_null_end)` is the matched set; the prefix
/// and any trailing NULL keys are unmatched. No bitmap is allocated.
///
/// Both stream kinds only ever mark a suffix, which is what makes one index enough:
/// - `ExistencePWMJStream` marks `[k, len)` for the first buffered row `k` matching
/// - `ExistencePWMJStream` marks `[k, non_null_end)` for the first row `k` matching
/// a streamed batch's extreme key.
/// - `ClassicPWMJStream` emits `buffered[k..] x streamed_row` on each match, so the
/// - `ClassicPWMJStream` emits `buffered[k..non_null_end] x streamed_row`, so the
/// rows it marks are exactly that same suffix.
///
/// Shared so each partition benefits from what the others have marked; it only ever
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -65,22 +65,22 @@
//! so a partition that starts late reads nothing at all.
//!
//! Rows whose join key is NULL never satisfy a comparison predicate. Buffered NULLs sort to
//! the front, so the scan starts past them and null-keyed buffered rows are left unmarked —
//! either end, outside the searched range, and null-keyed buffered rows are left unmarked —
//! correctly excluded from `LeftSemi` and included in `LeftAnti`. The extreme key is picked
//! from each streamed batch with NULLs ignored, so it is non-null unless the whole batch is.
//!
//! # Output
//!
//! Marking only ever covers a suffix, and each mark lowers the watermark to its own start,
//! so the matched set is always exactly `[min_marked, buffered_len)`. A bitmap would be a
//! so the matched set is always exactly `[min_marked, non_null_end)`. A bitmap would be a
//! less compact encoding of that one index, so none is allocated. `ClassicPWMJStream` marks
//! the same way, which is why the watermark lives in `BufferedSideData` rather than here.
//!
//! Once every streamed partition has been consumed, the last one to finish slices the
//! buffered batch: `LeftSemi` takes `[min_marked, len)`, `LeftAnti` the complementary
//! prefix `[0, min_marked)`, which is where the null-keyed rows live. `LeftMark` takes
//! buffered batch: `LeftSemi` takes `[min_marked, non_null_end)`, `LeftAnti` the complementary
//! prefix `[0, min_marked)` plus any trailing null-keyed rows. `LeftMark` takes
//! neither slice: every buffered row is preserved, with a `mark` column built from the same
//! watermark (`true` from `min_marked` on, `false` before it) appended instead of any row
//! watermark (`true` in the matched range, `false` elsewhere) appended instead of any row
//! being dropped. Only the buffered (left) columns, plus that `mark` column for `LeftMark`,
//! are produced.
//!
Expand All @@ -91,8 +91,12 @@ use std::sync::Arc;
use std::sync::atomic::Ordering as AtomicOrdering;
use std::task::{Poll, ready};

use arrow::array::{Array, ArrayRef, BooleanArray, BooleanBufferBuilder, RecordBatch};
use arrow::compute::BatchCoalescer;
use arrow::array::{
Array, ArrayRef, BooleanArray, BooleanBufferBuilder, RecordBatch, UInt32Array,
new_null_array,
};
use arrow::compute::{BatchCoalescer, take};
use arrow_ord::ord::make_comparator;
use arrow_schema::{SchemaRef, SortOptions};
use datafusion_common::{NullEquality, Result, internal_err};
use datafusion_execution::{RecordBatchStream, SendableRecordBatchStream};
Expand All @@ -101,6 +105,7 @@ use datafusion_functions_aggregate_common::min_max::{max_batch, min_batch};
use datafusion_physical_expr::PhysicalExprRef;
use futures::{Stream, StreamExt};

use super::utils::{non_null_range, unmatched_buffered_batch};
use crate::handle_state;
use crate::joins::piecewise_merge_join::exec::{BufferedSide, BufferedSideReadyState};
use crate::joins::utils::{
Expand Down Expand Up @@ -265,7 +270,8 @@ impl ExistencePWMJStream {
let min_marked = buffered_data.min_marked.load(AtomicOrdering::SeqCst);
let buffered_values = buffered_data.values();

Ok(min_marked.min(buffered_values.len()) <= buffered_values.null_count())
let (start, end) = non_null_range(buffered_values.as_ref(), self.sort_option);
Ok(min_marked.min(end) <= start)
}

/// Marks this partition done with the streamed side: releases the input pipeline and,
Expand Down Expand Up @@ -301,20 +307,18 @@ impl ExistencePWMJStream {
{
let buffered_data = &self.buffered_side.try_as_ready()?.buffered_data;
let buffered_values = buffered_data.values();
let buffered_len = buffered_values.len();

// NULL keys can never match, and `sort_options` uses `nulls_first` for every
// operator (see `try_new`), so buffered nulls sit at the front -- skip past them.
let first_non_null_buffered = buffered_values.null_count();
let (first_non_null_buffered, non_null_end) =
non_null_range(buffered_values.as_ref(), sort_option);

// `[min_marked, buffered_len)` was already marked, by this partition or
// `[min_marked, non_null_end)` was already marked, by this partition or
// another, so a match found there would write nothing. Stop the scan at the
// watermark: that bounds the comparisons this batch performs, not just the
// bits it writes.
let scan_limit = buffered_data
.min_marked
.load(AtomicOrdering::SeqCst)
.min(buffered_len);
.min(non_null_end);

// The extreme key is the only one that can decide anything: it reaches the
// smallest matching `buffer_idx`, and every other row in the batch matches a
Expand Down Expand Up @@ -370,7 +374,7 @@ impl ExistencePWMJStream {
if buffer_idx < scan_limit {
// Everything from `buffer_idx` on matches, so lowering the
// watermark to it records the match: the marked set is exactly
// `[min_marked, buffered_len)` and needs no bitmap.
// `[min_marked, non_null_end)` and needs no bitmap.
//
// INVARIANT: sound only because the buffered side and each
// streamed batch are sorted the same way for this operator
Expand Down Expand Up @@ -403,42 +407,43 @@ impl ExistencePWMJStream {
let buffered_batch = buffered_data.batch();
let buffered_len = buffered_batch.num_rows();

// The marked rows are always the contiguous suffix `[min_marked, len)`: each
// match covers `[k, previous min_marked)` and then lowers the watermark to
// `k`, so the union is `[k, len)`. The result is therefore a slice, with no
// index array to materialize and no `take`.
// The matched non-null suffix is `[min_marked, non_null_end)`.
// Each match lowers the watermark; trailing NULL keys stay unmarked.
let (_, non_null_end) =
non_null_range(buffered_data.values().as_ref(), self.sort_option);
let min_marked = buffered_data
.min_marked
.load(AtomicOrdering::SeqCst)
.min(buffered_len);
.min(non_null_end);

let (num_rows, columns) = match self.join_type {
JoinType::LeftSemi => {
let sliced =
buffered_batch.slice(min_marked, buffered_len - min_marked);
buffered_batch.slice(min_marked, non_null_end - min_marked);
(sliced.num_rows(), sliced.columns().to_vec())
}
// `LeftMark` keeps every buffered row -- nothing to slice -- and appends
// the watermark as a `mark` column instead of using it to drop rows: `false`
// for the unmatched prefix `[0, min_marked)`, `true` for the matched suffix
// `[min_marked, len)` -- the same split `LeftSemi`/`LeftAnti` slice the
// buffered batch on, just kept as one column instead of used to drop rows.
// the watermark as a `mark` column: true only within the matched
// non-null range, false for the prefix and any trailing NULL keys.
JoinType::LeftMark => {
let mut mark = BooleanBufferBuilder::new(buffered_len);
mark.append_n(min_marked, false);
mark.append_n(buffered_len - min_marked, true);
mark.append_n(non_null_end - min_marked, true);
mark.append_n(buffered_len - non_null_end, false);

let mut columns = buffered_batch.columns().to_vec();
columns.push(
Arc::new(BooleanArray::new(mark.finish(), None)) as ArrayRef
);
(buffered_len, columns)
}
// `LeftAnti`: the unmarked prefix, which includes every null-keyed row --
// nulls sort first and the watermark never drops below the buffered null
// count.
// `LeftAnti` includes the unmarked prefix and any trailing NULL keys.
JoinType::LeftAnti => {
let sliced = buffered_batch.slice(0, min_marked);
let sliced = unmatched_buffered_batch(
buffered_batch,
min_marked,
non_null_end,
)?;
(sliced.num_rows(), sliced.columns().to_vec())
}
other => {
Expand Down Expand Up @@ -478,10 +483,36 @@ impl ExistencePWMJStream {
/// Ordered the same way as [`JoinKeyComparator`]: both use IEEE 754 totalOrder for floats,
/// and the comparator normalizes `-0.0` on either side of it.
///
/// Numeric, temporal, string, binary and boolean keys get a typed arrow kernel -- a linear
/// scan that allocates nothing. Dictionary and nested keys fall to `min_max_batch_generic`, a
/// `ScalarValue`-per-row comparator loop; specializing those is left to a follow-up.
/// Nested keys use Arrow's comparator, so inner NULLs have the same ordering as
/// sorting and range comparisons. Other keys use the typed min/max kernels.
pub(super) fn extreme_key(values: &ArrayRef, descending: bool) -> Result<ArrayRef> {
if values.data_type().is_nested() {
// Reverse both value and inner-NULL ordering for max, matching SQL's
// ascending NULLS FIRST nested comparison order.
let cmp = make_comparator(
values.as_ref(),
values.as_ref(),
SortOptions::new(descending, !descending),
)?;
let nulls = values.logical_nulls();
let mut extreme = None;
for idx in 0..values.len() {
if nulls.as_ref().is_some_and(|n| n.is_null(idx)) {
continue;
}
if extreme.is_none_or(|current| cmp(idx, current) == Ordering::Less) {
extreme = Some(idx);
}
}
return match extreme {
Some(idx) => Ok(take(
values.as_ref(),
&UInt32Array::from(vec![idx as u32]),
None,
)?),
None => Ok(new_null_array(values.data_type(), 1)),
};
}
let extreme = if descending {
max_batch(values)?
} else {
Expand Down
28 changes: 28 additions & 0 deletions datafusion/physical-plan/src/joins/piecewise_merge_join/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,8 +15,36 @@
// specific language governing permissions and limitations
// under the License.

use arrow::array::{Array, RecordBatch};
use arrow::compute::concat_batches;
use arrow_schema::SortOptions;
use datafusion_common::Result;
use datafusion_expr::JoinType;

/// Bounds of the non-null keys in an array sorted with `options`.
pub(super) fn non_null_range(values: &dyn Array, options: SortOptions) -> (usize, usize) {
let null_count = values.logical_null_count();
if options.nulls_first {
(null_count, values.len())
} else {
(0, values.len() - null_count)
}
}

/// Keep the unmatched prefix and any trailing NULL keys outside the matching range.
pub(super) fn unmatched_buffered_batch(
batch: &RecordBatch,
min_marked: usize,
non_null_end: usize,
) -> Result<RecordBatch> {
let prefix = batch.slice(0, min_marked.min(non_null_end));
if non_null_end == batch.num_rows() {
return Ok(prefix);
}
let nulls = batch.slice(non_null_end, batch.num_rows() - non_null_end);
Ok(concat_batches(&batch.schema(), [&prefix, &nulls])?)
}

// Returns boolean for whether the join is a right existence join served by
// `RightExistencePWMJStream`, which reads nothing but a single min/max off the buffered side.
//
Expand Down
Loading
Loading