Skip to main content

xz_search/cache/
memory.rs

1use async_trait::async_trait;
2use std::collections::HashMap;
3use std::time::{Duration, Instant};
4use tokio::sync::RwLock;
5
6use crate::traits::{CacheStats, SearchCache};
7use crate::types::{SearchItem, SearchResult};
8
9#[derive(Debug, Clone)]
10struct CacheEntry {
11    result: SearchResult,
12    expires_at: Instant,
13    last_accessed: Instant,
14}
15
16#[derive(Debug)]
17pub struct MemorySearchCache {
18    entries: RwLock<HashMap<String, CacheEntry>>,
19    max_entries: usize,
20    hits: RwLock<u64>,
21    misses: RwLock<u64>,
22}
23
24impl MemorySearchCache {
25    pub fn new(max_entries: usize) -> Self {
26        Self {
27            entries: RwLock::new(HashMap::new()),
28            max_entries,
29            hits: RwLock::new(0),
30            misses: RwLock::new(0),
31        }
32    }
33
34    async fn evict_expired(&self) {
35        let mut entries = self.entries.write().await;
36        let now = Instant::now();
37        entries.retain(|_, v| v.expires_at > now);
38    }
39
40    async fn evict_lru(&self) {
41        let mut entries = self.entries.write().await;
42        if entries.len() < self.max_entries {
43            return;
44        }
45
46        let now = Instant::now();
47        let mut oldest_key: Option<String> = None;
48        let mut oldest_time = now;
49
50        for (key, entry) in entries.iter() {
51            if entry.last_accessed < oldest_time {
52                oldest_time = entry.last_accessed;
53                oldest_key = Some(key.clone());
54            }
55        }
56
57        if let Some(key) = oldest_key {
58            entries.remove(&key);
59        }
60    }
61}
62
63#[async_trait]
64impl SearchCache for MemorySearchCache {
65    async fn get(&self, key: &str) -> Option<SearchResult> {
66        self.evict_expired().await;
67
68        let mut entries = self.entries.write().await;
69        if let Some(entry) = entries.get_mut(key) {
70            if entry.expires_at > Instant::now() {
71                entry.last_accessed = Instant::now();
72                *self.hits.write().await += 1;
73                let mut result = entry.result.clone();
74                result.cached = true;
75                return Some(result);
76            }
77        }
78
79        *self.misses.write().await += 1;
80        None
81    }
82
83    async fn set(&self, key: &str, result: &SearchResult, ttl: Duration) {
84        let now = Instant::now();
85        let mut entries = self.entries.write().await;
86
87        if entries.contains_key(key) {
88            entries.insert(
89                key.to_string(),
90                CacheEntry { result: result.clone(), expires_at: now + ttl, last_accessed: now },
91            );
92            return;
93        }
94
95        if entries.len() >= self.max_entries {
96            drop(entries);
97            self.evict_lru().await;
98            entries = self.entries.write().await;
99        }
100
101        entries.insert(
102            key.to_string(),
103            CacheEntry { result: result.clone(), expires_at: now + ttl, last_accessed: now },
104        );
105    }
106
107    async fn invalidate(&self, key: &str) {
108        self.entries.write().await.remove(key);
109    }
110
111    fn stats(&self) -> CacheStats {
112        let (size_bytes, entry_count) = match self.entries.try_read() {
113            Ok(entries) => {
114                let size_bytes: usize = entries
115                    .values()
116                    .map(|e| {
117                        e.result.query.len()
118                            + e.result
119                                .items
120                                .iter()
121                                .map(|i: &SearchItem| i.title.len() + i.url.len() + i.snippet.len())
122                                .sum::<usize>()
123                    })
124                    .sum();
125                (size_bytes as u64, entries.len())
126            }
127            Err(_) => (0, 0),
128        };
129
130        CacheStats {
131            hits: self.hits.try_read().map(|hits| *hits).unwrap_or(0),
132            misses: self.misses.try_read().map(|misses| *misses).unwrap_or(0),
133            size_bytes,
134            entry_count,
135        }
136    }
137}