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
16const 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#[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 pub ef: Option<usize>,
71 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
95pub enum Strategy {
96 Hnsw,
98 Exact,
100 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#[derive(Clone, Debug)]
116pub struct SearchReport {
117 pub hits: Vec<SearchHit>,
118 pub strategy: Strategy,
119 pub ef: Option<usize>,
121 pub visited: usize,
122 pub distance_computations: usize,
123 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 pub deleted: usize,
135 pub nodes_per_layer: Vec<usize>,
137 pub avg_degree_layer0: f64,
138 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 let distance = match self.config.metric {
672 Metric::Cosine => candidate.dist.max(0.0),
673 _ => candidate.dist,
674 };
675 SearchHit {
676 id: self.ids[node].clone(),
677 distance,
678 metadata: self.metadata[node].clone(),
679 }
680 }
681}
682
683fn default_threads() -> usize {
684 std::thread::available_parallelism().map_or(1, |n| n.get())
685}
686
687fn percentile(times: &mut [Duration], p: f64) -> Duration {
688 if times.is_empty() {
689 return Duration::ZERO;
690 }
691 times.sort_unstable();
692 times[((times.len() - 1) as f64 * p).round() as usize]
693}