pub mod local;
pub mod openai;
pub use local::LocalProvider;
pub use openai::OpenAIProvider;
use crate::error::Result;
pub const OPENAI_EMBEDDING_DIM: usize = 1536; pub const LOCAL_EMBEDDING_DIM: usize = 384;
#[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)
}
pub fn embed_missing_symbols<P: EmbeddingProvider + ?Sized>(
db: &crate::db::Database,
provider: &P,
batch_size: usize,
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 = provider.embed_batch(&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);
}
#[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);
}
}