tokio_prompt_orchestrator/
semantic_cache.rs1use std::collections::HashMap;
11use std::time::Instant;
12
13const DIM: usize = 64;
16
17fn 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
27pub 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 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
51pub 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
65pub struct CacheEntry {
69 pub prompt: String,
71 pub response: String,
73 pub embedding: Vec<f32>,
75 pub hit_count: u64,
77 pub timestamp: Instant,
79}
80
81#[derive(Debug, Clone, Default)]
85pub struct CacheStats {
86 pub hits: u64,
88 pub misses: u64,
90 pub evictions: u64,
92 pub size: usize,
94}
95
96pub struct SemanticCache {
105 capacity: usize,
106 threshold: f64,
107 order: Vec<String>,
109 entries: HashMap<String, CacheEntry>,
111 hits: u64,
113 misses: u64,
114 evictions: u64,
115}
116
117impl SemanticCache {
118 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 pub fn lookup(&mut self, prompt: &str) -> Option<&str> {
138 let query_emb = embed_prompt(prompt);
139
140 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 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 Some(self.entries[&key].response.as_str())
164 } else {
165 self.misses += 1;
166 None
167 }
168 }
169
170 pub fn insert(&mut self, prompt: &str, response: String) {
175 let key = prompt.to_string();
176
177 if self.entries.contains_key(&key) {
178 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 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 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#[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 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 let _ = cache.lookup("prompt one");
268 cache.insert("prompt three", "resp3".to_string());
270 assert_eq!(cache.stats().evictions, 1);
271 assert_eq!(cache.stats().size, 2);
272 assert!(cache.lookup("prompt two").is_none());
274 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 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 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 assert_eq!(cache.stats().size, 1);
334 let r = cache.lookup("hello world").unwrap();
335 assert_eq!(r, "v2");
336 }
337}