Skip to main content

tokio_prompt_orchestrator/
output_cache.rs

1//! Semantic output deduplication cache with TTL and pluggable eviction policies.
2
3use std::collections::HashMap;
4
5/// Key identifying a cached prompt/model/temperature combination.
6#[derive(Debug, Clone, PartialEq, Eq, Hash)]
7pub struct CacheKey {
8    /// FNV-1a hash of the prompt text.
9    pub prompt_hash: u64,
10    /// Model identifier string.
11    pub model: String,
12    /// Temperature rounded to tenths and stored as `(temperature * 10) as u8`.
13    pub temperature_bucket: u8,
14}
15
16/// A single cached LLM response entry.
17#[derive(Debug, Clone)]
18pub struct CachedOutput {
19    /// The response text.
20    pub response: String,
21    /// Number of tokens consumed by the request.
22    pub tokens_used: u64,
23    /// Unix epoch milliseconds when this entry was created.
24    pub created_at_ms: u64,
25    /// Time-to-live in milliseconds.
26    pub ttl_ms: u64,
27    /// How many times this entry has been read.
28    pub access_count: u64,
29}
30
31impl CachedOutput {
32    /// Returns true if this entry has expired relative to `now_ms`.
33    pub fn is_expired(&self, now_ms: u64) -> bool {
34        now_ms >= self.created_at_ms + self.ttl_ms
35    }
36}
37
38/// Policy used to select which entry to evict when the cache is full.
39#[derive(Debug, Clone)]
40pub enum EvictionPolicy {
41    /// Evict the entry with the smallest `access_count` (approximated by insertion order for
42    /// ties) — true LRU would require timestamps, but access_count serves as a proxy here.
43    LRU,
44    /// Evict the entry with the lowest `access_count`.
45    LFU,
46    /// Evict the entry whose TTL expires soonest.
47    TTLFirst,
48    /// Evict a pseudo-random entry using a seeded LCG.
49    Random(u64),
50}
51
52/// Statistics snapshot returned by [`OutputCache::stats`].
53#[derive(Debug, Clone)]
54pub struct CacheStats {
55    /// Current number of entries.
56    pub size: usize,
57    /// Maximum number of entries before eviction.
58    pub capacity: usize,
59    /// Total cache hits.
60    pub hits: u64,
61    /// Total cache misses.
62    pub misses: u64,
63    /// Ratio of hits to total lookups.
64    pub hit_rate: f64,
65    /// `created_at_ms` of the oldest (smallest timestamp) live entry.
66    pub oldest_entry_ms: Option<u64>,
67}
68
69/// FNV-1a 64-bit hash of `data`.
70pub fn fnv1a_hash(data: &[u8]) -> u64 {
71    const OFFSET: u64 = 14_695_981_039_346_656_037;
72    const PRIME: u64 = 1_099_511_628_211;
73    let mut hash = OFFSET;
74    for &byte in data {
75        hash ^= byte as u64;
76        hash = hash.wrapping_mul(PRIME);
77    }
78    hash
79}
80
81/// Build a [`CacheKey`] from a prompt string, model name, and temperature.
82pub fn cache_key(prompt: &str, model: &str, temperature: f32) -> CacheKey {
83    let prompt_hash = fnv1a_hash(prompt.as_bytes());
84    let temperature_bucket = (temperature * 10.0).round() as u8;
85    CacheKey {
86        prompt_hash,
87        model: model.to_string(),
88        temperature_bucket,
89    }
90}
91
92/// In-memory output deduplication cache.
93pub struct OutputCache {
94    /// Live entries.
95    pub entries: HashMap<CacheKey, CachedOutput>,
96    /// Maximum number of entries.
97    pub capacity: usize,
98    /// Eviction policy applied when `entries.len() == capacity`.
99    pub policy: EvictionPolicy,
100    /// Lifetime hit counter.
101    pub hits: u64,
102    /// Lifetime miss counter.
103    pub misses: u64,
104    /// Internal LCG state used by `EvictionPolicy::Random`.
105    pub lcg_state: u64,
106}
107
108impl OutputCache {
109    /// Create a new cache with the given capacity and eviction policy.
110    pub fn new(capacity: usize, policy: EvictionPolicy) -> Self {
111        let lcg_state = match &policy {
112            EvictionPolicy::Random(seed) => *seed,
113            _ => 6_364_136_223_846_793_005,
114        };
115        Self {
116            entries: HashMap::new(),
117            capacity,
118            policy,
119            hits: 0,
120            misses: 0,
121            lcg_state,
122        }
123    }
124
125    /// Advance the internal LCG and return a pseudo-random `u64`.
126    fn lcg_next(&mut self) -> u64 {
127        self.lcg_state = self
128            .lcg_state
129            .wrapping_mul(6_364_136_223_846_793_005)
130            .wrapping_add(1_442_695_040_888_963_407);
131        self.lcg_state
132    }
133
134    /// Look up `key`.  Returns `None` if the entry is absent or expired.
135    /// On a hit the entry's `access_count` is incremented and the hit stat updated.
136    pub fn get(&mut self, key: &CacheKey, now_ms: u64) -> Option<&str> {
137        if let Some(entry) = self.entries.get_mut(key) {
138            if entry.is_expired(now_ms) {
139                self.misses += 1;
140                return None;
141            }
142            entry.access_count += 1;
143            self.hits += 1;
144            return Some(entry.response.as_str());
145        }
146        self.misses += 1;
147        None
148    }
149
150    /// Insert a new entry, evicting one existing entry first if at capacity.
151    pub fn insert(
152        &mut self,
153        key: CacheKey,
154        response: String,
155        tokens: u64,
156        ttl_ms: u64,
157        now_ms: u64,
158    ) {
159        if self.entries.len() >= self.capacity && !self.entries.contains_key(&key) {
160            self.evict_one(now_ms);
161        }
162        self.entries.insert(
163            key,
164            CachedOutput {
165                response,
166                tokens_used: tokens,
167                created_at_ms: now_ms,
168                ttl_ms,
169                access_count: 0,
170            },
171        );
172    }
173
174    /// Evict one entry: expired entries take priority; otherwise apply `policy`.
175    pub fn evict_one(&mut self, now_ms: u64) {
176        // First try to evict an expired entry.
177        let expired_key = self
178            .entries
179            .iter()
180            .find(|(_, v)| v.is_expired(now_ms))
181            .map(|(k, _)| k.clone());
182        if let Some(k) = expired_key {
183            self.entries.remove(&k);
184            return;
185        }
186
187        if self.entries.is_empty() {
188            return;
189        }
190
191        let victim = match &self.policy {
192            EvictionPolicy::LRU => {
193                // Proxy: evict entry with smallest access_count (least recently used approx).
194                self.entries
195                    .iter()
196                    .min_by_key(|(_, v)| v.access_count)
197                    .map(|(k, _)| k.clone())
198            }
199            EvictionPolicy::LFU => {
200                self.entries
201                    .iter()
202                    .min_by_key(|(_, v)| v.access_count)
203                    .map(|(k, _)| k.clone())
204            }
205            EvictionPolicy::TTLFirst => {
206                // Evict the entry whose absolute expiry timestamp is smallest.
207                self.entries
208                    .iter()
209                    .min_by_key(|(_, v)| v.created_at_ms + v.ttl_ms)
210                    .map(|(k, _)| k.clone())
211            }
212            EvictionPolicy::Random(_) => {
213                let rnd = self.lcg_next();
214                let idx = (rnd as usize) % self.entries.len();
215                self.entries.keys().nth(idx).cloned()
216            }
217        };
218
219        if let Some(k) = victim {
220            self.entries.remove(&k);
221        }
222    }
223
224    /// Remove all expired entries and return how many were removed.
225    pub fn prune_expired(&mut self, now_ms: u64) -> usize {
226        let before = self.entries.len();
227        self.entries.retain(|_, v| !v.is_expired(now_ms));
228        before - self.entries.len()
229    }
230
231    /// Ratio of hits to total lookups, or `0.0` if no lookups have been made.
232    pub fn hit_rate(&self) -> f64 {
233        let total = self.hits + self.misses;
234        if total == 0 {
235            0.0
236        } else {
237            self.hits as f64 / total as f64
238        }
239    }
240
241    /// Return a statistics snapshot for the current moment `now_ms`.
242    pub fn stats(&self, now_ms: u64) -> CacheStats {
243        let hit_rate = self.hit_rate();
244        let oldest_entry_ms = self
245            .entries
246            .values()
247            .filter(|v| !v.is_expired(now_ms))
248            .map(|v| v.created_at_ms)
249            .min();
250        CacheStats {
251            size: self.entries.len(),
252            capacity: self.capacity,
253            hits: self.hits,
254            misses: self.misses,
255            hit_rate,
256            oldest_entry_ms,
257        }
258    }
259}
260
261#[cfg(test)]
262mod tests {
263    use super::*;
264
265    fn key(prompt: &str) -> CacheKey {
266        cache_key(prompt, "gpt-4", 0.7)
267    }
268
269    #[test]
270    fn test_cache_miss() {
271        let mut cache = OutputCache::new(10, EvictionPolicy::LRU);
272        let k = key("hello");
273        assert!(cache.get(&k, 1000).is_none());
274        assert_eq!(cache.misses, 1);
275        assert_eq!(cache.hits, 0);
276    }
277
278    #[test]
279    fn test_cache_hit() {
280        let mut cache = OutputCache::new(10, EvictionPolicy::LRU);
281        let k = key("hello");
282        cache.insert(k.clone(), "world".to_string(), 10, 60_000, 1000);
283        let result = cache.get(&k, 2000);
284        assert_eq!(result, Some("world"));
285        assert_eq!(cache.hits, 1);
286        assert_eq!(cache.misses, 0);
287    }
288
289    #[test]
290    fn test_ttl_expiry() {
291        let mut cache = OutputCache::new(10, EvictionPolicy::LRU);
292        let k = key("expire me");
293        // TTL of 500 ms, inserted at t=0
294        cache.insert(k.clone(), "resp".to_string(), 5, 500, 0);
295        // Not yet expired at t=499
296        assert!(cache.get(&k, 499).is_some());
297        // Expired at t=500
298        assert!(cache.get(&k, 500).is_none());
299        // Pruning removes expired entries
300        cache.insert(k.clone(), "resp".to_string(), 5, 500, 0);
301        let pruned = cache.prune_expired(600);
302        assert_eq!(pruned, 1);
303        assert!(cache.entries.is_empty());
304    }
305
306    #[test]
307    fn test_lru_evicts_least_accessed() {
308        let mut cache = OutputCache::new(2, EvictionPolicy::LRU);
309        let k1 = cache_key("p1", "m", 0.0);
310        let k2 = cache_key("p2", "m", 0.0);
311        cache.insert(k1.clone(), "r1".to_string(), 1, 60_000, 0);
312        cache.insert(k2.clone(), "r2".to_string(), 1, 60_000, 0);
313        // Access k2 to give it higher access_count
314        cache.get(&k2, 1);
315        // Insert k3 — should evict k1 (access_count=0)
316        let k3 = cache_key("p3", "m", 0.0);
317        cache.insert(k3.clone(), "r3".to_string(), 1, 60_000, 1);
318        assert!(!cache.entries.contains_key(&k1));
319        assert!(cache.entries.contains_key(&k2));
320        assert!(cache.entries.contains_key(&k3));
321    }
322
323    #[test]
324    fn test_lfu_evicts_least_used() {
325        let mut cache = OutputCache::new(2, EvictionPolicy::LFU);
326        let k1 = cache_key("p1", "m", 0.0);
327        let k2 = cache_key("p2", "m", 0.0);
328        cache.insert(k1.clone(), "r1".to_string(), 1, 60_000, 0);
329        cache.insert(k2.clone(), "r2".to_string(), 1, 60_000, 0);
330        // Access k1 multiple times
331        cache.get(&k1, 1);
332        cache.get(&k1, 2);
333        // Insert k3 — should evict k2 (access_count=0)
334        let k3 = cache_key("p3", "m", 0.0);
335        cache.insert(k3.clone(), "r3".to_string(), 1, 60_000, 3);
336        assert!(cache.entries.contains_key(&k1));
337        assert!(!cache.entries.contains_key(&k2));
338        assert!(cache.entries.contains_key(&k3));
339    }
340
341    #[test]
342    fn test_hit_rate_calculation() {
343        let mut cache = OutputCache::new(10, EvictionPolicy::LRU);
344        let k = key("q");
345        cache.insert(k.clone(), "a".to_string(), 1, 60_000, 0);
346        cache.get(&k, 1); // hit
347        cache.get(&key("other"), 1); // miss
348        // 1 hit / 2 total = 0.5
349        let rate = cache.hit_rate();
350        assert!((rate - 0.5).abs() < 1e-9);
351        let stats = cache.stats(1);
352        assert_eq!(stats.hits, 1);
353        assert_eq!(stats.misses, 1);
354        assert!((stats.hit_rate - 0.5).abs() < 1e-9);
355    }
356}