1use std::path::Path;
4
5use super::vector_index::{SearchHit, VectorIndex};
6use crate::error::{KernelError, Result};
7
8pub struct TurbovecIndex {
14 inner: turbovec::IdMapIndex,
15 dim: usize,
16 bit_width: u8,
17 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 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 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 pub fn meta(&self) -> Option<&IndexMeta> {
77 self.meta.as_ref()
78 }
79
80 pub fn bit_width(&self) -> u8 {
82 self.bit_width
83 }
84
85 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 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 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 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 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 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 std::fs::rename(&tmp_meta, path.with_extension("meta.json"))?;
272 std::fs::rename(&tmp_index, path)?;
273
274 Ok(())
275 }
276}
277
278#[derive(Debug, serde::Serialize, serde::Deserialize, Clone)]
284pub struct IndexMeta {
285 pub dim: usize,
287 pub bit_width: u8,
289 #[serde(default, skip_serializing_if = "Option::is_none")]
294 pub model_id: Option<String>,
295 #[serde(default, skip_serializing_if = "Option::is_none")]
298 pub prefix_policy: Option<String>,
299 #[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}