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 89763efc4d75c..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; @@ -146,6 +146,188 @@ 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) -> 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(()) + } + + /// 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); + } + + // 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; + } + } + + 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) { + // 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 +365,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) } @@ -281,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 @@ -340,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; @@ -457,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 d08bd55de91ca..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,6 +87,9 @@ pub(crate) struct SortPreservingMergeStream { /// Target batch size batch_size: usize, + /// 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>>, @@ -146,11 +150,20 @@ impl SortPreservingMergeStream { poll_reset_epochs: vec![0; stream_count], loser_tree: vec![], batch_size, + batch_memory_budget: None, fetch, produced: 0, } } + pub(super) fn with_batch_memory_budget( + mut self, + budget: Option, + ) -> Self { + self.batch_memory_budget = budget; + self + } + pub(crate) fn into_stream(self) -> SendableRecordBatchStream where C: 'static, @@ -198,7 +211,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 +297,20 @@ impl SortPreservingMergeStream { ); drop(timer); + 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(); } @@ -294,7 +321,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 c1fe893e0df57..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}; @@ -159,6 +161,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 +203,7 @@ impl MultiLevelMergeBuilder { reservation, merge_pool: None, reserve_replay_headroom: false, + widen_intermediate_merges: true, enable_round_robin_tie_breaker, fetch, } @@ -209,8 +214,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 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 @@ -226,13 +231,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 +280,55 @@ 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), 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, + reservation, + } = retry; + self.sorted_spill_files + .splice(0..0, spills.drain(original_count..)); + // 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) + .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 + .spill_record_batch_stream_and_return_max_batch_memory( + &mut stream, + "MultiLevelMergeBuilder intermediate spill retry", + ) + .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; }; @@ -314,6 +360,7 @@ impl MultiLevelMergeBuilder { Ok(MergeStep::Stream { stream: self.observe_output(empty_stream), batch_size_limit: self.batch_size, + retry: None, }) } @@ -323,6 +370,7 @@ impl MultiLevelMergeBuilder { Ok(MergeStep::Stream { stream: self.observe_output(output_stream), batch_size_limit: self.batch_size, + retry: None, }) } @@ -338,6 +386,7 @@ impl MultiLevelMergeBuilder { Ok(MergeStep::Stream { stream: self.observe_output(output_stream), batch_size_limit: batch_size, + retry: None, }) } @@ -354,8 +403,10 @@ impl MultiLevelMergeBuilder { true, true, self.batch_size, + None, )?, batch_size_limit: self.batch_size, + retry: None, }) } @@ -373,95 +424,163 @@ 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( - 2, - // 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) + let selection = self.get_sorted_spill_files_to_merge( + 2, + minimum_number_of_required_streams, + &mut memory_reservation, + allow_minimum_without_headroom, + )?; + let (mut spills, mut 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_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() + { + // 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() + .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), + ); + } } + let reservation = Arc::new(memory_reservation); + 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, + original_count, + buffer_size: original_buffer_size, + reservation, + }); + 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. + /// `bound_batch_memory` requires spill-only inputs and a one-batch read buffer. + fn merge_selected_runs( + &mut self, + sorted_spill_files: &[(SortedSpillFile, usize)], + buffer_size: usize, + memory_reservation: Arc, + 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(); + + // 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, + None, + )?, + 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( + 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, + batch_memory_budget, + )?; + + Ok(( + Box::pin(StreamAttachedReservation::new( + merge_sort_stream, + memory_reservation, + )), + output_batch_size, + )) } fn create_new_merge_sort( @@ -470,11 +589,13 @@ impl MultiLevelMergeBuilder { is_output: bool, all_in_memory: bool, output_batch_size: usize, + 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_batch_memory_budget(batch_memory_budget) .with_fetch(self.fetch) .with_metrics(if is_output { // Only add the metrics to the last run @@ -515,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. @@ -580,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. @@ -620,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. @@ -634,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 @@ -821,6 +946,15 @@ 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, + 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). @@ -832,6 +966,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), @@ -855,16 +990,34 @@ 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: 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), } } } @@ -884,12 +1037,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) } @@ -906,6 +1059,9 @@ impl RecordBatchStream for StreamAttachedReservation { } } +#[cfg(test)] +mod replay_headroom_tests; + #[cfg(test)] mod tests { use super::*; @@ -923,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), @@ -937,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, @@ -961,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 new file mode 100644 index 0000000000000..e86206c4b797a --- /dev/null +++ b/datafusion/physical-plan/src/sorts/multi_level_merge/replay_headroom_tests.rs @@ -0,0 +1,541 @@ +// 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 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_common::metrics::{ExecutionPlanMetricsSet, SpillMetrics}; +use std::sync::atomic::{AtomicBool, Ordering}; + +struct ReplayMergeFixture { + builder: MultiLevelMergeBuilder, + env: Arc, + pool: Arc, + pool_size: usize, + metrics: SpillMetrics, + input_bytes: usize, +} + +fn replay_merge_builder( + spill_manager: SpillManager, + schema: SchemaRef, + spills: Vec, + pool: &Arc, + batch_size: usize, +) -> MultiLevelMergeBuilder { + build_merge_builder(spill_manager, schema, spills, pool, batch_size) + .with_replay_headroom(true) +} + +/// 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 = 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)); + let builder = + replay_merge_builder(spill_manager, schema, spills, &pool, rows_per_run); + 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_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)] +#[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(()) +} + +#[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 { + 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(()) +} + +#[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_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 manager = build_spill_manager(&env, &schema); + let values = (0..run_count * batches_per_run) + .map(|value| format!("{value:04}{}", "x".repeat(1024))) + .collect::>(); + 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(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? { + 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); + } + 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 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. + // 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?; + assert_eq!( + spill_count, 4, + "intermediate spill: {spilled_rows} rows, {spilled_bytes} 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(()) +} + +#[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"); + let message = error.to_string(); + assert!( + 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 + // 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 minimum 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..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) => {{ + ($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,6 +64,7 @@ macro_rules! merge_helper { $reservation, $enable_round_robin_tie_breaker, ) + .with_batch_memory_budget($batch_memory_budget) .into_stream()); }}; } @@ -98,6 +108,7 @@ pub struct StreamingMergeBuilder<'a> { merge_pool: Option>, /// Leave memory for the aggregate consuming the merged spill rows. reserve_replay_headroom: bool, + batch_memory_budget: Option, enable_round_robin_tie_breaker: bool, } @@ -170,6 +181,15 @@ impl<'a> StreamingMergeBuilder<'a> { self } + /// 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 + } + /// See [SortPreservingMergeExec::with_round_robin_repartition] for more /// information. /// @@ -204,6 +224,7 @@ impl<'a> StreamingMergeBuilder<'a> { reservation, merge_pool, reserve_replay_headroom, + batch_memory_budget, fetch, expressions, enable_round_robin_tie_breaker, @@ -259,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), - 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, 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) _ => {} } } @@ -290,18 +317,22 @@ impl<'a> StreamingMergeBuilder<'a> { reservation, enable_round_robin_tie_breaker, ) + .with_batch_memory_budget(batch_memory_budget) .into_stream()) } } #[cfg(test)] mod tests { + use crate::spill::{get_record_batch_memory_size, 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, Int64Type, Schema}; use arrow_schema::SortOptions; use datafusion_common::Result; use datafusion_execution::TaskContext; @@ -309,6 +340,336 @@ 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_bounds_short_batches( + #[case] data_type: DataType, + #[values(false, true)] row_cursor: 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 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]; + 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!(), + }; + 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) + .with_metrics(BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0)) + .with_streams(streams) + .with_batch_size(8192) + .with_fetch(fetch) + .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?; + 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 budgeted { + for batch in &batches { + assert!(get_record_batch_memory_size(batch) <= max_output_bytes); + } + } else { + assert_eq!(batches.len(), 1, "ordinary merges retain their batching"); + } + } + 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![ + 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 max_input_memory = 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()?); + max_input_memory = + max_input_memory.max(get_record_batch_memory_size(&batch)); + 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_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?; + 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() {