tokio_prompt_orchestrator/
output_cache.rs1use std::collections::HashMap;
4
5#[derive(Debug, Clone, PartialEq, Eq, Hash)]
7pub struct CacheKey {
8 pub prompt_hash: u64,
10 pub model: String,
12 pub temperature_bucket: u8,
14}
15
16#[derive(Debug, Clone)]
18pub struct CachedOutput {
19 pub response: String,
21 pub tokens_used: u64,
23 pub created_at_ms: u64,
25 pub ttl_ms: u64,
27 pub access_count: u64,
29}
30
31impl CachedOutput {
32 pub fn is_expired(&self, now_ms: u64) -> bool {
34 now_ms >= self.created_at_ms + self.ttl_ms
35 }
36}
37
38#[derive(Debug, Clone)]
40pub enum EvictionPolicy {
41 LRU,
44 LFU,
46 TTLFirst,
48 Random(u64),
50}
51
52#[derive(Debug, Clone)]
54pub struct CacheStats {
55 pub size: usize,
57 pub capacity: usize,
59 pub hits: u64,
61 pub misses: u64,
63 pub hit_rate: f64,
65 pub oldest_entry_ms: Option<u64>,
67}
68
69pub 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
81pub 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
92pub struct OutputCache {
94 pub entries: HashMap<CacheKey, CachedOutput>,
96 pub capacity: usize,
98 pub policy: EvictionPolicy,
100 pub hits: u64,
102 pub misses: u64,
104 pub lcg_state: u64,
106}
107
108impl OutputCache {
109 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 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 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 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 pub fn evict_one(&mut self, now_ms: u64) {
176 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 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 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 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 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 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 cache.insert(k.clone(), "resp".to_string(), 5, 500, 0);
295 assert!(cache.get(&k, 499).is_some());
297 assert!(cache.get(&k, 500).is_none());
299 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 cache.get(&k2, 1);
315 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 cache.get(&k1, 1);
332 cache.get(&k1, 2);
333 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); cache.get(&key("other"), 1); 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}