This is an automated email from the ASF dual-hosted git repository.
github-merge-queue[bot] pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datafusion.git
The following commit(s) were added to refs/heads/main by this push:
new 6a59ce00c8 fix: account for memory that we still hold on when
splitting batch (#24852)
6a59ce00c8 is described below
commit 6a59ce00c8abb8ca24ccd0e3a46340277f16a04b
Author: Raz Luvaton <[email protected]>
AuthorDate: Thu Sep 10 10:45:17 2026 +0000
fix: account for memory that we still hold on when splitting batch (#24852)
## Which issue does this PR close?
N/A
## Rationale for this change
hold on batches between emit should be reserved
## What changes are included in this PR?
reserve memory + test
## What is the testing strategy for this PR?
unit test
## Are there any user-facing changes?
no
---
.../physical-plan/src/aggregates/hash_stream.rs | 329 ++++++++++++++++++++-
1 file changed, 320 insertions(+), 9 deletions(-)
diff --git a/datafusion/physical-plan/src/aggregates/hash_stream.rs
b/datafusion/physical-plan/src/aggregates/hash_stream.rs
index 796a2b2885..2bcff82c80 100644
--- a/datafusion/physical-plan/src/aggregates/hash_stream.rs
+++ b/datafusion/physical-plan/src/aggregates/hash_stream.rs
@@ -470,13 +470,7 @@ impl PartialHashAggregateStream {
break;
}
HandleInputResult::OOM => {
- let state_batch_result = hash_table.take_state_batch();
-
- // Emitting clears the aggregate table and releases its
- // accumulated memory. Update the reservation
accordingly.
- self.reservation.try_resize(hash_table.memory_size())?;
-
- let materialized_group_states =
state_batch_result?.ok_or_else(|| {
+ let materialized_group_states =
hash_table.take_state_batch()?.ok_or_else(|| {
internal_datafusion_err!(
"Partial hash aggregate ran out of memory with
no aggregated groups"
)
@@ -486,6 +480,7 @@ impl PartialHashAggregateStream {
self.emit_on_memory_pressure(
materialized_group_states,
&mut emitter,
+ hash_table.memory_size(),
)
.await?;
}
@@ -594,7 +589,37 @@ impl PartialHashAggregateStream {
// with batch slicing.
mut remaining_groups: RecordBatch,
emitter: &mut TryEmitter<RecordBatch, DataFusionError>,
+ hash_table_mem_size: usize,
) -> Result<()> {
+ let remaining_groups_memory = remaining_groups.get_array_memory_size();
+
+ // Emitting clears the aggregate table and releases its
+ // accumulated memory. Update the reservation accordingly.
+ // We account here for the remaining groups memory to see if we can
return batch size states
+ // if there is not enough memory, fallback to emit large batch
+ match self
+ .reservation
+ .try_resize(hash_table_mem_size + remaining_groups_memory)
+ {
+ Ok(_) => {
+ // Continue with slicing
+ }
+ Err(DataFusionError::ResourcesExhausted(_)) => {
+ // Fail to reserve memory for the hash table + state batch
while slicing so emit a huge batch
+
+ // Try resize without holding the state batch, if it fails
there is nothing we can do
+ self.reservation.try_resize(hash_table_mem_size)?;
+
+ self.reduction_factor.add_part(remaining_groups.num_rows());
+ emitter
+
.emit(remaining_groups.record_output(&self.baseline_metrics))
+ .await;
+
+ return Ok(());
+ }
+ Err(e) => return Err(e),
+ }
+
while remaining_groups.num_rows() > self.batch_size {
// More batch to output, continue in the current state.
let output = remaining_groups.slice(0, self.batch_size);
@@ -615,6 +640,10 @@ impl PartialHashAggregateStream {
self.reduction_factor.add_part(remaining_groups.num_rows());
debug_assert!(remaining_groups.num_rows() > 0);
+ // We are no longer holding on the batch while slicing, so release the
memory.
+ // The memory will now equal to the hash table size
+ self.reservation.try_shrink(remaining_groups_memory)?;
+
emitter
.emit(remaining_groups.record_output(&self.baseline_metrics))
.await;
@@ -981,14 +1010,16 @@ impl FinalHashAggregateStream {
#[cfg(test)]
mod tests {
use std::sync::Arc;
+ use std::time::Duration;
use super::*;
use crate::aggregates::{AggregateMode, PhysicalGroupBy};
use crate::execution_plan::ExecutionPlan;
use crate::test::TestMemoryExec;
+ use crate::test::exec::BarrierExec;
- use arrow::array::{Int32Array, Int64Array};
- use arrow::datatypes::{DataType, Field, Schema};
+ use arrow::array::{AsArray, Int32Array, Int64Array};
+ use arrow::datatypes::{DataType, Field, Int32Type, Schema};
use datafusion_common::Result;
use datafusion_execution::runtime_env::RuntimeEnvBuilder;
use datafusion_functions_aggregate::count::count_udaf;
@@ -1249,4 +1280,284 @@ mod tests {
Ok(())
}
+
+ /// Builds a partial hash aggregate stream over a single input batch of
+ /// `num_groups` distinct groups, running under `memory_limit` bytes.
+ ///
+ /// The input does not signal end-of-stream until `wait_finish` is called
+ /// on the returned [`BarrierExec`], so any output produced before that can
+ /// only come from the memory pressure emission path (normal output waits
+ /// for all input). Skip partial aggregation is disabled for the same
reason.
+ fn partial_stream_under_memory_limit(
+ memory_limit: usize,
+ batch_size: usize,
+ num_groups: usize,
+ ) -> Result<(
+ SendableRecordBatchStream,
+ Arc<BarrierExec>,
+ Arc<datafusion_execution::runtime_env::RuntimeEnv>,
+ )> {
+ let schema = Arc::new(Schema::new(vec![
+ Field::new("group_col", DataType::Int32, false),
+ Field::new("value_col", DataType::Int64, false),
+ ]));
+
+ let group_ids: Vec<i32> = (0..num_groups as i32).collect();
+ let values: Vec<i64> = vec![1; num_groups];
+
+ let batch = RecordBatch::try_new(
+ Arc::clone(&schema),
+ vec![
+ Arc::new(Int32Array::from(group_ids)),
+ Arc::new(Int64Array::from(values)),
+ ],
+ )?;
+ let input_partitions = vec![vec![batch]];
+
+ let runtime = RuntimeEnvBuilder::default()
+ .with_memory_limit(memory_limit, 1.0)
+ .build_arc()?;
+
+ let mut task_ctx =
TaskContext::default().with_runtime(Arc::clone(&runtime));
+ let session_config = task_ctx
+ .session_config()
+ .clone()
+ .set(
+ "datafusion.execution.batch_size",
+ &datafusion_common::ScalarValue::UInt64(Some(batch_size as
u64)),
+ )
+ .set(
+
"datafusion.execution.skip_partial_aggregation_probe_ratio_threshold",
+ &datafusion_common::ScalarValue::Float64(Some(1.0)),
+ );
+ task_ctx = task_ctx.with_session_config(session_config);
+ let task_ctx = Arc::new(task_ctx);
+
+ // Create aggregate: COUNT(*) GROUP BY group_col
+ let group_expr = vec![(col("group_col", &schema)?,
"group_col".to_string())];
+ let aggr_expr = vec![Arc::new(
+ AggregateExprBuilder::new(count_udaf(), vec![col("value_col",
&schema)?])
+ .schema(Arc::clone(&schema))
+ .alias("count_value")
+ .build()?,
+ )];
+
+ let input = Arc::new(
+ BarrierExec::new(input_partitions, Arc::clone(&schema))
+ .without_start_barrier()
+ .with_finish_barrier()
+ .with_log(false),
+ );
+
+ let aggregate_exec = AggregateExec::try_new(
+ AggregateMode::Partial,
+ PhysicalGroupBy::new_single(group_expr),
+ aggr_expr,
+ vec![None],
+ Arc::clone(&input) as Arc<dyn ExecutionPlan>,
+ Arc::clone(&schema),
+ )?;
+
+ let stream =
+ PartialHashAggregateStream::new(&aggregate_exec, &task_ctx,
0)?.into_stream();
+
+ Ok((stream, input, runtime))
+ }
+
+ #[tokio::test]
+ async fn
test_partial_hash_stream_accounts_held_batch_on_memory_pressure_while_slicing()
+ -> Result<()> {
+ // When memory pressure triggers early emission, the materialized state
+ // batch is held while it is sliced into `batch_size` outputs. The
+ // stream must keep that held batch accounted for in its memory
+ // reservation until the last slice is emitted; before the fix the
+ // reservation was resized down to just the (emptied) hash table size,
+ // leaving the held batch unaccounted.
+
+ let batch_size = 1024;
+ // One row per group so the state batch is emitted in 4 slices
+ let num_groups = 4 * batch_size;
+
+ // Smaller than the building hash table (so pressure triggers) but
large
+ // enough to hold the materialized state batch (so slicing can proceed)
+ let memory_limit = 100 * 1024;
+ let (mut stream, input, runtime) =
+ partial_stream_under_memory_limit(memory_limit, batch_size,
num_groups)?;
+
+ // The first output batch must be a pressure-emitted slice, with the
rest
+ // of the materialized state batch still held by the stream
+ let first = tokio::time::timeout(Duration::from_secs(5), stream.next())
+ .await
+ .expect(
+ "did not get early emit due to OOM, this probably means that
the \
+ memory limit is too high to trigger the OOM",
+ )
+ .expect("stream ended early")?;
+ assert_eq!(first.num_rows(), batch_size);
+
+ // The emitted slice shares buffers with the held state batch, so its
+ // array memory size reflects the full held allocation
+ let held_size = first.get_array_memory_size();
+ let reserved = runtime.memory_pool.reserved();
+ assert!(
+ reserved >= held_size,
+ "memory pool has {reserved} bytes reserved but the stream is \
+ holding a materialized state batch of {held_size} bytes"
+ );
+
+ let second = stream.next().await.expect("stream ended early")?;
+ assert_eq!(second.num_rows(), batch_size);
+
+ // Make sure the state batch is really being sliced (and not emitted
whole by the fallback path):
+ // the second output must share the same underlying buffer as the first
+ //
+ // If you changed the code and this fail because
+ // - you now deep copy `batch_size` from the full state batch, please
update this assertion to something else
+ // - you only take batch size from the hash table, you can remove the
test
+ assert_eq!(
+ first
+ .column(0)
+ .as_primitive::<Int32Type>()
+ .values()
+ .inner()
+ .data_ptr(),
+ second
+ .column(0)
+ .as_primitive::<Int32Type>()
+ .values()
+ .inner()
+ .data_ptr(),
+ "both batches should be slices of the same materialized state
batch"
+ );
+
+ // Let the input finish and drain the stream: no groups lost
+ input.wait_finish().await;
+ let mut total_rows = first.num_rows() + second.num_rows();
+ while let Some(batch) = stream.next().await {
+ total_rows += batch?.num_rows();
+ }
+ assert_eq!(total_rows, num_groups);
+
+ Ok(())
+ }
+
+ #[tokio::test]
+ async fn
test_partial_hash_stream_emits_whole_batch_when_held_batch_does_not_fit()
+ -> Result<()> {
+ // When memory pressure triggers early emission but the materialized
+ // state batch itself does not fit in the reservation, the stream must
+ // not fail with a resources exhausted error. Instead it gives up on
+ // slicing and emits the whole state batch at once.
+
+ let batch_size = 1024;
+ let num_groups = 4 * batch_size;
+
+ // Smaller than the materialized state batch (4096 rows of Int32 group
+ // keys plus Int64 counts is at least 48 KiB), so the reservation for
+ // hash table held batch fails. The emptied hash table itself is tiny
+ // and still fits.
+ let memory_limit = 32 * 1024;
+ let (mut stream, input, runtime) =
+ partial_stream_under_memory_limit(memory_limit, batch_size,
num_groups)?;
+
+ let first = tokio::time::timeout(Duration::from_secs(5), stream.next())
+ .await
+ .expect(
+ "did not get early emit due to OOM, this probably means that
the \
+ memory limit is too high to trigger the OOM",
+ )
+ .expect("stream ended early")?;
+
+ // The whole state batch is emitted at once instead of `batch_size`
slices
+ assert_eq!(first.num_rows(), num_groups);
+ assert!(
+ first.get_array_memory_size() > memory_limit,
+ "test setup is wrong: the state batch fits within the memory
limit, \
+ so the slicing path would have been taken"
+ );
+
+ // Unlike the slicing path, the stream does not hold on to the emitted
+ // batch, so it must not be accounted for in the reservation. Only the
+ // (emptied) hash table remains reserved
+ let emitted_size = first.get_array_memory_size();
+ let reserved = runtime.memory_pool.reserved();
+ assert!(
+ reserved < emitted_size,
+ "memory pool has {reserved} bytes reserved but the stream no
longer \
+ holds the emitted state batch of {emitted_size} bytes"
+ );
+
+ input.wait_finish().await;
+ let mut total_rows = first.num_rows();
+ while let Some(batch) = stream.next().await {
+ total_rows += batch?.num_rows();
+ }
+ assert_eq!(total_rows, num_groups);
+
+ Ok(())
+ }
+
+ #[tokio::test]
+ async fn test_partial_hash_stream_releases_held_batch_after_last_slice()
-> Result<()>
+ {
+ // While the pressure-emitted state batch is sliced, the stream holds
+ // the remaining groups and keeps them reserved. Once the last slice is
+ // handed out nothing is held anymore, so the reservation must drop
+ // back to just the (emptied) hash table before the input is resumed.
+
+ let batch_size = 1024;
+ let num_slices = 4;
+ let num_groups = num_slices * batch_size;
+
+ let memory_limit = 100 * 1024;
+ let (mut stream, input, runtime) =
+ partial_stream_under_memory_limit(memory_limit, batch_size,
num_groups)?;
+
+ // The input has not finished, so all of these are pressure-emitted
slices
+ let mut held_size = 0;
+ for slice_idx in 0..num_slices {
+ let slice = if slice_idx == 0 {
+ tokio::time::timeout(Duration::from_secs(5), stream.next())
+ .await
+ .expect(
+ "did not get early emit due to OOM, this probably
means that the \
+ memory limit is too high to trigger the OOM",
+ )
+ .expect("stream ended early")?
+ } else {
+ stream.next().await.expect("stream ended early")?
+ };
+
+ assert_eq!(slice.num_rows(), batch_size);
+
+ // Every slice shares buffers with the held state batch, so this is
+ // the size of the full held allocation
+ held_size = slice.get_array_memory_size();
+ let reserved = runtime.memory_pool.reserved();
+
+ if slice_idx + 1 < num_slices {
+ assert!(
+ reserved >= held_size,
+ "after slice {slice_idx} the stream still holds
{held_size} \
+ bytes but only {reserved} bytes are reserved"
+ );
+ } else {
+ assert!(
+ reserved < held_size,
+ "after the last slice nothing is held anymore but
{reserved} \
+ bytes are still reserved (held batch was {held_size}
bytes)"
+ );
+ }
+ }
+ assert!(held_size > 0);
+
+ input.wait_finish().await;
+ let mut total_rows = num_groups;
+ while let Some(batch) = stream.next().await {
+ total_rows += batch?.num_rows();
+ }
+ assert_eq!(total_rows, num_groups);
+
+ Ok(())
+ }
}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]