Skip to main content

recern_vector/
collection.rs

1use std::collections::{BinaryHeap, HashMap};
2use std::fmt;
3use std::time::{Duration, Instant};
4
5use serde_json::Value;
6
7use crate::error::{Error, Result};
8use crate::filter::Filter;
9use crate::hnsw::{Candidate, Counters, Graph, HnswParams, Vectors};
10use crate::metric::{self, Metric};
11use crate::rng::SplitMix64;
12
13pub(crate) const DEFAULT_SEED: u64 = 0x5245_4345_524E_5645;
14const MAX_DIM: usize = 65_536;
15
16/// Filters estimated to match fewer records than this share are answered by
17/// scanning the matching records exactly instead of walking the graph.
18const EXACT_FILTER_SELECTIVITY: f64 = 0.02;
19const SELECTIVITY_SAMPLE: usize = 512;
20
21#[derive(Clone, Copy, Debug, PartialEq, Eq)]
22pub struct CollectionConfig {
23    pub dim: usize,
24    pub metric: Metric,
25    pub hnsw: HnswParams,
26}
27
28impl CollectionConfig {
29    pub fn new(dim: usize, metric: Metric) -> Self {
30        Self {
31            dim,
32            metric,
33            hnsw: HnswParams::default(),
34        }
35    }
36
37    pub fn with_hnsw(mut self, hnsw: HnswParams) -> Self {
38        self.hnsw = hnsw;
39        self
40    }
41
42    fn validate(&self) -> Result<()> {
43        if self.dim == 0 || self.dim > MAX_DIM {
44            return Err(Error::InvalidArgument(format!(
45                "dim must be between 1 and {MAX_DIM}"
46            )));
47        }
48        self.hnsw.validate()
49    }
50}
51
52/// A stored record. For cosine collections `vector` is the normalized vector.
53#[derive(Clone, Debug, PartialEq)]
54pub struct Record {
55    pub id: String,
56    pub vector: Vec<f32>,
57    pub metadata: Option<Value>,
58}
59
60#[derive(Clone, Debug, PartialEq)]
61pub struct SearchHit {
62    pub id: String,
63    pub distance: f32,
64    pub metadata: Option<Value>,
65}
66
67#[derive(Clone, Debug, Default)]
68pub struct SearchOptions {
69    /// Candidate list size; defaults to the collection's `ef_search`.
70    pub ef: Option<usize>,
71    /// Scan every record instead of using the index.
72    pub exact: bool,
73    pub filter: Option<Filter>,
74}
75
76impl SearchOptions {
77    pub fn ef(mut self, ef: usize) -> Self {
78        self.ef = Some(ef);
79        self
80    }
81
82    pub fn exact(mut self) -> Self {
83        self.exact = true;
84        self
85    }
86
87    pub fn filter(mut self, filter: Filter) -> Self {
88        self.filter = Some(filter);
89        self
90    }
91}
92
93/// How a query was answered.
94#[derive(Clone, Copy, Debug, PartialEq, Eq)]
95pub enum Strategy {
96    /// HNSW graph search (filtered results are collected while traversing).
97    Hnsw,
98    /// Full scan, requested explicitly.
99    Exact,
100    /// Full scan chosen automatically because the filter is very selective.
101    FilteredExact,
102}
103
104impl fmt::Display for Strategy {
105    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
106        f.write_str(match self {
107            Strategy::Hnsw => "hnsw",
108            Strategy::Exact => "exact",
109            Strategy::FilteredExact => "exact (selective filter)",
110        })
111    }
112}
113
114/// Search results together with the work it took to produce them.
115#[derive(Clone, Debug)]
116pub struct SearchReport {
117    pub hits: Vec<SearchHit>,
118    pub strategy: Strategy,
119    /// Candidate list size used by an HNSW search.
120    pub ef: Option<usize>,
121    pub visited: usize,
122    pub distance_computations: usize,
123    /// Estimated share of live records matching the filter.
124    pub filter_selectivity: Option<f64>,
125    pub elapsed: Duration,
126}
127
128#[derive(Clone, Debug, PartialEq)]
129pub struct CollectionStats {
130    pub name: String,
131    pub config: CollectionConfig,
132    pub live: usize,
133    /// Deleted or replaced records still occupying space until `compact`.
134    pub deleted: usize,
135    /// Nodes on each HNSW layer, bottom first.
136    pub nodes_per_layer: Vec<usize>,
137    pub avg_degree_layer0: f64,
138    /// Live records that graph search can never reach from the entry point.
139    pub unreachable: usize,
140    pub vector_bytes: usize,
141    pub graph_bytes: usize,
142    pub metadata_bytes: usize,
143}
144
145#[derive(Clone, Debug)]
146pub struct RecallOptions {
147    /// Number of stored vectors used as queries.
148    pub sample: usize,
149    pub k: usize,
150    pub ef_values: Vec<usize>,
151    pub seed: u64,
152}
153
154impl Default for RecallOptions {
155    fn default() -> Self {
156        Self {
157            sample: 100,
158            k: 10,
159            ef_values: vec![16, 32, 64, 128, 256],
160            seed: 42,
161        }
162    }
163}
164
165#[derive(Clone, Debug)]
166pub struct RecallPoint {
167    pub ef: usize,
168    /// Share of the exact top-k found by the index, averaged over queries.
169    pub recall: f64,
170    pub p50: Duration,
171    pub p95: Duration,
172}
173
174#[derive(Clone, Debug)]
175pub struct RecallReport {
176    pub k: usize,
177    pub sample: usize,
178    pub exact_p50: Duration,
179    pub points: Vec<RecallPoint>,
180}
181
182pub struct Collection {
183    pub(crate) name: String,
184    pub(crate) config: CollectionConfig,
185    pub(crate) vectors: Vectors,
186    pub(crate) ids: Vec<String>,
187    pub(crate) metadata: Vec<Option<Value>>,
188    pub(crate) deleted: Vec<bool>,
189    pub(crate) graph: Graph,
190    index: HashMap<String, u32>,
191}
192
193impl Collection {
194    pub(crate) fn new(name: &str, config: CollectionConfig) -> Result<Self> {
195        config.validate()?;
196        Ok(Self {
197            name: name.to_owned(),
198            config,
199            vectors: Vectors::new(config.dim),
200            ids: Vec::new(),
201            metadata: Vec::new(),
202            deleted: Vec::new(),
203            graph: Graph::new(config.hnsw, DEFAULT_SEED),
204            index: HashMap::new(),
205        })
206    }
207
208    /// Reassembles a collection read from disk, rebuilding the id index.
209    pub(crate) fn from_parts(
210        name: String,
211        config: CollectionConfig,
212        vectors: Vectors,
213        ids: Vec<String>,
214        metadata: Vec<Option<Value>>,
215        deleted: Vec<bool>,
216        graph: Graph,
217    ) -> Result<Self> {
218        config
219            .validate()
220            .map_err(|e| Error::Corrupt(format!("collection '{name}': {e}")))?;
221        let mut index = HashMap::with_capacity(ids.len());
222        for (node, id) in ids.iter().enumerate() {
223            if !deleted[node] && index.insert(id.clone(), node as u32).is_some() {
224                return Err(Error::Corrupt(format!(
225                    "collection '{name}': duplicate id '{id}'"
226                )));
227            }
228        }
229        Ok(Self {
230            name,
231            config,
232            vectors,
233            ids,
234            metadata,
235            deleted,
236            graph,
237            index,
238        })
239    }
240
241    pub fn name(&self) -> &str {
242        &self.name
243    }
244
245    pub fn config(&self) -> &CollectionConfig {
246        &self.config
247    }
248
249    /// Number of live records.
250    pub fn len(&self) -> usize {
251        self.index.len()
252    }
253
254    pub fn is_empty(&self) -> bool {
255        self.index.is_empty()
256    }
257
258    pub fn contains(&self, id: &str) -> bool {
259        self.index.contains_key(id)
260    }
261
262    /// Inserts a record, replacing any record with the same id.
263    pub fn upsert(&mut self, id: &str, vector: &[f32], metadata: Option<Value>) -> Result<()> {
264        let vector = self.prepare(vector)?;
265        if self.ids.len() >= u32::MAX as usize {
266            return Err(Error::InvalidArgument("collection is full".into()));
267        }
268        if let Some(&old) = self.index.get(id) {
269            self.deleted[old as usize] = true;
270        }
271        let node = self.append(id.to_owned(), &vector, metadata);
272        self.graph.insert(node, &self.vectors, self.config.metric);
273        Ok(())
274    }
275
276    /// Inserts many records, replacing records with the same ids, and links
277    /// them into the index using all available cores. The batch is atomic: if
278    /// any vector is invalid, the collection is left unchanged. Returns the
279    /// number of records written.
280    pub fn upsert_many<I, S, V>(&mut self, records: I) -> Result<usize>
281    where
282        I: IntoIterator<Item = (S, V, Option<Value>)>,
283        S: Into<String>,
284        V: AsRef<[f32]>,
285    {
286        self.upsert_many_with_threads(records, default_threads())
287    }
288
289    /// [`Collection::upsert_many`] with an explicit thread count. With one
290    /// thread the graph is identical to upserting the records one by one;
291    /// with more, link order (and so the exact graph) depends on scheduling.
292    pub fn upsert_many_with_threads<I, S, V>(&mut self, records: I, threads: usize) -> Result<usize>
293    where
294        I: IntoIterator<Item = (S, V, Option<Value>)>,
295        S: Into<String>,
296        V: AsRef<[f32]>,
297    {
298        let start = self.ids.len();
299        let mut replaced = Vec::new();
300        for (id, vector, metadata) in records {
301            let id = id.into();
302            let vector = match self.prepare(vector.as_ref()) {
303                Ok(vector) if self.ids.len() < u32::MAX as usize => vector,
304                Ok(_) => {
305                    self.rollback(start, replaced);
306                    return Err(Error::InvalidArgument("collection is full".into()));
307                }
308                Err(err) => {
309                    self.rollback(start, replaced);
310                    return Err(err);
311                }
312            };
313            if let Some(&old) = self.index.get(&id) {
314                self.deleted[old as usize] = true;
315                replaced.push((id.clone(), old));
316            }
317            self.append(id, &vector, metadata);
318        }
319        let end = self.ids.len();
320        self.graph.insert_batch(
321            start as u32..end as u32,
322            &self.vectors,
323            self.config.metric,
324            threads.max(1),
325        );
326        Ok(end - start)
327    }
328
329    /// Removes a record. Returns whether it existed.
330    ///
331    /// The node stays in the graph as a waypoint until [`Collection::compact`].
332    pub fn delete(&mut self, id: &str) -> bool {
333        match self.index.remove(id) {
334            Some(node) => {
335                self.deleted[node as usize] = true;
336                true
337            }
338            None => false,
339        }
340    }
341
342    pub fn get(&self, id: &str) -> Option<Record> {
343        let node = *self.index.get(id)?;
344        Some(Record {
345            id: id.to_owned(),
346            vector: self.vectors.get(node).to_vec(),
347            metadata: self.metadata[node as usize].clone(),
348        })
349    }
350
351    pub fn search(
352        &self,
353        query: &[f32],
354        k: usize,
355        options: &SearchOptions,
356    ) -> Result<Vec<SearchHit>> {
357        Ok(self.explain(query, k, options)?.hits)
358    }
359
360    /// Runs a search and reports how it was executed.
361    pub fn explain(
362        &self,
363        query: &[f32],
364        k: usize,
365        options: &SearchOptions,
366    ) -> Result<SearchReport> {
367        let start = Instant::now();
368        if k == 0 {
369            return Err(Error::InvalidArgument("k must be greater than zero".into()));
370        }
371        let query = self.prepare(query)?;
372        let filter = options.filter.as_ref();
373        let selectivity = filter.map(|f| self.estimate_selectivity(f));
374        let strategy = if options.exact {
375            Strategy::Exact
376        } else if selectivity.is_some_and(|s| s < EXACT_FILTER_SELECTIVITY) {
377            Strategy::FilteredExact
378        } else {
379            Strategy::Hnsw
380        };
381
382        let accept = |node: u32| {
383            !self.deleted[node as usize]
384                && filter.is_none_or(|f| f.matches(self.metadata[node as usize].as_ref()))
385        };
386        let mut counters = Counters::default();
387        let (found, ef) = match strategy {
388            Strategy::Hnsw => {
389                let ef = options.ef.unwrap_or(self.config.hnsw.ef_search).max(k);
390                let found = self.graph.search(
391                    &query,
392                    k,
393                    ef,
394                    &self.vectors,
395                    self.config.metric,
396                    &accept,
397                    &mut counters,
398                );
399                (found, Some(ef))
400            }
401            Strategy::Exact | Strategy::FilteredExact => {
402                (self.scan(&query, k, &accept, &mut counters), None)
403            }
404        };
405
406        Ok(SearchReport {
407            hits: found.into_iter().map(|c| self.hit(c)).collect(),
408            strategy,
409            ef,
410            visited: counters.visited,
411            distance_computations: counters.distance_computations,
412            filter_selectivity: selectivity,
413            elapsed: start.elapsed(),
414        })
415    }
416
417    pub fn stats(&self) -> CollectionStats {
418        let nodes = self.graph.len();
419        let reachable = self.graph.reachable();
420        let unreachable = (0..nodes)
421            .filter(|&n| !self.deleted[n] && !reachable[n])
422            .count();
423        let layer0_links: usize = self.graph.links.iter().map(|layers| layers[0].len()).sum();
424        let graph_bytes: usize = self
425            .graph
426            .links
427            .iter()
428            .map(|layers| {
429                size_of::<Vec<Vec<u32>>>()
430                    + layers
431                        .iter()
432                        .map(|l| size_of::<Vec<u32>>() + l.len() * 4)
433                        .sum::<usize>()
434            })
435            .sum();
436        let metadata_bytes = self
437            .metadata
438            .iter()
439            .flatten()
440            .map(|m| serde_json::to_vec(m).map_or(0, |b| b.len()))
441            .sum();
442
443        CollectionStats {
444            name: self.name.clone(),
445            config: self.config,
446            live: self.len(),
447            deleted: nodes - self.len(),
448            nodes_per_layer: self.graph.nodes_per_layer(),
449            avg_degree_layer0: if nodes == 0 {
450                0.0
451            } else {
452                layer0_links as f64 / nodes as f64
453            },
454            unreachable,
455            vector_bytes: self.vectors.data.len() * size_of::<f32>(),
456            graph_bytes,
457            metadata_bytes,
458        }
459    }
460
461    /// Measures how many of the exact nearest neighbors the index finds.
462    ///
463    /// Stored vectors are sampled as queries; each query's own record is
464    /// excluded from both the exact and the approximate results.
465    pub fn estimate_recall(&self, options: &RecallOptions) -> Result<RecallReport> {
466        if options.k == 0 || options.sample == 0 || options.ef_values.is_empty() {
467            return Err(Error::InvalidArgument(
468                "k, sample and ef values must be non-empty and greater than zero".into(),
469            ));
470        }
471        let mut live: Vec<u32> = (0..self.graph.len() as u32)
472            .filter(|&n| !self.deleted[n as usize])
473            .collect();
474        if live.len() < 2 {
475            return Err(Error::InvalidArgument(
476                "recall needs at least two records".into(),
477            ));
478        }
479
480        let mut rng = SplitMix64::new(options.seed);
481        let sample = options.sample.min(live.len());
482        for i in 0..sample {
483            let j = i + rng.below(live.len() - i);
484            live.swap(i, j);
485        }
486        let queries = &live[..sample];
487        let k = options.k.min(live.len() - 1);
488        let metric = self.config.metric;
489
490        let mut exact_times = Vec::with_capacity(sample);
491        let truths: Vec<Vec<u32>> = queries
492            .iter()
493            .map(|&q| {
494                let accept = |n: u32| n != q && !self.deleted[n as usize];
495                let start = Instant::now();
496                let found = self.scan(self.vectors.get(q), k, &accept, &mut Counters::default());
497                exact_times.push(start.elapsed());
498                found.into_iter().map(|c| c.id).collect()
499            })
500            .collect();
501
502        let mut points = Vec::with_capacity(options.ef_values.len());
503        for &ef in &options.ef_values {
504            let mut times = Vec::with_capacity(sample);
505            let mut found_total = 0;
506            for (&q, truth) in queries.iter().zip(&truths) {
507                let accept = |n: u32| n != q && !self.deleted[n as usize];
508                let start = Instant::now();
509                let found = self.graph.search(
510                    self.vectors.get(q),
511                    k,
512                    ef.max(k),
513                    &self.vectors,
514                    metric,
515                    &accept,
516                    &mut Counters::default(),
517                );
518                times.push(start.elapsed());
519                found_total += found.iter().filter(|c| truth.contains(&c.id)).count();
520            }
521            points.push(RecallPoint {
522                ef,
523                recall: found_total as f64 / (k * sample) as f64,
524                p50: percentile(&mut times, 0.50),
525                p95: percentile(&mut times, 0.95),
526            });
527        }
528
529        Ok(RecallReport {
530            k,
531            sample,
532            exact_p50: percentile(&mut exact_times, 0.50),
533            points,
534        })
535    }
536
537    /// Rebuilds the collection without deleted records. Returns how many
538    /// were removed.
539    pub fn compact(&mut self) -> usize {
540        let removed = self.graph.len() - self.len();
541        if removed == 0 {
542            return 0;
543        }
544        let mut fresh = Collection::new(&self.name, self.config)
545            .expect("configuration was validated when the collection was created");
546        for node in 0..self.graph.len() {
547            if !self.deleted[node] {
548                let vector = self.vectors.get(node as u32);
549                fresh.append(self.ids[node].clone(), vector, self.metadata[node].clone());
550            }
551        }
552        let live = fresh.ids.len() as u32;
553        fresh.graph.insert_batch(
554            0..live,
555            &fresh.vectors,
556            fresh.config.metric,
557            default_threads(),
558        );
559        *self = fresh;
560        removed
561    }
562
563    fn prepare(&self, vector: &[f32]) -> Result<Vec<f32>> {
564        if vector.len() != self.config.dim {
565            return Err(Error::DimensionMismatch {
566                expected: self.config.dim,
567                actual: vector.len(),
568            });
569        }
570        if vector.iter().any(|x| !x.is_finite()) {
571            return Err(Error::InvalidVector(
572                "contains NaN or infinite values".into(),
573            ));
574        }
575        let mut vector = vector.to_vec();
576        if self.config.metric == Metric::Cosine && !metric::normalize(&mut vector) {
577            return Err(Error::InvalidVector(
578                "zero vector has no direction for cosine".into(),
579            ));
580        }
581        Ok(vector)
582    }
583
584    /// Appends an already prepared record without linking it into the graph.
585    fn append(&mut self, id: String, vector: &[f32], metadata: Option<Value>) -> u32 {
586        let node = self.ids.len() as u32;
587        self.vectors.push(vector);
588        self.ids.push(id.clone());
589        self.metadata.push(metadata);
590        self.deleted.push(false);
591        self.index.insert(id, node);
592        node
593    }
594
595    /// Undoes the appends of a failed batch that started at node `start`.
596    fn rollback(&mut self, start: usize, replaced: Vec<(String, u32)>) {
597        for node in start..self.ids.len() {
598            let id = &self.ids[node];
599            if self.index.get(id) == Some(&(node as u32)) {
600                self.index.remove(id);
601            }
602        }
603        self.vectors.data.truncate(start * self.config.dim);
604        self.ids.truncate(start);
605        self.metadata.truncate(start);
606        self.deleted.truncate(start);
607        for (id, old) in replaced.into_iter().rev() {
608            if (old as usize) < start {
609                self.deleted[old as usize] = false;
610                self.index.insert(id, old);
611            }
612        }
613    }
614
615    fn scan(
616        &self,
617        query: &[f32],
618        k: usize,
619        accept: &dyn Fn(u32) -> bool,
620        counters: &mut Counters,
621    ) -> Vec<Candidate> {
622        let mut heap = BinaryHeap::with_capacity(k + 1);
623        for node in 0..self.graph.len() as u32 {
624            if !accept(node) {
625                continue;
626            }
627            counters.visited += 1;
628            counters.distance_computations += 1;
629            let dist = self.config.metric.distance(query, self.vectors.get(node));
630            if heap.len() < k {
631                heap.push(Candidate { dist, id: node });
632            } else if heap.peek().is_some_and(|w: &Candidate| dist < w.dist) {
633                heap.pop();
634                heap.push(Candidate { dist, id: node });
635            }
636        }
637        let mut out = heap.into_vec();
638        out.sort();
639        out
640    }
641
642    /// Estimates the share of live records matching `filter`. Small
643    /// collections are checked fully; larger ones from a random sample, since
644    /// an evenly spaced one aliases with periodic insertion patterns.
645    fn estimate_selectivity(&self, filter: &Filter) -> f64 {
646        let nodes = self.graph.len();
647        let mut rng = SplitMix64::new(DEFAULT_SEED);
648        let sample: Box<dyn Iterator<Item = usize>> = if nodes <= SELECTIVITY_SAMPLE {
649            Box::new(0..nodes)
650        } else {
651            Box::new((0..SELECTIVITY_SAMPLE).map(move |_| rng.below(nodes)))
652        };
653        let (mut checked, mut matched) = (0usize, 0usize);
654        for node in sample.filter(|&n| !self.deleted[n]) {
655            checked += 1;
656            if filter.matches(self.metadata[node].as_ref()) {
657                matched += 1;
658            }
659        }
660        if checked == 0 {
661            1.0
662        } else {
663            matched as f64 / checked as f64
664        }
665    }
666
667    fn hit(&self, candidate: Candidate) -> SearchHit {
668        let node = candidate.id as usize;
669        SearchHit {
670            id: self.ids[node].clone(),
671            distance: candidate.dist,
672            metadata: self.metadata[node].clone(),
673        }
674    }
675}
676
677fn default_threads() -> usize {
678    std::thread::available_parallelism().map_or(1, |n| n.get())
679}
680
681fn percentile(times: &mut [Duration], p: f64) -> Duration {
682    if times.is_empty() {
683        return Duration::ZERO;
684    }
685    times.sort_unstable();
686    times[((times.len() - 1) as f64 * p).round() as usize]
687}