Skip to main content

navi_core/memory/
embedding.rs

1//! Local embedding generation for semantic memory search.
2//!
3//! When the `embeddings` feature is enabled, NAVI uses Qwen3-Embedding-0.6B
4//! (GGUF format) via candle to generate 1024-dim embeddings, truncated to
5//! 256 dims via Matryoshka representation.
6//!
7//! Without the feature, the module provides a no-op fallback and search
8//! falls back to text matching (LIKE).
9
10use anyhow::Result;
11use std::path::PathBuf;
12use std::sync::{Arc, OnceLock};
13
14/// Target embedding dimension after Matryoshka truncation.
15/// 256 dims × 4 bytes = 1KB per memory — negligible storage overhead.
16pub const EMBED_DIM: usize = 256;
17
18/// Full model embedding dimension (before truncation).
19pub const FULL_EMBED_DIM: usize = 1024;
20
21/// Default model repo on HuggingFace.
22pub const DEFAULT_MODEL_REPO: &str = "Qwen/Qwen3-Embedding-0.6B-GGUF";
23pub const DEFAULT_MODEL_FILE: &str = "Qwen3-Embedding-0.6B-Q8_0.gguf";
24pub const DEFAULT_TOKENIZER_REPO: &str = "Qwen/Qwen3-Embedding-0.6B";
25pub const DEFAULT_TOKENIZER_FILE: &str = "tokenizer.json";
26
27/// Global cache for the embedding model — loaded once, reused across calls.
28static EMBEDDER_CACHE: OnceLock<std::sync::Mutex<Option<Arc<dyn Embedder>>>> = OnceLock::new();
29
30/// Returns the cached embedder, or loads it if not yet loaded.
31/// Returns None if embeddings are not available or the model is missing.
32pub fn get_cached_embedder(
33    model_path: &PathBuf,
34    tokenizer_path: &PathBuf,
35) -> Option<Arc<dyn Embedder>> {
36    let cache = EMBEDDER_CACHE.get_or_init(|| std::sync::Mutex::new(None));
37    let mut guard = cache.lock().ok()?;
38
39    if let Some(ref embedder) = *guard {
40        return Some(embedder.clone());
41    }
42
43    // Load the embedder
44    let config = EmbeddingConfig {
45        model_path: model_path.clone(),
46        tokenizer_path: tokenizer_path.clone(),
47        ..Default::default()
48    };
49    let embedder = create_embedder(config);
50    let embedder_arc: Arc<dyn Embedder> = Arc::from(embedder);
51
52    // Check if it's a NoEmbedder (feature off or model missing)
53    // We can't downcast Box to check, so we test with a simple embed call
54    match embedder_arc.embed("test") {
55        Ok(_) => {
56            *guard = Some(embedder_arc.clone());
57            Some(embedder_arc)
58        }
59        Err(_) => None,
60    }
61}
62
63/// Trait for embedding generation — allows mocking in tests.
64pub trait Embedder: Send + Sync {
65    /// Generates an embedding for the given text.
66    /// Returns a vector of `EMBED_DIM` f32 values.
67    fn embed(&self, text: &str) -> Result<Vec<f32>>;
68}
69
70/// No-op embedder used when the `embeddings` feature is disabled.
71/// Always returns an error — callers should fall back to text search.
72pub struct NoEmbedder;
73
74impl Embedder for NoEmbedder {
75    fn embed(&self, _text: &str) -> Result<Vec<f32>> {
76        anyhow::bail!("embeddings feature is not enabled")
77    }
78}
79
80/// Configuration for the local embedding model.
81#[derive(Debug, Clone)]
82pub struct EmbeddingConfig {
83    /// Path to the GGUF model file on disk.
84    pub model_path: PathBuf,
85    /// Path to the tokenizer.json file on disk.
86    pub tokenizer_path: PathBuf,
87    /// Whether to normalize embeddings (L2 norm = 1.0).
88    pub normalize: bool,
89}
90
91impl Default for EmbeddingConfig {
92    fn default() -> Self {
93        Self {
94            model_path: PathBuf::new(),
95            tokenizer_path: PathBuf::new(),
96            normalize: true,
97        }
98    }
99}
100
101#[cfg(feature = "embeddings")]
102mod candle_embedder {
103    use super::*;
104    use anyhow::{Context, Result as AnyResult};
105    use candle_core::quantized::gguf_file;
106    use candle_core::{DType, Device, Tensor};
107    use candle_transformers::models::quantized_qwen2::ModelWeights;
108    use std::fs::File;
109    use std::io::BufReader;
110    use tokenizers::Tokenizer;
111
112    /// Local embedding model using candle (pure Rust, no C++ dependency).
113    /// Loads Qwen3-Embedding-0.6B in GGUF format and runs on CPU.
114    pub struct CandleEmbedder {
115        model: std::sync::Mutex<ModelWeights>,
116        tokenizer: Tokenizer,
117        device: Device,
118        config: EmbeddingConfig,
119    }
120
121    impl CandleEmbedder {
122        /// Loads the model from a GGUF file on disk.
123        pub fn load(config: EmbeddingConfig) -> AnyResult<Self> {
124            let device = Device::Cpu;
125
126            // Open GGUF file
127            let file = File::open(&config.model_path)
128                .with_context(|| format!("Failed to open GGUF file: {:?}", config.model_path))?;
129            let mut reader = BufReader::new(file);
130
131            // Read GGUF content
132            let ct =
133                gguf_file::Content::read(&mut reader).context("Failed to read GGUF content")?;
134
135            // Build quantized model weights from GGUF
136            let model = ModelWeights::from_gguf(ct, &mut reader, &device)
137                .context("Failed to build quantized Qwen2 model from GGUF")?;
138
139            // Load tokenizer
140            let tokenizer = Tokenizer::from_file(&config.tokenizer_path).map_err(|e| {
141                anyhow::anyhow!(
142                    "Failed to load tokenizer from {:?}: {}",
143                    config.tokenizer_path,
144                    e
145                )
146            })?;
147
148            Ok(Self {
149                model: std::sync::Mutex::new(model),
150                tokenizer,
151                device,
152                config,
153            })
154        }
155
156        /// Mean pooling over token embeddings, then optional L2 normalization.
157        fn mean_pool(
158            &self,
159            token_embeddings: &Tensor,
160            attention_mask: &Tensor,
161        ) -> AnyResult<Tensor> {
162            let mask = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?; // [1, seq_len, 1]
163
164            let masked = token_embeddings.broadcast_mul(&mask)?;
165            let sum = masked.sum(1)?; // [1, hidden_size]
166
167            let mask_sum = mask.sum(1)?; // [1, 1]
168            let pooled = sum.broadcast_div(&mask_sum)?;
169
170            if self.config.normalize {
171                let norm = pooled.sqr()?.sum(1)?.sqrt()?;
172                let pooled = pooled.broadcast_div(&norm.unsqueeze(1)?)?;
173                Ok(pooled)
174            } else {
175                Ok(pooled)
176            }
177        }
178    }
179
180    impl Embedder for CandleEmbedder {
181        fn embed(&self, text: &str) -> AnyResult<Vec<f32>> {
182            let encoding = self
183                .tokenizer
184                .encode(text, true)
185                .map_err(|e| anyhow::anyhow!("Tokenization failed: {}", e))?;
186
187            let input_ids = encoding.get_ids();
188            let attention_mask = encoding.get_attention_mask();
189
190            // Create input_ids tensor [1, seq_len] as u32
191            let input_ids_tensor = Tensor::from_slice(
192                input_ids
193                    .iter()
194                    .map(|&v| v as u32)
195                    .collect::<Vec<_>>()
196                    .as_slice(),
197                (1, input_ids.len()),
198                &self.device,
199            )?;
200
201            // Create attention mask tensor [1, seq_len] as u32
202            let attention_mask_tensor = Tensor::from_slice(
203                attention_mask
204                    .iter()
205                    .map(|&v| v as u32)
206                    .collect::<Vec<_>>()
207                    .as_slice(),
208                (1, attention_mask.len()),
209                &self.device,
210            )?;
211
212            // Forward pass — quantized model returns hidden states
213            let mut model = self
214                .model
215                .lock()
216                .map_err(|e| anyhow::anyhow!("model lock poisoned: {}", e))?;
217            let embedded = model.forward(&input_ids_tensor, 0)?;
218            // embedded: [1, seq_len, hidden_size]
219
220            // Mean pooling with attention mask
221            let pooled = self.mean_pool(&embedded, &attention_mask_tensor)?;
222            // pooled: [1, hidden_size]
223
224            // Extract to vec
225            let full_embedding = pooled
226                .to_vec2::<f32>()?
227                .into_iter()
228                .next()
229                .unwrap_or_default();
230
231            // Matryoshka truncation: take first EMBED_DIM dimensions
232            let truncated: Vec<f32> = full_embedding.into_iter().take(EMBED_DIM).collect();
233
234            // Re-normalize after truncation
235            if self.config.normalize && !truncated.is_empty() {
236                let norm: f32 = truncated.iter().map(|v| v * v).sum::<f32>().sqrt();
237                if norm > 0.0 {
238                    return Ok(truncated.iter().map(|v| v / norm).collect());
239                }
240            }
241
242            Ok(truncated)
243        }
244    }
245}
246
247#[cfg(feature = "embeddings")]
248pub use candle_embedder::CandleEmbedder;
249
250/// Creates an embedder based on whether the feature is enabled and the model exists.
251#[cfg(feature = "embeddings")]
252pub fn create_embedder(config: EmbeddingConfig) -> Box<dyn Embedder> {
253    if config.model_path.exists() && config.tokenizer_path.exists() {
254        match CandleEmbedder::load(config) {
255            Ok(embedder) => {
256                tracing::info!("Local embedding model loaded successfully");
257                return Box::new(embedder);
258            }
259            Err(e) => {
260                tracing::warn!(
261                    "Failed to load embedding model: {}, falling back to text search",
262                    e
263                );
264            }
265        }
266    } else {
267        if !config.model_path.exists() {
268            tracing::info!(
269                "Embedding model not found at {:?}. Download with: huggingface-cli download {} {}",
270                config.model_path,
271                DEFAULT_MODEL_REPO,
272                DEFAULT_MODEL_FILE
273            );
274        }
275        if !config.tokenizer_path.exists() {
276            tracing::info!(
277                "Tokenizer not found at {:?}. Download with: huggingface-cli download {} {}",
278                config.tokenizer_path,
279                DEFAULT_TOKENIZER_REPO,
280                DEFAULT_TOKENIZER_FILE
281            );
282        }
283    }
284
285    Box::new(NoEmbedder)
286}
287
288/// Creates an embedder based on whether the feature is enabled and the model exists.
289#[cfg(not(feature = "embeddings"))]
290pub fn create_embedder(_config: EmbeddingConfig) -> Box<dyn Embedder> {
291    Box::new(NoEmbedder)
292}
293
294/// Convenience: check if embeddings are available at runtime.
295#[cfg(feature = "embeddings")]
296pub fn embeddings_available() -> bool {
297    true
298}
299
300/// Convenience: check if embeddings are available at runtime.
301#[cfg(not(feature = "embeddings"))]
302pub fn embeddings_available() -> bool {
303    false
304}
305
306#[cfg(test)]
307mod tests {
308    use super::*;
309
310    #[test]
311    fn test_no_embedder_returns_error() {
312        let embedder = NoEmbedder;
313        assert!(embedder.embed("test").is_err());
314    }
315
316    #[test]
317    fn test_embed_dim() {
318        assert_eq!(EMBED_DIM, 256);
319        assert_eq!(FULL_EMBED_DIM, 1024);
320    }
321
322    #[test]
323    fn test_create_embedder_without_model_file() {
324        let config = EmbeddingConfig::default();
325        let embedder = create_embedder(config);
326        // Without a model file on disk, should be NoEmbedder
327        assert!(embedder.embed("test").is_err());
328    }
329}