Skip to main content

somatize_runtime/cache/
memory.rs

1//! [`MemoryCache`] — in-memory LRU [`CacheStore`] with byte-bounded
2//! eviction.
3
4use 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
11/// In-memory LRU cache store.
12///
13/// Enforces a maximum byte limit. When the limit is exceeded,
14/// the least recently accessed entries are evicted.
15/// Thread-safe via Mutex.
16pub struct MemoryCache {
17    store: Mutex<LruStore>,
18}
19
20struct LruStore {
21    entries: HashMap<CacheKey, CacheEntry>,
22    /// Access order: most recent at back, least recent at front.
23    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        // Remove old entry if exists
63        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        // Evict if needed
69        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    /// Create a new memory cache with a maximum byte limit.
86    pub fn new(max_bytes: usize) -> Self {
87        Self {
88            store: Mutex::new(LruStore::new(max_bytes)),
89        }
90    }
91
92    /// Number of entries currently in the cache.
93    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    /// Whether the cache is empty.
102    pub fn is_empty(&self) -> bool {
103        self.len() == 0
104    }
105
106    /// Current memory usage in bytes.
107    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    /// Clear all entries.
115    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) // 1GB
126    }
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    /// Overridden because the trait's default drops the origin, and
153    /// `MemoryCache` is the *default* store — so unless it records
154    /// provenance, most entries in most runs have none. It used to write
155    /// `Computed { node_id: "", run_id: "" }` for everything, which reads
156    /// as provenance while carrying none.
157    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    /// `MemoryCache` is the default store, and it inherited the trait's
227    /// origin-dropping default, writing `Computed { node_id: "", run_id: "" }`
228    /// for every entry — provenance-shaped, with no provenance in it.
229    #[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        // 100 f64 values * 8 bytes + 2 shape elements * 8 bytes = 816
305        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    // ── LRU eviction tests ──
354
355    #[test]
356    fn lru_evicts_oldest_when_full() {
357        // Cache with 100 bytes max
358        let cache = MemoryCache::new(100);
359
360        // Each tensor of 5 f64s = 5*8 + 1*8 = 48 bytes
361        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        // Adding third should evict first (48+48+48=144 > 100)
374        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        // Access k1, making k2 the least recently used
399        cache.get(&k1).unwrap();
400
401        // Adding k3 should evict k2 (LRU), not k1
402        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        // 10 f64s = 80 bytes data + 8 bytes shape = 88
418        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        // Replace with larger value
441        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}