From 33827bbb1e87eb24fa5c1b07f250b2f9b84443d4 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 10:10:27 -0600 Subject: [PATCH 1/5] fix: let a spilled final aggregate read its spill files back past its memory share DataFusion 55's FinalHashAggregateStream replays its merged spill files through a stream that can't spill, so a refused memory request there fails the task. The merge's read buffers are a sibling reservation of the same consumer and take as many spill files as fit, so the replay often finds the consumer's share already taken (#6254). Both Comet pools now record that request as overcommit instead of refusing it. They recognize it as a request from a FinalHashAggregateStream consumer while another of its reservations holds memory, which in DataFusion 55.1 happens only during the replay. The replay emits its finished groups after every batch, so the overcommit stays around one batch of groups, and releases repay it first. Refusals while the aggregate reads its input, and while the merge picks its files, are unchanged. Remove this once Comet's DataFusion includes apache/datafusion#25383. --- .../src/execution/memory_pools/fair_pool.rs | 129 +++++++++++++++-- native/core/src/execution/memory_pools/mod.rs | 1 + .../execution/memory_pools/spill_replay.rs | 47 ++++++ .../execution/memory_pools/unified_pool.rs | 134 ++++++++++++++++-- .../comet/exec/CometAggregateSuite.scala | 38 ++++- 5 files changed, 330 insertions(+), 19 deletions(-) create mode 100644 native/core/src/execution/memory_pools/spill_replay.rs diff --git a/native/core/src/execution/memory_pools/fair_pool.rs b/native/core/src/execution/memory_pools/fair_pool.rs index 56b09da7365..f13262d6aee 100644 --- a/native/core/src/execution/memory_pools/fair_pool.rs +++ b/native/core/src/execution/memory_pools/fair_pool.rs @@ -20,13 +20,14 @@ use std::{ fmt::{Debug, Display, Formatter, Result as FmtResult}, }; -use super::spark_memory::SparkMemory; -use datafusion::common::resources_err; +use super::{spark_memory::SparkMemory, spill_replay}; +use datafusion::common::resources_datafusion_err; use datafusion::execution::memory_pool::MemoryConsumer; use datafusion::{ common::DataFusionError, execution::memory_pool::{MemoryPool, MemoryReservation}, }; +use log::debug; use parking_lot::Mutex; /// A DataFusion fair `MemoryPool` implementation for Comet. Internally this is @@ -84,6 +85,42 @@ impl CometFairMemoryPool { pub(super) fn overcommit(&self) -> usize { self.spark.overcommit() } + + /// Records `additional` bytes for `reservation` whatever the fair and pool limits, and carries + /// what Spark doesn't grant as overcommit. See [`SparkMemory`]. + fn record( + &self, + state: &mut CometFairPoolState, + reservation: &MemoryReservation, + additional: usize, + ) { + self.spark.acquire(additional); + state.used = state.used.saturating_add(additional); + let consumer_used = state.consumer_used(reservation); + *consumer_used = consumer_used.saturating_add(additional); + } + + /// Refuses a `try_grow` with `err`, unless it comes from a final hash aggregate reading its + /// spill files back, which can't spill. That request is recorded instead; see + /// [`spill_replay`]. + fn refuse( + &self, + state: &mut CometFairPoolState, + reservation: &MemoryReservation, + additional: usize, + err: DataFusionError, + ) -> Result<(), DataFusionError> { + if !spill_replay::is_spill_replay(reservation, *state.consumer_used(reservation)) { + return Err(err); + } + debug!( + "Task {} records {additional} bytes for {} while it reads its spill files back: {err}", + self.spark.task_attempt_id(), + reservation.consumer().name() + ); + self.record(state, reservation, additional); + Ok(()) + } } impl Display for CometFairMemoryPool { @@ -122,11 +159,7 @@ impl MemoryPool for CometFairMemoryPool { if additional == 0 { return; } - let mut state = self.state.lock(); - self.spark.acquire(additional); - state.used = state.used.saturating_add(additional); - let consumer_used = state.consumer_used(reservation); - *consumer_used = consumer_used.saturating_add(additional); + self.record(&mut self.state.lock(), reservation, additional); } fn shrink(&self, reservation: &MemoryReservation, subtractive: usize) { @@ -160,31 +193,34 @@ impl MemoryPool for CometFairMemoryPool { .expect("overflow in checked_div"); let consumer_used = *state.consumer_used(reservation); if limit < consumer_used.saturating_add(additional) { - return resources_err!( + let err = resources_datafusion_err!( "Failed to acquire {additional} bytes where this consumer already holds {consumer_used} bytes and the fair limit is {limit} bytes, {num} registered ({} bytes overcommitted)", self.spark.overcommit() ); + return self.refuse(&mut state, reservation, additional, err); } // The shares alone do not bound the pool's total, because a consumer keeps what it // reserved before another consumer registered. let used = state.used; if self.pool_size < used.saturating_add(additional) { - return resources_err!( + let err = resources_datafusion_err!( "Failed to acquire {additional} bytes where {used} bytes already reserved ({} bytes overcommitted) and the pool limit is {} bytes", self.spark.overcommit(), self.pool_size ); + return self.refuse(&mut state, reservation, additional, err); } // A partial grant is handed back and refused, which triggers spilling in the caller. if let Err(refusal) = self.spark.try_acquire(additional)? { - return resources_err!( + let err = resources_datafusion_err!( "Failed to acquire {} bytes plus {} bytes overcommitted, only got {} bytes. Reserved: {} bytes", additional, refusal.overcommit, refusal.granted, state.used ); + return self.refuse(&mut state, reservation, additional, err); } state.used = state .used @@ -328,4 +364,77 @@ mod tests { assert_eq!(pool.reserved(), 50); assert_eq!(fake.held(), 50); } + + #[test] + fn a_final_aggregate_reading_its_spill_files_back_carries_what_spark_refuses() { + let fake = FakeSpark::with(100); + let pool = Arc::new(CometFairMemoryPool::with_spark(fake.memory(), 1000)); + let dyn_pool: Arc = Arc::clone(&pool) as _; + // Like DataFusion 55's final hash aggregate, whose spill merge and replay table are + // sibling reservations of one consumer. + let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&dyn_pool); + let replay = merge.new_empty(); + merge.try_grow(90).unwrap(); + + // Spark grants 10 of the replay's 30 bytes, and the other 20 are overcommit. + replay.try_grow(30).unwrap(); + assert_eq!(pool.reserved(), 120); + assert_eq!(fake.held(), 100); + assert_eq!(pool.overcommit(), 20); + + // The replay emits groups and shrinks, which repays the overcommit before Spark. + replay.shrink(25); + assert_eq!(pool.overcommit(), 0); + assert_eq!(fake.held(), 95); + drop(merge); + drop(replay); + assert_eq!(pool.reserved(), 0); + assert_eq!(fake.held(), 0); + } + + #[test] + fn a_final_aggregate_reading_its_spill_files_back_may_pass_its_fair_limit() { + // Spark grants everything, so only the pool's own checks refuse. + let fake = FakeSpark::with(usize::MAX); + let pool = Arc::new(CometFairMemoryPool::with_spark(fake.memory(), 100)); + let dyn_pool: Arc = Arc::clone(&pool) as _; + let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&dyn_pool); + let replay = merge.new_empty(); + merge.try_grow(90).unwrap(); + + replay.try_grow(30).unwrap(); + assert_eq!(pool.reserved(), 120); + assert_eq!(fake.held(), 120); + assert_eq!(pool.overcommit(), 0); + } + + #[test] + fn other_refusals_are_unchanged() { + let fake = FakeSpark::with(100); + let pool: Arc = + Arc::new(CometFairMemoryPool::with_spark(fake.memory(), 1000)); + + // While the aggregate reads its input, its table is the consumer's only reservation + // holding memory, so a refusal makes it spill. + let table = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); + let replay = table.new_empty(); + table.try_grow(90).unwrap(); + assert!(table.try_grow(30).is_err()); + drop(table); + + // The merge picks its spill files while nothing else is held, so a refusal still limits + // how many it opens. + let merge = replay.new_empty(); + merge.try_grow(90).unwrap(); + assert!(merge.try_grow(30).is_err()); + drop(merge); + + // Any other operator with a sibling holding memory is still refused. + let sort = MemoryConsumer::new("ExternalSorterMerge[0]").register(&pool); + let sibling = sort.new_empty(); + sort.try_grow(90).unwrap(); + assert!(sibling.try_grow(30).is_err()); + assert_eq!(pool.reserved(), 90); + assert_eq!(fake.held(), 90); + } } diff --git a/native/core/src/execution/memory_pools/mod.rs b/native/core/src/execution/memory_pools/mod.rs index 3efceb072c0..a8c0997081e 100644 --- a/native/core/src/execution/memory_pools/mod.rs +++ b/native/core/src/execution/memory_pools/mod.rs @@ -20,6 +20,7 @@ mod fair_pool; pub mod logging_pool; mod plan_pool; mod spark_memory; +mod spill_replay; mod task_shared; mod unified_pool; diff --git a/native/core/src/execution/memory_pools/spill_replay.rs b/native/core/src/execution/memory_pools/spill_replay.rs new file mode 100644 index 00000000000..d33426fd189 --- /dev/null +++ b/native/core/src/execution/memory_pools/spill_replay.rs @@ -0,0 +1,47 @@ +// 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. + +//! Lets a final hash aggregate read its spill files back past its share of the pool. +//! +//! Once DataFusion 55's `FinalHashAggregateStream` has spilled, it merges its sorted spill files +//! and replays them through an `OrderedFinalAggregateStream` that has no way to spill, so a +//! refused memory request there fails the task (#6254). The merge reserves read buffers for as +//! many spill files as fit, and those buffers belong to the same consumer, so the replay often +//! finds the consumer's share already taken. The pools record such a request as overcommit +//! instead of refusing it. The replay emits every finished group after each batch, so it holds +//! about one batch of groups, and releasing memory repays the overcommit first. +//! +//! Remove this once Comet's DataFusion has apache/datafusion#25383, which leaves the replay room +//! when the merge picks its files. + +use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation}; + +/// Whether `consumer` belongs to a `FinalHashAggregateStream`, whose spill replay can't spill. +pub(super) fn is_final_hash_aggregate(consumer: &MemoryConsumer) -> bool { + consumer.name().starts_with("FinalHashAggregateStream[") +} + +/// Whether a refused request from `reservation` should be recorded as overcommit instead. +/// `consumer_used` is what the reservation's consumer holds across all of its reservations. +/// +/// A final hash aggregate grows one reservation while another of its reservations holds memory +/// only during the replay, while the merge holds its read buffers. Before that, the aggregate's +/// table is its only reservation holding memory, so a refusal makes it spill. The merge picks its +/// files while nothing else is held, so a refusal there still limits how many it opens. +pub(super) fn is_spill_replay(reservation: &MemoryReservation, consumer_used: usize) -> bool { + consumer_used > reservation.size() && is_final_hash_aggregate(reservation.consumer()) +} diff --git a/native/core/src/execution/memory_pools/unified_pool.rs b/native/core/src/execution/memory_pools/unified_pool.rs index edc51e0dd92..adf384e9970 100644 --- a/native/core/src/execution/memory_pools/unified_pool.rs +++ b/native/core/src/execution/memory_pools/unified_pool.rs @@ -16,16 +16,18 @@ // under the License. use std::{ + collections::HashMap, fmt::{Debug, Display, Formatter, Result as FmtResult}, sync::atomic::{AtomicUsize, Ordering::Relaxed}, }; -use super::spark_memory::SparkMemory; +use super::{spark_memory::SparkMemory, spill_replay}; use datafusion::{ common::{resources_datafusion_err, DataFusionError}, - execution::memory_pool::{MemoryPool, MemoryReservation}, + execution::memory_pool::{MemoryConsumer, MemoryPool, MemoryReservation}, }; -use log::warn; +use log::{debug, warn}; +use parking_lot::Mutex; /// A DataFusion `MemoryPool` implementation for Comet that delegates to /// Spark's off-heap executor memory pool via JNI by calling @@ -33,6 +35,10 @@ use log::warn; pub struct CometUnifiedMemoryPool { spark: SparkMemory, used: AtomicUsize, + /// Bytes held by each final hash aggregate's consumer across its reservations, keyed by + /// [`MemoryConsumer::id`], which [`spill_replay`] needs. Other consumers aren't tracked, so + /// they never take the lock. + final_aggregates: Mutex>, } impl Debug for CometUnifiedMemoryPool { @@ -49,6 +55,7 @@ impl CometUnifiedMemoryPool { Self { spark, used: AtomicUsize::new(0), + final_aggregates: Mutex::new(HashMap::new()), } } @@ -56,6 +63,31 @@ impl CometUnifiedMemoryPool { pub(super) fn overcommit(&self) -> usize { self.spark.overcommit() } + + /// Applies `update` to what `reservation`'s consumer holds, if it is a final hash aggregate. + fn track(&self, reservation: &MemoryReservation, update: impl FnOnce(&mut usize)) { + if spill_replay::is_final_hash_aggregate(reservation.consumer()) { + if let Some(used) = self + .final_aggregates + .lock() + .get_mut(&reservation.consumer().id()) + { + update(used); + } + } + } + + /// Whether a refused request from `reservation` comes from a final hash aggregate reading its + /// spill files back, which can't spill; see [`spill_replay`]. + fn is_spill_replay(&self, reservation: &MemoryReservation) -> bool { + let consumer_used = self + .final_aggregates + .lock() + .get(&reservation.consumer().id()) + .copied() + .unwrap_or(0); + spill_replay::is_spill_replay(reservation, consumer_used) + } } impl Drop for CometUnifiedMemoryPool { @@ -86,11 +118,23 @@ impl MemoryPool for CometUnifiedMemoryPool { "CometUnifiedMemoryPool" } + fn register(&self, consumer: &MemoryConsumer) { + if spill_replay::is_final_hash_aggregate(consumer) { + self.final_aggregates.lock().insert(consumer.id(), 0); + } + } + + fn unregister(&self, consumer: &MemoryConsumer) { + if spill_replay::is_final_hash_aggregate(consumer) { + self.final_aggregates.lock().remove(&consumer.id()); + } + } + /// Records memory that already exists, so it must not fail; see [`SparkMemory`]. // Rust 1.99 deprecates `fetch_update` in favor of `try_update`, which needs Rust 1.95, newer // than the workspace `rust-version`. #[allow(deprecated)] - fn grow(&self, _: &MemoryReservation, additional: usize) { + fn grow(&self, reservation: &MemoryReservation, additional: usize) { if additional == 0 { return; } @@ -98,12 +142,13 @@ impl MemoryPool for CometUnifiedMemoryPool { self.used .fetch_update(Relaxed, Relaxed, |old| Some(old.saturating_add(additional))) .unwrap(); + self.track(reservation, |used| *used = used.saturating_add(additional)); } // Rust 1.99 deprecates `fetch_update` in favor of `try_update`, which needs Rust 1.95, newer // than the workspace `rust-version`. #[allow(deprecated)] - fn shrink(&self, _: &MemoryReservation, size: usize) { + fn shrink(&self, reservation: &MemoryReservation, size: usize) { if let Err(e) = self.spark.release(size) { panic!( "Task {} failed to return {size} bytes to Spark: {e:?}", @@ -119,23 +164,38 @@ impl MemoryPool for CometUnifiedMemoryPool { self.spark.task_attempt_id() ); } + self.track(reservation, |used| *used = used.saturating_sub(size)); } // Rust 1.99 deprecates `fetch_update` in favor of `try_update`, which needs Rust 1.95, newer // than the workspace `rust-version`. #[allow(deprecated)] - fn try_grow(&self, _: &MemoryReservation, additional: usize) -> Result<(), DataFusionError> { + fn try_grow( + &self, + reservation: &MemoryReservation, + additional: usize, + ) -> Result<(), DataFusionError> { if additional > 0 { // A partial grant is handed back and refused, which triggers spilling in the caller. if let Err(refusal) = self.spark.try_acquire(additional)? { - return Err(resources_datafusion_err!( + let err = resources_datafusion_err!( "Task {} failed to acquire {} bytes plus {} bytes overcommitted, only got {}. Reserved: {}", self.spark.task_attempt_id(), additional, refusal.overcommit, refusal.granted, self.reserved() - )); + ); + if !self.is_spill_replay(reservation) { + return Err(err); + } + debug!( + "Task {} records {additional} bytes for {} while it reads its spill files back: {err}", + self.spark.task_attempt_id(), + reservation.consumer().name() + ); + self.grow(reservation, additional); + return Ok(()); } if let Err(prev) = self .used @@ -148,6 +208,7 @@ impl MemoryPool for CometUnifiedMemoryPool { prev )); } + self.track(reservation, |used| *used = used.saturating_add(additional)); } Ok(()) } @@ -231,4 +292,61 @@ mod tests { assert_eq!(pool.spark.overcommit(), 0); assert_eq!(fake.held(), 0); } + + #[test] + fn a_final_aggregate_reading_its_spill_files_back_carries_what_spark_refuses() { + let fake = FakeSpark::with(100); + let pool = Arc::new(CometUnifiedMemoryPool::with_spark(fake.memory())); + let dyn_pool: Arc = Arc::clone(&pool) as _; + // Like DataFusion 55's final hash aggregate, whose spill merge and replay table are + // sibling reservations of one consumer. + let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&dyn_pool); + let replay = merge.new_empty(); + merge.try_grow(90).unwrap(); + + // Spark grants 10 of the replay's 30 bytes, and the other 20 are overcommit. + replay.try_grow(30).unwrap(); + assert_eq!(pool.reserved(), 120); + assert_eq!(fake.held(), 100); + assert_eq!(pool.overcommit(), 20); + + // The replay emits groups and shrinks, which repays the overcommit before Spark. + replay.shrink(25); + assert_eq!(pool.overcommit(), 0); + assert_eq!(fake.held(), 95); + drop(merge); + drop(replay); + assert_eq!(pool.reserved(), 0); + assert_eq!(fake.held(), 0); + assert!(pool.final_aggregates.lock().is_empty()); + } + + #[test] + fn other_refusals_are_unchanged() { + let fake = FakeSpark::with(100); + let pool: Arc = Arc::new(CometUnifiedMemoryPool::with_spark(fake.memory())); + + // While the aggregate reads its input, its table is the consumer's only reservation + // holding memory, so a refusal makes it spill. + let table = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); + let replay = table.new_empty(); + table.try_grow(90).unwrap(); + assert!(table.try_grow(30).is_err()); + drop(table); + + // The merge picks its spill files while nothing else is held, so a refusal still limits + // how many it opens. + let merge = replay.new_empty(); + merge.try_grow(90).unwrap(); + assert!(merge.try_grow(30).is_err()); + drop(merge); + + // Any other operator with a sibling holding memory is still refused. + let sort = MemoryConsumer::new("ExternalSorterMerge[0]").register(&pool); + let sibling = sort.new_empty(); + sort.try_grow(90).unwrap(); + assert!(sibling.try_grow(30).is_err()); + assert_eq!(pool.reserved(), 90); + assert_eq!(fake.held(), 90); + } } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala index ed16cbb1d4d..f4c3750b1d2 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -26,7 +26,7 @@ import scala.util.Random import org.apache.hadoop.fs.Path import org.apache.spark.{CometListenerBusUtils, SparkConf} import org.apache.spark.scheduler.{SparkListener, SparkListenerTaskEnd} -import org.apache.spark.sql.{CometTestBase, DataFrame, Row} +import org.apache.spark.sql.{CometTestBase, DataFrame, QueryTest, Row} import org.apache.spark.sql.catalyst.expressions.Cast import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, PartialMerge} import org.apache.spark.sql.catalyst.optimizer.EliminateSorts @@ -154,6 +154,42 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("final aggregate that has spilled reads its spill files back (issue #6254)") { + withTempPath { dir => + val path = dir.getCanonicalPath + spark + .range(0, 125000, 1, 1) + .selectExpr("lpad(cast(id % 62500 AS STRING), 128, 'k') AS k", "id AS v") + .write + .parquet(path) + // A 3 MiB pool (3/2048 of the suite's 2g off-heap) makes the final aggregate spill several + // runs of its 62,500 groups. Merging the runs then takes nearly all of its share, and + // reading them back must not fail the task. + withSQLConf( + SQLConf.SHUFFLE_PARTITIONS.key -> "1", + CometConf.COMET_BATCH_SIZE.key -> "1024", + CometConf.COMET_OFFHEAP_MEMORY_POOL_FRACTION.key -> (3.0 / 2048).toString) { + val df = spark.read.parquet(path).groupBy("k").agg(expr("count(1)"), expr("max(v)")) + val expected = (0 until 62500).map { i => + val key = i.toString + Row("k" * (128 - key.length) + key, 2L, i + 62500L) + } + // checkToRDD = false runs the query once, so the plan read below is the one that ran. + QueryTest.checkAnswer(df, expected, checkToRDD = false) + + val finalAggregates = collect(df.queryExecution.executedPlan) { + case aggregate: CometHashAggregateExec if aggregate.modes.contains(Final) => aggregate + } + assert( + finalAggregates.nonEmpty, + s"Expected a native final aggregate:\n${df.queryExecution.executedPlan}") + assert( + finalAggregates.forall(_.metrics("spill_count").value > 0), + "The final aggregate did not spill") + } + } + } + test("collect_list over struct with non-nullable fields") { // Building a struct from non-nullable columns yields non-nullable struct fields. Native // collect_list derives its declared output element type from the argument's declared type, From 20782d356af3293538889379217c6698a4891ce9 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 3 Oct 2026 13:08:13 -0600 Subject: [PATCH 2/5] fix: let an ordered final aggregate read its spill files back too DataFusion runs a final aggregate whose input is sorted on some of its grouping keys as an OrderedFinalAggregateStream. It merges and replays its spill files the same way FinalHashAggregateStream does, so its replay failed the task the same way. Treat its consumer as a final aggregate too. Explain in spill_replay.rs why recording the replay's request is safe: the replay asks for memory only after it has aggregated a batch, so the memory already exists, as it does for a grow. Describe the exception in the memory management guide, which said the fair pool always refuses a request that fails its local checks. --- .../contributor-guide/memory_management.md | 26 +++++++- .../src/execution/memory_pools/fair_pool.rs | 5 +- .../execution/memory_pools/spill_replay.rs | 66 ++++++++++++++----- .../execution/memory_pools/unified_pool.rs | 14 ++-- .../comet/exec/CometAggregateSuite.scala | 53 ++++++++++++++- 5 files changed, 134 insertions(+), 30 deletions(-) diff --git a/docs/source/contributor-guide/memory_management.md b/docs/source/contributor-guide/memory_management.md index d899d02029d..4f766d25c14 100644 --- a/docs/source/contributor-guide/memory_management.md +++ b/docs/source/contributor-guide/memory_management.md @@ -286,7 +286,8 @@ JNI, which goes through Spark's ordinary `TaskMemoryManager`. That means: `Display` output and their `try_grow` errors report the current overcommit. `CometFairMemoryPool` additionally applies two local checks before it asks Spark, and refuses the -request without calling Spark if either fails: +request without calling Spark if either fails, except for the request described in +[Final aggregates reading their spill files back](#final-aggregates-reading-their-spill-files-back): - **The requesting consumer against its share.** The share is `pool_size` divided by the number of consumers currently registered with the pool. What the consumer already holds plus the request @@ -314,6 +315,29 @@ This is why `fair_unified` can spill earlier than `greedy_unified`: a consumer a refused even when the rest of the pool is free, which keeps that memory for the task's other consumers. +### Final aggregates reading their spill files back + +Both pools make one exception to refusing a `try_grow` (`spill_replay.rs`). Once one of +DataFusion 55's final aggregates has spilled, it merges its spill files and replays them through +an aggregate that cannot spill, so a refused request during the replay fails the task. +`FinalHashAggregateStream` does this, and so does `OrderedFinalAggregateStream`, which DataFusion +uses when the input is sorted on some of the grouping keys. The merge reserves read buffers for as +many spill files as fit, in a sibling reservation of the same consumer, so the replay often finds +the consumer's share already taken. The replay asks for memory only after it has aggregated a +batch, so the memory already exists. The pools therefore record its request the way they record a +`grow`, past the share and the pool's total, and carry what Spark does not grant as overcommit. + +A pool treats a request as part of a replay when it comes from one of these consumers while +another of the consumer's reservations holds memory. In DataFusion 55.1 that happens only while the +replay grows and the merge holds its read buffers. Before the replay, the aggregate's table is its +only reservation holding memory, so a refusal still makes it spill. The merge picks its files while +nothing else is held, so a refusal still limits how many it opens. + +This works around [issue #6254](https://github.com/apache/datafusion-comet/issues/6254) until +Comet's DataFusion includes +[apache/datafusion#25383](https://github.com/apache/datafusion/pull/25383), which leaves the replay +room when the merge picks its files. + ### Task-shared pools and their lifetime A single Spark task can run more than one native plan at a time. A native shuffle is not one of diff --git a/native/core/src/execution/memory_pools/fair_pool.rs b/native/core/src/execution/memory_pools/fair_pool.rs index f13262d6aee..7517699e3f4 100644 --- a/native/core/src/execution/memory_pools/fair_pool.rs +++ b/native/core/src/execution/memory_pools/fair_pool.rs @@ -100,9 +100,8 @@ impl CometFairMemoryPool { *consumer_used = consumer_used.saturating_add(additional); } - /// Refuses a `try_grow` with `err`, unless it comes from a final hash aggregate reading its - /// spill files back, which can't spill. That request is recorded instead; see - /// [`spill_replay`]. + /// Refuses a `try_grow` with `err`, unless it comes from a final aggregate reading its spill + /// files back, which can't spill. That request is recorded instead; see [`spill_replay`]. fn refuse( &self, state: &mut CometFairPoolState, diff --git a/native/core/src/execution/memory_pools/spill_replay.rs b/native/core/src/execution/memory_pools/spill_replay.rs index d33426fd189..1ac94187a76 100644 --- a/native/core/src/execution/memory_pools/spill_replay.rs +++ b/native/core/src/execution/memory_pools/spill_replay.rs @@ -15,33 +15,63 @@ // specific language governing permissions and limitations // under the License. -//! Lets a final hash aggregate read its spill files back past its share of the pool. +//! Lets a final aggregate read its spill files back past its share of memory. //! -//! Once DataFusion 55's `FinalHashAggregateStream` has spilled, it merges its sorted spill files -//! and replays them through an `OrderedFinalAggregateStream` that has no way to spill, so a -//! refused memory request there fails the task (#6254). The merge reserves read buffers for as -//! many spill files as fit, and those buffers belong to the same consumer, so the replay often -//! finds the consumer's share already taken. The pools record such a request as overcommit -//! instead of refusing it. The replay emits every finished group after each batch, so it holds -//! about one batch of groups, and releasing memory repays the overcommit first. +//! Once one of DataFusion 55's final aggregates has spilled, it merges its sorted spill files and +//! replays them through an `OrderedFinalAggregateStream` that has no way to spill, so a refused +//! memory request there fails the task (#6254). `FinalHashAggregateStream` does this, and so does +//! `OrderedFinalAggregateStream` itself, which DataFusion uses when the input is sorted on some of +//! the grouping keys. The merge reserves read buffers for as many spill files as fit, and those +//! buffers belong to the same consumer, so the replay often finds the consumer's share already +//! taken. The replay only asks for memory once it has aggregated a batch, so like a `grow`, the +//! request is for memory that already exists. The pools record it the way they record a `grow`, +//! carrying what Spark doesn't grant as overcommit. The replay emits every finished group after +//! each batch, so it holds about one batch of groups, and releasing memory repays the overcommit +//! first. //! //! Remove this once Comet's DataFusion has apache/datafusion#25383, which leaves the replay room //! when the merge picks its files. use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation}; -/// Whether `consumer` belongs to a `FinalHashAggregateStream`, whose spill replay can't spill. -pub(super) fn is_final_hash_aggregate(consumer: &MemoryConsumer) -> bool { - consumer.name().starts_with("FinalHashAggregateStream[") +/// Whether `consumer` belongs to a final aggregate whose spill replay can't spill. +pub(super) fn is_final_aggregate(consumer: &MemoryConsumer) -> bool { + let name = consumer.name(); + name.starts_with("FinalHashAggregateStream[") + || name.starts_with("OrderedFinalAggregateStream[") } -/// Whether a refused request from `reservation` should be recorded as overcommit instead. -/// `consumer_used` is what the reservation's consumer holds across all of its reservations. +/// Whether a refused request from `reservation` should be recorded instead. `consumer_used` is +/// what the reservation's consumer holds across all of its reservations. /// -/// A final hash aggregate grows one reservation while another of its reservations holds memory -/// only during the replay, while the merge holds its read buffers. Before that, the aggregate's -/// table is its only reservation holding memory, so a refusal makes it spill. The merge picks its -/// files while nothing else is held, so a refusal there still limits how many it opens. +/// A final aggregate grows one reservation while another of its reservations holds memory only +/// during the replay, while the merge holds its read buffers. Before that, the aggregate's table +/// is its only reservation holding memory, so a refusal stands and makes it spill. The merge picks +/// its files while nothing else is held, so a refusal there still limits how many it opens. pub(super) fn is_spill_replay(reservation: &MemoryReservation, consumer_used: usize) -> bool { - consumer_used > reservation.size() && is_final_hash_aggregate(reservation.consumer()) + consumer_used > reservation.size() && is_final_aggregate(reservation.consumer()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn only_final_aggregates_replay_their_spill_files() { + for name in [ + "FinalHashAggregateStream[3]", + "OrderedFinalAggregateStream[0]", + ] { + assert!(is_final_aggregate(&MemoryConsumer::new(name)), "{name}"); + } + // Partial aggregates emit early instead of spilling, and Comet never plans single ones. + for name in [ + "PartialHashAggregateStream[0]", + "OrderedPartialAggregateStream[0]", + "SingleHashAggregateStream[0]", + "ExternalSorterMerge[0]", + ] { + assert!(!is_final_aggregate(&MemoryConsumer::new(name)), "{name}"); + } + } } diff --git a/native/core/src/execution/memory_pools/unified_pool.rs b/native/core/src/execution/memory_pools/unified_pool.rs index adf384e9970..f85a34aae64 100644 --- a/native/core/src/execution/memory_pools/unified_pool.rs +++ b/native/core/src/execution/memory_pools/unified_pool.rs @@ -35,7 +35,7 @@ use parking_lot::Mutex; pub struct CometUnifiedMemoryPool { spark: SparkMemory, used: AtomicUsize, - /// Bytes held by each final hash aggregate's consumer across its reservations, keyed by + /// Bytes held by each final aggregate's consumer across its reservations, keyed by /// [`MemoryConsumer::id`], which [`spill_replay`] needs. Other consumers aren't tracked, so /// they never take the lock. final_aggregates: Mutex>, @@ -64,9 +64,9 @@ impl CometUnifiedMemoryPool { self.spark.overcommit() } - /// Applies `update` to what `reservation`'s consumer holds, if it is a final hash aggregate. + /// Applies `update` to what `reservation`'s consumer holds, if it is a final aggregate. fn track(&self, reservation: &MemoryReservation, update: impl FnOnce(&mut usize)) { - if spill_replay::is_final_hash_aggregate(reservation.consumer()) { + if spill_replay::is_final_aggregate(reservation.consumer()) { if let Some(used) = self .final_aggregates .lock() @@ -77,8 +77,8 @@ impl CometUnifiedMemoryPool { } } - /// Whether a refused request from `reservation` comes from a final hash aggregate reading its - /// spill files back, which can't spill; see [`spill_replay`]. + /// Whether a refused request from `reservation` comes from a final aggregate reading its spill + /// files back, which can't spill; see [`spill_replay`]. fn is_spill_replay(&self, reservation: &MemoryReservation) -> bool { let consumer_used = self .final_aggregates @@ -119,13 +119,13 @@ impl MemoryPool for CometUnifiedMemoryPool { } fn register(&self, consumer: &MemoryConsumer) { - if spill_replay::is_final_hash_aggregate(consumer) { + if spill_replay::is_final_aggregate(consumer) { self.final_aggregates.lock().insert(consumer.id(), 0); } } fn unregister(&self, consumer: &MemoryConsumer) { - if spill_replay::is_final_hash_aggregate(consumer) { + if spill_replay::is_final_aggregate(consumer) { self.final_aggregates.lock().remove(&consumer.id()); } } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala index f4c3750b1d2..e80772b7300 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -31,7 +31,7 @@ import org.apache.spark.sql.catalyst.expressions.Cast import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, PartialMerge} import org.apache.spark.sql.catalyst.optimizer.EliminateSorts import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, RangePartitioning} -import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, CometNativeExec, CometProjectExec} +import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, CometNativeExec, CometProjectExec, CometSortExec} import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec import org.apache.spark.sql.execution.SQLExecution import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, ShuffleQueryStageExec} @@ -190,6 +190,57 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("ordered final aggregate that has spilled reads its spill files back (issue #6254)") { + withTempPath { dir => + val path = dir.getCanonicalPath + spark + .range(0, 125000, 1, 1) + .selectExpr("0L AS g", "lpad(cast(id % 62500 AS STRING), 128, 'k') AS k", "id AS v") + .write + .parquet(path) + // Sorting on g, the first grouping key, makes DataFusion run the final aggregate as an + // OrderedFinalAggregateStream, which spills and reads its spill files back the same way. + // Spark drops a sort under an aggregate unless EliminateSorts is excluded. With a 2.5 MiB + // pool, merging the runs takes nearly all of the aggregate's share. The sort's merge + // reservation is lowered so that the sort fits in the pool too. + withSQLConf( + SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> EliminateSorts.ruleName, + CometConf.COMET_BATCH_SIZE.key -> "1024", + CometConf.COMET_RESPECT_DATAFUSION_CONFIGS.key -> "true", + "spark.comet.datafusion.execution.sort_spill_reservation_bytes" -> "65536", + CometConf.COMET_OFFHEAP_MEMORY_POOL_FRACTION.key -> (2.5 / 2048).toString) { + val df = spark.read + .parquet(path) + .repartition(1, col("g")) + .sortWithinPartitions("g") + .groupBy("g", "k") + .agg(expr("count(1)"), expr("max(v)")) + val expected = (0 until 62500).map { i => + val key = i.toString + Row(0L, "k" * (128 - key.length) + key, 2L, i + 62500L) + } + // checkAnswer would compare the rows in order, because the plan has a sort. Collecting runs + // the query once, so the plan read below is the one that ran. + QueryTest.sameRows(expected, df.collect().toSeq).foreach(fail(_)) + + val plan = df.queryExecution.executedPlan + val finalAggregates = collect(plan) { + case aggregate: CometHashAggregateExec if aggregate.modes.contains(Final) => aggregate + } + // Only a sort in the same native plan keeps the final aggregate's input sorted. + assert( + finalAggregates.exists(_.child match { + case partial: CometHashAggregateExec => partial.child.isInstanceOf[CometSortExec] + case _ => false + }), + s"Expected a native final aggregate over a partial aggregate over a sort:\n$plan") + assert( + finalAggregates.forall(_.metrics("spill_count").value > 0), + "The final aggregate did not spill") + } + } + } + test("collect_list over struct with non-nullable fields") { // Building a struct from non-nullable columns yields non-nullable struct fields. Native // collect_list derives its declared output element type from the argument's declared type, From ee985caeb2679a61c5cc0971fee82a612ff1aaf2 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 3 Oct 2026 13:11:56 -0600 Subject: [PATCH 3/5] docs: tighten the spill replay comments and guide The fair pool is the only one with a share and a pool total to skip, so say so. The greedy pool takes its tracking lock for other consumers when Spark refuses them, not never. Point whoever removes the workaround at the tests that show whether the replay still needs it. --- docs/source/contributor-guide/memory_management.md | 3 ++- native/core/src/execution/memory_pools/spill_replay.rs | 3 ++- native/core/src/execution/memory_pools/unified_pool.rs | 2 +- 3 files changed, 5 insertions(+), 3 deletions(-) diff --git a/docs/source/contributor-guide/memory_management.md b/docs/source/contributor-guide/memory_management.md index 4f766d25c14..4b967e15e2c 100644 --- a/docs/source/contributor-guide/memory_management.md +++ b/docs/source/contributor-guide/memory_management.md @@ -325,7 +325,8 @@ uses when the input is sorted on some of the grouping keys. The merge reserves r many spill files as fit, in a sibling reservation of the same consumer, so the replay often finds the consumer's share already taken. The replay asks for memory only after it has aggregated a batch, so the memory already exists. The pools therefore record its request the way they record a -`grow`, past the share and the pool's total, and carry what Spark does not grant as overcommit. +`grow`. `CometFairMemoryPool` skips its two local checks for it, and both pools carry what Spark +does not grant as overcommit. A pool treats a request as part of a replay when it comes from one of these consumers while another of the consumer's reservations holds memory. In DataFusion 55.1 that happens only while the diff --git a/native/core/src/execution/memory_pools/spill_replay.rs b/native/core/src/execution/memory_pools/spill_replay.rs index 1ac94187a76..2d8568f6d06 100644 --- a/native/core/src/execution/memory_pools/spill_replay.rs +++ b/native/core/src/execution/memory_pools/spill_replay.rs @@ -30,7 +30,8 @@ //! first. //! //! Remove this once Comet's DataFusion has apache/datafusion#25383, which leaves the replay room -//! when the merge picks its files. +//! when the merge picks its files. The #6254 tests in `CometAggregateSuite` fail without this on +//! DataFusion 55.1, so they show whether the replay still needs it. use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation}; diff --git a/native/core/src/execution/memory_pools/unified_pool.rs b/native/core/src/execution/memory_pools/unified_pool.rs index f85a34aae64..b714678ae96 100644 --- a/native/core/src/execution/memory_pools/unified_pool.rs +++ b/native/core/src/execution/memory_pools/unified_pool.rs @@ -37,7 +37,7 @@ pub struct CometUnifiedMemoryPool { used: AtomicUsize, /// Bytes held by each final aggregate's consumer across its reservations, keyed by /// [`MemoryConsumer::id`], which [`spill_replay`] needs. Other consumers aren't tracked, so - /// they never take the lock. + /// they take the lock only when Spark refuses them. final_aggregates: Mutex>, } From 52bae81391beb4f37ce2163c1f973e6901bf9121 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 3 Oct 2026 15:23:39 -0600 Subject: [PATCH 4/5] test: share the spill replay tests between the pools fair_pool.rs and unified_pool.rs each had a copy of the same two spill replay scenarios. Run them once, in spill_replay.rs, against both pool types as createPlan builds them, which also checks that the wrappers pass register through to the greedy pool's tracking. Keep a greedy pool test for its per-consumer map, and drop a test import that the pool's own imports now cover. Share the final aggregate spill checks between the two #6254 tests in CometAggregateSuite, and use checkCometAnswer, which collects once and labels the Comet answer as Comet's. --- .../src/execution/memory_pools/fair_pool.rs | 57 ------------- .../execution/memory_pools/spill_replay.rs | 80 +++++++++++++++++++ .../execution/memory_pools/unified_pool.rs | 55 ++----------- .../comet/exec/CometAggregateSuite.scala | 47 +++++------ 4 files changed, 112 insertions(+), 127 deletions(-) diff --git a/native/core/src/execution/memory_pools/fair_pool.rs b/native/core/src/execution/memory_pools/fair_pool.rs index 7517699e3f4..b61af1ece80 100644 --- a/native/core/src/execution/memory_pools/fair_pool.rs +++ b/native/core/src/execution/memory_pools/fair_pool.rs @@ -364,33 +364,6 @@ mod tests { assert_eq!(fake.held(), 50); } - #[test] - fn a_final_aggregate_reading_its_spill_files_back_carries_what_spark_refuses() { - let fake = FakeSpark::with(100); - let pool = Arc::new(CometFairMemoryPool::with_spark(fake.memory(), 1000)); - let dyn_pool: Arc = Arc::clone(&pool) as _; - // Like DataFusion 55's final hash aggregate, whose spill merge and replay table are - // sibling reservations of one consumer. - let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&dyn_pool); - let replay = merge.new_empty(); - merge.try_grow(90).unwrap(); - - // Spark grants 10 of the replay's 30 bytes, and the other 20 are overcommit. - replay.try_grow(30).unwrap(); - assert_eq!(pool.reserved(), 120); - assert_eq!(fake.held(), 100); - assert_eq!(pool.overcommit(), 20); - - // The replay emits groups and shrinks, which repays the overcommit before Spark. - replay.shrink(25); - assert_eq!(pool.overcommit(), 0); - assert_eq!(fake.held(), 95); - drop(merge); - drop(replay); - assert_eq!(pool.reserved(), 0); - assert_eq!(fake.held(), 0); - } - #[test] fn a_final_aggregate_reading_its_spill_files_back_may_pass_its_fair_limit() { // Spark grants everything, so only the pool's own checks refuse. @@ -406,34 +379,4 @@ mod tests { assert_eq!(fake.held(), 120); assert_eq!(pool.overcommit(), 0); } - - #[test] - fn other_refusals_are_unchanged() { - let fake = FakeSpark::with(100); - let pool: Arc = - Arc::new(CometFairMemoryPool::with_spark(fake.memory(), 1000)); - - // While the aggregate reads its input, its table is the consumer's only reservation - // holding memory, so a refusal makes it spill. - let table = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); - let replay = table.new_empty(); - table.try_grow(90).unwrap(); - assert!(table.try_grow(30).is_err()); - drop(table); - - // The merge picks its spill files while nothing else is held, so a refusal still limits - // how many it opens. - let merge = replay.new_empty(); - merge.try_grow(90).unwrap(); - assert!(merge.try_grow(30).is_err()); - drop(merge); - - // Any other operator with a sibling holding memory is still refused. - let sort = MemoryConsumer::new("ExternalSorterMerge[0]").register(&pool); - let sibling = sort.new_empty(); - sort.try_grow(90).unwrap(); - assert!(sibling.try_grow(30).is_err()); - assert_eq!(pool.reserved(), 90); - assert_eq!(fake.held(), 90); - } } diff --git a/native/core/src/execution/memory_pools/spill_replay.rs b/native/core/src/execution/memory_pools/spill_replay.rs index 2d8568f6d06..4b0de9e3fe8 100644 --- a/native/core/src/execution/memory_pools/spill_replay.rs +++ b/native/core/src/execution/memory_pools/spill_replay.rs @@ -55,7 +55,87 @@ pub(super) fn is_spill_replay(reservation: &MemoryReservation, consumer_used: us #[cfg(test)] mod tests { + use super::super::spark_memory::fake::FakeSpark; + use super::super::{create_pool, overcommit, MemoryPoolConfig, MemoryPoolType}; use super::*; + use datafusion::execution::memory_pool::MemoryPool; + use std::sync::Arc; + + /// A pool of each type, built the way `createPlan` builds it and connected to a fake Spark + /// that grants at most 100 bytes. The fair pool's own limits are far above that, so only + /// Spark refuses. Task-shared pools are keyed by task attempt process-wide, so each pool needs + /// its own id. + fn each_pool_type( + task_attempt_ids: [i64; 2], + ) -> Vec<(&'static str, Arc, Arc)> { + [ + ("greedy_unified", MemoryPoolType::GreedyUnified), + ("fair_unified", MemoryPoolType::FairUnified), + ] + .into_iter() + .zip(task_attempt_ids) + .map(|((name, pool_type), task_attempt_id)| { + let fake = FakeSpark::with(100); + let config = MemoryPoolConfig::new(pool_type, 1000); + let pool = create_pool(&config, task_attempt_id, || fake.memory()); + (name, pool, fake) + }) + .collect() + } + + #[test] + fn a_final_aggregate_reading_its_spill_files_back_carries_what_spark_refuses() { + for (name, pool, fake) in each_pool_type([-3011, -3012]) { + // Like DataFusion 55's final hash aggregate, whose spill merge and replay table are + // sibling reservations of one consumer. + let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); + let replay = merge.new_empty(); + merge.try_grow(90).unwrap(); + + // Spark grants 10 of the replay's 30 bytes, and the other 20 are overcommit. + replay.try_grow(30).unwrap(); + assert_eq!(pool.reserved(), 120, "{name}"); + assert_eq!(fake.held(), 100, "{name}"); + assert_eq!(overcommit(&pool), 20, "{name}"); + + // The replay emits groups and shrinks, which repays the overcommit before Spark. + replay.shrink(25); + assert_eq!(overcommit(&pool), 0, "{name}"); + assert_eq!(fake.held(), 95, "{name}"); + drop(merge); + drop(replay); + assert_eq!(pool.reserved(), 0, "{name}"); + assert_eq!(fake.held(), 0, "{name}"); + } + } + + #[test] + fn other_refusals_are_unchanged() { + for (name, pool, fake) in each_pool_type([-3013, -3014]) { + // While the aggregate reads its input, its table is the consumer's only reservation + // holding memory, so a refusal makes it spill. + let table = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); + let replay = table.new_empty(); + table.try_grow(90).unwrap(); + assert!(table.try_grow(30).is_err(), "{name}"); + drop(table); + + // The merge picks its spill files while nothing else is held, so a refusal still + // limits how many it opens. + let merge = replay.new_empty(); + merge.try_grow(90).unwrap(); + assert!(merge.try_grow(30).is_err(), "{name}"); + drop(merge); + + // Any other operator with a sibling holding memory is still refused. + let sort = MemoryConsumer::new("ExternalSorterMerge[0]").register(&pool); + let sibling = sort.new_empty(); + sort.try_grow(90).unwrap(); + assert!(sibling.try_grow(30).is_err(), "{name}"); + assert_eq!(pool.reserved(), 90, "{name}"); + assert_eq!(fake.held(), 90, "{name}"); + } + } #[test] fn only_final_aggregates_replay_their_spill_files() { diff --git a/native/core/src/execution/memory_pools/unified_pool.rs b/native/core/src/execution/memory_pools/unified_pool.rs index b714678ae96..8b89b1aec1b 100644 --- a/native/core/src/execution/memory_pools/unified_pool.rs +++ b/native/core/src/execution/memory_pools/unified_pool.rs @@ -222,7 +222,6 @@ impl MemoryPool for CometUnifiedMemoryPool { mod tests { use super::super::spark_memory::fake::FakeSpark; use super::*; - use datafusion::execution::memory_pool::MemoryConsumer; use std::sync::Arc; #[test] @@ -294,59 +293,21 @@ mod tests { } #[test] - fn a_final_aggregate_reading_its_spill_files_back_carries_what_spark_refuses() { - let fake = FakeSpark::with(100); - let pool = Arc::new(CometUnifiedMemoryPool::with_spark(fake.memory())); + fn a_final_aggregate_is_tracked_across_its_reservations_until_it_unregisters() { + let pool = Arc::new(CometUnifiedMemoryPool::with_spark( + FakeSpark::with(100).memory(), + )); let dyn_pool: Arc = Arc::clone(&pool) as _; - // Like DataFusion 55's final hash aggregate, whose spill merge and replay table are - // sibling reservations of one consumer. let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&dyn_pool); let replay = merge.new_empty(); merge.try_grow(90).unwrap(); + replay.grow(20); + replay.shrink(5); + let id = merge.consumer().id(); + assert_eq!(pool.final_aggregates.lock().get(&id), Some(&105)); - // Spark grants 10 of the replay's 30 bytes, and the other 20 are overcommit. - replay.try_grow(30).unwrap(); - assert_eq!(pool.reserved(), 120); - assert_eq!(fake.held(), 100); - assert_eq!(pool.overcommit(), 20); - - // The replay emits groups and shrinks, which repays the overcommit before Spark. - replay.shrink(25); - assert_eq!(pool.overcommit(), 0); - assert_eq!(fake.held(), 95); drop(merge); drop(replay); - assert_eq!(pool.reserved(), 0); - assert_eq!(fake.held(), 0); assert!(pool.final_aggregates.lock().is_empty()); } - - #[test] - fn other_refusals_are_unchanged() { - let fake = FakeSpark::with(100); - let pool: Arc = Arc::new(CometUnifiedMemoryPool::with_spark(fake.memory())); - - // While the aggregate reads its input, its table is the consumer's only reservation - // holding memory, so a refusal makes it spill. - let table = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); - let replay = table.new_empty(); - table.try_grow(90).unwrap(); - assert!(table.try_grow(30).is_err()); - drop(table); - - // The merge picks its spill files while nothing else is held, so a refusal still limits - // how many it opens. - let merge = replay.new_empty(); - merge.try_grow(90).unwrap(); - assert!(merge.try_grow(30).is_err()); - drop(merge); - - // Any other operator with a sibling holding memory is still refused. - let sort = MemoryConsumer::new("ExternalSorterMerge[0]").register(&pool); - let sibling = sort.new_empty(); - sort.try_grow(90).unwrap(); - assert!(sibling.try_grow(30).is_err()); - assert_eq!(pool.reserved(), 90); - assert_eq!(fake.held(), 90); - } } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala index e80772b7300..b9a597197a0 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -154,6 +154,22 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + /** + * Returns the native final aggregates in the plan that `df` last ran, after checking that there + * is one and that each spilled. + */ + private def spilledFinalAggregates(df: DataFrame): Seq[CometHashAggregateExec] = { + val plan = df.queryExecution.executedPlan + val finalAggregates = collect(plan) { + case aggregate: CometHashAggregateExec if aggregate.modes.contains(Final) => aggregate + } + assert(finalAggregates.nonEmpty, s"Expected a native final aggregate:\n$plan") + assert( + finalAggregates.forall(_.metrics("spill_count").value > 0), + "The final aggregate did not spill") + finalAggregates + } + test("final aggregate that has spilled reads its spill files back (issue #6254)") { withTempPath { dir => val path = dir.getCanonicalPath @@ -174,18 +190,9 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { val key = i.toString Row("k" * (128 - key.length) + key, 2L, i + 62500L) } - // checkToRDD = false runs the query once, so the plan read below is the one that ran. - QueryTest.checkAnswer(df, expected, checkToRDD = false) - - val finalAggregates = collect(df.queryExecution.executedPlan) { - case aggregate: CometHashAggregateExec if aggregate.modes.contains(Final) => aggregate - } - assert( - finalAggregates.nonEmpty, - s"Expected a native final aggregate:\n${df.queryExecution.executedPlan}") - assert( - finalAggregates.forall(_.metrics("spill_count").value > 0), - "The final aggregate did not spill") + // checkCometAnswer runs the query once, so the plan read below is the one that ran. + checkCometAnswer(df, expected) + spilledFinalAggregates(df) } } } @@ -219,24 +226,18 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { val key = i.toString Row(0L, "k" * (128 - key.length) + key, 2L, i + 62500L) } - // checkAnswer would compare the rows in order, because the plan has a sort. Collecting runs - // the query once, so the plan read below is the one that ran. + // checkCometAnswer would compare the rows in order, because the plan has a sort. + // Collecting runs the query once, so the plan read below is the one that ran. QueryTest.sameRows(expected, df.collect().toSeq).foreach(fail(_)) - val plan = df.queryExecution.executedPlan - val finalAggregates = collect(plan) { - case aggregate: CometHashAggregateExec if aggregate.modes.contains(Final) => aggregate - } // Only a sort in the same native plan keeps the final aggregate's input sorted. assert( - finalAggregates.exists(_.child match { + spilledFinalAggregates(df).exists(_.child match { case partial: CometHashAggregateExec => partial.child.isInstanceOf[CometSortExec] case _ => false }), - s"Expected a native final aggregate over a partial aggregate over a sort:\n$plan") - assert( - finalAggregates.forall(_.metrics("spill_count").value > 0), - "The final aggregate did not spill") + "Expected a native final aggregate over a partial aggregate over a sort:\n" + + df.queryExecution.executedPlan) } } } From cc60fe783dd12cd44729ec35b5a06f9fee2467cd Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Mon, 5 Oct 2026 06:55:41 -0600 Subject: [PATCH 5/5] refactor: move the spill replay workaround into a pool wrapper The workaround for #6254 lived in both Comet pools. The fair pool sent its three refusals through a check, and the greedy pool tracked each final aggregate across its reservations. #5613 reworks the fair pool's try_grow, and #6583 has to take the workaround out again, so both would have had to rework the pools. SpillReplayPool, in spill_replay.rs, now wraps the tracked Comet pool instead. It keeps a total for each final aggregate's consumer. When the pool refuses a request from one while another of its reservations holds memory, it calls the pool's grow, which skips the fair pool's limits and carries what Spark doesn't grant as overcommit. fair_pool.rs and unified_pool.rs are back to main's versions, and overcommit() looks through the wrapper. The Rust tests run against both pool types as createPlan builds them. They now reach the fair pool's pool-limit refusal too, and check that a failed Spark call during the replay is returned rather than recorded. The guide adds the wrapper to the pool stack and keeps a short section. The DataFusion 55.1 facts that an upgrade has to re-check stay in spill_replay.rs, which names #6583 next to the removal condition. --- .../contributor-guide/memory_management.md | 52 ++-- .../src/execution/memory_pools/fair_pool.rs | 71 +---- native/core/src/execution/memory_pools/mod.rs | 12 +- .../execution/memory_pools/spill_replay.rs | 263 ++++++++++++++++-- .../execution/memory_pools/unified_pool.rs | 97 +------ 5 files changed, 287 insertions(+), 208 deletions(-) diff --git a/docs/source/contributor-guide/memory_management.md b/docs/source/contributor-guide/memory_management.md index 4b967e15e2c..2a48baa3f18 100644 --- a/docs/source/contributor-guide/memory_management.md +++ b/docs/source/contributor-guide/memory_management.md @@ -253,14 +253,17 @@ ignored and the pool is always `UnboundedMemoryPool`. the inside out, a Comet plan in the default configuration sees: ```text -[LoggingMemoryPool] <- only when spark.comet.debug.memory=true - [TaskSharedMemoryPool] <- RAII handle for the per-task registry - [TrackConsumersPool] <- DataFusion; names the top 10 consumers in error messages - [CometFairMemoryPool] <- delegates acquire/release to Spark over JNI +[LoggingMemoryPool] <- only when spark.comet.debug.memory=true + [TaskSharedMemoryPool] <- RAII handle for the per-task registry + [SpillReplayPool] <- lets a spilled final aggregate read its spill files back + [TrackConsumersPool] <- DataFusion; names the top 10 consumers in error messages + [CometFairMemoryPool] <- delegates acquire/release to Spark over JNI ``` Each decorator forwards every `MemoryPool` method to its inner pool, so `reserved()` at any level -reports the base pool's number. +reports the base pool's number. `SpillReplayPool` also turns one kind of refused `try_grow` into a +`grow`; see +[Final aggregates reading their spill files back](#final-aggregates-reading-their-spill-files-back). ### The unified pools @@ -286,8 +289,7 @@ JNI, which goes through Spark's ordinary `TaskMemoryManager`. That means: `Display` output and their `try_grow` errors report the current overcommit. `CometFairMemoryPool` additionally applies two local checks before it asks Spark, and refuses the -request without calling Spark if either fails, except for the request described in -[Final aggregates reading their spill files back](#final-aggregates-reading-their-spill-files-back): +request without calling Spark if either fails: - **The requesting consumer against its share.** The share is `pool_size` divided by the number of consumers currently registered with the pool. What the consumer already holds plus the request @@ -315,30 +317,6 @@ This is why `fair_unified` can spill earlier than `greedy_unified`: a consumer a refused even when the rest of the pool is free, which keeps that memory for the task's other consumers. -### Final aggregates reading their spill files back - -Both pools make one exception to refusing a `try_grow` (`spill_replay.rs`). Once one of -DataFusion 55's final aggregates has spilled, it merges its spill files and replays them through -an aggregate that cannot spill, so a refused request during the replay fails the task. -`FinalHashAggregateStream` does this, and so does `OrderedFinalAggregateStream`, which DataFusion -uses when the input is sorted on some of the grouping keys. The merge reserves read buffers for as -many spill files as fit, in a sibling reservation of the same consumer, so the replay often finds -the consumer's share already taken. The replay asks for memory only after it has aggregated a -batch, so the memory already exists. The pools therefore record its request the way they record a -`grow`. `CometFairMemoryPool` skips its two local checks for it, and both pools carry what Spark -does not grant as overcommit. - -A pool treats a request as part of a replay when it comes from one of these consumers while -another of the consumer's reservations holds memory. In DataFusion 55.1 that happens only while the -replay grows and the merge holds its read buffers. Before the replay, the aggregate's table is its -only reservation holding memory, so a refusal still makes it spill. The merge picks its files while -nothing else is held, so a refusal still limits how many it opens. - -This works around [issue #6254](https://github.com/apache/datafusion-comet/issues/6254) until -Comet's DataFusion includes -[apache/datafusion#25383](https://github.com/apache/datafusion/pull/25383), which leaves the replay -room when the merge picks its files. - ### Task-shared pools and their lifetime A single Spark task can run more than one native plan at a time. A native shuffle is not one of @@ -371,6 +349,18 @@ hands every native plan in a task the same one and drops it when the task comple covers the whole task, so `CometExecIterator.close()` warns about memory still in use only when the task's last open native plan closes. +### Final aggregates reading their spill files back + +`SpillReplayPool` (`spill_replay.rs`) wraps both Comet pools to work around +[issue #6254](https://github.com/apache/datafusion-comet/issues/6254). Once one of DataFusion 55's +final aggregates has spilled, it reads its spill files back through an aggregate that cannot spill, +so a refused `try_grow` there fails the task. When the pool refuses such a request, +`SpillReplayPool` records it with the pool's `grow` instead, which skips `CometFairMemoryPool`'s +local checks and carries what Spark does not grant as overcommit. Every other refusal is passed on +unchanged. `spill_replay.rs` describes how the wrapper recognizes these requests and what that relies +on in DataFusion. [Issue #6583](https://github.com/apache/datafusion-comet/issues/6583) tracks +removing it. + ## How DataFusion consumes the pool Native operators reserve through DataFusion's `MemoryConsumer` / `MemoryReservation` API: diff --git a/native/core/src/execution/memory_pools/fair_pool.rs b/native/core/src/execution/memory_pools/fair_pool.rs index b61af1ece80..56b09da7365 100644 --- a/native/core/src/execution/memory_pools/fair_pool.rs +++ b/native/core/src/execution/memory_pools/fair_pool.rs @@ -20,14 +20,13 @@ use std::{ fmt::{Debug, Display, Formatter, Result as FmtResult}, }; -use super::{spark_memory::SparkMemory, spill_replay}; -use datafusion::common::resources_datafusion_err; +use super::spark_memory::SparkMemory; +use datafusion::common::resources_err; use datafusion::execution::memory_pool::MemoryConsumer; use datafusion::{ common::DataFusionError, execution::memory_pool::{MemoryPool, MemoryReservation}, }; -use log::debug; use parking_lot::Mutex; /// A DataFusion fair `MemoryPool` implementation for Comet. Internally this is @@ -85,41 +84,6 @@ impl CometFairMemoryPool { pub(super) fn overcommit(&self) -> usize { self.spark.overcommit() } - - /// Records `additional` bytes for `reservation` whatever the fair and pool limits, and carries - /// what Spark doesn't grant as overcommit. See [`SparkMemory`]. - fn record( - &self, - state: &mut CometFairPoolState, - reservation: &MemoryReservation, - additional: usize, - ) { - self.spark.acquire(additional); - state.used = state.used.saturating_add(additional); - let consumer_used = state.consumer_used(reservation); - *consumer_used = consumer_used.saturating_add(additional); - } - - /// Refuses a `try_grow` with `err`, unless it comes from a final aggregate reading its spill - /// files back, which can't spill. That request is recorded instead; see [`spill_replay`]. - fn refuse( - &self, - state: &mut CometFairPoolState, - reservation: &MemoryReservation, - additional: usize, - err: DataFusionError, - ) -> Result<(), DataFusionError> { - if !spill_replay::is_spill_replay(reservation, *state.consumer_used(reservation)) { - return Err(err); - } - debug!( - "Task {} records {additional} bytes for {} while it reads its spill files back: {err}", - self.spark.task_attempt_id(), - reservation.consumer().name() - ); - self.record(state, reservation, additional); - Ok(()) - } } impl Display for CometFairMemoryPool { @@ -158,7 +122,11 @@ impl MemoryPool for CometFairMemoryPool { if additional == 0 { return; } - self.record(&mut self.state.lock(), reservation, additional); + let mut state = self.state.lock(); + self.spark.acquire(additional); + state.used = state.used.saturating_add(additional); + let consumer_used = state.consumer_used(reservation); + *consumer_used = consumer_used.saturating_add(additional); } fn shrink(&self, reservation: &MemoryReservation, subtractive: usize) { @@ -192,34 +160,31 @@ impl MemoryPool for CometFairMemoryPool { .expect("overflow in checked_div"); let consumer_used = *state.consumer_used(reservation); if limit < consumer_used.saturating_add(additional) { - let err = resources_datafusion_err!( + return resources_err!( "Failed to acquire {additional} bytes where this consumer already holds {consumer_used} bytes and the fair limit is {limit} bytes, {num} registered ({} bytes overcommitted)", self.spark.overcommit() ); - return self.refuse(&mut state, reservation, additional, err); } // The shares alone do not bound the pool's total, because a consumer keeps what it // reserved before another consumer registered. let used = state.used; if self.pool_size < used.saturating_add(additional) { - let err = resources_datafusion_err!( + return resources_err!( "Failed to acquire {additional} bytes where {used} bytes already reserved ({} bytes overcommitted) and the pool limit is {} bytes", self.spark.overcommit(), self.pool_size ); - return self.refuse(&mut state, reservation, additional, err); } // A partial grant is handed back and refused, which triggers spilling in the caller. if let Err(refusal) = self.spark.try_acquire(additional)? { - let err = resources_datafusion_err!( + return resources_err!( "Failed to acquire {} bytes plus {} bytes overcommitted, only got {} bytes. Reserved: {} bytes", additional, refusal.overcommit, refusal.granted, state.used ); - return self.refuse(&mut state, reservation, additional, err); } state.used = state .used @@ -363,20 +328,4 @@ mod tests { assert_eq!(pool.reserved(), 50); assert_eq!(fake.held(), 50); } - - #[test] - fn a_final_aggregate_reading_its_spill_files_back_may_pass_its_fair_limit() { - // Spark grants everything, so only the pool's own checks refuse. - let fake = FakeSpark::with(usize::MAX); - let pool = Arc::new(CometFairMemoryPool::with_spark(fake.memory(), 100)); - let dyn_pool: Arc = Arc::clone(&pool) as _; - let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&dyn_pool); - let replay = merge.new_empty(); - merge.try_grow(90).unwrap(); - - replay.try_grow(30).unwrap(); - assert_eq!(pool.reserved(), 120); - assert_eq!(fake.held(), 120); - assert_eq!(pool.overcommit(), 0); - } } diff --git a/native/core/src/execution/memory_pools/mod.rs b/native/core/src/execution/memory_pools/mod.rs index a8c0997081e..1fc4db2021c 100644 --- a/native/core/src/execution/memory_pools/mod.rs +++ b/native/core/src/execution/memory_pools/mod.rs @@ -28,6 +28,7 @@ use datafusion::execution::memory_pool::{MemoryPool, TrackConsumersPool, Unbound use fair_pool::CometFairMemoryPool; use jni::objects::{Global, JObject}; use spark_memory::SparkMemory; +use spill_replay::SpillReplayPool; use std::num::NonZeroUsize; use std::sync::Arc; use unified_pool::CometUnifiedMemoryPool; @@ -71,10 +72,16 @@ fn create_pool( match pool_type { MemoryPoolType::GreedyUnified => acquire_task_shared_pool(task_attempt_id, || { - tracked(CometUnifiedMemoryPool::with_spark(spark())) + Arc::new(SpillReplayPool::new( + task_attempt_id, + tracked(CometUnifiedMemoryPool::with_spark(spark())), + )) }), MemoryPoolType::FairUnified => acquire_task_shared_pool(task_attempt_id, || { - tracked(CometFairMemoryPool::with_spark(spark(), pool_size)) + Arc::new(SpillReplayPool::new( + task_attempt_id, + tracked(CometFairMemoryPool::with_spark(spark(), pool_size)), + )) }), MemoryPoolType::Unbounded => Arc::new(UnboundedMemoryPool::default()), } @@ -86,6 +93,7 @@ fn create_pool( /// takes nothing from Spark or that the function did not create. pub(crate) fn overcommit(pool: &Arc) -> usize { let pool = task_shared::unwrap_task_shared(pool).unwrap_or(pool); + let pool = spill_replay::unwrap_spill_replay(pool).unwrap_or(pool); if let Some(tracked) = pool.downcast_ref::>() { tracked.inner().overcommit() } else if let Some(tracked) = pool.downcast_ref::>() { diff --git a/native/core/src/execution/memory_pools/spill_replay.rs b/native/core/src/execution/memory_pools/spill_replay.rs index 4b0de9e3fe8..e7d947e38c1 100644 --- a/native/core/src/execution/memory_pools/spill_replay.rs +++ b/native/core/src/execution/memory_pools/spill_replay.rs @@ -15,51 +15,178 @@ // specific language governing permissions and limitations // under the License. -//! Lets a final aggregate read its spill files back past its share of memory. +//! Lets a final aggregate read its spill files back past its share of memory (#6254). //! //! Once one of DataFusion 55's final aggregates has spilled, it merges its sorted spill files and //! replays them through an `OrderedFinalAggregateStream` that has no way to spill, so a refused -//! memory request there fails the task (#6254). `FinalHashAggregateStream` does this, and so does +//! memory request there fails the task. `FinalHashAggregateStream` does this, and so does //! `OrderedFinalAggregateStream` itself, which DataFusion uses when the input is sorted on some of //! the grouping keys. The merge reserves read buffers for as many spill files as fit, and those //! buffers belong to the same consumer, so the replay often finds the consumer's share already -//! taken. The replay only asks for memory once it has aggregated a batch, so like a `grow`, the -//! request is for memory that already exists. The pools record it the way they record a `grow`, -//! carrying what Spark doesn't grant as overcommit. The replay emits every finished group after -//! each batch, so it holds about one batch of groups, and releasing memory repays the overcommit -//! first. +//! taken. //! -//! Remove this once Comet's DataFusion has apache/datafusion#25383, which leaves the replay room -//! when the merge picks its files. The #6254 tests in `CometAggregateSuite` fail without this on -//! DataFusion 55.1, so they show whether the replay still needs it. +//! [`SpillReplayPool`] wraps a Comet pool. When the pool refuses a request from the replay, it +//! records the request with the pool's `grow` instead, which ignores `CometFairMemoryPool`'s limits +//! and carries what Spark doesn't grant as overcommit. +//! +//! This relies on the following in DataFusion 55.1, which a DataFusion upgrade has to re-check: +//! +//! - The two aggregates name their consumers `FinalHashAggregateStream[{partition}]` and +//! `OrderedFinalAggregateStream[{partition}]`, and their merge and replay reserve through sibling +//! reservations of that consumer. +//! - A final aggregate grows one reservation while another of its reservations holds memory only +//! during the replay, while the merge holds its read buffers. Before that, the aggregate's table +//! is its only reservation holding memory, so a refusal stands and makes it spill. The merge +//! picks its files while nothing else is held, so a refusal there still limits how many it opens. +//! - The replay asks for memory only after it has aggregated a batch, so like a `grow`, the request +//! is for memory that already exists. +//! - The replay emits every finished group after each batch, so it holds about one batch of groups, +//! and releasing memory repays the overcommit first. +//! +//! Remove this, as #6583 describes, once Comet's DataFusion has apache/datafusion#25383, which +//! leaves the replay room when the merge picks its files. The #6254 tests in `CometAggregateSuite` +//! fail without this on DataFusion 55.1, so they show whether the replay still needs it. + +use std::collections::HashMap; +use std::fmt; +use std::sync::Arc; + +use datafusion::common::{DataFusionError, Result}; +use datafusion::execution::memory_pool::{ + MemoryConsumer, MemoryLimit, MemoryPool, MemoryReservation, +}; +use log::debug; +use parking_lot::Mutex; + +/// Wraps a Comet pool and records the spill replay requests that it refuses; see the +/// [module documentation](self). Every other call, and every other refusal, is passed through +/// unchanged. +#[derive(Debug)] +pub(super) struct SpillReplayPool { + task_attempt_id: i64, + inner: Arc, + /// Bytes held by each final aggregate's consumer across all of its reservations, keyed by + /// [`MemoryConsumer::id`], because `reservation.size()` covers only one reservation. Other + /// consumers aren't tracked, so they never take the lock. + final_aggregates: Mutex>, +} + +impl SpillReplayPool { + pub(super) fn new(task_attempt_id: i64, inner: Arc) -> Self { + Self { + task_attempt_id, + inner, + final_aggregates: Mutex::new(HashMap::new()), + } + } + + /// Applies `update` to what `reservation`'s consumer holds, if it is a final aggregate. + fn track(&self, reservation: &MemoryReservation, update: impl FnOnce(&mut usize)) { + if is_final_aggregate(reservation.consumer()) { + if let Some(used) = self + .final_aggregates + .lock() + .get_mut(&reservation.consumer().id()) + { + update(used); + } + } + } + + /// Whether a refused request from `reservation` comes from a final aggregate reading its spill + /// files back, the only time one of its reservations grows while another holds memory. + fn is_spill_replay(&self, reservation: &MemoryReservation) -> bool { + is_final_aggregate(reservation.consumer()) + && self + .final_aggregates + .lock() + .get(&reservation.consumer().id()) + .is_some_and(|&consumer_used| consumer_used > reservation.size()) + } +} + +impl fmt::Display for SpillReplayPool { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fmt::Display::fmt(self.inner.as_ref(), f) + } +} -use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation}; +impl MemoryPool for SpillReplayPool { + fn name(&self) -> &str { + self.inner.name() + } + + fn register(&self, consumer: &MemoryConsumer) { + if is_final_aggregate(consumer) { + self.final_aggregates.lock().insert(consumer.id(), 0); + } + self.inner.register(consumer) + } + + fn unregister(&self, consumer: &MemoryConsumer) { + if is_final_aggregate(consumer) { + self.final_aggregates.lock().remove(&consumer.id()); + } + self.inner.unregister(consumer) + } + + fn grow(&self, reservation: &MemoryReservation, additional: usize) { + self.inner.grow(reservation, additional); + self.track(reservation, |used| *used = used.saturating_add(additional)); + } + + fn shrink(&self, reservation: &MemoryReservation, shrink: usize) { + self.inner.shrink(reservation, shrink); + self.track(reservation, |used| *used = used.saturating_sub(shrink)); + } + + fn try_grow(&self, reservation: &MemoryReservation, additional: usize) -> Result<()> { + match self.inner.try_grow(reservation, additional) { + Ok(()) => {} + Err(DataFusionError::ResourcesExhausted(refusal)) + if self.is_spill_replay(reservation) => + { + debug!( + "Task {} records {additional} bytes for {} while it reads its spill files back: {refusal}", + self.task_attempt_id, + reservation.consumer().name() + ); + self.inner.grow(reservation, additional); + } + Err(e) => return Err(e), + } + self.track(reservation, |used| *used = used.saturating_add(additional)); + Ok(()) + } + + fn reserved(&self) -> usize { + self.inner.reserved() + } + + fn memory_limit(&self) -> MemoryLimit { + self.inner.memory_limit() + } +} + +/// The pool that `pool` wraps, if it is a [`SpillReplayPool`]. +pub(super) fn unwrap_spill_replay(pool: &Arc) -> Option<&Arc> { + pool.downcast_ref::() + .map(|replay| &replay.inner) +} /// Whether `consumer` belongs to a final aggregate whose spill replay can't spill. -pub(super) fn is_final_aggregate(consumer: &MemoryConsumer) -> bool { +fn is_final_aggregate(consumer: &MemoryConsumer) -> bool { let name = consumer.name(); name.starts_with("FinalHashAggregateStream[") || name.starts_with("OrderedFinalAggregateStream[") } -/// Whether a refused request from `reservation` should be recorded instead. `consumer_used` is -/// what the reservation's consumer holds across all of its reservations. -/// -/// A final aggregate grows one reservation while another of its reservations holds memory only -/// during the replay, while the merge holds its read buffers. Before that, the aggregate's table -/// is its only reservation holding memory, so a refusal stands and makes it spill. The merge picks -/// its files while nothing else is held, so a refusal there still limits how many it opens. -pub(super) fn is_spill_replay(reservation: &MemoryReservation, consumer_used: usize) -> bool { - consumer_used > reservation.size() && is_final_aggregate(reservation.consumer()) -} - #[cfg(test)] mod tests { use super::super::spark_memory::fake::FakeSpark; use super::super::{create_pool, overcommit, MemoryPoolConfig, MemoryPoolType}; use super::*; - use datafusion::execution::memory_pool::MemoryPool; - use std::sync::Arc; + use datafusion::execution::memory_pool::UnboundedMemoryPool; /// A pool of each type, built the way `createPlan` builds it and connected to a fake Spark /// that grants at most 100 bytes. The fair pool's own limits are far above that, so only @@ -83,6 +210,15 @@ mod tests { .collect() } + /// A `fair_unified` pool of `pool_size` bytes, built the way `createPlan` builds it and + /// connected to a fake Spark that grants everything, so only the pool's own limits refuse. + fn fair_pool(task_attempt_id: i64, pool_size: usize) -> (Arc, Arc) { + let fake = FakeSpark::with(usize::MAX); + let config = MemoryPoolConfig::new(MemoryPoolType::FairUnified, pool_size); + let pool = create_pool(&config, task_attempt_id, || fake.memory()); + (pool, fake) + } + #[test] fn a_final_aggregate_reading_its_spill_files_back_carries_what_spark_refuses() { for (name, pool, fake) in each_pool_type([-3011, -3012]) { @@ -109,9 +245,39 @@ mod tests { } } + #[test] + fn a_final_aggregate_reading_its_spill_files_back_may_pass_its_fair_limit() { + let (pool, fake) = fair_pool(-3013, 100); + let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); + let replay = merge.new_empty(); + merge.try_grow(90).unwrap(); + + // The aggregate is the only consumer, so its share is the whole pool. + replay.try_grow(30).unwrap(); + assert_eq!(pool.reserved(), 120); + assert_eq!(fake.held(), 120); + assert_eq!(overcommit(&pool), 0); + } + + #[test] + fn a_final_aggregate_reading_its_spill_files_back_may_pass_the_pool_limit() { + let (pool, fake) = fair_pool(-3014, 100); + let sort = MemoryConsumer::new("ExternalSorter[0]").register(&pool); + sort.try_grow(70).unwrap(); + let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); + let replay = merge.new_empty(); + merge.try_grow(25).unwrap(); + + // 25 + 20 is within the aggregate's share of 50, but 70 + 25 + 20 is over the pool's 100. + replay.try_grow(20).unwrap(); + assert_eq!(pool.reserved(), 115); + assert_eq!(fake.held(), 115); + assert_eq!(overcommit(&pool), 0); + } + #[test] fn other_refusals_are_unchanged() { - for (name, pool, fake) in each_pool_type([-3013, -3014]) { + for (name, pool, fake) in each_pool_type([-3015, -3016]) { // While the aggregate reads its input, its table is the consumer's only reservation // holding memory, so a refusal makes it spill. let table = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); @@ -137,6 +303,51 @@ mod tests { } } + #[test] + fn a_failed_spark_call_during_the_replay_is_not_recorded() { + for (name, pool_type, task_attempt_id) in [ + ("greedy_unified", MemoryPoolType::GreedyUnified, -3017), + ("fair_unified", MemoryPoolType::FairUnified, -3018), + ] { + let fake = FakeSpark::failing(); + let config = MemoryPoolConfig::new(pool_type, 1000); + let pool = create_pool(&config, task_attempt_id, || fake.memory()); + let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&pool); + let replay = merge.new_empty(); + merge.grow(90); + + // Only a refusal is recorded. A JNI error goes back to the aggregate. + let err = replay.try_grow(30).unwrap_err(); + assert!( + !matches!(err, DataFusionError::ResourcesExhausted(_)), + "{name}: {err}" + ); + assert_eq!(pool.reserved(), 90, "{name}"); + } + } + + #[test] + fn a_final_aggregate_is_tracked_across_its_reservations_until_it_unregisters() { + let pool = Arc::new(SpillReplayPool::new( + -3019, + Arc::new(UnboundedMemoryPool::default()), + )); + let dyn_pool: Arc = Arc::clone(&pool) as _; + let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&dyn_pool); + let replay = merge.new_empty(); + merge.try_grow(90).unwrap(); + replay.grow(20); + replay.shrink(5); + let sort = MemoryConsumer::new("ExternalSorter[0]").register(&dyn_pool); + sort.grow(10); + let id = merge.consumer().id(); + assert_eq!(*pool.final_aggregates.lock(), HashMap::from([(id, 105)])); + + drop(merge); + drop(replay); + assert!(pool.final_aggregates.lock().is_empty()); + } + #[test] fn only_final_aggregates_replay_their_spill_files() { for name in [ diff --git a/native/core/src/execution/memory_pools/unified_pool.rs b/native/core/src/execution/memory_pools/unified_pool.rs index 8b89b1aec1b..edc51e0dd92 100644 --- a/native/core/src/execution/memory_pools/unified_pool.rs +++ b/native/core/src/execution/memory_pools/unified_pool.rs @@ -16,18 +16,16 @@ // under the License. use std::{ - collections::HashMap, fmt::{Debug, Display, Formatter, Result as FmtResult}, sync::atomic::{AtomicUsize, Ordering::Relaxed}, }; -use super::{spark_memory::SparkMemory, spill_replay}; +use super::spark_memory::SparkMemory; use datafusion::{ common::{resources_datafusion_err, DataFusionError}, - execution::memory_pool::{MemoryConsumer, MemoryPool, MemoryReservation}, + execution::memory_pool::{MemoryPool, MemoryReservation}, }; -use log::{debug, warn}; -use parking_lot::Mutex; +use log::warn; /// A DataFusion `MemoryPool` implementation for Comet that delegates to /// Spark's off-heap executor memory pool via JNI by calling @@ -35,10 +33,6 @@ use parking_lot::Mutex; pub struct CometUnifiedMemoryPool { spark: SparkMemory, used: AtomicUsize, - /// Bytes held by each final aggregate's consumer across its reservations, keyed by - /// [`MemoryConsumer::id`], which [`spill_replay`] needs. Other consumers aren't tracked, so - /// they take the lock only when Spark refuses them. - final_aggregates: Mutex>, } impl Debug for CometUnifiedMemoryPool { @@ -55,7 +49,6 @@ impl CometUnifiedMemoryPool { Self { spark, used: AtomicUsize::new(0), - final_aggregates: Mutex::new(HashMap::new()), } } @@ -63,31 +56,6 @@ impl CometUnifiedMemoryPool { pub(super) fn overcommit(&self) -> usize { self.spark.overcommit() } - - /// Applies `update` to what `reservation`'s consumer holds, if it is a final aggregate. - fn track(&self, reservation: &MemoryReservation, update: impl FnOnce(&mut usize)) { - if spill_replay::is_final_aggregate(reservation.consumer()) { - if let Some(used) = self - .final_aggregates - .lock() - .get_mut(&reservation.consumer().id()) - { - update(used); - } - } - } - - /// Whether a refused request from `reservation` comes from a final aggregate reading its spill - /// files back, which can't spill; see [`spill_replay`]. - fn is_spill_replay(&self, reservation: &MemoryReservation) -> bool { - let consumer_used = self - .final_aggregates - .lock() - .get(&reservation.consumer().id()) - .copied() - .unwrap_or(0); - spill_replay::is_spill_replay(reservation, consumer_used) - } } impl Drop for CometUnifiedMemoryPool { @@ -118,23 +86,11 @@ impl MemoryPool for CometUnifiedMemoryPool { "CometUnifiedMemoryPool" } - fn register(&self, consumer: &MemoryConsumer) { - if spill_replay::is_final_aggregate(consumer) { - self.final_aggregates.lock().insert(consumer.id(), 0); - } - } - - fn unregister(&self, consumer: &MemoryConsumer) { - if spill_replay::is_final_aggregate(consumer) { - self.final_aggregates.lock().remove(&consumer.id()); - } - } - /// Records memory that already exists, so it must not fail; see [`SparkMemory`]. // Rust 1.99 deprecates `fetch_update` in favor of `try_update`, which needs Rust 1.95, newer // than the workspace `rust-version`. #[allow(deprecated)] - fn grow(&self, reservation: &MemoryReservation, additional: usize) { + fn grow(&self, _: &MemoryReservation, additional: usize) { if additional == 0 { return; } @@ -142,13 +98,12 @@ impl MemoryPool for CometUnifiedMemoryPool { self.used .fetch_update(Relaxed, Relaxed, |old| Some(old.saturating_add(additional))) .unwrap(); - self.track(reservation, |used| *used = used.saturating_add(additional)); } // Rust 1.99 deprecates `fetch_update` in favor of `try_update`, which needs Rust 1.95, newer // than the workspace `rust-version`. #[allow(deprecated)] - fn shrink(&self, reservation: &MemoryReservation, size: usize) { + fn shrink(&self, _: &MemoryReservation, size: usize) { if let Err(e) = self.spark.release(size) { panic!( "Task {} failed to return {size} bytes to Spark: {e:?}", @@ -164,38 +119,23 @@ impl MemoryPool for CometUnifiedMemoryPool { self.spark.task_attempt_id() ); } - self.track(reservation, |used| *used = used.saturating_sub(size)); } // Rust 1.99 deprecates `fetch_update` in favor of `try_update`, which needs Rust 1.95, newer // than the workspace `rust-version`. #[allow(deprecated)] - fn try_grow( - &self, - reservation: &MemoryReservation, - additional: usize, - ) -> Result<(), DataFusionError> { + fn try_grow(&self, _: &MemoryReservation, additional: usize) -> Result<(), DataFusionError> { if additional > 0 { // A partial grant is handed back and refused, which triggers spilling in the caller. if let Err(refusal) = self.spark.try_acquire(additional)? { - let err = resources_datafusion_err!( + return Err(resources_datafusion_err!( "Task {} failed to acquire {} bytes plus {} bytes overcommitted, only got {}. Reserved: {}", self.spark.task_attempt_id(), additional, refusal.overcommit, refusal.granted, self.reserved() - ); - if !self.is_spill_replay(reservation) { - return Err(err); - } - debug!( - "Task {} records {additional} bytes for {} while it reads its spill files back: {err}", - self.spark.task_attempt_id(), - reservation.consumer().name() - ); - self.grow(reservation, additional); - return Ok(()); + )); } if let Err(prev) = self .used @@ -208,7 +148,6 @@ impl MemoryPool for CometUnifiedMemoryPool { prev )); } - self.track(reservation, |used| *used = used.saturating_add(additional)); } Ok(()) } @@ -222,6 +161,7 @@ impl MemoryPool for CometUnifiedMemoryPool { mod tests { use super::super::spark_memory::fake::FakeSpark; use super::*; + use datafusion::execution::memory_pool::MemoryConsumer; use std::sync::Arc; #[test] @@ -291,23 +231,4 @@ mod tests { assert_eq!(pool.spark.overcommit(), 0); assert_eq!(fake.held(), 0); } - - #[test] - fn a_final_aggregate_is_tracked_across_its_reservations_until_it_unregisters() { - let pool = Arc::new(CometUnifiedMemoryPool::with_spark( - FakeSpark::with(100).memory(), - )); - let dyn_pool: Arc = Arc::clone(&pool) as _; - let merge = MemoryConsumer::new("FinalHashAggregateStream[0]").register(&dyn_pool); - let replay = merge.new_empty(); - merge.try_grow(90).unwrap(); - replay.grow(20); - replay.shrink(5); - let id = merge.consumer().id(); - assert_eq!(pool.final_aggregates.lock().get(&id), Some(&105)); - - drop(merge); - drop(replay); - assert!(pool.final_aggregates.lock().is_empty()); - } }