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]

Reply via email to