From 52b9fd5700d8005b885846c0131be725f0eee14d Mon Sep 17 00:00:00 2001 From: Jay Zhan Date: Sat, 12 Sep 2026 12:54:30 +0800 Subject: [PATCH 1/3] feat: sort-merge fallback for hash joins under memory pressure A hash join materializes its whole build side in memory and fails the query when the memory pool refuses that reservation. Sort-merge join already runs in bounded memory, so a partitioned hash join whose build side does not fit now sorts both of its inputs with an external (spilling) sort and finishes by merging them instead of failing. The switch is decided per output partition at runtime, after the join's own `try_grow` fails, so joins that fit in memory are untouched. It requires disk spilling to be enabled. A new option, `hash_join_max_build_size`, triggers the same switch once a partition's build side grows past a given size; that is the only trigger available under the default, unbounded memory pool. The fallback is restricted to `PartitionMode::Partitioned`, where every partition owns a disjoint slice of the key space on both sides and can therefore sort and merge independently of its siblings. `CollectLeft`, null-aware joins, and joins that promise to preserve their probe-side ordering keep the existing behavior. Dynamic filter pushdown needs a third report for a build side that holds rows but has no hash table to test membership against, so `PushdownStrategy::Unknown` pushes down bounds only. Reporting `Empty` there would prune every probe row routed to the partition and silently drop its matches. Also fixes a latent bug in the sort-merge join: `get_filter_columns` assembled filter columns as all-left-then-all-right, which is wrong for the layout `JoinFilter::swap` produces. Only this fallback could reach it, since the planner never swaps a `SortMergeJoinExec`. --- datafusion/common/src/config.rs | 26 + datafusion/core/tests/helper/plan_metrics.rs | 17 + .../hash_join_sort_merge_fallback.rs | 276 +++++++ datafusion/core/tests/memory_limit/mod.rs | 1 + .../physical-plan/src/joins/hash_join/exec.rs | 685 ++++++++++++++++-- .../physical-plan/src/joins/hash_join/mod.rs | 1 + .../src/joins/hash_join/shared_bounds.rs | 125 +++- .../joins/hash_join/sort_merge_fallback.rs | 361 +++++++++ .../src/joins/hash_join/stream.rs | 170 ++++- .../src/joins/sort_merge_join/exec.rs | 198 +++-- .../src/joins/sort_merge_join/filter.rs | 36 +- .../src/joins/sort_merge_join/mod.rs | 1 + .../src/joins/sort_merge_join/tests.rs | 60 ++ datafusion/physical-plan/src/sorts/sort.rs | 13 +- .../test_files/information_schema.slt | 2 + docs/source/user-guide/configs.md | 1 + 16 files changed, 1810 insertions(+), 163 deletions(-) create mode 100644 datafusion/core/tests/memory_limit/hash_join_sort_merge_fallback.rs create mode 100644 datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs diff --git a/datafusion/common/src/config.rs b/datafusion/common/src/config.rs index 360586b0e9bae..15e349711a52f 100644 --- a/datafusion/common/src/config.rs +++ b/datafusion/common/src/config.rs @@ -1050,6 +1050,32 @@ config_namespace! { /// unaffected and always keep the fallback. pub enable_nlj_coordinated_fallback: bool, default = true + /// Maximum build-side size, in bytes per output partition, that a hash + /// join keeps in memory. `NULL` (the default) means no limit. A join + /// whose build side grows past this size finishes as a sort-merge join + /// instead: both inputs are sorted with an external (spilling) sort and + /// then merged. + /// + /// This is not needed for memory safety: with + /// `datafusion.runtime.memory_limit` set, a hash join that cannot + /// reserve memory for its build side already falls back the same way. + /// Set it to switch over earlier, when no memory limit is configured or + /// to make the choice reproducible. + /// + /// The size is per output partition, so a join may hold it times + /// `datafusion.execution.target_partitions`: divide the memory you want + /// hash joins to use by that count, giving 1 GB here for an 8 GB budget + /// over 8 partitions. + /// + /// Requires disk spilling (see `DiskManager`); a partition that falls + /// back reserves about `2 * sort_spill_reservation_bytes` for its sorts, + /// so a budget too small to cover that still fails, inside the sort. + /// Only `PartitionMode::Partitioned` joins fall back. + /// + /// The sort-merge fallback is a short-term measure, so this option may + /// be deprecated and removed once hash joins spill natively. + pub hash_join_max_build_size: Option, default = None + /// Number of files to read in parallel when inferring schema and statistics pub meta_fetch_concurrency: ConfigNonZeroUsize, default = non_zero_usize_default(32) diff --git a/datafusion/core/tests/helper/plan_metrics.rs b/datafusion/core/tests/helper/plan_metrics.rs index 12d3eaba1ad96..fc2ce68e8cb31 100644 --- a/datafusion/core/tests/helper/plan_metrics.rs +++ b/datafusion/core/tests/helper/plan_metrics.rs @@ -52,3 +52,20 @@ pub fn plan_spilled_bytes(plan: &dyn ExecutionPlan) -> usize { .map(|child| plan_spilled_bytes(child.as_ref())) .sum::() } + +/// Sums the counter or gauge named `name` across every operator of `plan`. +/// +/// Returns 0 when no operator reports that metric. +pub fn plan_metric_sum(plan: &dyn ExecutionPlan, name: &str) -> usize { + let own = plan + .metrics() + .and_then(|m| m.sum_by_name(name)) + .map(|v| v.as_usize()) + .unwrap_or(0); + + own + plan + .children() + .into_iter() + .map(|child| plan_metric_sum(child.as_ref(), name)) + .sum::() +} diff --git a/datafusion/core/tests/memory_limit/hash_join_sort_merge_fallback.rs b/datafusion/core/tests/memory_limit/hash_join_sort_merge_fallback.rs new file mode 100644 index 0000000000000..b2573ce5d4afa --- /dev/null +++ b/datafusion/core/tests/memory_limit/hash_join_sort_merge_fallback.rs @@ -0,0 +1,276 @@ +// 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. + +//! A partitioned hash join whose build side does not fit in memory finishes +//! as a sort-merge join (see `datafusion.execution.hash_join_max_build_size`) + +use std::sync::Arc; + +use arrow::array::{Int32Array, RecordBatch}; +use arrow::datatypes::{DataType, Field, Schema}; +use datafusion::prelude::*; +use datafusion_common::assert_contains; +use datafusion_execution::disk_manager::{DiskManagerBuilder, DiskManagerMode}; +use datafusion_execution::runtime_env::RuntimeEnvBuilder; + +use crate::helper::plan_metrics::plan_metric_sum; + +fn table(rows: usize, value_column: &str, seed: u64) -> RecordBatch { + table_with_keys(rows, value_column, seed, 500) +} + +/// `rows` rows of a nullable `k` key drawn from `distinct_keys` values and a +/// nullable value column. A larger `distinct_keys` makes the join output +/// smaller without shrinking the build side. +fn table_with_keys( + rows: usize, + value_column: &str, + seed: u64, + distinct_keys: u64, +) -> RecordBatch { + let mut state = seed; + let mut next = || { + state = state + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + state >> 33 + }; + let schema = Arc::new(Schema::new(vec![ + Field::new("k", DataType::Int32, true), + Field::new(value_column, DataType::Int32, true), + ])); + let keys: Vec> = (0..rows) + .map(|_| (next() % 50 != 0).then(|| (next() % distinct_keys) as i32)) + .collect(); + let values: Vec> = + (0..rows).map(|_| Some((next() % 60) as i32)).collect(); + RecordBatch::try_new( + schema, + vec![ + Arc::new(Int32Array::from(keys)), + Arc::new(Int32Array::from(values)), + ], + ) + .unwrap() +} + +/// Both tables are joined by partitioned hash joins whatever their size +fn config() -> SessionConfig { + SessionConfig::new() + .with_target_partitions(2) + .with_batch_size(256) + .set_usize( + "datafusion.optimizer.hash_join_single_partition_threshold", + 0, + ) + .set_usize( + "datafusion.optimizer.hash_join_single_partition_threshold_rows", + 0, + ) + // Scaled down from the 10 MB default to match these small budgets. A + // falling-back partition runs two sorts, each pre-reserving this much + // for its merge, so the default would consume every budget small + // enough to trigger the fallback in the first place. + .with_sort_spill_reservation_bytes(64 * 1024) +} + +fn context(config: SessionConfig, runtime: RuntimeEnvBuilder) -> SessionContext { + let ctx = SessionContext::new_with_config_rt(config, runtime.build_arc().unwrap()); + ctx.register_batch("l", table(30_000, "v", 11)).unwrap(); + ctx.register_batch("r", table(20_000, "w", 16)).unwrap(); + ctx +} + +fn rows(batches: &[RecordBatch]) -> Vec { + let mut rows = vec![]; + for batch in batches { + for row in 0..batch.num_rows() { + let cells: Vec = (0..batch.num_columns()) + .map(|col| { + datafusion_common::ScalarValue::try_from_array(batch.column(col), row) + .unwrap() + .to_string() + }) + .collect(); + rows.push(cells.join("|")); + } + } + rows.sort(); + rows +} + +/// A context whose tables are big enough that each partition builds about 1 MB, +/// keyed widely so the join output stays small. +fn wide_context( + memory_limit: Option, + spilling: bool, + max_build_size: Option, +) -> SessionContext { + let mut cfg = config(); + if let Some(max_build_size) = max_build_size { + cfg = cfg.set_usize( + "datafusion.execution.hash_join_max_build_size", + max_build_size, + ); + } + let mut runtime = RuntimeEnvBuilder::new(); + if let Some(limit) = memory_limit { + runtime = runtime.with_memory_limit(limit, 1.0); + } + if !spilling { + runtime = runtime.with_disk_manager_builder( + DiskManagerBuilder::default().with_mode(DiskManagerMode::Disabled), + ); + } + let ctx = SessionContext::new_with_config_rt(cfg, runtime.build_arc().unwrap()); + ctx.register_batch("l", table_with_keys(240_000, "v", 11, 100_000)) + .unwrap(); + ctx.register_batch("r", table_with_keys(160_000, "w", 16, 100_000)) + .unwrap(); + ctx +} + +/// Runs `sql` and returns its rows together with how many partitions of the +/// hash join finished as a sort-merge join. +async fn run_counting_fallbacks(ctx: &SessionContext, sql: &str) -> (Vec, usize) { + let plan = ctx + .sql(sql) + .await + .unwrap() + .create_physical_plan() + .await + .unwrap(); + let batches = datafusion::physical_plan::collect(Arc::clone(&plan), ctx.task_ctx()) + .await + .unwrap(); + let fallbacks = plan_metric_sum(plan.as_ref(), "sort_merge_fallback_count"); + (rows(&batches), fallbacks) +} + +/// Runs `sql` as an ordinary hash join and again with a build-size cap far +/// below the build side, returning both results and how many partitions fell +/// back in the second run. +/// +/// Neither run is under a memory limit, so the cap alone decides the algorithm +/// and the comparison cannot be perturbed by whatever else is allocating. +async fn run_with_and_without_cap(sql: &str) -> (Vec, Vec, usize) { + let plain = context(config(), RuntimeEnvBuilder::new()); + let (plain_rows, fallbacks) = run_counting_fallbacks(&plain, sql).await; + assert_eq!( + fallbacks, 0, + "{sql}: without a cap nothing should fall back" + ); + + let capped = context( + config().set_usize("datafusion.execution.hash_join_max_build_size", 16 * 1024), + RuntimeEnvBuilder::new(), + ); + let (capped_rows, fallbacks) = run_counting_fallbacks(&capped, sql).await; + (plain_rows, capped_rows, fallbacks) +} + +/// The whole decision in one table: which configuration runs a plain hash join, +/// which switches to a sort-merge join, and which still fails outright. +/// +/// Every case runs the same query over the same data, so only the configuration +/// differs, and every successful case must return the same rows as an +/// unconstrained run. +#[tokio::test] +async fn config_matrix() { + #[derive(Debug, PartialEq)] + enum Expect { + /// Stays a hash join, holding its whole build side in memory + HashJoin, + /// Both partitions finish as a sort-merge join + Fallback, + /// Nothing to fall back to, so the query fails as it did before + Fails, + } + use Expect::*; + + // Aggregated so comparing answers is cheap, and keyed widely so the join + // emits few rows while each partition still builds about 1 MB. + let sql = "SELECT count(*), sum(l.v), sum(r.w) FROM l JOIN r ON l.k = r.k"; + const KB: usize = 1024; + const MB: usize = 1024 * KB; + + #[rustfmt::skip] + let cases = [ + // what the configuration is memory limit spilling cap expected + ("no memory limit and no cap", None, true, None, HashJoin), + ("no memory limit, cap under the build side", None, true, Some(16 * KB), Fallback), + ("memory limit under the build side, spilling on", Some(512 * KB), true, None, Fallback), + ("memory limit under the build side, spilling off", Some(512 * KB), false, None, Fails), + ("memory limit far above the build side", Some(64 * MB), true, None, HashJoin), + ]; + + let reference = { + let ctx = wide_context(None, true, None); + let (rows, _) = run_counting_fallbacks(&ctx, sql).await; + assert!(!rows.is_empty()); + rows + }; + + for (what, limit, spilling, cap, expect) in cases { + let ctx = wide_context(limit, spilling, cap); + let plan = ctx + .sql(sql) + .await + .unwrap() + .create_physical_plan() + .await + .unwrap(); + let result = + datafusion::physical_plan::collect(Arc::clone(&plan), ctx.task_ctx()).await; + let fallbacks = plan_metric_sum(plan.as_ref(), "sort_merge_fallback_count"); + + match expect { + Fails => { + let err = result.expect_err(&format!("{what}: expected failure")); + assert_contains!(err.to_string(), "Resources exhausted"); + } + HashJoin | Fallback => { + let batches = result.unwrap_or_else(|e| panic!("{what}: {e}")); + assert_eq!(rows(&batches), reference, "{what}: wrong answer"); + if expect == Fallback { + // Normally every partition switches. Only "at least one" is + // asserted because the `force_hash_collisions` test feature + // routes every row to a single partition, leaving the others + // with an empty build side and nothing to fall back for. + assert!(fallbacks >= 1, "{what}: expected the fallback"); + } else { + assert_eq!(fallbacks, 0, "{what}: expected a plain hash join"); + } + } + } + } +} + +#[tokio::test] +async fn semi_and_anti_joins_fall_back() { + for sql in [ + "SELECT l.k, l.v FROM l WHERE EXISTS (SELECT 1 FROM r WHERE l.k = r.k)", + "SELECT l.k, l.v FROM l WHERE NOT EXISTS (SELECT 1 FROM r WHERE l.k = r.k)", + "SELECT l.k, l.v FROM l WHERE l.k IN (SELECT r.k FROM r WHERE r.w > 30)", + ] { + let (plain, capped, fallbacks) = run_with_and_without_cap(sql).await; + // See `config_matrix` on why this is not an exact count. + assert!(fallbacks >= 1, "{sql}: expected the fallback"); + assert!(!plain.is_empty(), "{sql}"); + assert_eq!(capped, plain, "{sql}"); + } +} diff --git a/datafusion/core/tests/memory_limit/mod.rs b/datafusion/core/tests/memory_limit/mod.rs index 2cda54d678250..f2802372238dd 100644 --- a/datafusion/core/tests/memory_limit/mod.rs +++ b/datafusion/core/tests/memory_limit/mod.rs @@ -20,6 +20,7 @@ use std::num::NonZeroUsize; use std::sync::{Arc, LazyLock}; +mod hash_join_sort_merge_fallback; #[cfg(feature = "extended_tests")] mod memory_limit_validation; mod nlj_spill_unmatched; diff --git a/datafusion/physical-plan/src/joins/hash_join/exec.rs b/datafusion/physical-plan/src/joins/hash_join/exec.rs index b72e180543f9a..fc9d80d2ef733 100644 --- a/datafusion/physical-plan/src/joins/hash_join/exec.rs +++ b/datafusion/physical-plan/src/joins/hash_join/exec.rs @@ -36,6 +36,9 @@ use crate::joins::hash_join::probe_completion::{ProbeCompletion, ProbeSideSummar use crate::joins::hash_join::shared_bounds::{ ColumnBounds, PartitionBounds, PushdownStrategy, SharedBuildAccumulator, }; +use crate::joins::hash_join::sort_merge_fallback::{ + BuildSideOutcome, SortMergeFallbackContext, is_resources_exhausted, sort_build_side, +}; use crate::joins::hash_join::stream::{ BuildSide, BuildSideInitialState, HashJoinStream, HashJoinStreamState, }; @@ -45,7 +48,7 @@ use crate::joins::utils::{ is_existence_join, reorder_output_after_swap, swap_join_projection, update_hash, }; use crate::joins::{JoinOn, JoinOnRef, PartitionMode, SharedBitmapBuilder}; -use crate::metrics::{Count, MetricBuilder, MetricCategory}; +use crate::metrics::{Count, MetricBuilder, MetricCategory, SpillMetrics}; use crate::projection::{ EmbeddedProjection, JoinData, ProjectionExec, try_embed_projection, try_pushdown_through_join_with_column_indices, @@ -70,9 +73,10 @@ use crate::{ }; use arrow::array::{Array, ArrayRef, BooleanBufferBuilder, UInt64Array}; -use arrow::compute::concat_batches; +use arrow::compute::{SortOptions, concat_batches}; use arrow::datatypes::SchemaRef; use arrow::record_batch::RecordBatch; +use arrow::row::{RowConverter, SortField}; use arrow::util::bit_util; use arrow_schema::{DataType, Schema}; use datafusion_common::config::ConfigOptions; @@ -96,7 +100,7 @@ use datafusion_physical_expr::{PhysicalExpr, PhysicalExprRef}; use datafusion_common::hash_utils::{RandomState, create_hashes}; use datafusion_physical_expr_common::physical_expr::fmt_sql; use datafusion_physical_expr_common::utils::evaluate_expressions_to_arrays; -use futures::TryStreamExt; +use futures::StreamExt; use parking_lot::Mutex; use super::partitioned_hash_eval::SeededRandomState; @@ -106,6 +110,8 @@ pub(crate) const HASH_JOIN_SEED: SeededRandomState = SeededRandomState::with_seed(12210250226015887276); const ARRAY_MAP_CREATED_COUNT_METRIC_NAME: &str = "array_map_created_count"; +/// Number of partitions that fell back to a sort-merge join under memory pressure +const SORT_MERGE_FALLBACK_COUNT_METRIC_NAME: &str = "sort_merge_fallback_count"; #[expect(clippy::too_many_arguments)] fn try_create_array_map( @@ -862,7 +868,7 @@ pub struct HashJoinExec { /// /// Each output stream waits on the `OnceAsync` to signal the completion of /// the hash table creation. - left_fut: Arc>, + left_fut: Arc>, /// Shared the `SeededRandomState` for the hashing algorithm (seeds preserved for serialization) random_state: SeededRandomState, /// Partitioning mode to use @@ -971,6 +977,66 @@ impl HashJoinExec { self.into() } + /// Returns what a partition needs to fall back to a sort-merge join when + /// its build side does not fit in memory, or `None` when this join cannot + /// fall back (see the `sort_merge_fallback` module). + fn sort_merge_fallback_context( + &self, + partition: usize, + context: &Arc, + ) -> Result> { + let options = context.session_config().options(); + if !context.runtime_env().disk_manager.tmp_files_enabled() + // Only `Partitioned` joins are self-contained per partition. A + // `CollectLeft` build side is shared by every probe partition and + // the join types that emit unmatched build rows need all of them + // to agree on what was matched. + || self.mode != PartitionMode::Partitioned + // The sort-merge join streams do not implement `NOT IN` semantics. + || self.null_aware + // The sort-merge output is ordered by the join keys, not by the + // probe side; a promised probe-side ordering must be honored. + || self.cache.output_ordering().is_some() + { + return Ok(None); + } + + // Both sides are sorted with the row format, which does not support + // every type a hash join can hash. + let left_schema = self.left.schema(); + for (left_key, _) in &self.on { + let data_type = left_key.data_type(&left_schema)?; + if !RowConverter::supports_fields(&[SortField::new(data_type)]) { + return Ok(None); + } + } + + let (on_left, on_right) = self + .on + .iter() + .map(|(l, r)| (Arc::clone(l), Arc::clone(r))) + .unzip::<_, _, Vec<_>, Vec<_>>(); + Ok(Some(SortMergeFallbackContext { + context: Arc::clone(context), + partition, + metrics: self.metrics.clone(), + spill_metrics: SpillMetrics::new(&self.metrics, partition), + fallback_count: MetricBuilder::new(&self.metrics) + .counter(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME, partition), + sort_options: vec![SortOptions::default(); on_left.len()], + on_left, + on_right, + join_type: self.join_type, + filter: self.filter.clone(), + null_equality: self.null_equality, + join_schema: Arc::clone(&self.join_schema), + output_schema: self.schema(), + projection: self.projection.as_deref().map(|p| p.to_vec()), + fetch: self.fetch, + max_build_size: options.execution.hash_join_max_build_size, + })) + } + fn create_dynamic_filter(on: &JoinOn) -> Arc { // Extract the right-side keys (probe side keys) from the `on` clauses // Dynamic filter will be created from build side values (left side) and applied to probe side (right side) @@ -1618,6 +1684,8 @@ impl ExecutionPlan for HashJoinExec { .flatten(); let null_aware = self.null_aware_mode()?; + let sort_merge_fallback = + self.sort_merge_fallback_context(partition, &context)?; let left_fut = match self.mode { PartitionMode::CollectLeft => self.left_fut.try_once(|| { @@ -1639,6 +1707,7 @@ impl ExecutionPlan for HashJoinExec { self.null_equality, null_aware, array_map_created_count, + None, )) })?, PartitionMode::Partitioned => { @@ -1660,6 +1729,7 @@ impl ExecutionPlan for HashJoinExec { self.null_equality, null_aware, array_map_created_count, + sort_merge_fallback.clone(), )) } PartitionMode::Auto => { @@ -1711,6 +1781,7 @@ impl ExecutionPlan for HashJoinExec { self.mode, null_aware, self.fetch, + sort_merge_fallback, ))) } @@ -2539,7 +2610,7 @@ fn lr_is_preserved(join_type: JoinType) -> (bool, bool) { /// The bounds are used for dynamic filter pushdown optimization, where filters /// based on the actual data ranges can be pushed down to the probe side to /// eliminate unnecessary data early. -struct CollectLeftAccumulator { +pub(super) struct CollectLeftAccumulator { /// The physical expression to evaluate for each batch expr: Arc, /// Accumulator for tracking the minimum value across all batches @@ -2589,7 +2660,7 @@ impl CollectLeftAccumulator { /// /// # Returns /// Ok(()) if the update succeeds, or an error if expression evaluation fails - fn update_batch(&mut self, batch: &RecordBatch) -> Result<()> { + pub(super) fn update_batch(&mut self, batch: &RecordBatch) -> Result<()> { let array = self.expr.evaluate(batch)?.into_array(batch.num_rows())?; self.min.update_batch(std::slice::from_ref(&array))?; self.max.update_batch(std::slice::from_ref(&array))?; @@ -2602,7 +2673,7 @@ impl CollectLeftAccumulator { /// /// # Returns /// The `ColumnBounds` containing the minimum and maximum values observed - fn evaluate(mut self) -> Result { + pub(super) fn evaluate(mut self) -> Result { Ok(ColumnBounds::new( self.min.evaluate()?, self.max.evaluate()?, @@ -2735,7 +2806,8 @@ async fn collect_left_input( null_equality: NullEquality, null_aware: Option, array_map_created_count: Count, -) -> Result { + sort_merge_fallback: Option, +) -> Result { let schema = left_stream.schema(); // The extra scope maps + null bitmap are only built for correlated @@ -2747,7 +2819,7 @@ async fn collect_left_input( let is_phj_candidate = is_perfect_hash_join_candidate(&on_left, &schema)?; - let initial = BuildSideState::try_new( + let mut state = BuildSideState::try_new( metrics, reservation, on_left.clone(), @@ -2755,43 +2827,82 @@ async fn collect_left_input( should_compute_dynamic_filters || is_phj_candidate, )?; - let state = left_stream - .try_fold(initial, |mut state, batch| async move { - // Update accumulators if computing bounds - if let Some(ref mut accumulators) = state.bounds_accumulators { - for accumulator in accumulators { - accumulator.update_batch(&batch)?; - } + let mut left_stream = left_stream; + let mut sort_merge_fallback = sort_merge_fallback; + while let Some(batch) = left_stream.next().await { + let batch = batch?; + // Update accumulators if computing bounds + if let Some(ref mut accumulators) = state.bounds_accumulators { + for accumulator in accumulators { + accumulator.update_batch(&batch)?; } + } - // Decide if we spill or not - let batch_size = state.memory_counter.count_batch(&batch); - // Reserve memory for incoming batch - state.reservation.try_grow(batch_size)?; - // Update metrics - state.metrics.build_mem_used.add(batch_size); - state.metrics.build_input_batches.add(1); - state.metrics.build_input_rows.add(batch.num_rows()); - // Update row count - state.num_rows += batch.num_rows(); - // Push batch to output - state.batches.push(batch); - Ok(state) - }) - .await?; + // Decide if we spill or not + let batch_size = state.memory_counter.count_batch(&batch); + // Reserve memory for incoming batch + let fall_back = match state.reservation.try_grow(batch_size) { + // The build side grew past the configured size + Ok(()) => sort_merge_fallback.as_ref().is_some_and(|fallback| { + fallback + .max_build_size + .is_some_and(|max| state.reservation.size() > max) + }), + // Only the join's own exhausted reservation can be recovered from + // by sorting; anything else (including errors of the input) is + // reported as is. + Err(error) + if sort_merge_fallback.is_some() && is_resources_exhausted(&error) => + { + true + } + Err(error) => return Err(error), + }; + if let Some(fallback) = sort_merge_fallback.take_if(|_| fall_back) { + let BuildSideState { + batches, + reservation, + bounds_accumulators, + .. + } = state; + // Release what the collected batches reserved: the external sort + // accounts for what it keeps in memory itself. + drop(reservation); + let sorted = sort_build_side( + fallback, + schema, + batches, + Some(batch), + Some(left_stream), + bounds_accumulators.filter(|_| should_compute_dynamic_filters), + None, + ) + .await?; + return Ok(BuildSideOutcome::SortMerge(sorted)); + } + // Update metrics + state.metrics.build_mem_used.add(batch_size); + state.metrics.build_input_batches.add(1); + state.metrics.build_input_rows.add(batch.num_rows()); + // Update row count + state.num_rows += batch.num_rows(); + // Push batch to output + state.batches.push(batch); + } + drop(left_stream); // Extract fields from state let BuildSideState { batches, num_rows, metrics, - mut reservation, + reservation, bounds_accumulators, memory_counter: _, } = state; // Compute bounds - let mut bounds = match bounds_accumulators { + let bounds = match bounds_accumulators { Some(accumulators) if num_rows > 0 => { let bounds = accumulators .into_iter() @@ -2802,12 +2913,81 @@ async fn collect_left_input( _ => None, }; + let in_memory = build_in_memory( + &random_state, + &schema, + &batches, + num_rows, + &on_left, + &metrics, + reservation, + bounds.clone(), + with_visited_indices_bitmap, + with_null_aware_mark_state, + probe_threads_count, + should_compute_dynamic_filters, + &config, + null_equality, + null_aware, + &array_map_created_count, + ); + match in_memory { + Ok(data) => Ok(BuildSideOutcome::InMemory(Arc::new(data))), + Err(error) => { + // The batches fit, but the hash table (or a bitmap) on top of them + // did not. The reservation was released with the failed build. + let Some(fallback) = + sort_merge_fallback.filter(|_| is_resources_exhausted(&error)) + else { + return Err(error); + }; + let sorted = sort_build_side( + fallback, + schema, + batches, + None, + None, + None, + bounds.filter(|_| should_compute_dynamic_filters), + ) + .await?; + Ok(BuildSideOutcome::SortMerge(sorted)) + } + } +} + +/// Builds the hash table (or perfect-hash [`ArrayMap`]) and the bitmaps over +/// the fully collected build `batches`. +/// +/// `reservation` already covers `batches`; the structures built here grow it +/// further. On error the reservation is dropped, releasing everything. +#[expect(clippy::too_many_arguments)] +fn build_in_memory( + random_state: &RandomState, + schema: &SchemaRef, + batches: &[RecordBatch], + num_rows: usize, + on_left: &[PhysicalExprRef], + metrics: &BuildProbeJoinMetrics, + mut reservation: MemoryReservation, + mut bounds: Option, + with_visited_indices_bitmap: bool, + with_null_aware_mark_state: bool, + probe_threads_count: usize, + should_compute_dynamic_filters: bool, + config: &ConfigOptions, + null_equality: NullEquality, + null_aware: Option, + array_map_created_count: &Count, +) -> Result { + let is_phj_candidate = is_perfect_hash_join_candidate(on_left, schema)?; + let (join_hash_map, batch, left_values) = if let Some((array_map, batch, left_value)) = try_create_array_map( bounds.as_ref(), - &schema, - &batches, - &on_left, + schema, + batches, + on_left, &mut reservation, config.execution.perfect_hash_join_small_build_threshold, config.execution.perfect_hash_join_min_key_density, @@ -2823,7 +3003,7 @@ async fn collect_left_input( // Use `u32` indices for the JoinHashMap when num_rows ≤ u32::MAX, otherwise use the // `u64` indice variant // Arc is used instead of Box to allow sharing with SharedBuildAccumulator for hash map pushdown - let mut hashmap = new_join_hashmap(num_rows, &mut reservation, &metrics)?; + let mut hashmap = new_join_hashmap(num_rows, &mut reservation, metrics)?; let mut hashes_buffer = Vec::new(); let mut offset = 0; @@ -2835,11 +3015,11 @@ async fn collect_left_input( hashes_buffer.clear(); hashes_buffer.resize(batch.num_rows(), 0); update_hash( - &on_left, + on_left, batch, &mut *hashmap, offset, - &random_state, + random_state, &mut hashes_buffer, 0, true, @@ -2849,9 +3029,9 @@ async fn collect_left_input( } // Merge all batches into a single batch, so we can directly index into the arrays - let batch = concat_batches(&schema, batches_iter.clone())?; + let batch = concat_batches(schema, batches_iter.clone())?; - let left_values = evaluate_expressions_to_arrays(&on_left, &batch)?; + let left_values = evaluate_expressions_to_arrays(on_left, &batch)?; (Map::HashMap(hashmap), batch, left_values) }; @@ -2890,7 +3070,7 @@ async fn collect_left_input( ); // Scope-only NULL marking uses a HashMap (the primary join map may use // ArrayMap for full-key matches, but scope keys have arbitrary shape). - let mut scope_map = new_join_hashmap(num_rows, &mut reservation, &metrics)?; + let mut scope_map = new_join_hashmap(num_rows, &mut reservation, metrics)?; let mut hashes_buffer = vec![0; batch.num_rows()]; update_hash( @@ -2898,7 +3078,7 @@ async fn collect_left_input( &batch, &mut *scope_map, 0, - &random_state, + random_state, &mut hashes_buffer, 0, true, @@ -2930,9 +3110,9 @@ async fn collect_left_input( metrics.build_mem_used.add(retained_size); let null_rows = build_indices.len(); - let mut map = new_join_hashmap(null_rows, &mut reservation, &metrics)?; + let mut map = new_join_hashmap(null_rows, &mut reservation, metrics)?; let mut hashes_buffer = vec![0; null_rows]; - create_hashes(&scope_values, &random_state, &mut hashes_buffer)?; + create_hashes(&scope_values, random_state, &mut hashes_buffer)?; map.update_from_iter(Box::new(hashes_buffer.iter().enumerate().rev()), 0); Some(NullValueScopeMap { @@ -3055,8 +3235,8 @@ mod tests { }; use arrow::array::{ - Array, ArrayRef, AsArray, Date32Array, DictionaryArray, Int32Array, Int64Array, - StructArray, UInt32Array, UInt64Array, + Array, ArrayRef, AsArray, BooleanArray, Date32Array, DictionaryArray, Int32Array, + Int64Array, StructArray, UInt32Array, UInt64Array, }; use arrow::buffer::NullBuffer; use arrow::datatypes::{DataType, Field, Int32Type}; @@ -3067,6 +3247,7 @@ mod tests { exec_err, internal_err, }; use datafusion_execution::config::SessionConfig; + use datafusion_execution::disk_manager::{DiskManagerBuilder, DiskManagerMode}; use datafusion_execution::runtime_env::RuntimeEnvBuilder; use datafusion_expr::Operator; use datafusion_physical_expr::expressions::{BinaryExpr, Literal}; @@ -6841,6 +7022,10 @@ mod tests { for join_type in join_types { let runtime = RuntimeEnvBuilder::new() .with_memory_limit(100, 1.0) + // The join would otherwise finish as a sort-merge join + .with_disk_manager_builder( + DiskManagerBuilder::default().with_mode(DiskManagerMode::Disabled), + ) .build_arc()?; let session_config = SessionConfig::default().with_batch_size(50); let task_ctx = TaskContext::default() @@ -9208,4 +9393,410 @@ mod tests { assert!(join.set_dynamic_filter(df).is_err()); Ok(()) } + + /// Two sides of `batches` batches each, with duplicate keys, keys that + /// only exist on one side, and a non-key column to filter on. + fn sort_merge_fallback_inputs( + batches: usize, + ) -> (Arc, Arc) { + let rows_per_batch = 16; + let side = |modulus: i32, offset: i32, a: &str, b: &str, c: &str| { + let batches: Vec = (0..batches as i32) + .map(|batch| { + let start = batch * rows_per_batch; + let ids: Vec = (start..start + rows_per_batch).collect(); + // keys repeat within and across batches, and each side + // has keys the other side lacks + let keys: Vec = + ids.iter().map(|id| (id * 7) % modulus + offset).collect(); + let values: Vec = ids.iter().map(|id| (id * 13) % 17).collect(); + build_table_i32((a, &ids), (b, &keys), (c, &values)) + }) + .collect(); + let schema = batches[0].schema(); + TestMemoryExec::try_new_exec(&[batches], schema, None).unwrap() + as Arc + }; + // left keys are 0..47, right keys 5..58 + (side(47, 0, "a1", "b1", "c1"), side(53, 5, "a2", "b2", "c2")) + } + + /// `c1 < c2`, so it references both sides + fn sort_merge_fallback_filter( + left: &Arc, + right: &Arc, + ) -> JoinFilter { + let column_indices = vec![ + ColumnIndex { + index: 2, + side: JoinSide::Left, + }, + ColumnIndex { + index: 2, + side: JoinSide::Right, + }, + ]; + // Outer joins evaluate the filter over null-padded rows + let intermediate_schema = Schema::new(vec![ + left.schema().field(2).clone().with_nullable(true), + right.schema().field(2).clone().with_nullable(true), + ]); + let expression = Arc::new(BinaryExpr::new( + Arc::new(Column::new("c1", 0)), + Operator::Lt, + Arc::new(Column::new("c2", 1)), + )) as Arc; + JoinFilter::new(expression, column_indices, Arc::new(intermediate_schema)) + } + + /// A task context with `memory_limit` bytes and spilling enabled, sized + /// so a sort can still run (no pre-reserved merge memory) + fn sort_merge_fallback_task_ctx( + memory_limit: Option, + disk_manager_builder: DiskManagerBuilder, + ) -> Arc { + let mut runtime = + RuntimeEnvBuilder::new().with_disk_manager_builder(disk_manager_builder); + if let Some(memory_limit) = memory_limit { + runtime = runtime.with_memory_limit(memory_limit, 1.0); + } + let mut session_config = SessionConfig::default().with_batch_size(16); + session_config + .options_mut() + .execution + .sort_spill_reservation_bytes = 0; + Arc::new( + TaskContext::default() + .with_session_config(session_config) + .with_runtime(runtime.build_arc().unwrap()), + ) + } + + fn sorted_rows(batches: &[RecordBatch]) -> Vec { + batches_to_sort_string(batches) + .lines() + .map(str::to_string) + .collect() + } + + const ALL_JOIN_TYPES: [JoinType; 10] = [ + JoinType::Inner, + JoinType::Left, + JoinType::Right, + JoinType::Full, + JoinType::LeftSemi, + JoinType::LeftAnti, + JoinType::RightSemi, + JoinType::RightAnti, + JoinType::LeftMark, + JoinType::RightMark, + ]; + + /// Under memory pressure a partitioned hash join finishes as a sort-merge + /// join and produces the same rows as the in-memory join, for every join + /// type, with and without a join filter. + #[tokio::test] + async fn partitioned_join_sort_merge_fallback_matches_in_memory_join() -> Result<()> { + let (left, right) = sort_merge_fallback_inputs(32); + let on = vec![( + Arc::new(Column::new_with_schema("b1", &left.schema())?) as _, + Arc::new(Column::new_with_schema("b2", &right.schema())?) as _, + )]; + + for join_type in ALL_JOIN_TYPES { + for filtered in [false, true] { + let filter = filtered.then(|| sort_merge_fallback_filter(&left, &right)); + let join = || { + HashJoinExec::try_new( + Arc::clone(&left), + Arc::clone(&right), + on.clone(), + filter.clone(), + &join_type, + None, + PartitionMode::Partitioned, + NullEquality::NullEqualsNothing, + false, + ) + }; + + let in_memory = join()?; + let expected = common::collect(in_memory.execute( + 0, + sort_merge_fallback_task_ctx(None, DiskManagerBuilder::default()), + )?) + .await?; + assert_eq!( + in_memory + .metrics() + .unwrap() + .sum_by_name(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME) + .map(|v| v.as_usize()), + Some(0), + "{join_type} filtered={filtered}: no fallback without memory pressure" + ); + + let fallback = join()?; + let actual = common::collect(fallback.execute( + 0, + sort_merge_fallback_task_ctx( + Some(6 * 1024), + DiskManagerBuilder::default(), + ), + )?) + .await?; + let metrics = fallback.metrics().unwrap(); + assert_eq!( + metrics + .sum_by_name(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME) + .map(|v| v.as_usize()), + Some(1), + "{join_type} filtered={filtered}: the join must have fallen back" + ); + assert!( + metrics.spill_count().unwrap() > 0, + "{join_type} filtered={filtered}: the fallback sorts must have spilled" + ); + assert_eq!( + metrics.output_rows().unwrap(), + actual.iter().map(|b| b.num_rows()).sum::(), + "{join_type} filtered={filtered}: output rows are recorded once" + ); + + assert!( + !expected.is_empty(), + "{join_type} filtered={filtered}: the test data must produce rows" + ); + assert_eq!(fallback.schema(), actual[0].schema()); + assert_eq!( + sorted_rows(&actual), + sorted_rows(&expected), + "{join_type} filtered={filtered}" + ); + } + } + Ok(()) + } + + /// The fallback output is projected and limited like the hash join's + #[tokio::test] + async fn sort_merge_fallback_applies_projection_and_fetch() -> Result<()> { + let (left, right) = sort_merge_fallback_inputs(32); + let on = vec![( + Arc::new(Column::new_with_schema("b1", &left.schema())?) as _, + Arc::new(Column::new_with_schema("b2", &right.schema())?) as _, + )]; + let join = || { + HashJoinExecBuilder::new( + Arc::clone(&left), + Arc::clone(&right), + on.clone(), + JoinType::Inner, + ) + .with_partition_mode(PartitionMode::Partitioned) + .with_projection(Some(vec![4, 1])) + .with_fetch(Some(7)) + .build() + }; + + let in_memory = join()?; + let expected = common::collect(in_memory.execute( + 0, + sort_merge_fallback_task_ctx(None, DiskManagerBuilder::default()), + )?) + .await?; + let fallback = join()?; + let actual = common::collect(fallback.execute( + 0, + sort_merge_fallback_task_ctx(Some(6 * 1024), DiskManagerBuilder::default()), + )?) + .await?; + assert_eq!( + fallback + .metrics() + .unwrap() + .sum_by_name(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME) + .map(|v| v.as_usize()), + Some(1) + ); + + assert_eq!(actual[0].schema(), fallback.schema()); + assert_eq!(columns(&actual[0].schema()), vec!["b2", "b1"]); + // The fetched rows differ (a hash join keeps probe order, a sort-merge + // join key order), but their count and shape do not + assert_eq!(actual.iter().map(|b| b.num_rows()).sum::(), 7); + assert_eq!(expected.iter().map(|b| b.num_rows()).sum::(), 7); + let rows = sorted_rows(&actual); + // skip the table header and the closing border + for row in &rows[3..rows.len() - 1] { + // every row is an actual match: b2 == b1 + let cells: Vec<&str> = row.split('|').map(str::trim).collect(); + assert_eq!(cells[1], cells[2], "{row}"); + } + Ok(()) + } + + /// A falling-back partition has no hash table left to test membership + /// against, so the dynamic filter it pushes to the probe side must stay + /// permissive. + /// + /// This is the guard for [`PushdownStrategy::Unknown`]. Reporting `Empty` + /// instead would claim the partition holds no rows at all, and the probe + /// side would then discard every row routed to it, losing matches. The + /// pruning semantics of the two reports are covered by + /// `partitioned_unknown_membership_keeps_its_probe_rows` in `shared_bounds`. + #[tokio::test] + async fn sort_merge_fallback_keeps_the_dynamic_filter_permissive() -> Result<()> { + let (left, right) = sort_merge_fallback_inputs(32); + let on: JoinOn = vec![( + Arc::new(Column::new_with_schema("b1", &left.schema())?) as _, + Arc::new(Column::new_with_schema("b2", &right.schema())?) as _, + )]; + let probe_schema = right.schema(); + + let dynamic_filter = HashJoinExec::create_dynamic_filter(&on); + let join = HashJoinExecBuilder::new(left, right, on, JoinType::Inner) + .with_partition_mode(PartitionMode::Partitioned) + .build()? + .set_dynamic_filter(Arc::clone(&dynamic_filter))?; + + let batches = common::collect(join.execute( + 0, + sort_merge_fallback_task_ctx(Some(6 * 1024), DiskManagerBuilder::default()), + )?) + .await?; + assert_eq!( + join.metrics() + .unwrap() + .sum_by_name(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME) + .map(|v| v.as_usize()), + Some(1), + "the partition must have fallen back for this to test anything" + ); + assert!(!batches.is_empty(), "the join should have produced rows"); + + // Every probe row must survive the filter the fallback reported. + let probe = RecordBatch::try_new( + Arc::clone(&probe_schema), + vec![ + Arc::new(Int32Array::from(vec![0, 1, 2, 3])), + Arc::new(Int32Array::from(vec![5, 6, 7, 8])), + Arc::new(Int32Array::from(vec![0, 1, 2, 3])), + ], + )?; + let filter = dynamic_filter.current()?; + let kept = filter.evaluate(&probe)?.into_array(probe.num_rows())?; + let kept = kept + .as_any() + .downcast_ref::() + .expect("a filter evaluates to a BooleanArray"); + assert_eq!( + (0..kept.len()).filter(|i| kept.value(*i)).count(), + probe.num_rows(), + "a fallen-back partition must not prune probe rows, filter was {filter}" + ); + + Ok(()) + } + + /// Without disk the join fails as before + #[tokio::test] + async fn sort_merge_fallback_needs_disk() -> Result<()> { + let (left, right) = sort_merge_fallback_inputs(32); + let on = vec![( + Arc::new(Column::new_with_schema("b1", &left.schema())?) as _, + Arc::new(Column::new_with_schema("b2", &right.schema())?) as _, + )]; + let join = HashJoinExec::try_new( + left, + right, + on, + None, + &JoinType::Inner, + None, + PartitionMode::Partitioned, + NullEquality::NullEqualsNothing, + false, + )?; + + let without_disk = sort_merge_fallback_task_ctx( + Some(6 * 1024), + DiskManagerBuilder::default().with_mode(DiskManagerMode::Disabled), + ); + let err = common::collect(join.execute(0, without_disk)?) + .await + .unwrap_err(); + assert_contains!(err.to_string(), "Failed to allocate additional"); + assert_contains!(err.to_string(), "HashJoinInput[0]"); + assert_eq!( + join.metrics() + .unwrap() + .sum_by_name(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME) + .map(|v| v.as_usize()), + None + ); + Ok(()) + } + + /// `hash_join_max_build_size` triggers the fallback without any memory + /// limit, once the build side of the partition grows past it + #[tokio::test] + async fn sort_merge_fallback_on_max_build_size() -> Result<()> { + let (left, right) = sort_merge_fallback_inputs(32); + let on = vec![( + Arc::new(Column::new_with_schema("b1", &left.schema())?) as _, + Arc::new(Column::new_with_schema("b2", &right.schema())?) as _, + )]; + let join = || { + HashJoinExec::try_new( + Arc::clone(&left), + Arc::clone(&right), + on.clone(), + None, + &JoinType::Left, + None, + PartitionMode::Partitioned, + NullEquality::NullEqualsNothing, + false, + ) + }; + let task_ctx = |max_build_size: Option| { + let unlimited = + sort_merge_fallback_task_ctx(None, DiskManagerBuilder::default()); + let mut session_config = unlimited.session_config().clone(); + session_config + .options_mut() + .execution + .hash_join_max_build_size = max_build_size; + Arc::new( + TaskContext::default() + .with_session_config(session_config) + .with_runtime(unlimited.runtime_env()), + ) + }; + + let in_memory = join()?; + let expected = common::collect(in_memory.execute(0, task_ctx(None))?).await?; + assert_eq!( + in_memory + .metrics() + .unwrap() + .sum_by_name(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME) + .map(|v| v.as_usize()), + Some(0) + ); + + let fallback = join()?; + let actual = common::collect(fallback.execute(0, task_ctx(Some(1024)))?).await?; + assert_eq!( + fallback + .metrics() + .unwrap() + .sum_by_name(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME) + .map(|v| v.as_usize()), + Some(1) + ); + assert_eq!(sorted_rows(&actual), sorted_rows(&expected)); + Ok(()) + } } diff --git a/datafusion/physical-plan/src/joins/hash_join/mod.rs b/datafusion/physical-plan/src/joins/hash_join/mod.rs index 7c1b0d76c2f73..7010e6625fe06 100644 --- a/datafusion/physical-plan/src/joins/hash_join/mod.rs +++ b/datafusion/physical-plan/src/joins/hash_join/mod.rs @@ -25,4 +25,5 @@ mod inlist_builder; mod partitioned_hash_eval; mod probe_completion; mod shared_bounds; +mod sort_merge_fallback; mod stream; diff --git a/datafusion/physical-plan/src/joins/hash_join/shared_bounds.rs b/datafusion/physical-plan/src/joins/hash_join/shared_bounds.rs index 62087c14c5179..569f93148224b 100644 --- a/datafusion/physical-plan/src/joins/hash_join/shared_bounds.rs +++ b/datafusion/physical-plan/src/joins/hash_join/shared_bounds.rs @@ -142,8 +142,11 @@ fn create_membership_predicate( hash_map, "hash_lookup".to_string(), )) as Arc)), - // Empty partition - should not create a filter for this - PushdownStrategy::Empty => Ok(None), + // No membership term for either of the remaining variants. They look + // alike only here: the pruning difference between them is applied by + // the caller, which gives an `Empty` partition a `false` branch without + // consulting this function at all. + PushdownStrategy::Empty | PushdownStrategy::Unknown => Ok(None), } } @@ -273,14 +276,42 @@ pub(crate) struct SharedBuildAccumulator { } /// Strategy for filter pushdown (decided at collection time) +/// +/// The first two variants carry a structure a probe row can be tested against. +/// The last two carry none, for opposite reasons, and that difference decides +/// whether the partition's probe rows may be discarded: see [`Self::Empty`] and +/// [`Self::Unknown`]. #[derive(Clone)] pub(crate) enum PushdownStrategy { /// Use InList for small build sides (< 128MB) InList(ArrayRef), /// Use map lookup for large build sides Map(Arc), - /// There was no data in this partition, do not build a dynamic filter for it + /// This partition's build side is known to hold no rows at all, so nothing + /// routed to it can ever match and every such probe row may be discarded. + /// + /// This is a statement about the *data*, and the pushdown machinery acts on + /// it: [`SharedBuildAccumulator::build_partitioned_filter`] gives such a + /// partition a `false` branch in the routing `CASE`, and collapses the whole + /// filter to `false` when every partition reports it. Empty, + /// This partition's build side may hold rows, but no structure exists to + /// test membership against, so no probe row may be discarded on membership + /// grounds. + /// + /// This is a statement about our *knowledge*, not about the data, and it is + /// why the variant cannot be folded into [`Self::Empty`]: reporting `Empty` + /// for a partition that actually holds rows would push `false` for it and + /// silently drop every match it owns. Only the bounds, which remain exact, + /// are pushed down. + /// + /// Reported when the sort-merge fallback of a memory-pressured hash join + /// sorted the build side instead of hashing it (see the + /// `sort_merge_fallback` module). The existing `CanceledUnknown` partition + /// status is also permissive but not a substitute: it additionally assumes + /// the build side may contain NULL keys and discards the bounds, because a + /// canceled partition never reported either. + Unknown, } /// Build-side data reported by a single partition @@ -655,8 +686,13 @@ impl SharedBuildAccumulator { } /// Builds one routed probe-side filter from finalized partitioned build data. - /// Empty partitions reject their routed rows, while canceled partitions stay - /// permissive because their build contents are unknown. + /// + /// Only a partition that reported [`PushdownStrategy::Empty`] rejects the + /// probe rows routed to it, because only that report proves the partition + /// holds nothing to match. Every other partition keeps whatever it can + /// filter on and admits the rest: a canceled one is fully permissive since + /// it reported nothing at all, and a partition whose membership is + /// [`PushdownStrategy::Unknown`] is filtered by its bounds alone. fn build_partitioned_filter(&self, partitions: Vec) -> Result<()> { let mut partition_filters = Vec::with_capacity(partitions.len()); let mut real_partition_ids = Vec::new(); @@ -1126,6 +1162,85 @@ mod tests { assert_literal_bool(&expr, true); } + /// How many of `keys` the accumulator's finalized filter keeps. + fn kept_by_filter(acc: &SharedBuildAccumulator, keys: &[i32]) -> Result { + let expr = current_expr(acc); + let batch = RecordBatch::try_new( + test_probe_schema(), + vec![Arc::new(Int32Array::from(keys.to_vec()))], + )?; + let result = expr.evaluate(&batch)?.into_array(batch.num_rows())?; + let result = result + .as_any() + .downcast_ref::() + .expect("dynamic filter should evaluate to BooleanArray"); + Ok((0..result.len()).filter(|i| result.value(*i)).count()) + } + + /// A partition whose membership is unknown must keep the probe rows routed + /// to it, because it may hold rows that match them. + /// + /// Guards the distinction from [`PushdownStrategy::Empty`]: if this report + /// were folded into `Empty`, the other partition's `InList` would become the + /// whole filter and every probe row outside that list would be discarded, + /// silently dropping the matches the unknown partition owns. The `Empty` + /// case below is the control showing exactly that pruning. + #[test] + fn partitioned_unknown_membership_keeps_its_probe_rows() -> Result<()> { + let keys: Vec = (0..64).collect(); + + let unknown = make_partitioned_expr_accumulator_for_test(2); + unknown.build_filter(FinalizeInput::Partitioned(vec![ + reported(PushdownStrategy::Unknown, no_bounds()), + reported(in_list(&[2]), no_bounds()), + ]))?; + let kept_with_unknown = kept_by_filter(&unknown, &keys)?; + + let empty = make_partitioned_expr_accumulator_for_test(2); + empty.build_filter(FinalizeInput::Partitioned(vec![ + reported(PushdownStrategy::Empty, no_bounds()), + reported(in_list(&[2]), no_bounds()), + ]))?; + let kept_with_empty = kept_by_filter(&empty, &keys)?; + + // An empty partition proves nothing routed to it can match, so only the + // other partition's single in-list value survives. + assert_eq!( + kept_with_empty, 1, + "an empty partition should prune every probe row routed to it" + ); + // An unknown partition proves nothing, so its share of the keys stays. + assert!( + kept_with_unknown > kept_with_empty, + "unknown membership must keep its routed probe rows, kept {kept_with_unknown} of {} \ + but an empty partition keeps {kept_with_empty}", + keys.len() + ); + + Ok(()) + } + + /// Unknown membership still narrows the probe side by the bounds it did + /// report, so it is not simply permissive. + #[test] + fn partitioned_unknown_membership_still_applies_its_bounds() -> Result<()> { + let acc = make_partitioned_expr_accumulator_for_test(2); + acc.build_filter(FinalizeInput::Partitioned(vec![ + reported(PushdownStrategy::Unknown, bounds(10, 20)), + reported(in_list(&[2]), no_bounds()), + ]))?; + + // Every key routed to partition 0 outside 10..=20 is rejected by the + // bounds, so far from all 64 keys survive. + let kept = kept_by_filter(&acc, &(0..64).collect::>())?; + assert!( + kept < 32, + "bounds reported alongside unknown membership must still prune, kept {kept}" + ); + + Ok(()) + } + #[test] fn partitioned_one_real_partition_with_rest_empty_skips_case() { let acc = make_partitioned_expr_accumulator_for_test(3); diff --git a/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs b/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs new file mode 100644 index 0000000000000..5e60712604195 --- /dev/null +++ b/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs @@ -0,0 +1,361 @@ +// 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 fallback of [`HashJoinExec`] under memory pressure. +//! +//! A hash join materializes its whole build side in memory. When the memory +//! pool refuses that reservation the join cannot continue as a hash join, but +//! it can still finish: both inputs are sorted on the join keys with an +//! external (spilling) sort, and the join is completed by the sort-merge join +//! streams, which only keep one key group of the buffered side in memory. +//! +//! The fallback is decided per output partition at runtime, after the hash +//! join's own `try_grow` failed (or, with `hash_join_max_build_size` set, +//! once the build side grew past that size), so joins that fit in memory +//! never pay for it. +//! It is restricted to [`PartitionMode::Partitioned`] joins: there every +//! partition owns a disjoint slice of the key space on both sides, so each +//! partition can sort and merge its own inputs independently of its siblings +//! and every join type stays correct. +//! +//! # Limitations +//! +//! The fallback replaces a build side that must fit entirely in memory with a +//! much smaller requirement, but not with none. Each sort pre-reserves +//! `sort_spill_reservation_bytes` so that its merge phase can always run, and a +//! falling-back partition runs two sorts, so one join can hold several such +//! reservations at once on top of the sorts' working set. A budget that does not +//! cover that floor still fails, and it fails inside the sort rather than in the +//! hash join. With the 10 MB default reservation and two partitions, the +//! reservations alone come to tens of megabytes, so a budget below that cannot +//! support the fallback however large the data is. Sizing the fallback's sorts +//! from the budget actually available, instead of inheriting the global default, +//! is left as follow-up work. +//! +//! Partitions also share one pool without coordinating, and a partition that +//! falls back does not make a sibling that still fits release its hash table. +//! That narrows the usable range further. Whether one or both partitions fall +//! back is not by itself decisive: a single falling-back partition completes +//! when the budget covers its sorts. +//! +//! [`HashJoinExec`]: super::HashJoinExec +//! [`PartitionMode::Partitioned`]: crate::joins::PartitionMode::Partitioned + +use std::sync::Arc; + +use crate::SendableRecordBatchStream; +use crate::expressions::PhysicalSortExpr; +use crate::joins::hash_join::exec::{CollectLeftAccumulator, JoinLeftData}; +use crate::joins::hash_join::shared_bounds::PartitionBounds; +use crate::joins::sort_merge_join::{SortMergeJoinInputs, sort_merge_join_stream}; +use crate::joins::utils::JoinFilter; +use crate::limit::LimitStream; +use crate::metrics::{BaselineMetrics, Count, ExecutionPlanMetricsSet, SpillMetrics}; +use crate::sorts::sort::ExternalSorter; +use crate::stream::RecordBatchStreamAdapter; + +use arrow::array::Array; +use arrow::compute::SortOptions; +use arrow::datatypes::SchemaRef; +use arrow::record_batch::RecordBatch; +use datafusion_common::{DataFusionError, JoinType, NullEquality, Result, internal_err}; +use datafusion_execution::TaskContext; +use datafusion_physical_expr::PhysicalExprRef; +use datafusion_physical_expr_common::sort_expr::LexOrdering; +use datafusion_physical_expr_common::utils::evaluate_expressions_to_arrays; +use futures::StreamExt; +use parking_lot::Mutex; + +/// Everything the fallback needs from the join, captured once per partition +/// in `HashJoinExec::execute` and shared by the build future and the stream. +#[derive(Clone)] +pub(super) struct SortMergeFallbackContext { + pub(super) context: Arc, + pub(super) partition: usize, + /// The join's metrics; the sort-merge join stream registers its metrics + /// (its output rows included) here. + pub(super) metrics: ExecutionPlanMetricsSet, + /// Spill metrics of the join; the fallback's sorts report their spills here. + pub(super) spill_metrics: SpillMetrics, + /// Incremented once when a partition falls back. + pub(super) fallback_count: Count, + /// Join keys of the left (build) side. + pub(super) on_left: Vec, + /// Join keys of the right (probe) side. + pub(super) on_right: Vec, + /// Sort options both sides are sorted with, one per join key. + pub(super) sort_options: Vec, + pub(super) join_type: JoinType, + pub(super) filter: Option, + pub(super) null_equality: NullEquality, + /// The join schema before any projection: the schema the sort-merge join + /// streams produce. + pub(super) join_schema: SchemaRef, + /// The join's output schema, after `projection`. + pub(super) output_schema: SchemaRef, + pub(super) projection: Option>, + pub(super) fetch: Option, + /// Build side size, in reserved bytes, past which the partition falls + /// back even though the memory pool would allow more; `None` for no limit + /// (`datafusion.execution.hash_join_max_build_size`) + pub(super) max_build_size: Option, +} + +impl SortMergeFallbackContext { + fn sort_ordering(&self, on: &[PhysicalExprRef]) -> Result { + let exprs = on + .iter() + .zip(&self.sort_options) + .map(|(expr, options)| PhysicalSortExpr::new(Arc::clone(expr), *options)); + LexOrdering::new(exprs) + .ok_or_else(|| DataFusionError::Internal("join without keys".to_string())) + } +} + +/// Result of collecting the build side of one partition. +pub(super) enum BuildSideOutcome { + /// The build side fits in memory and is ready to be probed. + InMemory(Arc), + /// The build side did not fit in memory and was sorted instead; the + /// partition finishes as a sort-merge join. + SortMerge(SortedBuildSide), +} + +/// The sorted build side of a partition that fell back to a sort-merge join. +pub(super) struct SortedBuildSide { + /// The build side, sorted on the join keys. Taken by the single stream + /// of the partition. + stream: Mutex>, + /// Min/max of the join keys, reported to the dynamic filter when one is + /// pushed down. + pub(super) bounds: Option, + /// Whether any build-side join key is NULL. + pub(super) keys_have_null: bool, +} + +impl SortedBuildSide { + pub(super) fn take_stream(&self) -> Result { + match self.stream.lock().take() { + Some(stream) => Ok(stream), + None => internal_err!("sorted build side was already consumed"), + } + } +} + +/// Whether `error` is an exhausted memory pool that the fallback can recover from. +pub(super) fn is_resources_exhausted(error: &DataFusionError) -> bool { + matches!(error.find_root(), DataFusionError::ResourcesExhausted(_)) +} + +/// Sorts `buffered` followed by the rest of `input` on `ordering` with an +/// external sort, calling `observe` for every batch on its way in. +async fn sort_batches( + ctx: &SortMergeFallbackContext, + schema: SchemaRef, + ordering: LexOrdering, + buffered: Vec, + mut input: Option, + mut observe: impl FnMut(&RecordBatch) -> Result<()>, +) -> Result { + let session_config = ctx.context.session_config(); + let execution_options = &session_config.options().execution; + // The sorter's own metrics set: its baseline metrics would otherwise count + // the sorted rows as output rows of the join. Its spills are copied to the + // join's spill metrics below. + let sorter_metrics = ExecutionPlanMetricsSet::new(); + let mut sorter = ExternalSorter::new( + ctx.partition, + schema, + ordering, + session_config.batch_size(), + execution_options.sort_spill_reservation_bytes, + execution_options.sort_in_place_threshold_bytes, + session_config.spill_compression(), + &sorter_metrics, + ctx.context.runtime_env(), + )?; + + for batch in buffered { + insert_batch(&mut sorter, batch, &mut observe).await?; + } + if let Some(input) = input.as_mut() { + while let Some(batch) = input.next().await { + let batch = match batch { + Ok(batch) => batch, + Err(error) => { + sorter.abort_in_progress_spill().await; + return Err(error); + } + }; + insert_batch(&mut sorter, batch, &mut observe).await?; + } + } + drop(input); + + let sorted = sorter.sort().await?; + + let spills = sorter.spill_metrics(); + ctx.spill_metrics + .spill_file_count + .add(spills.spill_file_count.value()); + ctx.spill_metrics + .spilled_bytes + .add(spills.spilled_bytes.value()); + ctx.spill_metrics + .spilled_rows + .add(spills.spilled_rows.value()); + + Ok(sorted) +} + +/// Feeds one batch to `sorter`, abandoning its in-progress spill on error. +async fn insert_batch( + sorter: &mut ExternalSorter, + batch: RecordBatch, + observe: &mut impl FnMut(&RecordBatch) -> Result<()>, +) -> Result<()> { + observe(&batch)?; + if let Err(error) = sorter.insert_batch(batch).await { + sorter.abort_in_progress_spill().await; + return Err(error); + } + Ok(()) +} + +/// Sorts the build side of a partition after its in-memory collection ran +/// out of memory. +/// +/// `batches` are the build batches collected so far, `pending` the batch whose +/// reservation failed and `rest` the not yet consumed remainder of the build +/// input (both `None` when the input was fully consumed and the hash table +/// itself did not fit). The caller has already released the reservation held +/// for `batches`; the external sort reserves what it keeps in memory itself. +/// +/// The join key bounds needed by a pushed-down dynamic filter are computed +/// with `bounds_accumulators` over every batch (min/max are idempotent, so +/// re-feeding the already accumulated `batches` is harmless), unless the +/// caller already has the final `bounds`. +pub(super) async fn sort_build_side( + ctx: SortMergeFallbackContext, + schema: SchemaRef, + batches: Vec, + pending: Option, + rest: Option, + mut bounds_accumulators: Option>, + bounds: Option, +) -> Result { + ctx.fallback_count.add(1); + + let ordering = ctx.sort_ordering(&ctx.on_left)?; + let mut buffered = batches; + buffered.extend(pending); + + let on_left = ctx.on_left.clone(); + let mut keys_have_null = false; + let mut num_rows = 0; + let observe = |batch: &RecordBatch| -> Result<()> { + num_rows += batch.num_rows(); + if let Some(accumulators) = bounds_accumulators.as_mut() { + for accumulator in accumulators { + accumulator.update_batch(batch)?; + } + } + if !keys_have_null { + keys_have_null = evaluate_expressions_to_arrays(&on_left, batch)? + .iter() + .any(|array| array.logical_null_count() > 0); + } + Ok(()) + }; + + let stream = sort_batches(&ctx, schema, ordering, buffered, rest, observe) + .await + .map_err(|e| { + e.context("HashJoinExec sort-merge fallback: sorting the build side") + })?; + + let bounds = match bounds_accumulators { + Some(accumulators) if num_rows > 0 => Some(PartitionBounds::new( + accumulators + .into_iter() + .map(CollectLeftAccumulator::evaluate) + .collect::>>()?, + )), + Some(_) => None, + None => bounds, + }; + + Ok(SortedBuildSide { + stream: Mutex::new(Some(stream)), + bounds, + keys_have_null, + }) +} + +/// Sorts the probe side and joins it with the sorted build side as a +/// sort-merge join, returning the join's output stream (projected and limited +/// like the hash join's own output would be). +pub(super) async fn run_sort_merge_fallback( + ctx: SortMergeFallbackContext, + build: SendableRecordBatchStream, + probe: SendableRecordBatchStream, +) -> Result { + let ordering = ctx.sort_ordering(&ctx.on_right)?; + let probe = sort_batches(&ctx, probe.schema(), ordering, vec![], Some(probe), |_| { + Ok(()) + }) + .await + .map_err(|e| e.context("HashJoinExec sort-merge fallback: sorting the probe side"))?; + + let joined = sort_merge_join_stream( + SortMergeJoinInputs { + schema: Arc::clone(&ctx.join_schema), + sort_options: ctx.sort_options.clone(), + null_equality: ctx.null_equality, + left: build, + right: probe, + on_left: ctx.on_left.clone(), + on_right: ctx.on_right.clone(), + filter: ctx.filter.clone(), + join_type: ctx.join_type, + partition: ctx.partition, + }, + &ctx.metrics, + &ctx.context, + )?; + + let output: SendableRecordBatchStream = match ctx.projection.clone() { + Some(projection) => Box::pin(RecordBatchStreamAdapter::new( + Arc::clone(&ctx.output_schema), + joined.map(move |batch| Ok(batch?.project(&projection)?)), + )), + None => joined, + }; + + // The limit's own baseline metrics would count the output a second time. + let output = match ctx.fetch { + Some(_) => Box::pin(LimitStream::new( + output, + 0, + ctx.fetch, + BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), ctx.partition), + )), + None => output, + }; + + Ok(output) +} diff --git a/datafusion/physical-plan/src/joins/hash_join/stream.rs b/datafusion/physical-plan/src/joins/hash_join/stream.rs index 625661188d81c..9ed10aef86c7d 100644 --- a/datafusion/physical-plan/src/joins/hash_join/stream.rs +++ b/datafusion/physical-plan/src/joins/hash_join/stream.rs @@ -30,7 +30,10 @@ use crate::joins::PartitionMode; use crate::joins::hash_join::exec::{JoinLeftData, NullAwareMode}; use crate::joins::hash_join::probe_completion::ProbeSideSummary; use crate::joins::hash_join::shared_bounds::{ - PartitionBounds, PartitionBuildData, SharedBuildAccumulator, + PartitionBounds, PartitionBuildData, PushdownStrategy, SharedBuildAccumulator, +}; +use crate::joins::hash_join::sort_merge_fallback::{ + BuildSideOutcome, SortMergeFallbackContext, run_sort_merge_fallback, }; use crate::joins::utils::{OnceFut, equal_rows_arr, matchable_join_keys}; use crate::stream::EmptyRecordBatchStream; @@ -57,7 +60,8 @@ use datafusion_physical_expr::PhysicalExprRef; use datafusion_common::hash_utils::RandomState; use datafusion_physical_expr_common::utils::evaluate_expressions_to_arrays; -use futures::{Stream, StreamExt, ready}; +use futures::future::BoxFuture; +use futures::{FutureExt, Stream, StreamExt, ready}; /// Represents build-side of hash join. pub(super) enum BuildSide { @@ -65,12 +69,23 @@ pub(super) enum BuildSide { Initial(BuildSideInitialState), /// Indicates that build-side data has been collected Ready(BuildSideReadyState), + /// Indicates that the build side did not fit in memory and the partition + /// is finishing as a sort-merge join (see the `sort_merge_fallback` module) + SortMergeFallback(SortMergeFallbackState), } /// Container for BuildSide::Initial related data pub(super) struct BuildSideInitialState { /// Future for building hash table from build-side input - pub(super) left_fut: OnceFut, + pub(super) left_fut: OnceFut, +} + +/// Progress of the sort-merge fallback of one partition +pub(super) enum SortMergeFallbackState { + /// Sorting the probe side and setting up the sort-merge join + Preparing(BoxFuture<'static, Result>), + /// Forwarding the sort-merge join's output + Running(SendableRecordBatchStream), } /// Container for BuildSide::Ready related data @@ -116,7 +131,8 @@ impl BuildSide { /// /// WaitBuildSide /// │ -/// ▼ +/// ├───────────────────────────► SortMergeFallback ─┐ +/// ▼ ▼ /// ┌─► FetchProbeBatch ───► ExhaustedProbeSide ──────────► Completed /// │ │ │ ▲ /// │ ▼ ▼ │ @@ -124,6 +140,13 @@ impl BuildSide { /// └──────────┘ /// ``` /// +/// The `SortMergeFallback` branch is taken when the build side did not fit in +/// memory and the join can fall back (see the `sort_merge_fallback` module): that +/// state forwards the output of a sort-merge join over both sorted inputs, and +/// no probe-side state is ever entered. Both branches pass through +/// `WaitPartitionBoundsReport` first when a dynamic filter is pushed down, +/// which is omitted above. +/// /// `ExhaustedProbeSide` moves to `EmitUnmatchedBuildRows` only for join types /// that emit build-side rows after the probe side is exhausted (see /// [`need_produce_result_in_final`]), and only in the partition that finished @@ -145,6 +168,9 @@ pub(super) enum HashJoinStreamState { /// partition, and this stream is emitting the final build-side rows /// (unmatched rows, or matched rows for `LeftSemi`) in chunks EmitUnmatchedBuildRows(EmitUnmatchedBuildRowsState), + /// Indicates that the build side did not fit in memory and the output is + /// produced by a sort-merge join over both sorted inputs + SortMergeFallback, /// Indicates that HashJoinStream execution is completed Completed, } @@ -386,6 +412,9 @@ pub(super) struct HashJoinStream { output_buffer: LimitedBatchCoalescer, /// Null-aware (`NOT IN`) semantics of this join, if any null_aware: Option, + /// What this partition needs to finish as a sort-merge join when its + /// build side does not fit in memory; `None` when it cannot fall back + sort_merge_fallback: Option, } impl RecordBatchStream for HashJoinStream { @@ -537,6 +566,7 @@ impl HashJoinStream { mode: PartitionMode, null_aware: Option, fetch: Option, + sort_merge_fallback: Option, ) -> Self { // Create output buffer with coalescing and optional fetch limit. let output_buffer = @@ -569,6 +599,7 @@ impl HashJoinStream { mode, output_buffer, null_aware, + sort_merge_fallback, } } @@ -682,6 +713,9 @@ impl HashJoinStream { HashJoinStreamState::EmitUnmatchedBuildRows(_) => { handle_state!(self.emit_unmatched_build_rows()) } + HashJoinStreamState::SortMergeFallback => { + self.poll_sort_merge_fallback(cx) + } HashJoinStreamState::Completed if !self.output_buffer.is_empty() => { // Flush any remaining buffered data self.output_buffer.finish()?; @@ -707,9 +741,18 @@ impl HashJoinStream { cx: &mut std::task::Context<'_>, ) -> Poll>>> { ready!(self.build_report.poll_delivery(cx))?; - let build_side = self.build_side.try_as_ready()?; - self.state = - Self::state_after_build_ready(self.join_type, build_side.left_data.as_ref()); + self.state = match &self.build_side { + BuildSide::Ready(build_side) => Self::state_after_build_ready( + self.join_type, + build_side.left_data.as_ref(), + ), + BuildSide::SortMergeFallback(_) => HashJoinStreamState::SortMergeFallback, + BuildSide::Initial(_) => { + return Poll::Ready(internal_err!( + "Expected build side to be collected before its report" + )); + } + }; Poll::Ready(Ok(StatefulStreamResult::Continue)) } @@ -722,7 +765,7 @@ impl HashJoinStream { ) -> Poll>>> { let build_timer = self.join_metrics.build_time.timer(); // build hash table from left (build) side, if not yet done - let left_data = ready!( + let outcome = ready!( self.build_side .try_as_initial_mut()? .left_fut @@ -730,16 +773,117 @@ impl HashJoinStream { )?; build_timer.done(); - // Note: For null-aware anti join, we need to check the probe side (right) for NULLs, - // not the build side (left). The probe-side NULL check happens during process_probe_batch. - // The probe_side_has_null flag will be set there if any probe batch contains NULL. + match outcome.as_ref() { + BuildSideOutcome::InMemory(left_data) => { + let left_data = Arc::clone(left_data); + // Note: For null-aware anti join, we need to check the probe side (right) for NULLs, + // not the build side (left). The probe-side NULL check happens during process_probe_batch. + // The probe_side_has_null flag will be set there if any probe batch contains NULL. - self.state = self.transition_after_build_collected(&left_data); + self.state = self.transition_after_build_collected(&left_data); - self.build_side = BuildSide::Ready(BuildSideReadyState { left_data }); + self.build_side = BuildSide::Ready(BuildSideReadyState { left_data }); + } + BuildSideOutcome::SortMerge(sorted_build) => { + let Some(fallback) = self.sort_merge_fallback.clone() else { + return Poll::Ready(internal_err!( + "build side fell back to a sort-merge join, but the stream cannot" + )); + }; + let build = sorted_build.take_stream()?; + // The sort-merge join consumes the probe side from here on. + let right_schema = self.right.schema(); + let probe = std::mem::replace( + &mut self.right, + Box::pin(EmptyRecordBatchStream::new(right_schema)), + ); + self.state = self.transition_after_fallback( + sorted_build.bounds.clone(), + sorted_build.keys_have_null, + ); + self.build_side = + BuildSide::SortMergeFallback(SortMergeFallbackState::Preparing( + run_sort_merge_fallback(fallback, build, probe).boxed(), + )); + } + } Poll::Ready(Ok(StatefulStreamResult::Continue)) } + /// Forwards the output of the sort-merge join this partition fell back to, + /// first waiting for the probe side to be sorted. + fn poll_sort_merge_fallback( + &mut self, + cx: &mut std::task::Context<'_>, + ) -> Poll>> { + let BuildSide::SortMergeFallback(fallback) = &mut self.build_side else { + return Poll::Ready(Some(internal_err!( + "Expected build side in sort-merge fallback state" + ))); + }; + loop { + match fallback { + SortMergeFallbackState::Preparing(preparing) => { + match ready!(preparing.poll_unpin(cx)) { + Ok(stream) => { + *fallback = SortMergeFallbackState::Running(stream); + } + Err(e) => { + self.state = HashJoinStreamState::Completed; + return Poll::Ready(Some(Err(e))); + } + } + } + SortMergeFallbackState::Running(stream) => { + // The sort-merge join stream records its output in the + // join's metrics itself. + let poll = stream.poll_next_unpin(cx); + if matches!(poll, Poll::Ready(None)) { + self.state = HashJoinStreamState::Completed; + } + return poll; + } + } + } + } + + /// Transitions state after the build side was sorted instead of hashed, + /// reporting its bounds to the accumulator when one is present so sibling + /// partitions (and this one) can start probing. + fn transition_after_fallback( + &mut self, + bounds: Option, + keys_have_null: bool, + ) -> HashJoinStreamState { + if !self.build_report.has_accumulator() { + return HashJoinStreamState::SortMergeFallback; + } + + // No hash table exists to push down; the bounds still narrow the + // probe-side scan. + let pushdown = PushdownStrategy::Unknown; + let bounds = bounds.unwrap_or_else(|| PartitionBounds::new(vec![])); + let build_data = match self.mode { + PartitionMode::Partitioned => PartitionBuildData::Partitioned { + partition_id: self.partition, + pushdown, + bounds, + keys_have_null, + }, + PartitionMode::CollectLeft => PartitionBuildData::CollectLeft { + pushdown, + bounds, + keys_have_null, + }, + PartitionMode::Auto => unreachable!( + "PartitionMode::Auto should not be present at execution time. This is a bug in DataFusion, please report it!" + ), + }; + + self.build_report.schedule(build_data); + HashJoinStreamState::WaitPartitionBoundsReport + } + /// Fetches next batch from probe-side /// /// If non-empty batch has been fetched, updates state to `ProcessProbeBatchState`, 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 911eca0a97928..be12ef8349f4a 100644 --- a/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs +++ b/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs @@ -583,82 +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(); - let reservation = MemoryConsumer::new(format!("SMJStream[{partition}]")) - .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); @@ -945,3 +890,110 @@ 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(); + let reservation = MemoryConsumer::new(format!("SMJStream[{partition}]")) + .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(), + ) + } +} diff --git a/datafusion/physical-plan/src/joins/sort_merge_join/filter.rs b/datafusion/physical-plan/src/joins/sort_merge_join/filter.rs index 306a154666fae..8b6d50f08749f 100644 --- a/datafusion/physical-plan/src/joins/sort_merge_join/filter.rs +++ b/datafusion/physical-plan/src/joins/sort_merge_join/filter.rs @@ -153,27 +153,21 @@ pub fn get_filter_columns( 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 + let Some(f) = join_filter else { + return vec![]; + }; + + // The filter's intermediate schema lists its columns in `column_indices` + // order, which need not put every left column before every right one + // (a filter swapped along with the join's inputs does the opposite). + f.column_indices() + .iter() + .filter_map(|col_index| match col_index.side { + JoinSide::Left => Some(Arc::clone(&left_columns[col_index.index])), + JoinSide::Right => Some(Arc::clone(&right_columns[col_index.index])), + JoinSide::None => None, + }) + .collect() } /// Determines if current index is the last occurrence of a row diff --git a/datafusion/physical-plan/src/joins/sort_merge_join/mod.rs b/datafusion/physical-plan/src/joins/sort_merge_join/mod.rs index 2fdb0924e723d..d21b43b7dd411 100644 --- a/datafusion/physical-plan/src/joins/sort_merge_join/mod.rs +++ b/datafusion/physical-plan/src/joins/sort_merge_join/mod.rs @@ -18,6 +18,7 @@ //! Sort Merge Join Execution Plan Operator pub use exec::SortMergeJoinExec; +pub(crate) use exec::{SortMergeJoinInputs, sort_merge_join_stream}; pub(crate) mod bitwise_stream; mod exec; diff --git a/datafusion/physical-plan/src/joins/sort_merge_join/tests.rs b/datafusion/physical-plan/src/joins/sort_merge_join/tests.rs index 4d817df738e9c..5064152bd7f56 100644 --- a/datafusion/physical-plan/src/joins/sort_merge_join/tests.rs +++ b/datafusion/physical-plan/src/joins/sort_merge_join/tests.rs @@ -938,6 +938,66 @@ async fn join_left_different_columns_count_with_filter() -> Result<()> { Ok(()) } +/// A filter whose intermediate schema lists a right column before a left one +/// (the layout `JoinFilter::swap` produces when a join's inputs are swapped) +#[tokio::test] +async fn join_left_with_filter_columns_right_before_left() -> Result<()> { + // select * + // from t2 + // left join t1 on t2.b1 = t1.b1 and t2.a2 > t1.a1 + + let left = build_table_two_cols( + ("a2", &vec![10, 20, 30]), + ("b1", &vec![4, 5, 6]), // 6 does not exist on the right + ); + + let right = build_table( + ("a1", &vec![1, 21, 3]), // 20(t2.a2) > 1(t1.a1) + ("b1", &vec![4, 5, 7]), + ("c1", &vec![7, 8, 9]), + ); + + let on = vec![( + Arc::new(Column::new_with_schema("b1", &left.schema())?) as _, + Arc::new(Column::new_with_schema("b1", &right.schema())?) as _, + )]; + + let filter = JoinFilter::new( + Arc::new(BinaryExpr::new( + Arc::new(Column::new("a2", 1)), + Operator::Gt, + Arc::new(Column::new("a1", 0)), + )), + vec![ + ColumnIndex { + index: 0, + side: JoinSide::Right, + }, + ColumnIndex { + index: 0, + side: JoinSide::Left, + }, + ], + Arc::new(Schema::new(vec![ + Field::new("a1", DataType::Int32, true), + Field::new("a2", DataType::Int32, true), + ])), + ); + + let (_, batches) = join_collect_with_filter(left, right, on, filter, Left).await?; + + assert_snapshot!(batches_to_string(&batches), @r" + +----+----+----+----+----+ + | a2 | b1 | a1 | b1 | c1 | + +----+----+----+----+----+ + | 10 | 4 | 1 | 4 | 7 | + | 20 | 5 | | | | + | 30 | 6 | | | | + +----+----+----+----+----+ + "); + Ok(()) +} + #[tokio::test] async fn join_left_mark_different_columns_count_with_filter() -> Result<()> { // select * diff --git a/datafusion/physical-plan/src/sorts/sort.rs b/datafusion/physical-plan/src/sorts/sort.rs index 9a149cc7c38c1..e978b4f69785e 100644 --- a/datafusion/physical-plan/src/sorts/sort.rs +++ b/datafusion/physical-plan/src/sorts/sort.rs @@ -214,7 +214,7 @@ impl ExternalSorterMetrics { /// /// in_mem_batches /// ``` -struct ExternalSorter { +pub(crate) struct ExternalSorter { // ======================================================================== // PROPERTIES: // Fields that define the sorter's configuration and remain constant @@ -331,7 +331,7 @@ impl ExternalSorter { /// Appends an unsorted [`RecordBatch`] to `in_mem_batches` /// /// Updates memory usage metrics, and possibly triggers spilling to disk - async fn insert_batch(&mut self, input: RecordBatch) -> Result<()> { + pub(crate) async fn insert_batch(&mut self, input: RecordBatch) -> Result<()> { if input.num_rows() == 0 { return Ok(()); } @@ -357,7 +357,7 @@ impl ExternalSorter { /// /// 2. A combined streaming merge incorporating both in-memory /// batches and data from spill files on disk. - async fn sort(&mut self) -> Result { + pub(crate) async fn sort(&mut self) -> Result { if self.spilled_before() { // Sort `in_mem_batches` and spill it first. If there are many // `in_mem_batches` and the memory limit is almost reached, merging @@ -419,6 +419,11 @@ impl ExternalSorter { self.metrics.spill_metrics.spill_file_count.value() } + /// The spill metrics of this sorter. + pub(crate) fn spill_metrics(&self) -> &SpillMetrics { + &self.metrics.spill_metrics + } + /// Appending globally sorted batches to the in-progress spill file, and clears /// the `globally_sorted_batches` (also its memory reservation) afterwards. async fn consume_and_spill_append( @@ -479,7 +484,7 @@ impl ExternalSorter { Ok(()) } - async fn abort_in_progress_spill(&mut self) { + pub(crate) async fn abort_in_progress_spill(&mut self) { if let Some((in_progress_file, _)) = &mut self.in_progress_spill_file && let Err(error) = in_progress_file.abort_async().await { diff --git a/datafusion/sqllogictest/test_files/information_schema.slt b/datafusion/sqllogictest/test_files/information_schema.slt index b270eba99d7b0..bd51ac29b3885 100644 --- a/datafusion/sqllogictest/test_files/information_schema.slt +++ b/datafusion/sqllogictest/test_files/information_schema.slt @@ -223,6 +223,7 @@ datafusion.execution.enable_nlj_coordinated_fallback true datafusion.execution.enable_recursive_ctes true datafusion.execution.enforce_batch_size_in_joins false datafusion.execution.hash_join_buffering_capacity 0 +datafusion.execution.hash_join_max_build_size NULL datafusion.execution.keep_partition_by_columns false datafusion.execution.listing_table_factory_infer_partitions true datafusion.execution.listing_table_ignore_subdirectory true @@ -384,6 +385,7 @@ datafusion.execution.enable_nlj_coordinated_fallback true Enables the memory-lim datafusion.execution.enable_recursive_ctes true Should DataFusion support recursive CTEs datafusion.execution.enforce_batch_size_in_joins false Should DataFusion enforce batch size in joins or not. By default, DataFusion will not enforce batch size in joins. Enforcing batch size in joins can reduce memory usage when joining large tables with a highly-selective join filter, but is also slightly slower. Note: this option currently only applies to the symmetric hash join. datafusion.execution.hash_join_buffering_capacity 0 How many bytes to buffer in the probe side of hash joins while the build side is concurrently being built. Without this, hash joins will wait until the full materialization of the build side before polling the probe side. This is useful in scenarios where the query is not completely CPU bounded, allowing to do some early work concurrently and reducing the latency of the query. Note that when hash join buffering is enabled, the probe side will start eagerly polling data, not giving time for the producer side of dynamic filters to produce any meaningful predicate. Queries with dynamic filters might see performance degradation. Disabled by default, set to a number greater than 0 for enabling it. +datafusion.execution.hash_join_max_build_size NULL Maximum build-side size, in bytes per output partition, that a hash join keeps in memory. `NULL` (the default) means no limit. A join whose build side grows past this size finishes as a sort-merge join instead: both inputs are sorted with an external (spilling) sort and then merged. This is not needed for memory safety: with `datafusion.runtime.memory_limit` set, a hash join that cannot reserve memory for its build side already falls back the same way. Set it to switch over earlier, when no memory limit is configured or to make the choice reproducible. The size is per output partition, so a join may hold it times `datafusion.execution.target_partitions`: divide the memory you want hash joins to use by that count, giving 1 GB here for an 8 GB budget over 8 partitions. Requires disk spilling (see `DiskManager`); a partition that falls back reserves about `2 * sort_spill_reservation_bytes` for its sorts, so a budget too small to cover that still fails, inside the sort. Only `PartitionMode::Partitioned` joins fall back. The sort-merge fallback is a short-term measure, so this option may be deprecated and removed once hash joins spill natively. datafusion.execution.keep_partition_by_columns false Should DataFusion keep the columns used for partition_by in the output RecordBatches datafusion.execution.listing_table_factory_infer_partitions true Should a `ListingTable` created through the `ListingTableFactory` infer table partitions from Hive compliant directories. Defaults to true (partition columns are inferred and will be represented in the table schema). datafusion.execution.listing_table_ignore_subdirectory true Should sub directories be ignored when scanning directories for data files. Defaults to true (ignores subdirectories), consistent with Hive. Note that this setting does not affect reading partitioned tables (e.g. `/table/year=2021/month=01/data.parquet`). diff --git a/docs/source/user-guide/configs.md b/docs/source/user-guide/configs.md index 16750d79d750d..7c442b8e56480 100644 --- a/docs/source/user-guide/configs.md +++ b/docs/source/user-guide/configs.md @@ -128,6 +128,7 @@ The following configuration settings are available: | datafusion.execution.sort_pushdown_buffer_capacity | 1073741824 | Maximum buffer capacity (in bytes) per partition for BufferExec inserted during sort pushdown optimization. When PushdownSort eliminates a SortExec under SortPreservingMergeExec, a BufferExec is inserted to replace SortExec's buffering role. This prevents I/O stalls by allowing the scan to run ahead of the merge. This uses strictly less memory than the SortExec it replaces (which buffers the entire partition). The buffer respects the global memory pool limit. Setting this to a large value is safe — actual memory usage is bounded by partition size and global memory limits. | | datafusion.execution.max_spill_file_size_bytes | 134217728 | Maximum size in bytes for individual spill files before rotating to a new file. When operators spill data to disk (e.g., RepartitionExec), they write multiple batches to the same file until this size limit is reached, then rotate to a new file. This reduces syscall overhead compared to one-file-per-batch while preventing files from growing too large. A larger value reduces file creation overhead but may hold more disk space. A smaller value creates more files but allows finer-grained space reclamation as files can be deleted once fully consumed. Now only `RepartitionExec` supports this spill file rotation feature, other spilling operators may create spill files larger than the limit. Default: 128 MB | | datafusion.execution.enable_nlj_coordinated_fallback | true | Enables the memory-limited fallback for `NestedLoopJoinExec` join types that emit unmatched left rows in the final output (LEFT, LEFT SEMI, LEFT ANTI, LEFT MARK, FULL) when the right side has multiple partitions. This fallback coordinates per-chunk left state (visited bitmap and probe-thread counter) across all right-side partitions, which assumes every partition runs in the same process. Distributed engines that execute each output partition as an independent task (e.g. Ballista, datafusion-distributed) build a separate coordinator per task and poll only one partition, so the cross-partition counter never reaches zero and the fallback would stall. Such engines should set this to `false`: the coordinated fallback is then disabled for left-emitting multi-partition joins, which instead fail with a resource-exhaustion error under memory pressure rather than deadlocking. Single-partition and non-left-emitting joins are unaffected and always keep the fallback. | +| datafusion.execution.hash_join_max_build_size | NULL | Maximum build-side size, in bytes per output partition, that a hash join keeps in memory. `NULL` (the default) means no limit. A join whose build side grows past this size finishes as a sort-merge join instead: both inputs are sorted with an external (spilling) sort and then merged. This is not needed for memory safety: with `datafusion.runtime.memory_limit` set, a hash join that cannot reserve memory for its build side already falls back the same way. Set it to switch over earlier, when no memory limit is configured or to make the choice reproducible. The size is per output partition, so a join may hold it times `datafusion.execution.target_partitions`: divide the memory you want hash joins to use by that count, giving 1 GB here for an 8 GB budget over 8 partitions. Requires disk spilling (see `DiskManager`); a partition that falls back reserves about `2 * sort_spill_reservation_bytes` for its sorts, so a budget too small to cover that still fails, inside the sort. Only `PartitionMode::Partitioned` joins fall back. The sort-merge fallback is a short-term measure, so this option may be deprecated and removed once hash joins spill natively. | | datafusion.execution.meta_fetch_concurrency | 32 | Number of files to read in parallel when inferring schema and statistics | | datafusion.execution.minimum_parallel_output_files | 4 | Guarantees a minimum level of output files running in parallel. RecordBatches will be distributed in round robin fashion to each parallel writer. Each writer is closed and a new file opened once soft_max_rows_per_output_file is reached. | | datafusion.execution.soft_max_rows_per_output_file | 50000000 | Target number of rows in output files when writing multiple. This is a soft max, so it can be exceeded slightly. There also will be one file smaller than the limit if the total number of rows written is not roughly divisible by the soft max | From b3a5bf335d488bf14d76dcac60457f10cbac4c65 Mon Sep 17 00:00:00 2001 From: Jay Zhan Date: Sun, 13 Sep 2026 11:08:26 +0800 Subject: [PATCH 2/3] test: promised probe-side ordering disables the sort-merge fallback --- .../hash_join_sort_merge_fallback.rs | 53 ++++++- .../physical-plan/src/joins/hash_join/exec.rs | 144 +++++++++++++++--- .../joins/hash_join/sort_merge_fallback.rs | 8 + 3 files changed, 180 insertions(+), 25 deletions(-) diff --git a/datafusion/core/tests/memory_limit/hash_join_sort_merge_fallback.rs b/datafusion/core/tests/memory_limit/hash_join_sort_merge_fallback.rs index b2573ce5d4afa..36dc1f058f493 100644 --- a/datafusion/core/tests/memory_limit/hash_join_sort_merge_fallback.rs +++ b/datafusion/core/tests/memory_limit/hash_join_sort_merge_fallback.rs @@ -20,8 +20,10 @@ use std::sync::Arc; -use arrow::array::{Int32Array, RecordBatch}; -use arrow::datatypes::{DataType, Field, Schema}; +use arrow::array::{AsArray, Int32Array, RecordBatch}; +use arrow::compute::concat_batches; +use arrow::datatypes::{DataType, Field, Int32Type, Schema}; +use datafusion::physical_plan::displayable; use datafusion::prelude::*; use datafusion_common::assert_contains; use datafusion_execution::disk_manager::{DiskManagerBuilder, DiskManagerMode}; @@ -274,3 +276,50 @@ async fn semi_and_anti_joins_fall_back() { assert_eq!(capped, plain, "{sql}"); } } + +/// A join that promises its probe side's ordering keeps that promise instead +/// of falling back. The planner pushes `ORDER BY r.w` below the join because a +/// hash join emits inner-join rows in probe order, and nothing above it sorts +/// again, so a merge's join-key order would come out wrong. +#[tokio::test] +async fn a_promised_probe_ordering_wins_over_the_fallback() { + let sql = "SELECT l.k, r.w FROM l JOIN r ON l.k = r.k ORDER BY r.w"; + let ctx = context( + config().set_usize("datafusion.execution.hash_join_max_build_size", 1024), + RuntimeEnvBuilder::new(), + ); + let plan = ctx + .sql(sql) + .await + .unwrap() + .create_physical_plan() + .await + .unwrap(); + // Guard the premise: the sort must sit below the join, or this would pass + // without the join ever having promised anything. + let shape = displayable(plan.as_ref()).indent(true).to_string(); + let join_at = shape.find("HashJoinExec").expect("a hash join"); + let sort_at = shape.find("SortExec:").expect("a sort"); + assert!( + sort_at > join_at, + "the sort should sit below the join:\n{shape}" + ); + + let batches = datafusion::physical_plan::collect(Arc::clone(&plan), ctx.task_ctx()) + .await + .unwrap(); + let output = concat_batches(&plan.schema(), &batches).unwrap(); + let w = output + .column_by_name("w") + .unwrap() + .as_primitive::(); + assert!( + w.values().is_sorted(), + "the output must come out ordered by w" + ); + assert_eq!( + plan_metric_sum(plan.as_ref(), "sort_merge_fallback_count"), + 0, + "the join promised the probe order, so it must not have fallen back" + ); +} diff --git a/datafusion/physical-plan/src/joins/hash_join/exec.rs b/datafusion/physical-plan/src/joins/hash_join/exec.rs index fc9d80d2ef733..6278244d02fbd 100644 --- a/datafusion/physical-plan/src/joins/hash_join/exec.rs +++ b/datafusion/physical-plan/src/joins/hash_join/exec.rs @@ -9396,14 +9396,15 @@ mod tests { /// Two sides of `batches` batches each, with duplicate keys, keys that /// only exist on one side, and a non-key column to filter on. - fn sort_merge_fallback_inputs( + fn sort_merge_fallback_batches( batches: usize, - ) -> (Arc, Arc) { + ) -> (Vec, Vec) { let rows_per_batch = 16; let side = |modulus: i32, offset: i32, a: &str, b: &str, c: &str| { - let batches: Vec = (0..batches as i32) + (0..batches as i32) .map(|batch| { let start = batch * rows_per_batch; + // ids ascend within and across batches let ids: Vec = (start..start + rows_per_batch).collect(); // keys repeat within and across batches, and each side // has keys the other side lacks @@ -9412,13 +9413,22 @@ mod tests { let values: Vec = ids.iter().map(|id| (id * 13) % 17).collect(); build_table_i32((a, &ids), (b, &keys), (c, &values)) }) - .collect(); + .collect::>() + }; + // left keys are 0..47, right keys 5..58 + (side(47, 0, "a1", "b1", "c1"), side(53, 5, "a2", "b2", "c2")) + } + + fn sort_merge_fallback_inputs( + batches: usize, + ) -> (Arc, Arc) { + let (left, right) = sort_merge_fallback_batches(batches); + let exec = |batches: Vec| { let schema = batches[0].schema(); TestMemoryExec::try_new_exec(&[batches], schema, None).unwrap() as Arc }; - // left keys are 0..47, right keys 5..58 - (side(47, 0, "a1", "b1", "c1"), side(53, 5, "a2", "b2", "c2")) + (exec(left), exec(right)) } /// `c1 < c2`, so it references both sides @@ -9472,6 +9482,22 @@ mod tests { ) } + /// An unbounded pool with `hash_join_max_build_size` set, so that only the + /// cap, or a reason to decline, decides whether a partition falls back. + fn sort_merge_fallback_capped_ctx(max_build_size: Option) -> Arc { + let unlimited = sort_merge_fallback_task_ctx(None, DiskManagerBuilder::default()); + let mut session_config = unlimited.session_config().clone(); + session_config + .options_mut() + .execution + .hash_join_max_build_size = max_build_size; + Arc::new( + TaskContext::default() + .with_session_config(session_config) + .with_runtime(unlimited.runtime_env()), + ) + } + fn sorted_rows(batches: &[RecordBatch]) -> Vec { batches_to_sort_string(batches) .lines() @@ -9699,6 +9725,88 @@ mod tests { Ok(()) } + /// The fallback is declined whenever the join promises its probe side's + /// ordering, because a merge emits join-key order instead. The promise is + /// what matters, not the input: `maintains_input_order` makes it only for + /// the join types that emit every row while scanning the probe side, and + /// only an ordered probe input turns it into an advertised output ordering. + #[tokio::test] + async fn sort_merge_fallback_honors_a_promised_probe_ordering() -> Result<()> { + let (left, _) = sort_merge_fallback_inputs(32); + let (_, right) = sort_merge_fallback_batches(32); + let schema = right[0].schema(); + // `a2` is the probe side's id column, ascending across batches + let ordering = datafusion_physical_expr_common::sort_expr::LexOrdering::new([ + PhysicalSortExpr::new_default(Arc::new(Column::new_with_schema( + "a2", &schema, + )?)), + ]) + .unwrap(); + let right = TestMemoryExec::try_new(&[right], schema, None)? + .try_with_sort_information(vec![ordering])?; + let right: Arc = + Arc::new(TestMemoryExec::update_cache(&Arc::new(right))); + let on: JoinOn = vec![( + Arc::new(Column::new_with_schema("b1", &left.schema())?) as _, + Arc::new(Column::new_with_schema("b2", &right.schema())?) as _, + )]; + let join = |join_type: JoinType| { + HashJoinExec::try_new( + Arc::clone(&left), + Arc::clone(&right), + on.clone(), + None, + &join_type, + None, + PartitionMode::Partitioned, + NullEquality::NullEqualsNothing, + false, + ) + }; + // A join that declined up front never registers the counter at all. + let fallbacks = |join: &HashJoinExec| { + join.metrics() + .unwrap() + .sum_by_name(SORT_MERGE_FALLBACK_COUNT_METRIC_NAME) + .map_or(0, |v| v.as_usize()) + }; + + // An inner join promises the probe order, so under a cap that would + // otherwise switch it, it stays a hash join and keeps that order. + let inner = join(JoinType::Inner)?; + assert!(inner.properties().output_ordering().is_some()); + let batches = common::collect( + inner.execute(0, sort_merge_fallback_capped_ctx(Some(1024)))?, + ) + .await?; + let output = concat_batches(&inner.schema(), &batches)?; + let ids = output + .column_by_name("a2") + .unwrap() + .as_primitive::(); + assert!( + ids.values().is_sorted(), + "the output must stay in probe order" + ); + assert_eq!( + fallbacks(&inner), + 0, + "a promised ordering must stop the fallback" + ); + + // A left join never promises it, so the same ordered input does not + // stop the fallback. + let left_join = join(JoinType::Left)?; + assert!(left_join.properties().output_ordering().is_none()); + common::collect( + left_join.execute(0, sort_merge_fallback_capped_ctx(Some(1024)))?, + ) + .await?; + assert_eq!(fallbacks(&left_join), 1); + + Ok(()) + } + /// Without disk the join fails as before #[tokio::test] async fn sort_merge_fallback_needs_disk() -> Result<()> { @@ -9760,23 +9868,10 @@ mod tests { false, ) }; - let task_ctx = |max_build_size: Option| { - let unlimited = - sort_merge_fallback_task_ctx(None, DiskManagerBuilder::default()); - let mut session_config = unlimited.session_config().clone(); - session_config - .options_mut() - .execution - .hash_join_max_build_size = max_build_size; - Arc::new( - TaskContext::default() - .with_session_config(session_config) - .with_runtime(unlimited.runtime_env()), - ) - }; - let in_memory = join()?; - let expected = common::collect(in_memory.execute(0, task_ctx(None))?).await?; + let expected = + common::collect(in_memory.execute(0, sort_merge_fallback_capped_ctx(None))?) + .await?; assert_eq!( in_memory .metrics() @@ -9787,7 +9882,10 @@ mod tests { ); let fallback = join()?; - let actual = common::collect(fallback.execute(0, task_ctx(Some(1024)))?).await?; + let actual = common::collect( + fallback.execute(0, sort_merge_fallback_capped_ctx(Some(1024)))?, + ) + .await?; assert_eq!( fallback .metrics() diff --git a/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs b/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs index 5e60712604195..e9bd4d1237153 100644 --- a/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs +++ b/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs @@ -52,6 +52,14 @@ //! back is not by itself decisive: a single falling-back partition completes //! when the budget covers its sorts. //! +//! A join that promises its probe side's ordering never falls back, because a +//! merge emits join-key order instead. That covers more than inputs with a +//! declared ordering: the planner pushes an `ORDER BY` on probe-side columns +//! below an inner or right join precisely because the join keeps that order, +//! so such queries stay on the in-memory path and still fail under memory +//! pressure. Re-sorting the merge output to honor the promise is follow-up +//! work. +//! //! [`HashJoinExec`]: super::HashJoinExec //! [`PartitionMode::Partitioned`]: crate::joins::PartitionMode::Partitioned From 0ac5d60e9e4209f3c442bc7f45e71d0349157a04 Mon Sep 17 00:00:00 2001 From: Jay Zhan Date: Thu, 17 Sep 2026 22:42:46 +0800 Subject: [PATCH 3/3] refactor: simplify the sort-merge fallback's sort input, output stream and build report --- .../physical-plan/src/joins/hash_join/exec.rs | 85 ++++---- .../joins/hash_join/sort_merge_fallback.rs | 195 ++++++++++-------- .../src/joins/hash_join/stream.rs | 102 +++------ 3 files changed, 175 insertions(+), 207 deletions(-) diff --git a/datafusion/physical-plan/src/joins/hash_join/exec.rs b/datafusion/physical-plan/src/joins/hash_join/exec.rs index c7b865fc5c815..28304e29644a3 100644 --- a/datafusion/physical-plan/src/joins/hash_join/exec.rs +++ b/datafusion/physical-plan/src/joins/hash_join/exec.rs @@ -979,11 +979,14 @@ impl HashJoinExec { /// Returns what a partition needs to fall back to a sort-merge join when /// its build side does not fit in memory, or `None` when this join cannot - /// fall back (see the `sort_merge_fallback` module). + /// fall back (see the `sort_merge_fallback` module). `compute_bounds` + /// says whether a dynamic filter is pushed down, in which case a sorted + /// build side must report its join key bounds. fn sort_merge_fallback_context( &self, partition: usize, context: &Arc, + compute_bounds: bool, ) -> Result> { let options = context.session_config().options(); if !context.runtime_env().disk_manager.tmp_files_enabled() @@ -1011,11 +1014,7 @@ impl HashJoinExec { } } - let (on_left, on_right) = self - .on - .iter() - .map(|(l, r)| (Arc::clone(l), Arc::clone(r))) - .unzip::<_, _, Vec<_>, Vec<_>>(); + let (on_left, on_right): (Vec<_>, Vec<_>) = self.on.iter().cloned().unzip(); Ok(Some(SortMergeFallbackContext { context: Arc::clone(context), partition, @@ -1034,6 +1033,7 @@ impl HashJoinExec { projection: self.projection.as_deref().map(|p| p.to_vec()), fetch: self.fetch, max_build_size: options.execution.hash_join_max_build_size, + compute_bounds, })) } @@ -1751,8 +1751,11 @@ impl ExecutionPlan for HashJoinExec { .flatten(); let null_aware = self.null_aware_mode()?; - let sort_merge_fallback = - self.sort_merge_fallback_context(partition, &context)?; + let sort_merge_fallback = self.sort_merge_fallback_context( + partition, + &context, + enable_dynamic_filter_pushdown, + )?; let left_fut = match self.mode { PartitionMode::CollectLeft => self.left_fut.try_once(|| { @@ -2746,7 +2749,10 @@ impl CollectLeftAccumulator { /// /// # Returns /// A new `CollectLeftAccumulator` instance configured for the expression's data type - fn try_new(expr: Arc, schema: &SchemaRef) -> Result { + pub(super) fn try_new( + expr: Arc, + schema: &SchemaRef, + ) -> Result { /// Recursively unwraps dictionary types to get the underlying value type. fn dictionary_value_type(data_type: &DataType) -> DataType { match data_type { @@ -2913,7 +2919,7 @@ fn new_join_hashmap( #[expect(clippy::too_many_arguments)] async fn collect_left_input( random_state: RandomState, - left_stream: SendableRecordBatchStream, + mut left_stream: SendableRecordBatchStream, on_left: Vec, metrics: BuildProbeJoinMetrics, reservation: MemoryReservation, @@ -2924,7 +2930,7 @@ async fn collect_left_input( null_equality: NullEquality, null_aware: Option, array_map_created_count: Count, - sort_merge_fallback: Option, + mut sort_merge_fallback: Option, ) -> Result { let schema = left_stream.schema(); @@ -2945,8 +2951,6 @@ async fn collect_left_input( should_compute_dynamic_filters || is_phj_candidate, )?; - let mut left_stream = left_stream; - let mut sort_merge_fallback = sort_merge_fallback; while let Some(batch) = left_stream.next().await { let batch = batch?; // Update accumulators if computing bounds @@ -2958,10 +2962,11 @@ async fn collect_left_input( // Decide if we spill or not let batch_size = state.memory_counter.count_batch(&batch); - // Reserve memory for incoming batch - let fall_back = match state.reservation.try_grow(batch_size) { + // Reserve memory for incoming batch, taking the fallback instead when + // the build side should be sorted from here on + let fallback = match state.reservation.try_grow(batch_size) { // The build side grew past the configured size - Ok(()) => sort_merge_fallback.as_ref().is_some_and(|fallback| { + Ok(()) => sort_merge_fallback.take_if(|fallback| { fallback .max_build_size .is_some_and(|max| state.reservation.size() > max) @@ -2972,30 +2977,22 @@ async fn collect_left_input( Err(error) if sort_merge_fallback.is_some() && is_resources_exhausted(&error) => { - true + sort_merge_fallback.take() } Err(error) => return Err(error), }; - if let Some(fallback) = sort_merge_fallback.take_if(|_| fall_back) { + if let Some(fallback) = fallback { let BuildSideState { - batches, + mut batches, reservation, - bounds_accumulators, .. } = state; + batches.push(batch); // Release what the collected batches reserved: the external sort // accounts for what it keeps in memory itself. drop(reservation); - let sorted = sort_build_side( - fallback, - schema, - batches, - Some(batch), - Some(left_stream), - bounds_accumulators.filter(|_| should_compute_dynamic_filters), - None, - ) - .await?; + let sorted = + sort_build_side(fallback, schema, batches, Some(left_stream)).await?; return Ok(BuildSideOutcome::SortMerge(sorted)); } // Update metrics @@ -3059,16 +3056,7 @@ async fn collect_left_input( else { return Err(error); }; - let sorted = sort_build_side( - fallback, - schema, - batches, - None, - None, - None, - bounds.filter(|_| should_compute_dynamic_filters), - ) - .await?; + let sorted = sort_build_side(fallback, schema, batches, None).await?; Ok(BuildSideOutcome::SortMerge(sorted)) } } @@ -9961,13 +9949,15 @@ mod tests { ); assert!(!batches.is_empty(), "the join should have produced rows"); - // Every probe row must survive the filter the fallback reported. + // Every probe row within the build side's key range (0..47) must + // survive the filter the fallback reported; the bounds it reported + // still prune a key outside that range. let probe = RecordBatch::try_new( Arc::clone(&probe_schema), vec![ - Arc::new(Int32Array::from(vec![0, 1, 2, 3])), - Arc::new(Int32Array::from(vec![5, 6, 7, 8])), - Arc::new(Int32Array::from(vec![0, 1, 2, 3])), + Arc::new(Int32Array::from(vec![0, 1, 2, 3, 4])), + Arc::new(Int32Array::from(vec![5, 6, 7, 8, 100])), + Arc::new(Int32Array::from(vec![0, 1, 2, 3, 4])), ], )?; let filter = dynamic_filter.current()?; @@ -9976,10 +9966,11 @@ mod tests { .as_any() .downcast_ref::() .expect("a filter evaluates to a BooleanArray"); + let kept: Vec = (0..kept.len()).map(|i| kept.value(i)).collect(); assert_eq!( - (0..kept.len()).filter(|i| kept.value(*i)).count(), - probe.num_rows(), - "a fallen-back partition must not prune probe rows, filter was {filter}" + kept, + [true, true, true, true, false], + "a fallen-back partition prunes by bounds only, filter was {filter}" ); Ok(()) diff --git a/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs b/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs index e9bd4d1237153..45d6e63a21c47 100644 --- a/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs +++ b/datafusion/physical-plan/src/joins/hash_join/sort_merge_fallback.rs @@ -74,7 +74,7 @@ use crate::joins::utils::JoinFilter; use crate::limit::LimitStream; use crate::metrics::{BaselineMetrics, Count, ExecutionPlanMetricsSet, SpillMetrics}; use crate::sorts::sort::ExternalSorter; -use crate::stream::RecordBatchStreamAdapter; +use crate::stream::{EmptyRecordBatchStream, RecordBatchStreamAdapter}; use arrow::array::Array; use arrow::compute::SortOptions; @@ -85,7 +85,7 @@ use datafusion_execution::TaskContext; use datafusion_physical_expr::PhysicalExprRef; use datafusion_physical_expr_common::sort_expr::LexOrdering; use datafusion_physical_expr_common::utils::evaluate_expressions_to_arrays; -use futures::StreamExt; +use futures::{Stream, StreamExt, TryStreamExt, stream}; use parking_lot::Mutex; /// Everything the fallback needs from the join, captured once per partition @@ -121,6 +121,9 @@ pub(super) struct SortMergeFallbackContext { /// back even though the memory pool would allow more; `None` for no limit /// (`datafusion.execution.hash_join_max_build_size`) pub(super) max_build_size: Option, + /// Whether a dynamic filter is pushed down to the probe side, so that a + /// sorted build side must report its join key bounds. + pub(super) compute_bounds: bool, } impl SortMergeFallbackContext { @@ -169,15 +172,12 @@ pub(super) fn is_resources_exhausted(error: &DataFusionError) -> bool { matches!(error.find_root(), DataFusionError::ResourcesExhausted(_)) } -/// Sorts `buffered` followed by the rest of `input` on `ordering` with an -/// external sort, calling `observe` for every batch on its way in. +/// Sorts `input` on `ordering` with an external sort. async fn sort_batches( ctx: &SortMergeFallbackContext, schema: SchemaRef, ordering: LexOrdering, - buffered: Vec, - mut input: Option, - mut observe: impl FnMut(&RecordBatch) -> Result<()>, + input: impl Stream> + Unpin, ) -> Result { let session_config = ctx.context.session_config(); let execution_options = &session_config.options().execution; @@ -197,23 +197,10 @@ async fn sort_batches( ctx.context.runtime_env(), )?; - for batch in buffered { - insert_batch(&mut sorter, batch, &mut observe).await?; - } - if let Some(input) = input.as_mut() { - while let Some(batch) = input.next().await { - let batch = match batch { - Ok(batch) => batch, - Err(error) => { - sorter.abort_in_progress_spill().await; - return Err(error); - } - }; - insert_batch(&mut sorter, batch, &mut observe).await?; - } + if let Err(error) = insert_all(&mut sorter, input).await { + sorter.abort_in_progress_spill().await; + return Err(error); } - drop(input); - let sorted = sorter.sort().await?; let spills = sorter.spill_metrics(); @@ -230,16 +217,13 @@ async fn sort_batches( Ok(sorted) } -/// Feeds one batch to `sorter`, abandoning its in-progress spill on error. -async fn insert_batch( +/// Feeds every batch of `input` to `sorter`. +async fn insert_all( sorter: &mut ExternalSorter, - batch: RecordBatch, - observe: &mut impl FnMut(&RecordBatch) -> Result<()>, + mut input: impl Stream> + Unpin, ) -> Result<()> { - observe(&batch)?; - if let Err(error) = sorter.insert_batch(batch).await { - sorter.abort_in_progress_spill().await; - return Err(error); + while let Some(batch) = input.next().await { + sorter.insert_batch(batch?).await?; } Ok(()) } @@ -247,64 +231,69 @@ async fn insert_batch( /// Sorts the build side of a partition after its in-memory collection ran /// out of memory. /// -/// `batches` are the build batches collected so far, `pending` the batch whose -/// reservation failed and `rest` the not yet consumed remainder of the build -/// input (both `None` when the input was fully consumed and the hash table -/// itself did not fit). The caller has already released the reservation held -/// for `batches`; the external sort reserves what it keeps in memory itself. +/// `batches` are the build batches collected so far and `rest` the not yet +/// consumed remainder of the build input (`None` when the input was fully +/// consumed and the hash table itself did not fit). The caller has already +/// released the reservation held for `batches`; the external sort reserves +/// what it keeps in memory itself. /// -/// The join key bounds needed by a pushed-down dynamic filter are computed -/// with `bounds_accumulators` over every batch (min/max are idempotent, so -/// re-feeding the already accumulated `batches` is harmless), unless the -/// caller already has the final `bounds`. +/// When a dynamic filter is pushed down, the join key bounds it needs are +/// computed over every batch on its way into the sort. pub(super) async fn sort_build_side( ctx: SortMergeFallbackContext, schema: SchemaRef, batches: Vec, - pending: Option, rest: Option, - mut bounds_accumulators: Option>, - bounds: Option, ) -> Result { ctx.fallback_count.add(1); let ordering = ctx.sort_ordering(&ctx.on_left)?; - let mut buffered = batches; - buffered.extend(pending); - - let on_left = ctx.on_left.clone(); + let mut accumulators = ctx + .compute_bounds + .then(|| { + ctx.on_left + .iter() + .map(|expr| CollectLeftAccumulator::try_new(Arc::clone(expr), &schema)) + .collect::>>() + }) + .transpose()?; let mut keys_have_null = false; let mut num_rows = 0; - let observe = |batch: &RecordBatch| -> Result<()> { - num_rows += batch.num_rows(); - if let Some(accumulators) = bounds_accumulators.as_mut() { - for accumulator in accumulators { - accumulator.update_batch(batch)?; + + let rest = rest + .unwrap_or_else(|| Box::pin(EmptyRecordBatchStream::new(Arc::clone(&schema)))); + let input = stream::iter(batches.into_iter().map(Ok)).chain(rest).map( + |batch| -> Result { + let batch = batch?; + num_rows += batch.num_rows(); + if let Some(accumulators) = accumulators.as_mut() { + for accumulator in accumulators { + accumulator.update_batch(&batch)?; + } } - } - if !keys_have_null { - keys_have_null = evaluate_expressions_to_arrays(&on_left, batch)? - .iter() - .any(|array| array.logical_null_count() > 0); - } - Ok(()) - }; + if !keys_have_null { + keys_have_null = evaluate_expressions_to_arrays(&ctx.on_left, &batch)? + .iter() + .any(|array| array.logical_null_count() > 0); + } + Ok(batch) + }, + ); - let stream = sort_batches(&ctx, schema, ordering, buffered, rest, observe) + let stream = sort_batches(&ctx, schema, ordering, input) .await .map_err(|e| { e.context("HashJoinExec sort-merge fallback: sorting the build side") })?; - let bounds = match bounds_accumulators { + let bounds = match accumulators { Some(accumulators) if num_rows > 0 => Some(PartitionBounds::new( accumulators .into_iter() .map(CollectLeftAccumulator::evaluate) .collect::>>()?, )), - Some(_) => None, - None => bounds, + _ => None, }; Ok(SortedBuildSide { @@ -314,56 +303,80 @@ pub(super) async fn sort_build_side( }) } -/// Sorts the probe side and joins it with the sorted build side as a -/// sort-merge join, returning the join's output stream (projected and limited -/// like the hash join's own output would be). -pub(super) async fn run_sort_merge_fallback( +/// Joins the sorted build side with the probe side as a sort-merge join, +/// returning the join's output stream (projected and limited like the hash +/// join's own output would be). The probe side is sorted on first poll. +pub(super) fn run_sort_merge_fallback( + ctx: SortMergeFallbackContext, + build: SendableRecordBatchStream, + probe: SendableRecordBatchStream, +) -> SendableRecordBatchStream { + let schema = Arc::clone(&ctx.output_schema); + let output = stream::once(sort_probe_and_join(ctx, build, probe)).try_flatten(); + Box::pin(RecordBatchStreamAdapter::new(schema, output)) +} + +async fn sort_probe_and_join( ctx: SortMergeFallbackContext, build: SendableRecordBatchStream, probe: SendableRecordBatchStream, ) -> Result { let ordering = ctx.sort_ordering(&ctx.on_right)?; - let probe = sort_batches(&ctx, probe.schema(), ordering, vec![], Some(probe), |_| { - Ok(()) - }) - .await - .map_err(|e| e.context("HashJoinExec sort-merge fallback: sorting the probe side"))?; + let probe = sort_batches(&ctx, probe.schema(), ordering, probe) + .await + .map_err(|e| { + e.context("HashJoinExec sort-merge fallback: sorting the probe side") + })?; + let SortMergeFallbackContext { + context, + partition, + metrics, + on_left, + on_right, + sort_options, + join_type, + filter, + null_equality, + join_schema, + output_schema, + projection, + fetch, + .. + } = ctx; let joined = sort_merge_join_stream( SortMergeJoinInputs { - schema: Arc::clone(&ctx.join_schema), - sort_options: ctx.sort_options.clone(), - null_equality: ctx.null_equality, + schema: join_schema, + sort_options, + null_equality, left: build, right: probe, - on_left: ctx.on_left.clone(), - on_right: ctx.on_right.clone(), - filter: ctx.filter.clone(), - join_type: ctx.join_type, - partition: ctx.partition, + on_left, + on_right, + filter, + join_type, + partition, }, - &ctx.metrics, - &ctx.context, + &metrics, + &context, )?; - let output: SendableRecordBatchStream = match ctx.projection.clone() { + let output: SendableRecordBatchStream = match projection { Some(projection) => Box::pin(RecordBatchStreamAdapter::new( - Arc::clone(&ctx.output_schema), + output_schema, joined.map(move |batch| Ok(batch?.project(&projection)?)), )), None => joined, }; // The limit's own baseline metrics would count the output a second time. - let output = match ctx.fetch { + Ok(match fetch { Some(_) => Box::pin(LimitStream::new( output, 0, - ctx.fetch, - BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), ctx.partition), + fetch, + BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), partition), )), None => output, - }; - - Ok(output) + }) } diff --git a/datafusion/physical-plan/src/joins/hash_join/stream.rs b/datafusion/physical-plan/src/joins/hash_join/stream.rs index 09c8822f84e81..6d32a555e051f 100644 --- a/datafusion/physical-plan/src/joins/hash_join/stream.rs +++ b/datafusion/physical-plan/src/joins/hash_join/stream.rs @@ -60,8 +60,7 @@ use datafusion_physical_expr::PhysicalExprRef; use datafusion_common::hash_utils::RandomState; use datafusion_physical_expr_common::utils::evaluate_expressions_to_arrays; -use futures::future::BoxFuture; -use futures::{FutureExt, Stream, StreamExt, ready}; +use futures::{Stream, StreamExt, ready}; /// Represents build-side of hash join. pub(super) enum BuildSide { @@ -70,8 +69,9 @@ pub(super) enum BuildSide { /// Indicates that build-side data has been collected Ready(BuildSideReadyState), /// Indicates that the build side did not fit in memory and the partition - /// is finishing as a sort-merge join (see the `sort_merge_fallback` module) - SortMergeFallback(SortMergeFallbackState), + /// is finishing as a sort-merge join whose output this stream forwards + /// (see the `sort_merge_fallback` module) + SortMergeFallback(SendableRecordBatchStream), } /// Container for BuildSide::Initial related data @@ -80,14 +80,6 @@ pub(super) struct BuildSideInitialState { pub(super) left_fut: OnceFut, } -/// Progress of the sort-merge fallback of one partition -pub(super) enum SortMergeFallbackState { - /// Sorting the probe side and setting up the sort-merge join - Preparing(BoxFuture<'static, Result>), - /// Forwarding the sort-merge join's output - Running(SendableRecordBatchStream), -} - /// Container for BuildSide::Ready related data pub(super) struct BuildSideReadyState { /// Collected build-side data @@ -662,11 +654,6 @@ impl HashJoinStream { return Self::state_after_build_ready(self.join_type, left_data.as_ref()); } - let pushdown = left_data.membership().clone(); - let bounds = left_data - .bounds - .clone() - .unwrap_or_else(|| PartitionBounds::new(vec![])); // Use the logical null count: a dictionary key whose entry points at a // NULL dictionary value is a NULL key even though the key bitmap has no // physical nulls (`null_count() == 0` but `logical_null_count() > 0`). @@ -675,6 +662,22 @@ impl HashJoinStream { .iter() .any(|array| array.logical_null_count() > 0); + self.schedule_build_report( + left_data.membership().clone(), + left_data.bounds.clone(), + keys_have_null, + ) + } + + /// Reports this partition's build side to the shared accumulator and + /// waits for the dynamic filter built from every partition's report. + fn schedule_build_report( + &mut self, + pushdown: PushdownStrategy, + bounds: Option, + keys_have_null: bool, + ) -> HashJoinStreamState { + let bounds = bounds.unwrap_or_else(|| PartitionBounds::new(vec![])); let build_data = match self.mode { PartitionMode::Partitioned => PartitionBuildData::Partitioned { partition_id: self.partition, @@ -807,7 +810,7 @@ impl HashJoinStream { self.build_side = BuildSide::Ready(BuildSideReadyState { left_data }); } BuildSideOutcome::SortMerge(sorted_build) => { - let Some(fallback) = self.sort_merge_fallback.clone() else { + let Some(fallback) = self.sort_merge_fallback.take() else { return Poll::Ready(internal_err!( "build side fell back to a sort-merge join, but the stream cannot" )); @@ -823,50 +826,31 @@ impl HashJoinStream { sorted_build.bounds.clone(), sorted_build.keys_have_null, ); - self.build_side = - BuildSide::SortMergeFallback(SortMergeFallbackState::Preparing( - run_sort_merge_fallback(fallback, build, probe).boxed(), - )); + self.build_side = BuildSide::SortMergeFallback(run_sort_merge_fallback( + fallback, build, probe, + )); } } Poll::Ready(Ok(StatefulStreamResult::Continue)) } - /// Forwards the output of the sort-merge join this partition fell back to, - /// first waiting for the probe side to be sorted. + /// Forwards the output of the sort-merge join this partition fell back to. fn poll_sort_merge_fallback( &mut self, cx: &mut std::task::Context<'_>, ) -> Poll>> { - let BuildSide::SortMergeFallback(fallback) = &mut self.build_side else { + let BuildSide::SortMergeFallback(stream) = &mut self.build_side else { return Poll::Ready(Some(internal_err!( "Expected build side in sort-merge fallback state" ))); }; - loop { - match fallback { - SortMergeFallbackState::Preparing(preparing) => { - match ready!(preparing.poll_unpin(cx)) { - Ok(stream) => { - *fallback = SortMergeFallbackState::Running(stream); - } - Err(e) => { - self.state = HashJoinStreamState::Completed; - return Poll::Ready(Some(Err(e))); - } - } - } - SortMergeFallbackState::Running(stream) => { - // The sort-merge join stream records its output in the - // join's metrics itself. - let poll = stream.poll_next_unpin(cx); - if matches!(poll, Poll::Ready(None)) { - self.state = HashJoinStreamState::Completed; - } - return poll; - } - } + // The sort-merge join stream records its output in the join's metrics + // itself. + let poll = stream.poll_next_unpin(cx); + if matches!(poll, Poll::Ready(None | Some(Err(_)))) { + self.state = HashJoinStreamState::Completed; } + poll } /// Transitions state after the build side was sorted instead of hashed, @@ -883,27 +867,7 @@ impl HashJoinStream { // No hash table exists to push down; the bounds still narrow the // probe-side scan. - let pushdown = PushdownStrategy::Unknown; - let bounds = bounds.unwrap_or_else(|| PartitionBounds::new(vec![])); - let build_data = match self.mode { - PartitionMode::Partitioned => PartitionBuildData::Partitioned { - partition_id: self.partition, - pushdown, - bounds, - keys_have_null, - }, - PartitionMode::CollectLeft => PartitionBuildData::CollectLeft { - pushdown, - bounds, - keys_have_null, - }, - PartitionMode::Auto => unreachable!( - "PartitionMode::Auto should not be present at execution time. This is a bug in DataFusion, please report it!" - ), - }; - - self.build_report.schedule(build_data); - HashJoinStreamState::WaitPartitionBoundsReport + self.schedule_build_report(PushdownStrategy::Unknown, bounds, keys_have_null) } /// Fetches next batch from probe-side