Skip to main content

nodedb_vector/
flat.rs

1// SPDX-License-Identifier: Apache-2.0
2
3//! Flat (brute-force) vector index for small collections.
4//!
5//! Simple linear scan over all stored vectors. No graph overhead, exact
6//! results. Automatically used when a collection has fewer than
7//! `DEFAULT_FLAT_INDEX_THRESHOLD` vectors (default 10K). Also serves as the
8//! search method for growing segments before HNSW construction.
9//!
10//! Complexity: O(N × D) per query where N = vectors, D = dimensions.
11
12use roaring::RoaringBitmap;
13
14use crate::distance::{DistanceMetric, distance};
15use crate::hnsw::SearchResult;
16
17/// Default threshold below which collections use flat index instead of HNSW.
18pub const DEFAULT_FLAT_INDEX_THRESHOLD: usize = 10_000;
19
20/// Flat vector index: append-only buffer with brute-force search.
21pub struct FlatIndex {
22    dim: usize,
23    metric: DistanceMetric,
24    /// Vectors stored contiguously for cache-friendly sequential scan.
25    data: Vec<f32>,
26    /// Tombstone bitmap: `deleted[i]` = true means vector i is soft-deleted.
27    deleted: Vec<bool>,
28    /// Number of live (non-deleted) vectors.
29    live_count: usize,
30}
31
32impl FlatIndex {
33    /// Create a new empty flat index.
34    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    /// Insert a vector. Returns the assigned vector ID.
45    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    /// Soft-delete a vector by ID.
61    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    /// Un-delete (clear the soft-delete tombstone of) a vector by ID. Reverses
73    /// [`FlatIndex::delete`] — used for transaction rollback so a rolled-back
74    /// delete restores the vector to the searchable set. Returns `true` if a
75    /// tombstone was actually cleared, `false` if the id was out of range or
76    /// already live.
77    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    /// Brute-force k-NN search with an explicit distance metric override.
89    /// Overrides the `self.metric` configured at collection creation time.
90    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    /// Brute-force k-NN search. Exact results — no approximation.
133    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    /// Search with a pre-filter bitmap (byte-array format).
171    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    /// Filtered search with an explicit metric override.
176    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    /// Search with a pre-filter bitmap applying a global id offset.
229    ///
230    /// `bitmap` is a serialized `RoaringBitmap` (matching the HNSW filter
231    /// format). Bit `i + id_offset` tests local id `i`. Used by multi-segment
232    /// collections where the bitmap holds GLOBAL vector ids. If the bytes
233    /// fail to deserialize, the search degrades to unfiltered.
234    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    /// Raw access bypassing tombstone filter — used by snapshot/restore.
308    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    /// Whether the given local id has been tombstoned.
319    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    /// Insert a vector that is already tombstoned (for checkpoint restore).
325    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        // No live_count increment — it's dead on arrival.
337        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}