From eedede1a6361b7b733e8cdeff20c782807161b35 Mon Sep 17 00:00:00 2001 From: msaf Date: Wed, 7 Oct 2026 22:30:46 +0100 Subject: [PATCH 01/13] perf: evaluate a sort-merge join filter before materializing the pairs A sort-merge join with a join filter materialized every column of every candidate pair (take over the streamed columns, interleave over the buffered ones), picked the filter columns out of that batch, and for LEFT/RIGHT/FULL joins pushed the whole wide batch, failing pairs included, through the deferred-filtering pipeline (concat, corrected mask, filter). A validity-interval join over large key groups, where almost every pair fails, spent nearly all its time copying rows it then dropped. freeze_streamed_matched now gathers only the filter columns, evaluates the filter, updates the FULL join's buffered filter state from the full mask, and materializes the output columns only for the pairs that can reach the output, reusing the gathered filter columns when every pair is kept. INNER joins keep the passing pairs. Deferred joins keep, within each run of one streamed row's pairs in a freeze, the passing pairs, or the run's last pair when none passed: across the freezes a row spans this keeps all its passing pairs and its overall last pair, which is all get_corrected_filter_mask outputs from it, so the output and its order are unchanged. Semi/anti/mark joins already evaluate the filter on the filter columns alone. Validity-interval bench (100 streamed columns, 11000 buffered rows per key, one passing pair per streamed row, 4.4M pairs): LEFT 471 -> 8.5, FULL 378 -> 12.6, INNER 119 -> 8.0 ns per pair. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../benches/sort_merge_join.rs | 194 +++++- .../src/joins/sort_merge_join/filter.rs | 92 +-- .../sort_merge_join/materializing_stream.rs | 409 ++++++++++--- .../src/joins/sort_merge_join/tests.rs | 572 +++++++++++++++++- .../apache/comet/exec/CometJoinSuite.scala | 56 ++ 5 files changed, 1195 insertions(+), 128 deletions(-) diff --git a/native/vendor/datafusion-physical-plan/benches/sort_merge_join.rs b/native/vendor/datafusion-physical-plan/benches/sort_merge_join.rs index 82610b2a54..8b0c3a1d72 100644 --- a/native/vendor/datafusion-physical-plan/benches/sort_merge_join.rs +++ b/native/vendor/datafusion-physical-plan/benches/sort_merge_join.rs @@ -23,14 +23,16 @@ use std::sync::Arc; -use arrow::array::{Int64Array, RecordBatch, StringArray}; +use arrow::array::{ArrayRef, Int64Array, RecordBatch, StringArray}; use arrow::compute::SortOptions; use arrow::datatypes::{DataType, Field, Schema, SchemaRef}; use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main}; -use datafusion_common::NullEquality; +use datafusion_common::{JoinSide, NullEquality}; use datafusion_execution::TaskContext; -use datafusion_physical_expr::expressions::col; +use datafusion_expr::Operator; +use datafusion_physical_expr::expressions::{BinaryExpr, Column, col}; use datafusion_physical_plan::collect; +use datafusion_physical_plan::joins::utils::{ColumnIndex, JoinFilter}; use datafusion_physical_plan::joins::{SortMergeJoinExec, utils::JoinOn}; use datafusion_physical_plan::test::TestMemoryExec; use tokio::runtime::Runtime; @@ -200,5 +202,189 @@ fn bench_smj(c: &mut Criterion) { group.finish(); } -criterion_group!(benches, bench_smj); +/// Streamed side of the validity-interval join: `key`, the lookup time `t` +/// and `payload_cols` payload columns alternating Int64 and Utf8. +fn build_interval_streamed( + keys: usize, + rows_per_key: usize, + group_rows: usize, + payload_cols: usize, +) -> (SchemaRef, Vec) { + let mut fields = vec![ + Field::new("key", DataType::Int64, false), + Field::new("t", DataType::Int64, false), + ]; + for c in 0..payload_cols { + let data_type = if c % 2 == 0 { + DataType::Int64 + } else { + DataType::Utf8 + }; + fields.push(Field::new(format!("p{c}"), data_type, false)); + } + let schema = Arc::new(Schema::new(fields)); + + let num_rows = keys * rows_per_key; + let key: Vec = (0..num_rows).map(|i| (i / rows_per_key) as i64).collect(); + let t: Vec = (0..num_rows) + .map(|i| ((i * 7919) % (group_rows * 10)) as i64 + 1) + .collect(); + let mut columns: Vec = vec![ + Arc::new(Int64Array::from(key)), + Arc::new(Int64Array::from(t)), + ]; + for c in 0..payload_cols { + if c % 2 == 0 { + columns.push(Arc::new(Int64Array::from_iter_values( + (0..num_rows).map(|i| (i * c) as i64), + ))); + } else { + columns.push(Arc::new(StringArray::from_iter_values( + (0..num_rows).map(|i| format!("payload_{c}_{i:012}")), + ))); + } + } + let batch = RecordBatch::try_new(Arc::clone(&schema), columns).unwrap(); + (schema, split_batch(&batch, 8192)) +} + +/// Buffered side of the validity-interval join: `group_rows` consecutive +/// intervals `(eff, next_eff]` per key. +fn build_interval_buffered( + keys: usize, + group_rows: usize, +) -> (SchemaRef, Vec) { + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Int64, false), + Field::new("eff", DataType::Int64, false), + Field::new("next_eff", DataType::Int64, false), + Field::new("rate", DataType::Utf8, false), + ])); + let num_rows = keys * group_rows; + let key: Vec = (0..num_rows).map(|i| (i / group_rows) as i64).collect(); + let eff: Vec = (0..num_rows) + .map(|i| ((i % group_rows) * 10) as i64) + .collect(); + let next_eff: Vec = eff.iter().map(|e| e + 10).collect(); + let rate = StringArray::from_iter_values((0..num_rows).map(|i| format!("rate_{i}"))); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(Int64Array::from(key)), + Arc::new(Int64Array::from(eff)), + Arc::new(Int64Array::from(next_eff)), + Arc::new(rate), + ], + ) + .unwrap(); + (schema, split_batch(&batch, 8192)) +} + +fn split_batch(batch: &RecordBatch, batch_size: usize) -> Vec { + (0..batch.num_rows()) + .step_by(batch_size) + .map(|offset| batch.slice(offset, (batch.num_rows() - offset).min(batch_size))) + .collect() +} + +/// `streamed.t > buffered.eff AND streamed.t <= buffered.next_eff` +fn interval_filter(streamed: &Schema, buffered: &Schema) -> JoinFilter { + let t = Arc::new(Column::new("t", 0)); + let expression = Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new( + Arc::clone(&t) as _, + Operator::Gt, + Arc::new(Column::new("eff", 1)), + )), + Operator::And, + Arc::new(BinaryExpr::new( + t, + Operator::LtEq, + Arc::new(Column::new("next_eff", 2)), + )), + )); + JoinFilter::new( + expression, + vec![ + ColumnIndex { + index: 1, + side: JoinSide::Left, + }, + ColumnIndex { + index: 1, + side: JoinSide::Right, + }, + ColumnIndex { + index: 2, + side: JoinSide::Right, + }, + ], + Arc::new(Schema::new(vec![ + streamed.field(1).clone(), + buffered.field(1).clone(), + buffered.field(2).clone(), + ])), + ) +} + +/// Validity-interval join where every streamed row passes the filter for one +/// buffered row of its key group: over large key groups, where almost every +/// pair fails, and over unique keys, where every pair passes. +fn bench_smj_filter(c: &mut Criterion) { + let rt = Runtime::new().unwrap(); + let mut group = c.benchmark_group("sort_merge_join_filter"); + group.sample_size(10); + for (shape, keys, rows_per_key, group_rows) in [ + ("interval", 4, 100, 11_000), + ("interval_unique", 100_000, 1, 1), + ] { + let (streamed_schema, streamed_batches) = + build_interval_streamed(keys, rows_per_key, group_rows, 98); + let (buffered_schema, buffered_batches) = + build_interval_buffered(keys, group_rows); + let pairs = keys * rows_per_key * group_rows; + for join_type in [ + datafusion_common::JoinType::Left, + datafusion_common::JoinType::Full, + datafusion_common::JoinType::Inner, + ] { + group.bench_function( + BenchmarkId::new(format!("{shape}_{join_type:?}"), pairs), + |b| { + b.iter(|| { + let left = make_exec(&streamed_batches, &streamed_schema); + let right = make_exec(&buffered_batches, &buffered_schema); + let on: JoinOn = vec![( + col("key", &streamed_schema).unwrap(), + col("key", &buffered_schema).unwrap(), + )]; + let join = SortMergeJoinExec::try_new( + left, + right, + on, + Some(interval_filter(&streamed_schema, &buffered_schema)), + join_type, + vec![SortOptions::default()], + NullEquality::NullEqualsNothing, + ) + .unwrap(); + let task_ctx = Arc::new(TaskContext::default()); + let rows: usize = rt.block_on(async { + let batches = + collect(Arc::new(join), task_ctx).await.unwrap(); + batches.iter().map(|b| b.num_rows()).sum() + }); + if join_type != datafusion_common::JoinType::Full { + assert_eq!(rows, keys * rows_per_key); + } + rows + }) + }, + ); + } + } + group.finish(); +} + +criterion_group!(benches, bench_smj, bench_smj_filter); criterion_main!(benches); diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs index 4fc6cccaa8..de3832fb71 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs @@ -26,13 +26,13 @@ use std::sync::Arc; use arrow::array::{ - Array, ArrayBuilder, ArrayRef, BooleanArray, BooleanBuilder, RecordBatch, - RecordBatchOptions, UInt64Array, UInt64Builder, new_null_array, + Array, ArrayBuilder, ArrayRef, BooleanArray, BooleanBufferBuilder, BooleanBuilder, + RecordBatch, RecordBatchOptions, UInt64Array, UInt64Builder, new_null_array, }; use arrow::compute::kernels::zip::zip; use arrow::compute::{self, filter_record_batch}; use arrow::datatypes::SchemaRef; -use datafusion_common::{JoinSide, JoinType, Result}; +use datafusion_common::{JoinType, Result}; use crate::joins::utils::JoinFilter; @@ -145,37 +145,6 @@ pub fn needs_deferred_filtering( && matches!(join_type, JoinType::Left | JoinType::Right | JoinType::Full) } -/// Gets the arrays which join filters are applied on -/// -/// Extracts the columns needed for filter evaluation from left and right batch columns -pub fn get_filter_columns( - join_filter: &Option, - left_columns: &[ArrayRef], - right_columns: &[ArrayRef], -) -> Vec { - let mut filter_columns = vec![]; - - if let Some(f) = join_filter { - let left_columns: Vec = f - .column_indices() - .iter() - .filter(|col_index| col_index.side == JoinSide::Left) - .map(|i| Arc::clone(&left_columns[i.index])) - .collect(); - let right_columns: Vec = f - .column_indices() - .iter() - .filter(|col_index| col_index.side == JoinSide::Right) - .map(|i| Arc::clone(&right_columns[i.index])) - .collect(); - - filter_columns.extend(left_columns); - filter_columns.extend(right_columns); - } - - filter_columns -} - /// Determines if current index is the last occurrence of a row /// /// Used during filter mask correction to detect row boundaries when grouping @@ -294,6 +263,61 @@ pub fn get_corrected_filter_mask( } } +/// Selects the pairs of one freeze that deferred filtering may output +/// +/// For each streamed row, `get_corrected_filter_mask` keeps every pair that +/// passed the filter, or null-joins the row's last pair when none did, and +/// discards the rest. A row's pairs can span several freezes, so a freeze +/// cannot tell which of its failing pairs ends up null-joined. Within each +/// run of one row's pairs it therefore keeps the passing pairs, or only the +/// run's last pair when none passed. Over the runs of one row this keeps +/// all of its passing pairs and its overall last pair, which is all +/// `get_corrected_filter_mask` outputs from it, and the discarded failing +/// pairs never decide which pair is null-joined. +/// +/// # Arguments +/// * `row_indices` - Which streamed row produced each pair (no nulls) +/// * `filter_mask` - Whether each pair passed the filter (no nulls) +/// +/// # Returns +/// A mask that is `true` for the pairs to materialize and keep +pub fn deferred_filter_candidates( + row_indices: &UInt64Array, + filter_mask: &BooleanArray, +) -> BooleanArray { + debug_assert_eq!( + row_indices.len(), + filter_mask.len(), + "row_indices and filter_mask must have same length" + ); + debug_assert_eq!(row_indices.null_count(), 0); + debug_assert_eq!(filter_mask.null_count(), 0); + + let rows = row_indices.values(); + let passed = filter_mask.values(); + let mut keep = BooleanBufferBuilder::new(rows.len()); + let mut run_start = 0; + for i in 0..rows.len() { + if i + 1 < rows.len() && rows[i + 1] == rows[i] { + continue; + } + let run_len = i + 1 - run_start; + if run_len == 1 { + keep.append(true); + } else { + let run = passed.slice(run_start, run_len); + if run.count_set_bits() > 0 { + keep.append_buffer(&run); + } else { + keep.append_n(run_len - 1, false); + keep.append(true); + } + } + run_start = i + 1; + } + BooleanArray::new(keep.finish(), None) +} + /// Applies corrected filter mask to record batch based on join type /// /// The corrected mask has three possible values per row: diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs index 43306248d0..b9f48bac0d 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs @@ -30,8 +30,8 @@ use std::ops::Range; use std::sync::Arc; use crate::joins::sort_merge_join::filter::{ - FilterMetadata, filter_record_batch_by_join_type, get_corrected_filter_mask, - get_filter_columns, needs_deferred_filtering, + FilterMetadata, deferred_filter_candidates, filter_record_batch_by_join_type, + get_corrected_filter_mask, needs_deferred_filtering, }; use crate::joins::sort_merge_join::metrics::SortMergeJoinMetrics; use crate::joins::utils::{JoinFilter, JoinKeyComparator}; @@ -49,7 +49,7 @@ use arrow::datatypes::SchemaRef; use datafusion_common::cast::as_uint64_array; use datafusion_common::instant::Instant; use datafusion_common::{ - DataFusionError, JoinType, NullEquality, Result, exec_err, internal_err, + DataFusionError, JoinSide, JoinType, NullEquality, Result, exec_err, internal_err, }; use datafusion_execution::memory_pool::MemoryReservation; use datafusion_execution::runtime_env::RuntimeEnv; @@ -1504,10 +1504,15 @@ impl MaterializingSortMergeJoinStream { Ok(()) } - /// Materializes columns, evaluates the join filter, and pushes output + /// Evaluates the join filter, materializes columns, and pushes output /// for all matched chunks in a single batch. This avoids per-chunk /// RecordBatch construction and filter evaluation, which dominates /// cost when keys are near-unique (1 row per chunk). + /// + /// The filter is evaluated on the filter columns alone; the remaining + /// columns are materialized only for the pairs that can reach the output, + /// so a large key group whose pairs mostly fail the filter does not + /// gather every column for every pair. fn freeze_streamed_matched( &mut self, matched_chunks: &[(usize, UInt64Array, UInt64Array)], @@ -1541,98 +1546,275 @@ impl MaterializingSortMergeJoinStream { as_uint64_array(&compute::concat(&refs)?)?.clone() }; - let left_columns = - materialize_left_columns(&self.streamed_batch.batch, &combined_left_indices)?; + let Some(evaluation) = self.evaluate_join_filter( + &combined_left_indices, + matched_chunks, + total_matched_rows, + )? + else { + let output_batch = self.materialize_output_batch( + &combined_left_indices, + matched_chunks, + total_matched_rows, + None, + )?; + self.joined_record_batches + .push_batch_without_metadata(output_batch); + return Ok(()); + }; + let mask = &evaluation.mask; - let right_columns = - self.materialize_right_columns(matched_chunks, total_matched_rows)?; + // Track which buffered rows had all filter matches fail, + // so full join can emit them as null-joined later. + if self.join_type == JoinType::Full { + let mut offset = 0usize; + for (batch_idx, _left, right) in matched_chunks { + let chunk_len = right.len(); + let buffered_batch = &mut self.buffered_data.batches[*batch_idx]; - let filter_columns = if self.join_type == JoinType::Right { - get_filter_columns(&self.filter, &right_columns, &left_columns) + for i in 0..chunk_len { + if right.is_null(i) { + continue; + } + let idx = right.value(i) as usize; + match buffered_batch.join_filter_status[idx] { + FilterState::SomePassed => {} + _ if mask.value(offset + i) => { + buffered_batch.join_filter_status[idx] = + FilterState::SomePassed; + } + _ => { + buffered_batch.join_filter_status[idx] = + FilterState::AllFailed; + } + } + } + offset += chunk_len; + } + debug_assert_eq!( + offset, total_matched_rows, + "offset must advance through every chunk exactly once" + ); + } + + // Deferred filtering outputs a subset of the pairs that + // `deferred_filter_candidates` keeps; inner joins output the pairs + // that passed. + let keep = if self.deferred_filtering { + deferred_filter_candidates(&combined_left_indices, mask) } else { - get_filter_columns(&self.filter, &left_columns, &right_columns) + mask.clone() }; + let num_kept = keep.true_count(); + + if num_kept == total_matched_rows { + let output_batch = self.materialize_output_batch( + &combined_left_indices, + matched_chunks, + total_matched_rows, + Some(&evaluation), + )?; + return self.push_filtered_batch(output_batch, &combined_left_indices, mask); + } + if num_kept == 0 { + return Ok(()); + } - let columns = if self.join_type != JoinType::Right { - [left_columns, right_columns].concat() + let kept_left_indices = + as_uint64_array(&compute::filter(&combined_left_indices, &keep)?)?.clone(); + let mut kept_chunks = Vec::with_capacity(matched_chunks.len()); + let mut offset = 0usize; + for (batch_idx, left, right) in matched_chunks { + let chunk_keep = keep.slice(offset, left.len()); + offset += left.len(); + if chunk_keep.true_count() == 0 { + continue; + } + kept_chunks.push(( + *batch_idx, + as_uint64_array(&compute::filter(left, &chunk_keep)?)?.clone(), + as_uint64_array(&compute::filter(right, &chunk_keep)?)?.clone(), + )); + } + let output_batch = self.materialize_output_batch( + &kept_left_indices, + &kept_chunks, + num_kept, + None, + )?; + let kept_mask = compute::filter(mask, &keep)?; + self.push_filtered_batch(output_batch, &kept_left_indices, kept_mask.as_boolean()) + } + + /// Pushes the pairs of a filtered freeze: with their filter metadata for + /// deferred filtering, or only those that passed for inner joins. + fn push_filtered_batch( + &mut self, + output_batch: RecordBatch, + left_indices: &UInt64Array, + mask: &BooleanArray, + ) -> Result<()> { + if self.deferred_filtering { + self.joined_record_batches.push_batch_with_filter_metadata( + output_batch, + left_indices, + mask, + self.streamed_batch_counter, + self.join_type, + ); + } else if mask.false_count() == 0 { + self.joined_record_batches + .push_batch_without_metadata(output_batch); } else { - [right_columns, left_columns].concat() + let filtered_batch = filter_record_batch(&output_batch, mask)?; + self.joined_record_batches + .push_batch_without_metadata(filtered_batch); + } + Ok(()) + } + + /// Evaluates the join filter over the given pairs on an intermediate + /// batch holding only the filter columns. Returns `None` when the join + /// has no filter columns to evaluate. + fn evaluate_join_filter( + &self, + left_indices: &UInt64Array, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + total_matched_rows: usize, + ) -> Result> { + let Some(filter) = &self.filter else { + return Ok(None); + }; + let (streamed_side, buffered_side) = if self.join_type == JoinType::Right { + (JoinSide::Right, JoinSide::Left) + } else { + (JoinSide::Left, JoinSide::Right) }; - let output_batch = RecordBatch::try_new(Arc::clone(&self.schema), columns)?; - - if !filter_columns.is_empty() { - if let Some(f) = &self.filter { - let filter_batch = - RecordBatch::try_new(Arc::clone(f.schema()), filter_columns)?; - let filter_result = f - .expression() - .evaluate(&filter_batch)? - .into_array(filter_batch.num_rows())?; - - let filter_result_mask = - datafusion_common::cast::as_boolean_array(&filter_result)?; - - // Convert NULL filter results to false β€” NULL means "not satisfied" - // per SQL semantics, same as Left/Right outer joins. - let mask = if filter_result_mask.null_count() > 0 { - compute::prep_null_mask_filter(filter_result_mask) + let side_projection = |side: JoinSide| { + let mut projection: Vec = filter + .column_indices() + .iter() + .filter(|col_index| col_index.side == side) + .map(|col_index| col_index.index) + .collect(); + projection.sort_unstable(); + projection.dedup(); + projection + }; + let streamed_projection = side_projection(streamed_side); + let buffered_projection = side_projection(buffered_side); + if streamed_projection.is_empty() && buffered_projection.is_empty() { + return Ok(None); + } + + let streamed_columns = if streamed_projection.is_empty() { + vec![] + } else { + materialize_left_columns( + &self.streamed_batch.batch.project(&streamed_projection)?, + left_indices, + )? + }; + let buffered_columns = if buffered_projection.is_empty() { + vec![] + } else { + self.materialize_right_columns( + matched_chunks, + total_matched_rows, + Some(&buffered_projection), + )? + }; + + // Left-side filter columns first, then right-side ones. + let filter_columns = [JoinSide::Left, JoinSide::Right] + .into_iter() + .flat_map(|side| { + filter + .column_indices() + .iter() + .filter(move |col_index| col_index.side == side) + }) + .map(|col_index| { + let (projection, columns) = if col_index.side == streamed_side { + (&streamed_projection, &streamed_columns) } else { - filter_result_mask.clone() + (&buffered_projection, &buffered_columns) }; + let pos = projection.binary_search(&col_index.index).unwrap(); + Arc::clone(&columns[pos]) + }) + .collect::>(); - if self.deferred_filtering { - self.joined_record_batches.push_batch_with_filter_metadata( - output_batch, - &combined_left_indices, - &mask, - self.streamed_batch_counter, - self.join_type, - ); - } else { - let filtered_batch = filter_record_batch(&output_batch, &mask)?; - self.joined_record_batches - .push_batch_without_metadata(filtered_batch); - } + let filter_batch = + RecordBatch::try_new(Arc::clone(filter.schema()), filter_columns)?; + let filter_result = filter + .expression() + .evaluate(&filter_batch)? + .into_array(filter_batch.num_rows())?; + let filter_result_mask = + datafusion_common::cast::as_boolean_array(&filter_result)?; + + // Convert NULL filter results to false β€” NULL means "not satisfied" + // per SQL semantics, same as Left/Right outer joins. + let mask = if filter_result_mask.null_count() > 0 { + compute::prep_null_mask_filter(filter_result_mask) + } else { + filter_result_mask.clone() + }; - // Track which buffered rows had all filter matches fail, - // so full join can emit them as null-joined later. - if self.join_type == JoinType::Full { - let mut offset = 0usize; - for (batch_idx, _left, right) in matched_chunks { - let chunk_len = right.len(); - let buffered_batch = &mut self.buffered_data.batches[*batch_idx]; - - for i in 0..chunk_len { - if right.is_null(i) { - continue; - } - let idx = right.value(i) as usize; - match buffered_batch.join_filter_status[idx] { - FilterState::SomePassed => {} - _ if mask.value(offset + i) => { - buffered_batch.join_filter_status[idx] = - FilterState::SomePassed; - } - _ => { - buffered_batch.join_filter_status[idx] = - FilterState::AllFailed; - } - } - } - offset += chunk_len; - } - debug_assert_eq!( - offset, total_matched_rows, - "offset must advance through every chunk exactly once" - ); + Ok(Some(FilterEvaluation { + mask, + streamed_columns: streamed_projection + .into_iter() + .zip(streamed_columns) + .collect(), + buffered_columns: buffered_projection + .into_iter() + .zip(buffered_columns) + .collect(), + })) + } + + /// Materializes the output batch of the given pairs, reusing the filter + /// columns `evaluation` gathered for the same pairs. + fn materialize_output_batch( + &self, + left_indices: &UInt64Array, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + total_matched_rows: usize, + evaluation: Option<&FilterEvaluation>, + ) -> Result { + let left_columns = complete_columns( + self.streamed_schema.fields().len(), + evaluation.map_or(&[], |e| e.streamed_columns.as_slice()), + |projection| match projection { + None => { + materialize_left_columns(&self.streamed_batch.batch, left_indices) } - } - } else { - self.joined_record_batches - .push_batch_without_metadata(output_batch); - } + Some(projection) => materialize_left_columns( + &self.streamed_batch.batch.project(projection)?, + left_indices, + ), + }, + )?; + let right_columns = complete_columns( + self.buffered_schema.fields().len(), + evaluation.map_or(&[], |e| e.buffered_columns.as_slice()), + |projection| { + self.materialize_right_columns( + matched_chunks, + total_matched_rows, + projection, + ) + }, + )?; - Ok(()) + let columns = if self.join_type != JoinType::Right { + [left_columns, right_columns].concat() + } else { + [right_columns, left_columns].concat() + }; + Ok(RecordBatch::try_new(Arc::clone(&self.schema), columns)?) } /// Materializes right-side columns across all matched chunks. @@ -1640,11 +1822,13 @@ impl MaterializingSortMergeJoinStream { /// When chunks reference a single buffered batch, indices are concatenated /// for a single fetch. When multiple batches are involved, `interleave` /// gathers columns across sources. A null-row sentinel at source index 0 - /// handles null right indices (unmatched streamed rows). + /// handles null right indices (unmatched streamed rows). `projection` + /// selects the buffered columns to materialize (all when `None`). fn materialize_right_columns( &self, matched_chunks: &[(usize, UInt64Array, UInt64Array)], total_matched_rows: usize, + projection: Option<&[usize]>, ) -> Result> { let first_batch_idx = matched_chunks[0].0; let single_source = matched_chunks.iter().all(|c| c.0 == first_batch_idx); @@ -1662,6 +1846,7 @@ impl MaterializingSortMergeJoinStream { &self.buffered_data, first_batch_idx, &combined_right_indices, + projection, ); } @@ -1769,8 +1954,6 @@ impl MaterializingSortMergeJoinStream { } } - let num_right_cols = self.buffered_schema.fields().len(); - // Read each source batch once (spilled batches require disk I/O). let source_data: Vec<&RecordBatch> = source_batches .iter() @@ -1794,10 +1977,14 @@ impl MaterializingSortMergeJoinStream { vec![] }; + let col_indices: Vec = match projection { + Some(projection) => projection.to_vec(), + None => (0..self.buffered_schema.fields().len()).collect(), + }; let mut source_arrays: Vec<&dyn Array> = Vec::with_capacity(source_data.len() + source_offset); - let mut right_columns = Vec::with_capacity(num_right_cols); - for col_idx in 0..num_right_cols { + let mut right_columns = Vec::with_capacity(col_indices.len()); + for col_idx in col_indices { source_arrays.clear(); source_arrays.extend(null_arrays.get(col_idx).map(|a| a.as_ref())); source_arrays.extend(source_data.iter().map(|d| d.column(col_idx).as_ref())); @@ -1914,6 +2101,39 @@ fn materialize_left_columns( } } +/// Join filter result for the pairs of one freeze, with the filter columns +/// gathered for it as `(column index, array)` per side. +struct FilterEvaluation { + mask: BooleanArray, + streamed_columns: Vec<(usize, ArrayRef)>, + buffered_columns: Vec<(usize, ArrayRef)>, +} + +/// Builds `num_columns` columns from those already `gathered`, gathering +/// the rest through `gather` (all of them when its projection is `None`). +fn complete_columns( + num_columns: usize, + gathered: &[(usize, ArrayRef)], + gather: impl FnOnce(Option<&[usize]>) -> Result>, +) -> Result> { + if gathered.is_empty() { + return gather(None); + } + let mut columns: Vec> = vec![None; num_columns]; + for (index, array) in gathered { + columns[*index] = Some(Arc::clone(array)); + } + let missing: Vec = (0..num_columns) + .filter(|&index| columns[index].is_none()) + .collect(); + if !missing.is_empty() { + for (index, array) in missing.iter().zip(gather(Some(&missing))?) { + columns[*index] = Some(array); + } + } + Ok(columns.into_iter().map(Option::unwrap).collect()) +} + fn create_unmatched_columns(schema: &SchemaRef, size: usize) -> Vec { schema .fields() @@ -1934,7 +2154,7 @@ fn produce_buffered_null_batch( // Take buffered (right) columns let right_columns = - fetch_right_columns_from_batch_by_idxs(buffered_batch, buffered_indices)?; + fetch_right_columns_from_batch_by_idxs(buffered_batch, buffered_indices, None)?; // Create null streamed (left) columns let mut left_columns = streamed_schema @@ -1981,10 +2201,12 @@ fn fetch_right_columns_by_idxs( buffered_data: &BufferedData, buffered_batch_idx: usize, buffered_indices: &UInt64Array, + projection: Option<&[usize]>, ) -> Result> { fetch_right_columns_from_batch_by_idxs( &buffered_data.batches[buffered_batch_idx], buffered_indices, + projection, ) } @@ -1992,9 +2214,18 @@ fn fetch_right_columns_by_idxs( fn fetch_right_columns_from_batch_by_idxs( buffered_batch: &BufferedBatch, buffered_indices: &UInt64Array, + projection: Option<&[usize]>, ) -> Result> { match &buffered_batch.batch { BufferedBatchState::InMemory(batch) => { + let projected; + let batch = match projection { + Some(projection) => { + projected = batch.project(projection)?; + &projected + } + None => batch, + }; if let Some(range) = is_contiguous_range(buffered_indices) { Ok(batch.slice(range.start, range.len()).columns().to_vec()) } else { diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs index 2300059f6e..e1705ff693 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs @@ -39,7 +39,10 @@ use crate::test::exec::BarrierExec; use crate::test::{build_table_i32, build_table_i32_two_cols}; use crate::{ExecutionPlan, RecordBatchStream, common}; use crate::{ - expressions::Column, joins::sort_merge_join::filter::get_corrected_filter_mask, + expressions::Column, + joins::sort_merge_join::filter::{ + deferred_filter_candidates, get_corrected_filter_mask, + }, joins::sort_merge_join::materializing_stream::JoinedRecordBatches, }; use arrow::array::{ @@ -3182,6 +3185,113 @@ fn build_joined_record_batches() -> Result { Ok(batches) } +/// Output of `get_corrected_filter_mask` over deferred-filter metadata, as +/// `(pair, passed)` per output row: a passing pair, a pair null-joined +/// because its row found no passing pair (`false`), or a null-metadata row. +fn corrected_output( + pairs: &[usize], + rows: &[Option], + mask: &[Option], +) -> Vec<(usize, Option)> { + use arrow::array::Array; + let row_indices = UInt64Array::from(rows.to_vec()); + let filter_mask = BooleanArray::from(mask.to_vec()); + let batch_ids = vec![1; rows.len()]; + let corrected = get_corrected_filter_mask( + Left, + &row_indices, + &batch_ids, + &filter_mask, + rows.len(), + ) + .unwrap(); + (0..rows.len()) + .filter(|&i| corrected.is_valid(i)) + .map(|i| { + let passed = filter_mask.is_valid(i).then(|| corrected.value(i)); + (pairs[i], passed) + }) + .collect() +} + +/// Reducing each freeze's pairs to `deferred_filter_candidates` leaves the +/// corrected output unchanged, wherever the freeze boundaries cut the pairs +/// of one streamed row. +#[test] +fn deferred_filter_candidates_keep_corrected_output() { + let mut state = 11u64; + let mut next = move |n: u64| { + state = state + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + (state >> 33) % n + }; + for _ in 0..2000 { + // Pairs of consecutive streamed rows, with null-metadata rows + // (unmatched rows) between some of them. + let mut rows: Vec> = vec![]; + let mut mask: Vec> = vec![]; + let pass_rate = next(4); + for row in 0..1 + next(6) { + if next(4) == 0 { + rows.push(None); + mask.push(None); + } + for _ in 0..1 + next(12) { + rows.push(Some(row)); + mask.push(Some(next(4) < pass_rate)); + } + } + let pairs: Vec = (0..rows.len()).collect(); + let expected = corrected_output(&pairs, &rows, &mask); + + // Cut the pairs into freezes and reduce each freeze's matched pairs. + let mut kept_pairs = vec![]; + let mut start = 0; + while start < rows.len() { + let end = (start + 1 + next(10) as usize).min(rows.len()); + let mut segment_start = start; + while segment_start < end { + if rows[segment_start].is_none() { + kept_pairs.push(segment_start); + segment_start += 1; + continue; + } + let mut segment_end = segment_start; + while segment_end < end && rows[segment_end].is_some() { + segment_end += 1; + } + let segment = segment_start..segment_end; + let keep = deferred_filter_candidates( + &UInt64Array::from_iter_values( + rows[segment.clone()].iter().map(|r| r.unwrap()), + ), + &BooleanArray::from_iter(mask[segment.clone()].iter().copied()), + ); + kept_pairs.extend(segment.filter(|&i| keep.value(i - segment_start))); + segment_start = segment_end; + } + start = end; + } + let kept_rows: Vec<_> = kept_pairs.iter().map(|&i| rows[i]).collect(); + let kept_mask: Vec<_> = kept_pairs.iter().map(|&i| mask[i]).collect(); + let actual = corrected_output(&kept_pairs, &kept_rows, &kept_mask); + + // A null-joined row may come from any pair of its streamed row: they + // differ only in the buffered side, which is nulled. + let as_rows = |output: Vec<(usize, Option)>| { + output + .into_iter() + .map(|(pair, passed)| match passed { + Some(false) => (rows[pair], None), + _ => (rows[pair], Some(pair)), + }) + .collect::>() + }; + assert_eq!(as_rows(actual), as_rows(expected), "{rows:?} {mask:?}"); + } +} + #[tokio::test] async fn test_left_outer_join_filtered_mask() -> Result<()> { let mut joined_batches = build_joined_record_batches()?; @@ -4187,6 +4297,466 @@ async fn join_wrapped_multi_source_freeze_with_null_buffered_index() -> Result<( Ok(()) } +/// A streamed row of the validity-interval tests: key, row id, lookup time. +type IntervalStreamedRow = (i32, i32, Option); +/// A buffered row of the validity-interval tests: key, row id, `(eff, next_eff]`. +type IntervalBufferedRow = (i32, i32, i64, i64); + +const INTERVAL_PAYLOAD_COLS: usize = 24; + +fn interval_payload(col: usize, sid: i32) -> String { + format!("p{col}_{sid}") +} + +/// Streamed table `(key, sid, t, p0..)` with `INTERVAL_PAYLOAD_COLS` +/// payload columns alternating Int64 and Utf8, split into batches. +fn build_interval_streamed( + rows: &[IntervalStreamedRow], + batch_rows: usize, +) -> Arc { + use arrow::array::{ArrayRef, Int64Array, StringArray}; + let mut fields = vec![ + Field::new("key", DataType::Int32, false), + Field::new("sid", DataType::Int32, false), + Field::new("t", DataType::Int64, true), + ]; + for c in 0..INTERVAL_PAYLOAD_COLS { + let data_type = if c % 2 == 0 { + DataType::Int64 + } else { + DataType::Utf8 + }; + fields.push(Field::new(format!("p{c}"), data_type, false)); + } + let schema = Arc::new(Schema::new(fields)); + let batches = rows + .chunks(batch_rows) + .map(|chunk| { + let mut columns: Vec = vec![ + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.0))), + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.1))), + Arc::new(Int64Array::from_iter(chunk.iter().map(|r| r.2))), + ]; + for c in 0..INTERVAL_PAYLOAD_COLS { + if c % 2 == 0 { + columns.push(Arc::new(Int64Array::from_iter_values( + chunk.iter().map(|r| r.1 as i64 * 1000 + c as i64), + ))); + } else { + columns.push(Arc::new(StringArray::from_iter_values( + chunk.iter().map(|r| interval_payload(c, r.1)), + ))); + } + } + RecordBatch::try_new(Arc::clone(&schema), columns).unwrap() + }) + .collect::>(); + TestMemoryExec::try_new_exec(&[batches], schema, None).unwrap() +} + +/// Buffered table `(bkey, bid, eff, next_eff, val)`, split into batches. +fn build_interval_buffered( + rows: &[IntervalBufferedRow], + batch_rows: usize, +) -> Arc { + use arrow::array::{Int64Array, StringArray}; + let schema = Arc::new(Schema::new(vec![ + Field::new("bkey", DataType::Int32, false), + Field::new("bid", DataType::Int32, false), + Field::new("eff", DataType::Int64, false), + Field::new("next_eff", DataType::Int64, false), + Field::new("val", DataType::Utf8, false), + ])); + let batches = rows + .chunks(batch_rows) + .map(|chunk| { + RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.0))), + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.1))), + Arc::new(Int64Array::from_iter_values(chunk.iter().map(|r| r.2))), + Arc::new(Int64Array::from_iter_values(chunk.iter().map(|r| r.3))), + Arc::new(StringArray::from_iter_values( + chunk.iter().map(|r| format!("v{}", r.1)), + )), + ], + ) + .unwrap() + }) + .collect::>(); + TestMemoryExec::try_new_exec(&[batches], schema, None).unwrap() +} + +/// `t > eff AND t <= next_eff`, with the streamed table on `streamed_side`. +fn build_interval_filter( + streamed: &Schema, + buffered: &Schema, + streamed_side: JoinSide, +) -> JoinFilter { + let t_field = streamed + .field_with_name("t") + .unwrap() + .clone() + .with_nullable(true); + let eff_field = buffered + .field_with_name("eff") + .unwrap() + .clone() + .with_nullable(true); + let next_eff_field = buffered + .field_with_name("next_eff") + .unwrap() + .clone() + .with_nullable(true); + let t = ColumnIndex { + index: 2, + side: streamed_side, + }; + let eff = ColumnIndex { + index: 2, + side: streamed_side.negate(), + }; + let next_eff = ColumnIndex { + index: 3, + side: streamed_side.negate(), + }; + let (column_indices, fields, t_idx, eff_idx, next_eff_idx) = + if streamed_side == JoinSide::Left { + ( + vec![t, eff, next_eff], + vec![t_field, eff_field, next_eff_field], + 0, + 1, + 2, + ) + } else { + ( + vec![eff, next_eff, t], + vec![eff_field, next_eff_field, t_field], + 2, + 0, + 1, + ) + }; + let t_col: PhysicalExprRef = Arc::new(Column::new("t", t_idx)); + JoinFilter::new( + Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new( + Arc::clone(&t_col), + Operator::Gt, + Arc::new(Column::new("eff", eff_idx)), + )), + Operator::And, + Arc::new(BinaryExpr::new( + t_col, + Operator::LtEq, + Arc::new(Column::new("next_eff", next_eff_idx)), + )), + )), + column_indices, + Arc::new(Schema::new(fields)), + ) +} + +/// Expected `(sid, bid)` output of the validity-interval join, in the +/// streamed order the join preserves: each streamed row's passing buffered +/// rows in buffered order, or the row null-joined when none passed. FULL +/// joins append the buffered rows no streamed row passed with. +fn expected_interval_join( + join_type: JoinType, + streamed: &[IntervalStreamedRow], + buffered: &[IntervalBufferedRow], +) -> Vec<(Option, Option)> { + let mut expected = vec![]; + let mut buffered_passed = vec![false; buffered.len()]; + for &(key, sid, t) in streamed { + let mut any = false; + for (b, &(bkey, bid, eff, next_eff)) in buffered.iter().enumerate() { + if bkey == key && t.is_some_and(|t| t > eff && t <= next_eff) { + expected.push((Some(sid), Some(bid))); + buffered_passed[b] = true; + any = true; + } + } + if !any && join_type != Inner { + expected.push((Some(sid), None)); + } + } + if join_type == Full { + for (b, &(_, bid, _, _)) in buffered.iter().enumerate() { + if !buffered_passed[b] { + expected.push((None, Some(bid))); + } + } + } + expected +} + +/// Reads `(sid, bid)` from each output row, checking that every streamed +/// and buffered column belongs to that row (or is null with it). +fn interval_join_output( + batches: &[RecordBatch], + streamed_offset: usize, + buffered_offset: usize, +) -> Vec<(Option, Option)> { + use arrow::array::{Array, AsArray}; + use arrow::datatypes::{Int32Type, Int64Type}; + let mut output = vec![]; + for batch in batches { + let sid = batch + .column(streamed_offset + 1) + .as_primitive::(); + let bid = batch + .column(buffered_offset + 1) + .as_primitive::(); + for row in 0..batch.num_rows() { + let s = sid.is_valid(row).then(|| sid.value(row)); + let b = bid.is_valid(row).then(|| bid.value(row)); + for c in (0..3 + INTERVAL_PAYLOAD_COLS).filter(|&c| c != 2) { + let col = batch.column(streamed_offset + c); + assert_eq!(col.is_valid(row), s.is_some(), "column {c}"); + } + if let Some(s) = s { + for c in 0..INTERVAL_PAYLOAD_COLS { + let col = batch.column(streamed_offset + 3 + c); + if c % 2 == 0 { + assert_eq!( + col.as_primitive::().value(row), + s as i64 * 1000 + c as i64 + ); + } else { + assert_eq!( + col.as_string::().value(row), + interval_payload(c, s) + ); + } + } + } + for c in 0..5 { + assert_eq!(batch.column(buffered_offset + c).is_valid(row), b.is_some()); + } + if let Some(b) = b { + assert_eq!( + batch + .column(buffered_offset + 4) + .as_string::() + .value(row), + format!("v{b}") + ); + } + output.push((s, b)); + } + } + output +} + +/// Runs the validity-interval join with the streamed table on the outer +/// side of `join_type` (the right side for RIGHT joins) and checks it +/// against [`expected_interval_join`]: in order where the streamed order is +/// preserved, as a multiset for FULL joins. +#[expect(clippy::too_many_arguments)] +async fn check_interval_join( + join_type: JoinType, + streamed: &[IntervalStreamedRow], + buffered: &[IntervalBufferedRow], + streamed_batch_rows: usize, + buffered_batch_rows: usize, + batch_size: usize, + spill: bool, + case: &str, +) -> Result<()> { + let streamed_plan = build_interval_streamed(streamed, streamed_batch_rows); + let buffered_plan = build_interval_buffered(buffered, buffered_batch_rows); + let streamed_side = if join_type == Right { + JoinSide::Right + } else { + JoinSide::Left + }; + let filter = build_interval_filter( + &streamed_plan.schema(), + &buffered_plan.schema(), + streamed_side, + ); + let (left, right) = if join_type == Right { + (buffered_plan, streamed_plan) + } else { + (streamed_plan, buffered_plan) + }; + let on = vec![( + Arc::new(Column::new_with_schema( + if join_type == Right { "bkey" } else { "key" }, + &left.schema(), + )?) as _, + Arc::new(Column::new_with_schema( + if join_type == Right { "key" } else { "bkey" }, + &right.schema(), + )?) as _, + )]; + let (streamed_offset, buffered_offset) = if join_type == Right { + (5, 0) + } else { + (0, 3 + INTERVAL_PAYLOAD_COLS) + }; + + let mut task_ctx = TaskContext::default() + .with_session_config(SessionConfig::default().with_batch_size(batch_size)); + if spill { + task_ctx = task_ctx.with_runtime( + RuntimeEnvBuilder::new() + .with_memory_limit(100, 1.0) + .with_disk_manager_builder( + DiskManagerBuilder::default() + .with_mode(DiskManagerMode::OsTmpDirectory), + ) + .build_arc()?, + ); + } + let join = join_with_filter( + left, + right, + on, + filter, + join_type, + vec![SortOptions::default()], + NullEquality::NullEqualsNothing, + )?; + let batches = common::collect(join.execute(0, Arc::new(task_ctx))?).await?; + if spill && !buffered.is_empty() { + assert!( + join.metrics().unwrap().spill_count().unwrap() > 0, + "{case}: expected spilling" + ); + } + + let mut actual = interval_join_output(&batches, streamed_offset, buffered_offset); + let mut expected = expected_interval_join(join_type, streamed, buffered); + if join_type == Full { + actual.sort(); + expected.sort(); + } + assert_eq!(actual, expected, "{case}"); + Ok(()) +} + +/// Buffered key group of `rows` consecutive intervals `(10i, 10i + 10 * width]`. +fn interval_group( + key: i32, + first_bid: i32, + rows: i32, + width: i64, +) -> Vec { + (0..rows) + .map(|i| { + ( + key, + first_bid + i, + i as i64 * 10, + i as i64 * 10 + 10 * width, + ) + }) + .collect() +} + +/// One streamed row against a buffered key group much larger than the batch +/// size and spanning several buffered batches, passing the filter for none, +/// one or a few of its rows (or NULL when the streamed time is NULL). The +/// join filter is evaluated for every pair, but only the pairs that can be +/// output are materialized; the result must be the same either way. +#[tokio::test] +async fn join_filter_single_streamed_row_large_buffered_group() -> Result<()> { + let group = interval_group(1, 0, 1000, 1); + let overlapping = interval_group(1, 0, 1000, 3); + let cases: Vec<(&str, Option, &[IntervalBufferedRow])> = vec![ + ("none", Some(-5), &group), + ("null", None, &group), + ("first", Some(5), &group), + ("middle", Some(7775), &group), + ("last", Some(9999), &group), + ("several", Some(7775), &overlapping), + ]; + for (name, t, buffered) in cases { + let streamed = [(1, 0, t)]; + for join_type in [Inner, Left, Right, Full] { + for batch_size in [64, 333, 8192] { + for spill in [false, true] { + check_interval_join( + join_type, + &streamed, + buffered, + 1, + 128, + batch_size, + spill, + &format!( + "{name} {join_type} batch_size={batch_size} spill={spill}" + ), + ) + .await?; + } + } + } + } + Ok(()) +} + +/// Many streamed rows over key groups of various sizes, some larger than +/// the batch size: freezes cut through one streamed row's pairs and hold +/// the pairs of several rows, rows pass for none, one or several buffered +/// rows, and some keys exist on one side only. +#[tokio::test] +async fn join_filter_streamed_rows_across_freezes() -> Result<()> { + let mut buffered = vec![]; + for (key, rows, width) in [(0, 3, 1), (1, 150, 1), (2, 400, 2), (4, 1, 1), (5, 90, 4)] + { + let first_bid = buffered.len() as i32; + buffered.extend(interval_group(key, first_bid, rows, width)); + } + let mut state = 7u64; + let mut next = move || { + state = state + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + (state >> 33) as i64 + }; + let mut streamed = vec![]; + for (key, rows) in [(0, 2), (1, 7), (2, 5), (3, 2), (5, 9), (6, 1)] { + for _ in 0..rows { + let t = match next() % 6 { + 0 => None, + 1 => Some(-1 - next() % 10), + 2 => Some(5000 + next() % 10), + _ => Some(next() % 4100), + }; + streamed.push((key, streamed.len() as i32, t)); + } + } + for join_type in [Inner, Left, Right, Full] { + for (streamed_batch_rows, buffered_batch_rows) in [(4, 64), (1000, 37)] { + for batch_size in [50, 128, 1000] { + for spill in [false, true] { + check_interval_join( + join_type, + &streamed, + &buffered, + streamed_batch_rows, + buffered_batch_rows, + batch_size, + spill, + &format!( + "{join_type} streamed_batch_rows={streamed_batch_rows} \ + buffered_batch_rows={buffered_batch_rows} \ + batch_size={batch_size} spill={spill}" + ), + ) + .await?; + } + } + } + } + Ok(()) +} + /// Returns the column names on the schema fn columns(schema: &Schema) -> Vec { schema.fields().iter().map(|f| f.name().clone()).collect() diff --git a/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala index f18689c51d..35b170bf4e 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala @@ -1142,6 +1142,62 @@ class CometJoinSuite extends CometTestBase { } } + test("SortMergeJoin with a validity-interval join filter over wide rows and large key groups") { + withTempPath { dir => + withSQLConf( + CometConf.COMET_EXEC_SORT_MERGE_JOIN_ENABLED.key -> "true", + CometConf.COMET_EXEC_SORT_MERGE_JOIN_WITH_JOIN_FILTER_ENABLED.key -> "true", + CometConf.COMET_BATCH_SIZE.key -> "1000", + SQLConf.SHUFFLE_PARTITIONS.key -> "2", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { + // Each currency but the last has 3000 consecutive validity intervals, so an + // order meets a key group three times the batch size and passes for one + // interval (several for the overlapping currency 3), or for none when its + // time is out of range or NULL. Currency 4 has no rates, currency 5 no orders. + val payload = (0 until 30).map { c => + if (c % 2 == 0) s"id * $c AS p$c" else s"concat('payload_$c', '_', id) AS p$c" + } + val ordersPath = s"${dir.getCanonicalPath}/orders" + spark + .range(0, 600, 1, 1) + .selectExpr(Seq( + "id", + "CAST(id % 5 AS INT) AS cur", + "CASE WHEN id % 11 = 0 THEN NULL WHEN id % 13 = 0 THEN -5 " + + "WHEN id % 17 = 0 THEN 50000 ELSE (id * 7919) % 30000 + 1 END AS t") ++ + payload: _*) + .write + .parquet(ordersPath) + val ratesPath = s"${dir.getCanonicalPath}/rates" + spark + .range(0, 15000, 1, 1) + .selectExpr( + "CAST(CASE WHEN id < 12000 THEN id DIV 3000 ELSE 5 END AS INT) AS cur", + "(id % 3000) * 10 AS eff", + "(id % 3000) * 10 + CASE WHEN id DIV 3000 = 3 THEN 30 ELSE 10 END AS next_eff", + "concat('rate_', id) AS rate") + .write + .parquet(ratesPath) + + withParquetTable(ordersPath, "orders") { + withParquetTable(ratesPath, "rates") { + val condition = "o.cur = r.cur AND o.t > r.eff AND o.t <= r.next_eff" + for (joinType <- Seq("INNER", "LEFT", "FULL")) { + checkSparkAnswerAndOperator( + sql(s"SELECT o.*, r.* FROM orders o $joinType JOIN rates r ON $condition")) + } + checkSparkAnswerAndOperator( + sql(s"SELECT o.*, r.* FROM rates r RIGHT JOIN orders o ON $condition")) + checkSparkAnswerAndOperator( + sql(s"SELECT o.id, count(r.rate), sum(o.p28) " + + s"FROM orders o LEFT JOIN rates r ON $condition GROUP BY o.id")) + } + } + } + } + } + test("full outer join") { withTempView("`left`", "`right`", "allNulls") { allNulls.createOrReplaceTempView("allNulls") From 93dc75be803c4c713cf5e3dbffd8633f8570fdca Mon Sep 17 00:00:00 2001 From: msaf Date: Wed, 7 Oct 2026 22:45:41 +0100 Subject: [PATCH 02/13] test: differential fuzz suite for sort-merge joins with a join filter Compares Comet's sort-merge join with a join filter against Spark for inner, left, right, full, left semi and left anti joins over generated data with key groups of one row, about a batch and several batches on either side, NULL keys and NULL filter columns, and wide rows with strings, decimals and nested columns. Filters cover a validity interval bounded by one side, a band bounded by the other, one-sided conditions, always true and always false, and almost nothing or almost everything passing, at batch sizes 7, 64 and 1024 and under a memory pool small enough to make the join spill. Results are compared with Spark as multisets, and every case asserts that the join runs as CometSortMergeJoin. EXISTS and IN in a disjunction check the existence join that stays in Spark. The default part runs in CI; -Dcomet.test.smjFuzz.full=true runs the full matrix of shapes, side orders, two join keys, a second seed and groups of thousands of rows. Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + .../exec/CometSmjJoinFilterFuzzSuite.scala | 598 ++++++++++++++++++ 3 files changed, 600 insertions(+) create mode 100644 spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 830f888164..1c11bd8de6 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -540,6 +540,7 @@ jobs: org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometPartitionAggregateWindowSuite org.apache.comet.exec.CometJoinSuite + org.apache.comet.exec.CometSmjJoinFilterFuzzSuite org.apache.spark.sql.comet.CometMapInBatchSuite org.apache.spark.sql.execution.python.CometArrowPythonRunnerSuite org.apache.comet.CometNativeSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 20f815854e..1cf2f06d21 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -188,6 +188,7 @@ jobs: org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometPartitionAggregateWindowSuite org.apache.comet.exec.CometJoinSuite + org.apache.comet.exec.CometSmjJoinFilterFuzzSuite org.apache.spark.sql.comet.CometMapInBatchSuite org.apache.spark.sql.execution.python.CometArrowPythonRunnerSuite org.apache.comet.CometNativeSuite diff --git a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala new file mode 100644 index 0000000000..3dfec606c0 --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala @@ -0,0 +1,598 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.exec + +import java.io.File +import java.nio.file.Files +import java.sql.{Date, Timestamp} +import java.time.{Instant, LocalDate} + +import scala.collection.mutable +import scala.util.Random + +import org.apache.commons.io.FileUtils +import org.apache.spark.sql.{CometTestBase, Row} +import org.apache.spark.sql.comet.CometSortMergeJoinExec +import org.apache.spark.sql.execution.SparkPlan +import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.joins.SortMergeJoinExec +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types._ + +import org.apache.comet.CometConf + +/** + * Differential tests of Comet's sort-merge join with a join filter against Spark, over generated + * data with key groups of one row, about a batch and several batches on either side, NULL keys + * and NULL filter columns, and a wide streamed side with strings, decimals and nested columns. + * + * Every join type runs with filters of different shapes and selectivities (a validity interval + * bounded by one side, a band bounded by the other, one-sided conditions, always true and always + * false, almost nothing or almost everything passing) at several batch sizes and under a memory + * pool small enough to make the join spill. Results are compared with Spark as multisets. + * + * The default run covers each join type and filter at one batch size and a subset at the others. + * The full matrix runs only with `-Dcomet.test.smjFuzz.full=true`; + * `-Dcomet.test.smjFuzz.seed=` changes the data seed. + */ +class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHelper { + + private val seed: Long = + sys.props.get("comet.test.smjFuzz.seed").map(_.toLong).getOrElse(20261007L) + private val fullMatrix: Boolean = sys.props.get("comet.test.smjFuzz.full").contains("true") + + // ---------------------------------------------------------------- data generation + + private val leftSchema = StructType( + Seq( + StructField("id", LongType, nullable = false), + StructField("k", IntegerType), + StructField("kc", StringType), + StructField("t", LongType), + StructField("w", LongType), + StructField("lv", IntegerType), + StructField("s", StringType), + StructField("d38", DecimalType(38, 10)), + StructField("d10", DecimalType(10, 2)), + StructField("dbl", DoubleType), + StructField("ts", TimestampType), + StructField("dt", DateType), + StructField("b", BooleanType), + StructField( + "st", + StructType( + Seq( + StructField("a", IntegerType), + StructField("b", StringType), + StructField("c", ArrayType(IntegerType))))), + StructField("arr", ArrayType(StringType)), + StructField("m", MapType(StringType, IntegerType)))) + + private val rightSchema = StructType( + Seq( + StructField("id", LongType, nullable = false), + StructField("k", IntegerType), + StructField("kc", StringType), + StructField("eff", LongType), + StructField("next_eff", LongType), + StructField("rv", IntegerType), + StructField("rs", StringType), + StructField("rdec", DecimalType(18, 4)), + StructField("rarr", ArrayType(IntegerType)), + StructField( + "rst", + StructType(Seq(StructField("x", DoubleType), StructField("y", StringType)))))) + + private case class DataSet(name: String, left: String, right: String, stats: String) + + private var tempRoot: File = _ + private val dataSets = mutable.Map[String, DataSet]() + + override protected def afterAll(): Unit = { + try { + if (tempRoot != null) FileUtils.deleteQuietly(tempRoot) + } finally { + super.afterAll() + } + } + + private def withConfs[T](pairs: (String, String)*)(f: => T): T = { + var result: Option[T] = None + withSQLConf(pairs: _*) { result = Some(f) } + result.get + } + + private def pick[T](r: Random, xs: Seq[T]): T = xs(r.nextInt(xs.length)) + + private def orNull(r: Random, pct: Int)(v: => Any): Any = if (r.nextInt(100) < pct) null else v + + private def randomDecimal(r: Random, precision: Int, scale: Int): java.math.BigDecimal = { + val digits = (0 until 1 + r.nextInt(precision)).map(_ => ('0' + r.nextInt(10)).toChar) + val unscaled = new java.math.BigInteger(digits.mkString) + new java.math.BigDecimal(if (r.nextBoolean()) unscaled.negate() else unscaled, scale) + } + + private def randomString(r: Random): String = r.nextInt(30) match { + case 0 => "" + case 1 => "ΠΆβ‚¬πŸ˜€ ΓΌnΓ―" + case 2 => r.alphanumeric.take(1500 + r.nextInt(1000)).mkString + case _ => r.alphanumeric.take(1 + r.nextInt(12)).mkString + } + + private def randomDouble(r: Random): Double = + if (r.nextInt(4) == 0) { + pick(r, Seq(Double.NaN, -0.0d, 0.0d, Double.PositiveInfinity, Double.MinPositiveValue)) + } else r.nextDouble() * 2e6 - 1e6 + + private case class KeyGroup(k: Int, nl: Int, nr: Int, base: Long, mode: Int) + + private def leftRow(r: Random, g: KeyGroup): Seq[Any] = { + val k: Any = if (g.k < 0) null else g.k + val kc = + if (g.k < 0) "c0" else if (g.k % 23 == 0 && r.nextInt(3) == 0) null else s"c${g.k % 5}" + val span = math.max(g.nr, 1) * 10L + 40 + Seq( + k, + kc, + orNull(r, 8)(g.base - 20 + (r.nextLong() & Long.MaxValue) % span), + orNull(r, 5)(pick(r, Seq(0L, 5L, 15L, 30L))), + orNull(r, 10)(r.nextInt(100)), + orNull(r, 10)(randomString(r)), + orNull(r, 10)(randomDecimal(r, 38, 10)), + orNull(r, 10)(randomDecimal(r, 10, 2)), + orNull(r, 10)(randomDouble(r)), + orNull(r, 10)( + Timestamp.from(Instant.ofEpochSecond(r.nextInt(2000000000) * 2L - 1500000000L, 1000L))), + orNull(r, 10)(Date.valueOf(LocalDate.ofEpochDay(r.nextInt(60000) - 20000L))), + orNull(r, 10)(r.nextBoolean()), + orNull(r, 10)( + Row( + orNull(r, 20)(r.nextInt(1000)), + orNull(r, 20)(randomString(r)), + orNull(r, 20)(Seq.fill(r.nextInt(4))(orNull(r, 20)(r.nextInt(50)))))), + orNull(r, 10)( + Seq.fill(r.nextInt(5))(orNull(r, 20)(r.alphanumeric.take(r.nextInt(8)).mkString))), + orNull(r, 10)((0 until r.nextInt(3)).map(i => s"k$i" -> orNull(r, 20)(r.nextInt(9))).toMap)) + } + + private def rightRow(r: Random, g: KeyGroup, j: Int): Seq[Any] = { + val k: Any = if (g.k < 0) null else g.k + val kc = if (g.k < 0) "c0" else s"c${g.k % 5}" + val eff = g.base + j * 10L + val width = g.mode match { + case 0 => 10L + case 1 => 25L + case _ => 5L + } + Seq( + k, + kc, + orNull(r, 4)(eff), + orNull(r, 4)(eff + width), + orNull(r, 10)(r.nextInt(100)), + orNull(r, 10)(randomString(r)), + orNull(r, 10)(randomDecimal(r, 18, 4)), + orNull(r, 10)(Seq.fill(r.nextInt(4))(orNull(r, 20)(r.nextInt(50)))), + orNull(r, 10)(Row(orNull(r, 20)(randomDouble(r)), orNull(r, 20)(randomString(r))))) + } + + private val mainCombos: Seq[(Int, Int)] = Seq( + 1 -> 1, + 1 -> 7, + 7 -> 1, + 7 -> 7, + 8 -> 6, + 6 -> 8, + 63 -> 64, + 64 -> 64, + 65 -> 1, + 1 -> 65, + 130 -> 2, + 2 -> 130, + 140 -> 140, + 200 -> 70, + 1023 -> 1, + 1 -> 1024, + 1025 -> 3, + 3 -> 1025, + 2100 -> 1, + 1 -> 2100, + 0 -> 10, + 10 -> 0, + 0 -> 1100, + 1100 -> 0) + + private val bigCombos: Seq[(Int, Int)] = + Seq(5000 -> 2, 2 -> 5000, 3000 -> 40, 40 -> 3000, 600 -> 600, 9000 -> 1, 1 -> 9000) + + private def dataSet(name: String): DataSet = synchronized { + dataSets.getOrElseUpdate( + name, + name match { + case "main" => createDataSet(name, seed, mainCombos, 300) + case "alt" => createDataSet(name, seed + 1, mainCombos, 300) + case "big" => createDataSet(name, seed, bigCombos, 200) + }) + } + + private def createDataSet( + name: String, + dataSeed: Long, + combos: Seq[(Int, Int)], + numRandom: Int): DataSet = { + val r = new Random(dataSeed ^ name.hashCode) + def small(): Int = r.nextInt(10) match { + case 0 => 0 + case 1 | 2 | 3 => 1 + case 9 => 6 + r.nextInt(15) + case _ => 2 + r.nextInt(4) + } + val sized = r.shuffle(combos ++ Seq.fill(numRandom)((small(), small()))) + val groups = sized.zipWithIndex.map { case ((nl, nr), i) => + KeyGroup(i * 3 + r.nextInt(3), nl, nr, r.nextInt(100000) * 10L, r.nextInt(3)) + } ++ Seq(KeyGroup(-1, 40 + r.nextInt(40), 40 + r.nextInt(40), 5000L, 0)) + val leftRows = r.shuffle(groups.flatMap(g => Seq.fill(g.nl)(leftRow(r, g)))) + val rightRows = + r.shuffle(groups.flatMap(g => (0 until g.nr).map(j => rightRow(r, g, j)))) + + if (tempRoot == null) tempRoot = Files.createTempDirectory("comet-smj-fuzz").toFile + val leftPath = new File(tempRoot, s"$name-left").getCanonicalPath + val rightPath = new File(tempRoot, s"$name-right").getCanonicalPath + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + val lr = leftRows.zipWithIndex.map { case (v, i) => Row.fromSeq(i.toLong +: v) } + val rr = rightRows.zipWithIndex.map { case (v, i) => Row.fromSeq(i.toLong +: v) } + spark + .createDataFrame(spark.sparkContext.parallelize(lr, 4), leftSchema) + .write + .option("parquet.block.size", 65536) + .parquet(leftPath) + spark + .createDataFrame(spark.sparkContext.parallelize(rr, 4), rightSchema) + .write + .option("parquet.block.size", 65536) + .parquet(rightPath) + } + val ds = DataSet( + name, + s"smj_fuzz_${name}_l", + s"smj_fuzz_${name}_r", + s"left rows ${leftRows.size}, right rows ${rightRows.size}, keys ${groups.size}, " + + s"max group ${groups.map(_.nl).max}/${groups.map(_.nr).max}, " + + s"key pairs ${groups.filter(_.k >= 0).map(g => g.nl.toLong * g.nr).sum}") + spark.read.parquet(leftPath).createOrReplaceTempView(ds.left) + spark.read.parquet(rightPath).createOrReplaceTempView(ds.right) + ds + } + + // ---------------------------------------------------------------- query shapes + + private case class JoinKind(name: String, sql: String, existence: Boolean = false) + + private val inner = JoinKind("inner", "INNER JOIN") + private val leftOuter = JoinKind("left", "LEFT JOIN") + private val rightOuter = JoinKind("right", "RIGHT JOIN") + private val fullOuter = JoinKind("full", "FULL JOIN") + private val leftSemi = JoinKind("semi", "LEFT SEMI JOIN", existence = true) + private val leftAnti = JoinKind("anti", "LEFT ANTI JOIN", existence = true) + private val joinKinds = Seq(inner, leftOuter, rightOuter, fullOuter, leftSemi, leftAnti) + + private case class Filter(name: String, sql: String) + + private val interval = Filter("interval", "l.t > r.eff AND l.t <= r.next_eff") + private val band = Filter("band", "r.eff >= l.t - l.w AND r.eff < l.t + l.w") + private val leftOnly = Filter("left_only", "l.lv % 3 = 0") + private val rightOnly = Filter("right_only", "r.rv % 3 = 0") + private val alwaysFalse = Filter("false", "l.id + r.id < 0") + private val alwaysTrue = Filter("true", "l.id + r.id >= 0") + private val rare = Filter("rare", "(l.id * 31 + r.id * 17) % 997 = 0") + private val most = Filter("most", "(l.id + r.id) % 10 <> 0") + private val nullable = Filter("nullable_cmp", "l.lv < r.rv") + private val typed = Filter("typed", "l.d10 * 2 >= r.rdec OR l.s < r.rs") + private val filters = + Seq(interval, band, leftOnly, rightOnly, alwaysFalse, alwaysTrue, rare, most, nullable, typed) + + private case class Shape(name: String, confs: Seq[(String, String)], spill: Boolean = false) + + private def batch(n: Int) = CometConf.COMET_BATCH_SIZE.key -> n.toString + private def aqe(on: Boolean) = SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> on.toString + private def partitions(n: Int) = SQLConf.SHUFFLE_PARTITIONS.key -> n.toString + private val tinyPool = CometConf.COMET_OFFHEAP_MEMORY_POOL_FRACTION.key -> "0.0005" + private val manySplits = Seq( + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "65536", + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> "1") + + private val b7 = Shape("batch7", Seq(batch(7), aqe(false), partitions(1))) + private val b64 = Shape("batch64", Seq(batch(64), aqe(true), partitions(2))) + private val b1024 = Shape("batch1024", Seq(batch(1024), aqe(true), partitions(2))) + private val b8192 = Shape("batch8192", Seq(batch(8192), aqe(false), partitions(3))) + private val spill64 = + Shape( + "spill64", + Seq(batch(64), aqe(true), partitions(1), tinyPool) ++ manySplits, + spill = true) + private val spill7 = + Shape( + "spill7", + Seq(batch(7), aqe(false), partitions(2), tinyPool) ++ manySplits, + spill = true) + private val allShapes = Seq(b7, b64, b1024, b8192, spill7, spill64) + + private val baseConfs: Seq[(String, String)] = Seq( + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.PREFER_SORTMERGEJOIN.key -> "true", + CometConf.COMET_EXEC_SORT_MERGE_JOIN_ENABLED.key -> "true", + CometConf.COMET_EXEC_SORT_MERGE_JOIN_WITH_JOIN_FILTER_ENABLED.key -> "true", + CometConf.COMET_FORCE_SHJ.key -> "false") + + private case class Q( + label: String, + sql: String, + expectNative: Boolean = true, + expectCondition: Boolean = true) + + private def joinQuery( + ds: DataSet, + kind: JoinKind, + filter: Filter, + flipped: Boolean, + twoKeys: Boolean): Q = { + val on = (if (twoKeys) "l.k = r.k AND l.kc = r.kc" else "l.k = r.k") + s" AND (${filter.sql})" + val (from, select) = + if (flipped) (s"${ds.right} r ${kind.sql} ${ds.left} l", if (kind.existence) "r.*" else "*") + else (s"${ds.left} l ${kind.sql} ${ds.right} r", if (kind.existence) "l.*" else "*") + val pushable = (filter, kind.name, flipped) match { + case (`leftOnly`, "inner" | "right", false) => true + case (`leftOnly`, "inner" | "left", true) => true + case (`leftOnly`, "semi" | "anti", true) => true + case (`leftOnly`, "semi", false) => true + case (`rightOnly`, "inner" | "left", false) => true + case (`rightOnly`, "inner" | "right", true) => true + case (`rightOnly`, "semi" | "anti", false) => true + case (`rightOnly`, "semi", true) => true + case _ => false + } + val label = s"${kind.name}${if (flipped) "_rl" else ""}${if (twoKeys) "_2keys" else ""}" + + s"/${filter.name}" + Q(label, s"SELECT $select FROM $from ON $on", expectCondition = !pushable) + } + + private def existenceQueries(ds: DataSet): Seq[Q] = Seq( + Q( + "exists_or", + s"SELECT l.* FROM ${ds.left} l WHERE l.lv > 90 OR EXISTS (SELECT 1 FROM ${ds.right} r " + + "WHERE r.k = l.k AND l.t > r.eff AND l.t <= r.next_eff)", + expectNative = false), + Q( + "in_or", + s"SELECT l.* FROM ${ds.left} l WHERE l.lv > 90 OR l.k IN (SELECT r.k FROM ${ds.right} r " + + "WHERE l.t > r.eff AND l.t <= r.next_eff)", + expectNative = false), + Q( + "not_exists", + s"SELECT l.* FROM ${ds.left} l WHERE NOT EXISTS (SELECT 1 FROM ${ds.right} r " + + "WHERE r.k = l.k AND r.eff >= l.t - l.w AND r.eff < l.t + l.w)")) + + // ---------------------------------------------------------------- comparison + + private def norm(v: Any): Any = v match { + case null => null + case b: Array[Byte] => ("bin", b.toList) + case d: Double => ("d", java.lang.Double.doubleToLongBits(d)) + case f: Float => ("f", java.lang.Float.floatToIntBits(f)) + case r: Row => r.toSeq.map(norm).toList + case m: scala.collection.Map[_, _] => m.map { case (k, x) => (norm(k), norm(x)) }.toMap + case s: scala.collection.Seq[_] => s.map(norm).toList + case other => other + } + + private def counts(rows: Seq[Row]): Map[Any, Int] = + rows.groupBy(r => norm(r)).map { case (k, v) => k -> v.size } + + private def showRow(v: Any): String = { + val s = String.valueOf(v) + if (s.length > 300) s.take(300) + "..." else s + } + + private case class Outcome( + label: String, + shape: String, + data: String, + sparkRows: Int, + cometRows: Int, + native: Boolean, + condition: Boolean, + spills: Long, + error: Option[String]) { + def line: String = + f"SMJFUZZ ${if (error.isEmpty) "PASS" else "FAIL"} seed=$seed data=$data%-4s " + + f"shape=$shape%-9s case=$label%-28s spark=$sparkRows%7d comet=$cometRows%7d " + + s"native=$native cond=$condition spills=$spills" + } + + private def cometJoins(plan: SparkPlan): Seq[CometSortMergeJoinExec] = + collect(plan) { case j: CometSortMergeJoinExec => j } + + private def runCase(ds: DataSet, shape: Shape, q: Q): Outcome = { + val ctx = s"[seed=$seed data=${ds.name} shape=${shape.name} case=${q.label}]" + withConfs(baseConfs ++ shape.confs: _*) { + val expected = withConfs(CometConf.COMET_ENABLED.key -> "false") { + sql(q.sql).collect().toSeq + } + val df = sql(q.sql) + val actual = + try Right(df.collect().toSeq) + catch { case e: Throwable => Left(e) } + val plan = df.queryExecution.executedPlan + val joins = cometJoins(plan) + val sparkJoins = collect(plan) { case j: SortMergeJoinExec => j } + val native = joins.nonEmpty && sparkJoins.isEmpty + val condition = joins.exists(_.condition.isDefined) + val spills = joins.flatMap(_.metrics.get("spill_count")).map(_.value).sum + val errors = mutable.ArrayBuffer[String]() + actual match { + case Left(e) => + errors += s"Comet failed: $e\n${e.getStackTrace.take(15).mkString("\n")}" + case Right(rows) => + val exp = counts(expected) + val act = counts(rows) + if (exp != act) { + val missing = exp.toSeq.flatMap { case (k, n) => + val d = n - act.getOrElse(k, 0) + if (d > 0) Some(k -> d) else None + } + val extra = act.toSeq.flatMap { case (k, n) => + val d = n - exp.getOrElse(k, 0) + if (d > 0) Some(k -> d) else None + } + errors += s"results differ: Spark ${expected.size} rows, Comet ${rows.size} rows, " + + s"${missing.map(_._2).sum} missing in Comet, ${extra.map(_._2).sum} extra in " + + "Comet\nmissing (row x count):\n" + + missing.take(5).map { case (k, d) => s" ${showRow(k)} x$d" }.mkString("\n") + + "\nextra (row x count):\n" + + extra.take(5).map { case (k, d) => s" ${showRow(k)} x$d" }.mkString("\n") + } + } + if (native != q.expectNative) { + errors += s"expected native=${q.expectNative}, got $native" + } + if (q.expectNative && condition != q.expectCondition) { + errors += s"expected a join filter in CometSortMergeJoin=${q.expectCondition}" + } + val error = + if (errors.isEmpty) None + else Some(s"$ctx ${errors.mkString("\n")}\nquery: ${q.sql}\nplan:\n$plan") + val o = Outcome( + q.label, + shape.name, + ds.name, + expected.size, + actual.fold(_ => -1, _.size), + native, + condition, + spills, + error) + // scalastyle:off println + println(o.line) + // scalastyle:on println + o + } + } + + private def runAll(ds: DataSet, shape: Shape, qs: Seq[Q]): Seq[Outcome] = { + val outcomes = qs.map(q => runCase(ds, shape, q)) + val failed = outcomes.filter(_.error.nonEmpty) + if (failed.nonEmpty) { + fail( + s"[seed=$seed data=${ds.name} (${ds.stats}) shape=${shape.name}] " + + s"${failed.size} of ${outcomes.size} cases failed: " + + failed.map(_.label).mkString(", ") + "\n\n" + + failed.take(3).flatMap(_.error).mkString("\n\n")) + } + outcomes + } + + private def matrix( + ds: DataSet, + kinds: Seq[JoinKind], + fs: Seq[Filter], + flipped: Boolean = false, + twoKeys: Boolean = false): Seq[Q] = + for (k <- kinds; f <- fs) yield joinQuery(ds, k, f, flipped, twoKeys) + + private def assumeFull(): Unit = + assume(fullMatrix, "the full matrix runs only with -Dcomet.test.smjFuzz.full=true") + + // ---------------------------------------------------------------- default tests + + test("every join type and filter shape, batch 64") { + val ds = dataSet("main") + runAll(ds, b64, matrix(ds, joinKinds, filters)) + } + + test("flipped sides and two join keys, batch 64") { + val ds = dataSet("main") + runAll( + ds, + b64, + matrix(ds, Seq(rightOuter, leftSemi, leftAnti), Seq(interval, band), flipped = true) ++ + matrix(ds, joinKinds, Seq(interval), twoKeys = true)) + } + + test("batch 7: key groups spanning many batches") { + val ds = dataSet("main") + runAll(ds, b7, matrix(ds, joinKinds, Seq(interval, band, most))) + } + + test("batch 1024: key groups around and over one batch") { + val ds = dataSet("main") + runAll(ds, b1024, matrix(ds, joinKinds, Seq(interval, rare, alwaysTrue))) + } + + test("spilling join under a tiny memory pool") { + val ds = dataSet("main") + val outcomes = runAll(ds, spill64, matrix(ds, joinKinds, Seq(interval, most))) + assert(outcomes.exists(_.spills > 0), s"[seed=$seed] the join did not spill") + } + + test("existence joins from EXISTS and IN in a disjunction") { + val ds = dataSet("main") + runAll(ds, b64, existenceQueries(ds)) + } + + // ---------------------------------------------------------------- full matrix + + for (shape <- allShapes) { + test(s"full: ${shape.name}, every join type, filter, side order and key count") { + assumeFull() + val ds = dataSet("main") + val outcomes = runAll( + ds, + shape, + matrix(ds, joinKinds, filters) ++ + matrix(ds, joinKinds, filters, flipped = true) ++ + matrix(ds, joinKinds, filters, twoKeys = true) ++ + existenceQueries(ds)) + if (shape.spill) { + assert(outcomes.exists(_.spills > 0), s"[seed=$seed] the join did not spill") + } + } + } + + for (shape <- Seq(b7, b64, spill64)) { + test(s"full: second seed, ${shape.name}") { + assumeFull() + val ds = dataSet("alt") + runAll( + ds, + shape, + matrix(ds, joinKinds, filters) ++ matrix(ds, joinKinds, filters, flipped = true)) + } + } + + for (shape <- Seq(b64, b1024, b8192, spill64)) { + test(s"full: groups of thousands of rows, ${shape.name}") { + assumeFull() + val ds = dataSet("big") + val fs = Seq(interval, band, rare, alwaysFalse, nullable) + runAll(ds, shape, matrix(ds, joinKinds, fs) ++ matrix(ds, joinKinds, fs, flipped = true)) + } + } +} From 600f88ee676dd96466fa0129c8dad713939714ee Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 00:11:06 +0100 Subject: [PATCH 03/13] fix: keep a FULL join from null-joining a streamed row twice get_corrected_filter_mask groups the deferred-filter entries of one streamed row by (batch id, row index) and treated an entry without a row index (a null-joined row) as the end of the row before it. A FULL join stages the null-joined rows of an unmatched buffered key group in freeze_buffered, ahead of the pairs freeze_streamed materializes in the same freeze. When a streamed row's pairs span two freezes and such a group was passed in between, its null-joined rows land between the row's two runs: the first run ended the row, so a row whose first run had no passing pair was emitted null-joined there and again (or with its passing pairs) after the second run. Before eedede1a6 every freeze pushed batch_size entries, so the deferred output was flushed at the next loop iteration, before the streamed row's remaining pairs could be separated. Keeping only the candidate pairs pushes far fewer entries, the flush comes later and the latent split shows: extra null-joined rows in FULL joins over key groups larger than the batch on both sides. Entries without a row index now neither end a row's group nor reset it, which is all they can be: buffered rows null-joined on the streamed side, or streamed rows that matched no key and form no group. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../src/joins/sort_merge_join/filter.rs | 63 ++++++---------- .../src/joins/sort_merge_join/tests.rs | 72 +++++++++++++++++++ 2 files changed, 93 insertions(+), 42 deletions(-) diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs index de3832fb71..b3cafc1df3 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs @@ -145,51 +145,31 @@ pub fn needs_deferred_filtering( && matches!(join_type, JoinType::Left | JoinType::Right | JoinType::Full) } -/// Determines if current index is the last occurrence of a row +/// Marks the entries that are the last of their input row /// /// Used during filter mask correction to detect row boundaries when grouping -/// output rows by input row. -fn last_index_for_row( - row_index: usize, - indices: &UInt64Array, - batch_ids: &[usize], - indices_len: usize, -) -> bool { - debug_assert_eq!( - indices.len(), - indices_len, - "indices.len() should match indices_len parameter" - ); +/// output rows by input row. Entries without a row index (null-joined rows +/// that belong to no input row group) are not marked and do not end a +/// group: a FULL join stages the null-joined rows of unmatched buffered rows +/// ahead of the pairs a freeze materializes, so they can fall between two +/// runs of one streamed row's pairs. +fn last_entries_of_rows(indices: &UInt64Array, batch_ids: &[usize]) -> Vec { debug_assert_eq!( batch_ids.len(), - indices_len, - "batch_ids.len() should match indices_len" - ); - debug_assert!( - row_index < indices_len, - "row_index {row_index} should be < indices_len {indices_len}", + indices.len(), + "batch_ids.len() should match indices.len()" ); - - // If this is the last index overall, it's definitely the last for this row - if row_index == indices_len - 1 { - return true; - } - - // Check if next row has different (batch_id, index) pair - let current_batch_id = batch_ids[row_index]; - let next_batch_id = batch_ids[row_index + 1]; - - if current_batch_id != next_batch_id { - return true; - } - - // Same batch_id, check if row index is different - // Both current and next should be non-null (already joined rows) - if indices.is_null(row_index) || indices.is_null(row_index + 1) { - return true; + let mut last = vec![false; indices.len()]; + let mut next_row: Option<(usize, u64)> = None; + for i in (0..indices.len()).rev() { + if indices.is_null(i) { + continue; + } + let row = (batch_ids[i], indices.value(i)); + last[i] = next_row != Some(row); + next_row = Some(row); } - - indices.value(row_index) != indices.value(row_index + 1) + last } /// Corrects the filter mask for joins with deferred filtering @@ -229,9 +209,8 @@ pub fn get_corrected_filter_mask( // discard (null) remaining matches, null-join if none passed. // Null metadata entries are already-null-joined rows that // flow through unchanged to preserve output ordering. - for i in 0..row_indices_length { - let last_index = - last_index_for_row(i, row_indices, batch_ids, row_indices_length); + let last_entries = last_entries_of_rows(row_indices, batch_ids); + for (i, &last_index) in last_entries.iter().enumerate() { if filter_mask.is_null(i) { corrected_mask.append_value(true); } else if filter_mask.value(i) { diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs index e1705ff693..b2d9d91452 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs @@ -4757,6 +4757,78 @@ async fn join_filter_streamed_rows_across_freezes() -> Result<()> { Ok(()) } +/// Key groups larger than the batch size on both sides, with streamed rows +/// passing the filter for none, one or several buffered rows, and buffered +/// key groups no streamed row matches between them. A FULL join stages the +/// null-joined rows of such a group at the next freeze, which can fall +/// between two runs of one streamed row's pairs; the row must still be +/// emitted null-joined at most once. +#[tokio::test] +async fn join_filter_large_groups_on_both_sides() -> Result<()> { + let mut state = 11u64; + let mut next = move || { + state = state + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + (state >> 33) as i64 + }; + let mut buffered = vec![]; + let mut streamed = vec![]; + for (key, streamed_rows, buffered_rows, width) in [ + (0, 140, 140, 1), + (1, 0, 50, 1), + (2, 200, 70, 2), + (3, 0, 3, 1), + (4, 3, 140, 1), + (5, 140, 2, 1), + (6, 5, 0, 1), + (7, 70, 70, 1), + ] { + let first_bid = buffered.len() as i32; + buffered.extend(interval_group(key, first_bid, buffered_rows, width)); + for _ in 0..streamed_rows { + let t = match next() % 4 { + 0 => None, + 1 => Some(-1 - next() % 10), + _ => Some(next() % (buffered_rows as i64 * 12 + 1)), + }; + streamed.push((key, streamed.len() as i32, t)); + } + } + let never: Vec = streamed + .iter() + .map(|&(key, sid, _)| (key, sid, Some(-1))) + .collect(); + for (name, streamed) in [("mixed", &streamed), ("never", &never)] { + for join_type in [Inner, Left, Right, Full] { + for (streamed_batch_rows, buffered_batch_rows) in + [(64, 64), (1000, 37), (7, 1000)] + { + for batch_size in [64, 100, 8192] { + for spill in [false, true] { + check_interval_join( + join_type, + streamed, + &buffered, + streamed_batch_rows, + buffered_batch_rows, + batch_size, + spill, + &format!( + "{name} {join_type} streamed_batch_rows={streamed_batch_rows} \ + buffered_batch_rows={buffered_batch_rows} \ + batch_size={batch_size} spill={spill}" + ), + ) + .await?; + } + } + } + } + } + Ok(()) +} + /// Returns the column names on the schema fn columns(schema: &Schema) -> Vec { schema.fields().iter().map(|f| f.name().clone()).collect() From 2692c3edce5385a234d9829821301ecf5cdc4233 Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 01:29:31 +0100 Subject: [PATCH 04/13] perf: evaluate buffered-only join filter subexpressions once per buffered row A sort-merge join filter like Spark's `o.t > cast(r.eff as timestamp) AND o.t <= cast(r.next_eff as timestamp)` re-evaluated the casts for every candidate pair, and Comet's date to timestamp cast resolves the session time zone per value: over a large key group almost all of the join's time went into casting the same buffered rows once per streamed row. The largest non-volatile subexpressions of the filter that read buffered columns alone are now lifted out (HoistedJoinFilter). A freeze whose pairs all fall in the current key group of their buffered batches evaluates them once over that group's rows of each batch, caches the results on the batch (accounted in its memory reservation) and gathers them per pair; other freezes, e.g. ones holding pairs of an earlier key group or null-joined pairs, evaluate the original filter as before. Also, for large key groups: - a streamed row is paired with a run of buffered rows at once instead of pair by pair; - buffered columns of pairs that form runs of consecutive rows across several batches are copied run by run instead of interleaved row by row; - the FULL join filter state is updated without a per-pair branch. New bench spark-expr/benches/smj_interval_filter.rs: string key, 4 keys x 100 streamed rows (98 payload columns) x 11000 buffered rows with date bounds and the Spark cast, one passing pair per streamed row, 4.4M pairs, batch_size 8192. ns per pair, eedede1a6 -> this commit: input batches 8192: LEFT 59.6 -> 5.2, INNER 59.6 -> 4.2, FULL 64.5 -> 8.0 input batches 1024: LEFT 61.2 -> 5.8, INNER 60.2 -> 4.6, FULL 64.2 -> 9.6 Unique keys change by +1..4%, the no-filter 1:1 shapes of sort_merge_join.rs by about +3%, inner_1to10 -20%. Co-Authored-By: Claude Opus 5.5 (1M context) --- native/spark-expr/Cargo.toml | 4 + .../spark-expr/benches/smj_interval_filter.rs | 262 ++++++++++ .../src/joins/sort_merge_join/filter.rs | 144 +++++- .../sort_merge_join/materializing_stream.rs | 469 ++++++++++++++++-- .../src/joins/sort_merge_join/tests.rs | 120 ++++- 5 files changed, 937 insertions(+), 62 deletions(-) create mode 100644 native/spark-expr/benches/smj_interval_filter.rs diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index 3e07bdda05..c86ac91696 100644 --- a/native/spark-expr/Cargo.toml +++ b/native/spark-expr/Cargo.toml @@ -442,3 +442,7 @@ harness = false [[bench]] name = "nested_comparison" harness = false + +[[bench]] +name = "smj_interval_filter" +harness = false diff --git a/native/spark-expr/benches/smj_interval_filter.rs b/native/spark-expr/benches/smj_interval_filter.rs new file mode 100644 index 0000000000..85062a45f0 --- /dev/null +++ b/native/spark-expr/benches/smj_interval_filter.rs @@ -0,0 +1,262 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Sort-merge join with a validity-interval join filter, shaped like a Spark +//! `o.t > cast(r.eff as timestamp) AND o.t <= cast(r.next_eff as timestamp)` +//! lookup: a string key, large buffered key groups spread over several input +//! batches, a wide streamed side, and Spark's date-to-timestamp casts. + +use arrow::array::{ + ArrayRef, Date32Array, Decimal128Array, Int64Array, RecordBatch, StringArray, + TimestampMicrosecondArray, +}; +use arrow::compute::SortOptions; +use arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit}; +use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; +use datafusion::common::{JoinSide, JoinType, NullEquality}; +use datafusion::datasource::memory::MemorySourceConfig; +use datafusion::execution::TaskContext; +use datafusion::logical_expr::Operator; +use datafusion::physical_expr::expressions::{BinaryExpr, Column}; +use datafusion::physical_expr::PhysicalExpr; +use datafusion::physical_plan::joins::utils::{ColumnIndex, JoinFilter}; +use datafusion::physical_plan::joins::SortMergeJoinExec; +use datafusion::physical_plan::{collect, ExecutionPlan}; +use datafusion::prelude::SessionConfig; +use datafusion_comet_spark_expr::{Cast, EvalMode, SparkCastOptions}; +use std::sync::Arc; +use tokio::runtime::Runtime; + +const MICROS_PER_DAY: i64 = 86_400_000_000; + +fn timestamp_type() -> DataType { + DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())) +} + +fn key_of(k: usize) -> String { + format!("C{k:06}") +} + +/// Streamed side: `key`, the lookup time `t` and `payload_cols` payload +/// columns alternating Int64 and Utf8. Every row falls into one interval of +/// its key group. +fn build_streamed( + keys: usize, + rows_per_key: usize, + group_rows: usize, + payload_cols: usize, + batch_size: usize, +) -> (SchemaRef, Vec) { + let mut fields = vec![ + Field::new("key", DataType::Utf8, false), + Field::new("t", timestamp_type(), true), + ]; + for c in 0..payload_cols { + let data_type = if c % 2 == 0 { + DataType::Int64 + } else { + DataType::Utf8 + }; + fields.push(Field::new(format!("p{c}"), data_type, true)); + } + let schema = Arc::new(Schema::new(fields)); + + let num_rows = keys * rows_per_key; + let mut rows: Vec<(usize, i64)> = (0..num_rows) + .map(|i| { + let day = (i * 7919) % group_rows; + let offset = ((i * 104_729) as i64 % MICROS_PER_DAY) + 1; + (i / rows_per_key, day as i64 * MICROS_PER_DAY + offset) + }) + .collect(); + rows.sort(); + let mut columns: Vec = vec![ + Arc::new(StringArray::from_iter_values( + rows.iter().map(|(k, _)| key_of(*k)), + )), + Arc::new( + TimestampMicrosecondArray::from_iter_values(rows.iter().map(|(_, t)| *t)) + .with_timezone("UTC"), + ), + ]; + for c in 0..payload_cols { + if c % 2 == 0 { + columns.push(Arc::new(Int64Array::from_iter_values( + (0..num_rows).map(|i| (i * c) as i64), + ))); + } else { + columns.push(Arc::new(StringArray::from_iter_values( + (0..num_rows).map(|i| format!("payload_{c}_{i:08}")), + ))); + } + } + let batch = RecordBatch::try_new(Arc::clone(&schema), columns).unwrap(); + (schema, split_batch(&batch, batch_size)) +} + +/// Buffered side: `group_rows` consecutive one-day intervals +/// `(eff, next_eff]` per key, as dates. +fn build_buffered( + keys: usize, + group_rows: usize, + batch_size: usize, +) -> (SchemaRef, Vec) { + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Utf8, false), + Field::new("eff", DataType::Date32, true), + Field::new("next_eff", DataType::Date32, true), + Field::new("rate", DataType::Decimal128(18, 6), true), + ])); + let num_rows = keys * group_rows; + let key = StringArray::from_iter_values((0..num_rows).map(|i| key_of(i / group_rows))); + let eff = Date32Array::from_iter_values((0..num_rows).map(|i| (i % group_rows) as i32)); + let next_eff = + Date32Array::from_iter_values((0..num_rows).map(|i| (i % group_rows) as i32 + 1)); + let rate = Decimal128Array::from_iter_values((0..num_rows).map(|i| i as i128)) + .with_precision_and_scale(18, 6) + .unwrap(); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(key), + Arc::new(eff), + Arc::new(next_eff), + Arc::new(rate), + ], + ) + .unwrap(); + (schema, split_batch(&batch, batch_size)) +} + +fn split_batch(batch: &RecordBatch, batch_size: usize) -> Vec { + (0..batch.num_rows()) + .step_by(batch_size) + .map(|offset| batch.slice(offset, (batch.num_rows() - offset).min(batch_size))) + .collect() +} + +fn to_timestamp(column: &str, index: usize) -> Arc { + Arc::new(Cast::new( + Arc::new(Column::new(column, index)), + timestamp_type(), + SparkCastOptions::new(EvalMode::Legacy, "UTC", false), + None, + None, + )) +} + +/// `t > cast(eff as timestamp) AND t <= cast(next_eff as timestamp)` +fn interval_filter(streamed: &Schema, buffered: &Schema) -> JoinFilter { + let t: Arc = Arc::new(Column::new("t", 0)); + let expression = Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new( + Arc::clone(&t), + Operator::Gt, + to_timestamp("eff", 1), + )), + Operator::And, + Arc::new(BinaryExpr::new( + t, + Operator::LtEq, + to_timestamp("next_eff", 2), + )), + )); + JoinFilter::new( + expression, + vec![ + ColumnIndex { + index: 1, + side: JoinSide::Left, + }, + ColumnIndex { + index: 1, + side: JoinSide::Right, + }, + ColumnIndex { + index: 2, + side: JoinSide::Right, + }, + ], + Arc::new(Schema::new(vec![ + streamed.field(1).clone(), + buffered.field(1).clone(), + buffered.field(2).clone(), + ])), + ) +} + +fn make_exec(batches: &[RecordBatch], schema: &SchemaRef) -> Arc { + MemorySourceConfig::try_new_exec(&[batches.to_vec()], Arc::clone(schema), None).unwrap() +} + +fn bench_smj_interval_filter(c: &mut Criterion) { + let rt = Runtime::new().unwrap(); + let mut group = c.benchmark_group("smj_interval_filter"); + group.sample_size(10); + for (shape, keys, rows_per_key, group_rows, input_batch_size) in [ + ("group11k_b8192", 4, 100, 11_000, 8192), + ("group11k_b1024", 4, 100, 11_000, 1024), + ("unique", 100_000, 1, 1, 8192), + ] { + let (streamed_schema, streamed_batches) = + build_streamed(keys, rows_per_key, group_rows, 98, input_batch_size); + let (buffered_schema, buffered_batches) = + build_buffered(keys, group_rows, input_batch_size); + let pairs = keys * rows_per_key * group_rows; + for join_type in [JoinType::Left, JoinType::Full, JoinType::Inner] { + group.bench_function( + BenchmarkId::new(format!("{shape}_{join_type:?}"), pairs), + |b| { + b.iter(|| { + let left = make_exec(&streamed_batches, &streamed_schema); + let right = make_exec(&buffered_batches, &buffered_schema); + let on = vec![( + Arc::new(Column::new("key", 0)) as Arc, + Arc::new(Column::new("key", 0)) as Arc, + )]; + let join = SortMergeJoinExec::try_new( + left, + right, + on, + Some(interval_filter(&streamed_schema, &buffered_schema)), + join_type, + vec![SortOptions::default()], + NullEquality::NullEqualsNothing, + ) + .unwrap(); + let task_ctx = Arc::new( + TaskContext::default() + .with_session_config(SessionConfig::new().with_batch_size(8192)), + ); + let rows: usize = rt.block_on(async { + let batches = collect(Arc::new(join), task_ctx).await.unwrap(); + batches.iter().map(|b| b.num_rows()).sum() + }); + if join_type != JoinType::Full { + assert_eq!(rows, keys * rows_per_key); + } + rows + }) + }, + ); + } + } + group.finish(); +} + +criterion_group!(benches, bench_smj_interval_filter); +criterion_main!(benches); diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs index b3cafc1df3..4d6398a6a6 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs @@ -31,10 +31,14 @@ use arrow::array::{ }; use arrow::compute::kernels::zip::zip; use arrow::compute::{self, filter_record_batch}; -use arrow::datatypes::SchemaRef; -use datafusion_common::{JoinType, Result}; +use arrow::datatypes::{Field, Schema, SchemaRef}; +use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRecursion}; +use datafusion_common::{JoinSide, JoinType, Result}; +use datafusion_physical_expr::expressions::Column; +use datafusion_physical_expr::utils::collect_columns; +use datafusion_physical_expr_common::physical_expr::{PhysicalExprRef, is_volatile}; -use crate::joins::utils::JoinFilter; +use crate::joins::utils::{ColumnIndex, JoinFilter}; /// Metadata for tracking filter results during deferred filtering /// @@ -389,3 +393,137 @@ pub fn filter_record_batch_by_join_type( JoinType::Inner => Ok(filter_record_batch(record_batch, corrected_mask)?), } } + +/// A join filter whose largest subexpressions over buffered columns alone +/// are lifted out, so they can be evaluated once per buffered row instead of +/// once per pair. +#[derive(Debug)] +pub(super) struct HoistedJoinFilter { + /// The filter over `schema`, reading each lifted subexpression's result + /// from a column + pub expression: PhysicalExprRef, + /// Schema of the batch `expression` is evaluated against + pub schema: SchemaRef, + /// Where each column of `schema` comes from + pub inputs: Vec, + /// The lifted subexpressions, over the buffered input's columns + pub buffered_exprs: Vec, +} + +#[derive(Debug)] +pub(super) enum HoistedFilterInput { + /// A column of the streamed or the buffered input + Column(ColumnIndex), + /// The result of `buffered_exprs[i]` + Buffered(usize), +} + +impl HoistedJoinFilter { + /// Lifts the non-volatile subexpressions of `filter` that read buffered + /// columns and no streamed ones. Returns `None` when there is none + /// besides bare columns. + pub fn try_new( + filter: &JoinFilter, + buffered_side: JoinSide, + buffered_schema: &Schema, + ) -> Result> { + let column_indices = filter.column_indices(); + let num_filter_columns = column_indices.len(); + let mut lifted: Vec = vec![]; + let expression = Arc::clone(filter.expression()) + .transform_down(|expr| { + if expr.downcast_ref::().is_some() || is_volatile(&expr) { + return Ok(Transformed::no(expr)); + } + let columns = collect_columns(&expr); + if columns.is_empty() + || columns + .iter() + .any(|c| column_indices[c.index()].side != buffered_side) + { + return Ok(Transformed::no(expr)); + } + let position = match lifted.iter().position(|e| **e == *expr) { + Some(position) => position, + None => { + lifted.push(Arc::clone(&expr)); + lifted.len() - 1 + } + }; + Ok(Transformed::new( + Arc::new(Column::new( + &format!("__buffered_{position}"), + num_filter_columns + position, + )), + true, + TreeNodeRecursion::Jump, + )) + })? + .data; + if lifted.is_empty() { + return Ok(None); + } + + let mut used: Vec = collect_columns(&expression) + .iter() + .map(|c| c.index()) + .filter(|&index| index < num_filter_columns) + .collect(); + used.sort_unstable(); + used.dedup(); + let mut fields: Vec = used + .iter() + .map(|&index| filter.schema().field(index).clone()) + .collect(); + let mut inputs: Vec = used + .iter() + .map(|&index| HoistedFilterInput::Column(column_indices[index].clone())) + .collect(); + for (position, expr) in lifted.iter().enumerate() { + fields.push(Field::new( + format!("__buffered_{position}"), + expr.data_type(filter.schema())?, + expr.nullable(filter.schema())?, + )); + inputs.push(HoistedFilterInput::Buffered(position)); + } + + let expression = expression + .transform_up(|expr| { + let Some(column) = expr.downcast_ref::() else { + return Ok(Transformed::no(expr)); + }; + let index = match used.binary_search(&column.index()) { + Ok(position) => position, + Err(_) => used.len() + column.index() - num_filter_columns, + }; + Ok(Transformed::yes( + Arc::new(Column::new(column.name(), index)) as PhysicalExprRef, + )) + })? + .data; + let buffered_exprs = lifted + .into_iter() + .map(|expr| { + expr.transform_up(|expr| { + let Some(column) = expr.downcast_ref::() else { + return Ok(Transformed::no(expr)); + }; + let index = column_indices[column.index()].index; + Ok(Transformed::yes(Arc::new(Column::new( + buffered_schema.field(index).name(), + index, + )) as PhysicalExprRef)) + }) + .map(|transformed| transformed.data) + }) + .collect::>>()?; + + Ok(Some(Self { + expression, + schema: Arc::new(Schema::new(fields)), + inputs, + buffered_exprs, + })) + } +} diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs index b9f48bac0d..d436b2d02a 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs @@ -30,8 +30,9 @@ use std::ops::Range; use std::sync::Arc; use crate::joins::sort_merge_join::filter::{ - FilterMetadata, deferred_filter_candidates, filter_record_batch_by_join_type, - get_corrected_filter_mask, needs_deferred_filtering, + FilterMetadata, HoistedFilterInput, HoistedJoinFilter, deferred_filter_candidates, + filter_record_batch_by_join_type, get_corrected_filter_mask, + needs_deferred_filtering, }; use crate::joins::sort_merge_join::metrics::SortMergeJoinMetrics; use crate::joins::utils::{JoinFilter, JoinKeyComparator}; @@ -156,6 +157,50 @@ impl StreamedBatch { } self.num_output_rows += 1; } + + /// Appends the pairs of the current streamed index with each buffered + /// index in `buffered_indices` of the buffered batch with + /// `buffered_batch_idx` index. + #[inline(never)] + fn append_output_pairs( + &mut self, + buffered_batch_idx: usize, + buffered_indices: Range, + batch_size: usize, + ) { + if buffered_indices.is_empty() { + return; + } + if self.output_indices.is_empty() + || self.buffered_batch_idx != Some(buffered_batch_idx) + { + debug_assert!( + batch_size > self.num_output_rows, + "batch_size ({batch_size}) must be > num_output_rows ({})", + self.num_output_rows + ); + let capacity = batch_size - self.num_output_rows; + self.output_indices.push(StreamedJoinedChunk { + buffered_batch_idx: Some(buffered_batch_idx), + streamed_indices: UInt64Builder::with_capacity(capacity), + buffered_indices: UInt64Builder::with_capacity(capacity), + }); + self.buffered_batch_idx = Some(buffered_batch_idx); + } + let current_chunk = self.output_indices.last_mut().unwrap(); + let num_pairs = buffered_indices.len(); + current_chunk + .streamed_indices + .append_value_n(self.idx as u64, num_pairs); + let buffered_builder = &mut current_chunk.buffered_indices; + buffered_builder.append_value_n(0, num_pairs); + let values = buffered_builder.values_slice_mut(); + let appended = values.len() - num_pairs..; + for (value, idx) in values[appended].iter_mut().zip(buffered_indices) { + *value = idx as u64; + } + self.num_output_rows += num_pairs; + } } /// Per-row filter outcome tracking for full outer joins. @@ -166,7 +211,7 @@ impl StreamedBatch { /// cannot distinguish "never matched" (handled by [`BufferedBatch::null_joined`]) /// from "matched but all filters failed" (must be emitted as null-joined). #[repr(u8)] -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] pub(super) enum FilterState { /// Row never appeared in a matched pair. Unvisited = 0, @@ -210,6 +255,18 @@ pub(super) struct BufferedBatch { /// but if batch is spilled to disk this property is preferable /// and less expensive pub num_rows: usize, + /// Results of the hoisted join filter subexpressions over a range of + /// this batch, tracked in `reserved_amount` + pub filter_cache: Option, +} + +/// Results of [`HoistedJoinFilter::buffered_exprs`] over the rows `range` +/// of a buffered batch. +#[derive(Debug)] +pub(super) struct BufferedFilterCache { + range: Range, + columns: Vec, + mem: usize, } impl BufferedBatch { @@ -248,6 +305,7 @@ impl BufferedBatch { reserved_amount: 0, join_filter_status: vec![FilterState::Unvisited; num_rows], num_rows, + filter_cache: None, }) } } @@ -290,6 +348,9 @@ pub(super) struct MaterializingSortMergeJoinStream { pub sort_options: Vec, /// optional join filter pub filter: Option, + /// `filter` with its buffered-only subexpressions lifted out, when it + /// has any + pub hoisted_filter: Option, /// How the join is performed pub join_type: JoinType, /// Cached `needs_deferred_filtering(filter, join_type)` β€” both inputs @@ -551,6 +612,18 @@ impl MaterializingSortMergeJoinStream { semi/anti/mark joins use BitwiseSortMergeJoinStream" ); let join_time = join_metrics.join_time(); + let buffered_side = if join_type == JoinType::Right { + JoinSide::Left + } else { + JoinSide::Right + }; + let hoisted_filter = filter + .as_ref() + .map(|filter| { + HoistedJoinFilter::try_new(filter, buffered_side, &buffered_schema) + }) + .transpose()? + .flatten(); let mut this = Self { sort_options, null_equality, @@ -568,6 +641,7 @@ impl MaterializingSortMergeJoinStream { on_buffered, deferred_filtering: needs_deferred_filtering(&filter, join_type), filter, + hoisted_filter, joined_record_batches: JoinedRecordBatches { joined_batches: new_output_coalescer(Arc::clone(&schema), batch_size), filter_metadata: FilterMetadata::new(), @@ -694,13 +768,24 @@ impl MaterializingSortMergeJoinStream { while !self.buffered_data.scanning_finished() && self.num_unfrozen_pairs() < self.batch_size { - let scanning_idx = self.buffered_data.scanning_idx(); - self.streamed_batch.append_output_pair( - Some(self.buffered_data.scanning_batch_idx), - Some(scanning_idx), - self.batch_size, - ); - self.buffered_data.scanning_advance(); + let range = &self.buffered_data.scanning_batch().range; + let scanning_idx = range.start + self.buffered_data.scanning_offset; + let num_pairs = (range.end - scanning_idx) + .min(self.batch_size - self.num_unfrozen_pairs()); + if num_pairs == 1 { + self.streamed_batch.append_output_pair( + Some(self.buffered_data.scanning_batch_idx), + Some(scanning_idx), + self.batch_size, + ); + } else { + self.streamed_batch.append_output_pairs( + self.buffered_data.scanning_batch_idx, + scanning_idx..scanning_idx + num_pairs, + self.batch_size, + ); + } + self.buffered_data.scanning_advance_by(num_pairs); } if self.num_unfrozen_pairs() >= self.batch_size { return false; @@ -1546,12 +1631,20 @@ impl MaterializingSortMergeJoinStream { as_uint64_array(&compute::concat(&refs)?)?.clone() }; - let Some(evaluation) = self.evaluate_join_filter( - &combined_left_indices, - matched_chunks, - total_matched_rows, - )? - else { + let evaluation = if self.cache_hoisted_filter(matched_chunks)? { + Some(self.evaluate_hoisted_join_filter( + &combined_left_indices, + matched_chunks, + total_matched_rows, + )?) + } else { + self.evaluate_join_filter( + &combined_left_indices, + matched_chunks, + total_matched_rows, + )? + }; + let Some(evaluation) = evaluation else { let output_batch = self.materialize_output_batch( &combined_left_indices, matched_chunks, @@ -1570,22 +1663,34 @@ impl MaterializingSortMergeJoinStream { let mut offset = 0usize; for (batch_idx, _left, right) in matched_chunks { let chunk_len = right.len(); - let buffered_batch = &mut self.buffered_data.batches[*batch_idx]; - - for i in 0..chunk_len { - if right.is_null(i) { - continue; - } - let idx = right.value(i) as usize; - match buffered_batch.join_filter_status[idx] { - FilterState::SomePassed => {} - _ if mask.value(offset + i) => { - buffered_batch.join_filter_status[idx] = - FilterState::SomePassed; + let status = + &mut self.buffered_data.batches[*batch_idx].join_filter_status; + let chunk_mask = mask.values().slice(offset, chunk_len); + if right.null_count() == 0 { + // A row's state only rises: every pair makes it at least + // AllFailed, and a passing one SomePassed. + if let Some(range) = is_contiguous_range(right) { + for state in &mut status[range] { + *state = (*state).max(FilterState::AllFailed); } - _ => { - buffered_batch.join_filter_status[idx] = - FilterState::AllFailed; + } else { + for &idx in right.values() { + let state = &mut status[idx as usize]; + *state = (*state).max(FilterState::AllFailed); + } + } + for i in chunk_mask.set_indices() { + status[right.value(i) as usize] = FilterState::SomePassed; + } + } else { + for (idx, passed) in right.iter().zip(chunk_mask.iter()) { + if let Some(idx) = idx { + let state = &mut status[idx as usize]; + *state = (*state).max(if passed { + FilterState::SomePassed + } else { + FilterState::AllFailed + }); } } } @@ -1747,20 +1852,7 @@ impl MaterializingSortMergeJoinStream { let filter_batch = RecordBatch::try_new(Arc::clone(filter.schema()), filter_columns)?; - let filter_result = filter - .expression() - .evaluate(&filter_batch)? - .into_array(filter_batch.num_rows())?; - let filter_result_mask = - datafusion_common::cast::as_boolean_array(&filter_result)?; - - // Convert NULL filter results to false β€” NULL means "not satisfied" - // per SQL semantics, same as Left/Right outer joins. - let mask = if filter_result_mask.null_count() > 0 { - compute::prep_null_mask_filter(filter_result_mask) - } else { - filter_result_mask.clone() - }; + let mask = evaluate_filter_mask(filter.expression(), &filter_batch)?; Ok(Some(FilterEvaluation { mask, @@ -1775,6 +1867,206 @@ impl MaterializingSortMergeJoinStream { })) } + /// Makes sure every buffered row the given pairs reference has its + /// hoisted filter subexpressions cached, computing them over the current + /// key group's rows of a batch the first time. Returns false when the + /// join has no hoisted filter or a pair falls outside the rows that can + /// be cached: a null-joined pair, or one from an earlier key group. + fn cache_hoisted_filter( + &mut self, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + ) -> Result { + let Some(hoisted) = &self.hoisted_filter else { + return Ok(false); + }; + for (batch_idx, _, right) in matched_chunks { + if right.null_count() > 0 { + return Ok(false); + } + let (min, max) = right + .values() + .iter() + .fold((u64::MAX, 0u64), |(lo, hi), &index| { + (lo.min(index), hi.max(index)) + }); + let covers = |range: &Range| { + range.start as u64 <= min && max < range.end as u64 + }; + let buffered_batch = &mut self.buffered_data.batches[*batch_idx]; + if buffered_batch + .filter_cache + .as_ref() + .is_some_and(|cache| covers(&cache.range)) + { + continue; + } + if !covers(&buffered_batch.range) { + return Ok(false); + } + let BufferedBatchState::InMemory(batch) = &buffered_batch.batch else { + return internal_err!( + "Buffered batch should have been unspilled before evaluating the join filter" + ); + }; + let range = buffered_batch.range.clone(); + let rows = batch.slice(range.start, range.len()); + let columns = hoisted + .buffered_exprs + .iter() + .map(|expr| expr.evaluate(&rows)?.into_array(rows.num_rows())) + .collect::>>()?; + let mem = columns.iter().map(|c| c.get_array_memory_size()).sum(); + if let Some(stale) = buffered_batch.filter_cache.take() { + self.reservation.shrink(stale.mem); + buffered_batch.reserved_amount -= stale.mem; + } + self.reservation.grow(mem); + buffered_batch.reserved_amount += mem; + buffered_batch.filter_cache = Some(BufferedFilterCache { + range, + columns, + mem, + }); + self.join_metrics + .peak_mem_used() + .set_max(self.reservation.size()); + } + Ok(true) + } + + /// Evaluates the hoisted join filter over the given pairs, whose + /// buffered rows [`Self::cache_hoisted_filter`] cached. + fn evaluate_hoisted_join_filter( + &self, + left_indices: &UInt64Array, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + total_matched_rows: usize, + ) -> Result { + let hoisted = self.hoisted_filter.as_ref().unwrap(); + let streamed_side = if self.join_type == JoinType::Right { + JoinSide::Right + } else { + JoinSide::Left + }; + let side_projection = |streamed: bool| { + let mut projection: Vec = hoisted + .inputs + .iter() + .filter_map(|input| match input { + HoistedFilterInput::Column(column) + if (column.side == streamed_side) == streamed => + { + Some(column.index) + } + _ => None, + }) + .collect(); + projection.sort_unstable(); + projection.dedup(); + projection + }; + let streamed_projection = side_projection(true); + let buffered_projection = side_projection(false); + + let streamed_columns = if streamed_projection.is_empty() { + vec![] + } else { + materialize_left_columns( + &self.streamed_batch.batch.project(&streamed_projection)?, + left_indices, + )? + }; + let buffered_columns = if buffered_projection.is_empty() { + vec![] + } else { + self.materialize_right_columns( + matched_chunks, + total_matched_rows, + Some(&buffered_projection), + )? + }; + let cached_columns = self + .gather_cached_filter_columns(matched_chunks, hoisted.buffered_exprs.len())?; + + let filter_columns = hoisted + .inputs + .iter() + .map(|input| match input { + HoistedFilterInput::Column(column) => { + let (projection, columns) = if column.side == streamed_side { + (&streamed_projection, &streamed_columns) + } else { + (&buffered_projection, &buffered_columns) + }; + let pos = projection.binary_search(&column.index).unwrap(); + Arc::clone(&columns[pos]) + } + HoistedFilterInput::Buffered(position) => { + Arc::clone(&cached_columns[*position]) + } + }) + .collect::>(); + let filter_batch = RecordBatch::try_new_with_options( + Arc::clone(&hoisted.schema), + filter_columns, + &RecordBatchOptions::new().with_row_count(Some(total_matched_rows)), + )?; + let mask = evaluate_filter_mask(&hoisted.expression, &filter_batch)?; + + Ok(FilterEvaluation { + mask, + streamed_columns: streamed_projection + .into_iter() + .zip(streamed_columns) + .collect(), + buffered_columns: buffered_projection + .into_iter() + .zip(buffered_columns) + .collect(), + }) + } + + /// Gathers the cached hoisted filter subexpression results of the + /// buffered rows the given pairs reference. + fn gather_cached_filter_columns( + &self, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + num_columns: usize, + ) -> Result> { + let mut pieces: Vec> = + vec![Vec::with_capacity(matched_chunks.len()); num_columns]; + for (batch_idx, _, right) in matched_chunks { + let Some(cache) = &self.buffered_data.batches[*batch_idx].filter_cache else { + return internal_err!("Hoisted join filter results were not cached"); + }; + let offset = cache.range.start; + if let Some(range) = is_contiguous_range(right) { + for (column, piece) in cache.columns.iter().zip(pieces.iter_mut()) { + piece.push(column.slice(range.start - offset, range.len())); + } + } else { + let indices = UInt64Array::from_iter_values( + right.values().iter().map(|&index| index - offset as u64), + ); + for (column, piece) in cache.columns.iter().zip(pieces.iter_mut()) { + piece.push(compute::take(column, &indices, None)?); + } + } + } + pieces + .into_iter() + .map(|piece| { + if piece.len() == 1 { + Ok(Arc::clone(&piece[0])) + } else { + let refs: Vec<&dyn Array> = + piece.iter().map(|a| a.as_ref()).collect(); + Ok(compute::concat(&refs)?) + } + }) + .collect() + } + /// Materializes the output batch of the given pairs, reusing the filter /// columns `evaluation` gathered for the same pairs. fn materialize_output_batch( @@ -1850,6 +2142,15 @@ impl MaterializingSortMergeJoinStream { ); } + // Multiple source batches each contributing runs of consecutive + // rows, as a large key group spanning several batches does: copy the + // runs instead of gathering row by row. + if matched_chunks.len() * MIN_AVG_RUN_LEN <= total_matched_rows + && let Some(columns) = self.copy_buffered_runs(matched_chunks, projection)? + { + return Ok(columns); + } + // Multiple source batches: map each buffered_batch_idx to a // contiguous source index. A null sentinel array is prepended as // source 0 only when some right index is actually null (an @@ -1995,6 +2296,54 @@ impl MaterializingSortMergeJoinStream { Ok(right_columns) } + /// Copies the buffered columns of the given pairs run by run when every + /// chunk is a run of consecutive rows of its batch. Returns `None` + /// otherwise. + #[inline(never)] + fn copy_buffered_runs( + &self, + matched_chunks: &[(usize, UInt64Array, UInt64Array)], + projection: Option<&[usize]>, + ) -> Result>> { + let Some(ranges) = matched_chunks + .iter() + .map(|(_, _, right)| is_contiguous_range(right)) + .collect::>>() + else { + return Ok(None); + }; + let batches = matched_chunks + .iter() + .map(|(batch_idx, _, _)| { + match &self.buffered_data.batches[*batch_idx].batch { + BufferedBatchState::InMemory(batch) => Ok(batch), + BufferedBatchState::Spilled(_) => internal_err!( + "Buffered batch should have been unspilled before fetching columns" + ), + } + }) + .collect::>>()?; + let col_indices: Vec = match projection { + Some(projection) => projection.to_vec(), + None => (0..self.buffered_schema.fields().len()).collect(), + }; + col_indices + .into_iter() + .map(|col_idx| { + let slices: Vec = batches + .iter() + .zip(&ranges) + .map(|(batch, range)| { + batch.column(col_idx).slice(range.start, range.len()) + }) + .collect(); + let refs: Vec<&dyn Array> = slices.iter().map(|a| a.as_ref()).collect(); + Ok(compute::concat(&refs)?) + }) + .collect::>>() + .map(Some) + } + fn filter_joined_batch(&mut self) -> Result { // Metadata should be aligned before processing self.joined_record_batches @@ -2101,6 +2450,30 @@ fn materialize_left_columns( } } +/// Average length of the runs of consecutive buffered rows from which +/// copying the runs beats gathering row by row. +const MIN_AVG_RUN_LEN: usize = 16; + +/// Evaluates a join filter expression into a mask, NULL results counting +/// as not satisfied. +fn evaluate_filter_mask( + expression: &Arc, + filter_batch: &RecordBatch, +) -> Result { + let filter_result = expression + .evaluate(filter_batch)? + .into_array(filter_batch.num_rows())?; + let filter_result_mask = datafusion_common::cast::as_boolean_array(&filter_result)?; + + // Convert NULL filter results to false β€” NULL means "not satisfied" + // per SQL semantics, same as Left/Right outer joins. + Ok(if filter_result_mask.null_count() > 0 { + compute::prep_null_mask_filter(filter_result_mask) + } else { + filter_result_mask.clone() + }) +} + /// Join filter result for the pairs of one freeze, with the filter columns /// gathered for it as `(column index, array)` per side. struct FilterEvaluation { @@ -2274,7 +2647,11 @@ impl BufferedData { } pub fn scanning_advance(&mut self) { - self.scanning_offset += 1; + self.scanning_advance_by(1); + } + + pub fn scanning_advance_by(&mut self, n: usize) { + self.scanning_offset += n; while !self.scanning_finished() && self.scanning_batch_finished() { self.scanning_batch_idx += 1; self.scanning_offset = 0; diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs index b2d9d91452..863367b71e 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs @@ -4389,10 +4389,14 @@ fn build_interval_buffered( } /// `t > eff AND t <= next_eff`, with the streamed table on `streamed_side`. +/// When `hoisted`, `t > eff + 0 AND t <= next_eff + 0 AND next_eff - eff > 0` +/// instead: the same filter with subexpressions over buffered columns alone, +/// which the join evaluates once per buffered row. fn build_interval_filter( streamed: &Schema, buffered: &Schema, streamed_side: JoinSide, + hoisted: bool, ) -> JoinFilter { let t_field = streamed .field_with_name("t") @@ -4440,23 +4444,37 @@ fn build_interval_filter( ) }; let t_col: PhysicalExprRef = Arc::new(Column::new("t", t_idx)); - JoinFilter::new( + let eff: PhysicalExprRef = Arc::new(Column::new("eff", eff_idx)); + let next_eff: PhysicalExprRef = Arc::new(Column::new("next_eff", next_eff_idx)); + let plus_zero = |expr: &PhysicalExprRef| -> PhysicalExprRef { Arc::new(BinaryExpr::new( - Arc::new(BinaryExpr::new( - Arc::clone(&t_col), - Operator::Gt, - Arc::new(Column::new("eff", eff_idx)), - )), + Arc::clone(expr), + Operator::Plus, + Arc::new(Literal::new(ScalarValue::Int64(Some(0)))), + )) + }; + let (lower, upper) = if hoisted { + (plus_zero(&eff), plus_zero(&next_eff)) + } else { + (Arc::clone(&eff), Arc::clone(&next_eff)) + }; + let mut expression: PhysicalExprRef = Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new(Arc::clone(&t_col), Operator::Gt, lower)), + Operator::And, + Arc::new(BinaryExpr::new(t_col, Operator::LtEq, upper)), + )); + if hoisted { + expression = Arc::new(BinaryExpr::new( + expression, Operator::And, Arc::new(BinaryExpr::new( - t_col, - Operator::LtEq, - Arc::new(Column::new("next_eff", next_eff_idx)), + Arc::new(BinaryExpr::new(next_eff, Operator::Minus, eff)), + Operator::Gt, + Arc::new(Literal::new(ScalarValue::Int64(Some(0)))), )), - )), - column_indices, - Arc::new(Schema::new(fields)), - ) + )); + } + JoinFilter::new(expression, column_indices, Arc::new(Schema::new(fields))) } /// Expected `(sid, bid)` output of the validity-interval join, in the @@ -4565,6 +4583,37 @@ async fn check_interval_join( batch_size: usize, spill: bool, case: &str, +) -> Result<()> { + for hoisted in [false, true] { + check_interval_join_with_filter( + join_type, + streamed, + buffered, + streamed_batch_rows, + buffered_batch_rows, + batch_size, + spill, + hoisted, + &format!("{case} hoisted={hoisted}"), + ) + .await?; + } + Ok(()) +} + +/// [`check_interval_join`] with the plain or the hoisted form of the filter +/// (see [`build_interval_filter`]). +#[expect(clippy::too_many_arguments)] +async fn check_interval_join_with_filter( + join_type: JoinType, + streamed: &[IntervalStreamedRow], + buffered: &[IntervalBufferedRow], + streamed_batch_rows: usize, + buffered_batch_rows: usize, + batch_size: usize, + spill: bool, + hoisted: bool, + case: &str, ) -> Result<()> { let streamed_plan = build_interval_streamed(streamed, streamed_batch_rows); let buffered_plan = build_interval_buffered(buffered, buffered_batch_rows); @@ -4577,6 +4626,7 @@ async fn check_interval_join( &streamed_plan.schema(), &buffered_plan.schema(), streamed_side, + hoisted, ); let (left, right) = if join_type == Right { (buffered_plan, streamed_plan) @@ -4757,6 +4807,50 @@ async fn join_filter_streamed_rows_across_freezes() -> Result<()> { Ok(()) } +/// The hoisted interval filter lifts `eff + 0`, `next_eff + 0` and the whole +/// `next_eff - eff > 0`, keeping `t` as the only input column; the plain one +/// has nothing to lift. +#[test] +fn hoisted_join_filter_lifts_buffered_subexpressions() -> Result<()> { + use super::filter::{HoistedFilterInput, HoistedJoinFilter}; + let streamed = build_interval_streamed(&[(1, 0, Some(1))], 1).schema(); + let buffered = build_interval_buffered(&interval_group(1, 0, 1, 1), 1).schema(); + for streamed_side in [JoinSide::Left, JoinSide::Right] { + let plain = build_interval_filter(&streamed, &buffered, streamed_side, false); + assert!( + HoistedJoinFilter::try_new(&plain, streamed_side.negate(), &buffered)? + .is_none() + ); + let filter = build_interval_filter(&streamed, &buffered, streamed_side, true); + let hoisted = + HoistedJoinFilter::try_new(&filter, streamed_side.negate(), &buffered)? + .unwrap(); + let lifted: Vec = hoisted + .buffered_exprs + .iter() + .map(|e| e.to_string()) + .collect(); + assert_eq!( + lifted, + vec!["eff@2 + 0", "next_eff@3 + 0", "next_eff@3 - eff@2 > 0"] + ); + assert!(matches!( + hoisted.inputs.as_slice(), + [ + HoistedFilterInput::Column(ColumnIndex { index: 2, side }), + HoistedFilterInput::Buffered(0), + HoistedFilterInput::Buffered(1), + HoistedFilterInput::Buffered(2), + ] if *side == streamed_side + )); + assert_eq!( + hoisted.expression.to_string(), + "t@0 > __buffered_0@1 AND t@0 <= __buffered_1@2 AND __buffered_2@3" + ); + } + Ok(()) +} + /// Key groups larger than the batch size on both sides, with streamed rows /// passing the filter for none, one or several buffered rows, and buffered /// key groups no streamed row matches between them. A FULL join stages the From 8b4fc571d6a6cedd01cf213b6c1042b36ad6b783 Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 02:05:54 +0100 Subject: [PATCH 05/13] feat: drop the sort-merge join condition cost after the join filter fix With the join filter tested on each pair of equal keys before the output is built, a native sort-merge join with a condition costs about what Spark does for any condition shape and any number of pairs per key: on the cluster, a left join under a validity interval or a band cost 5.7 to 22 ns per pair natively against 9.6 to 15 in Spark at 15 to 519 output leaves, an inner one 7.4 against 19.3, and one with 100 pairs per key about Spark's price per output row. The smjCondition class, its Line(3000, 80, 0, 850, 23) and the validity-interval exemption (JoinConditionShape) are removed, so a sort-merge join is priced by smj alone with or without a condition. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/source/user-guide/latest/tuning.md | 6 +- .../scala/org/apache/comet/CometConf.scala | 10 +- .../comet/rules/CostBasedEngineChoice.scala | 20 -- .../apache/comet/rules/EngineCostTable.scala | 23 +- .../comet/rules/JoinConditionShape.scala | 218 ------------------ .../rules/CostBasedEngineChoiceSuite.scala | 85 +++---- 6 files changed, 44 insertions(+), 318 deletions(-) delete mode 100644 spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala diff --git a/docs/source/user-guide/latest/tuning.md b/docs/source/user-guide/latest/tuning.md index 13fabf00bb..4b6d929ee9 100644 --- a/docs/source/user-guide/latest/tuning.md +++ b/docs/source/user-guide/latest/tuning.md @@ -602,11 +602,7 @@ prices per row from a table of measurements: `c0 + k0*L + k1*L*min(L, 600)` ns n Spark one. - `sort` over every leaf of the sorted rows (Comet also 0.15 ns per byte beyond 12 per leaf); `sortSpill` adds the price of a spill to a fraction `sortSpillFraction` of rows, none by default. `smj` prices a sort-merge join and `bhj` - the probe side of a broadcast hash join, over every output leaf. A sort-merge join with a join condition adds - `smjCondition`, also over every output leaf: Comet builds every pair of rows of equal keys before the condition drops - them, about 3.5 times Spark's price on band joins. A condition that is one validity interval, `L <= V < U` with `V` - from one input and `L` and `U` from the other (casts, date truncations, `COALESCE(U, literal)` and `U IS NULL OR` - allowed), adds nothing when `U` is not `L` shifted by a constant through the projections below the join. + the probe side of a broadcast hash join, over every output leaf; a join condition adds nothing. - `predicate` for filters, over the leaves their predicate references, plus a pass-through of 1.5 ns per output leaf natively and none in Spark. A native filter over a native scan, and the native projects over it, stay native whatever their prices (`keepFiltersOverNativeScans=false` lets them move): the rows a filter drops are not estimated, diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 9c10c3e88e..acc6e57215 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -689,11 +689,11 @@ object CometConf extends ShimCometConf { "(or `k0,k1`, keeping c0), or spark, with `c0,k` for c0 + k*L ns per row, where L is " + "the number of leaf columns the class prices (for the functions of an aggregate or a " + "window, the number of functions). The classes are shuffleWrite, shuffleRead, sort, " + - "sortSpill, smj, smjCondition, bhj, predicate, projectPassThrough, expression, agg, " + - "aggObjectHash, aggDeclarative, aggCollectList, aggCollectSet, aggPercentile, " + - "aggPercentileApprox, aggOther, aggArrayKey, codegenDispatch, window, " + - "windowAggregate, windowOffset, windowRank, wglPartial, wglFinal, expand, generate, " + - "rowLocal, the comet-only c2r and r2c, and the spark-only " + + "sortSpill, smj, bhj, predicate, projectPassThrough, expression, agg, aggObjectHash, " + + "aggDeclarative, aggCollectList, aggCollectSet, aggPercentile, aggPercentileApprox, " + + "aggOther, aggArrayKey, codegenDispatch, window, windowAggregate, windowOffset, " + + "windowRank, wglPartial, wglFinal, expand, generate, rowLocal, the comet-only c2r and " + + "r2c, and the spark-only " + "expressionOverScan, aggDeclarativeNoCodegen, expandNoCodegen and " + "generateNoCodegen. A row whose leaves are a fraction f inside structs, arrays or " + "maps costs (1 - f) times the flat price plus f times the nested one. The scalars " + diff --git a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala index 0b66d06b94..27aa5960a1 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala @@ -33,7 +33,6 @@ import org.apache.spark.sql.comet.{CometExec, CometFilterExec, CometHashAggregat import org.apache.spark.sql.execution.{ColumnarToRowTransition, ExpandExec, FilterExec, ProjectExec, SortExec, SparkPlan} import org.apache.spark.sql.execution.aggregate.BaseAggregateExec import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike -import org.apache.spark.sql.execution.joins.SortMergeJoinExec import org.apache.spark.sql.execution.window.WindowExec import org.apache.spark.sql.internal.SQLConf @@ -124,19 +123,6 @@ class EngineCostModel( private def aggregateShare(agg: BaseAggregateExec): Double = if (agg.aggregateExpressions.exists(_.mode == Complete)) 1.0 else 0.5 - /** Whether `condition` of `join`, run as `plan`, is one validity interval. */ - def validityInterval( - condition: Expression, - join: SortMergeJoinExec, - plan: SparkPlan): Boolean = { - val sides = if (plan.children.size == 2) plan.children else join.children - JoinConditionShape.isValidityInterval( - condition, - sides.head, - sides(1), - JoinConditionShape.aliases(plan.children ++ join.children)) - } - private def classTerms(costClass: CostClass, plan: SparkPlan, engine: Engine): Seq[Term] = { val op = sparkOperator(plan) lazy val out = widthOf(op.output) @@ -184,12 +170,6 @@ class EngineCostModel( functions.distinct.map { c => Term(c, Width(functions.size, 0), share * functions.count(_ == c)) } - case (SmjCondition, join: SortMergeJoinExec) => - if (join.condition.exists(c => !validityInterval(c, join, plan))) { - Seq(Term(SmjCondition, out)) - } else { - Nil - } case (AggObjectHash, agg: BaseAggregateExec) => Seq(Term(AggObjectHash, Width(0, 0), aggregateShare(agg))) case _ => Seq(Term(costClass, out)) diff --git a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala index d900df08d2..34761af94a 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala @@ -172,7 +172,6 @@ object EngineCostTable { case object Sort extends CostClass("sort") case object SortSpill extends CostClass("sortSpill") case object Smj extends CostClass("smj") - case object SmjCondition extends CostClass("smjCondition") case object Bhj extends CostClass("bhj") case object Predicate extends CostClass("predicate") case object ProjectPassThrough extends CostClass("projectPassThrough") @@ -208,7 +207,6 @@ object EngineCostTable { Sort, SortSpill, Smj, - SmjCondition, Bhj, Predicate, ProjectPassThrough, @@ -322,19 +320,11 @@ object EngineCostTable { * cannot be predicted. * - `smj`: the join over its sorted inputs, every output leaf, noisy. `bhj`: the probe side, * every output leaf; the nested `bhj` is noisy (Comet 0 to 350, Spark 50 to 4200 ns) and - * takes the flat line. `smjCondition`: what a join condition adds to `smj`, every output - * leaf. Comet joins every pair of rows of equal keys into a batch of all its output columns - * before the condition drops most of them; Spark tests the condition on the pair first. In - * a local micro-benchmark of left joins on a 14-day band (1 to 4% of 20 to 100 million - * pairs passing), Comet cost what the join without the condition cost, 0.44 to 1.2 us per - * output row and 21 ns per output leaf, Spark 3.6 to 4 times less (0.12 to 0.30 us, 4.5 ns - * per leaf); inner joins and small key groups ran 1 to 1.8 times slower. fbj_order_type's - * left band joins took 23 to 28 us per output row natively against 6 to 8 in Spark. The - * pairs per output row are not estimated, so the line keeps the ratio of 3.5 at a price - * between the two, high enough to outweigh `smj` and the conversions around the join at any - * width. A condition that is one validity interval ([[JoinConditionShape]]) with both - * bounds from one row of one input, where Comet costs about what Spark does, adds nothing; - * bounds from two inputs joined below make a band again. + * takes the flat line. A join condition adds nothing to `smj`: the native join tests it on + * each pair of rows of equal keys before building the output, and on the cluster a left + * join under a validity interval or a band cost 5.7 to 22 ns per pair natively against 9.6 + * to 15 in Spark at 15 to 519 output leaves, an inner one 7.4 against 19.3, and one with + * 100 pairs per key about what Spark costs per output row. * - `predicate`: a filter over the leaves its predicate reads, Spark noisy (maxrel 0.5 to * 0.8); passing rows costs the filter scalars. * - `projectPassThrough`: a project over every output leaf. Spark copies the row only when @@ -418,7 +408,6 @@ object EngineCostTable { (R2C, Flat) -> Line(3.4, 7.67, 0.036, 0, 0), (R2C, Nested) -> Line(0, 9.74, 0.040, 0, 0)) ++ both(ExprOverScan, Line(0, 0, 0, 2.5, 0)) ++ - both(SmjCondition, Line(3000, 80, 0, 850, 23)) ++ both(AggObjectHash, Line(1000, 0, 0, 2500, 0)) ++ both(AggDeclarative, Line(6, 0, 0, 15, 0)) ++ both(AggDeclarativeNoCodegen, Line(0, 0, 0, 170, 0)) ++ @@ -481,7 +470,7 @@ object EngineCostTable { */ val operatorClasses: Map[String, Seq[CostClass]] = Map( "SortExec" -> Seq(Sort), - "SortMergeJoinExec" -> Seq(Smj, SmjCondition), + "SortMergeJoinExec" -> Seq(Smj), "BroadcastHashJoinExec" -> Seq(Bhj), "WindowExec" -> Seq(Window), "WindowGroupLimitExec" -> Seq(WglPartial, WglFinal), diff --git a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala b/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala deleted file mode 100644 index de0d07eff9..0000000000 --- a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala +++ /dev/null @@ -1,218 +0,0 @@ -/* - * Licensed to the Apache Software Foundation (ASF) under one - * or more contributor license agreements. See the NOTICE file - * distributed with this work for additional information - * regarding copyright ownership. The ASF licenses this file - * to you under the Apache License, Version 2.0 (the - * "License"); you may not use this file except in compliance - * with the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, - * software distributed under the License is distributed on an - * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY - * KIND, either express or implied. See the License for the - * specific language governing permissions and limitations - * under the License. - */ - -package org.apache.comet.rules - -import scala.collection.mutable - -import org.apache.spark.sql.catalyst.expressions.{Add, AddMonths, Alias, And, Attribute, AttributeSet, Cast, Coalesce, DateAdd, DateAddInterval, DateAddYMInterval, DateSub, Expression, ExprId, GreaterThan, GreaterThanOrEqual, IsNull, LessThan, LessThanOrEqual, Or, Subtract, TimestampAddYMInterval, TruncDate, TruncTimestamp, WindowExpression} -import org.apache.spark.sql.comet.CometExec -import org.apache.spark.sql.execution.{ProjectExec, SparkPlan, UnionExec} -import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, QueryStageExec} -import org.apache.spark.sql.execution.exchange.ReusedExchangeExec - -/** - * Recognizes a join condition that is a single validity interval: `L <= V < U` with `V` from one - * side, `L` and `U` from the other, `U` not derived from `L` by a constant offset, and `L` and - * `U` read from one row of one input of the other side, with no join or union between them. - */ -object JoinConditionShape { - - private case class Bound(small: Expression, big: Expression, nullCheck: Option[Expression]) - - private def unwrap(e: Expression): Expression = e match { - case c: Cast => unwrap(c.child) - case t: TruncTimestamp => unwrap(t.timestamp) - case t: TruncDate => unwrap(t.date) - case c: Coalesce if c.children.size > 1 && c.children.tail.forall(_.foldable) => - unwrap(c.children.head) - case other => other - } - - private val timeAddClasses = Set("TimeAdd", "TimestampAddInterval") - - private def offsetBase(e: Expression): Option[Expression] = e match { - case a: Add if a.right.foldable => Some(a.left) - case a: Add if a.left.foldable => Some(a.right) - case s: Subtract if s.right.foldable => Some(s.left) - case d: DateAdd if d.days.foldable => Some(d.startDate) - case d: DateSub if d.days.foldable => Some(d.startDate) - case t - if timeAddClasses(t.getClass.getSimpleName) && t.children.size >= 2 && - t.children(1).foldable => - Some(t.children.head) - case d: DateAddInterval if d.interval.foldable => Some(d.start) - case d: DateAddYMInterval if d.interval.foldable => Some(d.date) - case t: TimestampAddYMInterval if t.interval.foldable => Some(t.timestamp) - case m: AddMonths if m.numMonths.foldable => Some(m.startDate) - case _ => None - } - - private def strip(e: Expression): Expression = { - val u = unwrap(e) - offsetBase(u).map(strip).getOrElse(u) - } - - private def comparison(e: Expression): Option[(Expression, Expression)] = e match { - case LessThan(a, b) => Some((a, b)) - case LessThanOrEqual(a, b) => Some((a, b)) - case GreaterThan(a, b) => Some((b, a)) - case GreaterThanOrEqual(a, b) => Some((b, a)) - case _ => None - } - - private def bound(e: Expression): Option[Bound] = e match { - case Or(IsNull(x), c) => comparison(c).map { case (s, b) => Bound(s, b, Some(x)) } - case Or(c, IsNull(x)) => comparison(c).map { case (s, b) => Bound(s, b, Some(x)) } - case c => comparison(c).map { case (s, b) => Bound(s, b, None) } - } - - private def conjuncts(e: Expression): Seq[Expression] = e match { - case And(a, b) => conjuncts(a) ++ conjuncts(b) - case other => Seq(other) - } - - private def same(a: Expression, b: Expression): Boolean = unwrap(a).semanticEquals(unwrap(b)) - - /** The aliases defined in `plans` and below, by the id of the attribute each defines. */ - def aliases(plans: Seq[SparkPlan]): Map[ExprId, Expression] = { - val result = mutable.Map[ExprId, Expression]() - def visit(node: SparkPlan): Unit = { - val operators = node match { - case c: CometExec => Seq(c, c.originalPlan) - case other => Seq(other) - } - operators - .flatMap(_.expressions) - .foreach(_.foreach { - case a: Alias => result.getOrElseUpdate(a.exprId, a.child) - case _ => - }) - node match { - case s: QueryStageExec => visit(s.plan) - case _ => node.children.foreach(visit) - } - } - plans.foreach(visit) - result.toMap - } - - private def sources(e: Expression, aliases: Map[ExprId, Expression], depth: Int): Set[ExprId] = - strip(e) match { - case a: Attribute => - aliases.get(a.exprId) match { - case Some(_: WindowExpression) => Set(a.exprId) - case Some(child) if depth < 64 => sources(child, aliases, depth + 1) - case _ => Set(a.exprId) - } - case other => other.references.map(_.exprId).toSet - } - - private def inputs(node: SparkPlan): Seq[SparkPlan] = node match { - case s: QueryStageExec => Seq(s.plan) - case a: AdaptiveSparkPlanExec => Seq(a.executedPlan) - case other => other.children - } - - private def original(node: SparkPlan): SparkPlan = node match { - case c: CometExec => c.originalPlan - case other => other - } - - private def outputs(node: SparkPlan, id: ExprId): Boolean = node.output.exists(_.exprId == id) - - private def producers(node: SparkPlan, id: ExprId, depth: Int): Option[Seq[List[SparkPlan]]] = - if (depth > 256 || !outputs(node, id)) { - None - } else { - node match { - case r: ReusedExchangeExec => - val i = r.output.indexWhere(_.exprId == id) - producers(r.child, r.child.output(i).exprId, depth + 1).map(_.map(r :: _)) - case _ => - inputs(node).find(outputs(_, id)) match { - case Some(child) => producers(child, id, depth + 1).map(_.map(node :: _)) - case None => - val projected = original(node) match { - case p: ProjectExec => - p.projectList.collectFirst { case a: Alias if a.exprId == id => a.child } - case _ => None - } - (projected, inputs(node)) match { - case (Some(e), Seq(child)) if e.references.nonEmpty => - val traced = e.references.toSeq.map(a => producers(child, a.exprId, depth + 1)) - if (traced.forall(_.isDefined)) { - Some(traced.flatMap(_.get).map(node :: _)) - } else { - None - } - case _ => Some(Seq(List(node))) - } - } - } - } - - private def oneRow(side: SparkPlan, refs: AttributeSet): Boolean = { - val traced = refs.toSeq.map(a => producers(side, a.exprId, 0)) - traced.nonEmpty && traced.forall(_.isDefined) && { - val paths = traced.flatMap(_.get) - val common = (0 until paths.map(_.size).min) - .takeWhile(i => paths.forall(_(i) eq paths.head(i))) - .size - val below = paths.map(_.drop(common)) - val heads = below.flatMap(_.headOption) - heads.forall(_ eq heads.head) && - paths.forall(_.forall(n => !original(n).isInstanceOf[UnionExec])) && - below.forall(_.forall(n => inputs(n).size <= 1)) - } - } - - /** - * Whether `condition` of a join between `left` and `right` is one validity interval whose - * bounds `aliases` does not derive from one another and that read one row of one input. - */ - def isValidityInterval( - condition: Expression, - left: SparkPlan, - right: SparkPlan, - aliases: Map[ExprId, Expression]): Boolean = { - conjuncts(condition).map(bound) match { - case Seq(Some(first), Some(second)) => - val shapes = Seq((first, second), (second, first)).collect { - case (lower, upper) if same(lower.big, upper.small) && lower.nullCheck.isEmpty => - (upper.small, lower.small, upper.big, upper.nullCheck) - } - shapes.exists { case (v, l, u, nullCheck) => - val (vRefs, lRefs, uRefs) = (v.references, l.references, u.references) - def onOtherSides(side: SparkPlan, other: SparkPlan): Boolean = - vRefs.subsetOf(side.outputSet) && lRefs.subsetOf(other.outputSet) && - uRefs.subsetOf(other.outputSet) - val boundsSide = - if (onOtherSides(left, right)) Some(right) - else if (onOtherSides(right, left)) Some(left) - else None - vRefs.nonEmpty && lRefs.nonEmpty && uRefs.nonEmpty && - nullCheck.forall(same(_, u)) && - sources(l, aliases, 0).intersect(sources(u, aliases, 0)).isEmpty && - boundsSide.exists(oneRow(_, lRefs ++ uRefs)) - } - case _ => false - } - } -} diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index 2ccc7f6f19..b38dc54b57 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -888,7 +888,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } } - test("a sort-merge join adds smjCondition only with a join condition") { + test("a sort-merge join with a join condition costs what one without does") { withTables { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", @@ -902,28 +902,23 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } val out = Width(4, 0) val equi = join("SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k") - assert(model.terms(equi, Engine.Comet) == Seq(Term(Smj, out))) - assert(model.terms(equi, Engine.Spark) == Seq(Term(Smj, out))) - same(model.operatorPrice(equi, Engine.Comet), EngineCostTable.default.comet(Smj, out)) - val band = join( "SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k " + "AND a.v <= b.v AND b.v <= a.v + 13") - assert(model.terms(band, Engine.Comet) == Seq(Term(Smj, out), Term(SmjCondition, out))) - assert(model.terms(band, Engine.Spark) == Seq(Term(Smj, out), Term(SmjCondition, out))) - val table = EngineCostTable.default + assert(band.originalPlan.asInstanceOf[SortMergeJoinExec].condition.isDefined) + for (engine <- Engine.all; j <- Seq(equi, band)) { + assert(model.terms(j, engine) == Seq(Term(Smj, out))) + } + same(model.operatorPrice(equi, Engine.Comet), EngineCostTable.default.comet(Smj, out)) for (engine <- Engine.all) { - val price: (CostClass, Width) => Double = - if (engine == Engine.Comet) table.comet else table.spark - same(model.operatorPrice(band, engine), price(Smj, out) + price(SmjCondition, out)) + same(model.operatorPrice(band, engine), model.operatorPrice(equi, engine)) } } } } for (aqe <- Seq("false", "true")) { - test( - s"a sort-merge join with a join condition runs in Spark, without one natively (AQE=$aqe)") { + test(s"a sort-merge join with a join condition stays native, as one without (AQE=$aqe)") { withTables { withAqe( aqe, @@ -935,18 +930,11 @@ class CostBasedEngineChoiceSuite extends CometTestBase { count(plan) { case j: SortMergeJoinExec => j }) val equi = "SELECT a.k, a.v, b.s FROM t a JOIN t b ON a.k = b.k" val band = equi + " AND a.v <= b.v AND b.v <= a.v + 13" - val (equiOff, equiOn) = offAndOn(run(equi)) - assert(joins(equiOff) == (1, 0), s"plan:\n$equiOff") - assert(joins(equiOn) == (1, 0), s"plan:\n$equiOn") - assert(cometOperatorNames(equiOff) == cometOperatorNames(equiOn), s"$equiOff\n$equiOn") - val (bandOff, bandOn) = offAndOn(run(band)) - assert(joins(bandOff) == (1, 0), s"plan:\n$bandOff") - assert(joins(bandOn) == (0, 1), s"plan:\n$bandOn") - withSQLConf( - flag -> "true", - costTable -> "smjCondition.comet=0,0,0;smjCondition.spark=0,0") { - val plan = run(band) - assert(joins(plan) == (1, 0), s"plan:\n$plan") + for (query <- Seq(equi, band)) { + val (off, on) = offAndOn(run(query)) + assert(joins(off) == (1, 0), s"plan:\n$off") + assert(joins(on) == (1, 0), s"plan:\n$on") + assert(cometOperatorNames(off) == cometOperatorNames(on), s"$off\n$on") } } } @@ -1007,7 +995,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { "SELECT o.k, o.v, g.v AS gv FROM ev o JOIN dim p ON o.k = p.k " + "LEFT JOIN ev g ON o.k = g.k AND p.completed_dt < g.v AND o.v > g.v" - private val chargedConditions = Seq( + private val otherConditions = Seq( crossInputBounds, "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k " + "AND e.vs BETWEEN d.ls - 2592000 AND d.ls", @@ -1032,34 +1020,28 @@ class CostBasedEngineChoiceSuite extends CometTestBase { case j: SortMergeJoinExec if j.condition.isDefined => j } assert(joins.nonEmpty, s"plan:\n$plan") - joins.map(j => model.terms(j, Engine.Comet).filter(_.costClass == SmjCondition)) + joins.map(j => model.terms(j, Engine.Comet)) } for (aqe <- Seq("false", "true")) { - test( - s"a validity interval adds no smjCondition, a band or another condition does (AQE=$aqe)") { + test(s"a join condition of any shape adds nothing to smj (AQE=$aqe)") { withIntervals { withAqe( aqe, flag -> "false", SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { - for (condition <- validityIntervals; joinType <- Seq("JOIN", "LEFT JOIN")) { - val query = - s"SELECT e.k, e.v, d.l FROM ev e $joinType dim d ON e.k = d.k AND $condition" - assert(conditionTerms(query).forall(_.isEmpty), query) - } - for (query <- intervalQueries) { - assert(conditionTerms(query).forall(_.isEmpty), query) - } - for (query <- chargedConditions) { - assert(conditionTerms(query).exists(_.nonEmpty), query) + val intervals = + for (condition <- validityIntervals; joinType <- Seq("JOIN", "LEFT JOIN")) + yield s"SELECT e.k, e.v, d.l FROM ev e $joinType dim d ON e.k = d.k AND $condition" + for (query <- intervals ++ intervalQueries ++ otherConditions) { + assert(conditionTerms(query).forall(_.map(_.costClass) == Seq(Smj)), query) } } } } - test(s"a validity interval join stays native, a band join runs in Spark (AQE=$aqe)") { + test(s"validity interval, band and cross-input joins stay native (AQE=$aqe)") { withIntervals { withAqe( aqe, @@ -1071,17 +1053,14 @@ class CostBasedEngineChoiceSuite extends CometTestBase { count(plan) { case j: SortMergeJoinExec => j }) val interval = "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + "AND e.v >= d.l AND e.v < COALESCE(d.u, timestamp '9999-12-31')" - val (intervalOff, intervalOn) = offAndOn(run(interval)) - assert(joins(intervalOff) == (1, 0), s"plan:\n$intervalOff") - assert(joins(intervalOn) == (1, 0), s"plan:\n$intervalOn") - val (bandOff, bandOn) = offAndOn(run(chargedConditions(1))) - assert(joins(bandOff) == (1, 0), s"plan:\n$bandOff") - assert(joins(bandOn) == (0, 1), s"plan:\n$bandOn") - val (crossOff, crossOn) = offAndOn(run(crossInputBounds)) - assert(joins(crossOff) == (2, 0), s"plan:\n$crossOff") - assert( - count(crossOn) { case j: SortMergeJoinExec if j.condition.isDefined => j } == 1, - s"plan:\n$crossOn") + for ((query, native) <- Seq( + interval -> 1, + otherConditions(1) -> 1, + crossInputBounds -> 2)) { + val (off, on) = offAndOn(run(query)) + assert(joins(off) == (native, 0), s"plan:\n$off") + assert(joins(on) == (native, 0), s"plan:\n$on") + } } } } @@ -1405,7 +1384,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { |""".stripMargin for (aqe <- Seq("false", "true")) { - test(s"fbj_order_type's band join under a root repartition runs in Spark (AQE=$aqe)") { + test(s"fbj_order_type's band join under a root repartition stays native (AQE=$aqe)") { withFbj { withAqe( aqe, @@ -1434,7 +1413,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { conditions(off) == (Seq("band", "forward", "replenishment", "replenishment"), Nil), s"plan:\n$off") val (native, spark) = conditions(on) - assert(spark.contains("band"), s"plan:\n$on") + assert(native.contains("band") && !spark.contains("band"), s"plan:\n$on") assert(native.contains("replenishment"), s"plan:\n$on") } } From f5eab1948dba41a2a3454b154de84594abc58c5c Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 02:40:20 +0100 Subject: [PATCH 06/13] test: drop an unneeded string interpolator Co-Authored-By: Claude Opus 5.5 (1M context) --- spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala index 35b170bf4e..8bb8d5ddf7 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala @@ -1190,7 +1190,7 @@ class CometJoinSuite extends CometTestBase { checkSparkAnswerAndOperator( sql(s"SELECT o.*, r.* FROM rates r RIGHT JOIN orders o ON $condition")) checkSparkAnswerAndOperator( - sql(s"SELECT o.id, count(r.rate), sum(o.p28) " + + sql("SELECT o.id, count(r.rate), sum(o.p28) " + s"FROM orders o LEFT JOIN rates r ON $condition GROUP BY o.id")) } } From fb34a5fbdd37a429b1648ac4dbed7aaac2670447 Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 08:25:39 +0100 Subject: [PATCH 07/13] bench: sort-merge join whose filter casts a streamed column New bench spark-expr/benches/smj_streamed_filter.rs, shaped like the fbj_order_type join `cast(o.completed as timestamp) < b.t AND o.t > b.t`: string key, 4 keys x 1000 (200) streamed rows with 16 payload columns x 2000 buffered rows of 3 columns, batch_size 8192. `completed` is a `yyyy-MM-dd HH:mm:ss` string (as in the model) or a date. ns per candidate pair at f5eab1948: few_string (1% pass, 8M pairs): LEFT 167, FULL 168, INNER 167 few_date (1% pass, 8M pairs): LEFT 32, FULL 32, INNER 31 all_string (all pass, 1.6M): LEFT 211, FULL 211, INNER 192 Co-Authored-By: Claude Opus 5.5 (1M context) --- native/spark-expr/Cargo.toml | 4 + .../spark-expr/benches/smj_streamed_filter.rs | 297 ++++++++++++++++++ 2 files changed, 301 insertions(+) create mode 100644 native/spark-expr/benches/smj_streamed_filter.rs diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index c86ac91696..bb85467b9d 100644 --- a/native/spark-expr/Cargo.toml +++ b/native/spark-expr/Cargo.toml @@ -446,3 +446,7 @@ harness = false [[bench]] name = "smj_interval_filter" harness = false + +[[bench]] +name = "smj_streamed_filter" +harness = false diff --git a/native/spark-expr/benches/smj_streamed_filter.rs b/native/spark-expr/benches/smj_streamed_filter.rs new file mode 100644 index 0000000000..9d6455faca --- /dev/null +++ b/native/spark-expr/benches/smj_streamed_filter.rs @@ -0,0 +1,297 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Sort-merge join whose filter casts a streamed column, shaped like a Spark +//! `cast(o.completed as timestamp) < b.t AND o.t > b.t` lookup of earlier +//! events: a string key, large key groups on both sides, a wide streamed +//! side and a narrow buffered one. `completed` is either a +//! `yyyy-MM-dd HH:mm:ss` string or a date. + +use arrow::array::{ + ArrayRef, Date32Array, Int64Array, RecordBatch, StringArray, TimestampMicrosecondArray, +}; +use arrow::compute::SortOptions; +use arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit}; +use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; +use datafusion::common::{JoinSide, JoinType, NullEquality}; +use datafusion::datasource::memory::MemorySourceConfig; +use datafusion::execution::TaskContext; +use datafusion::logical_expr::Operator; +use datafusion::physical_expr::expressions::{BinaryExpr, Column}; +use datafusion::physical_expr::PhysicalExpr; +use datafusion::physical_plan::joins::utils::{ColumnIndex, JoinFilter}; +use datafusion::physical_plan::joins::SortMergeJoinExec; +use datafusion::physical_plan::{collect, ExecutionPlan}; +use datafusion::prelude::SessionConfig; +use datafusion_comet_spark_expr::{Cast, EvalMode, SparkCastOptions}; +use std::sync::Arc; +use tokio::runtime::Runtime; + +const MICROS_PER_DAY: i64 = 86_400_000_000; +const BASE_MICROS: i64 = 20_000 * MICROS_PER_DAY; +const BATCH_SIZE: usize = 8192; + +fn timestamp_type() -> DataType { + DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())) +} + +fn key_of(k: usize) -> String { + format!("V{k:08}") +} + +#[derive(Clone, Copy)] +enum Completed { + String, + Date, +} + +/// Streamed side: `key`, the event time `t`, `completed` and +/// `payload_cols` payload columns alternating Int64 and Utf8. +/// +/// With `all_pass`, `completed` precedes and `t` follows every buffered +/// time of the key; otherwise `completed` lies `window_days` before `t`. +fn build_streamed( + keys: usize, + rows_per_key: usize, + group_days: i64, + window_days: i64, + all_pass: bool, + completed: Completed, + payload_cols: usize, +) -> (SchemaRef, Vec) { + let completed_type = match completed { + Completed::String => DataType::Utf8, + Completed::Date => DataType::Date32, + }; + let mut fields = vec![ + Field::new("key", DataType::Utf8, false), + Field::new("t", timestamp_type(), true), + Field::new("completed", completed_type, true), + ]; + for c in 0..payload_cols { + let data_type = if c % 2 == 0 { + DataType::Int64 + } else { + DataType::Utf8 + }; + fields.push(Field::new(format!("p{c}"), data_type, true)); + } + let schema = Arc::new(Schema::new(fields)); + + let num_rows = keys * rows_per_key; + let rows: Vec<(usize, i64, i64)> = (0..num_rows) + .map(|i| { + if all_pass { + let t = BASE_MICROS + (group_days + 1) * MICROS_PER_DAY + i as i64; + (i / rows_per_key, t, BASE_MICROS - MICROS_PER_DAY) + } else { + let day = window_days + (i as i64 * 7919) % (group_days - window_days); + let t = BASE_MICROS + day * MICROS_PER_DAY + (i as i64 * 104_729) % MICROS_PER_DAY; + (i / rows_per_key, t, t - window_days * MICROS_PER_DAY) + } + }) + .collect(); + let completed_column: ArrayRef = match completed { + Completed::String => Arc::new(StringArray::from_iter_values(rows.iter().map(|r| { + chrono::DateTime::from_timestamp_micros(r.2) + .unwrap() + .format("%Y-%m-%d %H:%M:%S") + .to_string() + }))), + Completed::Date => Arc::new(Date32Array::from_iter_values( + rows.iter().map(|r| r.2.div_euclid(MICROS_PER_DAY) as i32), + )), + }; + let mut columns: Vec = vec![ + Arc::new(StringArray::from_iter_values( + rows.iter().map(|r| key_of(r.0)), + )), + Arc::new( + TimestampMicrosecondArray::from_iter_values(rows.iter().map(|r| r.1)) + .with_timezone("UTC"), + ), + completed_column, + ]; + for c in 0..payload_cols { + if c % 2 == 0 { + columns.push(Arc::new(Int64Array::from_iter_values( + (0..num_rows).map(|i| (i * c) as i64), + ))); + } else { + columns.push(Arc::new(StringArray::from_iter_values( + (0..num_rows).map(|i| format!("payload_{c}_{i:08}")), + ))); + } + } + let batch = RecordBatch::try_new(Arc::clone(&schema), columns).unwrap(); + (schema, split_batch(&batch, BATCH_SIZE)) +} + +/// Buffered side: `key`, `t` spread over `group_days` days and `qty`. +fn build_buffered( + keys: usize, + group_rows: usize, + group_days: i64, + batch_size: usize, +) -> (SchemaRef, Vec) { + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Utf8, false), + Field::new("t", timestamp_type(), true), + Field::new("qty", DataType::Int64, true), + ])); + let num_rows = keys * group_rows; + let span = group_days * MICROS_PER_DAY; + let key = StringArray::from_iter_values((0..num_rows).map(|i| key_of(i / group_rows))); + let t = TimestampMicrosecondArray::from_iter_values( + (0..num_rows).map(|i| BASE_MICROS + (i % group_rows) as i64 * span / group_rows as i64), + ) + .with_timezone("UTC"); + let qty = Int64Array::from_iter_values((0..num_rows).map(|i| i as i64)); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(key), Arc::new(t), Arc::new(qty)], + ) + .unwrap(); + (schema, split_batch(&batch, batch_size)) +} + +fn split_batch(batch: &RecordBatch, batch_size: usize) -> Vec { + (0..batch.num_rows()) + .step_by(batch_size) + .map(|offset| batch.slice(offset, (batch.num_rows() - offset).min(batch_size))) + .collect() +} + +/// `cast(completed as timestamp) < b.t AND o.t > b.t` +fn streamed_cast_filter(streamed: &Schema, buffered: &Schema) -> JoinFilter { + let completed = Arc::new(Cast::new( + Arc::new(Column::new("completed", 0)), + timestamp_type(), + SparkCastOptions::new(EvalMode::Legacy, "UTC", false), + None, + None, + )); + let buffered_t: Arc = Arc::new(Column::new("t", 2)); + let expression = Arc::new(BinaryExpr::new( + Arc::new(BinaryExpr::new( + completed, + Operator::Lt, + Arc::clone(&buffered_t), + )), + Operator::And, + Arc::new(BinaryExpr::new( + Arc::new(Column::new("t", 1)), + Operator::Gt, + buffered_t, + )), + )); + JoinFilter::new( + expression, + vec![ + ColumnIndex { + index: 2, + side: JoinSide::Left, + }, + ColumnIndex { + index: 1, + side: JoinSide::Left, + }, + ColumnIndex { + index: 1, + side: JoinSide::Right, + }, + ], + Arc::new(Schema::new(vec![ + streamed.field(2).clone(), + streamed.field(1).clone(), + buffered.field(1).clone(), + ])), + ) +} + +fn make_exec(batches: &[RecordBatch], schema: &SchemaRef) -> Arc { + MemorySourceConfig::try_new_exec(&[batches.to_vec()], Arc::clone(schema), None).unwrap() +} + +fn bench_smj_streamed_filter(c: &mut Criterion) { + let rt = Runtime::new().unwrap(); + let mut group = c.benchmark_group("smj_streamed_filter"); + group.sample_size(10); + for (shape, keys, rows_per_key, group_rows, all_pass, completed) in [ + ("few_string", 4, 1000, 2000, false, Completed::String), + ("few_date", 4, 1000, 2000, false, Completed::Date), + ("all_string", 4, 200, 2000, true, Completed::String), + ] { + let group_days = 400; + let window_days = 4; + let (streamed_schema, streamed_batches) = build_streamed( + keys, + rows_per_key, + group_days, + window_days, + all_pass, + completed, + 16, + ); + let (buffered_schema, buffered_batches) = + build_buffered(keys, group_rows, group_days, BATCH_SIZE); + let pairs = keys * rows_per_key * group_rows; + for join_type in [JoinType::Left, JoinType::Full, JoinType::Inner] { + group.bench_function( + BenchmarkId::new(format!("{shape}_{join_type:?}"), pairs), + |b| { + b.iter(|| { + let left = make_exec(&streamed_batches, &streamed_schema); + let right = make_exec(&buffered_batches, &buffered_schema); + let on = vec![( + Arc::new(Column::new("key", 0)) as Arc, + Arc::new(Column::new("key", 0)) as Arc, + )]; + let join = SortMergeJoinExec::try_new( + left, + right, + on, + Some(streamed_cast_filter(&streamed_schema, &buffered_schema)), + join_type, + vec![SortOptions::default()], + NullEquality::NullEqualsNothing, + ) + .unwrap(); + let task_ctx = + Arc::new(TaskContext::default().with_session_config( + SessionConfig::new().with_batch_size(BATCH_SIZE), + )); + let rows: usize = rt.block_on(async { + let batches = collect(Arc::new(join), task_ctx).await.unwrap(); + batches.iter().map(|b| b.num_rows()).sum() + }); + if all_pass { + assert_eq!(rows, pairs); + } else { + assert!(rows > keys * rows_per_key && rows < pairs / 20); + } + rows + }) + }, + ); + } + } + group.finish(); +} + +criterion_group!(benches, bench_smj_streamed_filter); +criterion_main!(benches); From 4bc36982c6a55de1afcdcc467aad4747ecd5d523 Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 09:20:17 +0100 Subject: [PATCH 08/13] perf: evaluate streamed-only join filter subexpressions once per streamed row The fbj_order_type join `cast(l.completed_dt as timestamp) < r.odt AND l.odt > r.odt` spent about 21 hours in one sort-merge join: completed_dt is a string of the streamed side, so HoistedJoinFilter, which lifted only buffered-only subexpressions, left the string to timestamp cast to run once per candidate pair, over key groups of millions of pairs. HoistedJoinFilter now lifts the largest non-volatile subexpressions that read the columns of either input alone. Streamed ones are evaluated per freeze once per run of a streamed row's pairs (over that row only, so no row the join never pairs is evaluated) and spread to the pairs with a take; when the runs average under two pairs they are evaluated per pair. A filter with streamed-only subexpressions and no buffered-only ones always takes the hoisted path; one with both still falls back to the original filter when the buffered results cannot be cached. Deferred filtering also skips the corrected mask when no pair of the accumulated batch failed the filter: every row is kept as is. smj_streamed_filter bench, ns per candidate pair, f5eab1948 -> this: few_string (1% pass, 8M pairs): LEFT 165 -> 5.2, FULL 164 -> 5.5, INNER 164 -> 4.2 few_date (1% pass, 8M pairs): LEFT 34 -> 4.9, FULL 35 -> 5.3, INNER 33 -> 4.2 all_string (all pass, 1.6M): LEFT 206 -> 46, FULL 206 -> 48, INNER 182 -> 29 smj_interval_filter (buffered-side casts) unchanged within noise: group11k_b8192 LEFT 5.0 -> 5.2, FULL 7.9 -> 8.0, INNER 4.0 -> 4.0; group11k_b1024 LEFT 5.5 -> 5.5, FULL 9.4 -> 9.3, INNER 4.5 -> 4.5. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../src/joins/sort_merge_join/filter.rs | 128 +++++++++++++----- .../sort_merge_join/materializing_stream.rs | 79 ++++++++++- .../src/joins/sort_merge_join/tests.rs | 19 ++- 3 files changed, 185 insertions(+), 41 deletions(-) diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs index 4d6398a6a6..fa86b9f2f0 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/filter.rs @@ -394,9 +394,9 @@ pub fn filter_record_batch_by_join_type( } } -/// A join filter whose largest subexpressions over buffered columns alone -/// are lifted out, so they can be evaluated once per buffered row instead of -/// once per pair. +/// A join filter whose largest subexpressions over the columns of one input +/// alone are lifted out, so they can be evaluated once per buffered or +/// streamed row instead of once per pair. #[derive(Debug)] pub(super) struct HoistedJoinFilter { /// The filter over `schema`, reading each lifted subexpression's result @@ -406,8 +406,14 @@ pub(super) struct HoistedJoinFilter { pub schema: SchemaRef, /// Where each column of `schema` comes from pub inputs: Vec, - /// The lifted subexpressions, over the buffered input's columns + /// The lifted subexpressions over buffered columns, over the buffered + /// input's columns pub buffered_exprs: Vec, + /// The lifted subexpressions over streamed columns, over the streamed + /// input's columns in `streamed_projection` + pub streamed_exprs: Vec, + /// The streamed input's columns `streamed_exprs` read + pub streamed_projection: Vec, } #[derive(Debug)] @@ -416,43 +422,52 @@ pub(super) enum HoistedFilterInput { Column(ColumnIndex), /// The result of `buffered_exprs[i]` Buffered(usize), + /// The result of `streamed_exprs[i]` + Streamed(usize), } impl HoistedJoinFilter { - /// Lifts the non-volatile subexpressions of `filter` that read buffered - /// columns and no streamed ones. Returns `None` when there is none - /// besides bare columns. + /// Lifts the non-volatile subexpressions of `filter` that read the + /// columns of one input and none of the other. Returns `None` when there + /// is none besides bare columns. pub fn try_new( filter: &JoinFilter, buffered_side: JoinSide, + streamed_schema: &Schema, buffered_schema: &Schema, ) -> Result> { let column_indices = filter.column_indices(); let num_filter_columns = column_indices.len(); - let mut lifted: Vec = vec![]; + let mut lifted: Vec<(JoinSide, PhysicalExprRef)> = vec![]; let expression = Arc::clone(filter.expression()) .transform_down(|expr| { if expr.downcast_ref::().is_some() || is_volatile(&expr) { return Ok(Transformed::no(expr)); } let columns = collect_columns(&expr); - if columns.is_empty() - || columns - .iter() - .any(|c| column_indices[c.index()].side != buffered_side) + let Some(side) = columns + .iter() + .next() + .map(|c| column_indices[c.index()].side) + else { + return Ok(Transformed::no(expr)); + }; + if columns + .iter() + .any(|c| column_indices[c.index()].side != side) { return Ok(Transformed::no(expr)); } - let position = match lifted.iter().position(|e| **e == *expr) { + let position = match lifted.iter().position(|(_, e)| **e == *expr) { Some(position) => position, None => { - lifted.push(Arc::clone(&expr)); + lifted.push((side, Arc::clone(&expr))); lifted.len() - 1 } }; Ok(Transformed::new( Arc::new(Column::new( - &format!("__buffered_{position}"), + &lifted_name(&lifted, position, buffered_side), num_filter_columns + position, )), true, @@ -479,13 +494,21 @@ impl HoistedJoinFilter { .iter() .map(|&index| HoistedFilterInput::Column(column_indices[index].clone())) .collect(); - for (position, expr) in lifted.iter().enumerate() { + let mut buffered_exprs = vec![]; + let mut streamed_exprs = vec![]; + for (position, (side, expr)) in lifted.iter().enumerate() { fields.push(Field::new( - format!("__buffered_{position}"), + lifted_name(&lifted, position, buffered_side), expr.data_type(filter.schema())?, expr.nullable(filter.schema())?, )); - inputs.push(HoistedFilterInput::Buffered(position)); + if *side == buffered_side { + inputs.push(HoistedFilterInput::Buffered(buffered_exprs.len())); + buffered_exprs.push(Arc::clone(expr)); + } else { + inputs.push(HoistedFilterInput::Streamed(streamed_exprs.len())); + streamed_exprs.push(Arc::clone(expr)); + } } let expression = expression @@ -502,28 +525,67 @@ impl HoistedJoinFilter { )) })? .data; - let buffered_exprs = lifted - .into_iter() - .map(|expr| { - expr.transform_up(|expr| { - let Some(column) = expr.downcast_ref::() else { - return Ok(Transformed::no(expr)); - }; - let index = column_indices[column.index()].index; - Ok(Transformed::yes(Arc::new(Column::new( - buffered_schema.field(index).name(), - index, - )) as PhysicalExprRef)) + let rebind = |exprs: Vec, + schema: &Schema, + projection: Option<&[usize]>| { + exprs + .into_iter() + .map(|expr| { + expr.transform_up(|expr| { + let Some(column) = expr.downcast_ref::() else { + return Ok(Transformed::no(expr)); + }; + let index = column_indices[column.index()].index; + let position = match projection { + Some(projection) => projection.binary_search(&index).unwrap(), + None => index, + }; + Ok(Transformed::yes(Arc::new(Column::new( + schema.field(index).name(), + position, + )) + as PhysicalExprRef)) + }) + .map(|transformed| transformed.data) }) - .map(|transformed| transformed.data) - }) - .collect::>>()?; + .collect::>>() + }; + let mut streamed_projection: Vec = streamed_exprs + .iter() + .flat_map(collect_columns) + .map(|c| column_indices[c.index()].index) + .collect(); + streamed_projection.sort_unstable(); + streamed_projection.dedup(); + let buffered_exprs = rebind(buffered_exprs, buffered_schema, None)?; + let streamed_exprs = + rebind(streamed_exprs, streamed_schema, Some(&streamed_projection))?; Ok(Some(Self { expression, schema: Arc::new(Schema::new(fields)), inputs, buffered_exprs, + streamed_exprs, + streamed_projection, })) } } + +/// Name of the column holding the result of `lifted[position]` +fn lifted_name( + lifted: &[(JoinSide, PhysicalExprRef)], + position: usize, + buffered_side: JoinSide, +) -> String { + let side = lifted[position].0; + let side_position = lifted[..position] + .iter() + .filter(|(s, _)| *s == side) + .count(); + if side == buffered_side { + format!("__buffered_{side_position}") + } else { + format!("__streamed_{side_position}") + } +} diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs index d436b2d02a..ebd976fba4 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/materializing_stream.rs @@ -620,7 +620,12 @@ impl MaterializingSortMergeJoinStream { let hoisted_filter = filter .as_ref() .map(|filter| { - HoistedJoinFilter::try_new(filter, buffered_side, &buffered_schema) + HoistedJoinFilter::try_new( + filter, + buffered_side, + &streamed_schema, + &buffered_schema, + ) }) .transpose()? .flatten(); @@ -1631,7 +1636,12 @@ impl MaterializingSortMergeJoinStream { as_uint64_array(&compute::concat(&refs)?)?.clone() }; - let evaluation = if self.cache_hoisted_filter(matched_chunks)? { + let hoisted = self + .hoisted_filter + .as_ref() + .is_some_and(|hoisted| hoisted.buffered_exprs.is_empty()) + || self.cache_hoisted_filter(matched_chunks)?; + let evaluation = if hoisted { Some(self.evaluate_hoisted_join_filter( &combined_left_indices, matched_chunks, @@ -1935,7 +1945,8 @@ impl MaterializingSortMergeJoinStream { } /// Evaluates the hoisted join filter over the given pairs, whose - /// buffered rows [`Self::cache_hoisted_filter`] cached. + /// buffered rows [`Self::cache_hoisted_filter`] cached when the filter + /// lifts buffered subexpressions. fn evaluate_hoisted_join_filter( &self, left_indices: &UInt64Array, @@ -1987,6 +1998,11 @@ impl MaterializingSortMergeJoinStream { }; let cached_columns = self .gather_cached_filter_columns(matched_chunks, hoisted.buffered_exprs.len())?; + let streamed_lifted_columns = evaluate_streamed_filter_exprs( + hoisted, + &self.streamed_batch.batch, + left_indices, + )?; let filter_columns = hoisted .inputs @@ -2004,6 +2020,9 @@ impl MaterializingSortMergeJoinStream { HoistedFilterInput::Buffered(position) => { Arc::clone(&cached_columns[*position]) } + HoistedFilterInput::Streamed(position) => { + Arc::clone(&streamed_lifted_columns[*position]) + } }) .collect::>(); let filter_batch = RecordBatch::try_new_with_options( @@ -2033,6 +2052,9 @@ impl MaterializingSortMergeJoinStream { matched_chunks: &[(usize, UInt64Array, UInt64Array)], num_columns: usize, ) -> Result> { + if num_columns == 0 { + return Ok(vec![]); + } let mut pieces: Vec> = vec![Vec::with_capacity(matched_chunks.len()); num_columns]; for (batch_idx, _, right) in matched_chunks { @@ -2389,6 +2411,13 @@ impl MaterializingSortMergeJoinStream { return Ok(record_batch); } + // No pair failed the filter: the corrected mask keeps every row. + if !out_mask.has_false() { + self.joined_record_batches + .clear(&self.schema, self.batch_size); + return Ok(record_batch); + } + // Validate inputs to get_corrected_filter_mask debug_assert_eq!( out_indices.len(), @@ -2450,6 +2479,50 @@ fn materialize_left_columns( } } +/// Evaluates the lifted streamed subexpressions of `hoisted` for the +/// streamed rows `indices` of `batch`: once per run of equal indices, as a +/// streamed row's pairs form, unless the runs are short. +fn evaluate_streamed_filter_exprs( + hoisted: &HoistedJoinFilter, + batch: &RecordBatch, + indices: &UInt64Array, +) -> Result> { + if hoisted.streamed_exprs.is_empty() { + return Ok(vec![]); + } + let evaluate = |rows: &UInt64Array| { + let projected = batch.project(&hoisted.streamed_projection)?; + let rows = RecordBatch::try_new_with_options( + projected.schema(), + materialize_left_columns(&projected, rows)?, + &RecordBatchOptions::new().with_row_count(Some(rows.len())), + )?; + hoisted + .streamed_exprs + .iter() + .map(|expr| expr.evaluate(&rows)?.into_array(rows.num_rows())) + .collect::>>() + }; + let values = indices.values(); + let num_runs = values.windows(2).filter(|w| w[0] != w[1]).count() + 1; + if num_runs * 2 > values.len() { + return evaluate(indices); + } + let mut run_rows = Vec::with_capacity(num_runs); + let mut positions = Vec::with_capacity(values.len()); + for (i, &row) in values.iter().enumerate() { + if i == 0 || row != values[i - 1] { + run_rows.push(row); + } + positions.push((run_rows.len() - 1) as u32); + } + let positions = UInt32Array::from(positions); + evaluate(&UInt64Array::from(run_rows))? + .iter() + .map(|column| Ok(compute::take(column, &positions, None)?)) + .collect() +} + /// Average length of the runs of consecutive buffered rows from which /// copying the runs beats gathering row by row. const MIN_AVG_RUN_LEN: usize = 16; diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs index 863367b71e..618355eb28 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs @@ -4818,13 +4818,22 @@ fn hoisted_join_filter_lifts_buffered_subexpressions() -> Result<()> { for streamed_side in [JoinSide::Left, JoinSide::Right] { let plain = build_interval_filter(&streamed, &buffered, streamed_side, false); assert!( - HoistedJoinFilter::try_new(&plain, streamed_side.negate(), &buffered)? - .is_none() + HoistedJoinFilter::try_new( + &plain, + streamed_side.negate(), + &streamed, + &buffered + )? + .is_none() ); let filter = build_interval_filter(&streamed, &buffered, streamed_side, true); - let hoisted = - HoistedJoinFilter::try_new(&filter, streamed_side.negate(), &buffered)? - .unwrap(); + let hoisted = HoistedJoinFilter::try_new( + &filter, + streamed_side.negate(), + &streamed, + &buffered, + )? + .unwrap(); let lifted: Vec = hoisted .buffered_exprs .iter() From d531476ccf68a82cee3e61e76a9fecf8f53d5974 Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 09:40:46 +0100 Subject: [PATCH 09/13] test: sort-merge join filters with streamed-only and both-side subexpressions The validity-interval Rust tests run each case with four equivalent forms of the filter: plain, with buffered-only subexpressions, with streamed-only ones (`t + 0 > eff AND t - 0 <= next_eff AND t + 0 >= t - 0`) and with both, over NULL lookup times, outer and FULL joins, spilling, and key groups across batches on both sides. The hoisting test checks what each form lifts and how the residual filter reads it. CometSmjJoinFilterFuzzSuite gets filters casting a streamed column (string to bigint, and the fbj_order_type shape: a timestamp rendered as a string cast back to timestamp, compared with a buffered value) and casting both sides; they run in the default matrix, at batch 7 and in the large-group sets. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../src/joins/sort_merge_join/tests.rs | 177 +++++++++++++----- .../exec/CometSmjJoinFilterFuzzSuite.scala | 33 +++- 2 files changed, 154 insertions(+), 56 deletions(-) diff --git a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs index 618355eb28..6f6dd75b0d 100644 --- a/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs +++ b/native/vendor/datafusion-physical-plan/src/joins/sort_merge_join/tests.rs @@ -4388,15 +4388,39 @@ fn build_interval_buffered( TestMemoryExec::try_new_exec(&[batches], schema, None).unwrap() } -/// `t > eff AND t <= next_eff`, with the streamed table on `streamed_side`. -/// When `hoisted`, `t > eff + 0 AND t <= next_eff + 0 AND next_eff - eff > 0` -/// instead: the same filter with subexpressions over buffered columns alone, -/// which the join evaluates once per buffered row. +/// Forms of the validity-interval filter `t > eff AND t <= next_eff`, all +/// with the same result. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum IntervalFilter { + /// `t > eff AND t <= next_eff` + Plain, + /// `t > eff + 0 AND t <= next_eff + 0 AND next_eff - eff > 0`: with + /// subexpressions over buffered columns alone, which the join evaluates + /// once per buffered row + Buffered, + /// `t + 0 > eff AND t - 0 <= next_eff AND t + 0 >= t - 0`: with + /// subexpressions over streamed columns alone, which the join evaluates + /// once per streamed row + Streamed, + /// `t + 0 > eff + 0 AND t - 0 <= next_eff + 0 AND next_eff - eff > 0 + /// AND t + 0 >= t - 0`: with both + Both, +} + +const INTERVAL_FILTERS: [IntervalFilter; 4] = [ + IntervalFilter::Plain, + IntervalFilter::Buffered, + IntervalFilter::Streamed, + IntervalFilter::Both, +]; + +/// The validity-interval filter in the given form, with the streamed table +/// on `streamed_side`. fn build_interval_filter( streamed: &Schema, buffered: &Schema, streamed_side: JoinSide, - hoisted: bool, + form: IntervalFilter, ) -> JoinFilter { let t_field = streamed .field_with_name("t") @@ -4446,24 +4470,31 @@ fn build_interval_filter( let t_col: PhysicalExprRef = Arc::new(Column::new("t", t_idx)); let eff: PhysicalExprRef = Arc::new(Column::new("eff", eff_idx)); let next_eff: PhysicalExprRef = Arc::new(Column::new("next_eff", next_eff_idx)); - let plus_zero = |expr: &PhysicalExprRef| -> PhysicalExprRef { + let zero = |expr: &PhysicalExprRef, op: Operator| -> PhysicalExprRef { Arc::new(BinaryExpr::new( Arc::clone(expr), - Operator::Plus, + op, Arc::new(Literal::new(ScalarValue::Int64(Some(0)))), )) }; - let (lower, upper) = if hoisted { - (plus_zero(&eff), plus_zero(&next_eff)) + let lift_buffered = matches!(form, IntervalFilter::Buffered | IntervalFilter::Both); + let lift_streamed = matches!(form, IntervalFilter::Streamed | IntervalFilter::Both); + let (lower, upper) = if lift_buffered { + (zero(&eff, Operator::Plus), zero(&next_eff, Operator::Plus)) } else { (Arc::clone(&eff), Arc::clone(&next_eff)) }; + let (t_lower, t_upper) = if lift_streamed { + (zero(&t_col, Operator::Plus), zero(&t_col, Operator::Minus)) + } else { + (Arc::clone(&t_col), Arc::clone(&t_col)) + }; let mut expression: PhysicalExprRef = Arc::new(BinaryExpr::new( - Arc::new(BinaryExpr::new(Arc::clone(&t_col), Operator::Gt, lower)), + Arc::new(BinaryExpr::new(Arc::clone(&t_lower), Operator::Gt, lower)), Operator::And, - Arc::new(BinaryExpr::new(t_col, Operator::LtEq, upper)), + Arc::new(BinaryExpr::new(Arc::clone(&t_upper), Operator::LtEq, upper)), )); - if hoisted { + if lift_buffered { expression = Arc::new(BinaryExpr::new( expression, Operator::And, @@ -4474,6 +4505,13 @@ fn build_interval_filter( )), )); } + if lift_streamed { + expression = Arc::new(BinaryExpr::new( + expression, + Operator::And, + Arc::new(BinaryExpr::new(t_lower, Operator::GtEq, t_upper)), + )); + } JoinFilter::new(expression, column_indices, Arc::new(Schema::new(fields))) } @@ -4584,7 +4622,7 @@ async fn check_interval_join( spill: bool, case: &str, ) -> Result<()> { - for hoisted in [false, true] { + for form in INTERVAL_FILTERS { check_interval_join_with_filter( join_type, streamed, @@ -4593,16 +4631,15 @@ async fn check_interval_join( buffered_batch_rows, batch_size, spill, - hoisted, - &format!("{case} hoisted={hoisted}"), + form, + &format!("{case} filter={form:?}"), ) .await?; } Ok(()) } -/// [`check_interval_join`] with the plain or the hoisted form of the filter -/// (see [`build_interval_filter`]). +/// [`check_interval_join`] with one form of the filter. #[expect(clippy::too_many_arguments)] async fn check_interval_join_with_filter( join_type: JoinType, @@ -4612,7 +4649,7 @@ async fn check_interval_join_with_filter( buffered_batch_rows: usize, batch_size: usize, spill: bool, - hoisted: bool, + form: IntervalFilter, case: &str, ) -> Result<()> { let streamed_plan = build_interval_streamed(streamed, streamed_batch_rows); @@ -4626,7 +4663,7 @@ async fn check_interval_join_with_filter( &streamed_plan.schema(), &buffered_plan.schema(), streamed_side, - hoisted, + form, ); let (left, right) = if join_type == Right { (buffered_plan, streamed_plan) @@ -4807,55 +4844,93 @@ async fn join_filter_streamed_rows_across_freezes() -> Result<()> { Ok(()) } -/// The hoisted interval filter lifts `eff + 0`, `next_eff + 0` and the whole -/// `next_eff - eff > 0`, keeping `t` as the only input column; the plain one -/// has nothing to lift. +/// Each form of the interval filter lifts its subexpressions over one input +/// alone, and the bare columns it still reads stay input columns; the plain +/// form has nothing to lift. #[test] -fn hoisted_join_filter_lifts_buffered_subexpressions() -> Result<()> { +fn hoisted_join_filter_lifts_single_side_subexpressions() -> Result<()> { use super::filter::{HoistedFilterInput, HoistedJoinFilter}; + use HoistedFilterInput::{Buffered, Column as Input, Streamed}; let streamed = build_interval_streamed(&[(1, 0, Some(1))], 1).schema(); let buffered = build_interval_buffered(&interval_group(1, 0, 1, 1), 1).schema(); for streamed_side in [JoinSide::Left, JoinSide::Right] { - let plain = build_interval_filter(&streamed, &buffered, streamed_side, false); - assert!( - HoistedJoinFilter::try_new( - &plain, - streamed_side.negate(), - &streamed, - &buffered - )? - .is_none() - ); - let filter = build_interval_filter(&streamed, &buffered, streamed_side, true); - let hoisted = HoistedJoinFilter::try_new( - &filter, - streamed_side.negate(), - &streamed, - &buffered, - )? - .unwrap(); - let lifted: Vec = hoisted - .buffered_exprs - .iter() - .map(|e| e.to_string()) - .collect(); + let buffered_side = streamed_side.negate(); + let hoist = |form| { + let filter = build_interval_filter(&streamed, &buffered, streamed_side, form); + HoistedJoinFilter::try_new(&filter, buffered_side, &streamed, &buffered) + }; + let to_strings = |exprs: &[PhysicalExprRef]| -> Vec { + exprs.iter().map(|e| e.to_string()).collect() + }; + assert!(hoist(IntervalFilter::Plain)?.is_none()); + + let hoisted = hoist(IntervalFilter::Buffered)?.unwrap(); assert_eq!( - lifted, + to_strings(&hoisted.buffered_exprs), vec!["eff@2 + 0", "next_eff@3 + 0", "next_eff@3 - eff@2 > 0"] ); + assert!(hoisted.streamed_exprs.is_empty()); assert!(matches!( hoisted.inputs.as_slice(), [ - HoistedFilterInput::Column(ColumnIndex { index: 2, side }), - HoistedFilterInput::Buffered(0), - HoistedFilterInput::Buffered(1), - HoistedFilterInput::Buffered(2), + Input(ColumnIndex { index: 2, side }), + Buffered(0), + Buffered(1), + Buffered(2), ] if *side == streamed_side )); assert_eq!( hoisted.expression.to_string(), "t@0 > __buffered_0@1 AND t@0 <= __buffered_1@2 AND __buffered_2@3" ); + + let hoisted = hoist(IntervalFilter::Streamed)?.unwrap(); + assert!(hoisted.buffered_exprs.is_empty()); + assert_eq!( + to_strings(&hoisted.streamed_exprs), + vec!["t@0 + 0", "t@0 - 0", "t@0 + 0 >= t@0 - 0"] + ); + assert_eq!(hoisted.streamed_projection, vec![2]); + assert!(matches!( + hoisted.inputs.as_slice(), + [ + Input(ColumnIndex { index: 2, side: eff_side }), + Input(ColumnIndex { index: 3, side: next_eff_side }), + Streamed(0), + Streamed(1), + Streamed(2), + ] if *eff_side == buffered_side && *next_eff_side == buffered_side + )); + assert_eq!( + hoisted.expression.to_string(), + "__streamed_0@2 > eff@0 AND __streamed_1@3 <= next_eff@1 AND __streamed_2@4" + ); + + let hoisted = hoist(IntervalFilter::Both)?.unwrap(); + assert_eq!( + to_strings(&hoisted.buffered_exprs), + vec!["eff@2 + 0", "next_eff@3 + 0", "next_eff@3 - eff@2 > 0"] + ); + assert_eq!( + to_strings(&hoisted.streamed_exprs), + vec!["t@0 + 0", "t@0 - 0", "t@0 + 0 >= t@0 - 0"] + ); + assert!(matches!( + hoisted.inputs.as_slice(), + [ + Streamed(0), + Buffered(0), + Streamed(1), + Buffered(1), + Buffered(2), + Streamed(2), + ] + )); + assert_eq!( + hoisted.expression.to_string(), + "__streamed_0@0 > __buffered_0@1 AND __streamed_1@2 <= __buffered_1@3 \ + AND __buffered_2@4 AND __streamed_2@5" + ); } Ok(()) } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala index 3dfec606c0..29068c0af5 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala @@ -45,8 +45,9 @@ import org.apache.comet.CometConf * * Every join type runs with filters of different shapes and selectivities (a validity interval * bounded by one side, a band bounded by the other, one-sided conditions, always true and always - * false, almost nothing or almost everything passing) at several batch sizes and under a memory - * pool small enough to make the join spill. Results are compared with Spark as multisets. + * false, almost nothing or almost everything passing, casts of one or both sides) at several + * batch sizes and under a memory pool small enough to make the join spill. Results are compared + * with Spark as multisets. * * The default run covers each join type and filter at one batch size and a subset at the others. * The full matrix runs only with `-Dcomet.test.smjFuzz.full=true`; @@ -305,8 +306,30 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe private val most = Filter("most", "(l.id + r.id) % 10 <> 0") private val nullable = Filter("nullable_cmp", "l.lv < r.rv") private val typed = Filter("typed", "l.d10 * 2 >= r.rdec OR l.s < r.rs") + private val leftCast = + Filter("left_cast", "CAST(CAST(l.t AS STRING) AS BIGINT) > r.eff AND l.t <= r.next_eff") + private val leftTimestampCast = Filter( + "left_ts_cast", + "CAST(CAST(l.ts AS STRING) AS TIMESTAMP) < CAST(r.eff * 2000 AS TIMESTAMP) AND l.t > r.eff") + private val bothCast = Filter( + "both_cast", + "CAST(CAST(l.t AS STRING) AS BIGINT) > CAST(CAST(r.eff AS STRING) AS BIGINT) AND " + + "CAST(l.t AS DECIMAL(20, 0)) <= CAST(r.next_eff AS DECIMAL(20, 0))") private val filters = - Seq(interval, band, leftOnly, rightOnly, alwaysFalse, alwaysTrue, rare, most, nullable, typed) + Seq( + interval, + band, + leftOnly, + rightOnly, + alwaysFalse, + alwaysTrue, + rare, + most, + nullable, + typed, + leftCast, + leftTimestampCast, + bothCast) private case class Shape(name: String, confs: Seq[(String, String)], spill: Boolean = false) @@ -538,7 +561,7 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe test("batch 7: key groups spanning many batches") { val ds = dataSet("main") - runAll(ds, b7, matrix(ds, joinKinds, Seq(interval, band, most))) + runAll(ds, b7, matrix(ds, joinKinds, Seq(interval, band, most, leftCast, bothCast))) } test("batch 1024: key groups around and over one batch") { @@ -591,7 +614,7 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe test(s"full: groups of thousands of rows, ${shape.name}") { assumeFull() val ds = dataSet("big") - val fs = Seq(interval, band, rare, alwaysFalse, nullable) + val fs = Seq(interval, band, rare, alwaysFalse, nullable, leftCast, bothCast) runAll(ds, shape, matrix(ds, joinKinds, fs) ++ matrix(ds, joinKinds, fs, flipped = true)) } } From 4575593c79f26a84280f9fb27b723621fd9410c6 Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 11:19:47 +0100 Subject: [PATCH 10/13] test: fuzz sort-merge join filters with boolean and conditional operators around lifted subexpressions Adds string, int, boolean and date columns to both sides (generated from separate random streams, so existing data is unchanged) with valid, NULL and invalid values, and 26 filter shapes that wrap single-side subexpressions in OR, CASE, IF, NOT, IS [NOT] NULL, IN / NOT IN, COALESCE and null-safe equality. A few run by default; the rest run in the full matrix over every join type, side order, batch size and spill mode. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../exec/CometSmjJoinFilterFuzzSuite.scala | 213 +++++++++++++++++- 1 file changed, 207 insertions(+), 6 deletions(-) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala index 29068c0af5..220757b080 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala @@ -22,7 +22,8 @@ package org.apache.comet.exec import java.io.File import java.nio.file.Files import java.sql.{Date, Timestamp} -import java.time.{Instant, LocalDate} +import java.time.{Instant, LocalDate, LocalDateTime, ZoneOffset} +import java.time.format.DateTimeFormatter import scala.collection.mutable import scala.util.Random @@ -84,7 +85,12 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe StructField("b", StringType), StructField("c", ArrayType(IntegerType))))), StructField("arr", ArrayType(StringType)), - StructField("m", MapType(StringType, IntegerType)))) + StructField("m", MapType(StringType, IntegerType)), + StructField("ss", StringType), + StructField("li", StringType), + StructField("k2", IntegerType), + StructField("flag", BooleanType), + StructField("ldt", DateType))) private val rightSchema = StructType( Seq( @@ -99,7 +105,11 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe StructField("rarr", ArrayType(IntegerType)), StructField( "rst", - StructType(Seq(StructField("x", DoubleType), StructField("y", StringType)))))) + StructType(Seq(StructField("x", DoubleType), StructField("y", StringType)))), + StructField("rt", TimestampType), + StructField("rss", StringType), + StructField("rk2", IntegerType), + StructField("rdt", DateType))) private case class DataSet(name: String, left: String, right: String, stats: String) @@ -194,6 +204,60 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe orNull(r, 10)(Row(orNull(r, 20)(randomDouble(r)), orNull(r, 20)(randomString(r))))) } + private val tsFormat = DateTimeFormatter.ofPattern("yyyy-MM-dd HH:mm:ss") + + private def formatSeconds(sec: Long): String = + LocalDateTime.ofEpochSecond(sec, 0, ZoneOffset.UTC).format(tsFormat) + + // A string column cast to TIMESTAMP in the filters: mostly valid timestamps near the key + // group's values, but also NULL, empty, malformed, out of range, date-only, padded and + // year-only strings, which Legacy casts turn into NULL or partial values. + private def timestampString(r: Random, sec: Long): Any = r.nextInt(24) match { + case 0 | 1 | 2 => null + case 3 => "" + case 4 => "not a timestamp" + case 5 => "2026-13-45 99:00:00" + case 6 => formatSeconds(sec).replace(' ', 'T').dropRight(3) + ":75" + case 7 => " " + formatSeconds(sec) + " " + case 8 => formatSeconds(sec).take(10) + case 9 => formatSeconds(sec).take(4) + case 10 => String.valueOf(sec) + case 11 => "ΠΆβ‚¬πŸ˜€" + case _ => formatSeconds(sec) + } + + // A string column cast to INT in the filters: numbers, padded numbers, fractions, NULL and + // strings that Legacy casts turn into NULL. + private def intString(r: Random): Any = r.nextInt(16) match { + case 0 | 1 => null + case 2 => "" + case 3 => "x1" + case 4 => " " + r.nextInt(100) + " " + case 5 => s"${r.nextInt(100)}.${r.nextInt(10)}" + case 6 => "99999999999" + case _ => String.valueOf(r.nextInt(100)) + } + + private def leftExtra(r: Random, g: KeyGroup): Seq[Any] = { + val span = math.max(g.nr, 1) * 10L + 40 + val sec = g.base - 20 + r.nextInt(span.toInt) + Seq( + timestampString(r, sec), + intString(r), + orNull(r, 10)(r.nextInt(6)), + orNull(r, 15)(r.nextBoolean()), + orNull(r, 30)(Date.valueOf(LocalDate.ofEpochDay(sec / 86400 + r.nextInt(3) - 1)))) + } + + private def rightExtra(r: Random, g: KeyGroup, j: Int): Seq[Any] = { + val sec = g.base + j * 10L + r.nextInt(31) - 15 + Seq( + orNull(r, 10)(Timestamp.from(Instant.ofEpochSecond(sec))), + timestampString(r, sec + r.nextInt(21) - 10), + orNull(r, 10)(r.nextInt(6)), + orNull(r, 30)(Date.valueOf(LocalDate.ofEpochDay(sec / 86400 + r.nextInt(3) - 1)))) + } + private val mainCombos: Seq[(Int, Int)] = Seq( 1 -> 1, 1 -> 7, @@ -249,9 +313,12 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe val groups = sized.zipWithIndex.map { case ((nl, nr), i) => KeyGroup(i * 3 + r.nextInt(3), nl, nr, r.nextInt(100000) * 10L, r.nextInt(3)) } ++ Seq(KeyGroup(-1, 40 + r.nextInt(40), 40 + r.nextInt(40), 5000L, 0)) - val leftRows = r.shuffle(groups.flatMap(g => Seq.fill(g.nl)(leftRow(r, g)))) - val rightRows = - r.shuffle(groups.flatMap(g => (0 until g.nr).map(j => rightRow(r, g, j)))) + val lx = new Random(dataSeed ^ name.hashCode ^ 0x5eed1L) + val rx = new Random(dataSeed ^ name.hashCode ^ 0x5eed2L) + val leftRows = + r.shuffle(groups.flatMap(g => Seq.fill(g.nl)(leftRow(r, g) ++ leftExtra(lx, g)))) + val rightRows = r.shuffle( + groups.flatMap(g => (0 until g.nr).map(j => rightRow(r, g, j) ++ rightExtra(rx, g, j)))) if (tempRoot == null) tempRoot = Files.createTempDirectory("comet-smj-fuzz").toFile val leftPath = new File(tempRoot, s"$name-left").getCanonicalPath @@ -331,6 +398,98 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe leftTimestampCast, bothCast) + // Boolean and conditional operators around subexpressions that read one side only, which + // the native join lifts out of the filter and evaluates once per row: OR, CASE, IF, NOT, + // IS [NOT] NULL, IN and NOT IN, COALESCE and null-safe equality, with casts of strings that + // are often invalid, so the lifted parts are often NULL and three-valued logic matters. + private val lts = "CAST(l.ss AS TIMESTAMP)" + private val rts = "CAST(r.rss AS TIMESTAMP)" + private val lint = "CAST(l.li AS INT)" + private val orTs = Filter("or_ts", s"l.t > r.eff OR $lts < r.rt") + private val orTsNull = Filter("or_ts_null", s"$lts < r.rt OR l.lv IS NULL") + private val orMostlyTrue = + Filter("or_mostly_true", s"(l.id + r.id) % 10 <> 0 OR $lts >= r.rt") + private val orBufferedLifted = + Filter("or_buffered_lifted", s"l.lv < r.rv OR $rts > CAST(l.t AS TIMESTAMP)") + private val orSingleSides = + Filter("or_single_sides", s"(l.flag AND $lts IS NOT NULL) OR r.rv % 3 = 0") + private val orBothLifted = + Filter("or_both_lifted", s"$lts < $rts OR ($lts IS NULL AND $rts IS NULL)") + private val caseFlag = + Filter("case_flag", s"CASE WHEN l.flag THEN $lts < r.rt ELSE r.eff <= l.t END") + private val caseRightWhen = Filter( + "case_right_when", + s"CASE WHEN r.rv % 2 = 0 THEN $lts < TIMESTAMP '1970-01-06 00:00:00' " + + s"WHEN r.rv IS NULL THEN l.flag ELSE $lint > 50 END") + private val caseLeftWhen = Filter( + "case_left_when", + s"CASE WHEN l.flag THEN $rts > TIMESTAMP '1970-01-06 00:00:00' ELSE r.rk2 > 2 END") + private val caseNested = Filter( + "case_nested", + s"CASE WHEN l.k2 IS NULL THEN r.rv IS NULL WHEN l.k2 > 2 THEN " + + s"CASE WHEN $lts IS NULL THEN r.rk2 = l.k2 ELSE $lts < r.rt END " + + s"ELSE CASE WHEN r.rk2 IS NULL THEN l.flag ELSE $lint < r.rv END END") + private val caseValue = + Filter("case_value", s"CASE WHEN l.flag THEN $lts ELSE CAST(l.t AS TIMESTAMP) END < r.rt") + private val ifNull = + Filter("if_null", s"IF($lint IS NULL, r.rk2 > 2, $lint <= r.rv)") + private val notAndUpper = Filter("not_and_upper", "NOT (l.lv > r.rv AND upper(l.s) = r.rs)") + private val notOr = Filter("not_or", s"NOT ($lts >= r.rt OR l.flag)") + private val isNullMix = + Filter("is_null_mix", s"($lts IS NULL) = (r.rt IS NULL) OR ($lint + r.rv) IS NULL") + private val isNotNullCmp = + Filter("is_not_null_cmp", s"($lts < r.rt) IS NOT NULL AND NOT ($lts < r.rt)") + private val inList = + Filter("in_list", "l.k2 IN (1, 2, 3) AND l.t > r.eff OR r.rk2 IN (0, 4) AND l.lv < r.rv") + private val notInList = Filter( + "not_in_list", + s"l.k2 NOT IN (1, 3) AND r.rk2 NOT IN (2) OR $lint NOT IN (1, 2, 50) AND l.t <= r.next_eff") + private val inCross = + Filter("in_cross", s"r.rk2 IN (l.k2, l.k2 + 1, $lint) OR l.k2 IN (r.rk2, 3)") + private val coalesceTs = + Filter("coalesce_ts", s"coalesce($lts, r.rt) > CAST(r.eff AS TIMESTAMP)") + private val coalesceDate = + Filter("coalesce_date", "coalesce(l.ldt, r.rdt) >= CAST(r.rt AS DATE)") + private val threeValued = + Filter("three_valued", s"($lint > r.rv OR $lts < r.rt) AND NOT ($rts > $lts AND l.flag)") + private val notNullPart = + Filter("not_null_part", "NOT (CAST(l.ss AS INT) > 0) OR r.rv > 50") + private val nullSafeEq = + Filter("null_safe_eq", s"$lts <=> r.rt OR $lint <=> r.rk2") + private val liftedTwice = + Filter("lifted_twice", s"$lts >= r.rt AND $lts <= CAST(r.eff + 20 AS TIMESTAMP)") + private val upperInvalid = + Filter("upper_invalid", "upper(l.ss) = upper(r.rss) OR upper(l.s) < r.rs") + private val boolFilters = Seq( + orTs, + orTsNull, + orMostlyTrue, + orBufferedLifted, + orSingleSides, + orBothLifted, + caseFlag, + caseRightWhen, + caseLeftWhen, + caseNested, + caseValue, + ifNull, + notAndUpper, + notOr, + isNullMix, + isNotNullCmp, + inList, + notInList, + inCross, + coalesceTs, + coalesceDate, + threeValued, + notNullPart, + nullSafeEq, + liftedTwice, + upperInvalid) + private val defaultBoolFilters = + Seq(orTs, caseFlag, notAndUpper, inList, coalesceTs, threeValued) + private case class Shape(name: String, confs: Seq[(String, String)], spill: Boolean = false) private def batch(n: Int) = CometConf.COMET_BATCH_SIZE.key -> n.toString @@ -569,6 +728,17 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe runAll(ds, b1024, matrix(ds, joinKinds, Seq(interval, rare, alwaysTrue))) } + test("boolean and conditional operators around lifted subexpressions, batch 64") { + val ds = dataSet("main") + runAll(ds, b64, matrix(ds, joinKinds, defaultBoolFilters)) + } + + test("boolean and conditional operators, batch 7 and spilling") { + val ds = dataSet("main") + runAll(ds, b7, matrix(ds, joinKinds, Seq(caseFlag, threeValued))) + runAll(ds, spill64, matrix(ds, Seq(inner, fullOuter, leftAnti), Seq(orTs, coalesceTs))) + } + test("spilling join under a tiny memory pool") { val ds = dataSet("main") val outcomes = runAll(ds, spill64, matrix(ds, joinKinds, Seq(interval, most))) @@ -599,6 +769,37 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe } } + for (shape <- allShapes) { + test(s"full: ${shape.name}, boolean and conditional operators, both side orders") { + assumeFull() + val ds = dataSet("main") + val outcomes = runAll( + ds, + shape, + matrix(ds, joinKinds, boolFilters) ++ matrix(ds, joinKinds, boolFilters, flipped = true)) + if (shape.spill) { + assert(outcomes.exists(_.spills > 0), s"[seed=$seed] the join did not spill") + } + } + } + + for (shape <- Seq(b7, spill64)) { + test(s"full: second seed, boolean and conditional operators, ${shape.name}") { + assumeFull() + val ds = dataSet("alt") + runAll(ds, shape, matrix(ds, joinKinds, boolFilters)) + } + } + + for (shape <- Seq(b1024, spill64)) { + test(s"full: groups of thousands of rows, boolean and conditional operators, ${shape.name}") { + assumeFull() + val ds = dataSet("big") + val fs = Seq(orTs, caseFlag, caseNested, notOr, inCross, coalesceTs, threeValued) + runAll(ds, shape, matrix(ds, joinKinds, fs) ++ matrix(ds, joinKinds, fs, flipped = true)) + } + } + for (shape <- Seq(b7, b64, spill64)) { test(s"full: second seed, ${shape.name}") { assumeFull() From 60bbf76147566c270d911076386dca049cf605fd Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 11:25:27 +0100 Subject: [PATCH 11/13] feat: charge a sort-merge join condition holding a CASE or IF over both inputs On the cluster (fix4d), left joins with about one passing pair per streamed row under `CASE WHEN l.c = r.c THEN l.t - r.eff ELSE 10 - (r.nxt - l.t) END BETWEEN 1 AND 10` cost 28 to 51 ns per pair natively against 8.7 to 17.6 in Spark, 1.7 to 4.5 times. Cross-input datediff costs about what Spark does, cross-input instr and an OR over a cast of one input 4 to 10 times less, and conditions over one input at most Spark's price. A sort-merge join whose condition holds a CaseWhen or an If referencing columns of both inputs adds smjCrossCondition over every output leaf, at the former smjCondition line Line(3000, 80, 0, 850, 23); any other condition adds nothing. The rule stays disabled by default. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/source/user-guide/latest/tuning.md | 4 +- .../scala/org/apache/comet/CometConf.scala | 10 +-- .../comet/rules/CostBasedEngineChoice.scala | 18 ++++- .../apache/comet/rules/EngineCostTable.scala | 16 ++++- .../rules/CostBasedEngineChoiceSuite.scala | 67 +++++++++++++++++-- 5 files changed, 99 insertions(+), 16 deletions(-) diff --git a/docs/source/user-guide/latest/tuning.md b/docs/source/user-guide/latest/tuning.md index 4b6d929ee9..e29e45477c 100644 --- a/docs/source/user-guide/latest/tuning.md +++ b/docs/source/user-guide/latest/tuning.md @@ -602,7 +602,9 @@ prices per row from a table of measurements: `c0 + k0*L + k1*L*min(L, 600)` ns n Spark one. - `sort` over every leaf of the sorted rows (Comet also 0.15 ns per byte beyond 12 per leaf); `sortSpill` adds the price of a spill to a fraction `sortSpillFraction` of rows, none by default. `smj` prices a sort-merge join and `bhj` - the probe side of a broadcast hash join, over every output leaf; a join condition adds nothing. + the probe side of a broadcast hash join, over every output leaf. A sort-merge join condition holding a `CASE` or an + `IF` over columns of both inputs adds `smjCrossCondition`, also over every output leaf: natively such a condition + cost 1.7 to 4.5 times Spark's price per pair of rows of equal keys. Any other condition adds nothing. - `predicate` for filters, over the leaves their predicate references, plus a pass-through of 1.5 ns per output leaf natively and none in Spark. A native filter over a native scan, and the native projects over it, stay native whatever their prices (`keepFiltersOverNativeScans=false` lets them move): the rows a filter drops are not estimated, diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index acc6e57215..627be899e4 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -689,11 +689,11 @@ object CometConf extends ShimCometConf { "(or `k0,k1`, keeping c0), or spark, with `c0,k` for c0 + k*L ns per row, where L is " + "the number of leaf columns the class prices (for the functions of an aggregate or a " + "window, the number of functions). The classes are shuffleWrite, shuffleRead, sort, " + - "sortSpill, smj, bhj, predicate, projectPassThrough, expression, agg, aggObjectHash, " + - "aggDeclarative, aggCollectList, aggCollectSet, aggPercentile, aggPercentileApprox, " + - "aggOther, aggArrayKey, codegenDispatch, window, windowAggregate, windowOffset, " + - "windowRank, wglPartial, wglFinal, expand, generate, rowLocal, the comet-only c2r and " + - "r2c, and the spark-only " + + "sortSpill, smj, smjCrossCondition, bhj, predicate, projectPassThrough, expression, " + + "agg, aggObjectHash, aggDeclarative, aggCollectList, aggCollectSet, aggPercentile, " + + "aggPercentileApprox, aggOther, aggArrayKey, codegenDispatch, window, " + + "windowAggregate, windowOffset, windowRank, wglPartial, wglFinal, expand, generate, " + + "rowLocal, the comet-only c2r and r2c, and the spark-only " + "expressionOverScan, aggDeclarativeNoCodegen, expandNoCodegen and " + "generateNoCodegen. A row whose leaves are a fraction f inside structs, arrays or " + "maps costs (1 - f) times the flat price plus f times the nested one. The scalars " + diff --git a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala index 27aa5960a1..cef4dd8b2a 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala @@ -25,7 +25,7 @@ import scala.collection.mutable import org.apache.spark.internal.Logging import org.apache.spark.sql.SparkSession -import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, Expression, NamedExpression, ScalaUDF} +import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, CaseWhen, Expression, If, NamedExpression, ScalaUDF} import org.apache.spark.sql.catalyst.expressions.aggregate.{Complete, Partial} import org.apache.spark.sql.catalyst.plans.logical.statsEstimation.EstimationUtils import org.apache.spark.sql.catalyst.rules.Rule @@ -33,6 +33,7 @@ import org.apache.spark.sql.comet.{CometExec, CometFilterExec, CometHashAggregat import org.apache.spark.sql.execution.{ColumnarToRowTransition, ExpandExec, FilterExec, ProjectExec, SortExec, SparkPlan} import org.apache.spark.sql.execution.aggregate.BaseAggregateExec import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike +import org.apache.spark.sql.execution.joins.SortMergeJoinExec import org.apache.spark.sql.execution.window.WindowExec import org.apache.spark.sql.internal.SQLConf @@ -123,6 +124,15 @@ class EngineCostModel( private def aggregateShare(agg: BaseAggregateExec): Double = if (agg.aggregateExpressions.exists(_.mode == Complete)) 1.0 else 0.5 + /** Whether `condition` of `join` holds a CASE or an IF over columns of both its inputs. */ + def crossInputConditional(condition: Expression, join: SortMergeJoinExec): Boolean = + condition.exists { + case e @ (_: CaseWhen | _: If) => + e.references.exists(join.left.outputSet.contains) && + e.references.exists(join.right.outputSet.contains) + case _ => false + } + private def classTerms(costClass: CostClass, plan: SparkPlan, engine: Engine): Seq[Term] = { val op = sparkOperator(plan) lazy val out = widthOf(op.output) @@ -170,6 +180,12 @@ class EngineCostModel( functions.distinct.map { c => Term(c, Width(functions.size, 0), share * functions.count(_ == c)) } + case (SmjCrossCondition, join: SortMergeJoinExec) => + if (join.condition.exists(crossInputConditional(_, join))) { + Seq(Term(SmjCrossCondition, out)) + } else { + Nil + } case (AggObjectHash, agg: BaseAggregateExec) => Seq(Term(AggObjectHash, Width(0, 0), aggregateShare(agg))) case _ => Seq(Term(costClass, out)) diff --git a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala index 34761af94a..3d2e428452 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala @@ -172,6 +172,7 @@ object EngineCostTable { case object Sort extends CostClass("sort") case object SortSpill extends CostClass("sortSpill") case object Smj extends CostClass("smj") + case object SmjCrossCondition extends CostClass("smjCrossCondition") case object Bhj extends CostClass("bhj") case object Predicate extends CostClass("predicate") case object ProjectPassThrough extends CostClass("projectPassThrough") @@ -207,6 +208,7 @@ object EngineCostTable { Sort, SortSpill, Smj, + SmjCrossCondition, Bhj, Predicate, ProjectPassThrough, @@ -324,7 +326,16 @@ object EngineCostTable { * each pair of rows of equal keys before building the output, and on the cluster a left * join under a validity interval or a band cost 5.7 to 22 ns per pair natively against 9.6 * to 15 in Spark at 15 to 519 output leaves, an inner one 7.4 against 19.3, and one with - * 100 pairs per key about what Spark costs per output row. + * 100 pairs per key about what Spark costs per output row. The exception is a condition + * holding a CASE or an IF over columns of both inputs, which adds `smjCrossCondition`, + * every output leaf: on the cluster, left joins with about one passing pair per streamed + * row under `CASE WHEN l.c = r.c THEN l.t - r.eff ELSE 10 - (r.nxt - l.t) END BETWEEN 1 AND + * 10` cost 28 to 51 ns per pair natively against 8.7 to 17.6 in Spark (1.7 to 4.5 times), + * at 10k and 100 keys and 32 and 128 pairs per key. Cross-input `datediff` cost about what + * Spark does, `instr` and an OR over a cast of one input 4 to 10 times less, and conditions + * over one input at most Spark's price. The pairs per output row are not estimated, so the + * line keeps the former `smjCondition`'s, Comet 3.5 times Spark at a price high enough to + * outweigh `smj` and the conversions around the join at any width. * - `predicate`: a filter over the leaves its predicate reads, Spark noisy (maxrel 0.5 to * 0.8); passing rows costs the filter scalars. * - `projectPassThrough`: a project over every output leaf. Spark copies the row only when @@ -408,6 +419,7 @@ object EngineCostTable { (R2C, Flat) -> Line(3.4, 7.67, 0.036, 0, 0), (R2C, Nested) -> Line(0, 9.74, 0.040, 0, 0)) ++ both(ExprOverScan, Line(0, 0, 0, 2.5, 0)) ++ + both(SmjCrossCondition, Line(3000, 80, 0, 850, 23)) ++ both(AggObjectHash, Line(1000, 0, 0, 2500, 0)) ++ both(AggDeclarative, Line(6, 0, 0, 15, 0)) ++ both(AggDeclarativeNoCodegen, Line(0, 0, 0, 170, 0)) ++ @@ -470,7 +482,7 @@ object EngineCostTable { */ val operatorClasses: Map[String, Seq[CostClass]] = Map( "SortExec" -> Seq(Sort), - "SortMergeJoinExec" -> Seq(Smj), + "SortMergeJoinExec" -> Seq(Smj, SmjCrossCondition), "BroadcastHashJoinExec" -> Seq(Bhj), "WindowExec" -> Seq(Window), "WindowGroupLimitExec" -> Seq(WglPartial, WglFinal), diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index b38dc54b57..25ec4c2e0e 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -888,7 +888,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } } - test("a sort-merge join with a join condition costs what one without does") { + test("a sort-merge join adds smjCrossCondition only with a CASE or IF over both inputs") { withTables { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", @@ -905,20 +905,36 @@ class CostBasedEngineChoiceSuite extends CometTestBase { val band = join( "SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k " + "AND a.v <= b.v AND b.v <= a.v + 13") + val oneSide = join( + "SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k " + + "AND a.v <= CASE WHEN b.v > 500 THEN b.v ELSE b.v + 13 END") + val crossCase = join(crossCaseOnT) assert(band.originalPlan.asInstanceOf[SortMergeJoinExec].condition.isDefined) - for (engine <- Engine.all; j <- Seq(equi, band)) { + for (engine <- Engine.all; j <- Seq(equi, band, oneSide)) { assert(model.terms(j, engine) == Seq(Term(Smj, out))) } same(model.operatorPrice(equi, Engine.Comet), EngineCostTable.default.comet(Smj, out)) + val table = EngineCostTable.default for (engine <- Engine.all) { same(model.operatorPrice(band, engine), model.operatorPrice(equi, engine)) + assert( + model.terms(crossCase, engine) == Seq(Term(Smj, out), Term(SmjCrossCondition, out))) + val price: (CostClass, Width) => Double = + if (engine == Engine.Comet) table.comet else table.spark + same( + model.operatorPrice(crossCase, engine), + price(Smj, out) + price(SmjCrossCondition, out)) } } } } + private val crossCaseOnT = + "SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k " + + "AND CASE WHEN a.v % 2 = b.v % 2 THEN a.v - b.v ELSE 10 - (b.v - a.v) END BETWEEN 1 AND 10" + for (aqe <- Seq("false", "true")) { - test(s"a sort-merge join with a join condition stays native, as one without (AQE=$aqe)") { + test(s"a sort-merge join runs in Spark only under a cross-input CASE or IF (AQE=$aqe)") { withTables { withAqe( aqe, @@ -936,6 +952,15 @@ class CostBasedEngineChoiceSuite extends CometTestBase { assert(joins(on) == (1, 0), s"plan:\n$on") assert(cometOperatorNames(off) == cometOperatorNames(on), s"$off\n$on") } + val (crossOff, crossOn) = offAndOn(run(crossCaseOnT)) + assert(joins(crossOff) == (1, 0), s"plan:\n$crossOff") + assert(joins(crossOn) == (0, 1), s"plan:\n$crossOn") + withSQLConf( + flag -> "true", + costTable -> "smjCrossCondition.comet=0,0,0;smjCrossCondition.spark=0,0") { + val plan = run(crossCaseOnT) + assert(joins(plan) == (1, 0), s"plan:\n$plan") + } } } } @@ -1011,6 +1036,24 @@ class CostBasedEngineChoiceSuite extends CometTestBase { "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k AND e.v >= d.l AND e.v < d.u " + "AND e.s <> d.s") + private val crossConditionals = Seq( + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k AND CASE WHEN e.s = d.s " + + "THEN e.vs - d.ls ELSE 10 - (d.ls - e.vs) END BETWEEN 1 AND 10", + "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k AND e.v >= d.l " + + "AND IF(d.u IS NULL, e.vs < d.ls + 2592000, e.v < d.u)") + + private val unchargedConditions = Seq( + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND e.v >= CASE WHEN d.s = 's1' THEN d.l ELSE d.created_at END AND e.v < d.u", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND IF(e.s = 's1', e.vs, e.vs + 3600) BETWEEN d.ls AND d.ls + 86400", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND datediff(e.v, d.l) BETWEEN 0 AND 30", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND e.vs - d.ls BETWEEN 1 AND 3600", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND (e.s = d.s OR CAST(d.ls AS string) = e.s)") + private def conditionTerms(query: String): Seq[Seq[Term]] = { val plan = runUnordered(sql(query)) val joins = nodes(plan).collect { @@ -1024,7 +1067,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } for (aqe <- Seq("false", "true")) { - test(s"a join condition of any shape adds nothing to smj (AQE=$aqe)") { + test(s"only a CASE or IF over both inputs adds smjCrossCondition (AQE=$aqe)") { withIntervals { withAqe( aqe, @@ -1034,14 +1077,19 @@ class CostBasedEngineChoiceSuite extends CometTestBase { val intervals = for (condition <- validityIntervals; joinType <- Seq("JOIN", "LEFT JOIN")) yield s"SELECT e.k, e.v, d.l FROM ev e $joinType dim d ON e.k = d.k AND $condition" - for (query <- intervals ++ intervalQueries ++ otherConditions) { + for (query <- intervals ++ intervalQueries ++ otherConditions ++ unchargedConditions) { assert(conditionTerms(query).forall(_.map(_.costClass) == Seq(Smj)), query) } + for (query <- crossConditionals) { + assert( + conditionTerms(query).exists(_.map(_.costClass) == Seq(Smj, SmjCrossCondition)), + query) + } } } } - test(s"validity interval, band and cross-input joins stay native (AQE=$aqe)") { + test(s"a join under a CASE or IF over both inputs runs in Spark, others native (AQE=$aqe)") { withIntervals { withAqe( aqe, @@ -1056,11 +1104,16 @@ class CostBasedEngineChoiceSuite extends CometTestBase { for ((query, native) <- Seq( interval -> 1, otherConditions(1) -> 1, - crossInputBounds -> 2)) { + crossInputBounds -> 2) ++ unchargedConditions.map(_ -> 1)) { val (off, on) = offAndOn(run(query)) assert(joins(off) == (native, 0), s"plan:\n$off") assert(joins(on) == (native, 0), s"plan:\n$on") } + for (query <- crossConditionals) { + val (off, on) = offAndOn(run(query)) + assert(joins(off) == (1, 0), s"plan:\n$off") + assert(joins(on) == (0, 1), s"plan:\n$on") + } } } } From 780cb99a851c27bdf62b86fb1ed97e3c566b92ab Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 11:54:35 +0100 Subject: [PATCH 12/13] test: drop an unneeded string interpolator in the SMJ fuzz suite Co-Authored-By: Claude Opus 5.5 (1M context) --- .../org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala index 220757b080..d5e9c79fca 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala @@ -426,7 +426,7 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe s"CASE WHEN l.flag THEN $rts > TIMESTAMP '1970-01-06 00:00:00' ELSE r.rk2 > 2 END") private val caseNested = Filter( "case_nested", - s"CASE WHEN l.k2 IS NULL THEN r.rv IS NULL WHEN l.k2 > 2 THEN " + + "CASE WHEN l.k2 IS NULL THEN r.rv IS NULL WHEN l.k2 > 2 THEN " + s"CASE WHEN $lts IS NULL THEN r.rk2 = l.k2 ELSE $lts < r.rt END " + s"ELSE CASE WHEN r.rk2 IS NULL THEN l.flag ELSE $lint < r.rv END END") private val caseValue = From 5bf9f6e10f42488c3cd0079ef2fe50473fe3dc78 Mon Sep 17 00:00:00 2001 From: msaf Date: Thu, 8 Oct 2026 13:49:48 +0100 Subject: [PATCH 13/13] test: run the SMJ fuzz suite with ANSI off on every Spark version Spark 4 enables ANSI by default, where casting the suite's malformed strings throws in Spark itself instead of returning NULL. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala | 1 + 1 file changed, 1 insertion(+) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala index d5e9c79fca..496bfd3c52 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometSmjJoinFilterFuzzSuite.scala @@ -520,6 +520,7 @@ class CometSmjJoinFilterFuzzSuite extends CometTestBase with AdaptiveSparkPlanHe SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", SQLConf.PREFER_SORTMERGEJOIN.key -> "true", + SQLConf.ANSI_ENABLED.key -> "false", CometConf.COMET_EXEC_SORT_MERGE_JOIN_ENABLED.key -> "true", CometConf.COMET_EXEC_SORT_MERGE_JOIN_WITH_JOIN_FILTER_ENABLED.key -> "true", CometConf.COMET_FORCE_SHJ.key -> "false")