1use crate::model::ModelConfig;
11use crate::service::EmbeddingRole;
12use lru::LruCache;
13use parking_lot::RwLock;
14use std::num::NonZeroUsize;
15use std::sync::Arc;
16use std::sync::atomic::{AtomicU64, Ordering};
17use tracing::debug;
18
19pub type CacheKey = [u8; 32];
21
22pub const DEFAULT_CACHE_CAPACITY: usize = 4000;
26
27const NUM_SHARDS: usize = 16;
30
31const SHARD_MASK: usize = NUM_SHARDS - 1;
33
34const _: () = assert!(
36 NUM_SHARDS.is_power_of_two(),
37 "NUM_SHARDS must be a power of 2"
38);
39
40struct CacheShard {
42 lru: RwLock<LruCache<CacheKey, Arc<[f32]>>>,
43 hits: AtomicU64,
44 misses: AtomicU64,
45}
46
47impl CacheShard {
48 fn new(capacity: NonZeroUsize) -> Self {
49 Self {
50 lru: RwLock::new(LruCache::new(capacity)),
51 hits: AtomicU64::new(0),
52 misses: AtomicU64::new(0),
53 }
54 }
55
56 #[inline]
57 fn get(&self, key: &CacheKey) -> Option<Arc<[f32]>> {
58 let mut lru = self.lru.write();
59 let result = lru.get(key).cloned();
60 if result.is_some() {
61 self.hits.fetch_add(1, Ordering::Relaxed);
62 } else {
63 self.misses.fetch_add(1, Ordering::Relaxed);
64 }
65 result
66 }
67
68 #[inline]
69 fn put(&self, key: CacheKey, embedding: Arc<[f32]>) {
70 let mut lru = self.lru.write();
71 lru.put(key, embedding);
72 }
73
74 fn len(&self) -> usize {
75 self.lru.read().len()
76 }
77
78 fn clear(&self) {
79 self.lru.write().clear();
80 }
81
82 fn hits(&self) -> u64 {
83 self.hits.load(Ordering::Relaxed)
84 }
85
86 fn misses(&self) -> u64 {
87 self.misses.load(Ordering::Relaxed)
88 }
89}
90
91pub struct EmbeddingCache {
97 shards: Vec<CacheShard>,
98 enabled: bool,
99 capacity: usize,
100}
101
102#[inline(always)]
105fn shard_index(key: &CacheKey) -> usize {
106 key[0] as usize & SHARD_MASK
107}
108
109impl EmbeddingCache {
110 pub fn new(capacity: usize) -> Self {
116 let enabled = capacity != 0;
117
118 let per_shard = if enabled {
120 let base = capacity.div_ceil(NUM_SHARDS);
121 if base == 0 { 1 } else { base }
122 } else {
123 1 };
125
126 let per_shard_nz = NonZeroUsize::new(per_shard).expect("per_shard is always >= 1");
127
128 let shards = (0..NUM_SHARDS)
129 .map(|_| CacheShard::new(per_shard_nz))
130 .collect();
131
132 Self {
133 shards,
134 enabled,
135 capacity,
136 }
137 }
138
139 pub fn with_default_capacity() -> Self {
141 Self::new(DEFAULT_CACHE_CAPACITY)
142 }
143
144 pub fn compute_key(
149 &self,
150 text: &str,
151 model_config: ModelConfig,
152 role: EmbeddingRole,
153 ) -> CacheKey {
154 let mut hasher = blake3::Hasher::new();
155 hasher.update(text.as_bytes());
156 let model_key = format!(
158 "{}:{}:{}:{}",
159 model_config.model,
160 model_config.model.key_version(),
161 model_config.dimensions(),
162 role.cache_tag(),
163 );
164 hasher.update(model_key.as_bytes());
165 *hasher.finalize().as_bytes()
166 }
167
168 pub fn get(&self, key: &CacheKey) -> Option<Arc<[f32]>> {
173 if !self.enabled {
174 return None;
175 }
176
177 let idx = shard_index(key);
178 let result = self.shards[idx].get(key);
179
180 if result.is_some() {
181 debug!("cache hit for key {:?}", &key[..8]);
182 }
183
184 result
185 }
186
187 pub fn put(&self, key: CacheKey, embedding: Vec<f32>) {
192 if !self.enabled {
193 return;
194 }
195
196 let idx = shard_index(&key);
197 self.shards[idx].put(key, Arc::from(embedding));
198 debug!("cached embedding for key {:?}", &key[..8]);
199 }
200
201 pub fn get_many(&self, keys: &[CacheKey]) -> Vec<Option<Arc<[f32]>>> {
206 if !self.enabled {
207 return vec![None; keys.len()];
208 }
209
210 keys.iter()
211 .map(|key| {
212 let idx = shard_index(key);
213 self.shards[idx].get(key)
214 })
215 .collect()
216 }
217
218 pub fn put_many(&self, entries: Vec<(CacheKey, Vec<f32>)>) {
222 if !self.enabled {
223 return;
224 }
225
226 for (key, embedding) in entries {
227 let idx = shard_index(&key);
228 self.shards[idx].put(key, Arc::from(embedding));
229 }
230 }
231
232 pub fn stats(&self) -> CacheStats {
236 if !self.enabled {
237 let (hits, misses) = self.aggregate_counters();
238 return CacheStats {
239 size: 0,
240 capacity: 0,
241 hits,
242 misses,
243 };
244 }
245
246 let size: usize = self.shards.iter().map(CacheShard::len).sum();
247 let (hits, misses) = self.aggregate_counters();
248
249 CacheStats {
250 size,
251 capacity: self.capacity,
252 hits,
253 misses,
254 }
255 }
256
257 pub fn per_shard_stats(&self) -> Vec<ShardStats> {
261 self.shards
262 .iter()
263 .enumerate()
264 .map(|(i, s)| ShardStats {
265 shard_id: i,
266 size: s.len(),
267 hits: s.hits(),
268 misses: s.misses(),
269 })
270 .collect()
271 }
272
273 pub fn clear(&self) {
275 if !self.enabled {
276 return;
277 }
278
279 for shard in &self.shards {
280 shard.clear();
281 }
282 debug!("cache cleared");
283 }
284
285 #[inline]
287 pub fn is_enabled(&self) -> bool {
288 self.enabled
289 }
290
291 fn aggregate_counters(&self) -> (u64, u64) {
293 let hits: u64 = self.shards.iter().map(CacheShard::hits).sum();
294 let misses: u64 = self.shards.iter().map(CacheShard::misses).sum();
295 (hits, misses)
296 }
297}
298
299impl Default for EmbeddingCache {
300 fn default() -> Self {
301 Self::with_default_capacity()
302 }
303}
304
305#[derive(Debug, Clone, Copy)]
309pub struct CacheStats {
310 pub size: usize,
312 pub capacity: usize,
314 pub hits: u64,
316 pub misses: u64,
318}
319
320impl CacheStats {
321 pub fn hit_rate(&self) -> f64 {
323 let total = self.hits + self.misses;
324 if total == 0 {
325 0.0
326 } else {
327 self.hits as f64 / total as f64
328 }
329 }
330}
331
332#[derive(Debug, Clone, Copy)]
336pub struct ShardStats {
337 pub shard_id: usize,
339 pub size: usize,
341 pub hits: u64,
343 pub misses: u64,
345}
346
347impl ShardStats {
348 pub fn hit_rate(&self) -> f64 {
350 let total = self.hits + self.misses;
351 if total == 0 {
352 0.0
353 } else {
354 self.hits as f64 / total as f64
355 }
356 }
357}
358
359#[cfg(test)]
360mod tests {
361 use super::*;
362 use crate::model::EmbeddingModel;
363
364 #[test]
365 fn test_cache_basic_operations() {
366 let cache = EmbeddingCache::new(100);
367 let key = cache.compute_key(
368 "hello",
369 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
370 EmbeddingRole::Generic,
371 );
372
373 assert!(cache.get(&key).is_none());
374
375 let embedding = vec![0.1, 0.2, 0.3];
376 cache.put(key, embedding.clone());
377
378 let cached = cache.get(&key).unwrap();
379 assert_eq!(&*cached, &embedding[..]);
380 }
381
382 #[test]
383 fn test_cache_eviction() {
384 let cache = EmbeddingCache::new(16);
388
389 let mut keys = Vec::new();
393 for i in 0..32u32 {
394 let text = format!("text_{i}");
395 let key = cache.compute_key(
396 &text,
397 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
398 EmbeddingRole::Generic,
399 );
400 keys.push(key);
401 cache.put(key, vec![i as f32]);
402 }
403
404 let stats = cache.stats();
406 assert!(stats.size <= 16, "size {} exceeds capacity 16", stats.size);
407 }
408
409 #[test]
410 fn test_cache_lru_eviction_within_shard() {
411 let cache = EmbeddingCache::new(32);
413
414 let mut same_shard_keys = Vec::new();
416 let mut i = 0u32;
417
418 let first_key = cache.compute_key(
420 "probe_0",
421 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
422 EmbeddingRole::Generic,
423 );
424 let target_shard = shard_index(&first_key);
425
426 loop {
427 let key = cache.compute_key(
428 &format!("lru_test_{i}"),
429 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
430 EmbeddingRole::Generic,
431 );
432 if shard_index(&key) == target_shard {
433 same_shard_keys.push((key, i));
434 }
435 if same_shard_keys.len() == 3 {
436 break;
437 }
438 i += 1;
439 }
440
441 let (k1, v1) = same_shard_keys[0];
442 let (k2, v2) = same_shard_keys[1];
443 let (k3, v3) = same_shard_keys[2];
444
445 cache.put(k1, vec![v1 as f32]);
447 cache.put(k2, vec![v2 as f32]);
448
449 assert!(cache.get(&k1).is_some());
451
452 cache.put(k3, vec![v3 as f32]);
454
455 assert!(
456 cache.get(&k1).is_some(),
457 "k1 should survive (recently accessed)"
458 );
459 assert!(cache.get(&k2).is_none(), "k2 should be evicted (LRU)");
460 assert!(cache.get(&k3).is_some(), "k3 should exist (just inserted)");
461 }
462
463 #[test]
464 fn test_cache_different_models_different_keys() {
465 let cache = EmbeddingCache::new(100);
466
467 let key_small = cache.compute_key(
468 "text",
469 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
470 EmbeddingRole::Generic,
471 );
472 let key_base = cache.compute_key(
473 "text",
474 ModelConfig::new(EmbeddingModel::BgeBaseEnV15),
475 EmbeddingRole::Generic,
476 );
477
478 assert_ne!(key_small, key_base);
480 }
481
482 #[test]
483 fn test_cache_stats() {
484 let cache = EmbeddingCache::new(100);
485 let key = cache.compute_key(
486 "hello",
487 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
488 EmbeddingRole::Generic,
489 );
490
491 cache.get(&key); cache.put(key, vec![0.1]);
493 cache.get(&key); let stats = cache.stats();
496 assert_eq!(stats.size, 1);
497 assert_eq!(stats.hits, 1);
498 assert_eq!(stats.misses, 1);
499 assert!((stats.hit_rate() - 0.5).abs() < 0.001);
500 }
501
502 #[test]
503 fn test_cache_get_many() {
504 let cache = EmbeddingCache::new(100);
506
507 let key1 = cache.compute_key(
508 "one",
509 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
510 EmbeddingRole::Generic,
511 );
512 let key2 = cache.compute_key(
513 "two",
514 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
515 EmbeddingRole::Generic,
516 );
517 let key3 = cache.compute_key(
518 "three",
519 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
520 EmbeddingRole::Generic,
521 );
522
523 cache.put(key1, vec![1.0]);
524 cache.put(key3, vec![3.0]);
525
526 let results = cache.get_many(&[key1, key2, key3]);
527 assert_eq!(results.len(), 3);
528 assert_eq!(&**results[0].as_ref().unwrap(), &[1.0f32]);
529 assert!(results[1].is_none());
530 assert_eq!(&**results[2].as_ref().unwrap(), &[3.0f32]);
531 }
532
533 #[test]
534 fn test_cache_put_many() {
535 let cache = EmbeddingCache::new(100);
537
538 let key1 = cache.compute_key(
539 "one",
540 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
541 EmbeddingRole::Generic,
542 );
543 let key2 = cache.compute_key(
544 "two",
545 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
546 EmbeddingRole::Generic,
547 );
548
549 cache.put_many(vec![(key1, vec![1.0]), (key2, vec![2.0])]);
550
551 let v1 = cache.get(&key1).unwrap();
552 assert_eq!(&*v1, [1.0f32].as_slice());
553 let v2 = cache.get(&key2).unwrap();
554 assert_eq!(&*v2, [2.0f32].as_slice());
555 }
556
557 #[test]
558 fn test_cache_clear() {
559 let cache = EmbeddingCache::new(100);
560 let key = cache.compute_key(
561 "hello",
562 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
563 EmbeddingRole::Generic,
564 );
565
566 cache.put(key, vec![0.1]);
567 assert!(cache.get(&key).is_some());
568
569 cache.clear();
570 assert!(cache.get(&key).is_none());
571 assert_eq!(cache.stats().size, 0);
572 }
573
574 #[test]
575 fn test_cache_default_capacity() {
576 let cache = EmbeddingCache::with_default_capacity();
577 assert_eq!(cache.stats().capacity, DEFAULT_CACHE_CAPACITY);
578 }
579
580 #[test]
581 fn test_cache_disabled_is_noop() {
582 let cache = EmbeddingCache::new(0);
583 assert!(!cache.is_enabled());
584
585 let key = cache.compute_key(
586 "hello",
587 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
588 EmbeddingRole::Generic,
589 );
590 cache.put(key, vec![0.1]);
591 assert!(cache.get(&key).is_none());
592
593 let stats = cache.stats();
594 assert_eq!(stats.capacity, 0);
595 assert_eq!(stats.size, 0);
596 }
597
598 #[test]
599 fn test_concurrent_access() {
600 use std::thread;
601
602 let cache = Arc::new(EmbeddingCache::new(4000));
605 let mut handles = Vec::new();
606
607 for t in 0..8 {
609 let cache = Arc::clone(&cache);
610 handles.push(thread::spawn(move || {
611 for i in 0..100 {
612 let text = format!("thread_{t}_item_{i}");
614 let key = cache.compute_key(
615 &text,
616 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
617 EmbeddingRole::Generic,
618 );
619 let embedding = vec![t as f32; 384];
620 cache.put(key, embedding.clone());
621
622 let result = cache.get(&key);
623 assert!(result.is_some(), "put followed by get must succeed");
624 assert_eq!(result.unwrap().len(), 384);
625 }
626 }));
627 }
628
629 for h in handles {
630 h.join().expect("thread panicked");
631 }
632
633 let stats = cache.stats();
635 assert_eq!(stats.size, 800);
636 assert!(stats.hits >= 800, "at least 800 hits expected");
637 }
638
639 #[test]
640 fn test_shard_distribution() {
641 let cache = EmbeddingCache::new(4000);
644
645 let n = 800;
646 for i in 0..n {
647 let key = cache.compute_key(
648 &format!("item_{i}"),
649 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
650 EmbeddingRole::Generic,
651 );
652 cache.put(key, vec![i as f32]);
653 }
654
655 let shard_stats = cache.per_shard_stats();
656 assert_eq!(shard_stats.len(), NUM_SHARDS);
657
658 for ss in &shard_stats {
660 assert!(
661 ss.size > 0,
662 "shard {} is empty — distribution is pathological",
663 ss.shard_id
664 );
665 }
666
667 let total: usize = shard_stats.iter().map(|s| s.size).sum();
669 assert_eq!(total, n);
670
671 let avg = n / NUM_SHARDS; for ss in &shard_stats {
674 assert!(
675 ss.size <= avg * 3,
676 "shard {} has {} entries (avg {}), distribution too skewed",
677 ss.shard_id,
678 ss.size,
679 avg
680 );
681 }
682 }
683
684 #[test]
685 fn test_per_shard_stats_hit_tracking() {
686 let cache = EmbeddingCache::new(100);
687
688 let key1 = cache.compute_key(
690 "hello",
691 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
692 EmbeddingRole::Generic,
693 );
694 let key2 = cache.compute_key(
695 "world",
696 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
697 EmbeddingRole::Generic,
698 );
699
700 cache.put(key1, vec![1.0]);
701 cache.put(key2, vec![2.0]);
702
703 cache.get(&key1);
705 cache.get(&key1);
706 cache.get(&key1);
707 cache.get(&key2);
708
709 let shard_stats = cache.per_shard_stats();
710 let total_hits: u64 = shard_stats.iter().map(|s| s.hits).sum();
711 assert_eq!(total_hits, 4, "total hits should be 4");
712
713 let stats = cache.stats();
714 assert_eq!(stats.hits, 4);
715 assert_eq!(stats.misses, 0);
716 }
717
718 #[test]
719 fn test_small_capacity_rounds_up() {
720 let cache = EmbeddingCache::new(3);
722 assert!(cache.is_enabled());
723
724 let key = cache.compute_key(
725 "x",
726 ModelConfig::new(EmbeddingModel::BgeSmallEnV15),
727 EmbeddingRole::Generic,
728 );
729 cache.put(key, vec![42.0]);
730 assert!(cache.get(&key).is_some());
731 }
732
733 #[test]
740 fn test_role_query_vs_passage_different_keys() {
741 let cache = EmbeddingCache::new(100);
742 let model = ModelConfig::new(EmbeddingModel::MultilingualE5Small);
743 let text = "hello world";
744
745 let key_query = cache.compute_key(text, model, EmbeddingRole::Query);
746 let key_passage = cache.compute_key(text, model, EmbeddingRole::Passage);
747 let key_generic = cache.compute_key(text, model, EmbeddingRole::Generic);
748
749 assert_ne!(key_query, key_passage, "query vs passage must differ");
750 assert_ne!(key_query, key_generic, "query vs generic must differ");
751 assert_ne!(key_passage, key_generic, "passage vs generic must differ");
752 }
753
754 #[test]
756 fn test_role_key_deterministic() {
757 let cache = EmbeddingCache::new(100);
758 let model = ModelConfig::new(EmbeddingModel::BgeSmallEnV15);
759
760 let k1 = cache.compute_key("test", model, EmbeddingRole::Query);
761 let k2 = cache.compute_key("test", model, EmbeddingRole::Query);
762 assert_eq!(k1, k2, "identical inputs must produce identical key");
763 }
764
765 #[test]
767 fn test_role_cache_isolation() {
768 let cache = EmbeddingCache::new(100);
769 let model = ModelConfig::new(EmbeddingModel::MultilingualE5Small);
770
771 let key_query = cache.compute_key("embed me", model, EmbeddingRole::Query);
772 let key_passage = cache.compute_key("embed me", model, EmbeddingRole::Passage);
773
774 cache.put(key_query, vec![1.0, 2.0]);
776
777 assert!(
779 cache.get(&key_passage).is_none(),
780 "passage key must miss after storing under query key"
781 );
782
783 assert!(
785 cache.get(&key_query).is_some(),
786 "query key must hit after storing under query key"
787 );
788 }
789}