1pub mod local;
25pub mod openai;
26
27pub use local::LocalProvider;
29pub use openai::OpenAIProvider;
30
31use crate::error::Result;
32
33pub const OPENAI_EMBEDDING_DIM: usize = 1536; pub const LOCAL_EMBEDDING_DIM: usize = 384; #[derive(Debug, Clone)]
39pub struct Embedding {
40 pub vector: Vec<f32>,
42 #[allow(dead_code)]
44 pub token_count: Option<usize>,
45}
46
47impl Embedding {
48 pub fn new(vector: Vec<f32>) -> Self {
50 Self {
51 vector,
52 token_count: None,
53 }
54 }
55
56 #[allow(dead_code)]
58 pub fn dim(&self) -> usize {
59 self.vector.len()
60 }
61
62 #[allow(dead_code)]
64 pub fn cosine_similarity(&self, other: &Embedding) -> f32 {
65 cosine_similarity(&self.vector, &other.vector)
66 }
67
68 #[allow(dead_code)]
70 pub fn to_json(&self) -> Result<String> {
71 Ok(serde_json::to_string(&self.vector)?)
72 }
73
74 #[allow(dead_code)]
76 pub fn from_json(json: &str) -> Result<Self> {
77 let vector: Vec<f32> = serde_json::from_str(json)?;
78 Ok(Self::new(vector))
79 }
80}
81
82pub trait EmbeddingProvider: Send + Sync {
84 fn name(&self) -> &str;
86
87 fn dimension(&self) -> usize;
89
90 fn embed(&self, text: &str) -> Result<Embedding>;
92
93 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
96 texts.iter().map(|t| self.embed(t)).collect()
97 }
98}
99
100pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
102 if a.len() != b.len() {
103 return 0.0;
104 }
105
106 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
107 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
108 let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
109
110 if norm_a == 0.0 || norm_b == 0.0 {
111 return 0.0;
112 }
113
114 dot / (norm_a * norm_b)
115}
116
117#[allow(dead_code)]
119pub fn dot_product(a: &[f32], b: &[f32]) -> f32 {
120 if a.len() != b.len() {
121 return 0.0;
122 }
123 a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
124}
125
126#[allow(dead_code)]
128pub fn normalize(v: &mut [f32]) {
129 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
130 if norm > 0.0 {
131 for x in v.iter_mut() {
132 *x /= norm;
133 }
134 }
135}
136
137#[derive(Debug, Clone)]
139pub struct SearchResult {
140 pub symbol_id: String,
142 pub score: f32,
144 pub name: String,
146 pub kind: String,
148 pub file_path: String,
150 pub line: u32,
152}
153
154pub fn semantic_search(
159 db: &crate::db::Database,
160 query_embedding: &Embedding,
161 limit: usize,
162) -> Result<Vec<SearchResult>> {
163 if db.has_vector_embeddings() {
165 if let Ok(results) = db.vector_search(&query_embedding.vector, limit) {
166 if !results.is_empty() {
167 return Ok(results
171 .into_iter()
172 .map(
173 |(symbol_id, name, kind, file_path, line, distance)| SearchResult {
174 symbol_id,
175 score: 1.0 / (1.0 + distance),
176 name,
177 kind,
178 file_path,
179 line,
180 },
181 )
182 .collect());
183 }
184 }
185 }
186
187 semantic_search_slow(db, query_embedding, limit)
189}
190
191fn semantic_search_slow(
194 db: &crate::db::Database,
195 query_embedding: &Embedding,
196 limit: usize,
197) -> Result<Vec<SearchResult>> {
198 let all_embeddings = db.get_all_embeddings()?;
200
201 let mut scored: Vec<_> = all_embeddings
203 .into_iter()
204 .map(|(symbol_id, name, kind, file_path, line, vector)| {
205 let score = cosine_similarity(&query_embedding.vector, &vector);
206 SearchResult {
207 symbol_id,
208 score,
209 name,
210 kind,
211 file_path,
212 line,
213 }
214 })
215 .collect();
216
217 scored.sort_by(|a, b| {
219 b.score
220 .partial_cmp(&a.score)
221 .unwrap_or(std::cmp::Ordering::Equal)
222 });
223 scored.truncate(limit);
224
225 Ok(scored)
226}
227
228pub fn embed_missing_symbols<P: EmbeddingProvider + ?Sized>(
230 db: &crate::db::Database,
231 provider: &P,
232 batch_size: usize,
233 progress_callback: Option<&dyn Fn(usize, usize)>,
234) -> Result<usize> {
235 let mut total_embedded = 0;
236
237 loop {
238 let symbols = db.get_symbols_without_embeddings(batch_size as i64)?;
240
241 if symbols.is_empty() {
242 break;
243 }
244
245 let texts: Vec<String> = symbols.iter().map(|s| s.to_embedding_text()).collect();
247
248 let text_refs: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
249
250 let embeddings = provider.embed_batch(&text_refs)?;
252
253 for (symbol, embedding) in symbols.iter().zip(embeddings.iter()) {
255 db.store_embedding(&symbol.id, provider.name(), "default", &embedding.vector)?;
256 }
257
258 total_embedded += symbols.len();
259
260 if let Some(callback) = progress_callback {
261 callback(total_embedded, 0); }
263
264 if symbols.len() < batch_size {
265 break;
266 }
267 }
268
269 Ok(total_embedded)
270}
271
272#[cfg(test)]
273mod tests {
274 use super::*;
275
276 #[test]
277 fn test_cosine_similarity() {
278 let a = vec![1.0, 0.0, 0.0];
279 let b = vec![1.0, 0.0, 0.0];
280 assert!((cosine_similarity(&a, &b) - 1.0).abs() < 1e-6);
281
282 let c = vec![0.0, 1.0, 0.0];
283 assert!(cosine_similarity(&a, &c).abs() < 1e-6);
284
285 let d = vec![-1.0, 0.0, 0.0];
286 assert!((cosine_similarity(&a, &d) + 1.0).abs() < 1e-6);
287 }
288
289 #[test]
290 fn test_normalize() {
291 let mut v = vec![3.0, 4.0];
292 normalize(&mut v);
293 assert!((v[0] - 0.6).abs() < 1e-6);
294 assert!((v[1] - 0.8).abs() < 1e-6);
295 }
296
297 #[test]
298 fn test_embedding_json_roundtrip() {
299 let emb = Embedding::new(vec![0.1, 0.2, 0.3]);
300 let json = emb.to_json().unwrap();
301 let restored = Embedding::from_json(&json).unwrap();
302 assert_eq!(emb.vector, restored.vector);
303 }
304}