Skip to main content

sz_orm_ai/embedding/
mod.rs

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
103/// Simple embedding model based on token frequency statistics.
104///
105/// This is a deterministic, in-memory embedding implementation:
106/// - Maintains a fixed-size vocabulary (`dimension` slots).
107/// - Each token is hashed into one of the `dimension` slots using FNV-1a.
108/// - The embedding vector is the L2-normalized token-frequency histogram.
109///
110/// This is NOT a neural embedding (no semantic knowledge), but it is a real,
111/// deterministic, reproducible vector representation suitable for testing
112/// similarity-search pipelines end-to-end.
113pub 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    /// Registers a token in the vocabulary, returning its slot index.
133    /// New tokens are appended; existing tokens keep their slot.
134    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        // Hash into the dimension space to keep vector length stable
140        // regardless of vocabulary size (bucket collisions are acceptable
141        // for a simple model and keep memory bounded).
142        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        // L2 normalize so cosine similarity is well-defined.
170        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
189// ==================== Embedding 模型适配器 ====================
190
191/// 缓存型 Embedding 模型适配器
192///
193/// 包装一个内部 EmbeddingModel,对相同输入文本返回缓存的向量,
194/// 避免重复计算。适用于嵌入计算成本高的场景(如调用远程 API)。
195///
196/// # 泛型参数
197/// - `M`: 内部嵌入模型
198pub struct CachingEmbeddingModel<M>
199where
200    M: EmbeddingModel,
201{
202    /// 内部嵌入模型
203    inner: M,
204    /// 缓存:文本 → 向量
205    cache: RwLock<HashMap<String, Vec<f32>>>,
206    /// 缓存命中次数(用于统计)
207    hits: std::sync::atomic::AtomicU64,
208    /// 缓存未命中次数
209    misses: std::sync::atomic::AtomicU64,
210}
211
212impl<M> CachingEmbeddingModel<M>
213where
214    M: EmbeddingModel,
215{
216    /// 创建缓存型适配器
217    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    /// 获取缓存大小
227    pub fn cache_size(&self) -> usize {
228        self.cache.read().len()
229    }
230
231    /// 获取缓存命中次数
232    pub fn cache_hits(&self) -> u64 {
233        self.hits.load(std::sync::atomic::Ordering::Relaxed)
234    }
235
236    /// 获取缓存未命中次数
237    pub fn cache_misses(&self) -> u64 {
238        self.misses.load(std::sync::atomic::Ordering::Relaxed)
239    }
240
241    /// 缓存命中率
242    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    /// 清空缓存
253    pub fn clear_cache(&self) {
254        let mut cache = self.cache.write();
255        cache.clear();
256    }
257
258    /// 获取内部模型引用
259    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        // 先查缓存
271        {
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        // 缓存未命中,调用内部模型
280        self.misses
281            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
282        let vector = self.inner.embed(text).await?;
283
284        // 写入缓存
285        {
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        // 先查缓存
299        {
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()); // 占位
311                }
312            }
313        }
314
315        // 批量计算未缓存的
316        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
339/// 归一化 Embedding 模型适配器
340///
341/// 包装一个内部 EmbeddingModel,对输出向量进行 L2 归一化,
342/// 确保所有向量都是单位向量。适用于需要使用余弦相似度的场景。
343pub struct NormalizedEmbeddingModel<M>
344where
345    M: EmbeddingModel,
346{
347    /// 内部嵌入模型
348    inner: M,
349}
350
351impl<M> NormalizedEmbeddingModel<M>
352where
353    M: EmbeddingModel,
354{
355    /// 创建归一化适配器
356    pub fn new(inner: M) -> Self {
357        Self { inner }
358    }
359
360    /// 对向量进行 L2 归一化
361    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    /// 获取内部模型引用
371    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
404/// 降维 Embedding 模型适配器
405///
406/// 包装一个内部 EmbeddingModel,通过截断或平均池化将输出向量降维到指定维度。
407/// 适用于需要将高维向量适配到低维索引的场景。
408pub struct DimReductionEmbeddingModel<M>
409where
410    M: EmbeddingModel,
411{
412    /// 内部嵌入模型
413    inner: M,
414    /// 目标维度
415    target_dimension: usize,
416    /// 降维策略
417    strategy: DimReductionStrategy,
418}
419
420/// 降维策略
421#[derive(Debug, Clone, Copy, PartialEq, Eq)]
422pub enum DimReductionStrategy {
423    /// 截断:只保留前 N 维
424    Truncate,
425    /// 平均池化:将向量分块后取平均
426    Average,
427}
428
429impl<M> DimReductionEmbeddingModel<M>
430where
431    M: EmbeddingModel,
432{
433    /// 创建降维适配器
434    pub fn new(inner: M, target_dimension: usize, strategy: DimReductionStrategy) -> Self {
435        Self {
436            inner,
437            target_dimension,
438            strategy,
439        }
440    }
441
442    /// 创建截断降维适配器
443    pub fn truncate(inner: M, target_dimension: usize) -> Self {
444        Self::new(inner, target_dimension, DimReductionStrategy::Truncate)
445    }
446
447    /// 创建平均池化降维适配器
448    pub fn average(inner: M, target_dimension: usize) -> Self {
449        Self::new(inner, target_dimension, DimReductionStrategy::Average)
450    }
451
452    /// 降维单个向量
453    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    /// 获取内部模型引用
484    pub fn inner(&self) -> &M {
485        &self.inner
486    }
487
488    /// 获取目标维度
489    pub fn target_dimension(&self) -> usize {
490        self.target_dimension
491    }
492
493    /// 获取降维策略
494    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
523/// 日志型 Embedding 模型适配器
524///
525/// 包装一个内部 EmbeddingModel,在调用时记录调用的文本和耗时。
526/// 适用于调试和性能分析场景。
527pub struct LoggingEmbeddingModel<M>
528where
529    M: EmbeddingModel,
530{
531    /// 内部嵌入模型
532    inner: M,
533    /// 调用次数
534    call_count: std::sync::atomic::AtomicU64,
535    /// 总文本数
536    total_texts: std::sync::atomic::AtomicU64,
537}
538
539impl<M> LoggingEmbeddingModel<M>
540where
541    M: EmbeddingModel,
542{
543    /// 创建日志型适配器
544    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    /// 获取调用次数
553    pub fn call_count(&self) -> u64 {
554        self.call_count.load(std::sync::atomic::Ordering::Relaxed)
555    }
556
557    /// 获取处理的文本总数
558    pub fn total_texts(&self) -> u64 {
559        self.total_texts.load(std::sync::atomic::Ordering::Relaxed)
560    }
561
562    /// 获取内部模型引用
563    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        // L2 norm should be ~1 (or 0 for empty input)
627        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    // ============ Embedding 适配器测试 ============
692
693    // ---- CachingEmbeddingModel 测试 ----
694
695    #[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        // 第二次应命中缓存
704        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        // 不同输入应产生不同向量
719        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(); // hit
734        let v2_miss = caching.embed("b").await.unwrap();
735        let v1_hit2 = caching.embed("a").await.unwrap(); // hit
736
737        // 缓存命中必须返回与首次 miss 相同的向量
738        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        // 4 calls: 2 misses, 2 hits
746        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        // 再次调用应 miss,且必须返回与首次相同的向量(确定性)
768        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        // 先缓存一个
782        let seed = caching.embed("hello").await.unwrap();
783
784        // 批量调用:hello 已缓存,world 未缓存
785        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        // hello 必须返回与 seed 相同的向量(缓存命中)
790        assert_eq!(results[0], seed, "cached entry must match seed vector");
791        // world 是新向量,维度必须一致
792        assert_eq!(results[1].len(), 16);
793        assert_ne!(results[0], results[1], "different inputs must differ");
794        assert_eq!(caching.cache_hits(), 1); // hello
795        assert_eq!(caching.cache_misses(), 2); // 初始 hello + world
796    }
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        // 先缓存全部
804        let seed_hello = caching.embed("hello").await.unwrap();
805        let seed_world = caching.embed("world").await.unwrap();
806
807        // 批量调用:全部命中
808        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        // 批量结果必须与 seed 向量逐一匹配
813        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    // ---- NormalizedEmbeddingModel 测试 ----
835
836    #[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        // L2 范数应接近 1(或 0 对于空输入)
844        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]; // norm = 5
864        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        // 零向量归一化后仍为零
874        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    // ---- DimReductionEmbeddingModel 测试 ----
893
894    #[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        // 4 维降到 2 维:chunk_size = 4/2 = 2
950        // [1,2] -> avg=1.5, [3,4] -> avg=3.5
951        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        // 原始 3 维,目标 10 维:截断后只有 3 维
977        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    // ---- LoggingEmbeddingModel 测试 ----
1004
1005    #[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        // 返回向量必须与底层 SimpleEmbeddingModel 直接计算的结果一致
1014        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        // 验证批量返回值与底层模型直接计算结果一致
1031        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    // ---- 适配器组合测试 ----
1064
1065    #[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        // 缓存应生效(第二次命中)
1075        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(); // cache hit
1088
1089        // 两次调用返回必须一致(缓存命中)
1090        assert_eq!(v1, v2, "cache hit must return identical vector");
1091        // 验证向量与底层 SimpleEmbeddingModel 直接计算结果一致
1092        let baseline = SimpleEmbeddingModel::new("composed", 16);
1093        assert_eq!(v1, baseline.embed_text("hello"));
1094        // 缓存命中时不会调用内部模型,所以日志层只记录 1 次
1095        assert_eq!(caching.inner().call_count(), 1);
1096        assert_eq!(caching.cache_hits(), 1);
1097    }
1098}