1pub mod local;
25pub mod ollama;
26pub mod openai;
27
28pub use local::LocalProvider;
30pub use ollama::OllamaProvider;
31pub use openai::OpenAIProvider;
32
33use rayon::prelude::*;
34
35use crate::error::{CtxError, Result};
36
37pub const OPENAI_EMBEDDING_DIM: usize = 1536; pub const LOCAL_EMBEDDING_DIM: usize = 384; #[derive(clap::ValueEnum, serde::Deserialize, Clone, Copy, Debug, Default, PartialEq, Eq)]
48#[value(rename_all = "lowercase")]
49#[serde(rename_all = "lowercase")]
50pub enum Provider {
51 #[default]
53 Local,
54 Openai,
56 Ollama,
58}
59
60impl Provider {
61 pub fn resolve(
65 provider: Option<Provider>,
66 openai_flag: bool,
67 config_default: Option<Provider>,
68 ) -> Provider {
69 match provider {
70 Some(p) => p,
71 None if openai_flag => Provider::Openai,
72 None => config_default.unwrap_or_default(),
73 }
74 }
75
76 pub fn as_str(&self) -> &'static str {
78 match self {
79 Provider::Local => "local",
80 Provider::Openai => "openai",
81 Provider::Ollama => "ollama",
82 }
83 }
84}
85
86pub fn build_provider(
93 provider: Provider,
94 embedding: &crate::config::EmbeddingConfig,
95) -> Result<Box<dyn EmbeddingProvider>> {
96 match provider {
97 Provider::Local => Ok(Box::new(local::LocalProvider::new()?)),
98 Provider::Openai => {
99 let p = openai::OpenAIProvider::from_env().map_err(|_| {
100 CtxError::embedding(
101 "OPENAI_API_KEY environment variable not set.\n\
102 Set it with: export OPENAI_API_KEY=sk-...",
103 )
104 })?;
105 Ok(Box::new(p))
106 }
107 Provider::Ollama => Ok(Box::new(ollama::OllamaProvider::from_config(
108 embedding.model.as_deref(),
109 embedding.host.as_deref(),
110 )?)),
111 }
112}
113
114pub fn warn_index_mismatch(db: &crate::db::Database, provider: &dyn EmbeddingProvider) {
119 let query_dim = provider.dimension();
120 let query_name = provider.name();
121 if let Ok(metadata) = db.get_embedding_metadata() {
122 for (stored_provider, _model, stored_dim, count) in &metadata {
123 let stored_dim = *stored_dim as usize;
124 if stored_dim != query_dim || stored_provider != query_name {
125 eprintln!("Warning: embedding provider/dimension mismatch with the index!");
126 eprintln!(
127 " Index: {count} embeddings from '{stored_provider}' (dim {stored_dim})"
128 );
129 eprintln!(" Query: '{query_name}' (dim {query_dim})");
130 eprintln!(
131 " Results may be inaccurate. Re-run `ctx embed --provider {query_name}` \
132 to regenerate embeddings."
133 );
134 eprintln!();
135 break;
136 }
137 }
138 }
139}
140
141#[derive(Debug, Clone)]
143pub struct Embedding {
144 pub vector: Vec<f32>,
146 #[allow(dead_code)]
148 pub token_count: Option<usize>,
149}
150
151impl Embedding {
152 pub fn new(vector: Vec<f32>) -> Self {
154 Self {
155 vector,
156 token_count: None,
157 }
158 }
159
160 #[allow(dead_code)]
162 pub fn dim(&self) -> usize {
163 self.vector.len()
164 }
165
166 #[allow(dead_code)]
168 pub fn cosine_similarity(&self, other: &Embedding) -> f32 {
169 cosine_similarity(&self.vector, &other.vector)
170 }
171
172 #[allow(dead_code)]
174 pub fn to_json(&self) -> Result<String> {
175 Ok(serde_json::to_string(&self.vector)?)
176 }
177
178 #[allow(dead_code)]
180 pub fn from_json(json: &str) -> Result<Self> {
181 let vector: Vec<f32> = serde_json::from_str(json)?;
182 Ok(Self::new(vector))
183 }
184}
185
186pub trait EmbeddingProvider: Send + Sync {
188 fn name(&self) -> &str;
190
191 fn dimension(&self) -> usize;
193
194 fn embed(&self, text: &str) -> Result<Embedding>;
196
197 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
200 texts.iter().map(|t| self.embed(t)).collect()
201 }
202}
203
204pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
206 if a.len() != b.len() {
207 return 0.0;
208 }
209
210 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
211 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
212 let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
213
214 if norm_a == 0.0 || norm_b == 0.0 {
215 return 0.0;
216 }
217
218 dot / (norm_a * norm_b)
219}
220
221#[allow(dead_code)]
223pub fn dot_product(a: &[f32], b: &[f32]) -> f32 {
224 if a.len() != b.len() {
225 return 0.0;
226 }
227 a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
228}
229
230#[allow(dead_code)]
232pub fn normalize(v: &mut [f32]) {
233 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
234 if norm > 0.0 {
235 for x in v.iter_mut() {
236 *x /= norm;
237 }
238 }
239}
240
241#[derive(Debug, Clone)]
243pub struct SearchResult {
244 pub symbol_id: String,
246 pub score: f32,
248 pub name: String,
250 pub kind: String,
252 pub file_path: String,
254 pub line: u32,
256}
257
258pub fn semantic_search(
263 db: &crate::db::Database,
264 query_embedding: &Embedding,
265 limit: usize,
266) -> Result<Vec<SearchResult>> {
267 if db.has_vector_embeddings() {
269 if let Ok(results) = db.vector_search(&query_embedding.vector, limit) {
270 if !results.is_empty() {
271 return Ok(results
275 .into_iter()
276 .map(
277 |(symbol_id, name, kind, file_path, line, distance)| SearchResult {
278 symbol_id,
279 score: 1.0 / (1.0 + distance),
280 name,
281 kind,
282 file_path,
283 line,
284 },
285 )
286 .collect());
287 }
288 }
289 }
290
291 semantic_search_slow(db, query_embedding, limit)
293}
294
295fn semantic_search_slow(
298 db: &crate::db::Database,
299 query_embedding: &Embedding,
300 limit: usize,
301) -> Result<Vec<SearchResult>> {
302 let all_embeddings = db.get_all_embeddings()?;
304
305 let mut scored: Vec<_> = all_embeddings
307 .into_iter()
308 .map(|(symbol_id, name, kind, file_path, line, vector)| {
309 let score = cosine_similarity(&query_embedding.vector, &vector);
310 SearchResult {
311 symbol_id,
312 score,
313 name,
314 kind,
315 file_path,
316 line,
317 }
318 })
319 .collect();
320
321 scored.sort_by(|a, b| {
323 b.score
324 .partial_cmp(&a.score)
325 .unwrap_or(std::cmp::Ordering::Equal)
326 });
327 scored.truncate(limit);
328
329 Ok(scored)
330}
331
332fn embed_texts_parallel<P: EmbeddingProvider + ?Sized>(
342 provider: &P,
343 texts: &[&str],
344) -> Result<Vec<Embedding>> {
345 let num_chunks = rayon::current_num_threads().max(1);
348 let chunk_size = texts.len().div_ceil(num_chunks).max(1);
349
350 let per_chunk: Vec<Vec<Embedding>> = texts
351 .par_chunks(chunk_size)
352 .map(|chunk| provider.embed_batch(chunk))
353 .collect::<Result<Vec<_>>>()?;
354
355 Ok(per_chunk.into_iter().flatten().collect())
356}
357
358pub fn embed_missing_symbols<P: EmbeddingProvider + ?Sized>(
366 db: &crate::db::Database,
367 provider: &P,
368 batch_size: usize,
369 serial: bool,
370 progress_callback: Option<&dyn Fn(usize, usize)>,
371) -> Result<usize> {
372 let mut total_embedded = 0;
373
374 loop {
375 let symbols = db.get_symbols_without_embeddings(batch_size as i64)?;
377
378 if symbols.is_empty() {
379 break;
380 }
381
382 let texts: Vec<String> = symbols.iter().map(|s| s.to_embedding_text()).collect();
384
385 let text_refs: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
386
387 let embeddings = if serial {
390 provider.embed_batch(&text_refs)?
391 } else {
392 embed_texts_parallel(provider, &text_refs)?
393 };
394
395 for (symbol, embedding) in symbols.iter().zip(embeddings.iter()) {
397 db.store_embedding(&symbol.id, provider.name(), "default", &embedding.vector)?;
398 }
399
400 total_embedded += symbols.len();
401
402 if let Some(callback) = progress_callback {
403 callback(total_embedded, 0); }
405
406 if symbols.len() < batch_size {
407 break;
408 }
409 }
410
411 Ok(total_embedded)
412}
413
414#[cfg(test)]
415mod tests {
416 use super::*;
417
418 #[test]
419 fn test_cosine_similarity() {
420 let a = vec![1.0, 0.0, 0.0];
421 let b = vec![1.0, 0.0, 0.0];
422 assert!((cosine_similarity(&a, &b) - 1.0).abs() < 1e-6);
423
424 let c = vec![0.0, 1.0, 0.0];
425 assert!(cosine_similarity(&a, &c).abs() < 1e-6);
426
427 let d = vec![-1.0, 0.0, 0.0];
428 assert!((cosine_similarity(&a, &d) + 1.0).abs() < 1e-6);
429 }
430
431 #[test]
432 fn test_normalize() {
433 let mut v = vec![3.0, 4.0];
434 normalize(&mut v);
435 assert!((v[0] - 0.6).abs() < 1e-6);
436 assert!((v[1] - 0.8).abs() < 1e-6);
437 }
438
439 struct LenProvider;
442
443 impl EmbeddingProvider for LenProvider {
444 fn name(&self) -> &str {
445 "len"
446 }
447 fn dimension(&self) -> usize {
448 1
449 }
450 fn embed(&self, text: &str) -> Result<Embedding> {
451 Ok(Embedding::new(vec![text.len() as f32]))
452 }
453 }
454
455 #[test]
456 fn test_embed_texts_parallel_preserves_order() {
457 let texts = ["a", "bb", "ccc", "dddd", "eeeee", "ffffff", "g", "hh"];
459 let refs: Vec<&str> = texts.to_vec();
460
461 let serial = LenProvider.embed_batch(&refs).unwrap();
462 let parallel = embed_texts_parallel(&LenProvider, &refs).unwrap();
463
464 assert_eq!(serial.len(), texts.len());
465 assert_eq!(parallel.len(), texts.len());
466 for (i, text) in texts.iter().enumerate() {
467 let expected = text.len() as f32;
468 assert_eq!(serial[i].vector, vec![expected]);
469 assert_eq!(parallel[i].vector, vec![expected], "order mismatch at {i}");
470 }
471 }
472
473 #[test]
474 fn test_embed_texts_parallel_empty() {
475 let parallel = embed_texts_parallel(&LenProvider, &[]).unwrap();
476 assert!(parallel.is_empty());
477 }
478
479 #[test]
480 fn test_embedding_json_roundtrip() {
481 let emb = Embedding::new(vec![0.1, 0.2, 0.3]);
482 let json = emb.to_json().unwrap();
483 let restored = Embedding::from_json(&json).unwrap();
484 assert_eq!(emb.vector, restored.vector);
485 }
486}