Skip to main content

llm_kernel/embedding/
turbovec.rs

1//! TurboQuant-backed vector index implementation.
2
3use std::path::Path;
4
5use super::vector_index::{SearchHit, VectorIndex};
6use crate::error::{KernelError, Result};
7
8/// Compressed vector index backed by TurboQuant.
9///
10/// Wraps `turbovec::IdMapIndex` with dimension validation and a consistent
11/// error-handling layer. Supports online ingest (no training step),
12/// filtered search with allowlists, and persistence via `save`/`load`.
13pub struct TurbovecIndex {
14    inner: turbovec::IdMapIndex,
15    dim: usize,
16    bit_width: u8,
17    /// Persisted model/policy metadata. `None` for in-memory indices created
18    /// without it, or for loaded indices whose meta predates these fields.
19    meta: Option<IndexMeta>,
20}
21
22impl std::fmt::Debug for TurbovecIndex {
23    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
24        f.debug_struct("TurbovecIndex")
25            .field("dim", &self.dim)
26            .field("bit_width", &self.bit_width)
27            .field("len", &self.inner.len())
28            .finish()
29    }
30}
31
32impl TurbovecIndex {
33    /// Create a new index for vectors of the given dimension.
34    ///
35    /// `bit_width` must be 2 or 4, controlling the quantization level:
36    /// - **2-bit**: 16x compression, lower recall at low k
37    /// - **4-bit**: 8x compression, higher recall (recommended default)
38    pub fn new(dim: usize, bit_width: u8) -> Result<Self> {
39        if bit_width != 2 && bit_width != 4 {
40            return Err(KernelError::Embedding(format!(
41                "bit_width must be 2 or 4, got {bit_width}"
42            )));
43        }
44        let inner = turbovec::IdMapIndex::new(dim, bit_width as usize)
45            .map_err(|e| KernelError::Embedding(format!("failed to create index: {e}")))?;
46        Ok(Self {
47            inner,
48            dim,
49            bit_width,
50            meta: None,
51        })
52    }
53
54    /// Create with model/policy metadata so it persists on [`save`](VectorIndex::save).
55    /// Consumers compare these on load to decide whether a rebuild is needed
56    /// (e.g. a prefix-policy change with the same `dim`).
57    pub fn with_meta(
58        dim: usize,
59        bit_width: u8,
60        model_id: Option<String>,
61        prefix_policy: Option<String>,
62        schema_version: Option<u32>,
63    ) -> Result<Self> {
64        let mut idx = Self::new(dim, bit_width)?;
65        idx.meta = Some(IndexMeta {
66            dim,
67            bit_width,
68            model_id,
69            prefix_policy,
70            schema_version,
71        });
72        Ok(idx)
73    }
74
75    /// Persisted metadata (model id, prefix policy, schema version), if any.
76    pub fn meta(&self) -> Option<&IndexMeta> {
77        self.meta.as_ref()
78    }
79
80    /// Quantization bit width (2 or 4).
81    pub fn bit_width(&self) -> u8 {
82        self.bit_width
83    }
84
85    /// Load a previously saved index from disk.
86    ///
87    /// This is an inherent method (not on the `VectorIndex` trait) so that the
88    /// trait remains fully object-safe. Callers must use the concrete type:
89    /// `TurbovecIndex::load(path)`.
90    pub fn load(path: &Path) -> Result<Self> {
91        let inner = turbovec::IdMapIndex::load(path)
92            .map_err(|e| KernelError::Embedding(format!("failed to load vector index: {e}")))?;
93        let meta_path = path.with_extension("meta.json");
94        let meta: IndexMeta = serde_json::from_str(&std::fs::read_to_string(&meta_path)?)
95            .map_err(KernelError::embedding)?;
96        if meta.bit_width != 2 && meta.bit_width != 4 {
97            return Err(KernelError::Embedding(format!(
98                "corrupted index meta: bit_width must be 2 or 4, got {}",
99                meta.bit_width
100            )));
101        }
102        if meta.dim == 0 {
103            return Err(KernelError::Embedding(
104                "corrupted index meta: dim must be positive, got 0".into(),
105            ));
106        }
107
108        // Cross-validate: loaded index vs sidecar metadata.
109        let inner_dim = inner.dim();
110        if inner_dim != 0 && inner_dim != meta.dim {
111            return Err(KernelError::Embedding(format!(
112                "index-meta mismatch: index dim={inner_dim}, meta dim={}",
113                meta.dim
114            )));
115        }
116        let inner_bw = inner.bit_width();
117        if inner_bw != meta.bit_width as usize {
118            return Err(KernelError::Embedding(format!(
119                "index-meta mismatch: index bit_width={inner_bw}, meta bit_width={}",
120                meta.bit_width
121            )));
122        }
123
124        Ok(Self {
125            inner,
126            dim: meta.dim,
127            bit_width: meta.bit_width,
128            meta: Some(meta),
129        })
130    }
131
132    fn validate_dim(&self, v: &[f32]) -> Result<()> {
133        if v.len() != self.dim {
134            return Err(KernelError::Embedding(format!(
135                "vector dimension mismatch: expected {}, got {}",
136                self.dim,
137                v.len()
138            )));
139        }
140        Ok(())
141    }
142
143    fn validate_dims(&self, vectors: &[Vec<f32>]) -> Result<()> {
144        for v in vectors {
145            self.validate_dim(v)?;
146        }
147        Ok(())
148    }
149}
150
151impl VectorIndex for TurbovecIndex {
152    fn add(&mut self, vectors: &[Vec<f32>]) -> Result<()> {
153        if vectors.is_empty() {
154            return Ok(());
155        }
156        self.validate_dims(vectors)?;
157        let start_id = self.inner.len() as u64;
158        let ids: Vec<u64> = (start_id..start_id + vectors.len() as u64).collect();
159        // Skip validation — already checked above.
160        let flat: Vec<f32> = vectors.iter().flat_map(|v| v.iter().copied()).collect();
161        self.inner
162            .add_with_ids_2d(&flat, self.dim, &ids)
163            .map_err(|e| KernelError::Embedding(format!("add failed: {e}")))?;
164        Ok(())
165    }
166
167    fn add_with_ids(&mut self, vectors: &[Vec<f32>], ids: &[u64]) -> Result<()> {
168        if vectors.len() != ids.len() {
169            return Err(KernelError::Embedding(format!(
170                "vectors ({} entries) and ids ({} entries) must have the same length",
171                vectors.len(),
172                ids.len()
173            )));
174        }
175        self.validate_dims(vectors)?;
176        let flat: Vec<f32> = vectors.iter().flat_map(|v| v.iter().copied()).collect();
177        self.inner
178            .add_with_ids_2d(&flat, self.dim, ids)
179            .map_err(|e| {
180                // Tag the failure kind so consumers can distinguish "id already
181                // present" (a routine re-ingest, not data loss) from a real
182                // backend failure, without brittle full-string matching.
183                let kind = if e.to_string().contains("already present") {
184                    "duplicate_id"
185                } else {
186                    "backend"
187                };
188                KernelError::Embedding(format!("add failed[{kind}]: {e}"))
189            })?;
190        Ok(())
191    }
192
193    fn remove(&mut self, ids: &[u64]) -> Result<()> {
194        for &id in ids {
195            self.inner.remove(id);
196        }
197        Ok(())
198    }
199
200    fn search(&self, query: &[f32], k: usize) -> Result<Vec<SearchHit>> {
201        self.validate_dim(query)?;
202        if self.inner.is_empty() {
203            return Ok(vec![]);
204        }
205        let (scores, ids) = self.inner.search(query, k);
206        Ok(scores
207            .into_iter()
208            .zip(ids)
209            .map(|(score, id)| SearchHit { id, score })
210            .collect())
211    }
212
213    fn search_filtered(
214        &self,
215        query: &[f32],
216        k: usize,
217        allowlist: &[u64],
218    ) -> Result<Vec<SearchHit>> {
219        self.validate_dim(query)?;
220        if self.inner.is_empty() || allowlist.is_empty() {
221            return Ok(vec![]);
222        }
223        let (scores, ids) = self.inner.search_with_allowlist(query, k, Some(allowlist));
224        Ok(scores
225            .into_iter()
226            .zip(ids)
227            .map(|(score, id)| SearchHit { id, score })
228            .collect())
229    }
230
231    fn len(&self) -> usize {
232        self.inner.len()
233    }
234
235    fn is_empty(&self) -> bool {
236        self.inner.is_empty()
237    }
238
239    fn dim(&self) -> usize {
240        self.dim
241    }
242
243    fn save(&self, path: &Path) -> Result<()> {
244        // Atomic save: write to temp files, fsync, then rename.
245        let tmp_index = path.with_extension("tvim.tmp");
246        let tmp_meta = path.with_extension("meta.tmp");
247
248        self.inner
249            .write(&tmp_index)
250            .map_err(|e| KernelError::Embedding(format!("failed to write vector index: {e}")))?;
251
252        let meta = self.meta.clone().unwrap_or(IndexMeta {
253            dim: self.dim,
254            bit_width: self.bit_width,
255            model_id: None,
256            prefix_policy: None,
257            schema_version: None,
258        });
259        let json = serde_json::to_string_pretty(&meta).map_err(KernelError::embedding)?;
260        std::fs::write(&tmp_meta, &json)?;
261
262        // Fsync temp files to ensure data is on disk.
263        if let Ok(f) = std::fs::File::open(&tmp_index) {
264            let _ = f.sync_all();
265        }
266        if let Ok(f) = std::fs::File::open(&tmp_meta) {
267            let _ = f.sync_all();
268        }
269
270        // Atomic rename — POSIX guarantees rename is atomic.
271        std::fs::rename(&tmp_meta, path.with_extension("meta.json"))?;
272        std::fs::rename(&tmp_index, path)?;
273
274        Ok(())
275    }
276}
277
278/// Persisted vector-index metadata sidecar.
279///
280/// Stored as `vectors.meta.json` next to the index file. Consumers compare these
281/// on load to decide whether a rebuild is needed (e.g. a prefix-policy change
282/// with the same `dim`).
283#[derive(Debug, serde::Serialize, serde::Deserialize, Clone)]
284pub struct IndexMeta {
285    /// Embedding dimensionality (must match the model).
286    pub dim: usize,
287    /// Quantization bit width (2 or 4).
288    pub bit_width: u8,
289    /// Model identifier (e.g. "intfloat/multilingual-e5-small"). `None` on
290    /// indices written before this field existed — consumers treat that as
291    /// "unknown, rebuild" so a prefix/policy change forces a rebuild even when
292    /// `dim` is unchanged.
293    #[serde(default, skip_serializing_if = "Option::is_none")]
294    pub model_id: Option<String>,
295    /// Embedding policy tag (e.g. "e5-query-doc-v1"). Lets callers detect that
296    /// the index was built with a different prefix scheme without bumping dim.
297    #[serde(default, skip_serializing_if = "Option::is_none")]
298    pub prefix_policy: Option<String>,
299    /// Caller-defined schema version for the vector store layout (chunking etc.).
300    #[serde(default, skip_serializing_if = "Option::is_none")]
301    pub schema_version: Option<u32>,
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307    use tempfile::TempDir;
308
309    fn make_index(dim: usize, bit_width: u8) -> TurbovecIndex {
310        TurbovecIndex::new(dim, bit_width).unwrap()
311    }
312
313    fn random_vector(dim: usize, seed: f32) -> Vec<f32> {
314        (0..dim).map(|i| (seed + i as f32 * 0.001).sin()).collect()
315    }
316
317    #[test]
318    fn new_valid_bit_widths() {
319        assert!(TurbovecIndex::new(128, 2).is_ok());
320        assert!(TurbovecIndex::new(128, 4).is_ok());
321    }
322
323    #[test]
324    fn new_invalid_bit_width() {
325        assert!(TurbovecIndex::new(128, 3).is_err());
326        assert!(TurbovecIndex::new(128, 8).is_err());
327        assert!(TurbovecIndex::new(128, 1).is_err());
328    }
329
330    #[test]
331    fn add_and_len() {
332        let mut idx = make_index(64, 4);
333        assert!(idx.is_empty());
334        idx.add(&[random_vector(64, 1.0), random_vector(64, 2.0)])
335            .unwrap();
336        assert_eq!(idx.len(), 2);
337    }
338
339    #[test]
340    fn add_empty() {
341        let mut idx = make_index(64, 4);
342        idx.add(&[]).unwrap();
343        assert!(idx.is_empty());
344    }
345
346    #[test]
347    fn add_with_explicit_ids() {
348        let mut idx = make_index(64, 4);
349        idx.add_with_ids(&[random_vector(64, 1.0)], &[42u64])
350            .unwrap();
351        assert_eq!(idx.len(), 1);
352    }
353
354    #[test]
355    fn add_dimension_mismatch() {
356        let mut idx = make_index(64, 4);
357        let result = idx.add(&[vec![0.0; 32]]);
358        assert!(result.is_err());
359        assert!(
360            result
361                .unwrap_err()
362                .to_string()
363                .contains("dimension mismatch")
364        );
365    }
366
367    #[test]
368    fn add_with_ids_length_mismatch() {
369        let mut idx = make_index(64, 4);
370        let result = idx.add_with_ids(&[random_vector(64, 1.0), random_vector(64, 2.0)], &[1u64]);
371        assert!(result.is_err());
372        assert!(result.unwrap_err().to_string().contains("same length"));
373    }
374
375    #[test]
376    fn search_empty_index() {
377        let idx = make_index(64, 4);
378        let hits = idx.search(&random_vector(64, 1.0), 5).unwrap();
379        assert!(hits.is_empty());
380    }
381
382    #[test]
383    fn search_returns_nearest() {
384        let mut idx = make_index(64, 4);
385        let target = random_vector(64, 3.0);
386        idx.add_with_ids(
387            &[
388                random_vector(64, 100.0),
389                target.clone(),
390                random_vector(64, 200.0),
391            ],
392            &[0u64, 1u64, 2u64],
393        )
394        .unwrap();
395        let hits = idx.search(&target, 1).unwrap();
396        assert_eq!(hits.len(), 1);
397        assert_eq!(hits[0].id, 1);
398    }
399
400    #[test]
401    fn search_dimension_mismatch() {
402        let mut idx = make_index(64, 4);
403        idx.add(&[random_vector(64, 1.0)]).unwrap();
404        let result = idx.search(&[0.0; 32], 1);
405        assert!(result.is_err());
406    }
407
408    #[test]
409    fn search_filtered_with_allowlist() {
410        let mut idx = make_index(64, 4);
411        idx.add_with_ids(
412            &[
413                random_vector(64, 1.0),
414                random_vector(64, 2.0),
415                random_vector(64, 3.0),
416            ],
417            &[10u64, 20u64, 30u64],
418        )
419        .unwrap();
420        let hits = idx
421            .search_filtered(&random_vector(64, 1.0), 10, &[20u64, 30u64])
422            .unwrap();
423        let ids: Vec<u64> = hits.iter().map(|h| h.id).collect();
424        assert!(ids.contains(&20));
425        assert!(ids.contains(&30));
426        assert!(!ids.contains(&10));
427    }
428
429    #[test]
430    fn search_filtered_empty_allowlist() {
431        let mut idx = make_index(64, 4);
432        idx.add(&[random_vector(64, 1.0)]).unwrap();
433        let hits = idx
434            .search_filtered(&random_vector(64, 1.0), 5, &[])
435            .unwrap();
436        assert!(hits.is_empty());
437    }
438
439    #[test]
440    fn save_load_roundtrip() {
441        let dir = TempDir::new().unwrap();
442        let path = dir.path().join("test.tvim");
443        let mut idx = make_index(64, 4);
444        idx.add_with_ids(
445            &[random_vector(64, 1.0), random_vector(64, 2.0)],
446            &[100u64, 200u64],
447        )
448        .unwrap();
449        idx.save(&path).unwrap();
450        let loaded = TurbovecIndex::load(&path).unwrap();
451        assert_eq!(loaded.dim(), 64);
452        assert_eq!(loaded.bit_width(), 4);
453        assert_eq!(loaded.len(), 2);
454    }
455
456    #[test]
457    fn load_rejects_corrupted_meta() {
458        let dir = TempDir::new().unwrap();
459        let path = dir.path().join("corrupt.tvim");
460        let mut idx = make_index(64, 4);
461        idx.add(&[random_vector(64, 1.0)]).unwrap();
462        idx.save(&path).unwrap();
463        let meta_path = path.with_extension("meta.json");
464        std::fs::write(&meta_path, r#"{"dim": 64, "bit_width": 7}"#).unwrap();
465        let result = TurbovecIndex::load(&path);
466        assert!(result.is_err());
467        assert!(result.unwrap_err().to_string().contains("bit_width"));
468    }
469
470    #[test]
471    fn load_rejects_zero_dim() {
472        let dir = TempDir::new().unwrap();
473        let path = dir.path().join("zero.tvim");
474        let mut idx = make_index(64, 4);
475        idx.add(&[random_vector(64, 1.0)]).unwrap();
476        idx.save(&path).unwrap();
477        let meta_path = path.with_extension("meta.json");
478        std::fs::write(&meta_path, r#"{"dim": 0, "bit_width": 4}"#).unwrap();
479        let result = TurbovecIndex::load(&path);
480        assert!(result.is_err());
481        assert!(result.unwrap_err().to_string().contains("dim"));
482    }
483
484    #[test]
485    fn dim_and_bit_width_accessors() {
486        let idx = make_index(128, 2);
487        assert_eq!(idx.dim(), 128);
488        assert_eq!(idx.bit_width(), 2);
489    }
490
491    #[test]
492    fn trait_object_compatibility() {
493        let mut idx: Box<dyn VectorIndex> = Box::new(make_index(64, 4));
494        idx.add(&[random_vector(64, 1.0)]).unwrap();
495        assert_eq!(idx.len(), 1);
496        assert!(!idx.is_empty());
497    }
498
499    #[test]
500    fn remove_existing_id() {
501        let mut idx = make_index(64, 4);
502        idx.add_with_ids(
503            &[
504                random_vector(64, 1.0),
505                random_vector(64, 2.0),
506                random_vector(64, 3.0),
507            ],
508            &[10u64, 20u64, 30u64],
509        )
510        .unwrap();
511        assert_eq!(idx.len(), 3);
512        idx.remove(&[20u64]).unwrap();
513        assert_eq!(idx.len(), 2);
514        let hits = idx.search(&random_vector(64, 2.0), 10).unwrap();
515        let ids: Vec<u64> = hits.iter().map(|h| h.id).collect();
516        assert!(!ids.contains(&20));
517    }
518
519    #[test]
520    fn remove_nonexistent_id() {
521        let mut idx = make_index(64, 4);
522        idx.add_with_ids(&[random_vector(64, 1.0)], &[1u64])
523            .unwrap();
524        idx.remove(&[999u64]).unwrap();
525        assert_eq!(idx.len(), 1);
526    }
527
528    #[test]
529    fn remove_empty_ids() {
530        let mut idx = make_index(64, 4);
531        idx.add(&[random_vector(64, 1.0)]).unwrap();
532        idx.remove(&[]).unwrap();
533        assert_eq!(idx.len(), 1);
534    }
535
536    #[test]
537    fn remove_via_trait_object() {
538        let mut idx: Box<dyn VectorIndex> = Box::new(make_index(64, 4));
539        idx.add_with_ids(&[random_vector(64, 1.0)], &[42u64])
540            .unwrap();
541        idx.remove(&[42u64]).unwrap();
542        assert!(idx.is_empty());
543    }
544
545    #[test]
546    fn load_detects_dim_mismatch() {
547        let dir = TempDir::new().unwrap();
548        let path = dir.path().join("mismatch.tvim");
549        let mut idx = make_index(64, 4);
550        idx.add(&[random_vector(64, 1.0)]).unwrap();
551        idx.save(&path).unwrap();
552        let meta_path = path.with_extension("meta.json");
553        std::fs::write(&meta_path, r#"{"dim": 128, "bit_width": 4}"#).unwrap();
554        let result = TurbovecIndex::load(&path);
555        assert!(result.is_err());
556        let msg = result.unwrap_err().to_string();
557        assert!(
558            msg.contains("mismatch") || msg.contains("dim"),
559            "expected mismatch error, got: {msg}"
560        );
561    }
562}