Skip to main content

lattice_embed/service/
cached.rs

1//! Native-only LRU wrapper for an [`EmbeddingService`].
2//!
3//! It preserves caller order across partial cache hits and uses role-aware keys for asymmetric
4//! retrieval. See `docs/service.md` for the lookup and fill algorithm.
5
6use super::{EmbeddingRole, EmbeddingService, ValidatedTextBatch};
7use crate::error::Result;
8use crate::model::EmbeddingModel;
9use async_trait::async_trait;
10use std::sync::Arc;
11use tracing::debug;
12
13/// **Unstable**: caching strategy and constructor API may change; foundation-internal use only.
14///
15/// LRU-caching wrapper around an embedding service.
16///
17/// It preserves input order while reusing embeddings with matching model configuration and role.
18/// See [`docs/service.md`](../../docs/service.md#cachedembeddingservice-cache-hit-behavior) for the lookup and fill algorithm.
19pub struct CachedEmbeddingService<S> {
20    inner: Arc<S>,
21    cache: crate::cache::EmbeddingCache,
22}
23
24impl<S: EmbeddingService> CachedEmbeddingService<S> {
25    /// **Unstable**: constructor signature may change when cache config becomes a struct.
26    ///
27    /// # Arguments
28    ///
29    /// * `inner` - The underlying embedding service
30    /// * `cache_capacity` - Maximum number of embeddings to cache
31    pub fn new(inner: Arc<S>, cache_capacity: usize) -> Self {
32        Self {
33            inner,
34            cache: crate::cache::EmbeddingCache::new(cache_capacity),
35        }
36    }
37
38    /// **Unstable**: constructor signature may change when cache config becomes a struct.
39    pub fn with_default_cache(inner: Arc<S>) -> Self {
40        Self {
41            inner,
42            cache: crate::cache::EmbeddingCache::with_default_capacity(),
43        }
44    }
45
46    /// **Unstable**: returns internal `CacheStats` type which is itself Unstable.
47    pub fn cache_stats(&self) -> crate::cache::CacheStats {
48        self.cache.stats()
49    }
50
51    /// **Unstable**: internal cache management; API subject to change.
52    pub fn clear_cache(&self) {
53        self.cache.clear();
54    }
55}
56
57#[async_trait]
58impl<S: EmbeddingService + 'static> EmbeddingService for CachedEmbeddingService<S> {
59    async fn embed(&self, texts: &[String], model: EmbeddingModel) -> Result<Vec<Vec<f32>>> {
60        // Generic has its own role tag — see docs/service.md.
61        let texts = ValidatedTextBatch::new(texts)?;
62        self.cache_and_embed(texts, model, EmbeddingRole::Generic)
63            .await
64    }
65
66    /// Override: cache under the role key rather than prefixing here.
67    ///
68    /// `embed_query` and `embed_passage` reach this through their trait defaults,
69    /// so all three role paths share one cache-aware implementation.
70    async fn embed_with_role(
71        &self,
72        texts: &[String],
73        model: EmbeddingModel,
74        role: EmbeddingRole,
75    ) -> Result<Vec<Vec<f32>>> {
76        let texts = ValidatedTextBatch::new(texts)?;
77        self.cache_and_embed(texts, model, role).await
78    }
79
80    async fn embed_with_role_prevalidated(
81        &self,
82        texts: ValidatedTextBatch<'_>,
83        model: EmbeddingModel,
84        role: EmbeddingRole,
85    ) -> Result<Vec<Vec<f32>>> {
86        self.cache_and_embed(texts, model, role).await
87    }
88
89    fn supports_model(&self, model: EmbeddingModel) -> bool {
90        self.inner.supports_model(model)
91    }
92
93    fn name(&self) -> &'static str {
94        "cached-embedding"
95    }
96}
97
98impl<S: EmbeddingService + 'static> CachedEmbeddingService<S> {
99    /// Core cache-and-embed implementation shared by every entry point.
100    ///
101    /// `texts` is caller text. The retrieval instruction is applied by the wrapped
102    /// service rather than here, so the published cap is checked against what the
103    /// caller actually passed. `role` selects that instruction downstream and
104    /// namespaces the cache key, so the same raw text under two roles never shares
105    /// an entry; keying on caller text is equivalent to keying on prepared text
106    /// because the instruction is a function of the role and model config already
107    /// in the key.
108    async fn cache_and_embed(
109        &self,
110        texts: ValidatedTextBatch<'_>,
111        model: EmbeddingModel,
112        role: EmbeddingRole,
113    ) -> Result<Vec<Vec<f32>>> {
114        use crate::error::EmbedError;
115
116        // Fast path: bypass cache entirely when disabled (no key computation, no locking)
117        if !self.cache.is_enabled() {
118            return self
119                .inner
120                .embed_with_role_prevalidated(texts, model, role)
121                .await;
122        }
123
124        // Compute cache keys — include the active dimension (for MRL models) and role.
125        let model_config = self.inner.model_config(model);
126        let keys: Vec<_> = (0..texts.len())
127            .map(|index| self.cache.compute_key(texts.get(index), model_config, role))
128            .collect();
129
130        // Check cache for all texts — returns Arc<[f32]> refs (O(1) per hit)
131        let cached = self.cache.get_many(&keys);
132
133        let mut to_embed = Vec::new();
134        let mut results: Vec<Option<Vec<f32>>> = vec![None; texts.len()];
135
136        for (i, cached_emb) in cached.into_iter().enumerate() {
137            if let Some(arc) = cached_emb {
138                results[i] = Some(arc.to_vec());
139            } else {
140                to_embed.push(i);
141            }
142        }
143
144        // If all cached, return immediately
145        if to_embed.is_empty() {
146            debug!("all {} texts found in cache", texts.len());
147            // SAFETY: All slots are Some because we only reach here when to_embed is empty,
148            // meaning every text was found in cache and had results[i] = Some(...) assigned.
149            return Ok(results.into_iter().flatten().collect());
150        }
151
152        debug!(
153            "{} texts cached, {} need embedding",
154            texts.len() - to_embed.len(),
155            to_embed.len()
156        );
157
158        // Embed missing texts; the wrapped service applies the role instruction.
159        let texts_to_embed: Vec<&str> = to_embed.iter().map(|&index| texts.get(index)).collect();
160        let new_embeddings = self
161            .inner
162            .embed_with_role_prevalidated(texts.borrowed_subset(&texts_to_embed), model, role)
163            .await?;
164
165        // FP-035: validate count before zipping — a count mismatch would silently
166        // drop slots via zip() and return fewer embeddings than requested.
167        if new_embeddings.len() != to_embed.len() {
168            return Err(EmbedError::InferenceFailed(format!(
169                "embedding service returned {} vectors for {} inputs",
170                new_embeddings.len(),
171                to_embed.len()
172            )));
173        }
174
175        let mut cache_entries = Vec::with_capacity(to_embed.len());
176        for (i, embedding) in to_embed.into_iter().zip(new_embeddings.into_iter()) {
177            cache_entries.push((keys[i], embedding.clone()));
178            results[i] = Some(embedding);
179        }
180        self.cache.put_many(cache_entries);
181
182        // SAFETY: All slots are guaranteed to be Some at this point:
183        // - Cached items were assigned via results[i] = Some(arc.to_vec())
184        // - Non-cached items were assigned via results[i] = Some(embedding) in the loop above
185        Ok(results.into_iter().flatten().collect())
186    }
187}
188
189#[cfg(test)]
190mod tests {
191    use super::*;
192    use crate::error::EmbedError;
193    use crate::service::{
194        MAX_TEXT_BYTES, NativeEmbeddingService, reset_validate_texts_calls, validate_texts_calls,
195    };
196    use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
197    use std::sync::{Arc, Mutex};
198
199    #[derive(Default)]
200    struct ProbeService {
201        calls: AtomicUsize,
202        requests: Mutex<Vec<Vec<String>>>,
203        fail: bool,
204    }
205
206    impl ProbeService {
207        fn failing() -> Self {
208            Self {
209                fail: true,
210                ..Self::default()
211            }
212        }
213
214        fn calls(&self) -> usize {
215            self.calls.load(Ordering::Relaxed)
216        }
217
218        fn requests(&self) -> Vec<Vec<String>> {
219            self.requests.lock().unwrap().clone()
220        }
221
222        fn record(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
223            self.calls.fetch_add(1, Ordering::Relaxed);
224            self.requests.lock().unwrap().push(texts.to_vec());
225            if self.fail {
226                return Err(EmbedError::InferenceFailed("probe failure".into()));
227            }
228            Ok(texts.iter().map(|text| vec![text.len() as f32]).collect())
229        }
230    }
231
232    #[async_trait]
233    impl EmbeddingService for ProbeService {
234        async fn embed(&self, texts: &[String], _model: EmbeddingModel) -> Result<Vec<Vec<f32>>> {
235            self.record(texts)
236        }
237
238        async fn embed_with_role_prevalidated(
239            &self,
240            texts: ValidatedTextBatch<'_>,
241            model: EmbeddingModel,
242            role: EmbeddingRole,
243        ) -> Result<Vec<Vec<f32>>> {
244            let prepared = texts.to_owned_with_prefix(role.instruction(model));
245            self.record(&prepared)
246        }
247
248        fn supports_model(&self, _model: EmbeddingModel) -> bool {
249            true
250        }
251
252        fn name(&self) -> &'static str {
253            "cache-probe"
254        }
255    }
256
257    #[derive(Default)]
258    struct BorrowProbe {
259        saw_borrowed: AtomicBool,
260    }
261
262    #[async_trait]
263    impl EmbeddingService for BorrowProbe {
264        async fn embed(&self, _texts: &[String], _model: EmbeddingModel) -> Result<Vec<Vec<f32>>> {
265            panic!("cache delegation must use the prevalidated hook")
266        }
267
268        async fn embed_with_role_prevalidated(
269            &self,
270            texts: ValidatedTextBatch<'_>,
271            _model: EmbeddingModel,
272            _role: EmbeddingRole,
273        ) -> Result<Vec<Vec<f32>>> {
274            self.saw_borrowed
275                .store(texts.borrowed().is_some(), Ordering::Relaxed);
276            Ok((0..texts.len())
277                .map(|index| vec![texts.get(index).len() as f32])
278                .collect())
279        }
280
281        fn supports_model(&self, _model: EmbeddingModel) -> bool {
282            true
283        }
284
285        fn name(&self) -> &'static str {
286            "borrow-probe"
287        }
288    }
289
290    #[tokio::test(flavor = "current_thread")]
291    async fn disabled_delegate_validates_once_and_preserves_role_preparation() {
292        let inner = Arc::new(ProbeService::default());
293        let service = CachedEmbeddingService::new(inner.clone(), 0);
294        let texts = vec!["hello".to_string()];
295        let model = EmbeddingModel::BgeSmallEnV15;
296
297        reset_validate_texts_calls();
298        let result = service.embed_query(&texts, model).await.unwrap();
299
300        assert_eq!(validate_texts_calls(), 1);
301        assert_eq!(inner.calls(), 1);
302        let requests = inner.requests();
303        assert_eq!(requests.len(), 1);
304        assert_eq!(
305            requests[0][0],
306            format!("{}hello", model.query_instruction().unwrap())
307        );
308        assert_eq!(result, vec![vec![requests[0][0].len() as f32]]);
309    }
310
311    #[tokio::test(flavor = "current_thread")]
312    async fn miss_validates_once_and_delegates_only_uncached_texts() {
313        let inner = Arc::new(ProbeService::default());
314        let service = CachedEmbeddingService::new(inner.clone(), 128);
315        let model = EmbeddingModel::AllMiniLmL6V2;
316
317        reset_validate_texts_calls();
318        service.embed(&["cached".to_string()], model).await.unwrap();
319        assert_eq!(validate_texts_calls(), 1);
320
321        let calls_before = inner.calls();
322        reset_validate_texts_calls();
323        let result = service
324            .embed(&["cached".to_string(), "missing".to_string()], model)
325            .await
326            .unwrap();
327
328        assert_eq!(validate_texts_calls(), 1);
329        assert_eq!(inner.calls(), calls_before + 1);
330        assert_eq!(
331            inner.requests().last().unwrap(),
332            &vec!["missing".to_string()]
333        );
334        assert_eq!(result, vec![vec![6.0], vec![7.0]]);
335    }
336
337    #[tokio::test(flavor = "current_thread")]
338    async fn miss_delegates_borrowed_text_views_after_one_validation() {
339        let inner = Arc::new(BorrowProbe::default());
340        let service = CachedEmbeddingService::new(inner.clone(), 128);
341
342        reset_validate_texts_calls();
343        let result = service
344            .embed(&["uncached".to_string()], EmbeddingModel::AllMiniLmL6V2)
345            .await
346            .unwrap();
347
348        assert_eq!(validate_texts_calls(), 1);
349        assert!(inner.saw_borrowed.load(Ordering::Relaxed));
350        assert_eq!(result, vec![vec![8.0]]);
351    }
352
353    #[tokio::test(flavor = "current_thread")]
354    async fn all_hit_validates_once_without_delegating() {
355        let inner = Arc::new(ProbeService::default());
356        let service = CachedEmbeddingService::new(inner.clone(), 128);
357        let texts = vec!["cached".to_string()];
358        let model = EmbeddingModel::AllMiniLmL6V2;
359
360        service.embed(&texts, model).await.unwrap();
361        let calls_before = inner.calls();
362
363        reset_validate_texts_calls();
364        let result = service.embed(&texts, model).await.unwrap();
365
366        assert_eq!(validate_texts_calls(), 1);
367        assert_eq!(inner.calls(), calls_before);
368        assert_eq!(result, vec![vec![6.0]]);
369    }
370
371    #[tokio::test(flavor = "current_thread")]
372    async fn invalid_input_validates_once_without_delegating() {
373        let inner = Arc::new(ProbeService::default());
374        let service = CachedEmbeddingService::new(inner.clone(), 128);
375        let texts = vec!["x".repeat(MAX_TEXT_BYTES + 1)];
376
377        reset_validate_texts_calls();
378        let error = service
379            .embed(&texts, EmbeddingModel::AllMiniLmL6V2)
380            .await
381            .unwrap_err();
382
383        assert_eq!(validate_texts_calls(), 1);
384        assert!(matches!(
385            error,
386            EmbedError::TextTooLong {
387                max: MAX_TEXT_BYTES,
388                ..
389            }
390        ));
391        assert_eq!(inner.calls(), 0);
392    }
393
394    #[tokio::test(flavor = "current_thread")]
395    async fn delegate_error_still_validates_once() {
396        let inner = Arc::new(ProbeService::failing());
397        let service = CachedEmbeddingService::new(inner.clone(), 0);
398
399        reset_validate_texts_calls();
400        let error = service
401            .embed(&["valid".to_string()], EmbeddingModel::AllMiniLmL6V2)
402            .await
403            .unwrap_err();
404
405        assert_eq!(validate_texts_calls(), 1);
406        assert_eq!(inner.calls(), 1);
407        assert!(matches!(error, EmbedError::InferenceFailed(_)));
408    }
409
410    #[tokio::test(flavor = "current_thread")]
411    async fn nested_cache_delegate_preserves_single_validation() {
412        let probe = Arc::new(ProbeService::default());
413        let inner = Arc::new(CachedEmbeddingService::new(probe.clone(), 0));
414        let service = CachedEmbeddingService::new(inner, 0);
415
416        reset_validate_texts_calls();
417        let result = service
418            .embed_passage(&["nested".to_string()], EmbeddingModel::AllMiniLmL6V2)
419            .await
420            .unwrap();
421
422        assert_eq!(validate_texts_calls(), 1);
423        assert_eq!(probe.calls(), 1);
424        assert_eq!(result, vec![vec![6.0]]);
425    }
426
427    #[tokio::test(flavor = "current_thread")]
428    async fn native_delegate_preserves_model_validation_without_rescanning_caller_text() {
429        let inner = Arc::new(NativeEmbeddingService::default());
430        let service = CachedEmbeddingService::new(inner, 0);
431
432        reset_validate_texts_calls();
433        let error = service
434            .embed(&["valid".to_string()], EmbeddingModel::BgeBaseEnV15)
435            .await
436            .unwrap_err();
437
438        assert_eq!(validate_texts_calls(), 1);
439        assert!(matches!(error, EmbedError::InvalidInput(_)));
440    }
441}