1use chrono::Utc;
5use somatize_core::cache::{CacheKey, CacheStore, EntryMeta, Origin};
6use somatize_core::error::Result;
7use somatize_core::value::Value;
8use std::collections::{HashMap, VecDeque};
9use std::sync::Mutex;
10
11pub struct MemoryCache {
17 store: Mutex<LruStore>,
18}
19
20struct LruStore {
21 entries: HashMap<CacheKey, CacheEntry>,
22 access_order: VecDeque<CacheKey>,
24 current_bytes: usize,
25 max_bytes: usize,
26}
27
28struct CacheEntry {
29 value: Value,
30 meta: EntryMeta,
31 size: usize,
32}
33
34impl LruStore {
35 fn new(max_bytes: usize) -> Self {
36 Self {
37 entries: HashMap::new(),
38 access_order: VecDeque::new(),
39 current_bytes: 0,
40 max_bytes,
41 }
42 }
43
44 fn touch(&mut self, key: &CacheKey) {
45 self.access_order.retain(|k| k != key);
46 self.access_order.push_back(key.clone());
47 }
48
49 fn evict_until_fits(&mut self, needed: usize) {
50 while self.current_bytes + needed > self.max_bytes && !self.access_order.is_empty() {
51 if let Some(oldest_key) = self.access_order.pop_front()
52 && let Some(entry) = self.entries.remove(&oldest_key)
53 {
54 self.current_bytes = self.current_bytes.saturating_sub(entry.size);
55 }
56 }
57 }
58
59 fn insert(&mut self, key: CacheKey, entry: CacheEntry) {
60 let size = entry.size;
61
62 if let Some(old) = self.entries.remove(&key) {
64 self.current_bytes = self.current_bytes.saturating_sub(old.size);
65 self.access_order.retain(|k| k != &key);
66 }
67
68 self.evict_until_fits(size);
70
71 self.current_bytes += size;
72 self.access_order.push_back(key.clone());
73 self.entries.insert(key, entry);
74 }
75
76 fn remove(&mut self, key: &CacheKey) {
77 if let Some(entry) = self.entries.remove(key) {
78 self.current_bytes = self.current_bytes.saturating_sub(entry.size);
79 self.access_order.retain(|k| k != key);
80 }
81 }
82}
83
84impl MemoryCache {
85 pub fn new(max_bytes: usize) -> Self {
87 Self {
88 store: Mutex::new(LruStore::new(max_bytes)),
89 }
90 }
91
92 pub fn len(&self) -> usize {
94 self.store
95 .lock()
96 .unwrap_or_else(|e| e.into_inner())
97 .entries
98 .len()
99 }
100
101 pub fn is_empty(&self) -> bool {
103 self.len() == 0
104 }
105
106 pub fn current_bytes(&self) -> usize {
108 self.store
109 .lock()
110 .unwrap_or_else(|e| e.into_inner())
111 .current_bytes
112 }
113
114 pub fn clear(&self) {
116 let mut store = self.store.lock().unwrap_or_else(|e| e.into_inner());
117 store.entries.clear();
118 store.access_order.clear();
119 store.current_bytes = 0;
120 }
121}
122
123impl Default for MemoryCache {
124 fn default() -> Self {
125 Self::new(1024 * 1024 * 1024) }
127}
128
129impl CacheStore for MemoryCache {
130 fn get(&self, key: &CacheKey) -> Result<Option<Value>> {
131 let mut store = self.store.lock().unwrap_or_else(|e| e.into_inner());
132 if store.entries.contains_key(key) {
133 store.touch(key);
134 if let Some(entry) = store.entries.get_mut(key) {
135 entry.meta.last_accessed = Utc::now();
136 return Ok(Some(entry.value.clone()));
137 }
138 }
139 Ok(None)
140 }
141
142 fn put(&self, key: &CacheKey, value: &Value) -> Result<()> {
143 self.put_with_origin(
144 key,
145 value,
146 &Origin::Ingested {
147 source: "unknown".into(),
148 },
149 )
150 }
151
152 fn put_with_origin(&self, key: &CacheKey, value: &Value, origin: &Origin) -> Result<()> {
158 let size = estimate_size(value);
159 let now = Utc::now();
160
161 let mut store = self.store.lock().unwrap_or_else(|e| e.into_inner());
162 store.insert(
163 key.clone(),
164 CacheEntry {
165 value: value.clone(),
166 meta: EntryMeta {
167 key: key.clone(),
168 size_bytes: size as u64,
169 created_at: now,
170 last_accessed: now,
171 ttl: None,
172 origin: origin.clone(),
173 },
174 size,
175 },
176 );
177 Ok(())
178 }
179
180 fn exists(&self, key: &CacheKey) -> Result<bool> {
181 Ok(self
182 .store
183 .lock()
184 .unwrap_or_else(|e| e.into_inner())
185 .entries
186 .contains_key(key))
187 }
188
189 fn remove(&self, key: &CacheKey) -> Result<()> {
190 self.store
191 .lock()
192 .unwrap_or_else(|e| e.into_inner())
193 .remove(key);
194 Ok(())
195 }
196
197 fn metadata(&self, key: &CacheKey) -> Result<Option<EntryMeta>> {
198 Ok(self
199 .store
200 .lock()
201 .unwrap_or_else(|e| e.into_inner())
202 .entries
203 .get(key)
204 .map(|e| e.meta.clone()))
205 }
206}
207
208fn estimate_size(value: &Value) -> usize {
209 match value {
210 Value::Tensor { values, shape } => {
211 values.len() * std::mem::size_of::<f64>() + shape.len() * std::mem::size_of::<usize>()
212 }
213 Value::Text(s) => s.len(),
214 Value::Json(v) => v.to_string().len(),
215 Value::Bytes(b) | Value::Object(b) => b.len(),
216 Value::Empty => 0,
217 _ => 0,
218 }
219}
220
221#[cfg(test)]
222mod tests {
223 use super::*;
224 use serde_json::json;
225
226 #[test]
230 fn provenance_survives_a_put() {
231 let cache = MemoryCache::default();
232 let key = CacheKey::hash_data(b"provenance");
233
234 cache
235 .put_computed(
236 &key,
237 &Value::tensor(vec![1.0], vec![1]),
238 &Origin::Computed {
239 node_id: "scaler".into(),
240 run_id: "run-7".into(),
241 },
242 std::time::Duration::from_millis(3),
243 true,
244 )
245 .unwrap();
246
247 match cache.metadata(&key).unwrap().unwrap().origin {
248 Origin::Computed { node_id, run_id } => {
249 assert_eq!(node_id, "scaler");
250 assert_eq!(run_id, "run-7");
251 }
252 other => panic!("expected a Computed origin, got {other:?}"),
253 }
254 }
255
256 #[test]
257 fn put_and_get() {
258 let cache = MemoryCache::default();
259 let key = CacheKey::hash_data(b"test");
260 let value = Value::tensor(vec![1.0, 2.0, 3.0], vec![3]);
261
262 cache.put(&key, &value).unwrap();
263 let retrieved = cache.get(&key).unwrap().unwrap();
264 assert_eq!(retrieved, value);
265 }
266
267 #[test]
268 fn get_missing_returns_none() {
269 let cache = MemoryCache::default();
270 let key = CacheKey::hash_data(b"nonexistent");
271 assert!(cache.get(&key).unwrap().is_none());
272 }
273
274 #[test]
275 fn exists_check() {
276 let cache = MemoryCache::default();
277 let key = CacheKey::hash_data(b"test");
278 assert!(!cache.exists(&key).unwrap());
279
280 cache.put(&key, &Value::Empty).unwrap();
281 assert!(cache.exists(&key).unwrap());
282 }
283
284 #[test]
285 fn remove_entry() {
286 let cache = MemoryCache::default();
287 let key = CacheKey::hash_data(b"test");
288 cache.put(&key, &Value::Empty).unwrap();
289 assert_eq!(cache.len(), 1);
290
291 cache.remove(&key).unwrap();
292 assert_eq!(cache.len(), 0);
293 assert!(!cache.exists(&key).unwrap());
294 }
295
296 #[test]
297 fn metadata_available() {
298 let cache = MemoryCache::default();
299 let key = CacheKey::hash_data(b"test");
300 let value = Value::tensor(vec![1.0; 100], vec![10, 10]);
301
302 cache.put(&key, &value).unwrap();
303 let meta = cache.metadata(&key).unwrap().unwrap();
304 assert_eq!(meta.size_bytes, 816);
306 }
307
308 #[test]
309 fn clear_empties_cache() {
310 let cache = MemoryCache::default();
311 cache
312 .put(&CacheKey::hash_data(b"a"), &Value::Empty)
313 .unwrap();
314 cache
315 .put(&CacheKey::hash_data(b"b"), &Value::Empty)
316 .unwrap();
317 assert_eq!(cache.len(), 2);
318
319 cache.clear();
320 assert!(cache.is_empty());
321 assert_eq!(cache.current_bytes(), 0);
322 }
323
324 #[test]
325 fn overwrite_existing_key() {
326 let cache = MemoryCache::default();
327 let key = CacheKey::hash_data(b"test");
328
329 cache.put(&key, &Value::json(json!(1))).unwrap();
330 cache.put(&key, &Value::json(json!(2))).unwrap();
331
332 let val = cache.get(&key).unwrap().unwrap();
333 assert_eq!(val, Value::json(json!(2)));
334 assert_eq!(cache.len(), 1);
335 }
336
337 #[test]
338 fn multiple_keys() {
339 let cache = MemoryCache::default();
340 for i in 0..10 {
341 let key = CacheKey::hash_data(format!("key_{i}").as_bytes());
342 let val = Value::tensor(vec![i as f64], vec![1]);
343 cache.put(&key, &val).unwrap();
344 }
345 assert_eq!(cache.len(), 10);
346
347 let key5 = CacheKey::hash_data(b"key_5");
348 let val = cache.get(&key5).unwrap().unwrap();
349 let (data, _) = val.as_tensor().unwrap();
350 assert_eq!(data, &[5.0]);
351 }
352
353 #[test]
356 fn lru_evicts_oldest_when_full() {
357 let cache = MemoryCache::new(100);
359
360 let k1 = CacheKey::hash_data(b"first");
362 let k2 = CacheKey::hash_data(b"second");
363 let k3 = CacheKey::hash_data(b"third");
364
365 cache
366 .put(&k1, &Value::tensor(vec![0.0; 5], vec![5]))
367 .unwrap();
368 cache
369 .put(&k2, &Value::tensor(vec![0.0; 5], vec![5]))
370 .unwrap();
371 assert_eq!(cache.len(), 2);
372
373 cache
375 .put(&k3, &Value::tensor(vec![0.0; 5], vec![5]))
376 .unwrap();
377
378 assert!(!cache.exists(&k1).unwrap(), "k1 should be evicted");
379 assert!(cache.exists(&k2).unwrap(), "k2 should remain");
380 assert!(cache.exists(&k3).unwrap(), "k3 should remain");
381 }
382
383 #[test]
384 fn lru_access_prevents_eviction() {
385 let cache = MemoryCache::new(100);
386
387 let k1 = CacheKey::hash_data(b"first");
388 let k2 = CacheKey::hash_data(b"second");
389 let k3 = CacheKey::hash_data(b"third");
390
391 cache
392 .put(&k1, &Value::tensor(vec![0.0; 5], vec![5]))
393 .unwrap();
394 cache
395 .put(&k2, &Value::tensor(vec![0.0; 5], vec![5]))
396 .unwrap();
397
398 cache.get(&k1).unwrap();
400
401 cache
403 .put(&k3, &Value::tensor(vec![0.0; 5], vec![5]))
404 .unwrap();
405
406 assert!(cache.exists(&k1).unwrap(), "k1 was accessed, should remain");
407 assert!(!cache.exists(&k2).unwrap(), "k2 was LRU, should be evicted");
408 assert!(cache.exists(&k3).unwrap(), "k3 is new, should remain");
409 }
410
411 #[test]
412 fn lru_tracks_byte_usage() {
413 let cache = MemoryCache::new(1024);
414
415 assert_eq!(cache.current_bytes(), 0);
416
417 cache
419 .put(
420 &CacheKey::hash_data(b"a"),
421 &Value::tensor(vec![0.0; 10], vec![10]),
422 )
423 .unwrap();
424 assert_eq!(cache.current_bytes(), 88);
425
426 cache.remove(&CacheKey::hash_data(b"a")).unwrap();
427 assert_eq!(cache.current_bytes(), 0);
428 }
429
430 #[test]
431 fn lru_overwrite_updates_size() {
432 let cache = MemoryCache::new(1024);
433
434 let key = CacheKey::hash_data(b"key");
435 cache
436 .put(&key, &Value::tensor(vec![0.0; 10], vec![10]))
437 .unwrap();
438 let size1 = cache.current_bytes();
439
440 cache
442 .put(&key, &Value::tensor(vec![0.0; 20], vec![20]))
443 .unwrap();
444 let size2 = cache.current_bytes();
445
446 assert!(size2 > size1, "larger value should use more bytes");
447 assert_eq!(cache.len(), 1, "should still be one entry");
448 }
449}