Skip to main content

ctx/embeddings/
mod.rs

1//! Semantic search via embeddings.
2//!
3//! This module provides embedding generation and vector similarity search
4//! for semantic code search. It supports multiple embedding providers:
5//!
6//! - **OpenAI**: Uses text-embedding-3-small for high-quality embeddings
7//! - **Local**: Uses fastembed for local, offline embeddings
8//!
9//! # Architecture
10//!
11//! Embeddings are stored in SQLite as JSON-encoded float arrays. Similarity
12//! search is performed in Rust using cosine similarity for accuracy and
13//! portability (no native vector extensions required).
14//!
15//! # Usage
16//!
17//! ```ignore
18//! let provider = OpenAIProvider::new("sk-...")?;
19//! let embedding = provider.embed("fn authenticate(user: &str)")?;
20//!
21//! let results = search_similar(&db, "authentication functions", 10)?;
22//! ```
23
24pub mod local;
25pub mod ollama;
26pub mod openai;
27
28// Re-export providers for convenience
29pub use local::LocalProvider;
30pub use ollama::OllamaProvider;
31pub use openai::OpenAIProvider;
32
33use rayon::prelude::*;
34
35use crate::error::{CtxError, Result};
36
37/// Embedding dimension for different models
38pub const OPENAI_EMBEDDING_DIM: usize = 1536; // text-embedding-3-small
39pub const LOCAL_EMBEDDING_DIM: usize = 384; // all-MiniLM-L6-v2
40
41/// Which embedding backend to use.
42///
43/// `local` (fastembed) is the zero-config default; `openai` needs `OPENAI_API_KEY`;
44/// `ollama` talks to a local/remote Ollama server (`OLLAMA_HOST`,
45/// `OLLAMA_EMBED_MODEL`). Embeddings from different providers/models live in
46/// different vector spaces, so switching provider requires re-embedding.
47#[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    /// fastembed, local + offline (all-MiniLM-L6-v2, 384-dim).
52    #[default]
53    Local,
54    /// OpenAI API (text-embedding-3-small, 1536-dim). Requires `OPENAI_API_KEY`.
55    Openai,
56    /// Ollama server (model-dependent dimension). Local + offline.
57    Ollama,
58}
59
60impl Provider {
61    /// Resolve the effective provider by precedence:
62    /// `--provider` flag > deprecated `--openai` flag > `.ctx/config.toml`
63    /// (`[embedding].provider`) > built-in default (`local`).
64    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    /// Human-readable name matching `EmbeddingProvider::name()`.
77    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
86/// Build the embedding provider for the given backend, applying any
87/// provider-specific settings from `.ctx/config.toml` (`embedding`). This is the
88/// single place providers are constructed, so a new backend wires in once.
89///
90/// Env vars still take precedence over the config values (see the Ollama
91/// resolvers); pass `&EmbeddingConfig::default()` when there is no config.
92pub 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
114/// Warn (to stderr) when the query provider/dimension differs from what the index
115/// was embedded with. Embeddings from different providers/models occupy different
116/// vector spaces, so mixing them yields meaningless similarities — the fix is to
117/// re-embed. No-op when the index is empty or consistent.
118pub 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/// A vector embedding.
142#[derive(Debug, Clone)]
143pub struct Embedding {
144    /// The embedding vector
145    pub vector: Vec<f32>,
146    /// Number of tokens in the input
147    #[allow(dead_code)]
148    pub token_count: Option<usize>,
149}
150
151impl Embedding {
152    /// Create a new embedding from a vector.
153    pub fn new(vector: Vec<f32>) -> Self {
154        Self {
155            vector,
156            token_count: None,
157        }
158    }
159
160    /// Get the dimension of this embedding.
161    #[allow(dead_code)]
162    pub fn dim(&self) -> usize {
163        self.vector.len()
164    }
165
166    /// Compute cosine similarity with another embedding.
167    #[allow(dead_code)]
168    pub fn cosine_similarity(&self, other: &Embedding) -> f32 {
169        cosine_similarity(&self.vector, &other.vector)
170    }
171
172    /// Serialize to JSON for storage.
173    #[allow(dead_code)]
174    pub fn to_json(&self) -> Result<String> {
175        Ok(serde_json::to_string(&self.vector)?)
176    }
177
178    /// Deserialize from JSON.
179    #[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
186/// Trait for embedding providers.
187pub trait EmbeddingProvider: Send + Sync {
188    /// Get the name of this provider.
189    fn name(&self) -> &str;
190
191    /// Get the embedding dimension for this provider.
192    fn dimension(&self) -> usize;
193
194    /// Generate an embedding for a single text.
195    fn embed(&self, text: &str) -> Result<Embedding>;
196
197    /// Generate embeddings for multiple texts (batch).
198    /// Default implementation calls embed() for each text.
199    fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
200        texts.iter().map(|t| self.embed(t)).collect()
201    }
202}
203
204/// Compute cosine similarity between two vectors.
205pub 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/// Compute dot product similarity between two vectors.
222#[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/// Normalize a vector to unit length.
231#[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/// Search result from similarity search.
242#[derive(Debug, Clone)]
243pub struct SearchResult {
244    /// Symbol ID
245    pub symbol_id: String,
246    /// Similarity score (0.0 to 1.0)
247    pub score: f32,
248    /// Symbol name
249    pub name: String,
250    /// Symbol kind
251    pub kind: String,
252    /// File path
253    pub file_path: String,
254    /// Line number
255    pub line: u32,
256}
257
258/// Perform semantic similarity search using embeddings.
259///
260/// This automatically uses the fast vector search (sqlite-vec) when available,
261/// falling back to O(n) cosine similarity search otherwise.
262pub fn semantic_search(
263    db: &crate::db::Database,
264    query_embedding: &Embedding,
265    limit: usize,
266) -> Result<Vec<SearchResult>> {
267    // Try fast vector search first (sqlite-vec, O(log n))
268    if db.has_vector_embeddings() {
269        if let Ok(results) = db.vector_search(&query_embedding.vector, limit) {
270            if !results.is_empty() {
271                // Convert L2 distance to similarity score (0-1 range)
272                // L2 distance 0 = identical, higher = less similar
273                // We use 1/(1+d) to convert to similarity
274                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    // Fallback to O(n) cosine similarity search
292    semantic_search_slow(db, query_embedding, limit)
293}
294
295/// O(n) semantic search using cosine similarity.
296/// This loads all embeddings and computes similarity for each.
297fn semantic_search_slow(
298    db: &crate::db::Database,
299    query_embedding: &Embedding,
300    limit: usize,
301) -> Result<Vec<SearchResult>> {
302    // Get all embeddings from database
303    let all_embeddings = db.get_all_embeddings()?;
304
305    // Compute similarity for each
306    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    // Sort by score descending
322    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
332/// Compute embeddings for a batch of texts, splitting the work across rayon
333/// threads. The batch is divided into chunks and each chunk's `embed_batch`
334/// call runs concurrently.
335///
336/// Ordering is preserved: rayon's parallel `collect` keeps the chunks in
337/// source order, and each chunk keeps its own internal order, so flattening
338/// yields a `Vec<Embedding>` that lines up 1:1 with `texts`. Each chunk goes
339/// through the provider's normal `embed_batch`, so the provider's retry/backoff
340/// (e.g. the OpenAI rate-limit handling) is fully preserved.
341fn embed_texts_parallel<P: EmbeddingProvider + ?Sized>(
342    provider: &P,
343    texts: &[&str],
344) -> Result<Vec<Embedding>> {
345    // One chunk per worker thread (at least one), so provider calls run
346    // concurrently without over-splitting small batches.
347    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
358/// Embed all symbols that don't have embeddings yet.
359///
360/// Embedding computation is parallelized across rayon threads by default; pass
361/// `serial = true` for the single-threaded path. Regardless of mode, every
362/// `db.store_embedding` call happens serially on the owning thread (the
363/// `Database` connection is not shared for writes), and embeddings are stored
364/// in symbol order.
365pub 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        // Get symbols without embeddings
376        let symbols = db.get_symbols_without_embeddings(batch_size as i64)?;
377
378        if symbols.is_empty() {
379            break;
380        }
381
382        // Generate embedding text for each symbol
383        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        // Generate embeddings (parallel compute by default, serial on opt-out).
388        // Both paths return embeddings in the same order as `text_refs`.
389        let embeddings = if serial {
390            provider.embed_batch(&text_refs)?
391        } else {
392            embed_texts_parallel(provider, &text_refs)?
393        };
394
395        // Store embeddings serially on the owning thread, in symbol order.
396        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); // 0 = unknown total
404        }
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    /// Deterministic provider: embeds each text as a 1-dim vector holding the
440    /// byte length of the text. Lets us assert ordering without any network.
441    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        // Distinct lengths so each text has a unique embedding.
458        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}