1use roaring::RoaringBitmap;
13
14use crate::distance::{DistanceMetric, distance};
15use crate::hnsw::SearchResult;
16
17pub const DEFAULT_FLAT_INDEX_THRESHOLD: usize = 10_000;
19
20pub struct FlatIndex {
22 dim: usize,
23 metric: DistanceMetric,
24 data: Vec<f32>,
26 deleted: Vec<bool>,
28 live_count: usize,
30}
31
32impl FlatIndex {
33 pub fn new(dim: usize, metric: DistanceMetric) -> Self {
35 Self {
36 dim,
37 metric,
38 data: Vec::new(),
39 deleted: Vec::new(),
40 live_count: 0,
41 }
42 }
43
44 pub fn insert(&mut self, vector: Vec<f32>) -> u32 {
46 assert_eq!(
47 vector.len(),
48 self.dim,
49 "dimension mismatch: expected {}, got {}",
50 self.dim,
51 vector.len()
52 );
53 let id = self.len() as u32;
54 self.data.extend_from_slice(&vector);
55 self.deleted.push(false);
56 self.live_count += 1;
57 id
58 }
59
60 pub fn delete(&mut self, id: u32) -> bool {
62 let idx = id as usize;
63 if idx < self.deleted.len() && !self.deleted[idx] {
64 self.deleted[idx] = true;
65 self.live_count -= 1;
66 true
67 } else {
68 false
69 }
70 }
71
72 pub fn undelete(&mut self, id: u32) -> bool {
78 let idx = id as usize;
79 if idx < self.deleted.len() && self.deleted[idx] {
80 self.deleted[idx] = false;
81 self.live_count += 1;
82 true
83 } else {
84 false
85 }
86 }
87
88 pub fn search_with_metric(
91 &self,
92 query: &[f32],
93 top_k: usize,
94 metric: DistanceMetric,
95 ) -> Vec<SearchResult> {
96 assert_eq!(query.len(), self.dim);
97 let n = self.len();
98 if n == 0 || top_k == 0 {
99 return Vec::new();
100 }
101
102 let mut candidates: Vec<SearchResult> = Vec::with_capacity(n.min(top_k * 2));
103 for i in 0..n {
104 if self.deleted[i] {
105 continue;
106 }
107 let start = i * self.dim;
108 let vec_slice = &self.data[start..start + self.dim];
109 let dist = distance(query, vec_slice, metric);
110 candidates.push(SearchResult {
111 id: i as u32,
112 distance: dist,
113 });
114 }
115
116 if candidates.len() > top_k {
117 candidates.select_nth_unstable_by(top_k, |a, b| {
118 a.distance
119 .partial_cmp(&b.distance)
120 .unwrap_or(std::cmp::Ordering::Equal)
121 });
122 candidates.truncate(top_k);
123 }
124 candidates.sort_by(|a, b| {
125 a.distance
126 .partial_cmp(&b.distance)
127 .unwrap_or(std::cmp::Ordering::Equal)
128 });
129 candidates
130 }
131
132 pub fn search(&self, query: &[f32], top_k: usize) -> Vec<SearchResult> {
134 assert_eq!(query.len(), self.dim);
135 let n = self.len();
136 if n == 0 || top_k == 0 {
137 return Vec::new();
138 }
139
140 let mut candidates: Vec<SearchResult> = Vec::with_capacity(n.min(top_k * 2));
141 for i in 0..n {
142 if self.deleted[i] {
143 continue;
144 }
145 let start = i * self.dim;
146 let vec_slice = &self.data[start..start + self.dim];
147 let dist = distance(query, vec_slice, self.metric);
148 candidates.push(SearchResult {
149 id: i as u32,
150 distance: dist,
151 });
152 }
153
154 if candidates.len() > top_k {
155 candidates.select_nth_unstable_by(top_k, |a, b| {
156 a.distance
157 .partial_cmp(&b.distance)
158 .unwrap_or(std::cmp::Ordering::Equal)
159 });
160 candidates.truncate(top_k);
161 }
162 candidates.sort_by(|a, b| {
163 a.distance
164 .partial_cmp(&b.distance)
165 .unwrap_or(std::cmp::Ordering::Equal)
166 });
167 candidates
168 }
169
170 pub fn search_filtered(&self, query: &[f32], top_k: usize, bitmap: &[u8]) -> Vec<SearchResult> {
172 self.search_filtered_offset(query, top_k, bitmap, 0)
173 }
174
175 pub fn search_filtered_offset_with_metric(
177 &self,
178 query: &[f32],
179 top_k: usize,
180 bitmap: &[u8],
181 id_offset: u32,
182 metric: DistanceMetric,
183 ) -> Vec<SearchResult> {
184 assert_eq!(query.len(), self.dim);
185 let n = self.len();
186 if n == 0 || top_k == 0 {
187 return Vec::new();
188 }
189
190 let parsed = RoaringBitmap::deserialize_from(bitmap).ok();
191
192 let mut candidates: Vec<SearchResult> = Vec::with_capacity(top_k * 2);
193 for i in 0..n {
194 if self.deleted[i] {
195 continue;
196 }
197 if let Some(ref bm) = parsed {
198 let global = (i as u32).saturating_add(id_offset);
199 if !bm.contains(global) {
200 continue;
201 }
202 }
203 let start = i * self.dim;
204 let vec_slice = &self.data[start..start + self.dim];
205 let dist = distance(query, vec_slice, metric);
206 candidates.push(SearchResult {
207 id: i as u32,
208 distance: dist,
209 });
210 }
211
212 if candidates.len() > top_k {
213 candidates.select_nth_unstable_by(top_k, |a, b| {
214 a.distance
215 .partial_cmp(&b.distance)
216 .unwrap_or(std::cmp::Ordering::Equal)
217 });
218 candidates.truncate(top_k);
219 }
220 candidates.sort_by(|a, b| {
221 a.distance
222 .partial_cmp(&b.distance)
223 .unwrap_or(std::cmp::Ordering::Equal)
224 });
225 candidates
226 }
227
228 pub fn search_filtered_offset(
235 &self,
236 query: &[f32],
237 top_k: usize,
238 bitmap: &[u8],
239 id_offset: u32,
240 ) -> Vec<SearchResult> {
241 assert_eq!(query.len(), self.dim);
242 let n = self.len();
243 if n == 0 || top_k == 0 {
244 return Vec::new();
245 }
246
247 let parsed = RoaringBitmap::deserialize_from(bitmap).ok();
248
249 let mut candidates: Vec<SearchResult> = Vec::with_capacity(top_k * 2);
250 for i in 0..n {
251 if self.deleted[i] {
252 continue;
253 }
254 if let Some(ref bm) = parsed {
255 let global = (i as u32).saturating_add(id_offset);
256 if !bm.contains(global) {
257 continue;
258 }
259 }
260 let start = i * self.dim;
261 let vec_slice = &self.data[start..start + self.dim];
262 let dist = distance(query, vec_slice, self.metric);
263 candidates.push(SearchResult {
264 id: i as u32,
265 distance: dist,
266 });
267 }
268
269 if candidates.len() > top_k {
270 candidates.select_nth_unstable_by(top_k, |a, b| {
271 a.distance
272 .partial_cmp(&b.distance)
273 .unwrap_or(std::cmp::Ordering::Equal)
274 });
275 candidates.truncate(top_k);
276 }
277 candidates.sort_by(|a, b| {
278 a.distance
279 .partial_cmp(&b.distance)
280 .unwrap_or(std::cmp::Ordering::Equal)
281 });
282 candidates
283 }
284
285 pub fn len(&self) -> usize {
286 self.deleted.len()
287 }
288
289 pub fn live_count(&self) -> usize {
290 self.live_count
291 }
292
293 pub fn is_empty(&self) -> bool {
294 self.live_count == 0
295 }
296
297 pub fn get_vector(&self, id: u32) -> Option<&[f32]> {
298 let idx = id as usize;
299 if idx < self.deleted.len() && !self.deleted[idx] {
300 let start = idx * self.dim;
301 Some(&self.data[start..start + self.dim])
302 } else {
303 None
304 }
305 }
306
307 pub fn get_vector_raw(&self, id: u32) -> Option<&[f32]> {
309 let idx = id as usize;
310 if idx < self.deleted.len() {
311 let start = idx * self.dim;
312 Some(&self.data[start..start + self.dim])
313 } else {
314 None
315 }
316 }
317
318 pub fn is_deleted(&self, id: u32) -> bool {
320 let idx = id as usize;
321 idx < self.deleted.len() && self.deleted[idx]
322 }
323
324 pub fn insert_tombstoned(&mut self, vector: Vec<f32>) -> u32 {
326 assert_eq!(
327 vector.len(),
328 self.dim,
329 "dimension mismatch: expected {}, got {}",
330 self.dim,
331 vector.len()
332 );
333 let id = self.len() as u32;
334 self.data.extend_from_slice(&vector);
335 self.deleted.push(true);
336 id
338 }
339
340 pub fn dim(&self) -> usize {
341 self.dim
342 }
343
344 pub fn metric(&self) -> DistanceMetric {
345 self.metric
346 }
347
348 pub fn tombstone_count(&self) -> usize {
349 self.len().saturating_sub(self.live_count)
350 }
351}
352
353#[cfg(test)]
354mod tests {
355 use super::*;
356
357 #[test]
358 fn insert_and_search() {
359 let mut idx = FlatIndex::new(3, DistanceMetric::L2);
360 for i in 0..100u32 {
361 idx.insert(vec![i as f32, 0.0, 0.0]);
362 }
363 assert_eq!(idx.len(), 100);
364 assert_eq!(idx.live_count(), 100);
365
366 let results = idx.search(&[50.0, 0.0, 0.0], 3);
367 assert_eq!(results.len(), 3);
368 assert_eq!(results[0].id, 50);
369 assert!(results[0].distance < 0.01);
370 }
371
372 #[test]
373 fn delete_excludes_from_search() {
374 let mut idx = FlatIndex::new(2, DistanceMetric::L2);
375 idx.insert(vec![0.0, 0.0]);
376 idx.insert(vec![1.0, 0.0]);
377 idx.insert(vec![2.0, 0.0]);
378
379 assert!(idx.delete(1));
380 assert_eq!(idx.live_count(), 2);
381
382 let results = idx.search(&[1.0, 0.0], 3);
383 assert_eq!(results.len(), 2);
384 assert!(results.iter().all(|r| r.id != 1));
385 }
386
387 #[test]
388 fn exact_results() {
389 let mut idx = FlatIndex::new(2, DistanceMetric::Cosine);
390 idx.insert(vec![1.0, 0.0]);
391 idx.insert(vec![0.0, 1.0]);
392 idx.insert(vec![1.0, 1.0]);
393
394 let results = idx.search(&[1.0, 0.0], 1);
395 assert_eq!(results.len(), 1);
396 assert_eq!(results[0].id, 0);
397 }
398
399 #[test]
400 fn empty_search() {
401 let idx = FlatIndex::new(3, DistanceMetric::L2);
402 let results = idx.search(&[1.0, 0.0, 0.0], 5);
403 assert!(results.is_empty());
404 }
405
406 #[test]
407 fn filtered_search() {
408 let mut idx = FlatIndex::new(2, DistanceMetric::L2);
409 for i in 0..8u32 {
410 idx.insert(vec![i as f32, 0.0]);
411 }
412 let bitmap = vec![0b11001100u8];
413 let results = idx.search_filtered(&[3.0, 0.0], 2, &bitmap);
414 assert_eq!(results.len(), 2);
415 assert_eq!(results[0].id, 3);
416 assert_eq!(results[1].id, 2);
417 }
418}