From 519e137c8bdb3be1f34155aca32a5c84be3cded0 Mon Sep 17 00:00:00 2001 From: Jay Zhan Date: Sun, 13 Sep 2026 11:35:01 +0800 Subject: [PATCH 1/2] refactor: factor SortMergeJoinExec::execute into a reusable sort_merge_join_stream --- .../src/joins/sort_merge_join/exec.rs | 206 +++++++++++------- 1 file changed, 129 insertions(+), 77 deletions(-) diff --git a/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs b/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs index b0433250c0c51..102e19624e92b 100644 --- a/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs +++ b/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs @@ -583,86 +583,27 @@ impl ExecutionPlan for SortMergeJoinExec { "Invalid SortMergeJoinExec, partition count mismatch {left_partitions}!={right_partitions},\ consider using RepartitionExec" ); - let (on_left, on_right) = self.on.iter().cloned().unzip(); - let (streamed, buffered, on_streamed, on_buffered) = - if SortMergeJoinExec::probe_side(&self.join_type) == JoinSide::Left { - ( - Arc::clone(&self.left), - Arc::clone(&self.right), - on_left, - on_right, - ) - } else { - ( - Arc::clone(&self.right), - Arc::clone(&self.left), - on_right, - on_left, - ) - }; - // execute children plans - let streamed = streamed.execute(partition, Arc::clone(&context))?; - let buffered = buffered.execute(partition, Arc::clone(&context))?; - - let batch_size = context.session_config().batch_size(); - // The stream spills its buffered batches when it cannot grow, so a - // pool that budgets spillable and unspillable consumers differently - // (`FairSpillPool`) has to know it can. - let reservation = MemoryConsumer::new(format!("SMJStream[{partition}]")) - .with_can_spill(true) - .register(context.memory_pool()); - let spill_manager = SpillManager::new( - context.runtime_env(), - SpillMetrics::new(&self.metrics, partition), - buffered.schema(), - ) - .with_compression_type(context.session_config().spill_compression()); + let left = self.left.execute(partition, Arc::clone(&context))?; + let right = self.right.execute(partition, Arc::clone(&context))?; + let (on_left, on_right) = self.on.iter().cloned().unzip(); - let joined = if matches!( - self.join_type, - JoinType::LeftSemi - | JoinType::LeftAnti - | JoinType::RightSemi - | JoinType::RightAnti - | JoinType::LeftMark - | JoinType::RightMark - ) { - BitwiseSortMergeJoinStream::try_new( - Arc::clone(&self.schema), - self.sort_options.clone(), - self.null_equality, - streamed, - buffered, - on_streamed, - on_buffered, - self.filter.clone(), - self.join_type, - batch_size, + let joined = sort_merge_join_stream( + SortMergeJoinInputs { + schema: Arc::clone(&self.schema), + sort_options: self.sort_options.clone(), + null_equality: self.null_equality, + left, + right, + on_left, + on_right, + filter: self.filter.clone(), + join_type: self.join_type, partition, - &self.metrics, - reservation, - spill_manager, - context.runtime_env(), - ) - } else { - MaterializingSortMergeJoinStream::try_new( - Arc::clone(&self.schema), - self.sort_options.clone(), - self.null_equality, - streamed, - buffered, - on_streamed, - on_buffered, - self.filter.clone(), - self.join_type, - batch_size, - SortMergeJoinMetrics::new(partition, &self.metrics), - reservation, - spill_manager, - context.runtime_env(), - ) - }?; + }, + &self.metrics, + &context, + )?; let Some(projection) = self.projection.clone() else { return Ok(joined); @@ -949,3 +890,114 @@ impl SortMergeJoinExec { )) } } + +/// The two sorted inputs of one partition of a sort-merge join, and how to +/// join them. +pub(crate) struct SortMergeJoinInputs { + /// The join schema (before any projection) + pub(crate) schema: SchemaRef, + /// Sort options of the join keys, one per key, that both inputs are sorted with + pub(crate) sort_options: Vec, + pub(crate) null_equality: NullEquality, + /// Left input, sorted on `on_left` with `sort_options` + pub(crate) left: SendableRecordBatchStream, + /// Right input, sorted on `on_right` with `sort_options` + pub(crate) right: SendableRecordBatchStream, + pub(crate) on_left: Vec, + pub(crate) on_right: Vec, + pub(crate) filter: Option, + pub(crate) join_type: JoinType, + pub(crate) partition: usize, +} + +/// Joins two sorted inputs with the sort-merge join algorithm. +/// +/// Picks the streamed and buffered side by join type and the join stream +/// implementation by join type family, exactly as [`SortMergeJoinExec`] does; +/// the hash join's sort-merge fallback uses this too. The stream's metrics are +/// registered in `metrics`; its buffered side spills through the context's +/// disk manager under the memory pool's control. +pub(crate) fn sort_merge_join_stream( + inputs: SortMergeJoinInputs, + metrics: &ExecutionPlanMetricsSet, + context: &Arc, +) -> Result { + let SortMergeJoinInputs { + schema, + sort_options, + null_equality, + left, + right, + on_left, + on_right, + filter, + join_type, + partition, + } = inputs; + + let (streamed, buffered, on_streamed, on_buffered) = + if SortMergeJoinExec::probe_side(&join_type) == JoinSide::Left { + (left, right, on_left, on_right) + } else { + (right, left, on_right, on_left) + }; + + let batch_size = context.session_config().batch_size(); + // The stream spills its buffered batches when it cannot grow, so a + // pool that budgets spillable and unspillable consumers differently + // (`FairSpillPool`) has to know it can. + let reservation = MemoryConsumer::new(format!("SMJStream[{partition}]")) + .with_can_spill(true) + .register(context.memory_pool()); + let spill_manager = SpillManager::new( + context.runtime_env(), + SpillMetrics::new(metrics, partition), + buffered.schema(), + ) + .with_compression_type(context.session_config().spill_compression()); + + if matches!( + join_type, + JoinType::LeftSemi + | JoinType::LeftAnti + | JoinType::RightSemi + | JoinType::RightAnti + | JoinType::LeftMark + | JoinType::RightMark + ) { + BitwiseSortMergeJoinStream::try_new( + schema, + sort_options, + null_equality, + streamed, + buffered, + on_streamed, + on_buffered, + filter, + join_type, + batch_size, + partition, + metrics, + reservation, + spill_manager, + context.runtime_env(), + ) + } else { + MaterializingSortMergeJoinStream::try_new( + schema, + sort_options, + null_equality, + streamed, + buffered, + on_streamed, + on_buffered, + filter, + join_type, + batch_size, + SortMergeJoinMetrics::new(partition, metrics), + reservation, + spill_manager, + context.runtime_env(), + ) + } +} From 925c52a3dbbdf2284ecac85639c10b99e466e34a Mon Sep 17 00:00:00 2001 From: Jay Zhan Date: Mon, 21 Sep 2026 20:16:57 +0800 Subject: [PATCH 2/2] docs: sort_merge_join_stream has a single caller for now --- datafusion/physical-plan/src/joins/sort_merge_join/exec.rs | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs b/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs index 102e19624e92b..c0e57a93a49a1 100644 --- a/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs +++ b/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs @@ -913,8 +913,9 @@ pub(crate) struct SortMergeJoinInputs { /// Joins two sorted inputs with the sort-merge join algorithm. /// /// Picks the streamed and buffered side by join type and the join stream -/// implementation by join type family, exactly as [`SortMergeJoinExec`] does; -/// the hash join's sort-merge fallback uses this too. The stream's metrics are +/// implementation by join type family. [`SortMergeJoinExec::execute`] is the +/// only caller today; it is a separate function so that an operator which +/// already holds two sorted streams can reuse it. The stream's metrics are /// registered in `metrics`; its buffered side spills through the context's /// disk manager under the memory pool's control. pub(crate) fn sort_merge_join_stream(