Skip to main content

tokio_prompt_orchestrator/
semantic_cache.rs

1//! # Semantic Cache
2//!
3//! LRU-evicting similarity cache for prompt/response pairs.
4//!
5//! Prompts are embedded into a 64-dimensional TF-IDF-style vector using FNV-1a
6//! token hashing, then L2-normalised.  Cache lookup returns a stored response
7//! if the cosine similarity between the query embedding and any stored embedding
8//! is at or above the configured threshold.
9
10use std::collections::HashMap;
11use std::time::Instant;
12
13// ── Embedding ─────────────────────────────────────────────────────────────────
14
15const DIM: usize = 64;
16
17/// FNV-1a 32-bit hash of a byte slice.
18fn fnv1a(s: &str) -> u32 {
19    let mut hash: u32 = 2_166_136_261;
20    for byte in s.bytes() {
21        hash ^= u32::from(byte);
22        hash = hash.wrapping_mul(16_777_619);
23    }
24    hash
25}
26
27/// Produce a deterministic TF-IDF-style 64-dim embedding for `prompt`.
28///
29/// Each whitespace-separated token is hashed with FNV-1a and its count is
30/// accumulated into the dimension `hash % DIM`.  The resulting vector is then
31/// L2-normalised.
32pub fn embed_prompt(prompt: &str) -> Vec<f32> {
33    let mut vec = vec![0.0_f32; DIM];
34
35    for token in prompt.split_whitespace() {
36        let h = fnv1a(token) as usize % DIM;
37        vec[h] += 1.0;
38    }
39
40    // L2 normalise
41    let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
42    if norm > 0.0 {
43        for v in &mut vec {
44            *v /= norm;
45        }
46    }
47
48    vec
49}
50
51/// Cosine similarity between two equal-length vectors.
52///
53/// Returns 0.0 if either vector is zero-length.
54pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
55    debug_assert_eq!(a.len(), b.len());
56    let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
57    let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
58    let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
59    if na == 0.0 || nb == 0.0 {
60        return 0.0;
61    }
62    dot / (na * nb)
63}
64
65// ── CacheEntry ────────────────────────────────────────────────────────────────
66
67/// A single entry stored in the [`SemanticCache`].
68pub struct CacheEntry {
69    /// The original prompt text.
70    pub prompt: String,
71    /// The stored response.
72    pub response: String,
73    /// Embedding vector for the prompt.
74    pub embedding: Vec<f32>,
75    /// Number of times this entry has been returned as a cache hit.
76    pub hit_count: u64,
77    /// Wall-clock time this entry was inserted.
78    pub timestamp: Instant,
79}
80
81// ── CacheStats ────────────────────────────────────────────────────────────────
82
83/// Aggregate statistics for a [`SemanticCache`].
84#[derive(Debug, Clone, Default)]
85pub struct CacheStats {
86    /// Number of successful cache lookups.
87    pub hits: u64,
88    /// Number of failed cache lookups.
89    pub misses: u64,
90    /// Number of entries evicted due to capacity limits.
91    pub evictions: u64,
92    /// Current number of entries in the cache.
93    pub size: usize,
94}
95
96// ── SemanticCache ─────────────────────────────────────────────────────────────
97
98/// LRU-evicting semantic similarity cache.
99///
100/// Internally maintains an ordered list of keys (front = most recently used,
101/// back = least recently used) and a `HashMap` of entries.  On lookup the
102/// matching key is moved to the front; on insert an LRU eviction is performed
103/// when the cache is full.
104pub struct SemanticCache {
105    capacity: usize,
106    threshold: f64,
107    /// Insertion-order list: front = MRU, back = LRU.
108    order: Vec<String>,
109    /// Key → entry storage.
110    entries: HashMap<String, CacheEntry>,
111    /// Aggregate statistics.
112    hits: u64,
113    misses: u64,
114    evictions: u64,
115}
116
117impl SemanticCache {
118    /// Create a new `SemanticCache` with the given `capacity` and cosine
119    /// similarity `threshold` in `[0.0, 1.0]`.
120    pub fn new(capacity: usize, similarity_threshold: f64) -> Self {
121        Self {
122            capacity,
123            threshold: similarity_threshold,
124            order: Vec::with_capacity(capacity),
125            entries: HashMap::with_capacity(capacity),
126            hits: 0,
127            misses: 0,
128            evictions: 0,
129        }
130    }
131
132    /// Look up `prompt` in the cache.
133    ///
134    /// Returns `Some(&str)` with the stored response if any entry's embedding
135    /// has cosine similarity ≥ `threshold` with `prompt`'s embedding.
136    /// The best match is returned and promoted to MRU position.
137    pub fn lookup(&mut self, prompt: &str) -> Option<&str> {
138        let query_emb = embed_prompt(prompt);
139
140        // Find the best matching key.
141        let mut best_key: Option<String> = None;
142        let mut best_sim = -1.0_f32;
143
144        for (key, entry) in &self.entries {
145            let sim = cosine_similarity(&query_emb, &entry.embedding);
146            if sim > best_sim {
147                best_sim = sim;
148                best_key = Some(key.clone());
149            }
150        }
151
152        if best_sim >= self.threshold as f32 {
153            let key = best_key?;
154            // Move to front (MRU).
155            if let Some(pos) = self.order.iter().position(|k| k == &key) {
156                self.order.remove(pos);
157                self.order.insert(0, key.clone());
158            }
159            let entry = self.entries.get_mut(&key)?;
160            entry.hit_count += 1;
161            self.hits += 1;
162            // Return a reference to the response.
163            Some(self.entries[&key].response.as_str())
164        } else {
165            self.misses += 1;
166            None
167        }
168    }
169
170    /// Insert a prompt/response pair into the cache.
171    ///
172    /// If the cache is full the LRU entry is evicted first.
173    /// If `prompt` already exists (exact key match) it is updated in place.
174    pub fn insert(&mut self, prompt: &str, response: String) {
175        let key = prompt.to_string();
176
177        if self.entries.contains_key(&key) {
178            // Update existing entry and promote to MRU.
179            if let Some(entry) = self.entries.get_mut(&key) {
180                entry.response = response;
181                entry.timestamp = Instant::now();
182            }
183            if let Some(pos) = self.order.iter().position(|k| k == &key) {
184                self.order.remove(pos);
185                self.order.insert(0, key);
186            }
187            return;
188        }
189
190        if self.entries.len() >= self.capacity {
191            self.evict_lru();
192        }
193
194        let embedding = embed_prompt(prompt);
195        let entry = CacheEntry {
196            prompt: key.clone(),
197            response,
198            embedding,
199            hit_count: 0,
200            timestamp: Instant::now(),
201        };
202        self.entries.insert(key.clone(), entry);
203        self.order.insert(0, key);
204    }
205
206    /// Remove the least recently used entry from the cache.
207    pub fn evict_lru(&mut self) {
208        if let Some(lru_key) = self.order.pop() {
209            self.entries.remove(&lru_key);
210            self.evictions += 1;
211        }
212    }
213
214    /// Return a statistics snapshot.
215    pub fn stats(&self) -> CacheStats {
216        CacheStats {
217            hits: self.hits,
218            misses: self.misses,
219            evictions: self.evictions,
220            size: self.entries.len(),
221        }
222    }
223}
224
225// ── Tests ─────────────────────────────────────────────────────────────────────
226
227#[cfg(test)]
228#[allow(clippy::unwrap_used, clippy::expect_used)]
229mod tests {
230    use super::*;
231
232    #[test]
233    fn test_exact_match_lookup() {
234        let mut cache = SemanticCache::new(10, 0.99);
235        cache.insert("hello world", "response A".to_string());
236        let result = cache.lookup("hello world");
237        assert_eq!(result, Some("response A"));
238        assert_eq!(cache.stats().hits, 1);
239        assert_eq!(cache.stats().misses, 0);
240    }
241
242    #[test]
243    fn test_similar_prompt_retrieval() {
244        let mut cache = SemanticCache::new(10, 0.5);
245        cache.insert("the quick brown fox", "fox response".to_string());
246        // Overlapping tokens → similar embedding.
247        let result = cache.lookup("quick brown fox");
248        assert!(result.is_some(), "Expected similar prompt to hit cache");
249        assert_eq!(cache.stats().hits, 1);
250    }
251
252    #[test]
253    fn test_dissimilar_prompt_is_miss() {
254        let mut cache = SemanticCache::new(10, 0.99);
255        cache.insert("hello world", "response A".to_string());
256        let result = cache.lookup("completely different zzzz");
257        assert!(result.is_none());
258        assert_eq!(cache.stats().misses, 1);
259    }
260
261    #[test]
262    fn test_lru_eviction() {
263        let mut cache = SemanticCache::new(2, 1.0);
264        cache.insert("prompt one", "resp1".to_string());
265        cache.insert("prompt two", "resp2".to_string());
266        // Access "prompt one" → becomes MRU; "prompt two" is now LRU.
267        let _ = cache.lookup("prompt one");
268        // Insert a third entry → "prompt two" (LRU) should be evicted.
269        cache.insert("prompt three", "resp3".to_string());
270        assert_eq!(cache.stats().evictions, 1);
271        assert_eq!(cache.stats().size, 2);
272        // "prompt two" was evicted.
273        assert!(cache.lookup("prompt two").is_none());
274        // "prompt one" still present.
275        assert!(cache.lookup("prompt one").is_some());
276    }
277
278    #[test]
279    fn test_evict_lru_explicit() {
280        let mut cache = SemanticCache::new(3, 1.0);
281        cache.insert("a a a", "ra".to_string());
282        cache.insert("b b b", "rb".to_string());
283        assert_eq!(cache.stats().size, 2);
284        cache.evict_lru();
285        // LRU is "a a a" (inserted first, never accessed since).
286        assert_eq!(cache.stats().size, 1);
287        assert_eq!(cache.stats().evictions, 1);
288    }
289
290    #[test]
291    fn test_cosine_similarity_identical() {
292        let v = vec![1.0_f32, 0.0, 0.0];
293        assert!((cosine_similarity(&v, &v) - 1.0).abs() < 1e-6);
294    }
295
296    #[test]
297    fn test_cosine_similarity_orthogonal() {
298        let a = vec![1.0_f32, 0.0, 0.0];
299        let b = vec![0.0_f32, 1.0, 0.0];
300        assert!(cosine_similarity(&a, &b).abs() < 1e-6);
301    }
302
303    #[test]
304    fn test_cosine_similarity_known_value() {
305        // a = (3,4,0), b = (4,3,0); dot=24, |a|=5, |b|=5 → sim=24/25=0.96
306        let a = vec![3.0_f32, 4.0, 0.0];
307        let b = vec![4.0_f32, 3.0, 0.0];
308        let sim = cosine_similarity(&a, &b);
309        assert!((sim - 0.96).abs() < 1e-5, "sim={sim}");
310    }
311
312    #[test]
313    fn test_embed_prompt_is_normalised() {
314        let emb = embed_prompt("hello world foo bar");
315        let norm: f32 = emb.iter().map(|x| x * x).sum::<f32>().sqrt();
316        assert!((norm - 1.0).abs() < 1e-6, "norm={norm}");
317    }
318
319    #[test]
320    fn test_stats_size_tracks_inserts() {
321        let mut cache = SemanticCache::new(5, 1.0);
322        cache.insert("p1", "r1".to_string());
323        cache.insert("p2", "r2".to_string());
324        assert_eq!(cache.stats().size, 2);
325    }
326
327    #[test]
328    fn test_update_existing_entry() {
329        let mut cache = SemanticCache::new(5, 1.0);
330        cache.insert("hello world", "v1".to_string());
331        cache.insert("hello world", "v2".to_string());
332        // Should still be 1 entry.
333        assert_eq!(cache.stats().size, 1);
334        let r = cache.lookup("hello world").unwrap();
335        assert_eq!(r, "v2");
336    }
337}