use std::num::NonZeroUsize;
use std::sync::Arc;
use ballista_core::RuntimeProducer;
use datafusion::execution::runtime_env::RuntimeEnv;
use datafusion::prelude::SessionConfig;
use lru::LruCache;
use parking_lot::Mutex;
pub type MemoryPoolPolicy = Arc<
dyn Fn(Arc<RuntimeEnv>, &SessionConfig) -> datafusion::error::Result<Arc<RuntimeEnv>>
+ Send
+ Sync,
>;
pub trait SessionRuntimeCache: Send + Sync {
fn produce_runtime(
&self,
session_id: &str,
config: &SessionConfig,
) -> datafusion::error::Result<Arc<RuntimeEnv>>;
}
pub struct DefaultSessionRuntimeCache {
base_producer: RuntimeProducer,
pool_policy: MemoryPoolPolicy,
cache: Option<Mutex<LruCache<String, Arc<RuntimeEnv>>>>,
}
impl DefaultSessionRuntimeCache {
pub fn new(
base_producer: RuntimeProducer,
pool_policy: MemoryPoolPolicy,
capacity: usize,
) -> Self {
let cache = NonZeroUsize::new(capacity).map(|cap| Mutex::new(LruCache::new(cap)));
Self {
base_producer,
pool_policy,
cache,
}
}
}
impl SessionRuntimeCache for DefaultSessionRuntimeCache {
fn produce_runtime(
&self,
session_id: &str,
config: &SessionConfig,
) -> datafusion::error::Result<Arc<RuntimeEnv>> {
let base = match &self.cache {
None => (self.base_producer)(config)?,
Some(cache) => {
if let Some(base) = cache.lock().get(session_id) {
base.clone()
} else {
let base = (self.base_producer)(config)?;
cache.lock().put(session_id.to_string(), base.clone());
base
}
}
};
(self.pool_policy)(base, config)
}
}
#[cfg(test)]
mod tests {
use super::*;
use datafusion::execution::memory_pool::GreedyMemoryPool;
use datafusion::execution::runtime_env::RuntimeEnvBuilder;
fn base_producer() -> RuntimeProducer {
Arc::new(|_| Ok(Arc::new(RuntimeEnv::default())))
}
fn identity_policy() -> MemoryPoolPolicy {
Arc::new(|base, _| Ok(base))
}
fn per_task_pool_policy() -> MemoryPoolPolicy {
Arc::new(|base, _| {
RuntimeEnvBuilder::from_runtime_env(&base)
.with_memory_pool(Arc::new(GreedyMemoryPool::new(1024)))
.build_arc()
})
}
#[test]
fn same_session_shares_cache_manager() {
let cache =
DefaultSessionRuntimeCache::new(base_producer(), identity_policy(), 4);
let cfg = SessionConfig::new();
let e1 = cache.produce_runtime("s1", &cfg).unwrap();
let e2 = cache.produce_runtime("s1", &cfg).unwrap();
assert!(Arc::ptr_eq(&e1.cache_manager, &e2.cache_manager));
}
#[test]
fn different_sessions_get_different_base() {
let cache =
DefaultSessionRuntimeCache::new(base_producer(), identity_policy(), 4);
let cfg = SessionConfig::new();
let e1 = cache.produce_runtime("s1", &cfg).unwrap();
let e2 = cache.produce_runtime("s2", &cfg).unwrap();
assert!(!Arc::ptr_eq(&e1.cache_manager, &e2.cache_manager));
}
#[test]
fn per_task_pool_shares_footer_cache_but_not_env() {
let cache =
DefaultSessionRuntimeCache::new(base_producer(), per_task_pool_policy(), 4);
let cfg = SessionConfig::new();
let e1 = cache.produce_runtime("s1", &cfg).unwrap();
let e2 = cache.produce_runtime("s1", &cfg).unwrap();
assert!(!Arc::ptr_eq(&e1, &e2));
assert!(Arc::ptr_eq(
&e1.object_store_registry,
&e2.object_store_registry
));
assert!(Arc::ptr_eq(
&e1.cache_manager.get_file_metadata_cache(),
&e2.cache_manager.get_file_metadata_cache(),
));
}
#[test]
fn capacity_zero_disables_cache() {
let cache =
DefaultSessionRuntimeCache::new(base_producer(), identity_policy(), 0);
let cfg = SessionConfig::new();
let e1 = cache.produce_runtime("s1", &cfg).unwrap();
let e2 = cache.produce_runtime("s1", &cfg).unwrap();
assert!(!Arc::ptr_eq(&e1.cache_manager, &e2.cache_manager));
}
#[test]
fn evicts_least_recently_used() {
let cache =
DefaultSessionRuntimeCache::new(base_producer(), identity_policy(), 2);
let cfg = SessionConfig::new();
let s1_a = cache.produce_runtime("s1", &cfg).unwrap();
cache.produce_runtime("s2", &cfg).unwrap();
cache.produce_runtime("s3", &cfg).unwrap();
let s1_b = cache.produce_runtime("s1", &cfg).unwrap();
assert!(!Arc::ptr_eq(&s1_a.cache_manager, &s1_b.cache_manager));
}
}