use crate::metrics::BaselineMetrics;
use crate::{EmptyRecordBatchStream, SpillManager};
use arrow::array::RecordBatch;
use std::fmt::{Debug, Formatter};
use std::mem;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use arrow::datatypes::SchemaRef;
use datafusion_common::{Result, internal_err, resources_err};
use datafusion_execution::memory_pool::MemoryReservation;
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::stream::{ObservedStream, RecordBatchStreamAdapter};
use datafusion_execution::{RecordBatchStream, SendableRecordBatchStream};
use datafusion_physical_expr_common::sort_expr::LexOrdering;
use futures::TryStreamExt;
use futures::{Stream, StreamExt};
pub(crate) struct MultiLevelMergeBuilder {
spill_manager: SpillManager,
schema: SchemaRef,
sorted_spill_files: Vec<(SortedSpillFile, usize)>,
sorted_streams: Vec<SendableRecordBatchStream>,
expr: LexOrdering,
metrics: BaselineMetrics,
batch_size: usize,
reservation: MemoryReservation,
fetch: Option<usize>,
enable_round_robin_tie_breaker: bool,
}
impl Debug for MultiLevelMergeBuilder {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "MultiLevelMergeBuilder")
}
}
impl MultiLevelMergeBuilder {
#[expect(clippy::too_many_arguments)]
pub(crate) fn new(
spill_manager: SpillManager,
schema: SchemaRef,
sorted_spill_files: Vec<SortedSpillFile>,
sorted_streams: Vec<SendableRecordBatchStream>,
expr: LexOrdering,
metrics: BaselineMetrics,
batch_size: usize,
reservation: MemoryReservation,
fetch: Option<usize>,
enable_round_robin_tie_breaker: bool,
) -> Self {
Self {
spill_manager,
schema,
sorted_spill_files: sorted_spill_files
.into_iter()
.map(|file| (file, batch_size))
.collect(),
sorted_streams,
expr,
metrics,
batch_size,
reservation,
enable_round_robin_tie_breaker,
fetch,
}
}
pub(crate) fn create_spillable_merge_stream(self) -> SendableRecordBatchStream {
Box::pin(RecordBatchStreamAdapter::new(
Arc::clone(&self.schema),
futures::stream::once(self.create_stream()).try_flatten(),
))
}
async fn create_stream(mut self) -> Result<SendableRecordBatchStream> {
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) => {
self.split_spill_file_in_half(index).await?;
continue;
}
};
if self.sorted_spill_files.is_empty() {
assert!(
self.sorted_streams.is_empty(),
"We should not have any sorted streams left"
);
return Ok(stream);
}
let Some((spill_file, max_record_batch_memory)) = self
.spill_manager
.spill_record_batch_stream_and_return_max_batch_memory(
&mut stream,
"MultiLevelMergeBuilder intermediate spill",
)
.await?
else {
continue;
};
self.sorted_spill_files.push((
SortedSpillFile {
file: spill_file,
max_record_batch_memory,
},
batch_size_limit,
));
}
}
fn merge_sorted_runs_within_mem_limit(&mut self) -> Result<MergeStep> {
match (self.sorted_spill_files.len(), self.sorted_streams.len()) {
(0, 0) => {
let empty_stream =
Box::pin(EmptyRecordBatchStream::new(Arc::clone(&self.schema)));
Ok(MergeStep::Stream {
stream: self.observe_output(empty_stream),
batch_size_limit: self.batch_size,
})
}
(0, 1) => {
let output_stream = self.sorted_streams.remove(0);
Ok(MergeStep::Stream {
stream: self.observe_output(output_stream),
batch_size_limit: self.batch_size,
})
}
(1, 0) => {
let (spill_file, batch_size) = self.sorted_spill_files.remove(0);
let output_stream = self
.spill_manager
.read_spill_as_stream(spill_file.file, None)?;
Ok(MergeStep::Stream {
stream: self.observe_output(output_stream),
batch_size_limit: batch_size,
})
}
(0, _) => {
let sorted_stream = mem::take(&mut self.sorted_streams);
Ok(MergeStep::Stream {
stream: self.create_new_merge_sort(
sorted_stream,
true,
true,
self.batch_size,
)?,
batch_size_limit: self.batch_size,
})
}
(_, _) => {
let mut memory_reservation = self.reservation.take();
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,
minimum_number_of_required_streams,
&mut memory_reservation,
)? {
SpillFilesToMerge::Ready(sorted_spill_files, buffer_size) => {
(sorted_spill_files, buffer_size)
}
SpillFilesToMerge::SplitThenRetry(index) => {
return Ok(MergeStep::SplitThenRetry(index));
}
};
let mut sorted_streams = mem::take(&mut self.sorted_streams);
let is_only_merging_memory_streams = sorted_spill_files.is_empty();
if is_only_merging_memory_streams {
mem::swap(&mut self.reservation, &mut memory_reservation);
}
let mut output_batch_size = self.batch_size;
for (spill, batch_size_limit) in sorted_spill_files {
let stream = self
.spill_manager
.clone()
.with_batch_read_buffer_capacity(buffer_size)
.read_spill_as_stream(
spill.file,
Some(spill.max_record_batch_memory),
)?;
output_batch_size = output_batch_size.min(batch_size_limit);
sorted_streams.push(stream);
}
let merge_sort_stream = self.create_new_merge_sort(
sorted_streams,
self.sorted_spill_files.is_empty(),
is_only_merging_memory_streams,
output_batch_size,
)?;
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"
);
Ok(MergeStep::Stream {
stream: merge_sort_stream,
batch_size_limit: output_batch_size,
})
} else {
Ok(MergeStep::Stream {
stream: Box::pin(StreamAttachedReservation::new(
merge_sort_stream,
memory_reservation,
)),
batch_size_limit: output_batch_size,
})
}
}
}
}
fn create_new_merge_sort(
&mut self,
streams: Vec<SendableRecordBatchStream>,
is_output: bool,
all_in_memory: bool,
output_batch_size: usize,
) -> Result<SendableRecordBatchStream> {
let mut builder = StreamingMergeBuilder::new()
.with_schema(Arc::clone(&self.schema))
.with_expressions(&self.expr)
.with_batch_size(output_batch_size)
.with_fetch(self.fetch)
.with_metrics(if is_output {
self.metrics.clone()
} else {
self.metrics.intermediate()
})
.with_round_robin_tie_breaker(self.enable_round_robin_tie_breaker)
.with_streams(streams);
if !all_in_memory {
builder = builder.with_bypass_mempool();
} else {
builder = builder.with_reservation(self.reservation.take());
}
builder.build()
}
fn get_sorted_spill_files_to_merge(
&mut self,
buffer_len: usize,
minimum_number_of_required_streams: usize,
reservation: &mut MemoryReservation,
) -> Result<SpillFilesToMerge> {
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);
let mut total_needed: usize = 0;
for (spill, _) in &self.sorted_spill_files {
if number_of_spills_to_read_for_current_phase >= max_spill_files {
break;
}
let per_spill = get_reserved_bytes_for_record_batch_size(
spill.max_record_batch_memory,
spill.max_record_batch_memory,
) * buffer_len;
total_needed += per_spill;
match try_grow_reservation_to_at_least(reservation, total_needed) {
Ok(_) => {
number_of_spills_to_read_for_current_phase += 1;
}
Err(err) => {
if minimum_number_of_required_streams
> number_of_spills_to_read_for_current_phase
{
reservation.free();
if buffer_len > 1 {
return self.get_sorted_spill_files_to_merge(
buffer_len - 1,
minimum_number_of_required_streams,
reservation,
);
}
if number_of_spills_to_read_for_current_phase == 0 {
return Err(err);
}
let split_index = usize::from(
self.sorted_spill_files[1].0.max_record_batch_memory
> self.sorted_spill_files[0].0.max_record_batch_memory,
);
return Ok(SpillFilesToMerge::SplitThenRetry(split_index));
}
break;
}
}
}
let spills = self
.sorted_spill_files
.drain(..number_of_spills_to_read_for_current_phase)
.collect::<Vec<_>>();
Ok(SpillFilesToMerge::Ready(spills, buffer_len))
}
async fn split_spill_file_in_half(&mut self, index: usize) -> 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 \
the batch size to reduce memory needs for the next merge attempt. the shrunk run carries \
a halved batch-size limit so only merges consuming it use the smaller batch size"
);
let last = self.sorted_spill_files.len() - 1;
self.sorted_spill_files.swap(index, last);
let (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 reservation = self.reservation.new_empty();
reservation
.try_grow(get_reserved_bytes_for_record_batch_size(old_max, old_max))?;
let source = self
.spill_manager
.read_spill_as_stream(target.file, Some(old_max))?;
let mut halved: SendableRecordBatchStream =
Box::pin(RecordBatchStreamAdapter::new(
Arc::clone(&self.schema),
source.flat_map(|batch| {
futures::stream::iter(match batch {
Ok(batch) => split_batch_in_half(batch)
.into_iter()
.map(Ok)
.collect::<Vec<_>>(),
Err(e) => vec![Err(e)],
})
}),
));
let result = self
.spill_manager
.spill_record_batch_stream_and_return_max_batch_memory(
&mut halved,
"MultiLevelMergeBuilder split skewed spill",
)
.await?;
reservation.free();
let Some((file, new_max)) = result else {
return internal_err!("re-spilling a skewed spill file produced no data");
};
if new_max >= old_max {
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"
);
}
let new_batch_size_limit = (old_batch_size / 2).max(1);
self.sorted_spill_files.push((
SortedSpillFile {
file,
max_record_batch_memory: new_max,
},
new_batch_size_limit,
));
let last = self.sorted_spill_files.len() - 1;
self.sorted_spill_files.swap(index, last);
Ok(())
}
fn observe_output(
&self,
stream: SendableRecordBatchStream,
) -> SendableRecordBatchStream {
Box::pin(ObservedStream::new(stream, self.metrics.clone(), None))
}
}
enum SpillFilesToMerge {
Ready(Vec<(SortedSpillFile, usize)>, usize),
SplitThenRetry(usize),
}
enum MergeStep {
Stream {
stream: SendableRecordBatchStream,
batch_size_limit: usize,
},
SplitThenRetry(usize),
}
fn split_batch_in_half(batch: RecordBatch) -> Vec<RecordBatch> {
let num_rows = batch.num_rows();
if num_rows <= 1 {
return vec![batch];
}
let mid = num_rows / 2;
vec![batch.slice(0, mid), batch.slice(mid, num_rows - mid)]
}
fn effective_spill_merge_fan_in(configured_fan_in: usize) -> usize {
if configured_fan_in == 0 {
usize::MAX
} else {
configured_fan_in.max(2)
}
}
struct StreamAttachedReservation {
stream: SendableRecordBatchStream,
reservation: MemoryReservation,
}
impl StreamAttachedReservation {
fn new(stream: SendableRecordBatchStream, reservation: MemoryReservation) -> Self {
Self {
stream,
reservation,
}
}
}
impl Stream for StreamAttachedReservation {
type Item = Result<RecordBatch>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
let res = self.stream.poll_next_unpin(cx);
match res {
Poll::Ready(res) => {
match res {
Some(Ok(batch)) => Poll::Ready(Some(Ok(batch))),
Some(Err(err)) => {
self.reservation.free();
Poll::Ready(Some(Err(err)))
}
None => {
self.reservation.free();
Poll::Ready(None)
}
}
}
Poll::Pending => Poll::Pending,
}
}
}
impl RecordBatchStream for StreamAttachedReservation {
fn schema(&self) -> SchemaRef {
self.stream.schema()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::expressions::PhysicalSortExpr;
use arrow::array::{AsArray, Int64Array};
use arrow::compute::concat_batches;
use arrow::datatypes::{DataType, Field, Int64Type, Schema};
use datafusion_execution::memory_pool::{
GreedyMemoryPool, MemoryConsumer, MemoryPool,
};
use datafusion_execution::runtime_env::{RuntimeEnv, RuntimeEnvBuilder};
use datafusion_physical_expr::expressions::{Column, col};
use datafusion_physical_expr_common::metrics::{
ExecutionPlanMetricsSet, SpillMetrics,
};
fn test_schema() -> SchemaRef {
Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)]))
}
fn build_spill_manager(env: &Arc<RuntimeEnv>, schema: &SchemaRef) -> SpillManager {
SpillManager::new(
Arc::clone(env),
SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0),
Arc::clone(schema),
)
}
fn make_sorted_spill_file(
spill_manager: &SpillManager,
schema: &SchemaRef,
values: Vec<i64>,
) -> SortedSpillFile {
let batch = RecordBatch::try_new(
Arc::clone(schema),
vec![Arc::new(Int64Array::from(values))],
)
.unwrap();
let batches: Vec<Result<RecordBatch>> = vec![Ok(batch)];
let (file, max_record_batch_memory) = spill_manager
.spill_record_batch_iter_and_return_max_batch_memory(
batches.into_iter(),
"test input run",
)
.unwrap()
.expect("spill should produce a file");
SortedSpillFile {
file,
max_record_batch_memory,
}
}
fn build_merge_builder(
spill_manager: SpillManager,
schema: SchemaRef,
sorted_spill_files: Vec<SortedSpillFile>,
pool: &Arc<dyn MemoryPool>,
batch_size: usize,
) -> MultiLevelMergeBuilder {
let reservation = MemoryConsumer::new("test merge").register(pool);
let expr: LexOrdering =
[PhysicalSortExpr::new_default(Arc::new(Column::new("x", 0)))].into();
MultiLevelMergeBuilder::new(
spill_manager,
schema,
sorted_spill_files,
vec![],
expr,
BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0),
batch_size,
reservation,
None,
false,
)
}
#[tokio::test]
async fn skewed_runs_are_respilled_so_the_merge_fits() -> Result<()> {
let env = Arc::new(RuntimeEnv::default());
let schema = test_schema();
let spill_manager = build_spill_manager(&env, &schema);
let n: i64 = 16384;
let f0 = make_sorted_spill_file(&spill_manager, &schema, (0..n).collect());
let f1 = make_sorted_spill_file(&spill_manager, &schema, (0..n).collect());
let m = f0.max_record_batch_memory.max(f1.max_record_batch_memory);
let pool: Arc<dyn MemoryPool> = Arc::new(GreedyMemoryPool::new(m * 7 / 2));
let builder = build_merge_builder(
spill_manager,
Arc::clone(&schema),
vec![f0, f1],
&pool,
8192,
);
let stream = builder.create_spillable_merge_stream();
let batches: Vec<RecordBatch> = stream.try_collect().await?;
let total_rows: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(
total_rows,
(2 * n) as usize,
"the merge must emit every input row"
);
let merged = concat_batches(&schema, &batches)?;
let col = merged.column(0).as_primitive::<Int64Type>();
for i in 1..col.len() {
assert!(
col.value(i - 1) <= col.value(i),
"merge output must be sorted: {} > {} at {i}",
col.value(i - 1),
col.value(i),
);
}
Ok(())
}
#[tokio::test]
async fn respilling_an_unsplittable_run_surfaces_resources_exhausted() -> Result<()> {
let env = Arc::new(RuntimeEnv::default());
let schema = test_schema();
let spill_manager = build_spill_manager(&env, &schema);
let f0 = make_sorted_spill_file(&spill_manager, &schema, vec![42]);
let pool: Arc<dyn MemoryPool> = Arc::new(GreedyMemoryPool::new(1024 * 1024));
let mut builder =
build_merge_builder(spill_manager, schema, vec![f0], &pool, 1024);
let err = builder
.split_spill_file_in_half(0)
.await
.expect_err("re-spilling a one-row run cannot shrink it");
assert!(
err.to_string().contains("cannot be split further"),
"expected the un-splittable guard error, got: {err}"
);
Ok(())
}
#[tokio::test]
async fn respill_halves_the_merge_output_batch_size() -> Result<()> {
let env = Arc::new(RuntimeEnv::default());
let schema = test_schema();
let spill_manager = build_spill_manager(&env, &schema);
let n: i64 = 16384;
let f0 = make_sorted_spill_file(&spill_manager, &schema, (0..n).collect());
let f1 = make_sorted_spill_file(&spill_manager, &schema, (0..n).collect());
let m = f0.max_record_batch_memory.max(f1.max_record_batch_memory);
let initial_batch_size = 8192;
let pool: Arc<dyn MemoryPool> = Arc::new(GreedyMemoryPool::new(m * 7 / 2));
let builder = build_merge_builder(
spill_manager,
Arc::clone(&schema),
vec![f0, f1],
&pool,
initial_batch_size,
);
let stream = builder.create_spillable_merge_stream();
let batches: Vec<RecordBatch> = stream.try_collect().await?;
let total_rows: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(total_rows, (2 * n) as usize);
let expected_batch_size = initial_batch_size / 2;
let max_batch_rows = batches.iter().map(|b| b.num_rows()).max().unwrap_or(0);
assert_eq!(
max_batch_rows, expected_batch_size,
"after one re-spill the merge must emit {expected_batch_size}-row \
batches, got a largest batch of {max_batch_rows} rows"
);
Ok(())
}
#[tokio::test]
async fn respilling_two_skewed_runs_halves_the_output_without_compounding()
-> Result<()> {
let env = Arc::new(RuntimeEnv::default());
let schema = test_schema();
let spill_manager = build_spill_manager(&env, &schema);
let n: i64 = 16384;
let f0 = make_sorted_spill_file(&spill_manager, &schema, (0..n).collect());
let f1 = make_sorted_spill_file(&spill_manager, &schema, (0..n).collect());
let m = f0.max_record_batch_memory.max(f1.max_record_batch_memory);
let initial_batch_size = 8192;
let pool: Arc<dyn MemoryPool> = Arc::new(GreedyMemoryPool::new(m * 5 / 2));
let builder = build_merge_builder(
spill_manager,
Arc::clone(&schema),
vec![f0, f1],
&pool,
initial_batch_size,
);
let stream = builder.create_spillable_merge_stream();
let batches: Vec<RecordBatch> = stream.try_collect().await?;
let total_rows: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(total_rows, (2 * n) as usize);
let expected_batch_size = initial_batch_size / 2;
let max_batch_rows = batches.iter().map(|b| b.num_rows()).max().unwrap_or(0);
assert_eq!(
max_batch_rows, expected_batch_size,
"two re-spills must halve (not quarter) the output: expected \
{expected_batch_size}-row batches, got a largest batch of \
{max_batch_rows} rows"
);
Ok(())
}
#[test]
fn spill_merge_fan_in_is_unlimited_by_default() {
assert_eq!(effective_spill_merge_fan_in(0), usize::MAX);
}
#[test]
fn spill_merge_fan_in_preserves_merge_progress() {
assert_eq!(effective_spill_merge_fan_in(1), 2);
assert_eq!(effective_spill_merge_fan_in(2), 2);
assert_eq!(effective_spill_merge_fan_in(8), 8);
}
#[test]
fn spill_merge_phase_respects_configured_fan_in() -> Result<()> {
let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)]));
let runtime = RuntimeEnvBuilder::new()
.with_max_spill_merge_fan_in(2)
.build_arc()?;
let spill_manager = SpillManager::new(
Arc::clone(&runtime),
SpillMetrics::new(&ExecutionPlanMetricsSet::new(), 0),
Arc::clone(&schema),
);
let sorted_spill_files = (0..4)
.map(|idx| {
Ok(SortedSpillFile {
file: runtime
.disk_manager
.create_tmp_file(&format!("spill fan-in test {idx}"))?,
max_record_batch_memory: 1,
})
})
.collect::<Result<Vec<_>>>()?;
let expr = LexOrdering::new([PhysicalSortExpr::new_default(col("a", &schema)?)])
.unwrap();
let reservation =
MemoryConsumer::new("spill_merge_phase_respects_configured_fan_in")
.register(&runtime.memory_pool);
let metrics = BaselineMetrics::new(&ExecutionPlanMetricsSet::new(), 0);
let mut builder = MultiLevelMergeBuilder::new(
spill_manager,
schema,
sorted_spill_files,
vec![],
expr,
metrics,
1024,
reservation,
None,
false,
);
let mut merge_reservation = MemoryConsumer::new("spill_merge_fan_in_phase")
.register(&runtime.memory_pool);
let (spills, buffer_len) = match builder.get_sorted_spill_files_to_merge(
1,
2,
&mut merge_reservation,
)? {
SpillFilesToMerge::Ready(spills, buffer_len) => (spills, buffer_len),
SpillFilesToMerge::SplitThenRetry(index) => {
panic!("expected ready spill files, got retry for index {index}")
}
};
assert_eq!(spills.len(), 2);
assert_eq!(buffer_len, 1);
assert_eq!(builder.sorted_spill_files.len(), 2);
Ok(())
}
}