1use std::collections::{HashMap, VecDeque};
35use std::sync::{Arc, Mutex};
36use std::time::{Duration, Instant};
37
38use sha2::{Digest, Sha256};
39
40fn compute_key(model_id: &str, prompt: &str) -> String {
44 let mut hasher = Sha256::new();
45 hasher.update(model_id.as_bytes());
46 hasher.update(b"\x00"); hasher.update(prompt.as_bytes());
48 let digest = hasher.finalize();
49 hex::encode(digest)
50}
51
52#[derive(Debug, Clone)]
56pub struct CacheConfig {
57 pub max_entries: usize,
62 pub default_ttl: Duration,
64 pub max_prompt_len: usize,
69}
70
71impl Default for CacheConfig {
72 fn default() -> Self {
73 Self {
74 max_entries: 1_024,
75 default_ttl: Duration::from_secs(300),
76 max_prompt_len: 16_384,
77 }
78 }
79}
80
81#[derive(Debug, Clone)]
85pub struct CacheEntry {
86 pub response: Vec<String>,
88 pub created_at: Instant,
90 pub hits: u64,
92 pub ttl: Duration,
94}
95
96impl CacheEntry {
97 pub fn is_expired(&self) -> bool {
99 self.created_at.elapsed() >= self.ttl
100 }
101}
102
103#[derive(Debug, Clone, Default)]
107pub struct CacheStats {
108 pub entries: usize,
110 pub total_hits: u64,
112 pub total_misses: u64,
114 pub hit_rate: f64,
117 pub evictions: u64,
120}
121
122struct CacheInner {
125 map: HashMap<String, CacheEntry>,
127 order: VecDeque<String>,
129 config: CacheConfig,
130 total_hits: u64,
131 total_misses: u64,
132 evictions: u64,
133}
134
135impl CacheInner {
136 fn new(config: CacheConfig) -> Self {
137 Self {
138 map: HashMap::new(),
139 order: VecDeque::new(),
140 config,
141 total_hits: 0,
142 total_misses: 0,
143 evictions: 0,
144 }
145 }
146
147 fn touch(&mut self, key: &str) {
149 if let Some(pos) = self.order.iter().position(|k| k == key) {
150 self.order.remove(pos);
151 }
152 self.order.push_back(key.to_owned());
153 }
154
155 fn evict_lru_to_fit(&mut self) {
157 while self.map.len() >= self.config.max_entries {
158 if let Some(lru_key) = self.order.pop_front() {
159 if self.map.remove(&lru_key).is_some() {
160 self.evictions += 1;
161 }
162 } else {
163 break;
164 }
165 }
166 }
167
168 fn get(&mut self, key: &str) -> Option<Vec<String>> {
169 if let Some(entry) = self.map.get(key) {
171 if entry.is_expired() {
172 self.map.remove(key);
173 if let Some(pos) = self.order.iter().position(|k| k == key) {
174 self.order.remove(pos);
175 }
176 self.evictions += 1;
177 self.total_misses += 1;
178 return None;
179 }
180 }
181 if let Some(entry) = self.map.get_mut(key) {
182 entry.hits += 1;
183 self.total_hits += 1;
184 let response = entry.response.clone();
185 self.touch(key);
186 Some(response)
187 } else {
188 self.total_misses += 1;
189 None
190 }
191 }
192
193 fn insert(&mut self, key: String, response: Vec<String>, ttl: Duration) {
194 if self.map.contains_key(&key) {
195 if let Some(entry) = self.map.get_mut(&key) {
197 entry.response = response;
198 entry.created_at = Instant::now();
199 entry.ttl = ttl;
200 }
201 self.touch(&key);
202 return;
203 }
204 self.evict_lru_to_fit();
205 self.order.push_back(key.clone());
206 self.map.insert(
207 key,
208 CacheEntry {
209 response,
210 created_at: Instant::now(),
211 hits: 0,
212 ttl,
213 },
214 );
215 }
216
217 fn evict_expired(&mut self) -> usize {
218 let expired: Vec<String> = self
219 .map
220 .iter()
221 .filter(|(_, e)| e.is_expired())
222 .map(|(k, _)| k.clone())
223 .collect();
224 let count = expired.len();
225 for key in &expired {
226 self.map.remove(key);
227 if let Some(pos) = self.order.iter().position(|k| k == key) {
228 self.order.remove(pos);
229 }
230 self.evictions += 1;
231 }
232 count
233 }
234
235 fn stats(&self) -> CacheStats {
236 let total = self.total_hits + self.total_misses;
237 CacheStats {
238 entries: self.map.len(),
239 total_hits: self.total_hits,
240 total_misses: self.total_misses,
241 hit_rate: if total == 0 {
242 0.0
243 } else {
244 self.total_hits as f64 / total as f64
245 },
246 evictions: self.evictions,
247 }
248 }
249
250 fn flush(&mut self) {
251 let count = self.map.len();
252 self.map.clear();
253 self.order.clear();
254 self.evictions += count as u64;
255 }
256}
257
258#[derive(Clone)]
264pub struct PromptCache {
265 inner: Arc<Mutex<CacheInner>>,
266}
267
268impl PromptCache {
269 pub fn new(config: CacheConfig) -> Self {
271 Self {
272 inner: Arc::new(Mutex::new(CacheInner::new(config))),
273 }
274 }
275
276 pub fn get(&self, model_id: &str, prompt: &str) -> Option<Vec<String>> {
283 let max_len = {
284 let guard = self.inner.lock().ok()?;
285 guard.config.max_prompt_len
286 };
287 if prompt.len() > max_len {
288 return None;
289 }
290 let key = compute_key(model_id, prompt);
291 self.inner.lock().ok()?.get(&key)
292 }
293
294 pub fn insert(
300 &self,
301 model_id: &str,
302 prompt: &str,
303 response: Vec<String>,
304 ttl: Option<Duration>,
305 ) {
306 let mut guard = match self.inner.lock() {
307 Ok(g) => g,
308 Err(_) => return,
309 };
310 if prompt.len() > guard.config.max_prompt_len {
311 return;
312 }
313 let ttl = ttl.unwrap_or(guard.config.default_ttl);
314 let key = compute_key(model_id, prompt);
315 guard.insert(key, response, ttl);
316 }
317
318 pub fn evict_expired(&self) -> usize {
322 self.inner
323 .lock()
324 .map(|mut g| g.evict_expired())
325 .unwrap_or(0)
326 }
327
328 pub fn stats(&self) -> CacheStats {
330 self.inner
331 .lock()
332 .map(|g| g.stats())
333 .unwrap_or_default()
334 }
335
336 pub fn flush(&self) {
338 if let Ok(mut g) = self.inner.lock() {
339 g.flush();
340 }
341 }
342}
343
344#[cfg(test)]
347#[allow(clippy::unwrap_used, clippy::expect_used)]
348mod tests {
349 use super::*;
350 use std::time::Duration;
351
352 fn make_cache(max: usize) -> PromptCache {
353 PromptCache::new(CacheConfig {
354 max_entries: max,
355 default_ttl: Duration::from_secs(60),
356 max_prompt_len: 1024,
357 })
358 }
359
360 #[test]
361 fn test_basic_insert_and_get() {
362 let cache = make_cache(10);
363 cache.insert("m1", "hello", vec!["world".to_string()], None);
364 let result = cache.get("m1", "hello");
365 assert_eq!(result, Some(vec!["world".to_string()]));
366 }
367
368 #[test]
369 fn test_miss_returns_none() {
370 let cache = make_cache(10);
371 assert!(cache.get("m1", "missing").is_none());
372 }
373
374 #[test]
375 fn test_different_models_different_keys() {
376 let cache = make_cache(10);
377 cache.insert("m1", "p", vec!["r1".to_string()], None);
378 cache.insert("m2", "p", vec!["r2".to_string()], None);
379 assert_eq!(cache.get("m1", "p"), Some(vec!["r1".to_string()]));
380 assert_eq!(cache.get("m2", "p"), Some(vec!["r2".to_string()]));
381 }
382
383 #[test]
384 fn test_hit_counter_increments() {
385 let cache = make_cache(10);
386 cache.insert("m", "p", vec!["r".to_string()], None);
387 cache.get("m", "p");
388 cache.get("m", "p");
389 let stats = cache.stats();
390 assert_eq!(stats.total_hits, 2);
391 assert_eq!(stats.total_misses, 0);
392 }
393
394 #[test]
395 fn test_miss_counter_increments() {
396 let cache = make_cache(10);
397 cache.get("m", "not_there");
398 let stats = cache.stats();
399 assert_eq!(stats.total_misses, 1);
400 assert_eq!(stats.total_hits, 0);
401 }
402
403 #[test]
404 fn test_hit_rate_calculation() {
405 let cache = make_cache(10);
406 cache.insert("m", "p", vec![], None);
407 cache.get("m", "p"); cache.get("m", "miss"); let stats = cache.stats();
410 assert!((stats.hit_rate - 0.5).abs() < f64::EPSILON);
411 }
412
413 #[test]
414 fn test_lru_eviction_when_full() {
415 let cache = make_cache(3);
416 cache.insert("m", "a", vec!["a".to_string()], None);
417 cache.insert("m", "b", vec!["b".to_string()], None);
418 cache.insert("m", "c", vec!["c".to_string()], None);
419 cache.get("m", "a");
421 cache.insert("m", "d", vec!["d".to_string()], None);
423 assert!(cache.get("m", "b").is_none()); assert!(cache.get("m", "a").is_some());
425 assert!(cache.get("m", "c").is_some());
426 assert!(cache.get("m", "d").is_some());
427 }
428
429 #[test]
430 fn test_eviction_counter() {
431 let cache = make_cache(2);
432 cache.insert("m", "a", vec![], None);
433 cache.insert("m", "b", vec![], None);
434 cache.insert("m", "c", vec![], None); let stats = cache.stats();
436 assert_eq!(stats.evictions, 1);
437 }
438
439 #[test]
440 fn test_ttl_expiry_on_get() {
441 let cache = PromptCache::new(CacheConfig {
442 max_entries: 10,
443 default_ttl: Duration::from_millis(1),
444 max_prompt_len: 1024,
445 });
446 cache.insert("m", "p", vec!["r".to_string()], None);
447 std::thread::sleep(Duration::from_millis(5));
448 assert!(cache.get("m", "p").is_none());
449 }
450
451 #[test]
452 fn test_evict_expired_removes_stale() {
453 let cache = PromptCache::new(CacheConfig {
454 max_entries: 10,
455 default_ttl: Duration::from_millis(1),
456 max_prompt_len: 1024,
457 });
458 cache.insert("m", "a", vec![], None);
459 cache.insert("m", "b", vec![], None);
460 std::thread::sleep(Duration::from_millis(5));
461 let removed = cache.evict_expired();
462 assert_eq!(removed, 2);
463 assert_eq!(cache.stats().entries, 0);
464 }
465
466 #[test]
467 fn test_prompt_too_long_not_cached() {
468 let cache = PromptCache::new(CacheConfig {
469 max_entries: 10,
470 default_ttl: Duration::from_secs(60),
471 max_prompt_len: 5,
472 });
473 cache.insert("m", "toolongprompt", vec!["r".to_string()], None);
474 assert!(cache.get("m", "toolongprompt").is_none());
475 }
476
477 #[test]
478 fn test_flush_clears_all() {
479 let cache = make_cache(10);
480 cache.insert("m", "a", vec![], None);
481 cache.insert("m", "b", vec![], None);
482 cache.flush();
483 assert_eq!(cache.stats().entries, 0);
484 }
485
486 #[test]
487 fn test_flush_increments_evictions() {
488 let cache = make_cache(10);
489 cache.insert("m", "a", vec![], None);
490 cache.insert("m", "b", vec![], None);
491 cache.flush();
492 assert_eq!(cache.stats().evictions, 2);
493 }
494
495 #[test]
496 fn test_update_existing_entry() {
497 let cache = make_cache(10);
498 cache.insert("m", "p", vec!["old".to_string()], None);
499 cache.insert("m", "p", vec!["new".to_string()], None);
500 assert_eq!(cache.get("m", "p"), Some(vec!["new".to_string()]));
501 assert_eq!(cache.stats().entries, 1);
503 }
504
505 #[test]
506 fn test_custom_ttl_overrides_default() {
507 let cache = PromptCache::new(CacheConfig {
508 max_entries: 10,
509 default_ttl: Duration::from_secs(3600), max_prompt_len: 1024,
511 });
512 cache.insert("m", "p", vec!["r".to_string()], Some(Duration::from_millis(1)));
514 std::thread::sleep(Duration::from_millis(5));
515 assert!(cache.get("m", "p").is_none()); }
517
518 #[test]
519 fn test_clone_shares_state() {
520 let cache = make_cache(10);
521 let clone = cache.clone();
522 cache.insert("m", "p", vec!["r".to_string()], None);
523 assert_eq!(clone.get("m", "p"), Some(vec!["r".to_string()]));
524 }
525
526 #[test]
527 fn test_stats_default_zero() {
528 let cache = make_cache(10);
529 let stats = cache.stats();
530 assert_eq!(stats.entries, 0);
531 assert_eq!(stats.total_hits, 0);
532 assert_eq!(stats.total_misses, 0);
533 assert_eq!(stats.evictions, 0);
534 assert_eq!(stats.hit_rate, 0.0);
535 }
536
537 #[test]
538 fn test_compute_key_deterministic() {
539 let k1 = compute_key("model", "prompt");
540 let k2 = compute_key("model", "prompt");
541 assert_eq!(k1, k2);
542 }
543
544 #[test]
545 fn test_compute_key_differs_by_model() {
546 let k1 = compute_key("model-a", "prompt");
547 let k2 = compute_key("model-b", "prompt");
548 assert_ne!(k1, k2);
549 }
550}