1use std::collections::HashMap;
30use std::fs;
31use std::path::{Path, PathBuf};
32
33use super::hnsw::{HNSWIndex, SearchRuntimeOptions};
34use crate::utils::errors::DBError;
35use crate::utils::types::DistanceMetric;
36
37pub struct VectorIndex {
40 inner: HNSWIndex,
41 ids: Vec<String>,
42 by_id: HashMap<String, u64>,
43 path: PathBuf,
44}
45
46impl VectorIndex {
47 pub fn build(
52 path: impl AsRef<Path>,
53 entries: Vec<(String, Vec<f32>)>,
54 m: usize,
55 ef_construct: usize,
56 ) -> Result<Self, DBError> {
57 let path = path.as_ref().to_owned();
58 let dim = entries.first().map_or(0, |(_, v)| v.len());
59 let mut hnsw = HNSWIndex::new(DistanceMetric::Dot, m, ef_construct, 16, dim);
60 let ids: Vec<String> = entries.iter().map(|(id, _)| id.clone()).collect();
61 let indexed: Vec<(u64, Vec<f32>)> = entries
62 .into_iter()
63 .enumerate()
64 .map(|(i, (_, v))| (i as u64, v))
65 .collect();
66 hnsw.par_insert_batch(&indexed)?;
67 if !ids.is_empty() {
68 hnsw.reorder_rcm();
69 }
70 let by_id = ids
71 .iter()
72 .enumerate()
73 .map(|(i, id)| (id.clone(), i as u64))
74 .collect();
75 let index = Self {
76 inner: hnsw,
77 ids,
78 by_id,
79 path,
80 };
81 index.save()?;
82 Ok(index)
83 }
84
85 pub fn open(path: impl AsRef<Path>) -> Result<Self, DBError> {
87 let path = path.as_ref().to_owned();
88 let hnsw = HNSWIndex::load_from_path(Self::ann_path(&path))?;
89 let bytes = fs::read(Self::ids_path(&path))?;
90 let ids: Vec<String> = bincode::deserialize(&bytes)
91 .map_err(|e| DBError::SerializationError(anyhow::anyhow!("id map: {e}")))?;
92 let by_id = ids
93 .iter()
94 .enumerate()
95 .map(|(i, id)| (id.clone(), i as u64))
96 .collect();
97 Ok(Self {
98 inner: hnsw,
99 ids,
100 by_id,
101 path,
102 })
103 }
104
105 pub fn save(&self) -> Result<(), DBError> {
107 fs::create_dir_all(&self.path)?;
108 self.inner.save_to_path(Self::ann_path(&self.path))?;
109 let tmp = Self::ids_path(&self.path).with_extension("ids.tmp");
110 let bytes = bincode::serialize(&self.ids)
111 .map_err(|e| DBError::SerializationError(anyhow::anyhow!(e)))?;
112 fs::write(&tmp, &bytes)?;
113 fs::rename(&tmp, Self::ids_path(&self.path))?;
114 Ok(())
115 }
116
117 pub fn search(&self, query: &[f32], top_k: usize) -> Vec<(String, f32)> {
121 self.search_with_ef(query, top_k, top_k.max(64))
122 }
123
124 pub fn search_with_ef(
126 &self,
127 query: &[f32],
128 top_k: usize,
129 ef_search: usize,
130 ) -> Vec<(String, f32)> {
131 let opts = SearchRuntimeOptions {
132 ef_search: Some(ef_search.max(top_k)),
133 ..SearchRuntimeOptions::default()
134 };
135 let Ok(points) = self
136 .inner
137 .search_with_options(&query.to_vec(), top_k, &opts)
138 else {
139 return Vec::new();
140 };
141 points
142 .into_iter()
143 .filter_map(|p| {
144 let id = self.ids.get(p.id as usize)?;
145 Some((id.clone(), -p.sort_key))
146 })
147 .collect()
148 }
149
150 pub fn len(&self) -> usize {
152 self.ids.len()
153 }
154
155 pub fn is_empty(&self) -> bool {
156 self.ids.is_empty()
157 }
158
159 pub fn id_at(&self, i: usize) -> Option<&str> {
161 self.ids.get(i).map(String::as_str)
162 }
163
164 pub fn position_of(&self, id: &str) -> Option<u64> {
166 self.by_id.get(id).copied()
167 }
168
169 fn ann_path(root: &Path) -> PathBuf {
170 root.join("index.ann")
171 }
172
173 fn ids_path(root: &Path) -> PathBuf {
174 root.join("index.ids")
175 }
176}
177
178#[cfg(test)]
179mod tests {
180 use super::*;
181
182 #[test]
183 fn build_search_persist_round_trip() {
184 let dir = tempfile::tempdir().unwrap();
185 let entries = vec![
186 ("a".into(), vec![1.0_f32, 0.0, 0.0]),
187 ("b".into(), vec![0.0, 1.0, 0.0]),
188 ("c".into(), vec![0.0, 0.0, 1.0]),
189 ("d".into(), vec![0.707, 0.707, 0.0]),
190 ];
191 let idx = VectorIndex::build(dir.path(), entries, 4, 50).unwrap();
192 assert_eq!(idx.len(), 4);
193 assert_eq!(idx.position_of("a"), Some(0));
194
195 let results = idx.search(&[1.0, 0.0, 0.0], 2);
196 assert_eq!(results.len(), 2);
197 assert_eq!(results[0].0, "a");
198
199 let idx2 = VectorIndex::open(dir.path()).unwrap();
201 assert_eq!(idx2.len(), 4);
202 let results2 = idx2.search(&[1.0, 0.0, 0.0], 2);
203 assert_eq!(results, results2, "results must be identical after reload");
204 }
205}