1use 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
13pub struct CachedEmbeddingService<S> {
20 inner: Arc<S>,
21 cache: crate::cache::EmbeddingCache,
22}
23
24impl<S: EmbeddingService> CachedEmbeddingService<S> {
25 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 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 pub fn cache_stats(&self) -> crate::cache::CacheStats {
48 self.cache.stats()
49 }
50
51 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 let texts = ValidatedTextBatch::new(texts)?;
62 self.cache_and_embed(texts, model, EmbeddingRole::Generic)
63 .await
64 }
65
66 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 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 if !self.cache.is_enabled() {
118 return self
119 .inner
120 .embed_with_role_prevalidated(texts, model, role)
121 .await;
122 }
123
124 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 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 to_embed.is_empty() {
146 debug!("all {} texts found in cache", texts.len());
147 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 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 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 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}