pub mod local;
pub mod ollama;
pub mod openai;
pub use local::LocalProvider;
pub use ollama::OllamaProvider;
pub use openai::OpenAIProvider;
use rayon::prelude::*;
use crate::error::{CtxError, Result};
pub const OPENAI_EMBEDDING_DIM: usize = 1536; pub const LOCAL_EMBEDDING_DIM: usize = 384;
#[derive(clap::ValueEnum, serde::Deserialize, Clone, Copy, Debug, Default, PartialEq, Eq)]
#[value(rename_all = "lowercase")]
#[serde(rename_all = "lowercase")]
pub enum Provider {
#[default]
Local,
Openai,
Ollama,
}
impl Provider {
pub fn resolve(
provider: Option<Provider>,
openai_flag: bool,
config_default: Option<Provider>,
) -> Provider {
match provider {
Some(p) => p,
None if openai_flag => Provider::Openai,
None => config_default.unwrap_or_default(),
}
}
pub fn as_str(&self) -> &'static str {
match self {
Provider::Local => "local",
Provider::Openai => "openai",
Provider::Ollama => "ollama",
}
}
}
pub fn build_provider(
provider: Provider,
embedding: &crate::config::EmbeddingConfig,
) -> Result<Box<dyn EmbeddingProvider>> {
match provider {
Provider::Local => Ok(Box::new(local::LocalProvider::new()?)),
Provider::Openai => {
let p = openai::OpenAIProvider::from_env().map_err(|_| {
CtxError::embedding(
"OPENAI_API_KEY environment variable not set.\n\
Set it with: export OPENAI_API_KEY=sk-...",
)
})?;
Ok(Box::new(p))
}
Provider::Ollama => Ok(Box::new(ollama::OllamaProvider::from_config(
embedding.model.as_deref(),
embedding.host.as_deref(),
)?)),
}
}
pub fn warn_index_mismatch(db: &crate::db::Database, provider: &dyn EmbeddingProvider) {
let query_dim = provider.dimension();
let query_name = provider.name();
if let Ok(metadata) = db.get_embedding_metadata() {
for (stored_provider, _model, stored_dim, count) in &metadata {
let stored_dim = *stored_dim as usize;
if stored_dim != query_dim || stored_provider != query_name {
eprintln!("Warning: embedding provider/dimension mismatch with the index!");
eprintln!(
" Index: {count} embeddings from '{stored_provider}' (dim {stored_dim})"
);
eprintln!(" Query: '{query_name}' (dim {query_dim})");
eprintln!(
" Results may be inaccurate. Re-run `ctx embed --provider {query_name}` \
to regenerate embeddings."
);
eprintln!();
break;
}
}
}
}
#[derive(Debug, Clone)]
pub struct Embedding {
pub vector: Vec<f32>,
#[allow(dead_code)]
pub token_count: Option<usize>,
}
impl Embedding {
pub fn new(vector: Vec<f32>) -> Self {
Self {
vector,
token_count: None,
}
}
#[allow(dead_code)]
pub fn dim(&self) -> usize {
self.vector.len()
}
#[allow(dead_code)]
pub fn cosine_similarity(&self, other: &Embedding) -> f32 {
cosine_similarity(&self.vector, &other.vector)
}
#[allow(dead_code)]
pub fn to_json(&self) -> Result<String> {
Ok(serde_json::to_string(&self.vector)?)
}
#[allow(dead_code)]
pub fn from_json(json: &str) -> Result<Self> {
let vector: Vec<f32> = serde_json::from_str(json)?;
Ok(Self::new(vector))
}
}
pub trait EmbeddingProvider: Send + Sync {
fn name(&self) -> &str;
fn dimension(&self) -> usize;
fn embed(&self, text: &str) -> Result<Embedding>;
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
texts.iter().map(|t| self.embed(t)).collect()
}
}
pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() {
return 0.0;
}
let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm_a == 0.0 || norm_b == 0.0 {
return 0.0;
}
dot / (norm_a * norm_b)
}
#[allow(dead_code)]
pub fn dot_product(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() {
return 0.0;
}
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
#[allow(dead_code)]
pub fn normalize(v: &mut [f32]) {
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in v.iter_mut() {
*x /= norm;
}
}
}
#[derive(Debug, Clone)]
pub struct SearchResult {
pub symbol_id: String,
pub score: f32,
pub name: String,
pub kind: String,
pub file_path: String,
pub line: u32,
}
pub fn semantic_search(
db: &crate::db::Database,
query_embedding: &Embedding,
limit: usize,
) -> Result<Vec<SearchResult>> {
if db.has_vector_embeddings() {
if let Ok(results) = db.vector_search(&query_embedding.vector, limit) {
if !results.is_empty() {
return Ok(results
.into_iter()
.map(
|(symbol_id, name, kind, file_path, line, distance)| SearchResult {
symbol_id,
score: 1.0 / (1.0 + distance),
name,
kind,
file_path,
line,
},
)
.collect());
}
}
}
semantic_search_slow(db, query_embedding, limit)
}
fn semantic_search_slow(
db: &crate::db::Database,
query_embedding: &Embedding,
limit: usize,
) -> Result<Vec<SearchResult>> {
let all_embeddings = db.get_all_embeddings()?;
let mut scored: Vec<_> = all_embeddings
.into_iter()
.map(|(symbol_id, name, kind, file_path, line, vector)| {
let score = cosine_similarity(&query_embedding.vector, &vector);
SearchResult {
symbol_id,
score,
name,
kind,
file_path,
line,
}
})
.collect();
scored.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
scored.truncate(limit);
Ok(scored)
}
fn embed_texts_parallel<P: EmbeddingProvider + ?Sized>(
provider: &P,
texts: &[&str],
) -> Result<Vec<Embedding>> {
let num_chunks = rayon::current_num_threads().max(1);
let chunk_size = texts.len().div_ceil(num_chunks).max(1);
let per_chunk: Vec<Vec<Embedding>> = texts
.par_chunks(chunk_size)
.map(|chunk| provider.embed_batch(chunk))
.collect::<Result<Vec<_>>>()?;
Ok(per_chunk.into_iter().flatten().collect())
}
pub fn embed_missing_symbols<P: EmbeddingProvider + ?Sized>(
db: &crate::db::Database,
provider: &P,
batch_size: usize,
serial: bool,
progress_callback: Option<&dyn Fn(usize, usize)>,
) -> Result<usize> {
let mut total_embedded = 0;
loop {
let symbols = db.get_symbols_without_embeddings(batch_size as i64)?;
if symbols.is_empty() {
break;
}
let texts: Vec<String> = symbols.iter().map(|s| s.to_embedding_text()).collect();
let text_refs: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
let embeddings = if serial {
provider.embed_batch(&text_refs)?
} else {
embed_texts_parallel(provider, &text_refs)?
};
for (symbol, embedding) in symbols.iter().zip(embeddings.iter()) {
db.store_embedding(&symbol.id, provider.name(), "default", &embedding.vector)?;
}
total_embedded += symbols.len();
if let Some(callback) = progress_callback {
callback(total_embedded, 0); }
if symbols.len() < batch_size {
break;
}
}
Ok(total_embedded)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cosine_similarity() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![1.0, 0.0, 0.0];
assert!((cosine_similarity(&a, &b) - 1.0).abs() < 1e-6);
let c = vec![0.0, 1.0, 0.0];
assert!(cosine_similarity(&a, &c).abs() < 1e-6);
let d = vec![-1.0, 0.0, 0.0];
assert!((cosine_similarity(&a, &d) + 1.0).abs() < 1e-6);
}
#[test]
fn test_normalize() {
let mut v = vec![3.0, 4.0];
normalize(&mut v);
assert!((v[0] - 0.6).abs() < 1e-6);
assert!((v[1] - 0.8).abs() < 1e-6);
}
struct LenProvider;
impl EmbeddingProvider for LenProvider {
fn name(&self) -> &str {
"len"
}
fn dimension(&self) -> usize {
1
}
fn embed(&self, text: &str) -> Result<Embedding> {
Ok(Embedding::new(vec![text.len() as f32]))
}
}
#[test]
fn test_embed_texts_parallel_preserves_order() {
let texts = ["a", "bb", "ccc", "dddd", "eeeee", "ffffff", "g", "hh"];
let refs: Vec<&str> = texts.to_vec();
let serial = LenProvider.embed_batch(&refs).unwrap();
let parallel = embed_texts_parallel(&LenProvider, &refs).unwrap();
assert_eq!(serial.len(), texts.len());
assert_eq!(parallel.len(), texts.len());
for (i, text) in texts.iter().enumerate() {
let expected = text.len() as f32;
assert_eq!(serial[i].vector, vec![expected]);
assert_eq!(parallel[i].vector, vec![expected], "order mismatch at {i}");
}
}
#[test]
fn test_embed_texts_parallel_empty() {
let parallel = embed_texts_parallel(&LenProvider, &[]).unwrap();
assert!(parallel.is_empty());
}
#[test]
fn test_embedding_json_roundtrip() {
let emb = Embedding::new(vec![0.1, 0.2, 0.3]);
let json = emb.to_json().unwrap();
let restored = Embedding::from_json(&json).unwrap();
assert_eq!(emb.vector, restored.vector);
}
}