From a493834a003ea6a743ae71b07323cd62920b9db9 Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Thu, 17 Sep 2026 19:10:06 +0000 Subject: [PATCH 1/4] perf: skip replay headroom for intermediate spill merges --- .../src/sorts/multi_level_merge.rs | 84 ++++- .../replay_headroom_tests.rs | 288 ++++++++++++++++++ 2 files changed, 358 insertions(+), 14 deletions(-) create mode 100644 datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs diff --git a/datafusion/physical-plan/src/sorts/multi_level_merge.rs b/datafusion/physical-plan/src/sorts/multi_level_merge.rs index c1fe893e0df57..ceaec13ff35b6 100644 --- a/datafusion/physical-plan/src/sorts/multi_level_merge.rs +++ b/datafusion/physical-plan/src/sorts/multi_level_merge.rs @@ -209,8 +209,8 @@ impl MultiLevelMergeBuilder { self } - /// Leave replay headroom while selecting merge buffers. Temporary splitting - /// workspace can still use the full pool because replay has not started. + /// Leave replay headroom for the final merge. Intermediate merges and + /// splitting can use the full pool because replay has not started. pub(super) fn with_replay_headroom(mut self, reserve: bool) -> Self { self.reserve_replay_headroom = reserve; self @@ -373,14 +373,46 @@ impl MultiLevelMergeBuilder { let minimum_number_of_required_streams = 2_usize.saturating_sub(self.sorted_streams.len()); - let (sorted_spill_files, buffer_size) = match self - .get_sorted_spill_files_to_merge( + let total_spill_files = self.sorted_spill_files.len(); + let mut selection = self.get_sorted_spill_files_to_merge( + 2, + // we must have at least 2 streams to merge + minimum_number_of_required_streams, + &mut memory_reservation, + allow_minimum_without_headroom, + self.reserve_replay_headroom, + usize::MAX, + )?; + + // Prefer a final merge with replay headroom. Otherwise an + // intermediate merge writes back to disk and can use the full + // pool. Preserve split requests: splitting also caps the merge + // output size when input batches are shorter than batch_size. + // Hold one run back so this larger merge cannot feed replay, + // but only if enough inputs remain to make progress. + let is_intermediate = matches!( + &selection, + SpillFilesToMerge::Ready(spills, _) if spills.len() < total_spill_files + ); + if self.reserve_replay_headroom + && is_intermediate + && total_spill_files > minimum_number_of_required_streams + { + if let SpillFilesToMerge::Ready(spills, _) = selection { + self.sorted_spill_files.splice(0..0, spills); + } + memory_reservation.free(); + selection = self.get_sorted_spill_files_to_merge( 2, - // we must have at least 2 streams to merge minimum_number_of_required_streams, &mut memory_reservation, - allow_minimum_without_headroom, - )? { + false, + false, + total_spill_files - 1, + )?; + } + + let (sorted_spill_files, buffer_size) = match selection { SpillFilesToMerge::Ready(sorted_spill_files, buffer_size) => { (sorted_spill_files, buffer_size) } @@ -512,6 +544,8 @@ impl MultiLevelMergeBuilder { minimum_number_of_required_streams: usize, reservation: &mut MemoryReservation, allow_minimum_without_headroom: bool, + reserve_replay_headroom: bool, + max_files: usize, ) -> Result { assert_ne!(buffer_len, 0, "Buffer length must be greater than 0"); let mut number_of_spills_to_read_for_current_phase = 0; @@ -520,7 +554,8 @@ impl MultiLevelMergeBuilder { .env() .disk_manager .max_spill_merge_fan_in(); - let max_spill_files = effective_spill_merge_fan_in(configured_fan_in); + let max_spill_files = + effective_spill_merge_fan_in(configured_fan_in).min(max_files); // Track total memory needed for spill file buffers. When the // reservation has pre-reserved bytes (from sort_spill_reservation_bytes), // those bytes cover the first N spill files without additional pool @@ -551,7 +586,7 @@ impl MultiLevelMergeBuilder { && buffer_len == 1 && number_of_spills_to_read_for_current_phase < minimum_number_of_required_streams; - let check_headroom = self.reserve_replay_headroom && !skip_headroom; + let check_headroom = reserve_replay_headroom && !skip_headroom; let admission = if check_headroom { // Ask the pool for merge buffers plus equal replay space, then // return the spare bytes before exposing the merge stream. @@ -585,6 +620,8 @@ impl MultiLevelMergeBuilder { minimum_number_of_required_streams, reservation, allow_minimum_without_headroom, + reserve_replay_headroom, + max_files, ); } @@ -620,7 +657,7 @@ impl MultiLevelMergeBuilder { } } - if self.reserve_replay_headroom { + if reserve_replay_headroom { // `total_needed` may include a rejected candidate. Keep only the // buffers that were admitted, releasing temporary replay headroom. reservation.shrink(reservation.size() - accepted_memory); @@ -906,6 +943,9 @@ impl RecordBatchStream for StreamAttachedReservation { } } +#[cfg(test)] +mod replay_headroom_tests; + #[cfg(test)] mod tests { use super::*; @@ -1203,8 +1243,15 @@ mod tests { // Actual batches contain one row even though the nominal size is 8192. assert_eq!(builder.sorted_spill_files[0].1, 1); let mut reservation = builder.reservation.new_empty(); - let SpillFilesToMerge::Ready(spills, buffer_len) = - builder.get_sorted_spill_files_to_merge(2, 2, &mut reservation, true)? + let SpillFilesToMerge::Ready(spills, buffer_len) = builder + .get_sorted_spill_files_to_merge( + 2, + 2, + &mut reservation, + true, + true, + usize::MAX, + )? else { panic!("minimum merge should fit the pool"); }; @@ -1244,8 +1291,15 @@ mod tests { let mut builder = build_merge_builder(spill_manager, schema, spills, &pool, 1) .with_replay_headroom(true); let mut reservation = builder.reservation.new_empty(); - let SpillFilesToMerge::Ready(spills, buffer_len) = - builder.get_sorted_spill_files_to_merge(1, 2, &mut reservation, false)? + let SpillFilesToMerge::Ready(spills, buffer_len) = builder + .get_sorted_spill_files_to_merge( + 1, + 2, + &mut reservation, + false, + true, + usize::MAX, + )? else { panic!("two streams and replay headroom should fit the pool"); }; @@ -1670,6 +1724,8 @@ mod tests { 2, &mut merge_reservation, false, + false, + usize::MAX, )? { SpillFilesToMerge::Ready(spills, buffer_len) => (spills, buffer_len), SpillFilesToMerge::SplitThenRetry(index) => { diff --git a/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs b/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs new file mode 100644 index 0000000000000..6ba036e2302bd --- /dev/null +++ b/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs @@ -0,0 +1,288 @@ +// 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. + +use super::*; + +use crate::expressions::PhysicalSortExpr; +use arrow::array::{AsArray, Int64Array, StringArray}; +use arrow::compute::concat_batches; +use arrow::datatypes::{DataType, Field, Int64Type, Schema}; +use datafusion_execution::memory_pool::{GreedyMemoryPool, MemoryConsumer, MemoryPool}; +use datafusion_execution::runtime_env::{RuntimeEnv, RuntimeEnvBuilder}; +use datafusion_physical_expr::expressions::Column; +use datafusion_physical_expr_common::metrics::{ExecutionPlanMetricsSet, SpillMetrics}; + +struct ReplayMergeFixture { + builder: MultiLevelMergeBuilder, + env: Arc, + pool: Arc, + pool_size: usize, + metrics: SpillMetrics, + input_bytes: usize, +} + +/// Create equally sized, interleaved input runs. +fn replay_merge_fixture( + run_count: usize, + rows_per_run: usize, + memory_batches: usize, + max_fan_in: usize, +) -> Result { + let env = RuntimeEnvBuilder::new() + .with_max_spill_merge_fan_in(max_fan_in) + .build_arc()?; + let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)])); + let metrics = SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0); + let spill_manager = + SpillManager::new(Arc::clone(&env), metrics.clone(), Arc::clone(&schema)); + let mut spills = Vec::with_capacity(run_count); + for run in 0..run_count { + let values = Int64Array::from_iter_values( + (0..rows_per_run).map(|row| (row * run_count + run) as i64), + ); + let batch = RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(values)])?; + let (file, max_record_batch_memory) = spill_manager + .spill_record_batch_iter_and_return_max_batch_memory( + std::iter::once(Ok(batch)), + "replay headroom test input", + )? + .expect("a nonempty input must spill"); + spills.push(SortedSpillFile { + file, + max_record_batch_memory, + }); + } + let input_bytes = metrics.spilled_bytes.value(); + let pool_size = memory_batches * spills[0].max_record_batch_memory; + let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_size)); + let reservation = MemoryConsumer::new("replay headroom test").register(&pool); + let expr = [PhysicalSortExpr::new_default(Arc::new(Column::new("x", 0)))].into(); + let builder = MultiLevelMergeBuilder::new( + spill_manager, + Arc::clone(&schema), + spills, + vec![], + expr, + BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0), + rows_per_run, + reservation, + None, + false, + ) + .with_replay_headroom(true); + Ok(ReplayMergeFixture { + builder, + env, + pool, + pool_size, + metrics, + input_bytes, + }) +} + +/// Return additional spill metrics, checking final replay space and cleanup. +async fn merge_replay_runs( + run_count: usize, + rows_per_run: usize, + memory_batches: usize, +) -> Result<(usize, usize, usize)> { + let ReplayMergeFixture { + builder, + env, + pool, + pool_size, + metrics, + input_bytes, + } = replay_merge_fixture(run_count, rows_per_run, memory_batches, 0)?; + let schema = Arc::clone(&builder.schema); + let replay = MemoryConsumer::new("replay consumer").register(&pool); + let mut stream = builder.create_spillable_merge_stream(); + let mut batches = Vec::new(); + while let Some(batch) = stream.try_next().await? { + assert!(pool.reserved() <= pool_size / 2); + assert!(crate::spill::get_record_batch_memory_size(&batch) <= pool.reserved()); + // Check that another consumer can actually claim the replay allowance. + replay.try_grow(pool_size / 2)?; + replay.free(); + batches.push(batch); + } + let merged = concat_batches(&schema, &batches)?; + let expected = Int64Array::from_iter_values(0..(run_count * rows_per_run) as i64); + assert_eq!(merged.column(0).as_primitive::(), &expected); + assert_eq!(pool.reserved(), 0); + drop(stream); + assert_eq!(env.disk_manager.spilling_progress().active_files_count, 0); + assert_eq!(env.disk_manager.used_disk_space(), 0); + + Ok(( + metrics.spill_file_count.value() - run_count, + metrics.spilled_rows.value() - run_count * rows_per_run, + metrics.spilled_bytes.value() - input_bytes, + )) +} + +#[rstest::rstest] +#[case::intermediate_uses_full_pool(6, 16, 0, 2, 16)] +#[case::intermediate_holds_back_a_run(3, 16, 0, 1, 8)] +#[case::final_disables_read_ahead(3, 12, 0, 0, 6)] +#[case::fan_in_limited_intermediate(4, 8, 2, 2, 8)] +#[case::fan_in_limited_final(2, 8, 2, 0, 4)] +#[tokio::test] +async fn replay_headroom_depends_on_merge_phase( + #[case] run_count: usize, + #[case] memory_batches: usize, + #[case] max_fan_in: usize, + #[case] remaining_runs: usize, + #[case] reserved_batches: usize, +) -> Result<()> { + let ReplayMergeFixture { + mut builder, + env, + pool, + .. + } = replay_merge_fixture(run_count, 128, memory_batches, max_fan_in)?; + let batch_memory = builder.sorted_spill_files[0].0.max_record_batch_memory; + let MergeStep::Stream { stream, .. } = + builder.merge_sorted_runs_within_mem_limit(false)? + else { + panic!("the merge should fit without splitting a run"); + }; + assert_eq!(builder.sorted_spill_files.len(), remaining_runs); + assert_eq!(pool.reserved(), reserved_batches * batch_memory); + let batches: Vec = stream.try_collect().await?; + assert_eq!( + batches.iter().map(RecordBatch::num_rows).sum::(), + (run_count - remaining_runs) * 128 + ); + drop(builder); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.spilling_progress().active_files_count, 0); + assert_eq!(env.disk_manager.used_disk_space(), 0); + Ok(()) +} + +#[tokio::test] +async fn replay_headroom_keeps_split_retries_before_intermediate_merges() -> Result<()> { + let ReplayMergeFixture { + mut builder, pool, .. + } = replay_merge_fixture(3, 128, 6, 0)?; + assert!(matches!( + builder.merge_sorted_runs_within_mem_limit(false)?, + MergeStep::SplitThenRetry(_) + )); + assert_eq!(builder.sorted_spill_files.len(), 3); + assert_eq!(pool.reserved(), 0); + Ok(()) +} + +#[tokio::test] +async fn replay_headroom_preserves_indivisible_run_batch_limits() -> Result<()> { + let env = Arc::new(RuntimeEnv::default()); + let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)])); + let spill_manager = SpillManager::new( + Arc::clone(&env), + SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0), + Arc::clone(&schema), + ); + let values = (0..96) + .map(|value| format!("{value:03}{}", "x".repeat(1024))) + .collect::>(); + let mut spills = Vec::new(); + for run in 0..3 { + // Every batch is already indivisible, but the configured merge batch + // size is much larger. The split/retry path must discover the one-row + // limit before an intermediate pass can concatenate these batches. + let batches = values.iter().skip(run).step_by(3).map(|value| { + RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(StringArray::from(vec![value.as_str()]))], + ) + .map_err(Into::into) + }); + let (file, max_record_batch_memory) = spill_manager + .spill_record_batch_iter_and_return_max_batch_memory( + batches, + "indivisible replay input", + )? + .expect("a nonempty input must spill"); + spills.push(SortedSpillFile { + file, + max_record_batch_memory, + }); + } + let pool_size = 6 * spills[0].max_record_batch_memory; + let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_size)); + let reservation = MemoryConsumer::new("indivisible replay test").register(&pool); + let builder = MultiLevelMergeBuilder::new( + spill_manager, + Arc::clone(&schema), + spills, + vec![], + [PhysicalSortExpr::new_default(Arc::new(Column::new("x", 0)))].into(), + BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0), + 8192, + reservation, + None, + false, + ) + .with_replay_headroom(true); + let mut stream = builder.create_spillable_merge_stream(); + let mut batches = Vec::new(); + while let Some(batch) = stream.try_next().await? { + assert_eq!(batch.num_rows(), 1); + assert!(crate::spill::get_record_batch_memory_size(&batch) <= pool.reserved()); + assert!(pool.reserved() <= pool_size); + batches.push(batch); + } + let merged = concat_batches(&schema, &batches)?; + assert_eq!( + merged.column(0).as_string::(), + &StringArray::from(values) + ); + drop(stream); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.spilling_progress().active_files_count, 0); + assert_eq!(env.disk_manager.used_disk_space(), 0); + Ok(()) +} + +#[tokio::test] +async fn replay_headroom_does_not_rewrite_intermediate_runs_twice() -> Result<()> { + // The pool holds eight read-ahead inputs, or four plus equal replay headroom. + // Four intermediate merges of eight runs leave four runs for the final merge. + // Reserving headroom for every pass instead writes ten intermediate files, + // rewriting every input row twice before returning the final merge. + let (spill_count, spilled_rows, spilled_bytes) = + merge_replay_runs(32, 256, 32).await?; + println!( + "Intermediate spill: {spill_count} files, {spilled_rows} rows, {spilled_bytes} bytes" + ); + assert_eq!(spill_count, 4, "additional spill bytes: {spilled_bytes}"); + assert_eq!(spilled_rows, 32 * 256); + Ok(()) +} + +#[tokio::test] +async fn replay_headroom_is_restored_after_intermediate_split_retries() -> Result<()> { + // Two inputs need four batches of workspace. Three available batches force + // an intermediate split; the final merge then needs further splitting to + // leave replay headroom. Both phases must preserve every row and release + // all reservations and temporary files when the stream finishes. + let (spill_count, _, _) = merge_replay_runs(3, 128, 3).await?; + assert!(spill_count > 1, "the merge must spill and split runs"); + Ok(()) +} From 071e09a28c6496919c43cbdc94cf8a7378d4959b Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Thu, 17 Sep 2026 20:46:38 +0000 Subject: [PATCH 2/4] fix: preserve admitted spill merges when widening --- datafusion/physical-plan/src/sorts/builder.rs | 41 +- datafusion/physical-plan/src/sorts/merge.rs | 22 +- .../src/sorts/multi_level_merge.rs | 355 +++++++++++------- .../replay_headroom_tests.rs | 324 ++++++++++++++-- .../src/sorts/streaming_merge.rs | 202 +++++++++- 5 files changed, 758 insertions(+), 186 deletions(-) diff --git a/datafusion/physical-plan/src/sorts/builder.rs b/datafusion/physical-plan/src/sorts/builder.rs index 89763efc4d75c..1c38db9b6d162 100644 --- a/datafusion/physical-plan/src/sorts/builder.rs +++ b/datafusion/physical-plan/src/sorts/builder.rs @@ -146,6 +146,42 @@ impl BatchBuilder { &self.schema } + /// Release fully consumed batches after a merge drains at an input boundary. + /// Keeping their dictionaries can otherwise enlarge the next output even + /// though none of its rows refer to those batches. + pub(super) fn discard_consumed_batches(&mut self) { + assert!(self.indices.is_empty()); + self.retain_current_batches(false); + self.release_unused_memory(); + } + + fn retain_current_batches(&mut self, keep_consumed: bool) { + let mut batch_idx = 0; + let mut retained = 0; + self.batches.retain(|(stream_idx, batch)| { + let stream_cursor = &mut self.cursors[*stream_idx]; + let retain = stream_cursor.batch_idx == batch_idx + && (keep_consumed || stream_cursor.row_idx < batch.num_rows()); + batch_idx += 1; + + if retain { + stream_cursor.batch_idx = retained; + retained += 1; + } else { + self.batches_mem_used -= get_record_batch_memory_size(batch); + } + retain + }); + } + + fn release_unused_memory(&mut self) { + // Keep the initial grant to avoid re-admission between output batches. + let target = self.batches_mem_used.max(self.initial_reservation); + if self.reservation.size() > target { + self.reservation.shrink(self.reservation.size() - target); + } + } + /// Try to interleave all columns using the given index slice. fn try_interleave_columns( &self, @@ -183,10 +219,7 @@ impl BatchBuilder { // Release excess memory back to the pool, but never shrink below // initial_reservation to maintain the anti-starvation guarantee // for the merge phase. - let target = self.batches_mem_used.max(self.initial_reservation); - if self.reservation.size() > target { - self.reservation.shrink(self.reservation.size() - target); - } + self.release_unused_memory(); RecordBatch::try_new(Arc::clone(&self.schema), columns).map_err(Into::into) } diff --git a/datafusion/physical-plan/src/sorts/merge.rs b/datafusion/physical-plan/src/sorts/merge.rs index d08bd55de91ca..2cbd93751d9aa 100644 --- a/datafusion/physical-plan/src/sorts/merge.rs +++ b/datafusion/physical-plan/src/sorts/merge.rs @@ -86,6 +86,11 @@ pub(crate) struct SortPreservingMergeStream { /// Target batch size batch_size: usize, + /// Emit pending rows before replacing an exhausted input batch. Intermediate + /// spill merges use this to avoid collecting multiple short batches from an + /// input into an output batch larger than their reserved workspace. + flush_on_input_batch_boundary: bool, + /// Cursors for each input partition. `None` means the input is exhausted cursors: Vec>>, @@ -146,11 +151,17 @@ impl SortPreservingMergeStream { poll_reset_epochs: vec![0; stream_count], loser_tree: vec![], batch_size, + flush_on_input_batch_boundary: false, fetch, produced: 0, } } + pub(super) fn with_flush_on_input_batch_boundary(mut self, flush: bool) -> Self { + self.flush_on_input_batch_boundary = flush; + self + } + pub(crate) fn into_stream(self) -> SendableRecordBatchStream where C: 'static, @@ -198,7 +209,7 @@ impl SortPreservingMergeStream { async fn flush_in_progress( &mut self, - mut emitter: TryEmitter, + emitter: &mut TryEmitter, ) -> Result<()> { if self.in_progress.is_empty() { return Ok(()); @@ -284,6 +295,13 @@ impl SortPreservingMergeStream { ); drop(timer); + if self.flush_on_input_batch_boundary { + // Drain every pending row, including partial output + // from offset-overflow recovery, before accepting a + // replacement batch from this input. + self.flush_in_progress(&mut emitter).await?; + self.in_progress.discard_consumed_batches(); + } poll_fn(|cx| self.maybe_poll_stream(cx, winner_stream)).await?; timer = elapsed_compute.timer(); } @@ -294,7 +312,7 @@ impl SortPreservingMergeStream { } // 4. Flush any remaining rows in `self.in_progress` - self.flush_in_progress(emitter).await?; + self.flush_in_progress(&mut emitter).await?; Ok(()) }) diff --git a/datafusion/physical-plan/src/sorts/multi_level_merge.rs b/datafusion/physical-plan/src/sorts/multi_level_merge.rs index ceaec13ff35b6..80d029cd85837 100644 --- a/datafusion/physical-plan/src/sorts/multi_level_merge.rs +++ b/datafusion/physical-plan/src/sorts/multi_level_merge.rs @@ -159,6 +159,8 @@ pub(crate) struct MultiLevelMergeBuilder { merge_pool: Option>, /// Leave memory for the aggregate consuming the merged spill rows. reserve_replay_headroom: bool, + /// Stop widening after a failed intermediate write. + widen_intermediate_merges: bool, fetch: Option, enable_round_robin_tie_breaker: bool, } @@ -199,6 +201,7 @@ impl MultiLevelMergeBuilder { reservation, merge_pool: None, reserve_replay_headroom: false, + widen_intermediate_merges: true, enable_round_robin_tie_breaker, fetch, } @@ -226,13 +229,14 @@ impl MultiLevelMergeBuilder { async fn create_stream(mut self) -> Result { let mut allow_minimum_without_headroom = false; loop { - let (mut stream, batch_size_limit) = match self + let (mut stream, mut batch_size_limit, retry) = match self .merge_sorted_runs_within_mem_limit(allow_minimum_without_headroom)? { MergeStep::Stream { stream, batch_size_limit, - } => (stream, batch_size_limit), + retry, + } => (stream, batch_size_limit, retry), MergeStep::SplitThenRetry(index) => { // Couldn't reserve memory for the minimum of 2 streams. Re-spill // the larger of the two we're trying to merge with half its batch @@ -274,15 +278,46 @@ impl MultiLevelMergeBuilder { return Ok(stream); } - // Need to sort to a spill file - let Some((spill_file, max_record_batch_memory)) = self + // A wider merge can reduce total writes while increasing peak disk + // usage. Keep its inputs and admitted buffers until the writer has + // finished, so a failed write can retry the original smaller merge. + let mut result = self .spill_manager .spill_record_batch_stream_and_return_max_batch_memory( &mut stream, "MultiLevelMergeBuilder intermediate spill", ) - .await? - else { + .await; + drop(stream); + // A successful write drops the backups before the next selection. + if let Some(retry) = retry + && result.is_err() + { + self.widen_intermediate_merges = false; + let IntermediateMergeRetry { + mut spills, + original_count, + buffer_size, + original_memory, + reservation, + } = retry; + self.sorted_spill_files + .splice(0..0, spills.drain(original_count..)); + reservation.shrink(reservation.size() - original_memory); + // No pool re-admission is needed: the successful original + // grant has remained live, even if the writer failed at EOF. + let (mut stream, limit) = + self.merge_selected_runs(spills, buffer_size, reservation, false)?; + batch_size_limit = limit; + result = self + .spill_manager + .spill_record_batch_stream_and_return_max_batch_memory( + &mut stream, + "MultiLevelMergeBuilder intermediate spill retry", + ) + .await; + } + let Some((spill_file, max_record_batch_memory)) = result? else { continue; }; @@ -314,6 +349,7 @@ impl MultiLevelMergeBuilder { Ok(MergeStep::Stream { stream: self.observe_output(empty_stream), batch_size_limit: self.batch_size, + retry: None, }) } @@ -323,6 +359,7 @@ impl MultiLevelMergeBuilder { Ok(MergeStep::Stream { stream: self.observe_output(output_stream), batch_size_limit: self.batch_size, + retry: None, }) } @@ -338,6 +375,7 @@ impl MultiLevelMergeBuilder { Ok(MergeStep::Stream { stream: self.observe_output(output_stream), batch_size_limit: batch_size, + retry: None, }) } @@ -354,8 +392,10 @@ impl MultiLevelMergeBuilder { true, true, self.batch_size, + false, )?, batch_size_limit: self.batch_size, + retry: None, }) } @@ -373,127 +413,161 @@ impl MultiLevelMergeBuilder { let minimum_number_of_required_streams = 2_usize.saturating_sub(self.sorted_streams.len()); - let total_spill_files = self.sorted_spill_files.len(); - let mut selection = self.get_sorted_spill_files_to_merge( + let selection = self.get_sorted_spill_files_to_merge( 2, - // we must have at least 2 streams to merge minimum_number_of_required_streams, &mut memory_reservation, allow_minimum_without_headroom, - self.reserve_replay_headroom, - usize::MAX, )?; - - // Prefer a final merge with replay headroom. Otherwise an - // intermediate merge writes back to disk and can use the full - // pool. Preserve split requests: splitting also caps the merge - // output size when input batches are shorter than batch_size. - // Hold one run back so this larger merge cannot feed replay, - // but only if enough inputs remain to make progress. - let is_intermediate = matches!( - &selection, - SpillFilesToMerge::Ready(spills, _) if spills.len() < total_spill_files - ); - if self.reserve_replay_headroom - && is_intermediate - && total_spill_files > minimum_number_of_required_streams - { - if let SpillFilesToMerge::Ready(spills, _) = selection { - self.sorted_spill_files.splice(0..0, spills); - } - memory_reservation.free(); - selection = self.get_sorted_spill_files_to_merge( - 2, - minimum_number_of_required_streams, - &mut memory_reservation, - false, - false, - total_spill_files - 1, - )?; - } - - let (sorted_spill_files, buffer_size) = match selection { - SpillFilesToMerge::Ready(sorted_spill_files, buffer_size) => { - (sorted_spill_files, buffer_size) + let (mut spills, buffer_size) = match selection { + SpillFilesToMerge::Ready(spills, buffer_size) => { + (spills, buffer_size) } - // Not enough memory to seat 2 streams. Re-spill the blocking file - // smaller and retry. `get_sorted_spill_files_to_merge` already freed - // the reservation and `self.sorted_streams` is untouched, so the - // retry starts clean. SpillFilesToMerge::SplitThenRetry(index) => { return Ok(MergeStep::SplitThenRetry(index)); } }; - // Don't account for existing streams memory - // as we are not holding the memory for them - let mut sorted_streams = mem::take(&mut self.sorted_streams); - - let is_only_merging_memory_streams = sorted_spill_files.is_empty(); - - // If no spill files were selected (e.g. all too large for - // available memory but enough in-memory streams exist), - // return the pre-reserved bytes to self.reservation so - // create_new_merge_sort can transfer them to the merge - // stream's BatchBuilder. - if is_only_merging_memory_streams { - mem::swap(&mut self.reservation, &mut memory_reservation); + let original_count = spills.len(); + let original_memory = memory_reservation.size(); + if self.reserve_replay_headroom + && self.widen_intermediate_merges + && self.sorted_streams.is_empty() + && !spills.is_empty() + { + // Extend the admitted selection without releasing its grant + // or changing read-ahead. Keep one run for the final merge, + // which must still leave space for aggregate replay. + let max_fan_in = effective_spill_merge_fan_in( + self.spill_manager + .env() + .disk_manager + .max_spill_merge_fan_in(), + ); + let mut extra = 0; + for (spill, _) in self + .sorted_spill_files + .iter() + .take(self.sorted_spill_files.len().saturating_sub(1)) + { + if spills.len() + extra >= max_fan_in { + break; + } + let additional = get_reserved_bytes_for_record_batch_size( + spill.max_record_batch_memory, + spill.max_record_batch_memory, + ) * buffer_size; + if memory_reservation.try_grow(additional).is_err() { + break; + } + extra += 1; + } + spills.extend(self.sorted_spill_files.drain(..extra)); } + let reservation = Arc::new(memory_reservation); + let retry = + (spills.len() > original_count).then(|| IntermediateMergeRetry { + spills: spills + .iter() + .map(|(spill, limit)| { + ( + SortedSpillFile { + file: Arc::clone(&spill.file), + max_record_batch_memory: spill + .max_record_batch_memory, + }, + *limit, + ) + }) + .collect(), + original_count, + buffer_size, + original_memory, + reservation: Arc::clone(&reservation), + }); + let (stream, batch_size_limit) = self.merge_selected_runs( + spills, + buffer_size, + reservation, + retry.is_some(), + )?; + Ok(MergeStep::Stream { + stream, + batch_size_limit, + retry, + }) + } + } + } - // Cap the merge output at the smallest limit among the runs we're - // about to merge. Runs that were shrunk for skew carry a smaller limit, - // if none do, every run carries `self.batch_size` and the merge runs at - // the full batch size. The output stream is tagged with the same limit - // (see the `MergeStep::Stream` returns below) so a re-spilled - // intermediate run stays shrunk and won't rebuild an oversized batch on - // a later pass. - let mut output_batch_size = self.batch_size; - for (spill, batch_size_limit) in sorted_spill_files { - let stream = self - .spill_manager - .clone() - .with_batch_read_buffer_capacity(buffer_size) - .read_spill_as_stream( - spill.file, - Some(spill.max_record_batch_memory), - )?; - output_batch_size = output_batch_size.min(batch_size_limit); - sorted_streams.push(stream); - } - let merge_sort_stream = self.create_new_merge_sort( + /// Build a stream from an already admitted selection. The reservation can + /// also be held by a retry guard until an intermediate writer finishes. + fn merge_selected_runs( + &mut self, + sorted_spill_files: Vec<(SortedSpillFile, usize)>, + buffer_size: usize, + memory_reservation: Arc, + flush_on_input_batch_boundary: bool, + ) -> Result<(SendableRecordBatchStream, usize)> { + // Don't account for existing streams memory + // as we are not holding the memory for them + let mut sorted_streams = mem::take(&mut self.sorted_streams); + + let is_only_merging_memory_streams = sorted_spill_files.is_empty(); + + // If no spill files were selected (e.g. all too large for + // available memory but enough in-memory streams exist), + // return the pre-reserved bytes to self.reservation so + // create_new_merge_sort can transfer them to the merge + // stream's BatchBuilder. + if is_only_merging_memory_streams { + self.reservation = Arc::try_unwrap(memory_reservation) + .expect("in-memory merges do not retain a spill retry reservation"); + return Ok(( + self.create_new_merge_sort( sorted_streams, - // If we have no sorted spill files left, this is the last run self.sorted_spill_files.is_empty(), - is_only_merging_memory_streams, - output_batch_size, - )?; - - // If we're only merging memory streams, we don't need to attach the memory reservation - // as it's empty - if is_only_merging_memory_streams { - assert_eq!( - memory_reservation.size(), - 0, - "when only merging memory streams, we should not have any memory reservation and let the merge sort handle the memory" - ); + true, + self.batch_size, + false, + )?, + self.batch_size, + )); + } - Ok(MergeStep::Stream { - stream: merge_sort_stream, - batch_size_limit: output_batch_size, - }) - } else { - // Attach the memory reservation to the stream to make sure we have enough memory - // throughout the merge process as we bypassed the memory pool for the merge sort stream - Ok(MergeStep::Stream { - stream: Box::pin(StreamAttachedReservation::new( - merge_sort_stream, - memory_reservation, - )), - batch_size_limit: output_batch_size, - }) - } - } + // Cap the merge output at the smallest limit among the runs we're + // about to merge. Runs that were shrunk for skew carry a smaller limit, + // if none do, every run carries `self.batch_size` and the merge runs at + // the full batch size. The output stream is tagged with the same limit + // (see the `MergeStep::Stream` returns below) so a re-spilled + // intermediate run stays shrunk and won't rebuild an oversized batch on + // a later pass. + let mut output_batch_size = self.batch_size; + for (spill, batch_size_limit) in sorted_spill_files { + let stream = self + .spill_manager + .clone() + .with_batch_read_buffer_capacity(buffer_size) + .read_spill_as_stream(spill.file, Some(spill.max_record_batch_memory))?; + output_batch_size = output_batch_size.min(batch_size_limit); + sorted_streams.push(stream); } + let merge_sort_stream = self.create_new_merge_sort( + sorted_streams, + // If we have no sorted spill files left, this is the last run + self.sorted_spill_files.is_empty(), + is_only_merging_memory_streams, + output_batch_size, + flush_on_input_batch_boundary, + )?; + + Ok(( + Box::pin(StreamAttachedReservation::new( + merge_sort_stream, + memory_reservation, + )), + output_batch_size, + )) } fn create_new_merge_sort( @@ -502,11 +576,13 @@ impl MultiLevelMergeBuilder { is_output: bool, all_in_memory: bool, output_batch_size: usize, + flush_on_input_batch_boundary: bool, ) -> Result { let mut builder = StreamingMergeBuilder::new() .with_schema(Arc::clone(&self.schema)) .with_expressions(&self.expr) .with_batch_size(output_batch_size) + .with_flush_on_input_batch_boundary(flush_on_input_batch_boundary) .with_fetch(self.fetch) .with_metrics(if is_output { // Only add the metrics to the last run @@ -544,8 +620,6 @@ impl MultiLevelMergeBuilder { minimum_number_of_required_streams: usize, reservation: &mut MemoryReservation, allow_minimum_without_headroom: bool, - reserve_replay_headroom: bool, - max_files: usize, ) -> Result { assert_ne!(buffer_len, 0, "Buffer length must be greater than 0"); let mut number_of_spills_to_read_for_current_phase = 0; @@ -554,8 +628,7 @@ impl MultiLevelMergeBuilder { .env() .disk_manager .max_spill_merge_fan_in(); - let max_spill_files = - effective_spill_merge_fan_in(configured_fan_in).min(max_files); + let max_spill_files = effective_spill_merge_fan_in(configured_fan_in); // Track total memory needed for spill file buffers. When the // reservation has pre-reserved bytes (from sort_spill_reservation_bytes), // those bytes cover the first N spill files without additional pool @@ -586,7 +659,7 @@ impl MultiLevelMergeBuilder { && buffer_len == 1 && number_of_spills_to_read_for_current_phase < minimum_number_of_required_streams; - let check_headroom = reserve_replay_headroom && !skip_headroom; + let check_headroom = self.reserve_replay_headroom && !skip_headroom; let admission = if check_headroom { // Ask the pool for merge buffers plus equal replay space, then // return the spare bytes before exposing the merge stream. @@ -620,8 +693,6 @@ impl MultiLevelMergeBuilder { minimum_number_of_required_streams, reservation, allow_minimum_without_headroom, - reserve_replay_headroom, - max_files, ); } @@ -657,7 +728,7 @@ impl MultiLevelMergeBuilder { } } - if reserve_replay_headroom { + if self.reserve_replay_headroom { // `total_needed` may include a rejected candidate. Keep only the // buffers that were admitted, releasing temporary replay headroom. reservation.shrink(reservation.size() - accepted_memory); @@ -858,6 +929,16 @@ enum SpillFilesToMerge { SplitThenRetry(usize), } +/// Inputs and the original grant retained until a widened intermediate write +/// succeeds. Keeping the grant shared also covers writer finalization after EOF. +struct IntermediateMergeRetry { + spills: Vec<(SortedSpillFile, usize)>, + original_count: usize, + buffer_size: usize, + original_memory: usize, + reservation: Arc, +} + /// What one iteration of the multi-level merge loop should do next. enum MergeStep { /// A merged stream is ready to be consumed (and possibly spilled back). @@ -869,6 +950,7 @@ enum MergeStep { /// case it is that run's smaller limit so the re-spilled result stays capped /// and can't rebuild an oversized batch. batch_size_limit: usize, + retry: Option, }, /// Re-spill the spill file at this index smaller, then retry the merge step. SplitThenRetry(usize), @@ -894,14 +976,17 @@ fn effective_spill_merge_fan_in(configured_fan_in: usize) -> usize { struct StreamAttachedReservation { stream: SendableRecordBatchStream, - reservation: MemoryReservation, + reservation: Option>, } impl StreamAttachedReservation { - fn new(stream: SendableRecordBatchStream, reservation: MemoryReservation) -> Self { + fn new( + stream: SendableRecordBatchStream, + reservation: Arc, + ) -> Self { Self { stream, - reservation, + reservation: Some(reservation), } } } @@ -921,12 +1006,12 @@ impl Stream for StreamAttachedReservation { Some(Ok(batch)) => Poll::Ready(Some(Ok(batch))), Some(Err(err)) => { // Had an error so drop the data - self.reservation.free(); + self.reservation = None; Poll::Ready(Some(Err(err))) } None => { // Stream is done so free the memory - self.reservation.free(); + self.reservation = None; Poll::Ready(None) } @@ -1243,15 +1328,8 @@ mod tests { // Actual batches contain one row even though the nominal size is 8192. assert_eq!(builder.sorted_spill_files[0].1, 1); let mut reservation = builder.reservation.new_empty(); - let SpillFilesToMerge::Ready(spills, buffer_len) = builder - .get_sorted_spill_files_to_merge( - 2, - 2, - &mut reservation, - true, - true, - usize::MAX, - )? + let SpillFilesToMerge::Ready(spills, buffer_len) = + builder.get_sorted_spill_files_to_merge(2, 2, &mut reservation, true)? else { panic!("minimum merge should fit the pool"); }; @@ -1291,15 +1369,8 @@ mod tests { let mut builder = build_merge_builder(spill_manager, schema, spills, &pool, 1) .with_replay_headroom(true); let mut reservation = builder.reservation.new_empty(); - let SpillFilesToMerge::Ready(spills, buffer_len) = builder - .get_sorted_spill_files_to_merge( - 1, - 2, - &mut reservation, - false, - true, - usize::MAX, - )? + let SpillFilesToMerge::Ready(spills, buffer_len) = + builder.get_sorted_spill_files_to_merge(1, 2, &mut reservation, false)? else { panic!("two streams and replay headroom should fit the pool"); }; @@ -1724,8 +1795,6 @@ mod tests { 2, &mut merge_reservation, false, - false, - usize::MAX, )? { SpillFilesToMerge::Ready(spills, buffer_len) => (spills, buffer_len), SpillFilesToMerge::SplitThenRetry(index) => { diff --git a/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs b/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs index 6ba036e2302bd..0603389c52290 100644 --- a/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs +++ b/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs @@ -25,6 +25,7 @@ use datafusion_execution::memory_pool::{GreedyMemoryPool, MemoryConsumer, Memory use datafusion_execution::runtime_env::{RuntimeEnv, RuntimeEnvBuilder}; use datafusion_physical_expr::expressions::Column; use datafusion_physical_expr_common::metrics::{ExecutionPlanMetricsSet, SpillMetrics}; +use std::sync::atomic::{AtomicBool, Ordering}; struct ReplayMergeFixture { builder: MultiLevelMergeBuilder, @@ -35,6 +36,28 @@ struct ReplayMergeFixture { input_bytes: usize, } +fn replay_merge_builder( + spill_manager: SpillManager, + schema: SchemaRef, + spills: Vec, + pool: &Arc, + batch_size: usize, +) -> MultiLevelMergeBuilder { + MultiLevelMergeBuilder::new( + spill_manager, + schema, + spills, + vec![], + [PhysicalSortExpr::new_default(Arc::new(Column::new("x", 0)))].into(), + BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0), + batch_size, + MemoryConsumer::new("replay headroom test").register(pool), + None, + false, + ) + .with_replay_headroom(true) +} + /// Create equally sized, interleaved input runs. fn replay_merge_fixture( run_count: usize, @@ -69,21 +92,8 @@ fn replay_merge_fixture( let input_bytes = metrics.spilled_bytes.value(); let pool_size = memory_batches * spills[0].max_record_batch_memory; let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_size)); - let reservation = MemoryConsumer::new("replay headroom test").register(&pool); - let expr = [PhysicalSortExpr::new_default(Arc::new(Column::new("x", 0)))].into(); - let builder = MultiLevelMergeBuilder::new( - spill_manager, - Arc::clone(&schema), - spills, - vec![], - expr, - BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0), - rows_per_run, - reservation, - None, - false, - ) - .with_replay_headroom(true); + let builder = + replay_merge_builder(spill_manager, schema, spills, &pool, rows_per_run); Ok(ReplayMergeFixture { builder, env, @@ -139,7 +149,7 @@ async fn merge_replay_runs( #[case::intermediate_uses_full_pool(6, 16, 0, 2, 16)] #[case::intermediate_holds_back_a_run(3, 16, 0, 1, 8)] #[case::final_disables_read_ahead(3, 12, 0, 0, 6)] -#[case::fan_in_limited_intermediate(4, 8, 2, 2, 8)] +#[case::fan_in_limited_intermediate(4, 8, 2, 2, 4)] #[case::fan_in_limited_final(2, 8, 2, 0, 4)] #[tokio::test] async fn replay_headroom_depends_on_merge_phase( @@ -226,20 +236,8 @@ async fn replay_headroom_preserves_indivisible_run_batch_limits() -> Result<()> } let pool_size = 6 * spills[0].max_record_batch_memory; let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_size)); - let reservation = MemoryConsumer::new("indivisible replay test").register(&pool); - let builder = MultiLevelMergeBuilder::new( - spill_manager, - Arc::clone(&schema), - spills, - vec![], - [PhysicalSortExpr::new_default(Arc::new(Column::new("x", 0)))].into(), - BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0), - 8192, - reservation, - None, - false, - ) - .with_replay_headroom(true); + let builder = + replay_merge_builder(spill_manager, Arc::clone(&schema), spills, &pool, 8192); let mut stream = builder.create_spillable_merge_stream(); let mut batches = Vec::new(); while let Some(batch) = stream.try_next().await? { @@ -286,3 +284,269 @@ async fn replay_headroom_is_restored_after_intermediate_split_retries() -> Resul assert!(spill_count > 1, "the merge must spill and split runs"); Ok(()) } + +#[tokio::test] +async fn intermediate_merge_preserves_short_batch_replay() -> Result<()> { + let env = Arc::new(RuntimeEnv::default()); + let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)])); + let manager = SpillManager::new( + Arc::clone(&env), + SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0), + Arc::clone(&schema), + ); + let values = (0..18) + .map(|value| format!("{value:04}{}", "x".repeat(1024))) + .collect::>(); + let mut spills = Vec::new(); + for run in 0..6 { + let batches = values.iter().skip(run).step_by(6).map(|value| { + RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(StringArray::from(vec![value.as_str()]))], + ) + .map_err(Into::into) + }); + let (file, max_record_batch_memory) = manager + .spill_record_batch_iter_and_return_max_batch_memory( + batches, + "short replay input", + )? + .expect("a nonempty input must spill"); + spills.push(SortedSpillFile { + file, + max_record_batch_memory, + }); + } + // A two-input merge fits with headroom, so no initial split discovers the + // actual one-row input size. Widening to four runs must not combine their + // twelve rows into an intermediate batch too large to split or replay. + let pool: Arc = Arc::new(GreedyMemoryPool::new( + 16 * spills[0].max_record_batch_memory, + )); + let builder = replay_merge_builder(manager, Arc::clone(&schema), spills, &pool, 8192); + let batches: Vec = builder + .create_spillable_merge_stream() + .try_collect() + .await?; + let merged = concat_batches(&schema, &batches)?; + assert_eq!( + merged.column(0).as_string::(), + &StringArray::from(values) + ); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.used_disk_space(), 0); + Ok(()) +} + +#[rstest::rstest] +#[case::sufficient_quota(true)] +#[case::insufficient_quota(false)] +#[tokio::test] +async fn intermediate_merge_preserves_disk_quota( + #[case] sufficient_quota: bool, +) -> Result<()> { + check_intermediate_merge_disk_quota(sufficient_quota).await +} + +#[rstest::rstest] +#[case::sufficient_quota(true)] +#[case::insufficient_quota(false)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn intermediate_merge_preserves_disk_quota_with_read_ahead_tasks( + #[case] sufficient_quota: bool, +) -> Result<()> { + check_intermediate_merge_disk_quota(sufficient_quota).await +} + +async fn check_intermediate_merge_disk_quota(sufficient_quota: bool) -> Result<()> { + const RUN_COUNT: usize = 4; + const BATCHES_PER_RUN: usize = 100; + const BATCH_SIZE: usize = 128; + let env = Arc::new(RuntimeEnv::default()); + let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)])); + let manager = SpillManager::new( + Arc::clone(&env), + SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0), + Arc::clone(&schema), + ); + let mut spills = Vec::new(); + for run in 0..RUN_COUNT { + let batches = (0..BATCHES_PER_RUN).map(|batch_index| { + let values = + Int64Array::from_iter_values((0..BATCH_SIZE).map(|row| { + ((batch_index * BATCH_SIZE + row) * RUN_COUNT + run) as i64 + })); + RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(values)]) + .map_err(Into::into) + }); + let (file, max_record_batch_memory) = manager + .spill_record_batch_iter_and_return_max_batch_memory( + batches, + "disk quota replay input", + )? + .expect("a nonempty input must spill"); + spills.push(SortedSpillFile { + file, + max_record_batch_memory, + }); + } + // Interleaving keeps all inputs live while the output grows. This quota + // permits two-input intermediate merges but not a three-input merge. The + // smaller quota cannot accommodate either and must fail without leaking. + let initial_disk = env.disk_manager.used_disk_space(); + let quota_numerator = if sufficient_quota { 13 } else { 9 }; + env.disk_manager + .set_max_temp_directory_size(initial_disk * quota_numerator / 8)?; + let pool: Arc = Arc::new(GreedyMemoryPool::new( + 16 * spills[0].max_record_batch_memory, + )); + let builder = + replay_merge_builder(manager, Arc::clone(&schema), spills, &pool, BATCH_SIZE); + let result: Result> = + builder.create_spillable_merge_stream().try_collect().await; + if sufficient_quota { + let batches = result?; + let merged = concat_batches(&schema, &batches)?; + let expected = Int64Array::from_iter_values( + 0..(RUN_COUNT * BATCHES_PER_RUN * BATCH_SIZE) as i64, + ); + assert_eq!(merged.column(0).as_primitive::(), &expected); + } else { + let error = result.expect_err("the quota must reject even a two-input merge"); + assert!( + error.to_string().contains("max_temp_directory_size"), + "expected a disk quota error, got {error}" + ); + } + // On a multithreaded runtime, dropping a failed merge cancels background + // read-ahead tasks. Allow their cancellation to release input file handles. + tokio::time::timeout(std::time::Duration::from_secs(5), async { + while pool.reserved() != 0 + || env.disk_manager.used_disk_space() != 0 + || env.disk_manager.spilling_progress().active_files_count != 0 + { + tokio::task::yield_now().await; + } + }) + .await + .expect("merge completion must release reservations and temporary files"); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.used_disk_space(), 0); + Ok(()) +} + +#[rstest::rstest] +#[case::with_read_ahead(4, 16, 0, 15)] +#[case::without_read_ahead_under_pressure(5, 12, 6, 5)] +#[test] +fn intermediate_merge_keeps_admitted_buffers( + #[case] run_count: usize, + #[case] memory_batches: usize, + #[case] contender_headroom_batches: usize, + #[case] contender_handoff_batches: usize, +) -> Result<()> { + // Model another operator acquiring pool capacity immediately after the + // merge gives up its already-admitted buffers. All allocations use the + // underlying pool's normal fallible admission. + // With read-ahead disabled, the contender first takes the released replay + // headroom. Retrying admission with read-ahead enabled would then fail to + // seat two inputs and relinquish the usable three-input merge's grant. + #[derive(Debug)] + struct HandoffPool { + inner: Arc, + contender: MemoryReservation, + contender_headroom_bytes: usize, + contender_handoff_bytes: usize, + armed: AtomicBool, + handed_off: AtomicBool, + saw_headroom_release: AtomicBool, + } + + impl std::fmt::Display for HandoffPool { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "HandoffPool") + } + } + + impl MemoryPool for HandoffPool { + fn name(&self) -> &str { + "HandoffPool" + } + + fn grow(&self, reservation: &MemoryReservation, additional: usize) { + self.inner.grow(reservation, additional); + } + + fn try_grow( + &self, + reservation: &MemoryReservation, + additional: usize, + ) -> Result<()> { + self.inner.try_grow(reservation, additional) + } + + fn shrink(&self, reservation: &MemoryReservation, subtractive: usize) { + self.inner.shrink(reservation, subtractive); + if self.armed.load(Ordering::SeqCst) && subtractive != 0 { + if reservation.size() != 0 { + if !self.saw_headroom_release.swap(true, Ordering::SeqCst) { + self.contender + .try_grow(self.contender_headroom_bytes) + .unwrap(); + } + } else if self.saw_headroom_release.load(Ordering::SeqCst) + && self.armed.swap(false, Ordering::SeqCst) + { + // Releasing replay headroom leaves the admitted buffers + // intact. Only a later full free lets this allocation fit. + self.contender + .try_grow(self.contender_handoff_bytes) + .unwrap(); + self.handed_off.store(true, Ordering::SeqCst); + } + } + } + + fn reserved(&self) -> usize { + self.inner.reserved() + } + } + + let ReplayMergeFixture { + mut builder, + env, + pool: inner, + .. + } = replay_merge_fixture(run_count, 128, memory_batches, 0)?; + let batch_memory = builder.sorted_spill_files[0].0.max_record_batch_memory; + let contender = MemoryConsumer::new("competing operator").register(&inner); + let hooked = Arc::new(HandoffPool { + inner, + contender, + contender_headroom_bytes: contender_headroom_batches * batch_memory, + contender_handoff_bytes: contender_handoff_batches * batch_memory, + armed: AtomicBool::new(true), + handed_off: AtomicBool::new(false), + saw_headroom_release: AtomicBool::new(false), + }); + let pool: Arc = Arc::clone(&hooked) as Arc; + builder.reservation = MemoryConsumer::new("replay headroom test").register(&pool); + let admission = builder.merge_sorted_runs_within_mem_limit(false); + // Disable the scheduling hook before normal stream and builder cleanup. + hooked.armed.store(false, Ordering::SeqCst); + let MergeStep::Stream { stream, .. } = admission? else { + panic!("a previously admitted minimum merge must remain usable"); + }; + assert!(hooked.saw_headroom_release.load(Ordering::SeqCst)); + assert!(!hooked.handed_off.load(Ordering::SeqCst)); + assert_eq!( + hooked.contender.size(), + contender_headroom_batches * batch_memory + ); + drop(stream); + drop(builder); + hooked.contender.free(); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.used_disk_space(), 0); + Ok(()) +} diff --git a/datafusion/physical-plan/src/sorts/streaming_merge.rs b/datafusion/physical-plan/src/sorts/streaming_merge.rs index 61bb7742eecfa..9f170493b2aa1 100644 --- a/datafusion/physical-plan/src/sorts/streaming_merge.rs +++ b/datafusion/physical-plan/src/sorts/streaming_merge.rs @@ -43,7 +43,7 @@ macro_rules! primitive_merge_helper { } macro_rules! merge_helper { - ($t:ty, $sort:ident, $streams:ident, $schema:ident, $tracking_metrics:ident, $batch_size:ident, $fetch:ident, $reservation:ident, $enable_round_robin_tie_breaker:ident) => {{ + ($t:ty, $sort:ident, $streams:ident, $schema:ident, $tracking_metrics:ident, $batch_size:ident, $fetch:ident, $reservation:ident, $enable_round_robin_tie_breaker:ident, $flush_on_input_batch_boundary:ident) => {{ let streams = FieldCursorStream::<$t>::new($sort, $streams, $reservation.new_empty()); return Ok(SortPreservingMergeStream::new( @@ -55,6 +55,7 @@ macro_rules! merge_helper { $reservation, $enable_round_robin_tie_breaker, ) + .with_flush_on_input_batch_boundary($flush_on_input_batch_boundary) .into_stream()); }}; } @@ -98,6 +99,7 @@ pub struct StreamingMergeBuilder<'a> { merge_pool: Option>, /// Leave memory for the aggregate consuming the merged spill rows. reserve_replay_headroom: bool, + flush_on_input_batch_boundary: bool, enable_round_robin_tie_breaker: bool, } @@ -170,6 +172,12 @@ impl<'a> StreamingMergeBuilder<'a> { self } + /// Bound intermediate output to rows from one batch of each input. + pub(super) fn with_flush_on_input_batch_boundary(mut self, flush: bool) -> Self { + self.flush_on_input_batch_boundary = flush; + self + } + /// See [SortPreservingMergeExec::with_round_robin_repartition] for more /// information. /// @@ -204,6 +212,7 @@ impl<'a> StreamingMergeBuilder<'a> { reservation, merge_pool, reserve_replay_headroom, + flush_on_input_batch_boundary, fetch, expressions, enable_round_robin_tie_breaker, @@ -265,12 +274,12 @@ impl<'a> StreamingMergeBuilder<'a> { let sort = expressions[0].clone(); let data_type = sort.expr.data_type(schema.as_ref())?; downcast_primitive! { - data_type => (primitive_merge_helper, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker), - DataType::Utf8 => merge_helper!(StringArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker) - DataType::Utf8View => merge_helper!(StringViewArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker) - DataType::LargeUtf8 => merge_helper!(LargeStringArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker) - DataType::Binary => merge_helper!(BinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker) - DataType::LargeBinary => merge_helper!(LargeBinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker) + data_type => (primitive_merge_helper, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary), + DataType::Utf8 => merge_helper!(StringArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) + DataType::Utf8View => merge_helper!(StringViewArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) + DataType::LargeUtf8 => merge_helper!(LargeStringArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) + DataType::Binary => merge_helper!(BinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) + DataType::LargeBinary => merge_helper!(LargeBinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) _ => {} } } @@ -290,18 +299,22 @@ impl<'a> StreamingMergeBuilder<'a> { reservation, enable_round_robin_tie_breaker, ) + .with_flush_on_input_batch_boundary(flush_on_input_batch_boundary) .into_stream()) } } #[cfg(test)] mod tests { + use crate::spill::spill_manager::GetSlicedSize; use crate::{common::collect, stream::RecordBatchStreamAdapter}; use std::sync::Arc; use super::*; use arrow::array::{ArrayRef, RecordBatch}; + use arrow::compute::{cast, concat_batches}; + use arrow::datatypes::{Field, Int32Type, Schema}; use arrow_schema::SortOptions; use datafusion_common::Result; use datafusion_execution::TaskContext; @@ -310,6 +323,181 @@ mod tests { ExecutionPlanMetricsSet, SpillMetrics, }; + #[rstest::rstest] + #[case::primitive(DataType::Int32)] + #[case::strings(DataType::Utf8)] + #[case::views(DataType::Utf8View)] + #[tokio::test] + async fn intermediate_merge_flushes_before_replacing_input_batch( + #[case] data_type: DataType, + #[values(false, true)] row_cursor: bool, + #[values(false, true)] round_robin: bool, + #[values(None, Some(7))] fetch: Option, + ) -> Result<()> { + let schema = + Arc::new(Schema::new(vec![Field::new("x", data_type.clone(), false)])); + let sort = PhysicalSortExpr::new_default(col("x", &schema)?); + // Two keys exercise RowCursorStream as well as the specialized cursors. + let ordering = if row_cursor { + [sort.clone(), sort].into() + } else { + [sort].into() + }; + + for flush in [false, true] { + let streams = (0..2) + .map(|run| { + let batches = (0..3) + .map(|batch| { + let values = [4 * batch + run, 4 * batch + run + 2]; + let array: ArrayRef = match data_type { + DataType::Int32 => { + Arc::new(Int32Array::from_iter_values(values)) + } + DataType::Utf8 => { + Arc::new(StringArray::from_iter_values( + values.map(|value| format!("{value:024}")), + )) + } + DataType::Utf8View => { + Arc::new(StringViewArray::from_iter_values( + values.map(|value| format!("{value:024}")), + )) + } + _ => unreachable!(), + }; + Ok(RecordBatch::try_new(Arc::clone(&schema), vec![array])?) + }) + .collect::>>(); + Box::pin(RecordBatchStreamAdapter::new( + Arc::clone(&schema), + futures::stream::iter(batches), + )) as SendableRecordBatchStream + }) + .collect(); + let stream = StreamingMergeBuilder::new() + .with_schema(Arc::clone(&schema)) + .with_expressions(&ordering) + .with_metrics(BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0)) + .with_streams(streams) + .with_batch_size(8192) + .with_fetch(fetch) + .with_round_robin_tie_breaker(round_robin) + .with_flush_on_input_batch_boundary(flush) + .with_bypass_mempool() + .build()?; + let batches = collect(stream).await?; + let merged = concat_batches(&schema, &batches)?; + let actual = cast(merged.column(0), &DataType::Int32)?; + let expected = Int32Array::from_iter_values(0..fetch.unwrap_or(12) as i32); + assert_eq!(actual.as_primitive::(), &expected); + + if flush { + // Each output must contain rows from at most one source batch + // of either input, even though the target is 8192 rows. + for batch in &batches { + let values = cast(batch.column(0), &DataType::Int32)?; + for run in 0..2 { + let mut source_batches = values + .as_primitive::() + .values() + .iter() + .filter(|&&value| value % 2 == run) + .map(|value| value / 4); + if let Some(first) = source_batches.next() { + assert!(source_batches.all(|batch| batch == first)); + } + } + } + } else { + assert_eq!(batches.len(), 1, "ordinary merges retain their batching"); + } + } + Ok(()) + } + + #[tokio::test] + async fn intermediate_merge_releases_consumed_dictionary_batches() -> Result<()> { + let schema = Arc::new(Schema::new(vec![ + Field::new("x", DataType::Int32, false), + Field::new( + "payload", + DataType::Dictionary( + Box::new(DataType::Int32), + Box::new(DataType::Utf8View), + ), + false, + ), + ])); + let mut max_input_bytes = 0; + let mut streams = Vec::new(); + for run in 0..2 { + let mut batches = Vec::new(); + for batch_index in 0..3 { + let keys = Int32Array::from(vec![ + 4 * batch_index + run, + 4 * batch_index + run + 2, + ]); + let values = StringViewArray::from(vec![format!( + "{run}:{batch_index}:{}", + "x".repeat(1024), + )]); + let payload = DictionaryArray::::try_new( + Int32Array::from(vec![0, 0]), + Arc::new(values), + )?; + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(keys), Arc::new(payload)], + )?; + max_input_bytes = max_input_bytes.max(batch.get_sliced_size()?); + batches.push(Ok(batch)); + } + streams.push(Box::pin(RecordBatchStreamAdapter::new( + Arc::clone(&schema), + futures::stream::iter(batches), + )) as SendableRecordBatchStream); + } + let ordering = [PhysicalSortExpr::new_default(col("x", &schema)?)].into(); + let stream = StreamingMergeBuilder::new() + .with_schema(schema) + .with_expressions(&ordering) + .with_metrics(BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0)) + .with_streams(streams) + .with_batch_size(8192) + .with_flush_on_input_batch_boundary(true) + .with_bypass_mempool() + .build()?; + let batches = collect(stream).await?; + let mut expected_key = 0; + for batch in batches { + // Arrow's dictionary interleave can concatenate values from all + // buffered batches, including ones with no selected rows. Fully + // consumed inputs must be removed before accepting replacements. + assert!(batch.get_sliced_size()? <= 2 * max_input_bytes); + assert!(batch.column(1).as_dictionary::().values().len() <= 2); + let payload = cast(batch.column(1), &DataType::Utf8View)?; + for (key, value) in batch + .column(0) + .as_primitive::() + .values() + .iter() + .zip(payload.as_string_view().iter()) + { + assert_eq!(*key, expected_key); + assert_eq!( + value, + Some( + format!("{}:{}:{}", key % 2, key / 4, "x".repeat(1024)).as_str() + ) + ); + expected_key += 1; + } + } + assert_eq!(expected_key, 12); + Ok(()) + } + #[tokio::test] async fn test_sort_merge_fetch_zero_with_only_1_stream() { test_fetch_0_should_output_0_rows(1, 0).await.unwrap(); From fdc832bc49fb22389d6012ed1e03fb8ed90e172b Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Sat, 19 Sep 2026 05:27:34 +0000 Subject: [PATCH 3/4] test: bound merge fan-in for parallel aggregate spilling --- .../sqllogictest/test_files/aggregate_memory_spill.slt | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/datafusion/sqllogictest/test_files/aggregate_memory_spill.slt b/datafusion/sqllogictest/test_files/aggregate_memory_spill.slt index 53d515bf7fd13..740fc8802a3d6 100644 --- a/datafusion/sqllogictest/test_files/aggregate_memory_spill.slt +++ b/datafusion/sqllogictest/test_files/aggregate_memory_spill.slt @@ -198,6 +198,11 @@ FROM ( statement ok SET datafusion.execution.target_partitions = 4 +# Bound merge buffers so each partition can allocate replay state while the +# other partitions retain aggregate state in the shared greedy memory pool. +statement ok +SET datafusion.runtime.max_spill_merge_fan_in = 2 + query II SELECT count(*), sum(total) FROM ( @@ -227,6 +232,9 @@ FROM ( # Restore settings to slt runner defaults +statement ok +RESET datafusion.runtime.max_spill_merge_fan_in + statement ok RESET datafusion.runtime.memory_limit From ede48fcdd6240a73baf14d650d4219e5059de6e7 Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Sat, 19 Sep 2026 18:06:43 +0000 Subject: [PATCH 4/4] fix: keep intermediate spill merges within their admitted memory The amd64 SQL CI failure in aggregate_memory_spill Case G occurred when four final partitions shared a 1 MiB GreedyMemoryPool: replay requested 6.6 KiB with only 948 bytes free. Capping test fan-in removed default-path coverage without protecting competing replay allocations. Trade two-batch read-ahead for more intermediate inputs within the original admitted grant, restore the default fan-in test, and cover the peer-growth regression directly. Keep the original selection and read-ahead for a no-readmission retry and preserve the first error if that retry fails. Flush intermediate output according to retained-source and prospective output memory instead of every input boundary. Cover full batches, short and skewed inputs, dictionaries, ties, and seeded concurrent aggregations. Reuse test fixtures and share spill admission cost calculations. --- ...spilling_fuzz_in_memory_constrained_env.rs | 129 ++++++++- datafusion/physical-plan/src/sorts/builder.rs | 258 ++++++++++++++++-- datafusion/physical-plan/src/sorts/merge.rs | 35 ++- .../src/sorts/multi_level_merge.rs | 238 +++++++++------- .../replay_headroom_tests.rs | 257 +++++++++-------- .../src/sorts/streaming_merge.rs | 249 ++++++++++++++--- .../test_files/aggregate_memory_spill.slt | 8 - 7 files changed, 853 insertions(+), 321 deletions(-) diff --git a/datafusion/core/tests/fuzz_cases/spilling_fuzz_in_memory_constrained_env.rs b/datafusion/core/tests/fuzz_cases/spilling_fuzz_in_memory_constrained_env.rs index 761c492d8dde3..0138f0cfcfeba 100644 --- a/datafusion/core/tests/fuzz_cases/spilling_fuzz_in_memory_constrained_env.rs +++ b/datafusion/core/tests/fuzz_cases/spilling_fuzz_in_memory_constrained_env.rs @@ -22,7 +22,8 @@ use std::sync::Arc; use crate::fuzz_cases::aggregate_fuzz::assert_spill_count_metric; use crate::fuzz_cases::once_exec::OnceExec; -use arrow::array::UInt64Array; +use arrow::array::{AsArray, UInt64Array}; +use arrow::datatypes::UInt64Type; use arrow::row::{RowConverter, SortField}; use arrow::{array::StringArray, compute::SortOptions, record_batch::RecordBatch}; use arrow_schema::{DataType, Field, Schema}; @@ -34,7 +35,7 @@ use datafusion::physical_plan::sorts::sort::SortExec; use datafusion::prelude::SessionConfig; use datafusion_common::units::{KB, MB}; use datafusion_execution::memory_pool::{ - FairSpillPool, MemoryConsumer, MemoryReservation, + FairSpillPool, GreedyMemoryPool, MemoryConsumer, MemoryReservation, }; use datafusion_execution::{SendableRecordBatchStream, TaskContext}; use datafusion_functions_aggregate::array_agg::array_agg_udaf; @@ -48,7 +49,8 @@ use datafusion_physical_plan::aggregates::{ use datafusion_physical_plan::metrics::MetricValue; use datafusion_physical_plan::spill::get_record_batch_memory_size; use datafusion_physical_plan::stream::RecordBatchStreamAdapter; -use futures::StreamExt; +use futures::{StreamExt, TryStreamExt}; +use rand::{Rng, SeedableRng, rngs::StdRng}; use arrow::array::Int32Array; use datafusion::datasource::memory::MemorySourceConfig; @@ -826,9 +828,125 @@ async fn test_aggregate_with_high_cardinality_with_limited_memory_and_large_reco Ok(()) } +#[rstest::rstest] +#[case::greedy_seed_7(7, false)] +#[case::greedy_seed_91(91, false)] +#[case::fair_seed_7(7, true)] +#[case::fair_seed_91(91, true)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn aggregate_spill_with_seeded_batches_and_concurrent_peer( + #[case] seed: u64, + #[case] fair: bool, +) -> Result<()> { + const POOL_SIZE: usize = 2 * MB as usize; + const INPUT_BATCHES: usize = 96; + let inner: Arc = if fair { + Arc::new(FairSpillPool::new(POOL_SIZE)) + } else { + Arc::new(GreedyMemoryPool::new(POOL_SIZE)) + }; + let pool = Arc::new(PeakRecordingPool::new(inner)); + let runtime = RuntimeEnvBuilder::new() + .with_memory_pool(Arc::clone(&pool) as Arc) + .build_arc()?; + let mut rng = StdRng::seed_from_u64(seed); + let mut runs = Vec::new(); + // Construct both pipelines before polling either, so their reservations + // overlap. Each peer reads seeded narrow and oversized string batches, + // spills, and replays against the same pool with default merge fan-in. + for _ in 0..2 { + let batch_size = [64, 128][rng.random_range(0..2)]; + let sizes = Arc::new( + (0..INPUT_BATCHES) + .map(|index| { + if index % 19 == 0 { + rng.random_range(POOL_SIZE / 12..POOL_SIZE / 6) + } else { + rng.random_range(POOL_SIZE / 128..POOL_SIZE / 32) + } + }) + .collect::>(), + ); + let task_ctx = Arc::new( + TaskContext::default() + .with_session_config(SessionConfig::new().with_batch_size(batch_size)) + .with_runtime(Arc::clone(&runtime)), + ); + let generated_sizes = Arc::clone(&sizes); + let (args, plan, stream) = + build_high_cardinality_aggregate(RunTestWithLimitedMemoryArgs { + pool_size: POOL_SIZE, + task_ctx, + number_of_record_batches: INPUT_BATCHES, + get_size_of_record_batch_to_generate: Box::pin(move |index| { + generated_sizes[index] + }), + memory_behavior: MemoryBehavior::AsIs, + assert_all_output_batches_roughly_match_batch_size_conf: false, + })?; + let schema = stream.schema(); + let mut seen = vec![false; INPUT_BATCHES * batch_size]; + let checked = stream.map_ok(move |batch| { + let keys = batch.column(0).as_primitive::(); + let lists = batch.column(1).as_list::(); + for row in 0..batch.num_rows() { + let key = keys.value(row) as usize; + assert!(key < INPUT_BATCHES * batch_size, "seed={seed}, key={key}"); + assert!( + !std::mem::replace(&mut seen[key], true), + "duplicate key {key}, seed={seed}" + ); + let expected_length = sizes[key / batch_size] + .saturating_sub(size_of::() * batch_size) + / batch_size; + let values = lists.value(row); + let strings = values.as_string::(); + assert_eq!(values.len(), 1, "seed={seed}, key={key}"); + assert_eq!( + strings.value(0).len(), + expected_length, + "seed={seed}, key={key}" + ); + assert!(strings.value(0).bytes().all(|byte| byte == b'a')); + } + batch + }); + let stream = Box::pin(RecordBatchStreamAdapter::new(schema, checked)); + runs.push(run_test(args, plan, stream)); + } + tokio::time::timeout( + std::time::Duration::from_secs(30), + futures::future::try_join_all(runs), + ) + .await + .expect("concurrent seeded aggregates must complete")?; + assert!( + pool.peak_reserved() <= POOL_SIZE, + "seed={seed}, fair={fair}" + ); + assert_eq!(pool.reserved(), 0); + assert_eq!(runtime.disk_manager.used_disk_space(), 0); + assert_eq!( + runtime.disk_manager.spilling_progress().active_files_count, + 0 + ); + Ok(()) +} + async fn run_test_aggregate_with_high_cardinality( - mut args: RunTestWithLimitedMemoryArgs, + args: RunTestWithLimitedMemoryArgs, ) -> Result { + let (args, aggregate_final, result) = build_high_cardinality_aggregate(args)?; + run_test(args, aggregate_final, result).await +} + +fn build_high_cardinality_aggregate( + mut args: RunTestWithLimitedMemoryArgs, +) -> Result<( + RunTestWithLimitedMemoryArgs, + Arc, + SendableRecordBatchStream, +)> { let get_size_of_record_batch_to_generate = std::mem::replace( &mut args.get_size_of_record_batch_to_generate, Box::pin(move |_| unreachable!("should not be called after take")), @@ -909,8 +1027,7 @@ async fn run_test_aggregate_with_high_cardinality( )?); let result = aggregate_final.execute(0, Arc::clone(&args.task_ctx))?; - - run_test(args, aggregate_final, result).await + Ok((args, aggregate_final, result)) } async fn run_test( diff --git a/datafusion/physical-plan/src/sorts/builder.rs b/datafusion/physical-plan/src/sorts/builder.rs index 1c38db9b6d162..d63aaa5754826 100644 --- a/datafusion/physical-plan/src/sorts/builder.rs +++ b/datafusion/physical-plan/src/sorts/builder.rs @@ -18,10 +18,10 @@ use crate::spill::get_record_batch_memory_size; use arrow::array::ArrayRef; use arrow::compute::interleave; -use arrow::datatypes::SchemaRef; +use arrow::datatypes::{DataType, SchemaRef}; use arrow::error::ArrowError; use arrow::record_batch::RecordBatch; -use datafusion_common::{DataFusionError, Result}; +use datafusion_common::{DataFusionError, Result, assert_or_internal_err}; use datafusion_execution::memory_pool::MemoryReservation; use log::warn; use std::sync::Arc; @@ -149,29 +149,175 @@ impl BatchBuilder { /// Release fully consumed batches after a merge drains at an input boundary. /// Keeping their dictionaries can otherwise enlarge the next output even /// though none of its rows refer to those batches. - pub(super) fn discard_consumed_batches(&mut self) { - assert!(self.indices.is_empty()); - self.retain_current_batches(false); + pub(super) fn discard_consumed_batches(&mut self) -> Result<()> { + assert_or_internal_err!( + self.indices.is_empty(), + "pending merge rows must be emitted before discarding source batches" + ); + self.retain_cursor_batches(); + // Bypassed spill merges only update their local accounting here; their + // real pool reservation remains attached to the outer merge stream. self.release_unused_memory(); + Ok(()) } - fn retain_current_batches(&mut self, keep_consumed: bool) { - let mut batch_idx = 0; - let mut retained = 0; - self.batches.retain(|(stream_idx, batch)| { - let stream_cursor = &mut self.cursors[*stream_idx]; - let retain = stream_cursor.batch_idx == batch_idx - && (keep_consumed || stream_cursor.row_idx < batch.num_rows()); - batch_idx += 1; + /// Whether replacing an exhausted input would exceed the allowance for + /// retained source batches and materializing output together. This preserves + /// the caller's existing source/output estimate; cursor, read-ahead and IPC + /// allocations still depend on the merge's heuristic workspace reservation. + /// + /// This check runs only at input boundaries. A tight allowance falls back + /// to flushing at every boundary with pending rows, as before. An empty + /// builder skips flushing, so this policy cannot emit empty batches or + /// stall progress while waiting for a larger allowance. + pub(super) fn should_flush_before_input( + &self, + next_batch_bytes: usize, + memory_limit: usize, + batch_size: usize, + ) -> Result { + if self.is_empty() { + return Ok(false); + } - if retain { - stream_cursor.batch_idx = retained; - retained += 1; - } else { - self.batches_mem_used -= get_record_batch_memory_size(batch); + // Rows selected from each source batch form one contiguous range, even + // though the merged indices interleave those ranges. + let mut ranges = self + .schema + .fields() + .iter() + .any(|field| { + matches!( + field.data_type(), + DataType::Utf8 + | DataType::Binary + | DataType::LargeUtf8 + | DataType::LargeBinary + ) + }) + .then(|| vec![None; self.batches.len()]); + if let Some(ranges) = &mut ranges { + for &(batch, row) in &self.indices { + let range = ranges[batch].get_or_insert(row..row); + range.end = row + 1; } - retain + } + + let mut remaining_rows = 0usize; + for (batch_idx, (stream_idx, batch)) in self.batches.iter().enumerate() { + let cursor = &self.cursors[*stream_idx]; + // Released cursors have no matching batch and contribute no rows. + if cursor.batch_idx == batch_idx && cursor.row_idx < batch.num_rows() { + remaining_rows = + remaining_rows.saturating_add(batch.num_rows() - cursor.row_idx); + if let Some(ranges) = &mut ranges { + let range = + ranges[batch_idx].get_or_insert(cursor.row_idx..cursor.row_idx); + range.end = batch.num_rows(); + } + } + } + // A primitive column bounds the number of rows in the next input even + // when other columns are variable-width. Use the largest individual + // width, since columns may share their backing buffers. + let minimum_row_bytes = self + .schema + .fields() + .iter() + .filter_map(|field| field.data_type().primitive_width()) + .max(); + let future_rows = minimum_row_bytes.map_or(batch_size, |width| { + self.len() + .saturating_add(remaining_rows) + .saturating_add(next_batch_bytes / width) + .min(batch_size) }); + + let mut output_bytes = 0usize; + let mut next_batch_may_contribute_values = false; + for (column, field) in self.schema.fields().iter().enumerate() { + let data_type = field.data_type(); + let column_bytes = if let Some(width) = data_type.primitive_width() { + padded_buffer_size(future_rows.saturating_mul(width)) + } else { + match data_type { + DataType::Null => 0, + DataType::Boolean => padded_buffer_size(future_rows.div_ceil(8)), + DataType::Utf8 + | DataType::Binary + | DataType::LargeUtf8 + | DataType::LargeBinary => { + next_batch_may_contribute_values = true; + let offset_width = + if matches!(data_type, DataType::Utf8 | DataType::Binary) { + 4 + } else { + 8 + }; + let mut values_bytes = 0usize; + let ranges = ranges.as_ref().expect("byte arrays have ranges"); + for ((_, batch), range) in self.batches.iter().zip(ranges) { + if let Some(range) = range { + let data = batch.column(column).to_data(); + let slice = data.slice(range.start, range.len()); + // Remove the source's offsets and validity so + // the output accounts for those buffers once. + let validity = if slice.nulls().is_some() { + range.len().div_ceil(8) + } else { + 0 + }; + values_bytes = values_bytes.saturating_add( + slice.get_slice_memory_size()? + - (range.len() + 1) * offset_width + - validity, + ); + } + } + padded_buffer_size(values_bytes).saturating_add( + padded_buffer_size( + future_rows + .saturating_add(1) + .saturating_mul(offset_width), + ), + ) + } + _ => { + // Nested arrays, dictionaries and views can retain or + // concatenate buffers with no selected rows. Count the + // entire input buffers, including the next input, rather + // than estimating them from an average row width. + next_batch_may_contribute_values = true; + self.batches.iter().fold(0usize, |bytes, (_, batch)| { + bytes.saturating_add( + batch.column(column).get_buffer_memory_size(), + ) + }) + } + } + }; + output_bytes = output_bytes.saturating_add(column_bytes); + if field.is_nullable() + || self + .batches + .iter() + .any(|(_, batch)| batch.column(column).nulls().is_some()) + { + output_bytes = output_bytes + .saturating_add(padded_buffer_size(future_rows.div_ceil(8))); + } + } + if next_batch_may_contribute_values { + // The replacement's variable-width values are not available yet. + // Include their full maximum rather than their average row width. + output_bytes = output_bytes.saturating_add(next_batch_bytes); + } + + Ok(self + .batches_mem_used + .saturating_add(next_batch_bytes) + .saturating_add(output_bytes) + > memory_limit) } fn release_unused_memory(&mut self) { @@ -314,6 +460,11 @@ impl BatchBuilder { } } +/// Arrow rounds newly allocated buffers up to a multiple of 64 bytes. +fn padded_buffer_size(bytes: usize) -> usize { + bytes.checked_add(63).map_or(usize::MAX, |size| size & !63) +} + /// Try to grow `reservation` so it covers at least `needed` bytes. /// /// When a reservation has been pre-loaded with bytes (e.g. via @@ -373,7 +524,9 @@ where #[cfg(test)] mod tests { use super::*; - use arrow::array::{Array, ArrayDataBuilder, Int32Array, ListArray}; + use arrow::array::{ + Array, ArrayDataBuilder, Int32Array, Int64Array, ListArray, StringArray, + }; use arrow::buffer::Buffer; use arrow::datatypes::{DataType, Field, Schema}; use arrow::record_batch::RecordBatch; @@ -490,6 +643,71 @@ mod tests { assert_eq!(builder.reservation.size(), 0); } + #[test] + fn test_merge_budget_includes_future_replacement_rows() -> Result<()> { + let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)])); + let pool: Arc = Arc::new(UnboundedMemoryPool::default()); + let mut builder = BatchBuilder::new( + Arc::clone(&schema), + 2, + 8192, + MemoryConsumer::new("test").register(&pool), + ); + for stream in 0..2 { + builder.push_batch( + stream, + RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(Int64Array::from_iter_values(0..8))], + )?, + )?; + } + for _ in 0..8 { + builder.push_row(0); + } + // Current output (64 bytes), both inputs (128), and replacement (64) + // fit. Future output also includes the live input and replacement. + assert!(builder.should_flush_before_input(64, 256, 8192)?); + Ok(()) + } + + #[test] + fn test_merge_budget_includes_unselected_wide_values() -> Result<()> { + let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)])); + let pool: Arc = Arc::new(UnboundedMemoryPool::default()); + let mut builder = BatchBuilder::new( + Arc::clone(&schema), + 2, + 3, + MemoryConsumer::new("test").register(&pool), + ); + builder.push_batch( + 0, + RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(StringArray::from(vec!["a"]))], + )?, + )?; + let wide_value = "z".repeat(4096); + builder.push_batch( + 1, + RecordBatch::try_new( + schema, + vec![Arc::new(StringArray::from(vec!["b", wide_value.as_str()]))], + )?, + )?; + builder.push_row(0); + // Only a small value is pending, but the other live batch can add its + // large value before the next input boundary. Average widths or only + // the current pending slice would miss this materialization cost. + assert!(builder.should_flush_before_input( + 64, + builder.batches_mem_used + 64 + 256, + 3, + )?); + Ok(()) + } + #[test] fn test_retry_interleave_halves_rows_until_success() { let mut attempts = Vec::new(); diff --git a/datafusion/physical-plan/src/sorts/merge.rs b/datafusion/physical-plan/src/sorts/merge.rs index 2cbd93751d9aa..1fd122dfc708e 100644 --- a/datafusion/physical-plan/src/sorts/merge.rs +++ b/datafusion/physical-plan/src/sorts/merge.rs @@ -28,6 +28,7 @@ use crate::metrics::BaselineMetrics; use crate::sorts::builder::BatchBuilder; use crate::sorts::cursor::{Cursor, CursorValues}; use crate::sorts::stream::PartitionedStream; +use crate::sorts::streaming_merge::MergeBatchMemoryBudget; use crate::stream::{ObservedStream, RecordBatchStreamAdapter}; use arrow::datatypes::SchemaRef; @@ -86,10 +87,8 @@ pub(crate) struct SortPreservingMergeStream { /// Target batch size batch_size: usize, - /// Emit pending rows before replacing an exhausted input batch. Intermediate - /// spill merges use this to avoid collecting multiple short batches from an - /// input into an output batch larger than their reserved workspace. - flush_on_input_batch_boundary: bool, + /// Budget for retaining source batches while materializing merged output. + batch_memory_budget: Option, /// Cursors for each input partition. `None` means the input is exhausted cursors: Vec>>, @@ -151,14 +150,17 @@ impl SortPreservingMergeStream { poll_reset_epochs: vec![0; stream_count], loser_tree: vec![], batch_size, - flush_on_input_batch_boundary: false, + batch_memory_budget: None, fetch, produced: 0, } } - pub(super) fn with_flush_on_input_batch_boundary(mut self, flush: bool) -> Self { - self.flush_on_input_batch_boundary = flush; + pub(super) fn with_batch_memory_budget( + mut self, + budget: Option, + ) -> Self { + self.batch_memory_budget = budget; self } @@ -295,12 +297,19 @@ impl SortPreservingMergeStream { ); drop(timer); - if self.flush_on_input_batch_boundary { - // Drain every pending row, including partial output - // from offset-overflow recovery, before accepting a - // replacement batch from this input. - self.flush_in_progress(&mut emitter).await?; - self.in_progress.discard_consumed_batches(); + if let Some(budget) = &self.batch_memory_budget { + if self.in_progress.should_flush_before_input( + budget.input_batch_sizes[winner_stream], + budget.memory_limit, + self.batch_size, + )? { + // Drain partial output from offset-overflow + // recovery before remapping source batch indices. + self.flush_in_progress(&mut emitter).await?; + } + if self.in_progress.is_empty() { + self.in_progress.discard_consumed_batches()?; + } } poll_fn(|cx| self.maybe_poll_stream(cx, winner_stream)).await?; timer = elapsed_compute.timer(); diff --git a/datafusion/physical-plan/src/sorts/multi_level_merge.rs b/datafusion/physical-plan/src/sorts/multi_level_merge.rs index 80d029cd85837..53d96707e9643 100644 --- a/datafusion/physical-plan/src/sorts/multi_level_merge.rs +++ b/datafusion/physical-plan/src/sorts/multi_level_merge.rs @@ -32,7 +32,9 @@ use datafusion_execution::memory_pool::{MemoryReservation, MergeMemoryPool}; use crate::sorts::builder::try_grow_reservation_to_at_least; use crate::sorts::sort::get_reserved_bytes_for_record_batch_size; -use crate::sorts::streaming_merge::{SortedSpillFile, StreamingMergeBuilder}; +use crate::sorts::streaming_merge::{ + MergeBatchMemoryBudget, SortedSpillFile, StreamingMergeBuilder, +}; use crate::spill::gc_view_arrays; use crate::spill::spill_manager::GetSlicedSize; use crate::stream::{ObservedStream, RecordBatchStreamAdapter}; @@ -212,8 +214,8 @@ impl MultiLevelMergeBuilder { self } - /// Leave replay headroom for the final merge. Intermediate merges and - /// splitting can use the full pool because replay has not started. + /// Leave replay headroom when admitting a merge. Intermediate passes may + /// trade read-ahead for more inputs within the same admitted reservation. pub(super) fn with_replay_headroom(mut self, reserve: bool) -> Self { self.reserve_replay_headroom = reserve; self @@ -290,24 +292,28 @@ impl MultiLevelMergeBuilder { .await; drop(stream); // A successful write drops the backups before the next selection. - if let Some(retry) = retry - && result.is_err() - { + if let (Some(retry), Err(original_error)) = (retry, &result) { + let original_error = original_error.to_string(); self.widen_intermediate_merges = false; let IntermediateMergeRetry { mut spills, original_count, buffer_size, - original_memory, reservation, } = retry; self.sorted_spill_files .splice(0..0, spills.drain(original_count..)); - reservation.shrink(reservation.size() - original_memory); // No pool re-admission is needed: the successful original // grant has remained live, even if the writer failed at EOF. - let (mut stream, limit) = - self.merge_selected_runs(spills, buffer_size, reservation, false)?; + let (mut stream, limit) = self + .merge_selected_runs(&spills, buffer_size, reservation, false) + .map_err(|error| { + error.context(format!( + "Rebuilding the narrower intermediate merge after: {original_error}" + )) + })?; + // Only the stream must retain these inputs during the retry. + drop(spills); batch_size_limit = limit; result = self .spill_manager @@ -315,7 +321,12 @@ impl MultiLevelMergeBuilder { &mut stream, "MultiLevelMergeBuilder intermediate spill retry", ) - .await; + .await + .map_err(|error| { + error.context(format!( + "Retrying the narrower intermediate merge after: {original_error}" + )) + }); } let Some((spill_file, max_record_batch_memory)) = result? else { continue; @@ -392,7 +403,7 @@ impl MultiLevelMergeBuilder { true, true, self.batch_size, - false, + None, )?, batch_size_limit: self.batch_size, retry: None, @@ -419,7 +430,7 @@ impl MultiLevelMergeBuilder { &mut memory_reservation, allow_minimum_without_headroom, )?; - let (mut spills, buffer_size) = match selection { + let (mut spills, mut buffer_size) = match selection { SpillFilesToMerge::Ready(spills, buffer_size) => { (spills, buffer_size) } @@ -429,68 +440,50 @@ impl MultiLevelMergeBuilder { }; let original_count = spills.len(); + let original_buffer_size = buffer_size; let original_memory = memory_reservation.size(); if self.reserve_replay_headroom && self.widen_intermediate_merges + && !allow_minimum_without_headroom + && buffer_size > 1 && self.sorted_streams.is_empty() && !spills.is_empty() { - // Extend the admitted selection without releasing its grant - // or changing read-ahead. Keep one run for the final merge, - // which must still leave space for aggregate replay. - let max_fan_in = effective_spill_merge_fan_in( - self.spill_manager - .env() - .disk_manager - .max_spill_merge_fan_in(), - ); - let mut extra = 0; - for (spill, _) in self - .sorted_spill_files + // Trade read-ahead for fan-in without taking any more pool + // memory. Other partitions retain exactly the space left by + // the original admission, even while they replay aggregates. + // Keep one run for the final merge and its replay headroom. + let candidates = spills .iter() - .take(self.sorted_spill_files.len().saturating_sub(1)) - { - if spills.len() + extra >= max_fan_in { - break; - } - let additional = get_reserved_bytes_for_record_batch_size( - spill.max_record_batch_memory, - spill.max_record_batch_memory, - ) * buffer_size; - if memory_reservation.try_grow(additional).is_err() { - break; - } - extra += 1; + .chain(&self.sorted_spill_files) + .take(spills.len() + self.sorted_spill_files.len() - 1) + .map(|(spill, _)| spill); + let widened_count = spill_merge_memory_requirements( + candidates, + 1, + self.max_spill_merge_fan_in(), + ) + .take_while(|needed| *needed <= original_memory) + .count(); + if widened_count > original_count { + buffer_size = 1; + spills.extend( + self.sorted_spill_files + .drain(..widened_count - original_count), + ); } - spills.extend(self.sorted_spill_files.drain(..extra)); } let reservation = Arc::new(memory_reservation); - let retry = - (spills.len() > original_count).then(|| IntermediateMergeRetry { - spills: spills - .iter() - .map(|(spill, limit)| { - ( - SortedSpillFile { - file: Arc::clone(&spill.file), - max_record_batch_memory: spill - .max_record_batch_memory, - }, - *limit, - ) - }) - .collect(), - original_count, - buffer_size, - original_memory, - reservation: Arc::clone(&reservation), - }); - let (stream, batch_size_limit) = self.merge_selected_runs( + let widened = spills.len() > original_count; + let retry_reservation = widened.then(|| Arc::clone(&reservation)); + let (stream, batch_size_limit) = + self.merge_selected_runs(&spills, buffer_size, reservation, widened)?; + let retry = retry_reservation.map(|reservation| IntermediateMergeRetry { spills, - buffer_size, + original_count, + buffer_size: original_buffer_size, reservation, - retry.is_some(), - )?; + }); Ok(MergeStep::Stream { stream, batch_size_limit, @@ -502,16 +495,19 @@ impl MultiLevelMergeBuilder { /// Build a stream from an already admitted selection. The reservation can /// also be held by a retry guard until an intermediate writer finishes. + /// `bound_batch_memory` requires spill-only inputs and a one-batch read buffer. fn merge_selected_runs( &mut self, - sorted_spill_files: Vec<(SortedSpillFile, usize)>, + sorted_spill_files: &[(SortedSpillFile, usize)], buffer_size: usize, memory_reservation: Arc, - flush_on_input_batch_boundary: bool, + bound_batch_memory: bool, ) -> Result<(SendableRecordBatchStream, usize)> { // Don't account for existing streams memory // as we are not holding the memory for them let mut sorted_streams = mem::take(&mut self.sorted_streams); + debug_assert!(!bound_batch_memory || sorted_streams.is_empty()); + debug_assert!(!bound_batch_memory || buffer_size == 1); let is_only_merging_memory_streams = sorted_spill_files.is_empty(); @@ -529,7 +525,7 @@ impl MultiLevelMergeBuilder { self.sorted_spill_files.is_empty(), true, self.batch_size, - false, + None, )?, self.batch_size, )); @@ -548,17 +544,34 @@ impl MultiLevelMergeBuilder { .spill_manager .clone() .with_batch_read_buffer_capacity(buffer_size) - .read_spill_as_stream(spill.file, Some(spill.max_record_batch_memory))?; - output_batch_size = output_batch_size.min(batch_size_limit); + .read_spill_as_stream( + Arc::clone(&spill.file), + Some(spill.max_record_batch_memory), + )?; + output_batch_size = output_batch_size.min(*batch_size_limit); sorted_streams.push(stream); } + let batch_memory_budget = bound_batch_memory.then(|| { + let input_batch_sizes = sorted_spill_files + .iter() + .map(|(spill, _)| spill.max_record_batch_memory) + .collect::>(); + MergeBatchMemoryBudget { + // With buffer_size == 1, spill admission estimates two batches + // per selected run: one retained source and one materialized + // output. This can be smaller than the original grant retained + // by the widened merge. + memory_limit: input_batch_sizes.iter().sum::() * 2, + input_batch_sizes, + } + }); let merge_sort_stream = self.create_new_merge_sort( sorted_streams, // If we have no sorted spill files left, this is the last run self.sorted_spill_files.is_empty(), is_only_merging_memory_streams, output_batch_size, - flush_on_input_batch_boundary, + batch_memory_budget, )?; Ok(( @@ -576,13 +589,13 @@ impl MultiLevelMergeBuilder { is_output: bool, all_in_memory: bool, output_batch_size: usize, - flush_on_input_batch_boundary: bool, + batch_memory_budget: Option, ) -> Result { let mut builder = StreamingMergeBuilder::new() .with_schema(Arc::clone(&self.schema)) .with_expressions(&self.expr) .with_batch_size(output_batch_size) - .with_flush_on_input_batch_boundary(flush_on_input_batch_boundary) + .with_batch_memory_budget(batch_memory_budget) .with_fetch(self.fetch) .with_metrics(if is_output { // Only add the metrics to the last run @@ -623,35 +636,25 @@ impl MultiLevelMergeBuilder { ) -> Result { assert_ne!(buffer_len, 0, "Buffer length must be greater than 0"); let mut number_of_spills_to_read_for_current_phase = 0; - let configured_fan_in = self - .spill_manager - .env() - .disk_manager - .max_spill_merge_fan_in(); - let max_spill_files = effective_spill_merge_fan_in(configured_fan_in); // Track total memory needed for spill file buffers. When the // reservation has pre-reserved bytes (from sort_spill_reservation_bytes), // those bytes cover the first N spill files without additional pool // allocation, preventing starvation under memory pressure. - let mut total_needed: usize = 0; let mut accepted_memory: usize = 0; - - for (spill, _) in &self.sorted_spill_files { - if number_of_spills_to_read_for_current_phase >= max_spill_files - || (allow_minimum_without_headroom - && number_of_spills_to_read_for_current_phase - >= minimum_number_of_required_streams) + let mut reduce_buffer_len = false; + + for total_needed in spill_merge_memory_requirements( + self.sorted_spill_files.iter().map(|(spill, _)| spill), + buffer_len, + self.max_spill_merge_fan_in(), + ) { + if allow_minimum_without_headroom + && number_of_spills_to_read_for_current_phase + >= minimum_number_of_required_streams { break; } - let per_spill = get_reserved_bytes_for_record_batch_size( - spill.max_record_batch_memory, - // Size will be the same as the sliced size, bc it is a spilled batch. - spill.max_record_batch_memory, - ) * buffer_len; - total_needed += per_spill; - // If a run cannot shrink, allow only the minimum merge without // replay headroom. Disable read-ahead and still ask the pool for // every byte used by the merge. @@ -688,12 +691,8 @@ impl MultiLevelMergeBuilder { reservation.free(); if buffer_len > 1 { // Try again with smaller buffer size, it will be slower but at least we can merge - return self.get_sorted_spill_files_to_merge( - buffer_len - 1, - minimum_number_of_required_streams, - reservation, - allow_minimum_without_headroom, - ); + reduce_buffer_len = true; + break; } // buffer_len == 1 and we still can't seat the minimum of 2 streams. @@ -728,6 +727,15 @@ impl MultiLevelMergeBuilder { } } + if reduce_buffer_len { + return self.get_sorted_spill_files_to_merge( + buffer_len - 1, + minimum_number_of_required_streams, + reservation, + allow_minimum_without_headroom, + ); + } + if self.reserve_replay_headroom { // `total_needed` may include a rejected candidate. Keep only the // buffers that were admitted, releasing temporary replay headroom. @@ -742,6 +750,15 @@ impl MultiLevelMergeBuilder { Ok(SpillFilesToMerge::Ready(spills, buffer_len)) } + fn max_spill_merge_fan_in(&self) -> usize { + effective_spill_merge_fan_in( + self.spill_manager + .env() + .disk_manager + .max_spill_merge_fan_in(), + ) + } + /// Re-spill the spill file at `index` with half its batch size, putting it back /// at the same position. We read the file back and re-spill it through the normal /// spill API (which owns batch layout), slicing every batch in two, which halves @@ -935,7 +952,6 @@ struct IntermediateMergeRetry { spills: Vec<(SortedSpillFile, usize)>, original_count: usize, buffer_size: usize, - original_memory: usize, reservation: Arc, } @@ -974,6 +990,21 @@ fn effective_spill_merge_fan_in(configured_fan_in: usize) -> usize { } } +/// Cumulative buffer costs, shared by admission and fixed-budget widening. +fn spill_merge_memory_requirements<'a>( + spills: impl Iterator, + buffer_len: usize, + max_spill_files: usize, +) -> impl Iterator { + spills.take(max_spill_files).scan(0, move |total, spill| { + *total += get_reserved_bytes_for_record_batch_size( + spill.max_record_batch_memory, + spill.max_record_batch_memory, + ) * buffer_len; + Some(*total) + }) +} + struct StreamAttachedReservation { stream: SendableRecordBatchStream, reservation: Option>, @@ -1048,11 +1079,14 @@ mod tests { ExecutionPlanMetricsSet, SpillMetrics, }; - fn test_schema() -> SchemaRef { + pub(super) fn test_schema() -> SchemaRef { Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)])) } - fn build_spill_manager(env: &Arc, schema: &SchemaRef) -> SpillManager { + pub(super) fn build_spill_manager( + env: &Arc, + schema: &SchemaRef, + ) -> SpillManager { SpillManager::new( Arc::clone(env), SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0), @@ -1062,7 +1096,7 @@ mod tests { /// Spill `values` (which must already be sorted) as a single sorted run and /// return it as a `SortedSpillFile` carrying its recorded largest-batch memory. - fn make_sorted_spill_file( + pub(super) fn make_sorted_spill_file( spill_manager: &SpillManager, schema: &SchemaRef, values: Vec, @@ -1086,7 +1120,7 @@ mod tests { } } - fn build_merge_builder( + pub(super) fn build_merge_builder( spill_manager: SpillManager, schema: SchemaRef, sorted_spill_files: Vec, diff --git a/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs b/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs index 0603389c52290..e86206c4b797a 100644 --- a/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs +++ b/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs @@ -17,13 +17,14 @@ use super::*; -use crate::expressions::PhysicalSortExpr; +use super::tests::{ + build_merge_builder, build_spill_manager, make_sorted_spill_file, test_schema, +}; use arrow::array::{AsArray, Int64Array, StringArray}; use arrow::compute::concat_batches; use arrow::datatypes::{DataType, Field, Int64Type, Schema}; use datafusion_execution::memory_pool::{GreedyMemoryPool, MemoryConsumer, MemoryPool}; use datafusion_execution::runtime_env::{RuntimeEnv, RuntimeEnvBuilder}; -use datafusion_physical_expr::expressions::Column; use datafusion_physical_expr_common::metrics::{ExecutionPlanMetricsSet, SpillMetrics}; use std::sync::atomic::{AtomicBool, Ordering}; @@ -43,19 +44,8 @@ fn replay_merge_builder( pool: &Arc, batch_size: usize, ) -> MultiLevelMergeBuilder { - MultiLevelMergeBuilder::new( - spill_manager, - schema, - spills, - vec![], - [PhysicalSortExpr::new_default(Arc::new(Column::new("x", 0)))].into(), - BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0), - batch_size, - MemoryConsumer::new("replay headroom test").register(pool), - None, - false, - ) - .with_replay_headroom(true) + build_merge_builder(spill_manager, schema, spills, pool, batch_size) + .with_replay_headroom(true) } /// Create equally sized, interleaved input runs. @@ -68,27 +58,17 @@ fn replay_merge_fixture( let env = RuntimeEnvBuilder::new() .with_max_spill_merge_fan_in(max_fan_in) .build_arc()?; - let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)])); - let metrics = SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0); - let spill_manager = - SpillManager::new(Arc::clone(&env), metrics.clone(), Arc::clone(&schema)); - let mut spills = Vec::with_capacity(run_count); - for run in 0..run_count { - let values = Int64Array::from_iter_values( - (0..rows_per_run).map(|row| (row * run_count + run) as i64), - ); - let batch = RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(values)])?; - let (file, max_record_batch_memory) = spill_manager - .spill_record_batch_iter_and_return_max_batch_memory( - std::iter::once(Ok(batch)), - "replay headroom test input", - )? - .expect("a nonempty input must spill"); - spills.push(SortedSpillFile { - file, - max_record_batch_memory, - }); - } + let schema = test_schema(); + let spill_manager = build_spill_manager(&env, &schema); + let metrics = spill_manager.metrics.clone(); + let spills = (0..run_count) + .map(|run| { + let values = (0..rows_per_run) + .map(|row| (row * run_count + run) as i64) + .collect(); + make_sorted_spill_file(&spill_manager, &schema, values) + }) + .collect::>(); let input_bytes = metrics.spilled_bytes.value(); let pool_size = memory_batches * spills[0].max_record_batch_memory; let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_size)); @@ -146,7 +126,7 @@ async fn merge_replay_runs( } #[rstest::rstest] -#[case::intermediate_uses_full_pool(6, 16, 0, 2, 16)] +#[case::intermediate_reuses_admitted_buffers(6, 16, 0, 2, 8)] #[case::intermediate_holds_back_a_run(3, 16, 0, 1, 8)] #[case::final_disables_read_ahead(3, 12, 0, 0, 6)] #[case::fan_in_limited_intermediate(4, 8, 2, 2, 4)] @@ -185,6 +165,57 @@ async fn replay_headroom_depends_on_merge_phase( Ok(()) } +#[rstest::rstest] +#[case::original_read_ahead(false, 4)] +#[case::widened_intermediate(true, 8)] +#[tokio::test] +async fn intermediate_merge_preserves_competing_replay_budget( + #[case] widen: bool, + #[case] selected_runs: usize, +) -> Result<()> { + let ReplayMergeFixture { + mut builder, + env, + pool, + .. + } = replay_merge_fixture(10, 128, 40, 0)?; + let batch_memory = builder.sorted_spill_files[0].0.max_record_batch_memory; + let peer = MemoryConsumer::new("another partition replay").register(&pool); + peer.try_grow(8 * batch_memory)?; + builder.widen_intermediate_merges = widen; + + let MergeStep::Stream { stream, retry, .. } = + builder.merge_sorted_runs_within_mem_limit(false)? + else { + panic!("the merge should fit without splitting a run"); + }; + assert_eq!(builder.sorted_spill_files.len(), 10 - selected_runs); + assert_eq!(retry.is_some(), widen); + if let Some(retry) = &retry { + assert_eq!(retry.spills.len(), selected_runs); + assert_eq!(retry.original_count, 4); + assert_eq!(retry.buffer_size, 2); + assert_eq!(retry.reservation.size(), 16 * batch_memory); + } + // Both selections retain the same sixteen batches of memory. Widening + // used to grow this grant to thirty-two and consume all of the peer's + // remaining replay budget before its next allocation. + assert_eq!(pool.reserved() - peer.size(), 16 * batch_memory); + peer.try_grow(4 * batch_memory)?; + assert_eq!(pool.reserved(), 28 * batch_memory); + let batches: Vec = stream.try_collect().await?; + assert_eq!( + batches.iter().map(RecordBatch::num_rows).sum::(), + selected_runs * 128 + ); + drop(retry); + drop(builder); + peer.free(); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.used_disk_space(), 0); + Ok(()) +} + #[tokio::test] async fn replay_headroom_keeps_split_retries_before_intermediate_merges() -> Result<()> { let ReplayMergeFixture { @@ -199,50 +230,55 @@ async fn replay_headroom_keeps_split_retries_before_intermediate_merges() -> Res Ok(()) } +#[rstest::rstest] +#[case::indivisible_split_retry(3, 32, 6, true)] +#[case::widened_short_batches(6, 3, 16, false)] #[tokio::test] -async fn replay_headroom_preserves_indivisible_run_batch_limits() -> Result<()> { +async fn replay_headroom_preserves_short_run_batch_limits( + #[case] run_count: usize, + #[case] batches_per_run: usize, + #[case] memory_batches: usize, + #[case] expect_split: bool, +) -> Result<()> { let env = Arc::new(RuntimeEnv::default()); let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)])); - let spill_manager = SpillManager::new( - Arc::clone(&env), - SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0), - Arc::clone(&schema), - ); - let values = (0..96) - .map(|value| format!("{value:03}{}", "x".repeat(1024))) + let manager = build_spill_manager(&env, &schema); + let values = (0..run_count * batches_per_run) + .map(|value| format!("{value:04}{}", "x".repeat(1024))) .collect::>(); - let mut spills = Vec::new(); - for run in 0..3 { - // Every batch is already indivisible, but the configured merge batch - // size is much larger. The split/retry path must discover the one-row - // limit before an intermediate pass can concatenate these batches. - let batches = values.iter().skip(run).step_by(3).map(|value| { - RecordBatch::try_new( - Arc::clone(&schema), - vec![Arc::new(StringArray::from(vec![value.as_str()]))], - ) - .map_err(Into::into) - }); - let (file, max_record_batch_memory) = spill_manager - .spill_record_batch_iter_and_return_max_batch_memory( - batches, - "indivisible replay input", - )? - .expect("a nonempty input must spill"); - spills.push(SortedSpillFile { - file, - max_record_batch_memory, - }); - } - let pool_size = 6 * spills[0].max_record_batch_memory; + let spills = (0..run_count) + .map(|run| { + let batches = values.iter().skip(run).step_by(run_count).map(|value| { + RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(StringArray::from(vec![value.as_str()]))], + ) + .map_err(Into::into) + }); + let (file, max_record_batch_memory) = manager + .spill_record_batch_iter_and_return_max_batch_memory( + batches, + "short replay input", + )? + .expect("a nonempty input must spill"); + Ok(SortedSpillFile { + file, + max_record_batch_memory, + }) + }) + .collect::>>()?; + let pool_size = memory_batches * spills[0].max_record_batch_memory; let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_size)); - let builder = - replay_merge_builder(spill_manager, Arc::clone(&schema), spills, &pool, 8192); + let builder = replay_merge_builder(manager, Arc::clone(&schema), spills, &pool, 8192); let mut stream = builder.create_spillable_merge_stream(); let mut batches = Vec::new(); while let Some(batch) = stream.try_next().await? { - assert_eq!(batch.num_rows(), 1); - assert!(crate::spill::get_record_batch_memory_size(&batch) <= pool.reserved()); + if expect_split { + // Unlike the parent module's direct minimum-admission test, this + // carries many wide singleton batches through intermediate passes. + // Their discovered one-row limit must survive the nominal 8192 size. + assert_eq!(batch.num_rows(), 1); + } assert!(pool.reserved() <= pool_size); batches.push(batch); } @@ -260,16 +296,16 @@ async fn replay_headroom_preserves_indivisible_run_batch_limits() -> Result<()> #[tokio::test] async fn replay_headroom_does_not_rewrite_intermediate_runs_twice() -> Result<()> { - // The pool holds eight read-ahead inputs, or four plus equal replay headroom. + // The admitted grant holds four inputs with read-ahead, or eight without it. // Four intermediate merges of eight runs leave four runs for the final merge. - // Reserving headroom for every pass instead writes ten intermediate files, + // Keeping read-ahead for every pass instead writes ten intermediate files, // rewriting every input row twice before returning the final merge. let (spill_count, spilled_rows, spilled_bytes) = merge_replay_runs(32, 256, 32).await?; - println!( - "Intermediate spill: {spill_count} files, {spilled_rows} rows, {spilled_bytes} bytes" + assert_eq!( + spill_count, 4, + "intermediate spill: {spilled_rows} rows, {spilled_bytes} bytes" ); - assert_eq!(spill_count, 4, "additional spill bytes: {spilled_bytes}"); assert_eq!(spilled_rows, 32 * 256); Ok(()) } @@ -285,59 +321,6 @@ async fn replay_headroom_is_restored_after_intermediate_split_retries() -> Resul Ok(()) } -#[tokio::test] -async fn intermediate_merge_preserves_short_batch_replay() -> Result<()> { - let env = Arc::new(RuntimeEnv::default()); - let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)])); - let manager = SpillManager::new( - Arc::clone(&env), - SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0), - Arc::clone(&schema), - ); - let values = (0..18) - .map(|value| format!("{value:04}{}", "x".repeat(1024))) - .collect::>(); - let mut spills = Vec::new(); - for run in 0..6 { - let batches = values.iter().skip(run).step_by(6).map(|value| { - RecordBatch::try_new( - Arc::clone(&schema), - vec![Arc::new(StringArray::from(vec![value.as_str()]))], - ) - .map_err(Into::into) - }); - let (file, max_record_batch_memory) = manager - .spill_record_batch_iter_and_return_max_batch_memory( - batches, - "short replay input", - )? - .expect("a nonempty input must spill"); - spills.push(SortedSpillFile { - file, - max_record_batch_memory, - }); - } - // A two-input merge fits with headroom, so no initial split discovers the - // actual one-row input size. Widening to four runs must not combine their - // twelve rows into an intermediate batch too large to split or replay. - let pool: Arc = Arc::new(GreedyMemoryPool::new( - 16 * spills[0].max_record_batch_memory, - )); - let builder = replay_merge_builder(manager, Arc::clone(&schema), spills, &pool, 8192); - let batches: Vec = builder - .create_spillable_merge_stream() - .try_collect() - .await?; - let merged = concat_batches(&schema, &batches)?; - assert_eq!( - merged.column(0).as_string::(), - &StringArray::from(values) - ); - assert_eq!(pool.reserved(), 0); - assert_eq!(env.disk_manager.used_disk_space(), 0); - Ok(()) -} - #[rstest::rstest] #[case::sufficient_quota(true)] #[case::insufficient_quota(false)] @@ -413,9 +396,15 @@ async fn check_intermediate_merge_disk_quota(sufficient_quota: bool) -> Result<( assert_eq!(merged.column(0).as_primitive::(), &expected); } else { let error = result.expect_err("the quota must reject even a two-input merge"); + let message = error.to_string(); assert!( - error.to_string().contains("max_temp_directory_size"), - "expected a disk quota error, got {error}" + message.contains("Retrying the narrower intermediate merge after:"), + "the error must identify the failed wider attempt: {message}" + ); + assert_eq!( + message.matches("max_temp_directory_size").count(), + 2, + "both the original and narrower write errors must survive: {message}" ); } // On a multithreaded runtime, dropping a failed merge cancels background @@ -450,7 +439,7 @@ fn intermediate_merge_keeps_admitted_buffers( // underlying pool's normal fallible admission. // With read-ahead disabled, the contender first takes the released replay // headroom. Retrying admission with read-ahead enabled would then fail to - // seat two inputs and relinquish the usable three-input merge's grant. + // seat two inputs and relinquish the usable minimum merge's grant. #[derive(Debug)] struct HandoffPool { inner: Arc, diff --git a/datafusion/physical-plan/src/sorts/streaming_merge.rs b/datafusion/physical-plan/src/sorts/streaming_merge.rs index 9f170493b2aa1..e519f3d1726f6 100644 --- a/datafusion/physical-plan/src/sorts/streaming_merge.rs +++ b/datafusion/physical-plan/src/sorts/streaming_merge.rs @@ -36,6 +36,15 @@ use datafusion_execution::memory_pool::{ use datafusion_physical_expr_common::sort_expr::LexOrdering; use std::sync::Arc; +/// Allowance for simultaneously retained merge inputs and materialized output. +/// This preserves the source/output portion of the caller's existing heuristic +/// merge estimate, rather than enforcing a hard bound on all live allocations. +#[derive(Debug)] +pub(super) struct MergeBatchMemoryBudget { + pub memory_limit: usize, + pub input_batch_sizes: Vec, +} + macro_rules! primitive_merge_helper { ($t:ty, $($v:ident),+) => { merge_helper!(PrimitiveArray<$t>, $($v),+) @@ -43,7 +52,7 @@ macro_rules! primitive_merge_helper { } macro_rules! merge_helper { - ($t:ty, $sort:ident, $streams:ident, $schema:ident, $tracking_metrics:ident, $batch_size:ident, $fetch:ident, $reservation:ident, $enable_round_robin_tie_breaker:ident, $flush_on_input_batch_boundary:ident) => {{ + ($t:ty, $sort:ident, $streams:ident, $schema:ident, $tracking_metrics:ident, $batch_size:ident, $fetch:ident, $reservation:ident, $enable_round_robin_tie_breaker:ident, $batch_memory_budget:ident) => {{ let streams = FieldCursorStream::<$t>::new($sort, $streams, $reservation.new_empty()); return Ok(SortPreservingMergeStream::new( @@ -55,7 +64,7 @@ macro_rules! merge_helper { $reservation, $enable_round_robin_tie_breaker, ) - .with_flush_on_input_batch_boundary($flush_on_input_batch_boundary) + .with_batch_memory_budget($batch_memory_budget) .into_stream()); }}; } @@ -99,7 +108,7 @@ pub struct StreamingMergeBuilder<'a> { merge_pool: Option>, /// Leave memory for the aggregate consuming the merged spill rows. reserve_replay_headroom: bool, - flush_on_input_batch_boundary: bool, + batch_memory_budget: Option, enable_round_robin_tie_breaker: bool, } @@ -172,9 +181,12 @@ impl<'a> StreamingMergeBuilder<'a> { self } - /// Bound intermediate output to rows from one batch of each input. - pub(super) fn with_flush_on_input_batch_boundary(mut self, flush: bool) -> Self { - self.flush_on_input_batch_boundary = flush; + /// Bound intermediate source retention and output materialization together. + pub(super) fn with_batch_memory_budget( + mut self, + budget: Option, + ) -> Self { + self.batch_memory_budget = budget; self } @@ -212,7 +224,7 @@ impl<'a> StreamingMergeBuilder<'a> { reservation, merge_pool, reserve_replay_headroom, - flush_on_input_batch_boundary, + batch_memory_budget, fetch, expressions, enable_round_robin_tie_breaker, @@ -268,18 +280,24 @@ impl<'a> StreamingMergeBuilder<'a> { let metrics = metrics.expect("Metrics cannot be empty for streaming merge"); let reservation = reservation.expect("Reservation cannot be empty for streaming merge"); + if let Some(budget) = &batch_memory_budget { + assert_or_internal_err!( + budget.input_batch_sizes.len() == streams.len(), + "merge batch budget must provide one maximum size per input" + ); + } // Special case single column comparisons with optimized cursor implementations if expressions.len() == 1 { let sort = expressions[0].clone(); let data_type = sort.expr.data_type(schema.as_ref())?; downcast_primitive! { - data_type => (primitive_merge_helper, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary), - DataType::Utf8 => merge_helper!(StringArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) - DataType::Utf8View => merge_helper!(StringViewArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) - DataType::LargeUtf8 => merge_helper!(LargeStringArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) - DataType::Binary => merge_helper!(BinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) - DataType::LargeBinary => merge_helper!(LargeBinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, flush_on_input_batch_boundary) + data_type => (primitive_merge_helper, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, batch_memory_budget), + DataType::Utf8 => merge_helper!(StringArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, batch_memory_budget) + DataType::Utf8View => merge_helper!(StringViewArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, batch_memory_budget) + DataType::LargeUtf8 => merge_helper!(LargeStringArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, batch_memory_budget) + DataType::Binary => merge_helper!(BinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, batch_memory_budget) + DataType::LargeBinary => merge_helper!(LargeBinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation, enable_round_robin_tie_breaker, batch_memory_budget) _ => {} } } @@ -299,14 +317,14 @@ impl<'a> StreamingMergeBuilder<'a> { reservation, enable_round_robin_tie_breaker, ) - .with_flush_on_input_batch_boundary(flush_on_input_batch_boundary) + .with_batch_memory_budget(batch_memory_budget) .into_stream()) } } #[cfg(test)] mod tests { - use crate::spill::spill_manager::GetSlicedSize; + use crate::spill::{get_record_batch_memory_size, spill_manager::GetSlicedSize}; use crate::{common::collect, stream::RecordBatchStreamAdapter}; use std::sync::Arc; @@ -314,7 +332,7 @@ mod tests { use arrow::array::{ArrayRef, RecordBatch}; use arrow::compute::{cast, concat_batches}; - use arrow::datatypes::{Field, Int32Type, Schema}; + use arrow::datatypes::{Field, Int32Type, Int64Type, Schema}; use arrow_schema::SortOptions; use datafusion_common::Result; use datafusion_execution::TaskContext; @@ -322,16 +340,16 @@ mod tests { use datafusion_physical_expr_common::metrics::{ ExecutionPlanMetricsSet, SpillMetrics, }; + use futures::StreamExt; #[rstest::rstest] #[case::primitive(DataType::Int32)] #[case::strings(DataType::Utf8)] #[case::views(DataType::Utf8View)] #[tokio::test] - async fn intermediate_merge_flushes_before_replacing_input_batch( + async fn intermediate_merge_bounds_short_batches( #[case] data_type: DataType, #[values(false, true)] row_cursor: bool, - #[values(false, true)] round_robin: bool, #[values(None, Some(7))] fetch: Option, ) -> Result<()> { let schema = @@ -344,9 +362,11 @@ mod tests { [sort].into() }; - for flush in [false, true] { + for budgeted in [false, true] { + let mut input_batch_sizes = Vec::new(); let streams = (0..2) .map(|run| { + let mut max_batch_bytes = 0; let batches = (0..3) .map(|batch| { let values = [4 * batch + run, 4 * batch + run + 2]; @@ -366,15 +386,21 @@ mod tests { } _ => unreachable!(), }; - Ok(RecordBatch::try_new(Arc::clone(&schema), vec![array])?) + let batch = + RecordBatch::try_new(Arc::clone(&schema), vec![array])?; + max_batch_bytes = + max_batch_bytes.max(get_record_batch_memory_size(&batch)); + Ok(batch) }) .collect::>>(); + input_batch_sizes.push(max_batch_bytes); Box::pin(RecordBatchStreamAdapter::new( Arc::clone(&schema), futures::stream::iter(batches), )) as SendableRecordBatchStream }) .collect(); + let max_output_bytes = input_batch_sizes.iter().sum::(); let stream = StreamingMergeBuilder::new() .with_schema(Arc::clone(&schema)) .with_expressions(&ordering) @@ -382,8 +408,10 @@ mod tests { .with_streams(streams) .with_batch_size(8192) .with_fetch(fetch) - .with_round_robin_tie_breaker(round_robin) - .with_flush_on_input_batch_boundary(flush) + .with_batch_memory_budget(budgeted.then_some(MergeBatchMemoryBudget { + memory_limit: 2 * max_output_bytes, + input_batch_sizes, + })) .with_bypass_mempool() .build()?; let batches = collect(stream).await?; @@ -392,22 +420,9 @@ mod tests { let expected = Int32Array::from_iter_values(0..fetch.unwrap_or(12) as i32); assert_eq!(actual.as_primitive::(), &expected); - if flush { - // Each output must contain rows from at most one source batch - // of either input, even though the target is 8192 rows. + if budgeted { for batch in &batches { - let values = cast(batch.column(0), &DataType::Int32)?; - for run in 0..2 { - let mut source_batches = values - .as_primitive::() - .values() - .iter() - .filter(|&&value| value % 2 == run) - .map(|value| value / 4); - if let Some(first) = source_batches.next() { - assert!(source_batches.all(|batch| batch == first)); - } - } + assert!(get_record_batch_memory_size(batch) <= max_output_bytes); } } else { assert_eq!(batches.len(), 1, "ordinary merges retain their batching"); @@ -416,6 +431,158 @@ mod tests { Ok(()) } + #[rstest::rstest] + #[case::short_inputs(1000)] + #[case::full_inputs(8192)] + #[tokio::test] + async fn intermediate_merge_keeps_full_output_batches( + #[case] input_rows: usize, + #[values(false, true)] row_cursor: bool, + ) -> Result<()> { + const RUNS: usize = 8; + const BATCHES: usize = 10; + const OUTPUT_ROWS: usize = 8192; + let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)])); + let sort = PhysicalSortExpr::new_default(col("x", &schema)?); + let ordering = if row_cursor { + [sort.clone(), sort].into() + } else { + [sort].into() + }; + let mut input_batch_sizes = Vec::new(); + let mut streams = Vec::new(); + for run in 0..RUNS { + let mut batches = Vec::new(); + let mut max_batch_bytes = 0; + for batch_index in 0..BATCHES { + let values = + Int64Array::from_iter_values((0..input_rows).map(|row| { + ((batch_index * input_rows + row) * RUNS + run) as i64 + })); + let batch = + RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(values)])?; + max_batch_bytes = + max_batch_bytes.max(get_record_batch_memory_size(&batch)); + batches.push(Ok(batch)); + } + input_batch_sizes.push(max_batch_bytes); + streams.push(Box::pin(RecordBatchStreamAdapter::new( + Arc::clone(&schema), + futures::stream::iter(batches), + )) as SendableRecordBatchStream); + } + let memory_limit = 2 * input_batch_sizes.iter().sum::(); + let stream = StreamingMergeBuilder::new() + .with_schema(Arc::clone(&schema)) + .with_expressions(&ordering) + .with_metrics(BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0)) + .with_streams(streams) + .with_batch_size(OUTPUT_ROWS) + .with_batch_memory_budget(Some(MergeBatchMemoryBudget { + memory_limit, + input_batch_sizes, + })) + .with_bypass_mempool() + .build()?; + let batches = collect(stream).await?; + if input_rows == OUTPUT_ROWS { + assert_eq!(batches.len(), RUNS * BATCHES); + assert!(batches.iter().all(|batch| batch.num_rows() == OUTPUT_ROWS)); + } else { + // Unconditional boundary flushing produces 80 batches here, mostly + // one-row fragments. Future rows from every live input still need + // a conservative allowance, so short inputs can require early + // output, but should at least halve the number of batches. + assert!( + batches.len() <= RUNS * BATCHES / 2, + "{} output batches", + batches.len() + ); + } + let merged = concat_batches(&schema, &batches)?; + assert_eq!( + merged.column(0).as_primitive::(), + &Int64Array::from_iter_values(0..(RUNS * BATCHES * input_rows) as i64) + ); + Ok(()) + } + + #[rstest::rstest] + #[tokio::test] + async fn intermediate_merge_preserves_ties_across_pending_input_boundaries( + #[values(false, true)] row_cursor: bool, + #[values(false, true)] round_robin: bool, + ) -> Result<()> { + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Int32, false), + Field::new("source", DataType::Int32, false), + ])); + let sort = PhysicalSortExpr::new_default(col("key", &schema)?); + let ordering = if row_cursor { + [sort.clone(), sort].into() + } else { + [sort].into() + }; + let mut expected = None; + for budgeted in [false, true] { + let mut streams = Vec::new(); + let mut input_batch_sizes = Vec::new(); + for source in 0..2 { + let mut batches = Vec::new(); + let mut max_batch_bytes = 0; + for batch_index in 0..3 { + // Equal keys span two input batches, and source tags make + // changes to the tie-breaking order observable. + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(Int32Array::from(vec![batch_index / 2; 2])), + Arc::new(Int32Array::from(vec![source; 2])), + ], + )?; + max_batch_bytes = + max_batch_bytes.max(get_record_batch_memory_size(&batch)); + batches.push(Ok(batch)); + } + input_batch_sizes.push(max_batch_bytes); + let input = futures::stream::iter(batches).then(|batch| async { + tokio::task::yield_now().await; + batch + }); + streams.push(Box::pin(RecordBatchStreamAdapter::new( + Arc::clone(&schema), + input, + )) as SendableRecordBatchStream); + } + let memory_limit = 2 * input_batch_sizes.iter().sum::(); + let stream = StreamingMergeBuilder::new() + .with_schema(Arc::clone(&schema)) + .with_expressions(&ordering) + .with_metrics(BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0)) + .with_streams(streams) + .with_batch_size(8192) + .with_round_robin_tie_breaker(round_robin) + .with_batch_memory_budget(budgeted.then_some(MergeBatchMemoryBudget { + memory_limit, + input_batch_sizes, + })) + .with_bypass_mempool() + .build()?; + let batches = collect(stream).await?; + let merged = concat_batches(&schema, &batches)?; + assert_eq!( + merged.column(0).as_primitive::(), + &Int32Array::from(vec![0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1]) + ); + if let Some(expected) = &expected { + assert_eq!(&merged, expected); + } else { + expected = Some(merged); + } + } + Ok(()) + } + #[tokio::test] async fn intermediate_merge_releases_consumed_dictionary_batches() -> Result<()> { let schema = Arc::new(Schema::new(vec![ @@ -430,6 +597,7 @@ mod tests { ), ])); let mut max_input_bytes = 0; + let mut max_input_memory = 0; let mut streams = Vec::new(); for run in 0..2 { let mut batches = Vec::new(); @@ -451,6 +619,8 @@ mod tests { vec![Arc::new(keys), Arc::new(payload)], )?; max_input_bytes = max_input_bytes.max(batch.get_sliced_size()?); + max_input_memory = + max_input_memory.max(get_record_batch_memory_size(&batch)); batches.push(Ok(batch)); } streams.push(Box::pin(RecordBatchStreamAdapter::new( @@ -465,7 +635,10 @@ mod tests { .with_metrics(BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0)) .with_streams(streams) .with_batch_size(8192) - .with_flush_on_input_batch_boundary(true) + .with_batch_memory_budget(Some(MergeBatchMemoryBudget { + memory_limit: 4 * max_input_memory, + input_batch_sizes: vec![max_input_memory; 2], + })) .with_bypass_mempool() .build()?; let batches = collect(stream).await?; diff --git a/datafusion/sqllogictest/test_files/aggregate_memory_spill.slt b/datafusion/sqllogictest/test_files/aggregate_memory_spill.slt index 740fc8802a3d6..53d515bf7fd13 100644 --- a/datafusion/sqllogictest/test_files/aggregate_memory_spill.slt +++ b/datafusion/sqllogictest/test_files/aggregate_memory_spill.slt @@ -198,11 +198,6 @@ FROM ( statement ok SET datafusion.execution.target_partitions = 4 -# Bound merge buffers so each partition can allocate replay state while the -# other partitions retain aggregate state in the shared greedy memory pool. -statement ok -SET datafusion.runtime.max_spill_merge_fan_in = 2 - query II SELECT count(*), sum(total) FROM ( @@ -232,9 +227,6 @@ FROM ( # Restore settings to slt runner defaults -statement ok -RESET datafusion.runtime.max_spill_merge_fan_in - statement ok RESET datafusion.runtime.memory_limit