Skip to main content

tokio_prompt_orchestrator/
cache.rs

1//! # Prompt Cache
2//!
3//! Content-addressed, in-process LRU cache for LLM inference responses.
4//!
5//! Cache keys are the SHA-256 hash of `(model_id + prompt_text)`.  Entries
6//! carry a TTL; expired entries are evicted lazily on [`PromptCache::get`] and
7//! eagerly via [`PromptCache::evict_expired`].  When the cache reaches
8//! [`CacheConfig::max_entries`] the least-recently-used entry is evicted to
9//! make room (LRU via `VecDeque` order tracking + `HashMap` for O(1) lookup).
10//!
11//! The cache is fully thread-safe: the public handle is `Arc<Mutex<CacheInner>>`.
12//!
13//! ## Example
14//!
15//! ```
16//! use tokio_prompt_orchestrator::cache::{CacheConfig, PromptCache};
17//! use std::time::Duration;
18//!
19//! let cfg = CacheConfig {
20//!     max_entries: 100,
21//!     default_ttl: Duration::from_secs(300),
22//!     max_prompt_len: 8192,
23//! };
24//! let cache = PromptCache::new(cfg);
25//!
26//! // Cache a response.
27//! cache.insert("gpt-4o", "Hello, world!", vec!["Hi there!".to_string()], None);
28//!
29//! // Retrieve it.
30//! let result = cache.get("gpt-4o", "Hello, world!");
31//! assert!(result.is_some());
32//! ```
33
34use std::collections::{HashMap, VecDeque};
35use std::sync::{Arc, Mutex};
36use std::time::{Duration, Instant};
37
38use sha2::{Digest, Sha256};
39
40// ── Key helpers ───────────────────────────────────────────────────────────────
41
42/// Compute the SHA-256 cache key for `(model_id, prompt)`.
43fn 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"); // separator
47    hasher.update(prompt.as_bytes());
48    let digest = hasher.finalize();
49    hex::encode(digest)
50}
51
52// ── Config ────────────────────────────────────────────────────────────────────
53
54/// Configuration for [`PromptCache`].
55#[derive(Debug, Clone)]
56pub struct CacheConfig {
57    /// Maximum number of entries retained in the cache.
58    ///
59    /// When this limit is reached, the least-recently-used entry is evicted
60    /// before a new one is inserted.
61    pub max_entries: usize,
62    /// TTL applied to entries that do not supply their own TTL on insertion.
63    pub default_ttl: Duration,
64    /// Prompts longer than this (in bytes) are never cached.
65    ///
66    /// This prevents the cache from holding very large strings in memory.
67    /// Set to `usize::MAX` to disable the limit.
68    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// ── Entry ─────────────────────────────────────────────────────────────────────
82
83/// A single cached inference response.
84#[derive(Debug, Clone)]
85pub struct CacheEntry {
86    /// The cached response tokens/chunks.
87    pub response: Vec<String>,
88    /// Instant the entry was first inserted.
89    pub created_at: Instant,
90    /// Number of times this entry has been returned on a cache hit.
91    pub hits: u64,
92    /// How long this entry lives before it is considered expired.
93    pub ttl: Duration,
94}
95
96impl CacheEntry {
97    /// Return `true` when the entry has lived past its TTL.
98    pub fn is_expired(&self) -> bool {
99        self.created_at.elapsed() >= self.ttl
100    }
101}
102
103// ── Stats ─────────────────────────────────────────────────────────────────────
104
105/// Aggregate statistics for a [`PromptCache`] instance.
106#[derive(Debug, Clone, Default)]
107pub struct CacheStats {
108    /// Current number of live (non-expired) entries in the cache.
109    pub entries: usize,
110    /// Cumulative number of cache hits since the cache was created.
111    pub total_hits: u64,
112    /// Cumulative number of cache misses since the cache was created.
113    pub total_misses: u64,
114    /// Hit rate: `total_hits / (total_hits + total_misses)`, or `0.0` when no
115    /// requests have been made.
116    pub hit_rate: f64,
117    /// Number of entries evicted (LRU or TTL expiry) since the cache was
118    /// created.
119    pub evictions: u64,
120}
121
122// ── Inner state ───────────────────────────────────────────────────────────────
123
124struct CacheInner {
125    /// Map from SHA-256 key string → entry.
126    map: HashMap<String, CacheEntry>,
127    /// LRU order: front = least recently used, back = most recently used.
128    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    /// Move `key` to the back of the LRU queue (most recently used).
148    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    /// Evict LRU entries until `map.len() < max_entries`.
156    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        // Check for expiry first.
170        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            // Update in place, reset TTL.
196            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// ── Public handle ─────────────────────────────────────────────────────────────
259
260/// Thread-safe, content-addressed LRU prompt cache.
261///
262/// Clone-cheap: all clones share the same underlying data.
263#[derive(Clone)]
264pub struct PromptCache {
265    inner: Arc<Mutex<CacheInner>>,
266}
267
268impl PromptCache {
269    /// Create a new cache with the given configuration.
270    pub fn new(config: CacheConfig) -> Self {
271        Self {
272            inner: Arc::new(Mutex::new(CacheInner::new(config))),
273        }
274    }
275
276    /// Look up a cached response for `(model_id, prompt)`.
277    ///
278    /// Returns `None` on a cache miss or if the entry has expired.
279    /// Updates the LRU order and hit counter on a hit.
280    ///
281    /// Prompts longer than `max_prompt_len` are always treated as misses.
282    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    /// Insert a response for `(model_id, prompt)`.
295    ///
296    /// - `ttl`: when `None`, the cache's `default_ttl` is used.
297    /// - Prompts longer than `max_prompt_len` are silently dropped.
298    /// - If the cache is full the LRU entry is evicted first.
299    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    /// Evict all entries whose TTL has elapsed.
319    ///
320    /// Returns the number of entries removed.
321    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    /// Return a snapshot of current cache statistics.
329    pub fn stats(&self) -> CacheStats {
330        self.inner
331            .lock()
332            .map(|g| g.stats())
333            .unwrap_or_default()
334    }
335
336    /// Remove all entries from the cache.
337    pub fn flush(&self) {
338        if let Ok(mut g) = self.inner.lock() {
339            g.flush();
340        }
341    }
342}
343
344// ── Tests ─────────────────────────────────────────────────────────────────────
345
346#[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"); // hit
408        cache.get("m", "miss"); // miss
409        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        // Touch "a" so it becomes most recently used.
420        cache.get("m", "a");
421        // Insert "d" — "b" should be evicted (LRU).
422        cache.insert("m", "d", vec!["d".to_string()], None);
423        assert!(cache.get("m", "b").is_none()); // evicted
424        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); // evicts "a"
435        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        // Should not grow beyond 1 entry.
502        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), // long default
510            max_prompt_len: 1024,
511        });
512        // Short custom TTL.
513        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()); // expired via custom TTL
516    }
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}