diff --git a/datafusion/physical-plan/src/aggregates/grouped_hash_stream.rs b/datafusion/physical-plan/src/aggregates/grouped_hash_stream.rs index d8f2543996f0e..544d75a24393f 100644 --- a/datafusion/physical-plan/src/aggregates/grouped_hash_stream.rs +++ b/datafusion/physical-plan/src/aggregates/grouped_hash_stream.rs @@ -1367,6 +1367,7 @@ impl GroupedHashAggregateStream { .with_metrics(self.baseline_metrics.clone()) .with_batch_size(self.batch_size) .with_reservation(self.reservation.new_empty()) + .with_replay_headroom() .build()?; self.input_done = false; @@ -1394,6 +1395,9 @@ impl GroupedHashAggregateStream { // to ensure we don't spill the spilled data to disk again. self.oom_mode = OutOfMemoryMode::ReportError; + // Release unused initial capacity from recreated group values so it + // does not consume the memory available for spill replay. + self.group_values.clear_shrink(0); self.update_memory_reservation()?; ExecutionState::ReadingInput diff --git a/datafusion/physical-plan/src/aggregates/hash_stream.rs b/datafusion/physical-plan/src/aggregates/hash_stream.rs index 152082d88dd90..8e74e7bec78a1 100644 --- a/datafusion/physical-plan/src/aggregates/hash_stream.rs +++ b/datafusion/physical-plan/src/aggregates/hash_stream.rs @@ -348,6 +348,7 @@ impl FinalSpillContext { .with_metrics(baseline_metrics.intermediate()) .with_batch_size(batch_size) .with_reservation(merge_reservation) + .with_replay_headroom() .build()?; let replay = OrderedFinalAggregateStream::new_with_input_and_metrics( &final_agg, diff --git a/datafusion/physical-plan/src/aggregates/mod.rs b/datafusion/physical-plan/src/aggregates/mod.rs index 5b0d2c2d2ea6c..bc7e6cae3f85e 100644 --- a/datafusion/physical-plan/src/aggregates/mod.rs +++ b/datafusion/physical-plan/src/aggregates/mod.rs @@ -3209,16 +3209,19 @@ mod tests { BlockingExec, PanicExec, StatisticsExec, assert_strong_count_converges_to_zero, }; + use arrow::array::AsArray; use arrow::array::{ BooleanArray, DictionaryArray, Float32Array, Float64Array, Int32Array, Int64Array, NullArray, StructArray, UInt32Array, UInt64Array, }; use arrow::compute::{SortOptions, concat_batches}; - use arrow::datatypes::Int32Type; + use arrow::datatypes::{Int32Type, Int64Type}; use datafusion_common::test_util::{batches_to_sort_string, batches_to_string}; use datafusion_common::{DataFusionError, internal_err}; use datafusion_execution::config::SessionConfig; - use datafusion_execution::memory_pool::FairSpillPool; + use datafusion_execution::memory_pool::{ + FairSpillPool, MemoryPool, PeakRecordingPool, + }; use datafusion_execution::runtime_env::RuntimeEnvBuilder; use datafusion_expr::function::{AccumulatorArgs, StateFieldsArgs}; use datafusion_expr::{ @@ -3423,6 +3426,391 @@ mod tests { )) } + #[rstest::rstest] + #[case::single(AggregateMode::Single)] + #[case::final_stage(AggregateMode::Final)] + #[tokio::test] + async fn legacy_aggregate_spill_merge_leaves_memory_for_replay( + #[case] mode: AggregateMode, + ) -> Result<()> { + use arrow::array::ListArray; + use arrow::buffer::OffsetBuffer; + + const KEYS: usize = 64; + const VALUES_PER_KEY: i64 = 64; + const MEMORY_LIMIT: usize = 8192; + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Int64, false), + Field::new("value", DataType::Int64, false), + ])); + let group_by = + PhysicalGroupBy::new_single(vec![(col("key", &schema)?, "key".into())]); + let aggregates = vec![Arc::new( + AggregateExprBuilder::new(array_agg_udaf(), vec![col("value", &schema)?]) + .schema(Arc::clone(&schema)) + .alias("values") + .build()?, + )]; + let input_schema = if mode == AggregateMode::Single { + Arc::clone(&schema) + } else { + Arc::new(create_schema( + &schema, + &group_by, + &aggregates, + AggregateMode::Partial, + )?) + }; + let mut batches = vec![]; + // Each group grows across many one-row replay batches before it can be + // emitted. A merge that fills the allowance starves this state growth. + for value in 1..=VALUES_PER_KEY { + for key in (0..KEYS as i64).rev() { + let values: ArrayRef = Arc::new(Int64Array::from(vec![value])); + let values = if mode == AggregateMode::Single { + values + } else { + let DataType::List(field) = input_schema.field(1).data_type() else { + unreachable!("ARRAY_AGG state must be a list") + }; + Arc::new(ListArray::new( + Arc::clone(field), + OffsetBuffer::from_lengths([1]), + values, + None, + )) as ArrayRef + }; + batches.push(RecordBatch::try_new( + Arc::clone(&input_schema), + vec![Arc::new(Int64Array::from(vec![key])), values], + )?); + } + } + let input = TestMemoryExec::try_new_exec(&[batches], input_schema, None)?; + let aggregate = AggregateExec::try_new( + mode, + group_by, + aggregates, + vec![None], + input, + schema, + )?; + let pool = Arc::new(PeakRecordingPool::new(Arc::new(FairSpillPool::new( + MEMORY_LIMIT, + )))); + let context = Arc::new( + TaskContext::default() + .with_session_config( + SessionConfig::new().with_batch_size(1).set_bool( + "datafusion.execution.enable_migration_aggregate", + false, + ), + ) + .with_runtime( + RuntimeEnvBuilder::new() + .with_memory_pool(Arc::clone(&pool) as Arc) + .build_arc()?, + ), + ); + let stream = aggregate.execute_typed(0, &context)?; + assert!(matches!(stream, StreamType::GroupedHash(_))); + let output = collect(stream.into()).await?; + let mut seen = HashSet::new(); + for batch in output { + assert!( + batch + .columns() + .iter() + .all(|column| column.null_count() == 0) + ); + let keys = batch.column(0).as_primitive::(); + let values = batch.column(1).as_list::(); + for row in 0..batch.num_rows() { + let key = keys.value(row); + assert!((0..KEYS as i64).contains(&key)); + assert!(seen.insert(key), "duplicate group"); + let values = values.value(row); + assert_eq!(values.null_count(), 0); + let mut values = values.as_primitive::().values().to_vec(); + values.sort_unstable(); + assert_eq!(values, (1..=VALUES_PER_KEY).collect::>()); + } + } + assert_eq!(seen.len(), KEYS); + assert!(aggregate.metrics().unwrap().spill_count().unwrap() > 1); + assert!(pool.peak_reserved() <= MEMORY_LIMIT); + assert_eq!(pool.reserved(), 0); + let progress = context.runtime_env().disk_manager.spilling_progress(); + assert_eq!(progress.current_bytes, 0); + assert_eq!(progress.active_files_count, 0); + Ok(()) + } + + // This high-cardinality memory test would create quadratic collision scratch + // space; small group-value tests cover forced hash-collision correctness. + #[cfg(not(feature = "force_hash_collisions"))] + #[rstest::rstest] + #[case::final_hash(AggregateMode::Final, false)] + #[case::single_hash(AggregateMode::Single, false)] + #[case::ordered_final(AggregateMode::Final, true)] + #[case::ordered_single(AggregateMode::Single, true)] + #[tokio::test] + async fn migrated_aggregate_spill_merge_leaves_memory_for_replay( + #[case] mode: AggregateMode, + #[case] ordered: bool, + #[values(false, true)] with_peer: bool, + ) -> Result<()> { + use arrow::array::{ListArray, StringArray}; + use arrow::buffer::OffsetBuffer; + use datafusion_execution::memory_pool::MemoryConsumer; + + const BATCH_SIZE: usize = 8192; + const KEYS_PER_PREFIX: usize = 25 * BATCH_SIZE; + const MEMORY_LIMIT: usize = 2 * 1024 * 1024; + let schema = Arc::new(Schema::new(vec![ + Field::new("prefix", DataType::Int64, false), + Field::new("key", DataType::Int64, false), + Field::new("value", DataType::Utf8, false), + ])); + let group_by = PhysicalGroupBy::new_single(vec![ + (col("prefix", &schema)?, "prefix".to_string()), + (col("key", &schema)?, "key".to_string()), + ]); + let aggregates = vec![Arc::new( + AggregateExprBuilder::new(array_agg_udaf(), vec![col("value", &schema)?]) + .schema(Arc::clone(&schema)) + .alias("values") + .build()?, + )]; + let input_schema = if mode == AggregateMode::Single { + Arc::clone(&schema) + } else { + Arc::new(create_schema( + &schema, + &group_by, + &aggregates, + AggregateMode::Partial, + )?) + }; + let mut batches = vec![]; + for prefix in 0..2 { + // Repeat every key across spill runs. Only the prefix is ordered; + // the second pass restarts the descending key sequence. + for value in 1..=2 { + for start in (0..KEYS_PER_PREFIX).step_by(BATCH_SIZE).rev() { + let values: ArrayRef = + Arc::new(StringArray::from(vec![ + format!("{value:08}"); + BATCH_SIZE + ])); + // Final takes singleton ARRAY_AGG states rather than raw strings. + let values = if mode == AggregateMode::Single { + values + } else { + let DataType::List(field) = input_schema.field(2).data_type() + else { + unreachable!("ARRAY_AGG state must be a list") + }; + Arc::new(ListArray::new( + Arc::clone(field), + OffsetBuffer::from_lengths(std::iter::repeat_n( + 1, BATCH_SIZE, + )), + values, + None, + )) as ArrayRef + }; + batches.push(RecordBatch::try_new( + Arc::clone(&input_schema), + vec![ + Arc::new(Int64Array::from(vec![prefix; BATCH_SIZE])), + Arc::new(Int64Array::from_iter_values( + (start..start + BATCH_SIZE).rev().map(|key| key as i64), + )), + values, + ], + )?); + } + } + } + + let mut input = + TestMemoryExec::try_new(&[batches], Arc::clone(&input_schema), None)?; + if ordered { + input = input.try_with_sort_information(vec![ + LexOrdering::new([PhysicalSortExpr::new_default(col( + "prefix", &schema, + )?)]) + .unwrap(), + ])?; + } + let aggregate = AggregateExec::try_new( + mode, + group_by, + aggregates, + vec![None], + Arc::new(input), + Arc::clone(&schema), + )?; + // Keep the aggregate's allowance at 2 MiB even with a second consumer. + // Falling back to half the global limit would consume its entire share. + let pool_limit = MEMORY_LIMIT * if with_peer { 2 } else { 1 }; + let pool = Arc::new(PeakRecordingPool::new(Arc::new(FairSpillPool::new( + pool_limit, + )))); + let memory_pool = Arc::clone(&pool) as Arc; + let _peer = with_peer.then(|| { + MemoryConsumer::new("other spilling operator") + .with_can_spill(true) + .register(&memory_pool) + }); + let context = Arc::new( + TaskContext::default() + .with_session_config(migrated_hash_session_config(BATCH_SIZE)) + .with_runtime( + RuntimeEnvBuilder::new() + .with_memory_pool(memory_pool) + .build_arc()?, + ), + ); + let stream = aggregate.execute_typed(0, &context)?; + match (mode, ordered, &stream) { + (AggregateMode::Final, false, StreamType::FinalHash(_)) + | (AggregateMode::Single, false, StreamType::SingleHash(_)) => {} + (AggregateMode::Final, true, StreamType::OrderedFinalAggregate(_)) + | (AggregateMode::Single, true, StreamType::OrderedSingleAggregate(_)) => { + assert_eq!( + aggregate.input_order_mode(), + &InputOrderMode::PartiallySorted(vec![0]) + ); + } + _ => panic!("unexpected stream for {mode:?}, ordered={ordered}"), + } + let result = collect(stream.into()).await.unwrap_or_else(|error| { + panic!("{mode:?}, ordered={ordered}, with_peer={with_peer}: {error}") + }); + let mut seen = HashSet::new(); + for batch in &result { + assert!( + batch + .columns() + .iter() + .all(|column| column.null_count() == 0) + ); + let columns = batch + .columns() + .iter() + .take(2) + .map(|column| column.as_primitive::()) + .collect::>(); + let values = batch.column(2).as_list::(); + for row in 0..batch.num_rows() { + let prefix = columns[0].value(row); + let key = columns[1].value(row); + assert!((0..2).contains(&prefix)); + assert!((0..KEYS_PER_PREFIX as i64).contains(&key)); + let values = values.value(row); + assert_eq!(values.len(), 2); + assert_eq!(values.null_count(), 0); + let values = values.as_string::(); + let mut values = [values.value(0), values.value(1)]; + values.sort_unstable(); + assert_eq!(values, ["00000001", "00000002"]); + assert!(seen.insert((prefix, key)), "duplicate group"); + } + } + assert_eq!(seen.len(), 2 * KEYS_PER_PREFIX); + assert!(pool.peak_reserved() <= MEMORY_LIMIT); + let metrics = aggregate.metrics().unwrap(); + assert!(metrics.spill_count().unwrap() > 1); + assert!(metrics.spilled_rows().unwrap() > 0); + assert!(metrics.spilled_bytes().unwrap() > 0); + assert_eq!(context.memory_pool().reserved(), 0); + let progress = context.runtime_env().disk_manager.spilling_progress(); + assert_eq!(progress.current_bytes, 0); + assert_eq!(progress.active_files_count, 0); + Ok(()) + } + + #[tokio::test] + async fn migrated_aggregate_spill_merge_allows_indivisible_rows() -> Result<()> { + use arrow::array::StringArray; + use datafusion_execution::memory_pool::GreedyMemoryPool; + + const KEY_BYTES: usize = 350_000; + const GROUPS: usize = 24; + const MEMORY_LIMIT: usize = 5 * 1024 * 1024 / 2; + let schema = Arc::new(Schema::new(vec![ + Field::new("key", DataType::Utf8, false), + Field::new("value", DataType::Int64, false), + ])); + let mut batches = Vec::new(); + for _ in 0..2 { + for key in (0..GROUPS).rev() { + batches.push(RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(StringArray::from(vec![format!( + "{key:02}{}", + "x".repeat(KEY_BYTES - 2) + )])), + Arc::new(Int64Array::from(vec![1])), + ], + )?); + } + } + let input = TestMemoryExec::try_new(&[batches], Arc::clone(&schema), None)?; + let aggregate = AggregateExec::try_new( + AggregateMode::Single, + PhysicalGroupBy::new_single(vec![(col("key", &schema)?, "key".into())]), + vec![Arc::new( + AggregateExprBuilder::new(sum_udaf(), vec![col("value", &schema)?]) + .schema(Arc::clone(&schema)) + .alias("sum") + .build()?, + )], + vec![None], + Arc::new(input), + schema, + )?; + let pool = Arc::new(PeakRecordingPool::new(Arc::new(GreedyMemoryPool::new( + MEMORY_LIMIT, + )))); + let context = Arc::new( + TaskContext::default() + .with_session_config(migrated_hash_session_config(1)) + .with_runtime( + RuntimeEnvBuilder::new() + .with_memory_pool(Arc::clone(&pool) as Arc) + .build_arc()?, + ), + ); + + // Two one-row spill inputs need 1,400,064 bytes for merge buffers, more + // than half the pool. They cannot shrink, but merge plus replay fits. + let result = collect(aggregate.execute(0, Arc::clone(&context))?).await?; + let mut seen = HashSet::new(); + for batch in result { + let keys = batch.column(0).as_string::(); + let sums = batch.column(1).as_primitive::(); + for row in 0..batch.num_rows() { + let key = keys.value(row); + assert_eq!(key.len(), KEY_BYTES); + let group = key[..2].parse::().unwrap(); + assert!(group < GROUPS && seen.insert(group)); + assert_eq!(sums.value(row), 2); + } + } + assert_eq!(seen.len(), GROUPS); + assert!(aggregate.metrics().unwrap().spill_count().unwrap() > 1); + assert!(pool.peak_reserved() <= MEMORY_LIMIT); + assert_eq!(pool.reserved(), 0); + let progress = context.runtime_env().disk_manager.spilling_progress(); + assert_eq!(progress.current_bytes, 0); + assert_eq!(progress.active_files_count, 0); + Ok(()) + } + async fn check_grouping_sets( input: Arc, spill: bool, @@ -7391,32 +7779,22 @@ mod tests { Field::new("c", DataType::Int64, false), ])); - let batches = vec![vec![ - RecordBatch::try_new( - Arc::clone(&schema), - vec![ - Arc::new(Int64Array::from(vec![2])), - Arc::new(Int64Array::from(vec![2])), - Arc::new(Int64Array::from(vec![1])), - ], - )?, - RecordBatch::try_new( - Arc::clone(&schema), - vec![ - Arc::new(Int64Array::from(vec![1])), - Arc::new(Int64Array::from(vec![1])), - Arc::new(Int64Array::from(vec![1])), - ], - )?, - RecordBatch::try_new( - Arc::clone(&schema), - vec![ - Arc::new(Int64Array::from(vec![0])), - Arc::new(Int64Array::from(vec![0])), - Arc::new(Int64Array::from(vec![1])), - ], - )?, - ]]; + let mut descending_batches = Vec::new(); + for ordered_group in (0_i64..3).rev() { + // Multiple groups sharing the ordered prefix must remain in memory + // until its boundary, which forces an actual aggregation spill. + for unordered_group in 1_i64..=16 { + descending_batches.push(RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(Int64Array::from(vec![ordered_group])), + Arc::new(Int64Array::from(vec![ordered_group])), + Arc::new(Int64Array::from(vec![unordered_group])), + ], + )?); + } + } + let batches = vec![descending_batches]; let scan = TestMemoryExec::try_new(&batches, Arc::clone(&schema), None)?; let scan = scan.try_with_sort_information(vec![ LexOrdering::new([PhysicalSortExpr::new( @@ -7448,30 +7826,40 @@ mod tests { Arc::clone(&schema), )?); - let task_ctx = new_migrated_spill_ctx(1, 600); + // The merge input and replay aggregate now share one allowance. Keep + // enough space for both while the additional groups still force spilling. + let task_ctx = new_migrated_spill_ctx(1, 1024); let result = collect(aggr.execute(0, Arc::clone(&task_ctx))?).await?; + assert_eq!(task_ctx.memory_pool().reserved(), 0); assert_spill_count_metric(true, Arc::clone(&aggr)); - let metrics = aggr.metrics().unwrap(); - for phase in ["update", "state", "merge", "evaluate"] { - let time = metrics - .sum_by_name(&format!("agg_expr_0_{phase}_time")) - .unwrap_or_else(|| { - panic!("migrated single aggregate records {phase} time") - }); - assert!(time.as_usize() > 0); - } + assert_accumulator_phase_times(&aggr, &["update", "state", "merge", "evaluate"]); - allow_duplicates! { - assert_snapshot!(batches_to_string(&result), @r" - +---+---+--------+ - | b | c | SUM(c) | - +---+---+--------+ - | 2 | 1 | 1 | - | 1 | 1 | 1 | - | 0 | 1 | 1 | - +---+---+--------+ - "); - } + let batch = concat_batches(&result[0].schema(), &result)?; + assert!( + batch + .columns() + .iter() + .all(|column| column.null_count() == 0) + ); + let columns = batch + .columns() + .iter() + .map(|column| column.as_primitive::()) + .collect::>(); + let actual = (0..batch.num_rows()) + .map(|row| { + ( + columns[0].value(row), + columns[1].value(row), + columns[2].value(row), + ) + }) + .collect::>(); + let expected = (0_i64..3) + .rev() + .flat_map(|prefix| (1_i64..=16).map(move |key| (prefix, key, key))) + .collect::>(); + assert_eq!(actual, expected); Ok(()) } diff --git a/datafusion/physical-plan/src/aggregates/ordered_final_stream.rs b/datafusion/physical-plan/src/aggregates/ordered_final_stream.rs index ac6f2491b481e..cc992eaa51181 100644 --- a/datafusion/physical-plan/src/aggregates/ordered_final_stream.rs +++ b/datafusion/physical-plan/src/aggregates/ordered_final_stream.rs @@ -243,6 +243,7 @@ impl OrderedFinalSpillContext { .with_metrics(baseline_metrics.intermediate()) .with_batch_size(batch_size) .with_reservation(merge_reservation) + .with_replay_headroom() .build()?; let replay = OrderedFinalAggregateStream::new_with_input_and_metrics( &agg, diff --git a/datafusion/physical-plan/src/aggregates/ordered_single_stream.rs b/datafusion/physical-plan/src/aggregates/ordered_single_stream.rs index 701dee4e8146d..40ba90729b55f 100644 --- a/datafusion/physical-plan/src/aggregates/ordered_single_stream.rs +++ b/datafusion/physical-plan/src/aggregates/ordered_single_stream.rs @@ -303,6 +303,7 @@ impl OrderedSingleSpillContext { .with_metrics(baseline_metrics.intermediate()) .with_batch_size(batch_size) .with_reservation(merge_reservation) + .with_replay_headroom() .build()?; let replay = OrderedFinalAggregateStream::new_with_input_and_metrics( &final_agg, diff --git a/datafusion/physical-plan/src/aggregates/single_stream.rs b/datafusion/physical-plan/src/aggregates/single_stream.rs index cfa708c51129b..091e9fb940bce 100644 --- a/datafusion/physical-plan/src/aggregates/single_stream.rs +++ b/datafusion/physical-plan/src/aggregates/single_stream.rs @@ -329,6 +329,7 @@ impl SingleSpillContext { .with_metrics(baseline_metrics.intermediate()) .with_batch_size(batch_size) .with_reservation(merge_reservation) + .with_replay_headroom() .build()?; let replay = OrderedFinalAggregateStream::new_with_input_and_metrics( &final_agg, diff --git a/datafusion/physical-plan/src/sorts/multi_level_merge.rs b/datafusion/physical-plan/src/sorts/multi_level_merge.rs index b5aa5c4d54015..c1fe893e0df57 100644 --- a/datafusion/physical-plan/src/sorts/multi_level_merge.rs +++ b/datafusion/physical-plan/src/sorts/multi_level_merge.rs @@ -33,6 +33,8 @@ 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::spill::gc_view_arrays; +use crate::spill::spill_manager::GetSlicedSize; use crate::stream::{ObservedStream, RecordBatchStreamAdapter}; use datafusion_execution::{RecordBatchStream, SendableRecordBatchStream}; use datafusion_physical_expr_common::sort_expr::LexOrdering; @@ -155,6 +157,8 @@ pub(crate) struct MultiLevelMergeBuilder { reservation: MemoryReservation, /// Workspace retained across retries and intermediate spill passes. merge_pool: Option>, + /// Leave memory for the aggregate consuming the merged spill rows. + reserve_replay_headroom: bool, fetch: Option, enable_round_robin_tie_breaker: bool, } @@ -194,6 +198,7 @@ impl MultiLevelMergeBuilder { batch_size, reservation, merge_pool: None, + reserve_replay_headroom: false, enable_round_robin_tie_breaker, fetch, } @@ -204,6 +209,13 @@ impl MultiLevelMergeBuilder { self } + /// Leave replay headroom while selecting merge buffers. Temporary splitting + /// workspace can still 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 + } + pub(crate) fn create_spillable_merge_stream(self) -> SendableRecordBatchStream { Box::pin(RecordBatchStreamAdapter::new( Arc::clone(&self.schema), @@ -212,23 +224,35 @@ 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.merge_sorted_runs_within_mem_limit()? { - MergeStep::Stream { - stream, - batch_size_limit, - } => (stream, batch_size_limit), - 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 - // size so its largest batch shrinks, lowering the per-stream - // reservation, then retry. Makes the merge resilient to skewed - // (very wide) rows. - self.split_spill_file_in_half(index).await?; - continue; + let (mut stream, batch_size_limit) = match self + .merge_sorted_runs_within_mem_limit(allow_minimum_without_headroom)? + { + MergeStep::Stream { + stream, + batch_size_limit, + } => (stream, batch_size_limit), + 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 + // size so its largest batch shrinks, lowering the per-stream + // reservation, then retry. Makes the merge resilient to skewed + // (very wide) rows. + let retry_unsplittable = + self.reserve_replay_headroom && !allow_minimum_without_headroom; + if !self + .split_spill_file_in_half(index, retry_unsplittable) + .await? + { + // A single row may prevent leaving replay headroom while + // the minimum merge still fits the actual shared pool. + allow_minimum_without_headroom = true; } - }; + continue; + } + }; + allow_minimum_without_headroom = false; // TODO - add a threshold for number of files to disk even if empty and reading from disk so // we can avoid the memory reservation @@ -278,7 +302,10 @@ impl MultiLevelMergeBuilder { /// This tries to create a stream that merges the most sorted streams and sorted spill files /// as possible within the memory limit. - fn merge_sorted_runs_within_mem_limit(&mut self) -> Result { + fn merge_sorted_runs_within_mem_limit( + &mut self, + allow_minimum_without_headroom: bool, + ) -> Result { match (self.sorted_spill_files.len(), self.sorted_streams.len()) { // No data so empty batch (0, 0) => { @@ -352,6 +379,7 @@ impl MultiLevelMergeBuilder { // we must have at least 2 streams to merge minimum_number_of_required_streams, &mut memory_reservation, + allow_minimum_without_headroom, )? { SpillFilesToMerge::Ready(sorted_spill_files, buffer_size) => { (sorted_spill_files, buffer_size) @@ -483,6 +511,7 @@ impl MultiLevelMergeBuilder { buffer_len: usize, minimum_number_of_required_streams: usize, reservation: &mut MemoryReservation, + allow_minimum_without_headroom: bool, ) -> Result { assert_ne!(buffer_len, 0, "Buffer length must be greater than 0"); let mut number_of_spills_to_read_for_current_phase = 0; @@ -497,9 +526,14 @@ impl MultiLevelMergeBuilder { // 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 { + 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) + { break; } @@ -510,12 +544,30 @@ impl MultiLevelMergeBuilder { ) * buffer_len; total_needed += per_spill; - // For memory pools that are not shared this is good, for other - // this is not and there should be some upper limit to memory - // reservation so we won't starve the system. - match try_grow_reservation_to_at_least(reservation, total_needed) { + // 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. + let skip_headroom = allow_minimum_without_headroom + && 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 admission = if check_headroom { + // Ask the pool for merge buffers plus equal replay space, then + // return the spare bytes before exposing the merge stream. + match total_needed.checked_mul(2) { + Some(with_headroom) => { + try_grow_reservation_to_at_least(reservation, with_headroom) + } + None => resources_err!("Spill merge headroom exceeds usize::MAX"), + } + } else { + try_grow_reservation_to_at_least(reservation, total_needed) + }; + match admission { Ok(_) => { number_of_spills_to_read_for_current_phase += 1; + accepted_memory = total_needed; } // If we can't grow the reservation, we need to stop Err(err) => { @@ -532,11 +584,17 @@ impl MultiLevelMergeBuilder { buffer_len - 1, minimum_number_of_required_streams, reservation, + allow_minimum_without_headroom, ); } // buffer_len == 1 and we still can't seat the minimum of 2 streams. if number_of_spills_to_read_for_current_phase == 0 { + if check_headroom { + // Replay has not started, so splitting may use the + // full pool to leave room for replay afterward. + return Ok(SpillFilesToMerge::SplitThenRetry(0)); + } // We couldn't even reserve a single stream - one record batch // is larger than the whole merge budget. That's the lone-batch // case, not the 2-stream merge skew we rescue here - surface it. @@ -562,6 +620,12 @@ impl MultiLevelMergeBuilder { } } + 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); + } + let spills = self .sorted_spill_files .drain(..number_of_spills_to_read_for_current_phase) @@ -581,7 +645,14 @@ impl MultiLevelMergeBuilder { /// than one run is re-spilled), the shrunk run records its own smaller batch-size /// limit (tracked alongside the run in `sorted_spill_files`), so only merges that /// actually consume it pay the reduced batch size. - async fn split_spill_file_in_half(&mut self, index: usize) -> Result<()> { + /// + /// Returns whether the largest batch shrank. If `retry_unsplittable` is true, + /// restore an unchanged run for one minimum-merge admission attempt. + async fn split_spill_file_in_half( + &mut self, + index: usize, + retry_unsplittable: bool, + ) -> Result { log::debug!( "2 spilled streams could not be loaded into memory for merge \ (requires 2x of the largest batch from both), re-spilling the larger of the two with half \ @@ -598,17 +669,67 @@ impl MultiLevelMergeBuilder { // that reads this run so the merged output can't rebuild a full-size batch. let last = self.sorted_spill_files.len() - 1; self.sorted_spill_files.swap(index, last); - let (target, old_batch_size) = self + let (mut target, old_batch_size) = self .sorted_spill_files .pop() .expect("index is in bounds, so the vec is non-empty"); let old_max = target.max_record_batch_memory; + let mut max_batch_rows = old_batch_size; // Reserve enough to hold a single stream of this file while we re-spill it. let reservation = self.reservation.new_empty(); reservation .try_grow(get_reserved_bytes_for_record_batch_size(old_max, old_max))?; + if self.reserve_replay_headroom { + // A maximum-sized singleton cannot shrink. Find it without writing + // another file: the original runs may already fill the disk quota. + // Use an unbuffered reader so no background read outlives this guard. + let mut source = self.spill_manager.read_spill_as_stream_unbuffered( + Arc::clone(&target.file), + Some(old_max), + )?; + let mut all_singletons = true; + let mut max_is_singleton = false; + let mut decoded_max = 0; + max_batch_rows = 0; + while let Some(batch) = source.next().await { + let batch = batch?; + max_batch_rows = max_batch_rows.max(batch.num_rows()); + all_singletons &= batch.num_rows() == 1; + decoded_max = decoded_max.max(batch.get_sliced_size()?); + if batch.num_rows() == 1 + && gc_view_arrays(&batch)?.get_sliced_size()? >= old_max + { + max_is_singleton = true; + break; + } + } + // IPC can discard spare view-buffer capacity included in `old_max`. + // A complete scan can correct that estimate without rewriting the + // file. Use decoded buffers without GC, as the merge retains them. + let shrank = !max_is_singleton && decoded_max < old_max; + if shrank || all_singletons || max_is_singleton { + if !shrank && !retry_unsplittable { + return resources_err!( + "Cannot merge sorted runs: a single record batch of {old_max} bytes \ + exceeds the available merge memory and cannot be split further" + ); + } + if shrank { + target.max_record_batch_memory = decoded_max; + } + let batch_size_limit = if all_singletons || max_is_singleton { + 1 + } else { + old_batch_size.min(max_batch_rows).max(1) + }; + self.sorted_spill_files.push((target, batch_size_limit)); + self.sorted_spill_files.swap(index, last); + return Ok(shrank); + } + } + let source = self .spill_manager .read_spill_as_stream(target.file, Some(old_max))?; @@ -642,10 +763,11 @@ impl MultiLevelMergeBuilder { return internal_err!("re-spilling a skewed spill file produced no data"); }; - // If halving could not reduce the largest batch (e.g. a single row that is - // itself wider than the budget), there is nothing more we can do - surface - // the out-of-memory condition instead of looping forever. - if new_max >= old_max { + // If halving cannot reduce the largest batch, only a requested retry + // against the actual pool can make progress. The caller permits that + // retry once before surfacing this error. + let shrank = new_max < old_max; + if !shrank && !retry_unsplittable { return resources_err!( "Cannot merge sorted runs: a single record batch of {old_max} bytes \ exceeds the available merge memory and cannot be split further" @@ -655,8 +777,16 @@ impl MultiLevelMergeBuilder { // Record the halved batch size as a *per-run* limit rather than lowering the // global batch size. Merges that don't touch this run keep the full batch // size. a merge that reads it caps its output at this limit so the merged run - // can't rebuild a full-size batch and reintroduce the skew. - let new_batch_size_limit = (old_batch_size / 2).max(1); + // can't rebuild a full-size batch and reintroduce the skew. Actual batches + // can be shorter than the configured limit; use their size so the merge + // cannot accumulate extra split batches outside its reservation. + let new_batch_size_limit = if shrank { + (old_batch_size / 2).min(max_batch_rows.div_ceil(2)).max(1) + } else { + // Skipping headroom only covers indivisible input rows, not a + // larger output batch formed by concatenating those rows. + 1 + }; // Push the re-spilled (smaller) file and swap it back into `index`, undoing // the swap-to-back above so the order is preserved. @@ -670,7 +800,7 @@ impl MultiLevelMergeBuilder { let last = self.sorted_spill_files.len() - 1; self.sorted_spill_files.swap(index, last); - Ok(()) + Ok(shrank) } fn observe_output( @@ -900,13 +1030,13 @@ mod tests { // release the workspace and reacquire it from the parent pool. for _ in 0..expected_splits { let MergeStep::SplitThenRetry(index) = - builder.merge_sorted_runs_within_mem_limit()? + builder.merge_sorted_runs_within_mem_limit(false)? else { panic!("the merge must re-spill a skewed run"); }; assert_eq!(parent.reserved(), capacity); assert!(contender.try_grow(1).is_err()); - builder.split_spill_file_in_half(index).await?; + builder.split_spill_file_in_half(index, false).await?; assert_eq!(parent.reserved(), capacity); assert!(contender.try_grow(1).is_err()); } @@ -1002,6 +1132,350 @@ mod tests { Ok(()) } + /// A run can fit the pool during splitting but leave no room for replay. + /// Splitting must shrink the merge buffers to half the available memory. + #[rstest::rstest] + #[tokio::test] + async fn replay_headroom_splits_an_oversized_first_run( + #[values(128, 129, 4096)] rows: i64, + ) -> Result<()> { + let env = Arc::new(RuntimeEnv::default()); + let schema = test_schema(); + let spill_manager = build_spill_manager(&env, &schema); + let first = make_sorted_spill_file(&spill_manager, &schema, (0..rows).collect()); + let second = + make_sorted_spill_file(&spill_manager, &schema, (rows..2 * rows).collect()); + let merge_limit = first.max_record_batch_memory; + // One original run fills the pool. Splitting may use that space before + // replay starts, but the final merge must leave half the pool for replay. + let pool_size = 2 * merge_limit; + let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_size)); + let builder = build_merge_builder( + spill_manager, + Arc::clone(&schema), + vec![first, second], + &pool, + 8192, + ) + .with_replay_headroom(true); + let mut stream = builder.create_spillable_merge_stream(); + let mut batches = vec![]; + while let Some(batch) = stream.next().await { + let batch = batch?; + // Short input runs must not rebuild a larger output batch than + // the attached merge reservation can hold. + assert!( + crate::spill::get_record_batch_memory_size(&batch) <= pool.reserved() + ); + batches.push(batch); + assert!(pool.reserved() <= merge_limit); + } + let merged = concat_batches(&schema, &batches)?; + let values = merged.column(0).as_primitive::(); + assert_eq!(values.len(), (2 * rows) as usize); + for (expected, value) in values.values().iter().enumerate() { + assert_eq!(*value, expected as i64); + } + drop(stream); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.spilling_progress().current_bytes, 0); + assert_eq!(env.disk_manager.spilling_progress().active_files_count, 0); + Ok(()) + } + + #[tokio::test] + async fn replay_headroom_allows_only_an_indivisible_minimum() -> Result<()> { + let env = Arc::new(RuntimeEnv::default()); + let schema = test_schema(); + let spill_manager = build_spill_manager(&env, &schema); + let spills = (0..3) + .map(|value| make_sorted_spill_file(&spill_manager, &schema, vec![value])) + .collect::>(); + let batch_memory = spills[0].max_record_batch_memory; + let pool_size = 6 * batch_memory; + let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_size)); + let mut builder = + build_merge_builder(spill_manager, Arc::clone(&schema), spills, &pool, 8192) + .with_replay_headroom(true); + + assert!(!builder.split_spill_file_in_half(0, true).await?); + assert_eq!(builder.sorted_spill_files.len(), 3); + // 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)? + else { + panic!("minimum merge should fit the pool"); + }; + assert_eq!(buffer_len, 1); + assert_eq!(spills.len(), 2); + assert_eq!(builder.sorted_spill_files.len(), 1); + assert_eq!(reservation.size(), 4 * batch_memory); + builder.sorted_spill_files.splice(0..0, spills); + reservation.free(); + + let mut stream = builder.create_spillable_merge_stream(); + let mut values = Vec::new(); + while let Some(batch) = stream.next().await { + let batch = batch?; + assert_eq!(batch.num_rows(), 1); + values.push(batch.column(0).as_primitive::().value(0)); + } + assert_eq!(values, vec![0, 1, 2]); + drop(stream); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.spilling_progress().current_bytes, 0); + assert_eq!(env.disk_manager.spilling_progress().active_files_count, 0); + Ok(()) + } + + #[test] + fn replay_headroom_is_released_after_rejected_candidate() -> Result<()> { + let env = Arc::new(RuntimeEnv::default()); + let schema = test_schema(); + let spill_manager = build_spill_manager(&env, &schema); + let spills = (0..3) + .map(|value| make_sorted_spill_file(&spill_manager, &schema, vec![value])) + .collect::>(); + let batch_memory = spills[0].max_record_batch_memory; + let pool: Arc = + Arc::new(GreedyMemoryPool::new(10 * batch_memory)); + 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)? + else { + panic!("two streams and replay headroom should fit the pool"); + }; + assert_eq!(buffer_len, 1); + assert_eq!(spills.len(), 2); + assert_eq!(builder.sorted_spill_files.len(), 1); + // The third candidate failed its 12-batch reservation. Release the + // successful probe's spare four batches, retaining only two inputs. + assert_eq!(reservation.size(), 4 * batch_memory); + let replay = builder.reservation.new_empty(); + replay.try_grow(6 * batch_memory)?; + assert_eq!(pool.reserved(), 10 * batch_memory); + drop((replay, reservation, spills, builder)); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.spilling_progress().current_bytes, 0); + assert_eq!(env.disk_manager.spilling_progress().active_files_count, 0); + Ok(()) + } + + #[tokio::test] + async fn replay_headroom_still_enforces_the_actual_pool() -> Result<()> { + let env = Arc::new(RuntimeEnv::default()); + let schema = test_schema(); + let spill_manager = build_spill_manager(&env, &schema); + let first = make_sorted_spill_file(&spill_manager, &schema, vec![1]); + let second = make_sorted_spill_file(&spill_manager, &schema, vec![2]); + let batch_memory = first.max_record_batch_memory; + let pool: Arc = Arc::new(GreedyMemoryPool::new(3 * batch_memory)); + let builder = + build_merge_builder(spill_manager, schema, vec![first, second], &pool, 8192) + .with_replay_headroom(true); + let mut stream = builder.create_spillable_merge_stream(); + let error = stream.next().await.unwrap().unwrap_err(); + assert!(error.to_string().contains("cannot be split further")); + drop(stream); + assert_eq!(pool.reserved(), 0); + assert_eq!(env.disk_manager.spilling_progress().current_bytes, 0); + assert_eq!(env.disk_manager.spilling_progress().active_files_count, 0); + Ok(()) + } + + #[rstest::rstest] + #[case(1, false, DataType::Utf8)] + #[case(8192, true, DataType::Utf8)] + #[case(8192, true, DataType::Utf8View)] + #[tokio::test] + async fn indivisible_spill_merge_needs_no_extra_disk_space( + #[case] batch_size: usize, + #[case] mixed_batches: bool, + #[case] data_type: DataType, + ) -> Result<()> { + use arrow::array::{ArrayRef, StringArray, StringViewArray}; + + const KEY_BYTES: usize = 300_000; + const POOL_BYTES: usize = 2 * 1024 * 1024; + let schema = Arc::new(Schema::new(vec![Field::new("x", data_type, false)])); + let make_runs = |manager: &SpillManager| -> Result> { + (0..2) + .map(|run| { + let batches = (0..8).map(|batch| { + // A small two-row batch before the large singleton + // prevents treating the first batch as representative. + let keys = if mixed_batches && batch == 0 { + vec![format!("{run:04}"), format!("{:04}", run + 2)] + } else { + let key = run + 2 * (batch + usize::from(mixed_batches)); + vec![format!("{key:04}{}", "x".repeat(KEY_BYTES - 4))] + }; + let values: ArrayRef = match schema.field(0).data_type() { + DataType::Utf8 => Arc::new(StringArray::from(keys)), + DataType::Utf8View => Arc::new(StringViewArray::from(keys)), + _ => unreachable!(), + }; + RecordBatch::try_new(Arc::clone(&schema), vec![values]) + .map_err(Into::into) + }); + let (file, max_record_batch_memory) = manager + .spill_record_batch_iter_and_return_max_batch_memory( + batches, + "indivisible input run", + )? + .unwrap(); + Ok(SortedSpillFile { + file, + max_record_batch_memory, + }) + }) + .collect() + }; + + // Calibrate the quota to exactly the original IPC files, with no room + // for even a replacement header. Keep several batches in each run so + // read-ahead cannot retire the original before the unnecessary write. + let calibration = Arc::new(RuntimeEnv::default()); + let runs = make_runs(&build_spill_manager(&calibration, &schema))?; + let quota = runs.iter().map(|run| run.file.size().unwrap()).sum(); + drop(runs); + let env = RuntimeEnvBuilder::new() + .with_max_temp_directory_size(quota) + .build_arc()?; + let manager = build_spill_manager(&env, &schema); + let runs = make_runs(&manager)?; + assert_eq!(env.disk_manager.spilling_progress().current_bytes, quota); + assert!(4 * runs[0].max_record_batch_memory > POOL_BYTES / 2); + assert!(4 * runs[0].max_record_batch_memory <= POOL_BYTES); + let pool: Arc = Arc::new(GreedyMemoryPool::new(POOL_BYTES)); + let builder = build_merge_builder(manager, schema, runs, &pool, batch_size) + .with_replay_headroom(true); + let mut stream = builder.create_spillable_merge_stream(); + let mut expected = 0; + while let Some(batch) = stream.next().await { + let batch = batch?; + assert_eq!(batch.num_rows(), 1); + let values = arrow::compute::cast(batch.column(0), &DataType::Utf8)?; + let key = values.as_string::().value(0); + assert_eq!(key[..4].parse::().unwrap(), expected); + expected += 1; + } + assert_eq!(expected, if mixed_batches { 18 } else { 16 }); + drop(stream); + assert_eq!(pool.reserved(), 0); + let progress = env.disk_manager.spilling_progress(); + assert_eq!(progress.current_bytes, 0); + assert_eq!(progress.active_files_count, 0); + Ok(()) + } + + #[rstest::rstest] + #[tokio::test] + async fn small_view_spill_merge_needs_no_extra_disk_space( + #[values(3, 6)] pool_batches: usize, + #[values(false, true)] mixed_batches: bool, + ) -> Result<()> { + use arrow::array::StringViewBuilder; + + let schema = Arc::new(Schema::new(vec![Field::new( + "x", + DataType::Utf8View, + false, + )])); + let make_runs = |manager: &SpillManager| -> Result> { + (0..2) + .map(|run| { + let batches = (0..8).map(|batch| { + // A non-inline value retains the builder's 8 KiB block, + // but IPC stores only its used bytes. Small view buffers + // are below the spill writer's compaction threshold. + let key = format!("{:04}xxxxxxxxx", run + 2 * batch); + let mut values = StringViewBuilder::with_capacity(1) + .with_fixed_block_size(8192); + values.append_value(key); + // A larger final batch requires scanning the entire run + // before replacing its maximum decoded size. + if mixed_batches && batch == 7 { + values.append_value(format!( + "{:04}xxxxxxxxx", + run + 2 * (batch + 1) + )); + } + RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(values.finish())], + ) + .map_err(Into::into) + }); + let (file, max_record_batch_memory) = manager + .spill_record_batch_iter_and_return_max_batch_memory( + batches, + "small view singleton run", + )? + .unwrap(); + Ok(SortedSpillFile { + file, + max_record_batch_memory, + }) + }) + .collect() + }; + + // Fill the quota with the original runs. Multiple batches prevent + // read-ahead from retiring a run before an unnecessary replacement write. + let calibration = Arc::new(RuntimeEnv::default()); + let runs = make_runs(&build_spill_manager(&calibration, &schema))?; + let quota = runs.iter().map(|run| run.file.size().unwrap()).sum(); + drop(runs); + let env = RuntimeEnvBuilder::new() + .with_max_temp_directory_size(quota) + .build_arc()?; + let manager = build_spill_manager(&env, &schema); + let runs = make_runs(&manager)?; + assert_eq!(env.disk_manager.spilling_progress().current_bytes, quota); + let stored_max = runs[0].max_record_batch_memory; + let mut source = manager.read_spill_as_stream_unbuffered( + Arc::clone(&runs[0].file), + Some(stored_max), + )?; + let batch = source.next().await.unwrap()?; + assert_eq!(batch.num_rows(), 1); + assert!(gc_view_arrays(&batch)?.get_sliced_size()? < stored_max); + drop((source, batch)); + + // The stale estimate needs four batches for the minimum merge. Test + // limits below and above that estimate: both can fit the decoded rows. + let pool_bytes = pool_batches * stored_max; + let pool: Arc = Arc::new(GreedyMemoryPool::new(pool_bytes)); + let builder = build_merge_builder(manager, schema, runs, &pool, 8192) + .with_replay_headroom(true); + let mut stream = builder.create_spillable_merge_stream(); + let mut expected = 0; + while let Some(batch) = stream.next().await { + let batch = batch?; + if mixed_batches { + assert!(batch.num_rows() <= 2); + } else { + assert_eq!(batch.num_rows(), 1); + } + for key in batch.column(0).as_string_view().iter() { + assert_eq!(key.unwrap(), format!("{expected:04}xxxxxxxxx")); + expected += 1; + } + } + assert_eq!(expected, if mixed_batches { 18 } else { 16 }); + drop(stream); + assert_eq!(pool.reserved(), 0); + let progress = env.disk_manager.spilling_progress(); + assert_eq!(progress.current_bytes, 0); + assert_eq!(progress.active_files_count, 0); + Ok(()) + } + /// Tests the `new_max >= old_max` guard: a single-row run cannot be split /// any smaller, so re-spilling it does not shrink the largest batch and the /// rescue surfaces `ResourcesExhausted` rather than looping forever. @@ -1022,7 +1496,7 @@ mod tests { build_merge_builder(spill_manager, schema, vec![f0], &pool, 1024); let err = builder - .split_spill_file_in_half(0) + .split_spill_file_in_half(0, false) .await .expect_err("re-spilling a one-row run cannot shrink it"); assert!( @@ -1195,6 +1669,7 @@ mod tests { 1, 2, &mut merge_reservation, + false, )? { SpillFilesToMerge::Ready(spills, buffer_len) => (spills, buffer_len), SpillFilesToMerge::SplitThenRetry(index) => { diff --git a/datafusion/physical-plan/src/sorts/streaming_merge.rs b/datafusion/physical-plan/src/sorts/streaming_merge.rs index 726ae48b1c79c..61bb7742eecfa 100644 --- a/datafusion/physical-plan/src/sorts/streaming_merge.rs +++ b/datafusion/physical-plan/src/sorts/streaming_merge.rs @@ -96,6 +96,8 @@ pub struct StreamingMergeBuilder<'a> { fetch: Option, reservation: Option, merge_pool: Option>, + /// Leave memory for the aggregate consuming the merged spill rows. + reserve_replay_headroom: bool, enable_round_robin_tie_breaker: bool, } @@ -161,6 +163,13 @@ impl<'a> StreamingMergeBuilder<'a> { self } + /// Leave room for aggregate replay by checking that the pool can admit + /// merge buffers plus equal headroom. Release the headroom before replay. + pub(crate) fn with_replay_headroom(mut self) -> Self { + self.reserve_replay_headroom = true; + self + } + /// See [SortPreservingMergeExec::with_round_robin_repartition] for more /// information. /// @@ -194,6 +203,7 @@ impl<'a> StreamingMergeBuilder<'a> { batch_size, reservation, merge_pool, + reserve_replay_headroom, fetch, expressions, enable_round_robin_tie_breaker, @@ -235,6 +245,7 @@ impl<'a> StreamingMergeBuilder<'a> { enable_round_robin_tie_breaker, ) .with_merge_pool(merge_pool) + .with_replay_headroom(reserve_replay_headroom) .create_spillable_merge_stream()); }