Skip to main content

recern_vector/
hnsw.rs

1//! Hierarchical Navigable Small World graph (Malkov & Yashunin, 2016).
2//!
3//! The graph only stores links. Vectors live in [`Vectors`] and are passed in
4//! by the owning collection, which also decides which nodes are acceptable
5//! results (deleted and filtered-out nodes are still traversed as waypoints).
6//!
7//! Link access goes through [`ReadLinks`] / [`Links`] so the same insertion
8//! code runs on a plain `&mut` graph (single inserts, queries) and on a graph
9//! with one mutex per node (parallel batch build, as in hnswlib).
10
11use std::cell::RefCell;
12use std::cmp::{Ordering, Reverse};
13use std::collections::BinaryHeap;
14use std::ops::Range;
15use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering as AtomicOrdering};
16use std::sync::{Mutex, MutexGuard, PoisonError};
17
18use crate::error::{Error, Result};
19use crate::metric::Metric;
20use crate::rng::SplitMix64;
21
22/// Highest layer a node can be assigned to.
23pub(crate) const MAX_LEVEL: usize = 16;
24
25/// HNSW index parameters.
26#[derive(Clone, Copy, Debug, PartialEq, Eq)]
27pub struct HnswParams {
28    /// Maximum links per node on upper layers. Layer 0 allows `2 * m`.
29    pub m: usize,
30    /// Candidate list size while building the graph. Higher builds a better
31    /// graph more slowly.
32    pub ef_construction: usize,
33    /// Default candidate list size while searching. Higher improves recall at
34    /// the cost of latency.
35    pub ef_search: usize,
36}
37
38impl Default for HnswParams {
39    fn default() -> Self {
40        Self {
41            m: 16,
42            ef_construction: 200,
43            ef_search: 64,
44        }
45    }
46}
47
48impl HnswParams {
49    pub(crate) fn validate(&self) -> Result<()> {
50        if !(2..=256).contains(&self.m) {
51            return Err(Error::InvalidArgument("m must be between 2 and 256".into()));
52        }
53        if self.ef_construction == 0 || self.ef_search == 0 {
54            return Err(Error::InvalidArgument(
55                "ef values must be greater than zero".into(),
56            ));
57        }
58        Ok(())
59    }
60
61    pub(crate) fn max_links(&self, layer: usize) -> usize {
62        if layer == 0 { self.m * 2 } else { self.m }
63    }
64}
65
66/// Contiguous vector storage: node `i` occupies `data[i * dim..(i + 1) * dim]`.
67pub(crate) struct Vectors {
68    pub(crate) dim: usize,
69    pub(crate) data: Vec<f32>,
70}
71
72impl Vectors {
73    pub(crate) fn new(dim: usize) -> Self {
74        Self {
75            dim,
76            data: Vec::new(),
77        }
78    }
79
80    #[inline]
81    pub(crate) fn get(&self, node: u32) -> &[f32] {
82        let start = node as usize * self.dim;
83        &self.data[start..start + self.dim]
84    }
85
86    pub(crate) fn push(&mut self, vector: &[f32]) {
87        debug_assert_eq!(vector.len(), self.dim);
88        self.data.extend_from_slice(vector);
89    }
90}
91
92#[derive(Clone, Copy, Debug)]
93pub(crate) struct Candidate {
94    pub(crate) dist: f32,
95    pub(crate) id: u32,
96}
97
98impl PartialEq for Candidate {
99    fn eq(&self, other: &Self) -> bool {
100        self.cmp(other) == Ordering::Equal
101    }
102}
103
104impl Eq for Candidate {}
105
106impl PartialOrd for Candidate {
107    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
108        Some(self.cmp(other))
109    }
110}
111
112impl Ord for Candidate {
113    fn cmp(&self, other: &Self) -> Ordering {
114        self.dist
115            .total_cmp(&other.dist)
116            .then(self.id.cmp(&other.id))
117    }
118}
119
120/// Work done by a search, reported by `explain`.
121#[derive(Clone, Copy, Debug, Default)]
122pub(crate) struct Counters {
123    pub(crate) visited: usize,
124    pub(crate) distance_computations: usize,
125}
126
127/// Set of visited nodes: a bitset (one bit per node keeps it cache-resident
128/// on large graphs) that is reset by clearing only the words a search
129/// touched. Reused across searches on the same thread.
130#[derive(Default)]
131pub(crate) struct Visited {
132    bits: Vec<u64>,
133    touched: Vec<u32>,
134    /// Scratch buffer for the unvisited neighbors of the node being expanded.
135    fresh: Vec<u32>,
136}
137
138impl Visited {
139    fn reset(&mut self, len: usize) {
140        for word in self.touched.drain(..) {
141            self.bits[word as usize] = 0;
142        }
143        let words = len.div_ceil(64);
144        if self.bits.len() < words {
145            self.bits.resize(words, 0);
146        }
147    }
148
149    /// Marks `node` as visited and returns whether it was new.
150    #[inline]
151    fn insert(&mut self, node: u32) -> bool {
152        let (word, bit) = (node as usize / 64, 1u64 << (node % 64));
153        let bits = &mut self.bits[word];
154        if *bits & bit != 0 {
155            return false;
156        }
157        if *bits == 0 {
158            self.touched.push(word as u32);
159        }
160        *bits |= bit;
161        true
162    }
163}
164
165thread_local! {
166    static VISITED: RefCell<Visited> = RefCell::new(Visited::default());
167}
168
169/// Read access to neighbor lists.
170pub(crate) trait ReadLinks {
171    fn with_neighbors<R>(&self, node: u32, layer: usize, f: impl FnOnce(&[u32]) -> R) -> R;
172
173    /// Whether `node` is linked on every layer. Only nodes another thread is
174    /// still inserting are not.
175    fn is_ready(&self, _node: u32) -> bool {
176        true
177    }
178}
179
180/// Write access to neighbor lists.
181pub(crate) trait Links: ReadLinks {
182    /// Writes the links chosen for a newly inserted node. Links other inserts
183    /// added to it meanwhile are kept, pruned together to `max`.
184    fn set_neighbors(
185        &mut self,
186        node: u32,
187        layer: usize,
188        neighbors: Vec<u32>,
189        max: usize,
190        vectors: &Vectors,
191        metric: Metric,
192    );
193
194    /// Adds a link `node -> new`, pruning the list back to `max` links.
195    fn add_link(
196        &mut self,
197        node: u32,
198        layer: usize,
199        new: u32,
200        max: usize,
201        vectors: &Vectors,
202        metric: Metric,
203    );
204}
205
206impl ReadLinks for [Vec<Vec<u32>>] {
207    #[inline]
208    fn with_neighbors<R>(&self, node: u32, layer: usize, f: impl FnOnce(&[u32]) -> R) -> R {
209        f(&self[node as usize][layer])
210    }
211}
212
213/// Exclusive access for sequential inserts.
214struct Exclusive<'a>(&'a mut [Vec<Vec<u32>>]);
215
216impl ReadLinks for Exclusive<'_> {
217    fn with_neighbors<R>(&self, node: u32, layer: usize, f: impl FnOnce(&[u32]) -> R) -> R {
218        self.0.with_neighbors(node, layer, f)
219    }
220}
221
222impl Links for Exclusive<'_> {
223    fn set_neighbors(
224        &mut self,
225        node: u32,
226        layer: usize,
227        neighbors: Vec<u32>,
228        _: usize,
229        _: &Vectors,
230        _: Metric,
231    ) {
232        // Sequential inserts: nothing can have linked to `node` yet.
233        self.0[node as usize][layer] = neighbors;
234    }
235
236    fn add_link(
237        &mut self,
238        node: u32,
239        layer: usize,
240        new: u32,
241        max: usize,
242        vectors: &Vectors,
243        metric: Metric,
244    ) {
245        push_and_prune(
246            &mut self.0[node as usize][layer],
247            node,
248            new,
249            max,
250            vectors,
251            metric,
252        );
253    }
254}
255
256/// Shared access for parallel inserts: one lock per node.
257#[derive(Clone, Copy)]
258struct Shared<'a> {
259    links: &'a [Mutex<Vec<Vec<u32>>>],
260    ready: &'a [AtomicBool],
261}
262
263impl Shared<'_> {
264    fn lock(&self, node: u32) -> MutexGuard<'_, Vec<Vec<u32>>> {
265        self.links[node as usize]
266            .lock()
267            .unwrap_or_else(PoisonError::into_inner)
268    }
269}
270
271impl ReadLinks for Shared<'_> {
272    fn with_neighbors<R>(&self, node: u32, layer: usize, f: impl FnOnce(&[u32]) -> R) -> R {
273        f(&self.lock(node)[layer])
274    }
275
276    fn is_ready(&self, node: u32) -> bool {
277        self.ready[node as usize].load(AtomicOrdering::Acquire)
278    }
279}
280
281impl Links for Shared<'_> {
282    fn set_neighbors(
283        &mut self,
284        node: u32,
285        layer: usize,
286        neighbors: Vec<u32>,
287        max: usize,
288        vectors: &Vectors,
289        metric: Metric,
290    ) {
291        let mut links = self.lock(node);
292        let list = &mut links[layer];
293        if list.is_empty() {
294            *list = neighbors;
295            return;
296        }
297        for neighbor in neighbors {
298            if !list.contains(&neighbor) {
299                list.push(neighbor);
300            }
301        }
302        prune(list, node, max, vectors, metric);
303    }
304
305    fn add_link(
306        &mut self,
307        node: u32,
308        layer: usize,
309        new: u32,
310        max: usize,
311        vectors: &Vectors,
312        metric: Metric,
313    ) {
314        push_and_prune(&mut self.lock(node)[layer], node, new, max, vectors, metric);
315    }
316}
317
318pub(crate) struct Graph {
319    pub(crate) params: HnswParams,
320    /// `links[node][layer]` lists the neighbors of `node` on `layer`.
321    pub(crate) links: Vec<Vec<Vec<u32>>>,
322    pub(crate) entry: Option<u32>,
323    pub(crate) max_level: usize,
324    pub(crate) rng: SplitMix64,
325}
326
327impl Graph {
328    pub(crate) fn new(params: HnswParams, seed: u64) -> Self {
329        Self {
330            params,
331            links: Vec::new(),
332            entry: None,
333            max_level: 0,
334            rng: SplitMix64::new(seed),
335        }
336    }
337
338    pub(crate) fn len(&self) -> usize {
339        self.links.len()
340    }
341
342    fn random_level(&mut self) -> usize {
343        let ml = 1.0 / (self.params.m as f64).ln();
344        let level = (-self.rng.next_f64().ln() * ml).floor() as usize;
345        level.min(MAX_LEVEL)
346    }
347
348    /// Links node `id`, which must be the next node and whose vector must
349    /// already be in `vectors`.
350    pub(crate) fn insert(&mut self, id: u32, vectors: &Vectors, metric: Metric) {
351        debug_assert_eq!(id as usize, self.links.len());
352        self.insert_batch(id..id + 1, vectors, metric, 1);
353    }
354
355    /// Links nodes `nodes`, which must directly follow the existing nodes and
356    /// whose vectors must already be in `vectors`, using up to `threads`
357    /// threads. Levels are drawn in node order, so a single-threaded batch
358    /// builds exactly the same graph as inserting the nodes one by one.
359    pub(crate) fn insert_batch(
360        &mut self,
361        nodes: Range<u32>,
362        vectors: &Vectors,
363        metric: Metric,
364        threads: usize,
365    ) {
366        debug_assert_eq!(nodes.start as usize, self.links.len());
367        let levels: Vec<usize> = nodes.clone().map(|_| self.random_level()).collect();
368        for &level in &levels {
369            self.links.push(vec![Vec::new(); level + 1]);
370        }
371        let mut rest = nodes.clone();
372        if self.entry.is_none() {
373            let Some(first) = rest.next() else { return };
374            self.entry = Some(first);
375            self.max_level = levels[0];
376        }
377        let level_of = |node: u32| levels[(node - nodes.start) as usize];
378        let params = self.params;
379
380        // Small batches are not worth the lock overhead.
381        if threads <= 1 || rest.len() < 1000 {
382            VISITED.with_borrow_mut(|visited| {
383                for node in rest {
384                    let level = level_of(node);
385                    let entry = self.entry.expect("entry is set above");
386                    let mut links = Exclusive(&mut self.links);
387                    link_node(
388                        &mut links,
389                        &params,
390                        node,
391                        level,
392                        entry,
393                        self.max_level,
394                        vectors,
395                        metric,
396                        visited,
397                    );
398                    if level > self.max_level {
399                        self.max_level = level;
400                        self.entry = Some(node);
401                    }
402                }
403            });
404            return;
405        }
406
407        let locked: Vec<Mutex<Vec<Vec<u32>>>> = std::mem::take(&mut self.links)
408            .into_iter()
409            .map(Mutex::new)
410            .collect();
411        let top = Mutex::new((self.entry.expect("entry is set above"), self.max_level));
412        let next = AtomicUsize::new(rest.start as usize);
413        let ready: Vec<AtomicBool> = (0..locked.len())
414            .map(|node| AtomicBool::new(node < rest.start as usize))
415            .collect();
416        std::thread::scope(|scope| {
417            for _ in 0..threads {
418                scope.spawn(|| {
419                    let mut links = Shared {
420                        links: &locked,
421                        ready: &ready,
422                    };
423                    let mut visited = Visited::default();
424                    loop {
425                        let node = next.fetch_add(1, AtomicOrdering::Relaxed);
426                        if node >= rest.end as usize {
427                            break;
428                        }
429                        let (node, level) = (node as u32, level_of(node as u32));
430                        let mut top = top.lock().unwrap_or_else(PoisonError::into_inner);
431                        let (entry, max_level) = *top;
432                        if level > max_level {
433                            // Rare: the new node becomes the entry point, so other
434                            // inserts wait until it is fully linked (as in hnswlib).
435                            link_node(
436                                &mut links,
437                                &params,
438                                node,
439                                level,
440                                entry,
441                                max_level,
442                                vectors,
443                                metric,
444                                &mut visited,
445                            );
446                            *top = (node, level);
447                        } else {
448                            drop(top);
449                            link_node(
450                                &mut links,
451                                &params,
452                                node,
453                                level,
454                                entry,
455                                max_level,
456                                vectors,
457                                metric,
458                                &mut visited,
459                            );
460                        }
461                        ready[node as usize].store(true, AtomicOrdering::Release);
462                    }
463                });
464            }
465        });
466        self.links = locked
467            .into_iter()
468            .map(|m| m.into_inner().unwrap_or_else(PoisonError::into_inner))
469            .collect();
470        let (entry, max_level) = top.into_inner().unwrap_or_else(PoisonError::into_inner);
471        self.entry = Some(entry);
472        self.max_level = max_level;
473        self.repair_unreachable(vectors, metric);
474    }
475
476    /// Concurrent pruning can, rarely, remove every incoming link of a node,
477    /// making it invisible to search. Re-links each such node from its
478    /// nearest reachable neighbors on layer 0. Returns how many were found.
479    pub(crate) fn repair_unreachable(&mut self, vectors: &Vectors, metric: Metric) -> usize {
480        let reachable = self.reachable();
481        let orphans: Vec<u32> = (0..self.links.len() as u32)
482            .filter(|&node| !reachable[node as usize])
483            .collect();
484        let Some(entry) = self.entry else { return 0 };
485        let max = self.params.max_links(0);
486        VISITED.with_borrow_mut(|visited| {
487            for &orphan in &orphans {
488                let query = vectors.get(orphan);
489                let found = {
490                    let links: &[Vec<Vec<u32>>] = &self.links;
491                    let mut counters = Counters::default();
492                    let mut ep = Candidate {
493                        dist: metric.distance(query, vectors.get(entry)),
494                        id: entry,
495                    };
496                    for layer in (1..=self.max_level).rev() {
497                        ep = greedy(links, query, ep, layer, vectors, metric, &mut counters);
498                    }
499                    let ef = self.params.ef_construction;
500                    let accept = |n: u32| n != orphan;
501                    search_layer(
502                        links,
503                        query,
504                        &[ep],
505                        ef,
506                        0,
507                        vectors,
508                        metric,
509                        &accept,
510                        &mut counters,
511                        visited,
512                    )
513                };
514                let mut linked = false;
515                for candidate in found.iter().take(self.params.m) {
516                    let list = &mut self.links[candidate.id as usize][0];
517                    push_and_prune(list, candidate.id, orphan, max, vectors, metric);
518                    linked |= list.contains(&orphan);
519                }
520                if !linked {
521                    if let Some(nearest) = found.first() {
522                        // The heuristic rejected every link; keep one anyway.
523                        self.links[nearest.id as usize][0].push(orphan);
524                    }
525                }
526            }
527        });
528        orphans.len()
529    }
530
531    /// Returns up to `k` accepted nodes closest to `query`, nearest first.
532    #[allow(clippy::too_many_arguments)]
533    pub(crate) fn search(
534        &self,
535        query: &[f32],
536        k: usize,
537        ef: usize,
538        vectors: &Vectors,
539        metric: Metric,
540        accept: &dyn Fn(u32) -> bool,
541        counters: &mut Counters,
542    ) -> Vec<Candidate> {
543        let Some(entry) = self.entry else {
544            return Vec::new();
545        };
546        let links: &[Vec<Vec<u32>>] = &self.links;
547        counters.distance_computations += 1;
548        let mut ep = Candidate {
549            dist: metric.distance(query, vectors.get(entry)),
550            id: entry,
551        };
552        for layer in (1..=self.max_level).rev() {
553            ep = greedy(links, query, ep, layer, vectors, metric, counters);
554        }
555        let mut found = VISITED.with_borrow_mut(|visited| {
556            search_layer(
557                links,
558                query,
559                &[ep],
560                ef.max(k),
561                0,
562                vectors,
563                metric,
564                accept,
565                counters,
566                visited,
567            )
568        });
569        found.truncate(k);
570        found
571    }
572
573    /// Number of nodes present on each layer, bottom first.
574    pub(crate) fn nodes_per_layer(&self) -> Vec<usize> {
575        if self.links.is_empty() {
576            return Vec::new();
577        }
578        let mut counts = vec![0; self.max_level + 1];
579        for layers in &self.links {
580            for count in counts.iter_mut().take(layers.len()) {
581                *count += 1;
582            }
583        }
584        counts
585    }
586
587    /// Marks every node reachable from the entry point on layer 0.
588    pub(crate) fn reachable(&self) -> Vec<bool> {
589        let mut seen = vec![false; self.links.len()];
590        let Some(entry) = self.entry else { return seen };
591        let mut stack = vec![entry];
592        seen[entry as usize] = true;
593        while let Some(node) = stack.pop() {
594            for &neighbor in &self.links[node as usize][0] {
595                if !seen[neighbor as usize] {
596                    seen[neighbor as usize] = true;
597                    stack.push(neighbor);
598                }
599            }
600        }
601        seen
602    }
603}
604
605/// Inserts `node` into the graph starting from a snapshot of the entry point.
606/// The node's own links are written before any reverse link, so concurrent
607/// searches never reach a node without neighbors on a fully linked layer.
608#[allow(clippy::too_many_arguments)]
609fn link_node<L: Links>(
610    links: &mut L,
611    params: &HnswParams,
612    node: u32,
613    level: usize,
614    entry: u32,
615    max_level: usize,
616    vectors: &Vectors,
617    metric: Metric,
618    visited: &mut Visited,
619) {
620    let query = vectors.get(node);
621    let mut counters = Counters::default();
622    let mut ep = Candidate {
623        dist: metric.distance(query, vectors.get(entry)),
624        id: entry,
625    };
626    for layer in (level + 1..=max_level).rev() {
627        ep = greedy(&*links, query, ep, layer, vectors, metric, &mut counters);
628    }
629
630    let mut entry_points = vec![ep];
631    for layer in (0..=level.min(max_level)).rev() {
632        let found = search_layer(
633            &*links,
634            query,
635            &entry_points,
636            params.ef_construction,
637            layer,
638            vectors,
639            metric,
640            &|n| n != node,
641            &mut counters,
642            visited,
643        );
644        let max = params.max_links(layer);
645        let neighbors = select_neighbors(&found, max, vectors, metric);
646        links.set_neighbors(node, layer, neighbors.clone(), max, vectors, metric);
647        for neighbor in neighbors {
648            links.add_link(neighbor, layer, node, max, vectors, metric);
649        }
650        entry_points = found;
651    }
652}
653
654/// Hints the CPU to start loading `vector` into cache, one prefetch per
655/// 64-byte line.
656#[inline(always)]
657fn prefetch(vector: &[f32]) {
658    for line in vector.chunks(16) {
659        let ptr = line.as_ptr();
660        #[cfg(target_arch = "aarch64")]
661        // SAFETY: a prefetch is only a hint; it never faults and has no
662        // architectural side effects.
663        unsafe {
664            std::arch::asm!("prfm pldl1keep, [{0}]", in(reg) ptr, options(nostack, preserves_flags, readonly));
665        }
666        #[cfg(target_arch = "x86_64")]
667        // SAFETY: as above; `_mm_prefetch` is available on every x86_64 CPU.
668        unsafe {
669            std::arch::x86_64::_mm_prefetch(ptr.cast::<i8>(), std::arch::x86_64::_MM_HINT_T0);
670        }
671        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
672        let _ = ptr;
673    }
674}
675
676fn greedy<R: ReadLinks + ?Sized>(
677    links: &R,
678    query: &[f32],
679    mut best: Candidate,
680    layer: usize,
681    vectors: &Vectors,
682    metric: Metric,
683    counters: &mut Counters,
684) -> Candidate {
685    loop {
686        let current = best.id;
687        links.with_neighbors(current, layer, |neighbors| {
688            for &neighbor in neighbors {
689                prefetch(vectors.get(neighbor));
690            }
691            for &neighbor in neighbors {
692                counters.distance_computations += 1;
693                let dist = metric.distance(query, vectors.get(neighbor));
694                // A node still being inserted may have empty lower layers;
695                // descending into it would leave the search stranded.
696                if dist < best.dist && links.is_ready(neighbor) {
697                    best = Candidate { dist, id: neighbor };
698                }
699            }
700        });
701        if best.id == current {
702            return best;
703        }
704    }
705}
706
707/// Beam search on one layer. Every reachable node is traversed, but only
708/// nodes passing `accept` enter the result set, so a restrictive filter
709/// widens the search instead of returning fewer results.
710#[allow(clippy::too_many_arguments)]
711fn search_layer<R: ReadLinks + ?Sized>(
712    links: &R,
713    query: &[f32],
714    entry_points: &[Candidate],
715    ef: usize,
716    layer: usize,
717    vectors: &Vectors,
718    metric: Metric,
719    accept: &dyn Fn(u32) -> bool,
720    counters: &mut Counters,
721    visited: &mut Visited,
722) -> Vec<Candidate> {
723    visited.reset(vectors.data.len() / vectors.dim.max(1));
724    let mut candidates = BinaryHeap::new();
725    let mut results: BinaryHeap<Candidate> = BinaryHeap::new();
726
727    for &ep in entry_points {
728        if !visited.insert(ep.id) {
729            continue;
730        }
731        counters.visited += 1;
732        candidates.push(Reverse(ep));
733        if accept(ep.id) {
734            results.push(ep);
735            if results.len() > ef {
736                results.pop();
737            }
738        }
739    }
740
741    while let Some(Reverse(current)) = candidates.pop() {
742        if results.len() >= ef && results.peek().is_some_and(|w| current.dist > w.dist) {
743            break;
744        }
745        // Collect unvisited neighbors first and prefetch their vectors, so
746        // the memory loads overlap instead of stalling one distance at a time.
747        let mut fresh = std::mem::take(&mut visited.fresh);
748        fresh.clear();
749        links.with_neighbors(current.id, layer, |neighbors| {
750            for &neighbor in neighbors {
751                if visited.insert(neighbor) {
752                    fresh.push(neighbor);
753                }
754            }
755        });
756        for &neighbor in &fresh {
757            prefetch(vectors.get(neighbor));
758        }
759        for &neighbor in &fresh {
760            counters.visited += 1;
761            counters.distance_computations += 1;
762            let dist = metric.distance(query, vectors.get(neighbor));
763            let worst = results.peek().map_or(f32::INFINITY, |w| w.dist);
764            if results.len() < ef || dist < worst {
765                let candidate = Candidate { dist, id: neighbor };
766                candidates.push(Reverse(candidate));
767                if accept(neighbor) {
768                    results.push(candidate);
769                    if results.len() > ef {
770                        results.pop();
771                    }
772                }
773            }
774        }
775        visited.fresh = fresh;
776    }
777
778    let mut out = results.into_vec();
779    out.sort();
780    out
781}
782
783/// Neighbor selection heuristic: prefer candidates that are closer to the
784/// base node than to any already selected neighbor, which keeps links spread
785/// across directions. Remaining slots are filled with the nearest pruned
786/// candidates. `candidates` must be sorted nearest first.
787fn select_neighbors(
788    candidates: &[Candidate],
789    max: usize,
790    vectors: &Vectors,
791    metric: Metric,
792) -> Vec<u32> {
793    let mut selected: Vec<Candidate> = Vec::with_capacity(max);
794    let mut pruned = Vec::new();
795    for &candidate in candidates {
796        if selected.len() >= max {
797            break;
798        }
799        let vector = vectors.get(candidate.id);
800        let diverse = selected
801            .iter()
802            .all(|s| metric.distance(vector, vectors.get(s.id)) > candidate.dist);
803        if diverse {
804            selected.push(candidate);
805        } else {
806            pruned.push(candidate);
807        }
808    }
809    for candidate in pruned {
810        if selected.len() >= max {
811            break;
812        }
813        selected.push(candidate);
814    }
815    selected.into_iter().map(|c| c.id).collect()
816}
817
818fn push_and_prune(
819    list: &mut Vec<u32>,
820    node: u32,
821    new: u32,
822    max: usize,
823    vectors: &Vectors,
824    metric: Metric,
825) {
826    list.push(new);
827    prune(list, node, max, vectors, metric);
828}
829
830/// Shrinks `list` back to `max` links with the selection heuristic.
831fn prune(list: &mut Vec<u32>, node: u32, max: usize, vectors: &Vectors, metric: Metric) {
832    if list.len() <= max {
833        return;
834    }
835    let base = vectors.get(node);
836    let mut candidates: Vec<Candidate> = list
837        .iter()
838        .map(|&n| Candidate {
839            dist: metric.distance(base, vectors.get(n)),
840            id: n,
841        })
842        .collect();
843    candidates.sort();
844    *list = select_neighbors(&candidates, max, vectors, metric);
845}
846
847#[cfg(test)]
848mod tests {
849    use super::*;
850
851    fn random_graph(n: usize, dim: usize) -> (Graph, Vectors) {
852        let mut rng = SplitMix64::new(9);
853        let mut vectors = Vectors::new(dim);
854        for _ in 0..n {
855            let v: Vec<f32> = (0..dim).map(|_| rng.next_f64() as f32).collect();
856            vectors.push(&v);
857        }
858        let mut graph = Graph::new(HnswParams::default(), 1);
859        graph.insert_batch(0..n as u32, &vectors, Metric::L2, 1);
860        (graph, vectors)
861    }
862
863    #[test]
864    fn repair_relinks_orphaned_nodes() {
865        let (mut graph, vectors) = random_graph(500, 8);
866        let orphan = (0..500u32).find(|&n| Some(n) != graph.entry).unwrap();
867        for layers in &mut graph.links {
868            for list in layers.iter_mut() {
869                list.retain(|&n| n != orphan);
870            }
871        }
872        assert!(!graph.reachable()[orphan as usize]);
873        assert_eq!(graph.repair_unreachable(&vectors, Metric::L2), 1);
874        assert!(graph.reachable().iter().all(|&r| r));
875        assert_eq!(graph.repair_unreachable(&vectors, Metric::L2), 0);
876    }
877
878    #[test]
879    fn visited_set_resets_between_searches() {
880        let mut visited = Visited::default();
881        visited.reset(200);
882        assert!(visited.insert(3) && visited.insert(130));
883        assert!(!visited.insert(3));
884        visited.reset(200);
885        assert!(visited.insert(3) && visited.insert(130));
886    }
887}