Skip to content
Draft
Show file tree
Hide file tree
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
45 changes: 13 additions & 32 deletions native/core/src/execution/memory_pools/fair_pool.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,14 +15,9 @@
// specific language governing permissions and limitations
// under the License.

use std::{
fmt::{Debug, Display, Formatter, Result as FmtResult},
sync::Arc,
};

use jni::objects::{Global, JObject};
use std::fmt::{Debug, Display, Formatter, Result as FmtResult};

use crate::{errors::CometResult, jvm_bridge::JVMClasses};
use crate::execution::memory_pools::spark_client::SparkMemoryClient;
use datafusion::common::resources_err;
use datafusion::execution::memory_pool::MemoryConsumer;
use datafusion::{
Expand All @@ -34,7 +29,7 @@ use parking_lot::Mutex;
/// A DataFusion fair `MemoryPool` implementation for Comet. Internally this is
/// implemented via delegating calls to [`crate::jvm_bridge::CometTaskMemoryManager`].
pub struct CometFairMemoryPool {
task_memory_manager_handle: Arc<Global<JObject<'static>>>,
client: SparkMemoryClient,
pool_size: usize,
state: Mutex<CometFairPoolState>,
}
Expand All @@ -56,30 +51,18 @@ impl Debug for CometFairMemoryPool {
}

impl CometFairMemoryPool {
pub fn new(
task_memory_manager_handle: Arc<Global<JObject<'static>>>,
pool_size: usize,
) -> CometFairMemoryPool {
pub fn new(client: SparkMemoryClient, pool_size: usize) -> CometFairMemoryPool {
Self {
task_memory_manager_handle,
client,
pool_size,
state: Mutex::new(CometFairPoolState { used: 0, num: 0 }),
}
}
}

fn acquire(&self, additional: usize) -> CometResult<i64> {
let handle = self.task_memory_manager_handle.as_obj();
JVMClasses::with_env(|env| unsafe {
jni_call!(env,
comet_task_memory_manager(handle).acquire_memory(additional as i64) -> i64)
})
}

fn release(&self, size: usize) -> CometResult<()> {
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) -> ())
})
impl Drop for CometFairMemoryPool {
fn drop(&mut self) {
self.client.log_stats(self.name());
}
}

Expand All @@ -94,9 +77,6 @@ impl Display for CometFairMemoryPool {
}
}

unsafe impl Send for CometFairMemoryPool {}
unsafe impl Sync for CometFairMemoryPool {}

impl MemoryPool for CometFairMemoryPool {
fn name(&self) -> &str {
"CometFairMemoryPool"
Expand Down Expand Up @@ -134,7 +114,8 @@ impl MemoryPool for CometFairMemoryPool {
state.used
)
}
self.release(subtractive)
self.client
.release(subtractive)
.unwrap_or_else(|_| panic!("Failed to release {subtractive} bytes"));
state.used = state.used.checked_sub(subtractive).unwrap();
}
Expand Down Expand Up @@ -162,12 +143,12 @@ impl MemoryPool for CometFairMemoryPool {
);
}

let acquired = self.acquire(additional)?;
let acquired = self.client.acquire(additional)?;
// If the number of bytes we acquired is less than the requested, return an error,
// and hopefully will trigger spilling from the caller side.
if acquired < additional as i64 {
// Release the acquired bytes before throwing error
self.release(acquired as usize)?;
self.client.release(acquired as usize)?;

return resources_err!(
"Failed to acquire {} bytes, only got {} bytes. Reserved: {} bytes",
Expand Down
11 changes: 8 additions & 3 deletions native/core/src/execution/memory_pools/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
mod config;
mod fair_pool;
pub mod logging_pool;
mod spark_client;
mod task_shared;
mod unified_pool;

Expand All @@ -27,6 +28,7 @@ use datafusion::execution::memory_pool::{
use fair_pool::CometFairMemoryPool;
use jni::objects::{Global, JObject};
use once_cell::sync::OnceCell;
use spark_client::SparkMemoryClient;
use std::num::NonZeroUsize;
use std::sync::Arc;
use unified_pool::CometUnifiedMemoryPool;
Expand All @@ -46,10 +48,10 @@ pub(crate) fn create_memory_pool(
let per_task_memory_pool =
memory_pool_map.entry(task_attempt_id).or_insert_with(|| {
let pool: Arc<dyn MemoryPool> = Arc::new(TrackConsumersPool::new(
CometUnifiedMemoryPool::new(
CometUnifiedMemoryPool::new(SparkMemoryClient::new(
Arc::clone(&comet_task_memory_manager),
task_attempt_id,
),
)),
NonZeroUsize::new(NUM_TRACKED_CONSUMERS).unwrap(),
));
PerTaskMemoryPool::new(pool)
Expand All @@ -63,7 +65,10 @@ pub(crate) fn create_memory_pool(
memory_pool_map.entry(task_attempt_id).or_insert_with(|| {
let pool: Arc<dyn MemoryPool> = Arc::new(TrackConsumersPool::new(
CometFairMemoryPool::new(
Arc::clone(&comet_task_memory_manager),
SparkMemoryClient::new(
Arc::clone(&comet_task_memory_manager),
task_attempt_id,
),
memory_pool_config.pool_size,
),
NonZeroUsize::new(NUM_TRACKED_CONSUMERS).unwrap(),
Expand Down
Loading
Loading