use std::collections::{HashMap, VecDeque};
use crate::query::prepared::PreparedStatement;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct StatementCache {
cache: HashMap<String, PreparedStatement>,
capacity: usize,
order: VecDeque<String>,
}
impl StatementCache {
pub fn new(capacity: usize) -> Self {
Self {
cache: HashMap::with_capacity(capacity),
capacity,
order: VecDeque::with_capacity(capacity),
}
}
pub fn len(&self) -> usize {
self.cache.len()
}
pub fn is_empty(&self) -> bool {
self.cache.is_empty()
}
pub fn get(&mut self, sql: &str) -> Option<&PreparedStatement> {
if self.cache.contains_key(sql) {
self.order.retain(|s| s != sql);
self.order.push_front(sql.to_string());
self.cache.get(sql)
} else {
None
}
}
pub fn insert(&mut self, stmt: PreparedStatement) -> Option<PreparedStatement> {
let sql = stmt.sql().to_string();
self.order.retain(|s| s != &sql);
let evicted = if self.cache.len() >= self.capacity && !self.cache.contains_key(&sql) {
self.order
.pop_back()
.and_then(|old_sql| self.cache.remove(&old_sql))
} else {
None
};
self.order.push_front(sql.clone());
self.cache.insert(sql, stmt);
evicted
}
pub fn remove(&mut self, sql: &str) -> Option<PreparedStatement> {
self.order.retain(|s| s != sql);
self.cache.remove(sql)
}
pub fn clear(&mut self) -> Vec<PreparedStatement> {
self.order.clear();
self.cache.drain().map(|(_, v)| v).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
fn dummy_stmt(sql: &str, name: &str) -> PreparedStatement {
PreparedStatement {
name: name.into(),
sql: sql.into(),
param_types: vec![],
columns: Arc::new(vec![]),
}
}
#[test]
fn test_cache_insert_and_get() {
let mut cache = StatementCache::new(2);
let stmt = dummy_stmt("SELECT 1", "s1");
cache.insert(stmt);
assert_eq!(cache.len(), 1);
assert!(cache.get("SELECT 1").is_some());
assert!(cache.get("SELECT 2").is_none());
}
#[test]
fn test_cache_lru_eviction() {
let mut cache = StatementCache::new(2);
cache.insert(dummy_stmt("SELECT 1", "s1"));
cache.insert(dummy_stmt("SELECT 2", "s2"));
let _ = cache.get("SELECT 1");
let evicted = cache.insert(dummy_stmt("SELECT 3", "s3"));
assert!(evicted.is_some());
assert_eq!(evicted.unwrap().sql(), "SELECT 2");
assert!(cache.get("SELECT 1").is_some());
assert!(cache.get("SELECT 2").is_none());
assert!(cache.get("SELECT 3").is_some());
}
#[test]
fn test_cache_remove() {
let mut cache = StatementCache::new(2);
cache.insert(dummy_stmt("SELECT 1", "s1"));
let removed = cache.remove("SELECT 1");
assert!(removed.is_some());
assert_eq!(removed.unwrap().sql(), "SELECT 1");
assert!(cache.is_empty());
}
#[test]
fn test_cache_clear() {
let mut cache = StatementCache::new(2);
cache.insert(dummy_stmt("SELECT 1", "s1"));
cache.insert(dummy_stmt("SELECT 2", "s2"));
let cleared = cache.clear();
assert_eq!(cleared.len(), 2);
assert!(cache.is_empty());
}
}