Skip to content
Draft
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
301 changes: 259 additions & 42 deletions native/core/src/execution/memory_pools/fair_pool.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
// under the License.

use std::{
collections::HashMap,
fmt::{Debug, Display, Formatter, Result as FmtResult},
sync::Arc,
};
Expand All @@ -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;

Expand All @@ -37,11 +38,16 @@ pub struct CometFairMemoryPool {
task_memory_manager_handle: Arc<Global<JObject<'static>>>,
pool_size: usize,
state: Mutex<CometFairPoolState>,
#[cfg(test)]
test_memory_manager: Option<Arc<tests::TestMemoryManager>>,
}

#[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<usize, usize>,
}

impl Debug for CometFairMemoryPool {
Expand All @@ -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()
}
}
Expand All @@ -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<i64> {
#[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,
Expand All @@ -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) -> ())
Expand All @@ -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()
)
}
}
Expand All @@ -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
);
}

Expand All @@ -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)
Expand All @@ -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<i64> {
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<dyn MemoryPool>, Arc<TestMemoryManager>) {
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);
}
}
Loading