Skip to main content

ctx/embeddings/
local.rs

1//! Local embedding provider using fastembed.
2//!
3//! Uses all-MiniLM-L6-v2 (384 dimensions) for fast, offline embeddings.
4//! No API key required - models are downloaded once and cached locally.
5
6use fastembed::{EmbeddingModel, InitOptions, TextEmbedding};
7use std::sync::Mutex;
8
9use super::{Embedding, EmbeddingProvider, LOCAL_EMBEDDING_DIM};
10use crate::error::{CtxError, Result};
11
12/// Local embedding provider using fastembed.
13pub struct LocalProvider {
14    model: Mutex<TextEmbedding>,
15    #[allow(dead_code)]
16    model_name: String,
17}
18
19impl LocalProvider {
20    /// Create a new local provider with the default model (all-MiniLM-L6-v2).
21    pub fn new() -> Result<Self> {
22        Self::with_model(EmbeddingModel::AllMiniLML6V2)
23    }
24
25    /// Create a provider with a specific model.
26    pub fn with_model(model: EmbeddingModel) -> Result<Self> {
27        let model_name = format!("{:?}", model);
28
29        let text_embedding =
30            TextEmbedding::try_new(InitOptions::new(model).with_show_download_progress(true))
31                .map_err(|e| CtxError::ModelNotFound(e.to_string()))?;
32
33        Ok(Self {
34            model: Mutex::new(text_embedding),
35            model_name,
36        })
37    }
38
39    /// Create a provider with a larger, more accurate model.
40    /// Uses BGE-base-en-v1.5 (768 dimensions).
41    #[allow(dead_code)]
42    pub fn new_large() -> Result<Self> {
43        Self::with_model(EmbeddingModel::BGEBaseENV15)
44    }
45}
46
47impl Default for LocalProvider {
48    fn default() -> Self {
49        Self::new().expect("Failed to initialize local embedding model")
50    }
51}
52
53impl EmbeddingProvider for LocalProvider {
54    fn name(&self) -> &str {
55        "local"
56    }
57
58    fn dimension(&self) -> usize {
59        LOCAL_EMBEDDING_DIM
60    }
61
62    fn embed(&self, text: &str) -> Result<Embedding> {
63        let mut model = self
64            .model
65            .lock()
66            .map_err(|e| CtxError::embedding(format!("Lock error: {}", e)))?;
67
68        let embeddings = model
69            .embed(vec![text], None)
70            .map_err(|e| CtxError::embedding(e.to_string()))?;
71
72        let vector = embeddings
73            .into_iter()
74            .next()
75            .ok_or_else(|| CtxError::embedding("Empty embedding result"))?;
76
77        Ok(Embedding::new(vector))
78    }
79
80    fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
81        let mut model = self
82            .model
83            .lock()
84            .map_err(|e| CtxError::embedding(format!("Lock error: {}", e)))?;
85
86        let texts_owned: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
87
88        let embeddings = model
89            .embed(texts_owned, None)
90            .map_err(|e| CtxError::embedding(e.to_string()))?;
91
92        Ok(embeddings.into_iter().map(Embedding::new).collect())
93    }
94}
95
96#[cfg(test)]
97mod tests {
98    use super::*;
99
100    #[test]
101    #[ignore] // Requires model download
102    fn test_local_embedding() {
103        let provider = LocalProvider::new().expect("Failed to create provider");
104        let embedding = provider.embed("Hello, world!").expect("Embedding failed");
105
106        assert_eq!(embedding.dim(), LOCAL_EMBEDDING_DIM);
107
108        // Check that the embedding is normalized (approximately unit length)
109        let norm: f32 = embedding.vector.iter().map(|x| x * x).sum::<f32>().sqrt();
110        assert!((norm - 1.0).abs() < 0.1, "Embedding should be normalized");
111    }
112
113    #[test]
114    #[ignore] // Requires model download
115    fn test_batch_embedding() {
116        let provider = LocalProvider::new().expect("Failed to create provider");
117        let texts = vec!["Hello", "World", "Test"];
118        let embeddings = provider
119            .embed_batch(&texts)
120            .expect("Batch embedding failed");
121
122        assert_eq!(embeddings.len(), 3);
123        for emb in &embeddings {
124            assert_eq!(emb.dim(), LOCAL_EMBEDDING_DIM);
125        }
126    }
127
128    #[test]
129    #[ignore] // Requires model download
130    fn test_similarity() {
131        let provider = LocalProvider::new().expect("Failed to create provider");
132
133        let emb1 = provider
134            .embed("The cat sat on the mat")
135            .expect("Embedding failed");
136        let emb2 = provider
137            .embed("A feline rested on the rug")
138            .expect("Embedding failed");
139        let emb3 = provider
140            .embed("Python is a programming language")
141            .expect("Embedding failed");
142
143        let sim_similar = emb1.cosine_similarity(&emb2);
144        let sim_different = emb1.cosine_similarity(&emb3);
145
146        // Similar sentences should have higher similarity
147        assert!(
148            sim_similar > sim_different,
149            "Similar sentences should have higher similarity: {} vs {}",
150            sim_similar,
151            sim_different
152        );
153    }
154}