use crate::dictionary::Term;
use crate::error::Result;
use dashmap::DashMap;
use std::collections::VecDeque;
use std::hash::{Hash, Hasher};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct QueryCacheConfig {
pub max_entries: usize,
pub ttl: Duration,
pub enabled: bool,
pub max_result_size: usize,
}
impl Default for QueryCacheConfig {
fn default() -> Self {
Self {
max_entries: 1000,
ttl: Duration::from_secs(300), enabled: true,
max_result_size: 10000, }
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct QueryPattern {
pub subject: Option<String>,
pub predicate: Option<String>,
pub object: Option<String>,
}
impl QueryPattern {
pub fn new(subject: Option<&Term>, predicate: Option<&Term>, object: Option<&Term>) -> Self {
Self {
subject: subject.map(|t| format!("{:?}", t)),
predicate: predicate.map(|t| format!("{:?}", t)),
object: object.map(|t| format!("{:?}", t)),
}
}
pub fn is_cacheable(&self) -> bool {
self.subject.is_some() || self.predicate.is_some() || self.object.is_some()
}
}
#[derive(Debug, Clone)]
struct CachedResult {
results: Vec<(Term, Term, Term)>,
cached_at: Instant,
access_count: u64,
last_accessed: Instant,
}
impl CachedResult {
fn new(results: Vec<(Term, Term, Term)>) -> Self {
let now = Instant::now();
Self {
results,
cached_at: now,
access_count: 0,
last_accessed: now,
}
}
fn is_expired(&self, ttl: Duration) -> bool {
self.cached_at.elapsed() > ttl
}
fn access(&mut self) {
self.access_count += 1;
self.last_accessed = Instant::now();
}
}
#[derive(Debug, Clone)]
struct LruEntry {
pattern: QueryPattern,
last_accessed: Instant,
}
pub struct QueryCache {
config: QueryCacheConfig,
cache: Arc<DashMap<QueryPattern, CachedResult>>,
lru_queue: parking_lot::Mutex<VecDeque<LruEntry>>,
stats: QueryCacheStats,
}
#[derive(Debug, Default)]
pub struct QueryCacheStats {
pub hits: AtomicU64,
pub misses: AtomicU64,
pub evictions: AtomicU64,
pub invalidations: AtomicU64,
pub current_size: AtomicUsize,
}
impl QueryCacheStats {
pub fn hit_rate(&self) -> f64 {
let hits = self.hits.load(Ordering::Relaxed) as f64;
let misses = self.misses.load(Ordering::Relaxed) as f64;
let total = hits + misses;
if total == 0.0 {
0.0
} else {
hits / total
}
}
pub fn total_requests(&self) -> u64 {
self.hits.load(Ordering::Relaxed) + self.misses.load(Ordering::Relaxed)
}
}
impl QueryCache {
pub fn new(config: QueryCacheConfig) -> Self {
Self {
config,
cache: Arc::new(DashMap::new()),
lru_queue: parking_lot::Mutex::new(VecDeque::new()),
stats: QueryCacheStats::default(),
}
}
pub fn get(&self, pattern: &QueryPattern) -> Option<Vec<(Term, Term, Term)>> {
if !self.config.enabled || !pattern.is_cacheable() {
return None;
}
if let Some(mut entry) = self.cache.get_mut(pattern) {
if entry.is_expired(self.config.ttl) {
drop(entry);
self.cache.remove(pattern);
self.stats.misses.fetch_add(1, Ordering::Relaxed);
self.stats.current_size.fetch_sub(1, Ordering::Relaxed);
return None;
}
entry.access();
self.stats.hits.fetch_add(1, Ordering::Relaxed);
self.update_lru(pattern);
return Some(entry.results.clone());
}
self.stats.misses.fetch_add(1, Ordering::Relaxed);
None
}
pub fn put(&self, pattern: QueryPattern, results: Vec<(Term, Term, Term)>) -> Result<()> {
if !self.config.enabled || !pattern.is_cacheable() {
return Ok(());
}
if results.len() > self.config.max_result_size {
return Ok(());
}
while self.cache.len() >= self.config.max_entries {
self.evict_lru()?;
}
let entry = CachedResult::new(results);
self.cache.insert(pattern.clone(), entry);
self.stats.current_size.fetch_add(1, Ordering::Relaxed);
let mut lru = self.lru_queue.lock();
lru.push_back(LruEntry {
pattern,
last_accessed: Instant::now(),
});
Ok(())
}
pub fn invalidate_all(&self) {
let count = self.cache.len();
self.cache.clear();
self.lru_queue.lock().clear();
self.stats.current_size.store(0, Ordering::Relaxed);
self.stats
.invalidations
.fetch_add(count as u64, Ordering::Relaxed);
}
pub fn invalidate_pattern(
&self,
subject: Option<&str>,
predicate: Option<&str>,
object: Option<&str>,
) {
let mut invalidated = 0;
self.cache.retain(|pattern, _| {
let should_keep = !self.pattern_overlaps(pattern, subject, predicate, object);
if !should_keep {
invalidated += 1;
}
should_keep
});
self.stats
.current_size
.fetch_sub(invalidated, Ordering::Relaxed);
self.stats
.invalidations
.fetch_add(invalidated as u64, Ordering::Relaxed);
let mut lru = self.lru_queue.lock();
lru.retain(|entry| self.cache.contains_key(&entry.pattern));
}
fn pattern_overlaps(
&self,
cached: &QueryPattern,
subject: Option<&str>,
predicate: Option<&str>,
object: Option<&str>,
) -> bool {
let s_overlaps = match (&cached.subject, subject) {
(Some(cs), Some(s)) => cs == s,
_ => true, };
let p_overlaps = match (&cached.predicate, predicate) {
(Some(cp), Some(p)) => cp == p,
_ => true,
};
let o_overlaps = match (&cached.object, object) {
(Some(co), Some(o)) => co == o,
_ => true,
};
s_overlaps && p_overlaps && o_overlaps
}
fn evict_lru(&self) -> Result<()> {
let mut lru = self.lru_queue.lock();
while let Some(entry) = lru.pop_front() {
if self.cache.remove(&entry.pattern).is_some() {
self.stats.evictions.fetch_add(1, Ordering::Relaxed);
self.stats.current_size.fetch_sub(1, Ordering::Relaxed);
return Ok(());
}
}
Ok(())
}
fn update_lru(&self, pattern: &QueryPattern) {
let mut lru = self.lru_queue.lock();
if let Some(pos) = lru.iter().position(|e| &e.pattern == pattern) {
if let Some(mut entry) = lru.remove(pos) {
entry.last_accessed = Instant::now();
lru.push_back(entry);
}
}
}
pub fn stats(&self) -> &QueryCacheStats {
&self.stats
}
pub fn clear(&self) {
self.cache.clear();
self.lru_queue.lock().clear();
self.stats.current_size.store(0, Ordering::Relaxed);
}
pub fn len(&self) -> usize {
self.cache.len()
}
pub fn is_empty(&self) -> bool {
self.cache.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dictionary::Term;
fn create_test_pattern(s: Option<&str>, p: Option<&str>, o: Option<&str>) -> QueryPattern {
QueryPattern {
subject: s.map(String::from),
predicate: p.map(String::from),
object: o.map(String::from),
}
}
fn create_test_results(count: usize) -> Vec<(Term, Term, Term)> {
(0..count)
.map(|i| {
(
Term::Iri(format!("http://example.org/s{}", i)),
Term::Iri("http://example.org/knows".to_string()),
Term::Iri(format!("http://example.org/o{}", i)),
)
})
.collect()
}
#[test]
fn test_query_cache_creation() {
let config = QueryCacheConfig::default();
let cache = QueryCache::new(config);
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
}
#[test]
fn test_cache_put_and_get() {
let cache = QueryCache::new(QueryCacheConfig::default());
let pattern = create_test_pattern(Some("s1"), None, None);
let results = create_test_results(5);
cache.put(pattern.clone(), results.clone()).unwrap();
let cached = cache.get(&pattern);
assert!(cached.is_some());
assert_eq!(cached.unwrap().len(), 5);
}
#[test]
fn test_cache_miss() {
let cache = QueryCache::new(QueryCacheConfig::default());
let pattern = create_test_pattern(Some("s1"), None, None);
let cached = cache.get(&pattern);
assert!(cached.is_none());
let stats = cache.stats();
assert_eq!(stats.misses.load(Ordering::Relaxed), 1);
}
#[test]
fn test_cache_hit() {
let cache = QueryCache::new(QueryCacheConfig::default());
let pattern = create_test_pattern(Some("s1"), None, None);
let results = create_test_results(5);
cache.put(pattern.clone(), results).unwrap();
cache.get(&pattern);
let stats = cache.stats();
assert_eq!(stats.hits.load(Ordering::Relaxed), 1);
assert!(stats.hit_rate() > 0.0);
}
#[test]
fn test_cache_expiration() {
let config = QueryCacheConfig {
ttl: Duration::from_millis(10),
..Default::default()
};
let cache = QueryCache::new(config);
let pattern = create_test_pattern(Some("s1"), None, None);
let results = create_test_results(5);
cache.put(pattern.clone(), results).unwrap();
std::thread::sleep(Duration::from_millis(20));
let cached = cache.get(&pattern);
assert!(cached.is_none()); }
#[test]
fn test_lru_eviction() {
let config = QueryCacheConfig {
max_entries: 3,
..Default::default()
};
let cache = QueryCache::new(config);
for i in 0..3 {
let pattern = create_test_pattern(Some(&format!("s{}", i)), None, None);
cache.put(pattern, create_test_results(1)).unwrap();
}
assert_eq!(cache.len(), 3);
let pattern4 = create_test_pattern(Some("s4"), None, None);
cache.put(pattern4, create_test_results(1)).unwrap();
assert_eq!(cache.len(), 3);
let stats = cache.stats();
assert_eq!(stats.evictions.load(Ordering::Relaxed), 1);
}
#[test]
fn test_invalidate_all() {
let cache = QueryCache::new(QueryCacheConfig::default());
for i in 0..5 {
let pattern = create_test_pattern(Some(&format!("s{}", i)), None, None);
cache.put(pattern, create_test_results(1)).unwrap();
}
assert_eq!(cache.len(), 5);
cache.invalidate_all();
assert_eq!(cache.len(), 0);
let stats = cache.stats();
assert_eq!(stats.invalidations.load(Ordering::Relaxed), 5);
}
#[test]
fn test_invalidate_pattern() {
let cache = QueryCache::new(QueryCacheConfig::default());
cache
.put(
create_test_pattern(Some("s1"), Some("p1"), None),
create_test_results(1),
)
.unwrap();
cache
.put(
create_test_pattern(Some("s2"), Some("p1"), None),
create_test_results(1),
)
.unwrap();
cache
.put(
create_test_pattern(Some("s3"), Some("p2"), None),
create_test_results(1),
)
.unwrap();
assert_eq!(cache.len(), 3);
cache.invalidate_pattern(None, Some("p1"), None);
assert_eq!(cache.len(), 1); }
#[test]
fn test_max_result_size() {
let config = QueryCacheConfig {
max_result_size: 10,
..Default::default()
};
let cache = QueryCache::new(config);
let pattern = create_test_pattern(Some("s1"), None, None);
let large_results = create_test_results(100);
cache.put(pattern.clone(), large_results).unwrap();
assert_eq!(cache.len(), 0);
}
#[test]
fn test_pattern_is_cacheable() {
let pattern1 = create_test_pattern(None, None, None);
assert!(!pattern1.is_cacheable());
let pattern2 = create_test_pattern(Some("s1"), None, None);
assert!(pattern2.is_cacheable());
let pattern3 = create_test_pattern(None, Some("p1"), None);
assert!(pattern3.is_cacheable());
let pattern4 = create_test_pattern(None, None, Some("o1"));
assert!(pattern4.is_cacheable());
}
#[test]
fn test_cache_disabled() {
let config = QueryCacheConfig {
enabled: false,
..Default::default()
};
let cache = QueryCache::new(config);
let pattern = create_test_pattern(Some("s1"), None, None);
let results = create_test_results(5);
cache.put(pattern.clone(), results).unwrap();
assert_eq!(cache.len(), 0);
let cached = cache.get(&pattern);
assert!(cached.is_none());
}
#[test]
fn test_hit_rate_calculation() {
let cache = QueryCache::new(QueryCacheConfig::default());
let pattern = create_test_pattern(Some("s1"), None, None);
let results = create_test_results(5);
cache.put(pattern.clone(), results).unwrap();
cache.get(&pattern);
cache.get(&pattern);
cache.get(&pattern);
cache.get(&create_test_pattern(Some("s2"), None, None));
cache.get(&create_test_pattern(Some("s3"), None, None));
let stats = cache.stats();
assert_eq!(stats.total_requests(), 5);
assert_eq!(stats.hit_rate(), 0.6); }
#[test]
fn test_clear() {
let cache = QueryCache::new(QueryCacheConfig::default());
for i in 0..5 {
let pattern = create_test_pattern(Some(&format!("s{}", i)), None, None);
cache.put(pattern, create_test_results(1)).unwrap();
}
assert_eq!(cache.len(), 5);
cache.clear();
assert_eq!(cache.len(), 0);
assert_eq!(cache.stats().current_size.load(Ordering::Relaxed), 0);
}
}