diff --git a/native/core/src/execution/memory_pools/fair_pool.rs b/native/core/src/execution/memory_pools/fair_pool.rs index 347c3d8ef69..993194ee303 100644 --- a/native/core/src/execution/memory_pools/fair_pool.rs +++ b/native/core/src/execution/memory_pools/fair_pool.rs @@ -16,6 +16,7 @@ // under the License. use std::{ + collections::HashMap, fmt::{Debug, Display, Formatter, Result as FmtResult}, sync::Arc, }; @@ -27,7 +28,7 @@ use datafusion::common::resources_err; use datafusion::execution::memory_pool::MemoryConsumer; use datafusion::{ common::DataFusionError, - execution::memory_pool::{MemoryPool, MemoryReservation}, + execution::memory_pool::{MemoryLimit, MemoryPool, MemoryReservation}, }; use parking_lot::Mutex; @@ -37,11 +38,16 @@ pub struct CometFairMemoryPool { task_memory_manager_handle: Arc>>, pool_size: usize, state: Mutex, + #[cfg(test)] + test_memory_manager: Option>, } +#[derive(Default)] struct CometFairPoolState { used: usize, - num: usize, + // new_empty(), split(), and take() share a registered consumer but have independent sizes. + // All sibling reservations must count against the same fair share. + consumer_usage: HashMap, } impl Debug for CometFairMemoryPool { @@ -50,7 +56,7 @@ impl Debug for CometFairMemoryPool { f.debug_struct("CometFairMemoryPool") .field("pool_size", &self.pool_size) .field("used", &state.used) - .field("num", &state.num) + .field("num", &state.consumer_usage.len()) .finish() } } @@ -63,11 +69,17 @@ impl CometFairMemoryPool { Self { task_memory_manager_handle, pool_size, - state: Mutex::new(CometFairPoolState { used: 0, num: 0 }), + state: Mutex::new(CometFairPoolState::default()), + #[cfg(test)] + test_memory_manager: None, } } fn acquire(&self, additional: usize) -> CometResult { + #[cfg(test)] + if let Some(manager) = &self.test_memory_manager { + return manager.acquire(additional); + } let handle = self.task_memory_manager_handle.as_obj(); JVMClasses::with_env(|env| unsafe { jni_call!(env, @@ -76,6 +88,11 @@ impl CometFairMemoryPool { } fn release(&self, size: usize) -> CometResult<()> { + #[cfg(test)] + if let Some(manager) = &self.test_memory_manager { + manager.release(size); + return Ok(()); + } let handle = self.task_memory_manager_handle.as_obj(); JVMClasses::with_env(|env| unsafe { jni_call!(env, comet_task_memory_manager(handle).release_memory(size as i64) -> ()) @@ -89,7 +106,9 @@ impl Display for CometFairMemoryPool { write!( f, "CometFairMemoryPool(pool_size={}, used={}, num={})", - self.pool_size, state.used, state.num + self.pool_size, + state.used, + state.consumer_usage.len() ) } } @@ -102,63 +121,74 @@ impl MemoryPool for CometFairMemoryPool { "CometFairMemoryPool" } - fn register(&self, _: &MemoryConsumer) { - let mut state = self.state.lock(); - state.num = state - .num - .checked_add(1) - .expect("unexpected amount of register happened"); + fn register(&self, consumer: &MemoryConsumer) { + assert!( + self.state + .lock() + .consumer_usage + .insert(consumer.id(), 0) + .is_none(), + "memory consumer was registered more than once" + ); } - fn unregister(&self, _: &MemoryConsumer) { - let mut state = self.state.lock(); - state.num = state - .num - .checked_sub(1) - .expect("unexpected amount of unregister happened"); + fn unregister(&self, consumer: &MemoryConsumer) { + let usage = self.state.lock().consumer_usage.remove(&consumer.id()); + assert_eq!( + usage, + Some(0), + "consumer must release its reservations before unregistering" + ); } - fn grow(&self, _reservation: &MemoryReservation, additional: usize) { - self.try_grow(_reservation, additional).unwrap(); + fn grow(&self, reservation: &MemoryReservation, additional: usize) { + self.try_grow(reservation, additional).unwrap(); } - fn shrink(&self, _reservation: &MemoryReservation, subtractive: usize) { + fn shrink(&self, reservation: &MemoryReservation, subtractive: usize) { if subtractive > 0 { let mut state = self.state.lock(); - // We don't use reservation.size() here because DataFusion 53+ decrements - // the reservation's atomic size before calling pool.shrink(), so it would - // reflect the post-shrink value rather than the pre-shrink value. - if state.used < subtractive { - panic!( - "Failed to release {subtractive} bytes where only {} bytes tracked by pool", - state.used - ) - } + let usage = state.consumer_usage[&reservation.consumer().id()]; + assert!( + usage >= subtractive, + "consumer released more bytes than it reserved" + ); self.release(subtractive) .unwrap_or_else(|_| panic!("Failed to release {subtractive} bytes")); - state.used = state.used.checked_sub(subtractive).unwrap(); + *state + .consumer_usage + .get_mut(&reservation.consumer().id()) + .unwrap() -= subtractive; + state.used -= subtractive; } } 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 { + // Preserve the policy of sharing among all registered consumers. Spillability + // annotations and sharing only among spillable consumers are a separate change. + let num = state.consumer_usage.len(); + let limit = self.pool_size / num; + let used = state.consumer_usage[&reservation.consumer().id()]; + if used + .checked_add(additional) + .is_none_or(|requested| requested > limit) + { return resources_err!( - "Failed to acquire {additional} bytes where {used} bytes already reserved and the fair limit is {limit} bytes, {num} registered" + "Failed to acquire {additional} bytes where {used} bytes already reserved by this consumer and the fair limit is {limit} bytes, {num} registered" + ); + } + // Existing allocations may exceed their new fair share when another consumer + // registers. A per-consumer bound alone cannot enforce the configured pool size. + if additional > self.pool_size.saturating_sub(state.used) { + return resources_err!( + "Failed to acquire {additional} bytes where {} bytes already reserved pool-wide and the pool limit is {} bytes", + state.used, self.pool_size ); } @@ -176,6 +206,10 @@ impl MemoryPool for CometFairMemoryPool { state.used ); } + *state + .consumer_usage + .get_mut(&reservation.consumer().id()) + .unwrap() += additional; state.used = state .used .checked_add(additional) @@ -187,4 +221,187 @@ 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 std::sync::atomic::{AtomicBool, AtomicUsize, Ordering::SeqCst}; + + // Replace only JNI grants/releases so the tests exercise actual MemoryPool admission and + // DataFusion reservation lifetimes without requiring a running Spark task. + #[derive(Default)] + pub(super) struct TestMemoryManager { + used: AtomicUsize, + acquires: AtomicUsize, + partial_grant: AtomicBool, + fail_acquire: AtomicBool, + } + + impl TestMemoryManager { + pub(super) fn acquire(&self, requested: usize) -> CometResult { + self.acquires.fetch_add(1, SeqCst); + if self.fail_acquire.load(SeqCst) { + return Err(crate::errors::CometError::Internal( + "test acquire failure".into(), + )); + } + let granted = if self.partial_grant.load(SeqCst) { + requested / 2 + } else { + requested + }; + self.used.fetch_add(granted, SeqCst); + Ok(granted as i64) + } + + pub(super) fn release(&self, bytes: usize) { + self.used + .fetch_update(SeqCst, SeqCst, |used| used.checked_sub(bytes)) + .unwrap(); + } + } + + fn pool(size: usize) -> (Arc, Arc) { + let manager = Arc::new(TestMemoryManager::default()); + let mut pool = CometFairMemoryPool::new(Arc::new(Global::null()), size); + pool.test_memory_manager = Some(Arc::clone(&manager)); + (Arc::new(pool), manager) + } + + #[test] + fn each_consumer_can_use_its_fair_share() { + let (pool, manager) = pool(32); + let other = MemoryConsumer::new("other").register(&pool); + let requesting = MemoryConsumer::new("requesting").register(&pool); + other.try_grow(10).unwrap(); + requesting.try_grow(6).unwrap(); + requesting.try_grow(10).unwrap(); + assert!(requesting.try_grow(1).is_err()); + other.try_grow(6).unwrap(); + assert_eq!(pool.reserved(), 32); + assert_eq!(manager.used.load(SeqCst), 32); + assert!(matches!(pool.memory_limit(), MemoryLimit::Finite(32))); + } + + #[test] + fn sibling_reservations_share_their_consumers_limit() { + let (pool, manager) = pool(100); + let parent = MemoryConsumer::new("same name") + .with_can_spill(true) + .register(&pool); + let first = parent.new_empty(); + let second = parent.new_empty(); + let other = MemoryConsumer::new("same name") + .with_can_spill(true) + .register(&pool); + first.try_grow(30).unwrap(); + second.try_grow(20).unwrap(); + assert!(parent.try_grow(1).is_err()); + assert!(second.try_grow(1).is_err()); + // Rejected growth never asks Spark for memory, even though the pool has free capacity. + assert_eq!(manager.acquires.load(SeqCst), 2); + other.try_grow(50).unwrap(); + drop(parent); + drop(first); + second.try_grow(30).unwrap(); + assert_eq!(pool.reserved(), 100); + drop(second); + drop(other); + assert_eq!(pool.reserved(), 0); + assert_eq!(manager.used.load(SeqCst), 0); + } + + #[test] + fn split_and_take_do_not_create_another_allowance() { + let (pool, manager) = pool(100); + let mut parent = MemoryConsumer::new("consumer").register(&pool); + let other = MemoryConsumer::new("other").register(&pool); + parent.try_grow(50).unwrap(); + let split = parent.split(20); + let taken = parent.take(); + assert_eq!(parent.size(), 0); + assert_eq!(split.size(), 20); + assert_eq!(taken.size(), 30); + for reservation in [&parent, &split, &taken] { + assert!(reservation.try_grow(1).is_err()); + } + assert_eq!(manager.acquires.load(SeqCst), 1); + split.shrink(10); + parent.try_grow(10).unwrap(); + drop(parent); + drop(split); + taken.try_grow(20).unwrap(); + other.try_grow(50).unwrap(); + assert_eq!(pool.reserved(), 100); + drop(taken); + drop(other); + assert_eq!(pool.reserved(), 0); + assert_eq!(manager.used.load(SeqCst), 0); + } + + #[test] + fn registration_after_allocation_cannot_exceed_pool_capacity() { + let (pool, manager) = pool(100); + let first = MemoryConsumer::new("first").register(&pool); + first.try_grow(100).unwrap(); + let second = MemoryConsumer::new("second").register(&pool); + assert!(second.try_grow(1).is_err()); + assert_eq!(manager.acquires.load(SeqCst), 1); + assert_eq!(pool.reserved(), 100); + first.shrink(50); + second.try_grow(50).unwrap(); + assert_eq!(pool.reserved(), 100); + } + + #[test] + fn mixed_consumers_keep_the_existing_sharing_policy() { + let (pool, _) = pool(100); + let fixed = MemoryConsumer::new("fixed").register(&pool); + let spilling = MemoryConsumer::new("spilling") + .with_can_spill(true) + .register(&pool); + assert!(spilling.try_grow(51).is_err()); + spilling.try_grow(50).unwrap(); + fixed.try_grow(50).unwrap(); + assert_eq!(pool.reserved(), 100); + fixed.free(); + assert!(spilling.try_grow(1).is_err()); + drop(fixed); + spilling.try_grow(50).unwrap(); + assert_eq!(pool.reserved(), 100); + } + + #[test] + fn failed_acquisition_does_not_change_consumer_or_pool_usage() { + let (pool, manager) = pool(100); + let reservation = MemoryConsumer::new("consumer").register(&pool); + manager.partial_grant.store(true, SeqCst); + assert!(reservation.try_grow(100).is_err()); + assert_eq!(pool.reserved(), 0); + assert_eq!(manager.used.load(SeqCst), 0); + manager.partial_grant.store(false, SeqCst); + manager.fail_acquire.store(true, SeqCst); + assert!(reservation.try_grow(100).is_err()); + assert_eq!(pool.reserved(), 0); + manager.fail_acquire.store(false, SeqCst); + reservation.try_grow(100).unwrap(); + reservation.free(); + assert_eq!(pool.reserved(), 0); + assert_eq!(manager.used.load(SeqCst), 0); + } + + #[test] + fn overflowing_request_is_rejected_before_acquisition() { + let (pool, manager) = pool(100); + let reservation = MemoryConsumer::new("consumer").register(&pool); + reservation.try_grow(1).unwrap(); + assert!(reservation.try_grow(usize::MAX).is_err()); + assert_eq!(pool.reserved(), 1); + assert_eq!(manager.acquires.load(SeqCst), 1); + } }