1use std::cmp::Reverse;
9use std::collections::BinaryHeap;
10use std::sync::Arc;
11
12use crate::dsl::IvfRoutingMode;
13use rand::prelude::*;
14use serde::{Deserialize, Serialize};
15
16pub const HNSW_AUTO_THRESHOLD: usize = 4_096;
20
21const PARENT_BEAM_OVERSAMPLE: usize = 4;
25
26const HNSW_M: usize = 32;
27const HNSW_EF_CONSTRUCTION: usize = 200;
28const HNSW_QUERY_OVERSAMPLE: usize = 4;
29const HNSW_MIN_EF_SEARCH: usize = 128;
30
31#[derive(Clone, Copy, Debug)]
32struct GraphCandidate {
33 node: u32,
34 distance: f32,
35}
36
37impl PartialEq for GraphCandidate {
38 fn eq(&self, other: &Self) -> bool {
39 self.node == other.node && self.distance.to_bits() == other.distance.to_bits()
40 }
41}
42
43impl Eq for GraphCandidate {}
44
45impl PartialOrd for GraphCandidate {
46 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
47 Some(self.cmp(other))
48 }
49}
50
51impl Ord for GraphCandidate {
52 fn cmp(&self, other: &Self) -> std::cmp::Ordering {
53 self.distance
54 .total_cmp(&other.distance)
55 .then_with(|| self.node.cmp(&other.node))
56 }
57}
58
59struct VisitedNodes {
60 epochs: Vec<u32>,
61 current: u32,
62}
63
64impl VisitedNodes {
65 fn new(nodes: usize) -> Self {
66 Self {
67 epochs: vec![0; nodes],
68 current: 0,
69 }
70 }
71
72 fn reset(&mut self) {
73 self.current = self.current.wrapping_add(1);
74 if self.current == 0 {
75 self.epochs.fill(0);
76 self.current = 1;
77 }
78 }
79
80 fn ensure_nodes(&mut self, nodes: usize) {
81 if self.epochs.len() < nodes {
82 self.epochs.resize(nodes, 0);
83 }
84 }
85
86 fn insert(&mut self, node: u32) -> bool {
87 let slot = &mut self.epochs[node as usize];
88 if *slot == self.current {
89 false
90 } else {
91 *slot = self.current;
92 true
93 }
94 }
95}
96
97struct HnswQueryScratch {
98 visited: VisitedNodes,
99 candidates: BinaryHeap<Reverse<GraphCandidate>>,
100 best: BinaryHeap<GraphCandidate>,
101 ordered: Vec<GraphCandidate>,
102}
103
104impl HnswQueryScratch {
105 fn new() -> Self {
106 Self {
107 visited: VisitedNodes::new(0),
108 candidates: BinaryHeap::new(),
109 best: BinaryHeap::new(),
110 ordered: Vec::new(),
111 }
112 }
113}
114
115thread_local! {
116 static HNSW_QUERY_SCRATCH: std::cell::RefCell<HnswQueryScratch> =
120 std::cell::RefCell::new(HnswQueryScratch::new());
121}
122
123#[derive(Debug, Clone, Serialize, Deserialize)]
127pub struct HnswRoutingGraph {
128 m: u16,
129 ef_construction: u32,
130 entry_point: u32,
131 max_level: u8,
132 node_levels: Vec<u8>,
133 node_offsets: Vec<u32>,
136 level_offsets: Vec<u32>,
137 neighbors: Vec<u32>,
138}
139
140impl HnswRoutingGraph {
141 pub fn build(node_count: usize, distance: impl Fn(u32, u32) -> f32, seed: u64) -> Self {
142 assert!(node_count > 0 && node_count <= u32::MAX as usize);
143 let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
144 let level_multiplier = 1.0 / (HNSW_M as f64).ln();
145 let node_levels: Vec<u8> = (0..node_count)
146 .map(|_| {
147 let uniform = rng.random::<f64>().clamp(f64::MIN_POSITIVE, 1.0);
148 (-uniform.ln() * level_multiplier).floor().min(31.0) as u8
149 })
150 .collect();
151 let mut insertion_order: Vec<u32> = (0..node_count as u32).collect();
152 insertion_order.shuffle(&mut rng);
153 let mut links: Vec<Vec<Vec<u32>>> = node_levels
154 .iter()
155 .map(|&level| vec![Vec::new(); level as usize + 1])
156 .collect();
157 let mut visited = VisitedNodes::new(node_count);
158 let mut entry_point = insertion_order[0];
159 let mut max_level = node_levels[entry_point as usize];
160
161 for &node in insertion_order.iter().skip(1) {
162 let node_level = node_levels[node as usize];
163 let mut entry = entry_point;
164 let node_distance = |candidate| distance(node, candidate);
165
166 for level in ((node_level as usize + 1)..=max_level as usize).rev() {
167 entry = greedy_search_level(&links, entry, level, &node_distance);
168 }
169
170 for level in (0..=usize::min(node_level as usize, max_level as usize)).rev() {
171 let candidates = search_graph_layer(
172 &links,
173 entry,
174 level,
175 HNSW_EF_CONSTRUCTION,
176 &node_distance,
177 &mut visited,
178 );
179 if let Some(best) = candidates.first() {
180 entry = best.node;
181 }
182 let max_connections = if level == 0 { HNSW_M * 2 } else { HNSW_M };
183 let selected =
184 select_diverse_neighbors(node, candidates, max_connections, &distance);
185 links[node as usize][level] = selected.clone();
186 for neighbor in selected {
187 let adjacency = &mut links[neighbor as usize][level];
188 if !adjacency.contains(&node) {
189 adjacency.push(node);
190 }
191 if adjacency.len() > max_connections {
192 let candidates = adjacency
193 .iter()
194 .copied()
195 .map(|candidate| GraphCandidate {
196 node: candidate,
197 distance: distance(neighbor, candidate),
198 })
199 .collect();
200 *adjacency = select_diverse_neighbors(
201 neighbor,
202 candidates,
203 max_connections,
204 &distance,
205 );
206 }
207 }
208 }
209
210 if node_level > max_level {
211 entry_point = node;
212 max_level = node_level;
213 }
214 }
215
216 Self::compact(
217 HNSW_M,
218 HNSW_EF_CONSTRUCTION,
219 entry_point,
220 max_level,
221 node_levels,
222 links,
223 )
224 }
225
226 fn compact(
227 m: usize,
228 ef_construction: usize,
229 entry_point: u32,
230 max_level: u8,
231 node_levels: Vec<u8>,
232 links: Vec<Vec<Vec<u32>>>,
233 ) -> Self {
234 let mut node_offsets = Vec::with_capacity(links.len() + 1);
235 let level_count: usize = links.iter().map(|levels| levels.len() + 1).sum();
236 let neighbor_count: usize = links
237 .iter()
238 .flat_map(|levels| levels.iter())
239 .map(Vec::len)
240 .sum();
241 let mut level_offsets = Vec::with_capacity(level_count);
242 let mut neighbors = Vec::with_capacity(neighbor_count);
243 for levels in links {
244 node_offsets.push(level_offsets.len() as u32);
245 for mut adjacency in levels {
246 adjacency.sort_unstable();
247 adjacency.dedup();
248 level_offsets.push(neighbors.len() as u32);
249 neighbors.extend(adjacency);
250 }
251 level_offsets.push(neighbors.len() as u32);
252 }
253 node_offsets.push(level_offsets.len() as u32);
254 Self {
255 m: m as u16,
256 ef_construction: ef_construction as u32,
257 entry_point,
258 max_level,
259 node_levels,
260 node_offsets,
261 level_offsets,
262 neighbors,
263 }
264 }
265
266 #[inline]
267 pub fn neighbors(&self, node: u32, level: usize) -> &[u32] {
268 if (self.node_levels[node as usize] as usize) < level {
269 return &[];
270 }
271 let offset_index = self.node_offsets[node as usize] as usize + level;
272 let start = self.level_offsets[offset_index] as usize;
273 let end = self.level_offsets[offset_index + 1] as usize;
274 &self.neighbors[start..end]
275 }
276
277 pub fn search(&self, query_distance: impl Fn(u32) -> f32, take: usize) -> Vec<u32> {
278 let take = take.min(self.node_levels.len());
279 if take == 0 {
280 return Vec::new();
281 }
282 let mut entry = self.entry_point;
283 for level in (1..=self.max_level as usize).rev() {
284 entry = greedy_search_compact(self, entry, level, &query_distance);
285 }
286 let ef_search = take
287 .saturating_mul(HNSW_QUERY_OVERSAMPLE)
288 .max(HNSW_MIN_EF_SEARCH)
289 .min(self.node_levels.len());
290 HNSW_QUERY_SCRATCH.with(|scratch| {
291 let mut scratch = scratch.borrow_mut();
292 search_compact_layer_reusing(self, entry, ef_search, &query_distance, &mut scratch);
293 scratch
294 .ordered
295 .iter()
296 .take(take)
297 .map(|candidate| candidate.node)
298 .collect()
299 })
300 }
301
302 pub fn search_one(&self, query_distance: impl Fn(u32) -> f32) -> u32 {
303 let mut entry = self.entry_point;
304 for level in (1..=self.max_level as usize).rev() {
305 entry = greedy_search_compact(self, entry, level, &query_distance);
306 }
307 let ef_search = HNSW_MIN_EF_SEARCH.min(self.node_levels.len());
308 HNSW_QUERY_SCRATCH.with(|scratch| {
309 let mut scratch = scratch.borrow_mut();
310 search_compact_layer_reusing(self, entry, ef_search, &query_distance, &mut scratch);
311 scratch
312 .ordered
313 .first()
314 .map_or(entry, |candidate| candidate.node)
315 })
316 }
317
318 pub fn validate(&self, expected_nodes: usize) -> bool {
319 if self.m as usize != HNSW_M
320 || self.ef_construction as usize != HNSW_EF_CONSTRUCTION
321 || expected_nodes == 0
322 || self.node_levels.len() != expected_nodes
323 || self.node_offsets.len() != expected_nodes + 1
324 || self.node_offsets.first() != Some(&0)
325 || self.node_offsets.last().copied() != Some(self.level_offsets.len() as u32)
326 || self.node_offsets.windows(2).any(|pair| pair[0] > pair[1])
327 || self
328 .node_offsets
329 .iter()
330 .any(|&offset| offset as usize > self.level_offsets.len())
331 || self.entry_point as usize >= expected_nodes
332 || self.node_levels[self.entry_point as usize] != self.max_level
333 || self.node_levels.iter().copied().max() != Some(self.max_level)
334 || self.level_offsets.last().copied() != Some(self.neighbors.len() as u32)
335 || self.level_offsets.windows(2).any(|pair| pair[0] > pair[1])
336 || self
337 .neighbors
338 .iter()
339 .any(|&node| node as usize >= expected_nodes)
340 {
341 return false;
342 }
343 for node in 0..expected_nodes {
344 let start = self.node_offsets[node] as usize;
345 let end = self.node_offsets[node + 1] as usize;
346 if end.saturating_sub(start) != self.node_levels[node] as usize + 2 {
347 return false;
348 }
349 for level in 0..=self.node_levels[node] as usize {
350 let adjacency = self.neighbors(node as u32, level);
351 let max_connections = if level == 0 { HNSW_M * 2 } else { HNSW_M };
352 if adjacency.len() > max_connections
353 || adjacency.contains(&(node as u32))
354 || adjacency.windows(2).any(|pair| pair[0] >= pair[1])
355 {
356 return false;
357 }
358 }
359 }
360 true
361 }
362
363 pub fn size_bytes(&self) -> usize {
364 self.node_levels.len()
365 + self.node_offsets.len() * size_of::<u32>()
366 + self.level_offsets.len() * size_of::<u32>()
367 + self.neighbors.len() * size_of::<u32>()
368 + 32
369 }
370}
371
372fn greedy_search_level(
373 links: &[Vec<Vec<u32>>],
374 mut current: u32,
375 level: usize,
376 query_distance: &impl Fn(u32) -> f32,
377) -> u32 {
378 let mut current_distance = query_distance(current);
379 loop {
380 let mut changed = false;
381 for &candidate in &links[current as usize][level] {
382 let distance = query_distance(candidate);
383 if distance < current_distance || (distance == current_distance && candidate < current)
384 {
385 current = candidate;
386 current_distance = distance;
387 changed = true;
388 }
389 }
390 if !changed {
391 return current;
392 }
393 }
394}
395
396fn greedy_search_compact(
397 graph: &HnswRoutingGraph,
398 mut current: u32,
399 level: usize,
400 query_distance: &impl Fn(u32) -> f32,
401) -> u32 {
402 let mut current_distance = query_distance(current);
403 loop {
404 let mut changed = false;
405 for &candidate in graph.neighbors(current, level) {
406 let distance = query_distance(candidate);
407 if distance < current_distance || (distance == current_distance && candidate < current)
408 {
409 current = candidate;
410 current_distance = distance;
411 changed = true;
412 }
413 }
414 if !changed {
415 return current;
416 }
417 }
418}
419
420fn search_graph_layer(
421 links: &[Vec<Vec<u32>>],
422 entry: u32,
423 level: usize,
424 ef: usize,
425 query_distance: &impl Fn(u32) -> f32,
426 visited: &mut VisitedNodes,
427) -> Vec<GraphCandidate> {
428 search_layer_impl(entry, ef, query_distance, visited, |node| {
429 &links[node as usize][level]
430 })
431}
432
433fn search_compact_layer_reusing(
434 graph: &HnswRoutingGraph,
435 entry: u32,
436 ef: usize,
437 query_distance: &impl Fn(u32) -> f32,
438 scratch: &mut HnswQueryScratch,
439) {
440 scratch.visited.ensure_nodes(graph.node_levels.len());
441 scratch.visited.reset();
442 scratch.candidates.clear();
443 scratch.best.clear();
444 scratch.ordered.clear();
445 scratch.visited.insert(entry);
446 let first = GraphCandidate {
447 node: entry,
448 distance: query_distance(entry),
449 };
450 scratch.candidates.push(Reverse(first));
451 scratch.best.push(first);
452
453 while let Some(Reverse(current)) = scratch.candidates.pop() {
454 if scratch.best.len() >= ef
455 && scratch
456 .best
457 .peek()
458 .is_some_and(|worst| current.distance > worst.distance)
459 {
460 break;
461 }
462 for &neighbor in graph.neighbors(current.node, 0) {
463 if !scratch.visited.insert(neighbor) {
464 continue;
465 }
466 let candidate = GraphCandidate {
467 node: neighbor,
468 distance: query_distance(neighbor),
469 };
470 if scratch.best.len() < ef
471 || scratch.best.peek().is_some_and(|worst| candidate < *worst)
472 {
473 scratch.candidates.push(Reverse(candidate));
474 scratch.best.push(candidate);
475 if scratch.best.len() > ef {
476 scratch.best.pop();
477 }
478 }
479 }
480 }
481 scratch.ordered.extend(scratch.best.drain());
482 scratch.ordered.sort_unstable();
483}
484
485fn search_layer_impl<'a>(
486 entry: u32,
487 ef: usize,
488 query_distance: &impl Fn(u32) -> f32,
489 visited: &mut VisitedNodes,
490 neighbors: impl Fn(u32) -> &'a [u32],
491) -> Vec<GraphCandidate> {
492 visited.reset();
493 visited.insert(entry);
494 let first = GraphCandidate {
495 node: entry,
496 distance: query_distance(entry),
497 };
498 let mut candidates = BinaryHeap::new();
499 let mut best = BinaryHeap::new();
500 candidates.push(Reverse(first));
501 best.push(first);
502
503 while let Some(Reverse(current)) = candidates.pop() {
504 if best.len() >= ef
505 && best
506 .peek()
507 .is_some_and(|worst| current.distance > worst.distance)
508 {
509 break;
510 }
511 for &neighbor in neighbors(current.node) {
512 if !visited.insert(neighbor) {
513 continue;
514 }
515 let candidate = GraphCandidate {
516 node: neighbor,
517 distance: query_distance(neighbor),
518 };
519 if best.len() < ef || best.peek().is_some_and(|worst| candidate < *worst) {
520 candidates.push(Reverse(candidate));
521 best.push(candidate);
522 if best.len() > ef {
523 best.pop();
524 }
525 }
526 }
527 }
528 best.into_sorted_vec()
529}
530
531fn select_diverse_neighbors(
532 query_node: u32,
533 mut candidates: Vec<GraphCandidate>,
534 limit: usize,
535 distance: &impl Fn(u32, u32) -> f32,
536) -> Vec<u32> {
537 candidates.sort_unstable();
538 candidates.dedup_by_key(|candidate| candidate.node);
539 let mut selected = Vec::with_capacity(limit);
540 let mut deferred = Vec::new();
541 for candidate in candidates {
542 if candidate.node == query_node {
543 continue;
544 }
545 if selected
546 .iter()
547 .all(|&neighbor| distance(candidate.node, neighbor) > candidate.distance)
548 {
549 selected.push(candidate.node);
550 if selected.len() == limit {
551 return selected;
552 }
553 } else {
554 deferred.push(candidate.node);
555 }
556 }
557 for candidate in deferred {
558 if selected.len() == limit {
559 break;
560 }
561 selected.push(candidate);
562 }
563 selected
564}
565
566#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
570pub struct IvfRoutingTopology {
571 child_offsets: Vec<u32>,
572 leaf_ids: Vec<u32>,
573}
574
575impl IvfRoutingTopology {
576 pub fn from_children(children: &[Vec<u32>]) -> Self {
577 let mut child_offsets = Vec::with_capacity(children.len() + 1);
578 let mut leaf_ids = Vec::new();
579 child_offsets.push(0);
580 for child_list in children {
581 leaf_ids.extend_from_slice(child_list);
582 child_offsets.push(leaf_ids.len() as u32);
583 }
584 Self {
585 child_offsets,
586 leaf_ids,
587 }
588 }
589
590 pub fn parent_count(&self) -> usize {
591 self.child_offsets.len().saturating_sub(1)
592 }
593
594 pub fn children(&self, parent: usize) -> &[u32] {
595 let start = self.child_offsets[parent] as usize;
596 let end = self.child_offsets[parent + 1] as usize;
597 &self.leaf_ids[start..end]
598 }
599
600 pub fn validate(&self, num_leaves: usize) -> bool {
601 if self.parent_count() == 0 {
602 return self.child_offsets.is_empty() && self.leaf_ids.is_empty();
603 }
604 self.child_offsets.first() == Some(&0)
605 && self.child_offsets.last().copied() == Some(self.leaf_ids.len() as u32)
606 && self.child_offsets.windows(2).all(|pair| pair[0] <= pair[1])
607 && self.leaf_ids.len() == num_leaves
608 && self.leaf_ids.iter().all(|&leaf| leaf < num_leaves as u32)
609 && {
610 let mut leaves = self.leaf_ids.clone();
611 leaves.sort_unstable();
612 leaves.iter().copied().eq(0..num_leaves as u32)
613 }
614 }
615}
616
617pub fn routing_parent_count(num_leaves: usize) -> usize {
618 if num_leaves <= 1 {
619 return num_leaves;
620 }
621 ((num_leaves as f64).sqrt().ceil() as usize)
622 .clamp(2, 4_096)
623 .min(num_leaves)
624}
625
626pub fn allocate_child_clusters(group_sizes: &[usize], total_clusters: usize) -> Vec<usize> {
629 let mut allocated: Vec<usize> = group_sizes
630 .iter()
631 .map(|&size| usize::from(size > 0))
632 .collect();
633 let mut remaining = total_clusters.saturating_sub(allocated.iter().sum());
634 let total_points: usize = group_sizes.iter().sum();
635 if remaining == 0 || total_points == 0 {
636 return allocated;
637 }
638 for (allocation, &size) in allocated.iter_mut().zip(group_sizes) {
639 let capacity = size.saturating_sub(*allocation);
640 let share = remaining
641 .saturating_mul(size)
642 .checked_div(total_points)
643 .unwrap_or(0)
644 .min(capacity);
645 *allocation += share;
646 }
647 remaining = total_clusters.saturating_sub(allocated.iter().sum());
648 while remaining > 0 {
649 let Some((index, _)) = group_sizes
650 .iter()
651 .enumerate()
652 .filter(|(index, size)| allocated[*index] < **size)
653 .max_by_key(|(index, size)| (**size, std::cmp::Reverse(allocated[*index])))
654 else {
655 break;
656 };
657 allocated[index] += 1;
658 remaining -= 1;
659 }
660 allocated
661}
662
663#[derive(Debug, Clone, PartialEq, Eq)]
665pub struct IvfProbePlan {
666 pub quantizer_version: u64,
667 pub request_fingerprint: u64,
670 pub cluster_ids: Arc<[u32]>,
671}
672
673impl IvfProbePlan {
674 pub fn new(quantizer_version: u64, request_fingerprint: u64, cluster_ids: Vec<u32>) -> Self {
675 Self {
676 quantizer_version,
677 request_fingerprint,
678 cluster_ids: cluster_ids.into(),
679 }
680 }
681}
682
683fn fingerprint_words(
684 mode: IvfRoutingMode,
685 nprobe: usize,
686 words: impl IntoIterator<Item = u64>,
687) -> u64 {
688 let mut hash = 0xcbf2_9ce4_8422_2325u64;
691 let mode_tag = match mode {
692 IvfRoutingMode::Auto => 0u64,
693 IvfRoutingMode::Flat => 1,
694 IvfRoutingMode::TwoLevel => 2,
695 IvfRoutingMode::Hnsw => 3,
696 };
697 for word in std::iter::once(mode_tag)
698 .chain(std::iter::once(nprobe as u64))
699 .chain(words)
700 {
701 hash ^= word;
702 hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
703 }
704 hash ^= hash >> 33;
705 hash = hash.wrapping_mul(0xff51_afd7_ed55_8ccd);
706 hash ^ (hash >> 33)
707}
708
709pub fn float_probe_fingerprint(query: &[f32], nprobe: usize, mode: IvfRoutingMode) -> u64 {
710 fingerprint_words(
711 mode,
712 nprobe,
713 query.iter().map(|value| value.to_bits() as u64),
714 )
715}
716
717pub fn binary_probe_fingerprint(query: &[u8], nprobe: usize, mode: IvfRoutingMode) -> u64 {
718 fingerprint_words(mode, nprobe, query.iter().map(|&value| value as u64))
719}
720
721#[inline]
722pub fn effective_routing_mode(mode: IvfRoutingMode, num_leaves: usize) -> IvfRoutingMode {
723 match mode {
724 IvfRoutingMode::Auto if num_leaves >= HNSW_AUTO_THRESHOLD => IvfRoutingMode::Hnsw,
725 IvfRoutingMode::Auto => IvfRoutingMode::Flat,
726 explicit => explicit,
727 }
728}
729
730pub fn parent_probe_count(nprobe: usize, num_leaves: usize, num_parents: usize) -> usize {
732 if num_parents == 0 || num_leaves == 0 {
733 return 0;
734 }
735 let leaves_per_parent = num_leaves.div_ceil(num_parents).max(1);
736 nprobe
737 .saturating_mul(PARENT_BEAM_OVERSAMPLE)
738 .div_ceil(leaves_per_parent)
739 .clamp(1, num_parents)
740}
741
742pub fn select_best<const HIGHER_IS_BETTER: bool>(scores: &[f32], take: usize) -> Vec<u32> {
745 let take = take.min(scores.len());
746 if take == 0 {
747 return Vec::new();
748 }
749 let mut order: Vec<u32> = (0..scores.len() as u32).collect();
750 let compare = |left: &u32, right: &u32| {
751 let left_score = scores[*left as usize];
752 let right_score = scores[*right as usize];
753 let score_order = if HIGHER_IS_BETTER {
754 right_score.total_cmp(&left_score)
755 } else {
756 left_score.total_cmp(&right_score)
757 };
758 score_order.then_with(|| left.cmp(right))
759 };
760 if take < order.len() {
761 order.select_nth_unstable_by(take, compare);
762 order.truncate(take);
763 }
764 order.sort_unstable_by(compare);
765 order
766}
767
768pub fn select_best_candidates<const HIGHER_IS_BETTER: bool>(
772 candidates: &mut Vec<(u32, f32)>,
773 take: usize,
774) -> Vec<u32> {
775 let take = take.min(candidates.len());
776 if take == 0 {
777 return Vec::new();
778 }
779 let compare = |left: &(u32, f32), right: &(u32, f32)| {
780 let score_order = if HIGHER_IS_BETTER {
781 right.1.total_cmp(&left.1)
782 } else {
783 left.1.total_cmp(&right.1)
784 };
785 score_order.then_with(|| left.0.cmp(&right.0))
786 };
787 if take < candidates.len() {
788 candidates.select_nth_unstable_by(take, compare);
789 candidates.truncate(take);
790 }
791 candidates.sort_unstable_by(compare);
792 candidates
793 .iter()
794 .map(|(cluster_id, _)| *cluster_id)
795 .collect()
796}
797
798#[cfg(test)]
799mod tests {
800 use super::*;
801
802 #[test]
803 fn deterministic_selection_supports_both_metric_directions() {
804 let scores = [0.5, 0.9, 0.1, 0.9];
805 assert_eq!(select_best::<true>(&scores, 2), vec![1, 3]);
806 assert_eq!(select_best::<false>(&scores, 2), vec![2, 0]);
807 }
808
809 #[test]
810 fn two_level_beam_is_oversubscribed_but_bounded() {
811 assert_eq!(parent_probe_count(32, 65_536, 256), 1);
812 assert_eq!(parent_probe_count(256, 65_536, 256), 4);
813 assert_eq!(parent_probe_count(65_536, 65_536, 256), 256);
814 }
815
816 #[test]
817 fn compact_hnsw_routes_without_copying_points() {
818 let points: Vec<[f32; 2]> = (0..512)
819 .map(|index| {
820 let angle = index as f32 * std::f32::consts::TAU / 512.0;
821 [angle.cos(), angle.sin()]
822 })
823 .collect();
824 let distance = |left: u32, right: u32| {
825 let [lx, ly] = points[left as usize];
826 let [rx, ry] = points[right as usize];
827 (lx - rx).powi(2) + (ly - ry).powi(2)
828 };
829 let graph = HnswRoutingGraph::build(points.len(), distance, 42);
830 assert!(graph.validate(points.len()));
831 assert!(graph.size_bytes() < points.len() * 512);
832
833 let query = [0.37f32, -0.91];
834 let routed = graph.search(
835 |node| {
836 let [x, y] = points[node as usize];
837 (x - query[0]).powi(2) + (y - query[1]).powi(2)
838 },
839 10,
840 );
841 let mut exact: Vec<u32> = (0..points.len() as u32).collect();
842 exact.sort_unstable_by(|&left, &right| {
843 let score = |node: u32| {
844 let [x, y] = points[node as usize];
845 (x - query[0]).powi(2) + (y - query[1]).powi(2)
846 };
847 score(left)
848 .total_cmp(&score(right))
849 .then_with(|| left.cmp(&right))
850 });
851 assert_eq!(routed, exact[..10]);
852
853 let bytes = bincode::serde::encode_to_vec(&graph, bincode::config::standard()).unwrap();
854 let (decoded, consumed): (HnswRoutingGraph, usize) =
855 bincode::serde::decode_from_slice(&bytes, bincode::config::standard()).unwrap();
856 assert_eq!(consumed, bytes.len());
857 assert!(decoded.validate(points.len()));
858
859 let mut corrupted = decoded;
860 corrupted.node_offsets[1] = u32::MAX;
861 assert!(!corrupted.validate(points.len()));
862 }
863}