1use fastembed::{EmbeddingModel, InitOptions, TextEmbedding};
7use std::sync::Mutex;
8
9use super::{Embedding, EmbeddingProvider, LOCAL_EMBEDDING_DIM};
10use crate::error::{CtxError, Result};
11
12pub struct LocalProvider {
14 model: Mutex<TextEmbedding>,
15 #[allow(dead_code)]
16 model_name: String,
17}
18
19impl LocalProvider {
20 pub fn new() -> Result<Self> {
22 Self::with_model(EmbeddingModel::AllMiniLML6V2)
23 }
24
25 pub fn with_model(model: EmbeddingModel) -> Result<Self> {
27 let model_name = format!("{:?}", model);
28
29 let text_embedding =
30 TextEmbedding::try_new(InitOptions::new(model).with_show_download_progress(true))
31 .map_err(|e| CtxError::ModelNotFound(e.to_string()))?;
32
33 Ok(Self {
34 model: Mutex::new(text_embedding),
35 model_name,
36 })
37 }
38
39 #[allow(dead_code)]
42 pub fn new_large() -> Result<Self> {
43 Self::with_model(EmbeddingModel::BGEBaseENV15)
44 }
45}
46
47impl Default for LocalProvider {
48 fn default() -> Self {
49 Self::new().expect("Failed to initialize local embedding model")
50 }
51}
52
53impl EmbeddingProvider for LocalProvider {
54 fn name(&self) -> &str {
55 "local"
56 }
57
58 fn dimension(&self) -> usize {
59 LOCAL_EMBEDDING_DIM
60 }
61
62 fn embed(&self, text: &str) -> Result<Embedding> {
63 let mut model = self
64 .model
65 .lock()
66 .map_err(|e| CtxError::embedding(format!("Lock error: {}", e)))?;
67
68 let embeddings = model
69 .embed(vec![text], None)
70 .map_err(|e| CtxError::embedding(e.to_string()))?;
71
72 let vector = embeddings
73 .into_iter()
74 .next()
75 .ok_or_else(|| CtxError::embedding("Empty embedding result"))?;
76
77 Ok(Embedding::new(vector))
78 }
79
80 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
81 let mut model = self
82 .model
83 .lock()
84 .map_err(|e| CtxError::embedding(format!("Lock error: {}", e)))?;
85
86 let texts_owned: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
87
88 let embeddings = model
89 .embed(texts_owned, None)
90 .map_err(|e| CtxError::embedding(e.to_string()))?;
91
92 Ok(embeddings.into_iter().map(Embedding::new).collect())
93 }
94}
95
96#[cfg(test)]
97mod tests {
98 use super::*;
99
100 #[test]
101 #[ignore] fn test_local_embedding() {
103 let provider = LocalProvider::new().expect("Failed to create provider");
104 let embedding = provider.embed("Hello, world!").expect("Embedding failed");
105
106 assert_eq!(embedding.dim(), LOCAL_EMBEDDING_DIM);
107
108 let norm: f32 = embedding.vector.iter().map(|x| x * x).sum::<f32>().sqrt();
110 assert!((norm - 1.0).abs() < 0.1, "Embedding should be normalized");
111 }
112
113 #[test]
114 #[ignore] fn test_batch_embedding() {
116 let provider = LocalProvider::new().expect("Failed to create provider");
117 let texts = vec!["Hello", "World", "Test"];
118 let embeddings = provider
119 .embed_batch(&texts)
120 .expect("Batch embedding failed");
121
122 assert_eq!(embeddings.len(), 3);
123 for emb in &embeddings {
124 assert_eq!(emb.dim(), LOCAL_EMBEDDING_DIM);
125 }
126 }
127
128 #[test]
129 #[ignore] fn test_similarity() {
131 let provider = LocalProvider::new().expect("Failed to create provider");
132
133 let emb1 = provider
134 .embed("The cat sat on the mat")
135 .expect("Embedding failed");
136 let emb2 = provider
137 .embed("A feline rested on the rug")
138 .expect("Embedding failed");
139 let emb3 = provider
140 .embed("Python is a programming language")
141 .expect("Embedding failed");
142
143 let sim_similar = emb1.cosine_similarity(&emb2);
144 let sim_different = emb1.cosine_similarity(&emb3);
145
146 assert!(
148 sim_similar > sim_different,
149 "Similar sentences should have higher similarity: {} vs {}",
150 sim_similar,
151 sim_different
152 );
153 }
154}