use std::collections::HashMap;
use std::hash::{BuildHasher, BuildHasherDefault, Hasher};
use std::sync::Mutex;
use lora_compiler::CompiledQuery;
use std::sync::Arc;
const DEFAULT_CAPACITY: usize = 256;
pub(crate) struct PlanCache {
inner: Mutex<Inner>,
}
struct Inner {
entries: HashMap<u64, Entry>,
counter: u64,
capacity: usize,
}
struct Entry {
query: String,
plan: Arc<CompiledQuery>,
last_used: u64,
}
impl PlanCache {
pub(crate) fn new() -> Self {
Self::with_capacity(DEFAULT_CAPACITY)
}
pub(crate) fn with_capacity(capacity: usize) -> Self {
Self {
inner: Mutex::new(Inner {
entries: HashMap::with_capacity(capacity),
counter: 0,
capacity,
}),
}
}
pub(crate) fn get(&self, query: &str) -> Option<Arc<CompiledQuery>> {
let hash = hash_query(query);
let mut guard = self
.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let counter = guard.counter.wrapping_add(1);
guard.counter = counter;
let entry = guard.entries.get_mut(&hash)?;
if entry.query != query {
return None;
}
entry.last_used = counter;
Some(entry.plan.clone())
}
pub(crate) fn insert(&self, query: &str, plan: Arc<CompiledQuery>) {
let hash = hash_query(query);
let mut guard = self
.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let counter = guard.counter.wrapping_add(1);
guard.counter = counter;
if guard.entries.len() >= guard.capacity && !guard.entries.contains_key(&hash) {
evict_oldest(&mut guard.entries);
}
guard.entries.insert(
hash,
Entry {
query: query.to_owned(),
plan,
last_used: counter,
},
);
}
#[cfg(test)]
pub(crate) fn len(&self) -> usize {
self.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.entries
.len()
}
}
fn evict_oldest(entries: &mut HashMap<u64, Entry>) {
let oldest_key = entries
.iter()
.min_by_key(|(_, e)| e.last_used)
.map(|(k, _)| *k);
if let Some(k) = oldest_key {
entries.remove(&k);
}
}
fn hash_query(query: &str) -> u64 {
let hasher_builder = BuildHasherDefault::<std::collections::hash_map::DefaultHasher>::default();
let mut hasher = hasher_builder.build_hasher();
hasher.write(query.as_bytes());
hasher.finish()
}
#[cfg(test)]
mod tests {
use super::*;
use lora_compiler::{Compiler, PhysicalNodeId, PhysicalOp, PhysicalPlan};
fn dummy_plan() -> Arc<CompiledQuery> {
Arc::new(CompiledQuery {
physical: PhysicalPlan {
root: 0,
nodes: vec![PhysicalOp::Argument(lora_compiler::ArgumentExec)],
},
unions: Vec::new(),
})
}
#[test]
fn miss_then_hit() {
let cache = PlanCache::new();
let q = "MATCH (n) RETURN n";
assert!(cache.get(q).is_none());
cache.insert(q, dummy_plan());
let hit = cache.get(q).expect("expected cache hit");
cache.insert(q, hit.clone());
assert_eq!(cache.len(), 1);
}
#[test]
fn distinct_queries_are_independent() {
let cache = PlanCache::new();
cache.insert("MATCH (n) RETURN n", dummy_plan());
cache.insert("MATCH (m) RETURN m", dummy_plan());
assert_eq!(cache.len(), 2);
assert!(cache.get("MATCH (n) RETURN n").is_some());
assert!(cache.get("MATCH (m) RETURN m").is_some());
}
#[test]
fn lru_evicts_oldest() {
let cache = PlanCache::with_capacity(2);
cache.insert("a", dummy_plan());
cache.insert("b", dummy_plan());
let _ = cache.get("a");
cache.insert("c", dummy_plan());
assert_eq!(cache.len(), 2);
assert!(cache.get("a").is_some());
assert!(cache.get("b").is_none());
assert!(cache.get("c").is_some());
}
#[test]
fn unused_compiler_use_silences_warning() {
let _ = std::any::type_name::<Compiler>();
let _ = std::any::type_name::<PhysicalNodeId>();
}
}