1use crate::error::AiError;
2use async_trait::async_trait;
3use parking_lot::RwLock;
4use std::collections::HashMap;
5
6#[derive(Debug, Clone)]
7pub struct EmbeddingError {
8 pub message: String,
9 pub model: Option<String>,
10}
11
12impl EmbeddingError {
13 pub fn new(message: impl Into<String>) -> Self {
14 Self {
15 message: message.into(),
16 model: None,
17 }
18 }
19
20 pub fn with_model(message: impl Into<String>, model: impl Into<String>) -> Self {
21 Self {
22 message: message.into(),
23 model: Some(model.into()),
24 }
25 }
26}
27
28impl std::fmt::Display for EmbeddingError {
29 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
30 write!(f, "EmbeddingError: {}", self.message)?;
31 if let Some(ref model) = self.model {
32 write!(f, " (model: {})", model)?;
33 }
34 Ok(())
35 }
36}
37
38impl std::error::Error for EmbeddingError {}
39
40#[async_trait]
41pub trait EmbeddingModel: Send + Sync {
42 async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError>;
43 async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError>;
44 fn dimension(&self) -> usize;
45 fn model_name(&self) -> &str;
46}
47
48pub struct EmbeddingRecord {
49 pub id: String,
50 pub text: String,
51 pub vector: Vec<f32>,
52 pub metadata: Option<std::collections::HashMap<String, serde_json::Value>>,
53}
54
55impl EmbeddingRecord {
56 pub fn new(id: impl Into<String>, text: impl Into<String>, vector: Vec<f32>) -> Self {
57 Self {
58 id: id.into(),
59 text: text.into(),
60 vector,
61 metadata: None,
62 }
63 }
64
65 pub fn with_metadata(
66 id: impl Into<String>,
67 text: impl Into<String>,
68 vector: Vec<f32>,
69 metadata: std::collections::HashMap<String, serde_json::Value>,
70 ) -> Self {
71 Self {
72 id: id.into(),
73 text: text.into(),
74 vector,
75 metadata: Some(metadata),
76 }
77 }
78}
79
80pub struct EmbeddingBatch {
81 pub records: Vec<EmbeddingRecord>,
82 pub batch_size: usize,
83}
84
85impl EmbeddingBatch {
86 pub fn new(records: Vec<EmbeddingRecord>) -> Self {
87 Self {
88 records,
89 batch_size: 32,
90 }
91 }
92
93 pub fn with_batch_size(mut self, size: usize) -> Self {
94 self.batch_size = size;
95 self
96 }
97
98 pub fn batch_chunks(&self) -> Vec<&[EmbeddingRecord]> {
99 self.records.chunks(self.batch_size).collect()
100 }
101}
102
103pub struct SimpleEmbeddingModel {
114 name: String,
115 dimension: usize,
116 vocabulary: RwLock<HashMap<String, usize>>,
117}
118
119impl SimpleEmbeddingModel {
120 pub fn new(name: impl Into<String>, dimension: usize) -> Self {
121 Self {
122 name: name.into(),
123 dimension,
124 vocabulary: RwLock::new(HashMap::new()),
125 }
126 }
127
128 pub fn vocabulary_size(&self) -> usize {
129 self.vocabulary.read().len()
130 }
131
132 fn register_token(&self, token: &str) -> usize {
135 let mut vocab = self.vocabulary.write();
136 if let Some(&idx) = vocab.get(token) {
137 return idx;
138 }
139 let idx = fnv1a(token) % self.dimension.max(1);
143 vocab.insert(token.to_string(), idx);
144 idx
145 }
146
147 fn tokenize(text: &str) -> Vec<String> {
148 text.split(|c: char| !c.is_alphanumeric())
149 .filter(|s| !s.is_empty())
150 .map(|s| s.to_lowercase())
151 .collect()
152 }
153
154 fn embed_text(&self, text: &str) -> Vec<f32> {
155 let mut vec = vec![0.0f32; self.dimension];
156 if self.dimension == 0 {
157 return vec;
158 }
159 let tokens = Self::tokenize(text);
160 if tokens.is_empty() {
161 return vec;
162 }
163
164 for token in &tokens {
165 let idx = self.register_token(token);
166 vec[idx] += 1.0;
167 }
168
169 let norm: f32 = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
171 if norm > 0.0 {
172 for v in vec.iter_mut() {
173 *v /= norm;
174 }
175 }
176 vec
177 }
178}
179
180fn fnv1a(s: &str) -> usize {
181 let mut hash: u64 = 0xcbf29ce484222325;
182 for byte in s.as_bytes() {
183 hash ^= *byte as u64;
184 hash = hash.wrapping_mul(0x100000001b3);
185 }
186 hash as usize
187}
188
189pub struct CachingEmbeddingModel<M>
199where
200 M: EmbeddingModel,
201{
202 inner: M,
204 cache: RwLock<HashMap<String, Vec<f32>>>,
206 hits: std::sync::atomic::AtomicU64,
208 misses: std::sync::atomic::AtomicU64,
210}
211
212impl<M> CachingEmbeddingModel<M>
213where
214 M: EmbeddingModel,
215{
216 pub fn new(inner: M) -> Self {
218 Self {
219 inner,
220 cache: RwLock::new(HashMap::new()),
221 hits: std::sync::atomic::AtomicU64::new(0),
222 misses: std::sync::atomic::AtomicU64::new(0),
223 }
224 }
225
226 pub fn cache_size(&self) -> usize {
228 self.cache.read().len()
229 }
230
231 pub fn cache_hits(&self) -> u64 {
233 self.hits.load(std::sync::atomic::Ordering::Relaxed)
234 }
235
236 pub fn cache_misses(&self) -> u64 {
238 self.misses.load(std::sync::atomic::Ordering::Relaxed)
239 }
240
241 pub fn hit_rate(&self) -> f64 {
243 let hits = self.cache_hits();
244 let misses = self.cache_misses();
245 let total = hits + misses;
246 if total == 0 {
247 return 0.0;
248 }
249 hits as f64 / total as f64
250 }
251
252 pub fn clear_cache(&self) {
254 let mut cache = self.cache.write();
255 cache.clear();
256 }
257
258 pub fn inner(&self) -> &M {
260 &self.inner
261 }
262}
263
264#[async_trait]
265impl<M> EmbeddingModel for CachingEmbeddingModel<M>
266where
267 M: EmbeddingModel,
268{
269 async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
270 {
272 let cache = self.cache.read();
273 if let Some(vector) = cache.get(text) {
274 self.hits.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
275 return Ok(vector.clone());
276 }
277 }
278
279 self.misses
281 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
282 let vector = self.inner.embed(text).await?;
283
284 {
286 let mut cache = self.cache.write();
287 cache.insert(text.to_string(), vector.clone());
288 }
289
290 Ok(vector)
291 }
292
293 async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
294 let mut results = Vec::with_capacity(texts.len());
295 let mut uncached_indices = Vec::new();
296 let mut uncached_texts = Vec::new();
297
298 {
300 let cache = self.cache.read();
301 for (idx, text) in texts.iter().enumerate() {
302 if let Some(vector) = cache.get(text) {
303 self.hits.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
304 results.push(vector.clone());
305 } else {
306 self.misses
307 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
308 uncached_indices.push(idx);
309 uncached_texts.push(text.clone());
310 results.push(Vec::new()); }
312 }
313 }
314
315 if !uncached_texts.is_empty() {
317 let vectors = self.inner.embed_batch(&uncached_texts).await?;
318 let mut cache = self.cache.write();
319 for (i, idx) in uncached_indices.iter().enumerate() {
320 let text = &uncached_texts[i];
321 let vector = &vectors[i];
322 results[*idx] = vector.clone();
323 cache.insert(text.clone(), vector.clone());
324 }
325 }
326
327 Ok(results)
328 }
329
330 fn dimension(&self) -> usize {
331 self.inner.dimension()
332 }
333
334 fn model_name(&self) -> &str {
335 self.inner.model_name()
336 }
337}
338
339pub struct NormalizedEmbeddingModel<M>
344where
345 M: EmbeddingModel,
346{
347 inner: M,
349}
350
351impl<M> NormalizedEmbeddingModel<M>
352where
353 M: EmbeddingModel,
354{
355 pub fn new(inner: M) -> Self {
357 Self { inner }
358 }
359
360 pub fn l2_normalize(vector: &mut [f32]) {
362 let norm: f32 = vector.iter().map(|v| v * v).sum::<f32>().sqrt();
363 if norm > 0.0 {
364 for v in vector.iter_mut() {
365 *v /= norm;
366 }
367 }
368 }
369
370 pub fn inner(&self) -> &M {
372 &self.inner
373 }
374}
375
376#[async_trait]
377impl<M> EmbeddingModel for NormalizedEmbeddingModel<M>
378where
379 M: EmbeddingModel,
380{
381 async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
382 let mut vector = self.inner.embed(text).await?;
383 Self::l2_normalize(&mut vector);
384 Ok(vector)
385 }
386
387 async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
388 let mut vectors = self.inner.embed_batch(texts).await?;
389 for vector in vectors.iter_mut() {
390 Self::l2_normalize(vector);
391 }
392 Ok(vectors)
393 }
394
395 fn dimension(&self) -> usize {
396 self.inner.dimension()
397 }
398
399 fn model_name(&self) -> &str {
400 self.inner.model_name()
401 }
402}
403
404pub struct DimReductionEmbeddingModel<M>
409where
410 M: EmbeddingModel,
411{
412 inner: M,
414 target_dimension: usize,
416 strategy: DimReductionStrategy,
418}
419
420#[derive(Debug, Clone, Copy, PartialEq, Eq)]
422pub enum DimReductionStrategy {
423 Truncate,
425 Average,
427}
428
429impl<M> DimReductionEmbeddingModel<M>
430where
431 M: EmbeddingModel,
432{
433 pub fn new(inner: M, target_dimension: usize, strategy: DimReductionStrategy) -> Self {
435 Self {
436 inner,
437 target_dimension,
438 strategy,
439 }
440 }
441
442 pub fn truncate(inner: M, target_dimension: usize) -> Self {
444 Self::new(inner, target_dimension, DimReductionStrategy::Truncate)
445 }
446
447 pub fn average(inner: M, target_dimension: usize) -> Self {
449 Self::new(inner, target_dimension, DimReductionStrategy::Average)
450 }
451
452 pub fn reduce(&self, vector: Vec<f32>) -> Vec<f32> {
454 match self.strategy {
455 DimReductionStrategy::Truncate => {
456 vector.into_iter().take(self.target_dimension).collect()
457 }
458 DimReductionStrategy::Average => {
459 if vector.is_empty() || self.target_dimension == 0 {
460 return Vec::new();
461 }
462 let chunk_size = vector.len() / self.target_dimension;
463 if chunk_size == 0 {
464 return vector.into_iter().take(self.target_dimension).collect();
465 }
466 let mut result = Vec::with_capacity(self.target_dimension);
467 for i in 0..self.target_dimension {
468 let start = i * chunk_size;
469 let end = if i == self.target_dimension - 1 {
470 vector.len()
471 } else {
472 start + chunk_size
473 };
474 let chunk = &vector[start..end];
475 let avg: f32 = chunk.iter().sum::<f32>() / chunk.len() as f32;
476 result.push(avg);
477 }
478 result
479 }
480 }
481 }
482
483 pub fn inner(&self) -> &M {
485 &self.inner
486 }
487
488 pub fn target_dimension(&self) -> usize {
490 self.target_dimension
491 }
492
493 pub fn strategy(&self) -> DimReductionStrategy {
495 self.strategy
496 }
497}
498
499#[async_trait]
500impl<M> EmbeddingModel for DimReductionEmbeddingModel<M>
501where
502 M: EmbeddingModel,
503{
504 async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
505 let vector = self.inner.embed(text).await?;
506 Ok(self.reduce(vector))
507 }
508
509 async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
510 let vectors = self.inner.embed_batch(texts).await?;
511 Ok(vectors.into_iter().map(|v| self.reduce(v)).collect())
512 }
513
514 fn dimension(&self) -> usize {
515 self.target_dimension
516 }
517
518 fn model_name(&self) -> &str {
519 self.inner.model_name()
520 }
521}
522
523pub struct LoggingEmbeddingModel<M>
528where
529 M: EmbeddingModel,
530{
531 inner: M,
533 call_count: std::sync::atomic::AtomicU64,
535 total_texts: std::sync::atomic::AtomicU64,
537}
538
539impl<M> LoggingEmbeddingModel<M>
540where
541 M: EmbeddingModel,
542{
543 pub fn new(inner: M) -> Self {
545 Self {
546 inner,
547 call_count: std::sync::atomic::AtomicU64::new(0),
548 total_texts: std::sync::atomic::AtomicU64::new(0),
549 }
550 }
551
552 pub fn call_count(&self) -> u64 {
554 self.call_count.load(std::sync::atomic::Ordering::Relaxed)
555 }
556
557 pub fn total_texts(&self) -> u64 {
559 self.total_texts.load(std::sync::atomic::Ordering::Relaxed)
560 }
561
562 pub fn inner(&self) -> &M {
564 &self.inner
565 }
566}
567
568#[async_trait]
569impl<M> EmbeddingModel for LoggingEmbeddingModel<M>
570where
571 M: EmbeddingModel,
572{
573 async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
574 self.call_count
575 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
576 self.total_texts
577 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
578 self.inner.embed(text).await
579 }
580
581 async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
582 self.call_count
583 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
584 self.total_texts
585 .fetch_add(texts.len() as u64, std::sync::atomic::Ordering::Relaxed);
586 self.inner.embed_batch(texts).await
587 }
588
589 fn dimension(&self) -> usize {
590 self.inner.dimension()
591 }
592
593 fn model_name(&self) -> &str {
594 self.inner.model_name()
595 }
596}
597
598#[async_trait]
599impl EmbeddingModel for SimpleEmbeddingModel {
600 async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
601 Ok(self.embed_text(text))
602 }
603
604 async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
605 Ok(texts.iter().map(|t| self.embed_text(t)).collect())
606 }
607
608 fn dimension(&self) -> usize {
609 self.dimension
610 }
611
612 fn model_name(&self) -> &str {
613 &self.name
614 }
615}
616
617#[cfg(test)]
618mod tests {
619 use super::*;
620
621 #[tokio::test]
622 async fn test_embed_simple_text() {
623 let model = SimpleEmbeddingModel::new("test-model", 16);
624 let v = model.embed("hello world").await.unwrap();
625 assert_eq!(v.len(), 16);
626 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
628 assert!((norm - 1.0).abs() < 1e-5 || norm.abs() < 1e-5);
629 }
630
631 #[tokio::test]
632 async fn test_embed_empty_text() {
633 let model = SimpleEmbeddingModel::new("test-model", 8);
634 let v = model.embed("").await.unwrap();
635 assert!(v.iter().all(|x| *x == 0.0));
636 }
637
638 #[tokio::test]
639 async fn test_embed_deterministic() {
640 let model = SimpleEmbeddingModel::new("test-model", 32);
641 let v1 = model.embed("rust programming language").await.unwrap();
642 let v2 = model.embed("rust programming language").await.unwrap();
643 assert_eq!(v1, v2);
644 }
645
646 #[tokio::test]
647 async fn test_embed_similar_texts_closer_than_different() {
648 let model = SimpleEmbeddingModel::new("test-model", 64);
649 let v1 = model.embed("the quick brown fox jumps").await.unwrap();
650 let v2 = model.embed("the quick brown fox").await.unwrap();
651 let v3 = model
652 .embed("completely different words here")
653 .await
654 .unwrap();
655
656 let sim_close = cosine(&v1, &v2);
657 let sim_far = cosine(&v1, &v3);
658 assert!(
659 sim_close >= sim_far,
660 "similar texts should be at least as close"
661 );
662 }
663
664 #[tokio::test]
665 async fn test_embed_batch() {
666 let model = SimpleEmbeddingModel::new("test-model", 16);
667 let texts = vec!["hello".to_string(), "world".to_string()];
668 let vecs = model.embed_batch(&texts).await.unwrap();
669 assert_eq!(vecs.len(), 2);
670 assert_eq!(vecs[0].len(), 16);
671 assert_eq!(vecs[1].len(), 16);
672 }
673
674 #[test]
675 fn test_dimension_and_name() {
676 let model = SimpleEmbeddingModel::new("my-model", 128);
677 assert_eq!(model.dimension(), 128);
678 assert_eq!(model.model_name(), "my-model");
679 }
680
681 fn cosine(a: &[f32], b: &[f32]) -> f32 {
682 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
683 let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
684 let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
685 if na == 0.0 || nb == 0.0 {
686 return 0.0;
687 }
688 dot / (na * nb)
689 }
690
691 #[tokio::test]
696 async fn test_caching_model_caches_repeated_calls() {
697 let inner = SimpleEmbeddingModel::new("test", 16);
698 let caching = CachingEmbeddingModel::new(inner);
699
700 let v1 = caching.embed("hello world").await.unwrap();
701 let v2 = caching.embed("hello world").await.unwrap();
702
703 assert_eq!(v1, v2);
705 assert_eq!(caching.cache_hits(), 1);
706 assert_eq!(caching.cache_misses(), 1);
707 assert_eq!(caching.cache_size(), 1);
708 }
709
710 #[tokio::test]
711 async fn test_caching_model_different_texts() {
712 let inner = SimpleEmbeddingModel::new("test", 16);
713 let caching = CachingEmbeddingModel::new(inner);
714
715 let v1 = caching.embed("hello").await.unwrap();
716 let v2 = caching.embed("world").await.unwrap();
717
718 assert_eq!(v1.len(), 16);
720 assert_eq!(v2.len(), 16);
721 assert_ne!(v1, v2, "different inputs must yield different embeddings");
722 assert_eq!(caching.cache_misses(), 2);
723 assert_eq!(caching.cache_hits(), 0);
724 assert_eq!(caching.cache_size(), 2);
725 }
726
727 #[tokio::test]
728 async fn test_caching_model_hit_rate() {
729 let inner = SimpleEmbeddingModel::new("test", 16);
730 let caching = CachingEmbeddingModel::new(inner);
731
732 let v1_miss = caching.embed("a").await.unwrap();
733 let v1_hit = caching.embed("a").await.unwrap(); let v2_miss = caching.embed("b").await.unwrap();
735 let v1_hit2 = caching.embed("a").await.unwrap(); assert_eq!(v1_miss, v1_hit, "cache hit must return identical vector");
739 assert_eq!(v1_miss, v1_hit2, "cache hit must return identical vector");
740 assert_ne!(
741 v1_miss, v2_miss,
742 "different inputs must yield different vectors"
743 );
744
745 assert!((caching.hit_rate() - 0.5).abs() < 1e-6);
747 }
748
749 #[tokio::test]
750 async fn test_caching_model_hit_rate_zero_when_empty() {
751 let inner = SimpleEmbeddingModel::new("test", 16);
752 let caching = CachingEmbeddingModel::new(inner);
753 assert_eq!(caching.hit_rate(), 0.0);
754 }
755
756 #[tokio::test]
757 async fn test_caching_model_clear_cache() {
758 let inner = SimpleEmbeddingModel::new("test", 16);
759 let caching = CachingEmbeddingModel::new(inner);
760
761 let v1 = caching.embed("hello").await.unwrap();
762 assert_eq!(caching.cache_size(), 1);
763
764 caching.clear_cache();
765 assert_eq!(caching.cache_size(), 0);
766
767 let v2 = caching.embed("hello").await.unwrap();
769 assert_eq!(
770 v1, v2,
771 "embeddings must be deterministic across cache clears"
772 );
773 assert_eq!(caching.cache_misses(), 2);
774 }
775
776 #[tokio::test]
777 async fn test_caching_model_batch_mixed() {
778 let inner = SimpleEmbeddingModel::new("test", 16);
779 let caching = CachingEmbeddingModel::new(inner);
780
781 let seed = caching.embed("hello").await.unwrap();
783
784 let texts = vec!["hello".to_string(), "world".to_string()];
786 let results = caching.embed_batch(&texts).await.unwrap();
787
788 assert_eq!(results.len(), 2);
789 assert_eq!(results[0], seed, "cached entry must match seed vector");
791 assert_eq!(results[1].len(), 16);
793 assert_ne!(results[0], results[1], "different inputs must differ");
794 assert_eq!(caching.cache_hits(), 1); assert_eq!(caching.cache_misses(), 2); }
797
798 #[tokio::test]
799 async fn test_caching_model_batch_all_cached() {
800 let inner = SimpleEmbeddingModel::new("test", 16);
801 let caching = CachingEmbeddingModel::new(inner);
802
803 let seed_hello = caching.embed("hello").await.unwrap();
805 let seed_world = caching.embed("world").await.unwrap();
806
807 let texts = vec!["hello".to_string(), "world".to_string()];
809 let results = caching.embed_batch(&texts).await.unwrap();
810
811 assert_eq!(results.len(), 2);
812 assert_eq!(results[0], seed_hello, "cached hello must match seed");
814 assert_eq!(results[1], seed_world, "cached world must match seed");
815 assert_eq!(caching.cache_hits(), 2);
816 }
817
818 #[tokio::test]
819 async fn test_caching_model_preserves_dimension_and_name() {
820 let inner = SimpleEmbeddingModel::new("my-model", 32);
821 let caching = CachingEmbeddingModel::new(inner);
822
823 assert_eq!(caching.dimension(), 32);
824 assert_eq!(caching.model_name(), "my-model");
825 }
826
827 #[tokio::test]
828 async fn test_caching_model_inner_access() {
829 let inner = SimpleEmbeddingModel::new("inner", 8);
830 let caching = CachingEmbeddingModel::new(inner);
831 assert_eq!(caching.inner().model_name(), "inner");
832 }
833
834 #[tokio::test]
837 async fn test_normalized_model_produces_unit_vector() {
838 let inner = SimpleEmbeddingModel::new("test", 16);
839 let normalized = NormalizedEmbeddingModel::new(inner);
840
841 let v = normalized.embed("hello world").await.unwrap();
842 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
843 assert!(norm.abs() < 1e-5 || (norm - 1.0).abs() < 1e-5);
845 }
846
847 #[tokio::test]
848 async fn test_normalized_model_batch_produces_unit_vectors() {
849 let inner = SimpleEmbeddingModel::new("test", 16);
850 let normalized = NormalizedEmbeddingModel::new(inner);
851
852 let texts = vec!["hello".to_string(), "world".to_string()];
853 let vectors = normalized.embed_batch(&texts).await.unwrap();
854
855 for v in &vectors {
856 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
857 assert!(norm.abs() < 1e-5 || (norm - 1.0).abs() < 1e-5);
858 }
859 }
860
861 #[tokio::test]
862 async fn test_normalized_model_l2_normalize_static() {
863 let mut vector = vec![3.0, 4.0]; NormalizedEmbeddingModel::<SimpleEmbeddingModel>::l2_normalize(&mut vector);
865 assert!((vector[0] - 0.6).abs() < 1e-5);
866 assert!((vector[1] - 0.8).abs() < 1e-5);
867 }
868
869 #[tokio::test]
870 async fn test_normalized_model_l2_normalize_zero_vector() {
871 let mut vector = vec![0.0, 0.0, 0.0];
872 NormalizedEmbeddingModel::<SimpleEmbeddingModel>::l2_normalize(&mut vector);
873 assert!(vector.iter().all(|x| *x == 0.0));
875 }
876
877 #[tokio::test]
878 async fn test_normalized_model_preserves_dimension_and_name() {
879 let inner = SimpleEmbeddingModel::new("norm-model", 64);
880 let normalized = NormalizedEmbeddingModel::new(inner);
881 assert_eq!(normalized.dimension(), 64);
882 assert_eq!(normalized.model_name(), "norm-model");
883 }
884
885 #[tokio::test]
886 async fn test_normalized_model_inner_access() {
887 let inner = SimpleEmbeddingModel::new("inner", 8);
888 let normalized = NormalizedEmbeddingModel::new(inner);
889 assert_eq!(normalized.inner().model_name(), "inner");
890 }
891
892 #[test]
895 fn test_dim_reduction_strategy_variants() {
896 assert_eq!(
897 DimReductionStrategy::Truncate,
898 DimReductionStrategy::Truncate
899 );
900 assert_eq!(DimReductionStrategy::Average, DimReductionStrategy::Average);
901 assert_ne!(
902 DimReductionStrategy::Truncate,
903 DimReductionStrategy::Average
904 );
905 }
906
907 #[tokio::test]
908 async fn test_dim_reduction_truncate() {
909 let inner = SimpleEmbeddingModel::new("test", 16);
910 let reduced = DimReductionEmbeddingModel::truncate(inner, 8);
911
912 let v = reduced.embed("hello world").await.unwrap();
913 assert_eq!(v.len(), 8);
914 }
915
916 #[tokio::test]
917 async fn test_dim_reduction_truncate_batch() {
918 let inner = SimpleEmbeddingModel::new("test", 16);
919 let reduced = DimReductionEmbeddingModel::truncate(inner, 4);
920
921 let texts = vec!["hello".to_string(), "world".to_string()];
922 let vectors = reduced.embed_batch(&texts).await.unwrap();
923 for v in &vectors {
924 assert_eq!(v.len(), 4);
925 }
926 }
927
928 #[tokio::test]
929 async fn test_dim_reduction_average() {
930 let inner = SimpleEmbeddingModel::new("test", 16);
931 let reduced = DimReductionEmbeddingModel::average(inner, 4);
932
933 let v = reduced.embed("hello world").await.unwrap();
934 assert_eq!(v.len(), 4);
935 }
936
937 #[test]
938 fn test_dim_reduction_reduce_truncate() {
939 let inner = SimpleEmbeddingModel::new("test", 8);
940 let reduced = DimReductionEmbeddingModel::truncate(inner, 3);
941 let result = reduced.reduce(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
942 assert_eq!(result, vec![1.0, 2.0, 3.0]);
943 }
944
945 #[test]
946 fn test_dim_reduction_reduce_average() {
947 let inner = SimpleEmbeddingModel::new("test", 8);
948 let reduced = DimReductionEmbeddingModel::average(inner, 2);
949 let result = reduced.reduce(vec![1.0, 2.0, 3.0, 4.0]);
952 assert!((result[0] - 1.5).abs() < 1e-5);
953 assert!((result[1] - 3.5).abs() < 1e-5);
954 }
955
956 #[test]
957 fn test_dim_reduction_reduce_empty() {
958 let inner = SimpleEmbeddingModel::new("test", 8);
959 let reduced = DimReductionEmbeddingModel::average(inner, 2);
960 let result = reduced.reduce(vec![]);
961 assert!(result.is_empty());
962 }
963
964 #[test]
965 fn test_dim_reduction_reduce_target_zero() {
966 let inner = SimpleEmbeddingModel::new("test", 8);
967 let reduced = DimReductionEmbeddingModel::average(inner, 0);
968 let result = reduced.reduce(vec![1.0, 2.0, 3.0]);
969 assert!(result.is_empty());
970 }
971
972 #[test]
973 fn test_dim_reduction_reduce_truncate_smaller_than_target() {
974 let inner = SimpleEmbeddingModel::new("test", 8);
975 let reduced = DimReductionEmbeddingModel::truncate(inner, 10);
976 let result = reduced.reduce(vec![1.0, 2.0, 3.0]);
978 assert_eq!(result.len(), 3);
979 }
980
981 #[tokio::test]
982 async fn test_dim_reduction_dimension_returns_target() {
983 let inner = SimpleEmbeddingModel::new("test", 16);
984 let reduced = DimReductionEmbeddingModel::truncate(inner, 8);
985 assert_eq!(reduced.dimension(), 8);
986 }
987
988 #[tokio::test]
989 async fn test_dim_reduction_preserves_model_name() {
990 let inner = SimpleEmbeddingModel::new("original", 16);
991 let reduced = DimReductionEmbeddingModel::truncate(inner, 8);
992 assert_eq!(reduced.model_name(), "original");
993 }
994
995 #[test]
996 fn test_dim_reduction_target_dimension_and_strategy_accessors() {
997 let inner = SimpleEmbeddingModel::new("test", 16);
998 let reduced = DimReductionEmbeddingModel::new(inner, 8, DimReductionStrategy::Average);
999 assert_eq!(reduced.target_dimension(), 8);
1000 assert_eq!(reduced.strategy(), DimReductionStrategy::Average);
1001 }
1002
1003 #[tokio::test]
1006 async fn test_logging_model_counts_calls() {
1007 let inner = SimpleEmbeddingModel::new("test", 16);
1008 let logging = LoggingEmbeddingModel::new(inner);
1009
1010 let v1 = logging.embed("hello").await.unwrap();
1011 let v2 = logging.embed("world").await.unwrap();
1012
1013 let baseline = SimpleEmbeddingModel::new("test", 16);
1015 assert_eq!(v1, baseline.embed_text("hello"));
1016 assert_eq!(v2, baseline.embed_text("world"));
1017 assert_ne!(v1, v2);
1018 assert_eq!(logging.call_count(), 2);
1019 assert_eq!(logging.total_texts(), 2);
1020 }
1021
1022 #[tokio::test]
1023 async fn test_logging_model_batch_counts() {
1024 let inner = SimpleEmbeddingModel::new("test", 16);
1025 let logging = LoggingEmbeddingModel::new(inner);
1026
1027 let texts = vec!["hello".to_string(), "world".to_string(), "foo".to_string()];
1028 let results = logging.embed_batch(&texts).await.unwrap();
1029
1030 let baseline = SimpleEmbeddingModel::new("test", 16);
1032 assert_eq!(results.len(), 3);
1033 assert_eq!(results[0], baseline.embed_text("hello"));
1034 assert_eq!(results[1], baseline.embed_text("world"));
1035 assert_eq!(results[2], baseline.embed_text("foo"));
1036 assert_eq!(logging.call_count(), 1);
1037 assert_eq!(logging.total_texts(), 3);
1038 }
1039
1040 #[tokio::test]
1041 async fn test_logging_model_preserves_dimension_and_name() {
1042 let inner = SimpleEmbeddingModel::new("logged", 32);
1043 let logging = LoggingEmbeddingModel::new(inner);
1044 assert_eq!(logging.dimension(), 32);
1045 assert_eq!(logging.model_name(), "logged");
1046 }
1047
1048 #[tokio::test]
1049 async fn test_logging_model_inner_access() {
1050 let inner = SimpleEmbeddingModel::new("inner", 8);
1051 let logging = LoggingEmbeddingModel::new(inner);
1052 assert_eq!(logging.inner().model_name(), "inner");
1053 }
1054
1055 #[tokio::test]
1056 async fn test_logging_model_initial_counts_zero() {
1057 let inner = SimpleEmbeddingModel::new("test", 16);
1058 let logging = LoggingEmbeddingModel::new(inner);
1059 assert_eq!(logging.call_count(), 0);
1060 assert_eq!(logging.total_texts(), 0);
1061 }
1062
1063 #[tokio::test]
1066 async fn test_compose_caching_and_normalized() {
1067 let inner = SimpleEmbeddingModel::new("composed", 16);
1068 let caching = CachingEmbeddingModel::new(inner);
1069 let normalized = NormalizedEmbeddingModel::new(caching);
1070
1071 let v1 = normalized.embed("hello world").await.unwrap();
1072 let v2 = normalized.embed("hello world").await.unwrap();
1073
1074 assert_eq!(v1, v2);
1076 let norm: f32 = v1.iter().map(|x| x * x).sum::<f32>().sqrt();
1077 assert!(norm.abs() < 1e-5 || (norm - 1.0).abs() < 1e-5);
1078 }
1079
1080 #[tokio::test]
1081 async fn test_compose_logging_and_caching() {
1082 let inner = SimpleEmbeddingModel::new("composed", 16);
1083 let logging = LoggingEmbeddingModel::new(inner);
1084 let caching = CachingEmbeddingModel::new(logging);
1085
1086 let v1 = caching.embed("hello").await.unwrap();
1087 let v2 = caching.embed("hello").await.unwrap(); assert_eq!(v1, v2, "cache hit must return identical vector");
1091 let baseline = SimpleEmbeddingModel::new("composed", 16);
1093 assert_eq!(v1, baseline.embed_text("hello"));
1094 assert_eq!(caching.inner().call_count(), 1);
1096 assert_eq!(caching.cache_hits(), 1);
1097 }
1098}