diff --git a/datafusion/physical-plan/src/aggregates/ordered_final_stream.rs b/datafusion/physical-plan/src/aggregates/ordered_final_stream.rs index 8823dcb32f47c..ac6f2491b481e 100644 --- a/datafusion/physical-plan/src/aggregates/ordered_final_stream.rs +++ b/datafusion/physical-plan/src/aggregates/ordered_final_stream.rs @@ -900,3 +900,286 @@ impl RecordBatchStream for OrderedFinalAggregateStream { Arc::clone(&self.schema) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::ExecutionPlan; + use crate::aggregates::PhysicalGroupBy; + use crate::common::collect; + use crate::stream::RecordBatchStreamAdapter; + use crate::test::TestMemoryExec; + use arrow::array::{Int64Array, StringViewArray}; + use arrow::datatypes::{DataType, Field, Schema}; + use datafusion_execution::config::SessionConfig; + use datafusion_execution::memory_pool::{GreedyMemoryPool, MemoryPool}; + use datafusion_execution::runtime_env::RuntimeEnvBuilder; + use datafusion_functions_aggregate::{min_max::min_udaf, sum::sum_udaf}; + use datafusion_physical_expr::aggregate::AggregateExprBuilder; + use datafusion_physical_expr::expressions::col; + use futures::FutureExt; + use futures::channel::mpsc; + use std::collections::BTreeMap; + + #[derive(Clone, Copy)] + enum Finish { + Collect, + DropDuringMerge, + DropDuringReplay, + InputError, + } + + #[tokio::test] + async fn spill_replay_with_another_ordered_partition() -> Result<()> { + for input_batches in [28, 36, 55, 63] { + run_shared_pool_case(input_batches, 600 * 1024, Finish::Collect).await?; + } + // The same input also produces the reference results without spilling. + run_shared_pool_case(63, 10 * 1024 * 1024, Finish::Collect).await + } + + #[tokio::test] + async fn spill_replay_releases_memory_on_drop() -> Result<()> { + run_shared_pool_case(36, 600 * 1024, Finish::DropDuringMerge).await?; + run_shared_pool_case(36, 600 * 1024, Finish::DropDuringReplay).await + } + + #[tokio::test] + async fn ordered_spill_releases_memory_on_input_error() -> Result<()> { + run_shared_pool_case(36, 600 * 1024, Finish::InputError).await + } + + async fn run_shared_pool_case( + input_batches: i64, + limit: usize, + finish: Finish, + ) -> Result<()> { + let schema = Arc::new(Schema::new(vec![ + Field::new("a", DataType::Int64, false), + Field::new("b", DataType::Int64, false), + Field::new("v", DataType::Int64, false), + Field::new("s", DataType::Utf8View, false), + ])); + let groups = PhysicalGroupBy::new_single(vec![ + (col("a", &schema)?, "a".into()), + (col("b", &schema)?, "b".into()), + ]); + let expressions = vec![ + Arc::new( + AggregateExprBuilder::new(sum_udaf(), vec![col("v", &schema)?]) + .schema(Arc::clone(&schema)) + .alias("sum") + .build()?, + ), + Arc::new( + AggregateExprBuilder::new(min_udaf(), vec![col("s", &schema)?]) + .schema(Arc::clone(&schema)) + .alias("min") + .build()?, + ), + ]; + let empty = TestMemoryExec::try_new_exec(&[vec![]], Arc::clone(&schema), None)?; + let partial = AggregateExec::try_new( + AggregateMode::Partial, + groups.clone(), + expressions.clone(), + vec![None; 2], + empty, + Arc::clone(&schema), + )?; + let partial_schema = partial.schema(); + let ordering = + LexOrdering::new([PhysicalSortExpr::new_default(col("a", &partial_schema)?)]) + .unwrap(); + let input = TestMemoryExec::try_new( + &[vec![], vec![]], + Arc::clone(&partial_schema), + None, + )? + .try_with_sort_information(vec![ordering])?; + let aggregate = AggregateExec::try_new( + AggregateMode::FinalPartitioned, + groups.as_final(), + expressions, + vec![None; 2], + Arc::new(TestMemoryExec::update_cache(&Arc::new(input))), + schema, + )?; + assert_eq!( + aggregate.input_order_mode(), + &InputOrderMode::PartiallySorted(vec![0]) + ); + + let pool: Arc = Arc::new(GreedyMemoryPool::new(limit)); + let context = Arc::new( + TaskContext::default() + .with_session_config( + SessionConfig::new().with_batch_size(128).set_bool( + "datafusion.execution.enable_migration_aggregate", + true, + ), + ) + .with_runtime( + RuntimeEnvBuilder::new() + .with_max_spill_merge_fan_in(2) + .with_memory_pool(Arc::clone(&pool)) + .build_arc()?, + ), + ); + let mut streams = vec![]; + let mut senders = vec![]; + for partition in 0..2 { + let (sender, receiver) = mpsc::unbounded(); + let input = Box::pin(RecordBatchStreamAdapter::new( + Arc::clone(&partial_schema), + receiver, + )); + let stream = OrderedFinalAggregateStream::new_with_input( + &aggregate, + &context, + partition, + input, + aggregate.input_order_mode(), + )?; + senders.push(sender); + streams.push(stream); + } + let mut expected = BTreeMap::new(); + let mut make_batch = |partition: i64, start: i64| { + for value in start..start + 128 { + let entry = expected + .entry(( + partition, + if partition == 0 && value % 128 == 0 { + 0 + } else { + value + }, + )) + .or_insert_with(|| (0, (value % 2).to_string())); + entry.0 += value * 2; + } + RecordBatch::try_new( + Arc::clone(&partial_schema), + vec![ + Arc::new(Int64Array::from(vec![partition; 128])), + Arc::new(Int64Array::from_iter_values((start..start + 128).map( + |value| { + if partition == 0 && value % 128 == 0 { + 0 + } else { + value + } + }, + ))), + Arc::new(Int64Array::from_iter_values( + (start..start + 128).map(|v| v * 2), + )), + Arc::new(StringViewArray::from_iter_values( + (start..start + 128).map(|v| if v % 2 == 0 { "0" } else { "1" }), + )), + ], + ) + .unwrap() + }; + // Keep partition 1's incomplete ordered run live while partition 0 spills + // and replays. Channel inputs return Pending after each supplied batch, + // making this interleaving independent of task scheduling. + for batch in 0..55 { + senders[1] + .unbounded_send(Ok(make_batch(1, batch * 128))) + .unwrap(); + assert!(streams[1].next().now_or_never().is_none()); + } + let held = streams[1].reservation.size(); + assert!(held > 500 * 1024); + for batch in 0..input_batches { + // Repeated keys cross spill runs, so replay must merge their sums. + senders[0] + .unbounded_send(Ok(make_batch(0, batch * 128))) + .unwrap(); + assert!(streams[0].next().now_or_never().is_none()); + } + if limit == 600 * 1024 { + assert!(aggregate.metrics().unwrap().spill_count().unwrap() > 0); + } + let mut first = streams.remove(0); + match finish { + Finish::Collect => { + senders[0].close_channel(); + let mut output = collect(Box::pin(first)).await?; + assert_eq!(pool.reserved(), held); + senders[1].close_channel(); + output.extend(collect(Box::pin(streams.remove(0))).await?); + let mut actual = BTreeMap::new(); + for batch in output { + let a = batch + .column(0) + .as_any() + .downcast_ref::() + .unwrap(); + let b = batch + .column(1) + .as_any() + .downcast_ref::() + .unwrap(); + let sum = batch + .column(2) + .as_any() + .downcast_ref::() + .unwrap(); + let min = batch + .column(3) + .as_any() + .downcast_ref::() + .unwrap(); + for row in 0..batch.num_rows() { + assert!( + actual + .insert( + (a.value(row), b.value(row)), + (sum.value(row), min.value(row).to_string()) + ) + .is_none() + ); + } + } + assert_eq!(actual, expected); + assert_eq!( + aggregate.metrics().unwrap().spill_count().unwrap() > 0, + limit == 600 * 1024 + ); + } + Finish::DropDuringMerge => { + senders[0].close_channel(); + let _ = first.next().now_or_never(); + assert!(matches!( + first.state.as_ref(), + Some(OrderedFinalAggregateState::MergingSpills { .. }) + )); + drop(first); + } + Finish::DropDuringReplay => { + senders[0].close_channel(); + first.next().await.unwrap()?; + assert!(pool.reserved() > held); + drop(first); + assert_eq!(pool.reserved(), held); + } + Finish::InputError => { + senders[0] + .unbounded_send(datafusion_common::exec_err!( + "injected input failure" + )) + .unwrap(); + let error = first.next().await.unwrap().unwrap_err(); + assert!(error.to_string().contains("injected input failure")); + assert_eq!(pool.reserved(), held); + drop(first); + } + } + drop(streams); + assert_eq!(pool.reserved(), 0); + Ok(()) + } +} diff --git a/datafusion/sqllogictest/test_files/ordered_aggregate_spill.slt b/datafusion/sqllogictest/test_files/ordered_aggregate_spill.slt index 67fa20101dbfe..548dfc69b1b9a 100644 --- a/datafusion/sqllogictest/test_files/ordered_aggregate_spill.slt +++ b/datafusion/sqllogictest/test_files/ordered_aggregate_spill.slt @@ -35,6 +35,11 @@ SET datafusion.optimizer.prefer_existing_sort = true statement ok SET datafusion.execution.enable_migration_aggregate = true +# Limit merge buffers so replay can allocate while another partition retains +# its aggregate state in the shared greedy pool. +statement ok +SET datafusion.runtime.max_spill_merge_fan_in = 2 + statement ok SET datafusion.runtime.memory_limit = '1M' @@ -108,7 +113,7 @@ GROUP BY round(v1, -4), v1 % 5000 ---- 45000 values hashing to e6ece4b4b86e6152a1c785e90ab5ba12 -# Round 2: The same query spills five times with a 600 KB limit. +# Round 2: The same query spills with a 600 KB limit. statement ok SET datafusion.runtime.memory_limit = '600K' @@ -132,7 +137,7 @@ GROUP BY round(v1, -4), v1 % 5000 ---- 45000 values hashing to e6ece4b4b86e6152a1c785e90ab5ba12 -# Round 3: The same query spills six times with a 500 KB limit. +# Round 3: The same query spills with a 500 KB limit. statement ok SET datafusion.runtime.memory_limit = '500K' @@ -162,7 +167,7 @@ GROUP BY round(v1, -4), v1 % 5000 statement ok SET datafusion.runtime.memory_limit = '600K' -# Ensures final aggregate has spill_count > 0 +# Check spilled rows because extra merge passes can change the byte unit. query TT EXPLAIN ANALYZE SELECT round(v1, -4), v1 % 5000, @@ -171,7 +176,7 @@ FROM generate_series(20000) AS t1(v1) GROUP BY round(v1, -4), v1 % 5000 ---- Plan with Metrics -01)AggregateExec: mode=FinalPartitioned,aggr=[sum(t1.v1 * Int64(2)), min(t1.v1 % Int64(2))], ordering_mode=PartiallySorted([0]), metrics=[spilled_bytes=KB,] +01)AggregateExec: mode=FinalPartitioned,aggr=[sum(t1.v1 * Int64(2)), min(t1.v1 % Int64(2))], ordering_mode=PartiallySorted([0]), metrics=[spilled_rows= K,] 02)--RepartitionExec:input_partitions=1, maintains_sort_order=true 03)----AggregateExec: mode=Partial,aggr=[sum(t1.v1 * Int64(2)), min(t1.v1 % Int64(2))], ordering_mode=PartiallySorted([0]), metrics=[spill_count=0,] @@ -249,5 +254,8 @@ RESET datafusion.execution.batch_size statement ok SET datafusion.execution.target_partitions = 4 +statement ok +RESET datafusion.runtime.max_spill_merge_fan_in + statement ok RESET datafusion.catalog.create_default_catalog_and_schema