From a658c1960ea2ae27cfa4584e77800163fab16645 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Wed, 26 Aug 2026 00:08:39 +0800 Subject: [PATCH] fix: enforce fair memory share per reservation --- .../src/execution/memory_pools/fair_pool.rs | 53 ++++++++++++++----- 1 file changed, 41 insertions(+), 12 deletions(-) diff --git a/native/core/src/execution/memory_pools/fair_pool.rs b/native/core/src/execution/memory_pools/fair_pool.rs index 347c3d8ef69..1d231d00aba 100644 --- a/native/core/src/execution/memory_pools/fair_pool.rs +++ b/native/core/src/execution/memory_pools/fair_pool.rs @@ -24,10 +24,9 @@ use jni::objects::{Global, JObject}; use crate::{errors::CometResult, jvm_bridge::JVMClasses}; use datafusion::common::resources_err; -use datafusion::execution::memory_pool::MemoryConsumer; use datafusion::{ common::DataFusionError, - execution::memory_pool::{MemoryPool, MemoryReservation}, + execution::memory_pool::{MemoryConsumer, MemoryLimit, MemoryPool, MemoryReservation}, }; use parking_lot::Mutex; @@ -44,6 +43,17 @@ struct CometFairPoolState { num: usize, } +fn fair_limit_exceeded( + pool_size: usize, + num: usize, + reservation: &MemoryReservation, + additional: usize, +) -> Option<(usize, usize)> { + let used = reservation.size(); + let limit = pool_size.checked_div(num).expect("overflow in checked_div"); + (limit < used.saturating_add(additional)).then_some((used, limit)) +} + impl Debug for CometFairMemoryPool { fn fmt(&self, f: &mut Formatter<'_>) -> FmtResult { let state = self.state.lock(); @@ -142,21 +152,15 @@ impl MemoryPool for CometFairMemoryPool { fn try_grow( &self, - _reservation: &MemoryReservation, + reservation: &MemoryReservation, additional: usize, ) -> Result<(), DataFusionError> { if additional > 0 { let mut state = self.state.lock(); let num = state.num; - let limit = self - .pool_size - .checked_div(num) - .expect("overflow in checked_div"); - // We use state.used instead of reservation.size() because DataFusion 53+ - // calls pool.try_grow() before incrementing the reservation's atomic size, - // so reservation.size() would not include prior grows. - let used = state.used; - if limit < used + additional { + if let Some((used, limit)) = + fair_limit_exceeded(self.pool_size, num, reservation, additional) + { return resources_err!( "Failed to acquire {additional} bytes where {used} bytes already reserved and the fair limit is {limit} bytes, {num} registered" ); @@ -187,4 +191,29 @@ impl MemoryPool for CometFairMemoryPool { fn reserved(&self) -> usize { self.state.lock().used } + + fn memory_limit(&self) -> MemoryLimit { + MemoryLimit::Finite(self.pool_size) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use datafusion::execution::memory_pool::UnboundedMemoryPool; + + #[test] + fn fair_share_uses_requesting_reservation_and_reports_pool_limit() { + let backing: Arc = Arc::new(UnboundedMemoryPool::default()); + let other = MemoryConsumer::new("other").register(&backing); + let requesting = MemoryConsumer::new("requesting").register(&backing); + other.grow(10); + requesting.grow(6); + + assert_eq!(fair_limit_exceeded(32, 2, &requesting, 10), None); + assert_eq!(fair_limit_exceeded(32, 2, &requesting, 11), Some((6, 16))); + + let pool = CometFairMemoryPool::new(Arc::new(Global::null()), 32); + assert!(matches!(pool.memory_limit(), MemoryLimit::Finite(32))); + } }