use crate::{Document, Result, SearchResult};
use parking_lot::{Mutex, RwLock};
use rand::rngs::StdRng;
use rand::{RngExt, SeedableRng};
use rayon::prelude::*;
use std::cmp::{max, Reverse};
use std::collections::{BinaryHeap, HashSet};
#[inline(always)]
unsafe fn prefetch_read(ptr: *const u8) {
#[cfg(target_arch = "x86_64")]
{
std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(ptr as *const i8);
}
#[cfg(target_arch = "aarch64")]
{
std::arch::asm!("prfm pldl1keep, [{ptr}]", ptr = in(reg) ptr, options(nostack, preserves_flags));
}
#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
{
let _ = ptr;
}
}
#[inline(always)]
unsafe fn prefetch_embedding(ptr: *const u8, cache_lines: usize) {
for i in 0..cache_lines {
prefetch_read(ptr.add(i * 64));
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
struct OrderedFloat(f32);
impl Eq for OrderedFloat {}
impl PartialOrd for OrderedFloat {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for OrderedFloat {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.0.total_cmp(&other.0)
}
}
struct BitsetVisited {
bits: Vec<u64>,
}
impl BitsetVisited {
fn new(n: usize) -> Self {
Self {
bits: vec![0u64; (n + 63) / 64],
}
}
#[inline(always)]
fn is_visited(&self, node: usize) -> bool {
debug_assert!(
(node >> 6) < self.bits.len(),
"BitsetVisited::is_visited: node {} out of bounds (capacity {})",
node,
self.bits.len() * 64
);
let word = unsafe { *self.bits.get_unchecked(node >> 6) };
word & (1u64 << (node & 63)) != 0
}
#[inline(always)]
fn mark_visited(&mut self, node: usize) {
debug_assert!(
(node >> 6) < self.bits.len(),
"BitsetVisited::mark_visited: node {} out of bounds (capacity {})",
node,
self.bits.len() * 64
);
unsafe {
*self.bits.get_unchecked_mut(node >> 6) |= 1u64 << (node & 63);
}
}
#[inline]
fn clear(&mut self) {
self.bits.fill(0);
}
}
struct QueryPrep<'a> {
norm: f32,
rabitq: Option<&'a crate::vector::rabitq::PreparedQuery>,
}
struct SearchContext {
visited: BitsetVisited,
capacity: usize,
candidates: BinaryHeap<Reverse<(OrderedFloat, usize)>>,
best: BinaryHeap<(OrderedFloat, usize)>,
distance_calls: u64,
}
impl SearchContext {
fn new(n: usize) -> Self {
Self {
visited: BitsetVisited::new(n),
capacity: n,
candidates: BinaryHeap::with_capacity(256),
best: BinaryHeap::with_capacity(256),
distance_calls: 0,
}
}
#[inline]
fn reset(&mut self) {
self.visited.clear();
self.candidates.clear();
self.best.clear();
}
#[inline(always)]
fn is_visited(&self, node: usize) -> bool {
self.visited.is_visited(node)
}
#[inline(always)]
fn mark_visited(&mut self, node: usize) {
self.visited.mark_visited(node);
}
}
pub struct Searcher<'a> {
index: &'a HNSWIndex,
ctx: SearchContext,
}
impl Searcher<'_> {
pub fn search(&mut self, query: &[f32], k: usize) -> Result<Vec<SearchResult>> {
self.index.search_inner(query, k, &mut self.ctx)
}
pub fn distance_calls(&self) -> u64 {
self.ctx.distance_calls
}
pub fn reset_stats(&mut self) {
self.ctx.distance_calls = 0;
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum BuildStrategy {
#[default]
Parallel,
Sequential,
}
#[derive(Debug, Clone)]
pub struct HNSWConfig {
pub metric: DistanceMetric,
pub m: usize,
pub m0: usize,
pub ef_construction: usize,
pub ef_search: usize,
pub ml: f32,
pub use_heuristic: bool,
pub extend_candidates: bool,
pub keep_pruned_connections: bool,
pub build_strategy: BuildStrategy,
pub seed: Option<u64>,
pub storage: Storage,
pub rerank_candidates: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub enum DistanceMetric {
#[default]
Cosine,
L2,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub enum Storage {
#[default]
F32,
SQ8,
RaBitQ,
}
impl Default for HNSWConfig {
fn default() -> Self {
let m = 32; Self {
metric: DistanceMetric::default(),
m,
m0: m * 2,
ef_construction: 100,
ef_search: 100,
ml: 1.0 / (m as f32).ln(),
use_heuristic: true,
extend_candidates: false,
keep_pruned_connections: true,
build_strategy: BuildStrategy::default(),
seed: None,
storage: Storage::default(),
rerank_candidates: 100,
}
}
}
impl HNSWConfig {
pub fn with_simple_selection(mut self) -> Self {
self.use_heuristic = false;
self
}
pub fn with_extended_candidates(mut self) -> Self {
self.extend_candidates = true;
self
}
pub fn with_ef_search(mut self, ef: usize) -> Self {
self.ef_search = ef;
self
}
pub fn with_build_strategy(mut self, strategy: BuildStrategy) -> Self {
self.build_strategy = strategy;
self
}
pub fn with_seed(mut self, seed: u64) -> Self {
self.seed = Some(seed);
self
}
pub fn with_ef_construction(mut self, ef: usize) -> Self {
self.ef_construction = ef;
self
}
pub fn with_m(mut self, m: usize) -> Self {
self.m = m;
self.m0 = m * 2;
self.ml = 1.0 / (m as f32).ln();
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MemoryBreakdown {
pub embeddings: usize,
pub norms: usize,
pub layer0_links: usize,
pub upper_layer_links: usize,
pub payload: usize,
}
impl MemoryBreakdown {
pub fn total(&self) -> usize {
self.embeddings + self.norms + self.layer0_links + self.upper_layer_links + self.payload
}
}
#[inline(always)]
const fn node_hdr_len(m0: usize) -> usize {
(2 + m0 + 1).div_ceil(4) * 4
}
#[inline(always)]
const fn rabitq_bit_words(dim: usize) -> usize {
dim.div_ceil(8).div_ceil(4)
}
#[inline(always)]
const fn vec_words(storage: Storage, dim: usize) -> usize {
match storage {
Storage::F32 => dim,
Storage::SQ8 => dim.div_ceil(4),
Storage::RaBitQ => 2 + rabitq_bit_words(dim),
}
}
#[inline(always)]
const fn node_stride(m0: usize, dim: usize, storage: Storage) -> usize {
node_hdr_len(m0) + vec_words(storage, dim)
}
pub struct HNSWIndex {
embedding_dim: usize,
config: HNSWConfig,
nodes: Vec<u32>,
stride: usize,
hdr: usize,
q_min: Vec<f32>,
q_scale: Vec<f32>,
full: Vec<f32>,
rabitq: Option<crate::vector::rabitq::RaBitQuantizer>,
connections: Vec<Vec<Vec<u32>>>,
ids: Vec<String>,
contents: Vec<String>,
metadata: Vec<Option<serde_json::Value>>,
entry_point: Option<usize>,
max_layer: usize,
level_rng: StdRng,
}
#[allow(clippy::too_many_arguments)]
fn select_neighbors_core(
candidates: &[(f32, usize)],
m: usize,
use_heuristic: bool,
extend_candidates: bool,
keep_pruned: bool,
neighbours_of: impl Fn(usize) -> Vec<usize>,
dist_to_query: impl Fn(usize) -> f32,
dist_between: impl Fn(usize, usize) -> f32,
) -> Vec<(f32, usize)> {
if !use_heuristic {
return candidates.iter().take(m).copied().collect();
}
let mut pool: Vec<(f32, usize)> = candidates.to_vec();
if extend_candidates {
let mut seen: HashSet<usize> = candidates.iter().map(|&(_, id)| id).collect();
let mut extra: Vec<(f32, usize)> = Vec::new();
for &(_, cid) in candidates {
for n in neighbours_of(cid) {
if seen.insert(n) {
extra.push((dist_to_query(n), n));
}
}
}
pool.extend(extra);
pool.sort_by(|a, b| a.0.total_cmp(&b.0));
}
let mut selected: Vec<(f32, usize)> = Vec::with_capacity(m);
let mut pruned: Vec<(f32, usize)> = Vec::new();
for &(dq, cid) in &pool {
if selected.len() >= m {
break;
}
if selected.iter().all(|&(_, sid)| dist_between(cid, sid) >= dq) {
selected.push((dq, cid));
} else {
pruned.push((dq, cid));
}
}
if keep_pruned && selected.len() < m {
for c in pruned {
if selected.len() >= m {
break;
}
selected.push(c);
}
}
selected
}
impl HNSWIndex {
pub fn new(embedding_dim: usize, config: HNSWConfig) -> Self {
let level_rng = StdRng::seed_from_u64(config.seed.unwrap_or_else(rand::random));
Self {
level_rng,
embedding_dim,
stride: node_stride(config.m0, embedding_dim, config.storage),
hdr: node_hdr_len(config.m0),
config,
nodes: Vec::new(),
connections: Vec::new(),
q_min: Vec::new(),
q_scale: Vec::new(),
full: Vec::new(),
rabitq: None,
ids: Vec::new(),
contents: Vec::new(),
metadata: Vec::new(),
entry_point: None,
max_layer: 0,
}
}
fn fit_codebook(&mut self, embeddings: &[Vec<f32>]) {
if embeddings.is_empty() {
return;
}
match self.config.storage {
Storage::F32 => {}
Storage::SQ8 => {
let dim = self.embedding_dim;
let mut lo = vec![f32::INFINITY; dim];
let mut hi = vec![f32::NEG_INFINITY; dim];
for v in embeddings {
for d in 0..dim {
lo[d] = lo[d].min(v[d]);
hi[d] = hi[d].max(v[d]);
}
}
self.q_scale = (0..dim).map(|d| (hi[d] - lo[d]) / 255.0).collect();
self.q_min = lo;
}
Storage::RaBitQ => {
if self.config.metric == DistanceMetric::Cosine {
let normalized: Vec<Vec<f32>> = embeddings
.iter()
.map(|v| {
let mut n = v.clone();
crate::vector::ops::normalize(&mut n);
n
})
.collect();
self.rabitq = Some(crate::vector::rabitq::RaBitQuantizer::fit(&normalized));
} else {
self.rabitq = Some(crate::vector::rabitq::RaBitQuantizer::fit(embeddings));
}
}
}
}
fn rabitq_cosine_input<'a>(&self, v: &'a [f32]) -> std::borrow::Cow<'a, [f32]> {
if self.config.metric == DistanceMetric::Cosine {
let mut owned = v.to_vec();
crate::vector::ops::normalize(&mut owned);
std::borrow::Cow::Owned(owned)
} else {
std::borrow::Cow::Borrowed(v)
}
}
fn is_trained(&self) -> bool {
match self.config.storage {
Storage::F32 => true,
Storage::SQ8 => !self.q_scale.is_empty(),
Storage::RaBitQ => self.rabitq.is_some(),
}
}
pub fn train(&mut self, sample: &[Vec<f32>]) -> Result<()> {
if self.config.storage == Storage::F32 {
return Ok(());
}
if !self.is_empty() {
return Err(crate::RagError::InvalidInput(
"train() must be called before any add()/add_embedding() — retraining a \
non-empty index would desynchronize already-encoded vectors from the new \
codebook"
.into(),
));
}
if sample.is_empty() {
return Err(crate::RagError::InvalidInput(
"train() requires a non-empty sample to fit a codebook".into(),
));
}
for v in sample {
if v.len() != self.embedding_dim {
return Err(crate::RagError::DimensionMismatch {
expected: self.embedding_dim,
actual: v.len(),
});
}
}
self.fit_codebook(sample);
Ok(())
}
pub fn with_defaults(embedding_dim: usize) -> Self {
Self::new(embedding_dim, HNSWConfig::default())
}
pub fn set_ef_search(&mut self, ef: usize) {
self.config.ef_search = ef;
}
pub fn ef_search(&self) -> usize {
self.config.ef_search
}
pub fn set_rerank_candidates(&mut self, n: usize) -> Result<()> {
if n > 0 && self.config.storage != Storage::F32 && self.full.is_empty() && !self.is_empty()
{
return Err(crate::RagError::FullPrecisionDropped);
}
self.config.rerank_candidates = n;
Ok(())
}
pub fn rerank_candidates(&self) -> usize {
self.config.rerank_candidates
}
pub fn build(embeddings: Vec<Vec<f32>>, config: HNSWConfig) -> Self {
if embeddings.is_empty() {
return Self::new(0, config);
}
let expected_dim = embeddings[0].len();
for (i, embedding) in embeddings.iter().enumerate() {
assert!(
embedding.len() == expected_dim,
"All embeddings must have the same dimension: expected {}, got {} at index {}",
expected_dim,
embedding.len(),
i
);
}
match config.build_strategy {
BuildStrategy::Sequential => Self::build_sequential(embeddings, config),
BuildStrategy::Parallel => Self::build_parallel(embeddings, config),
}
}
fn build_sequential(embeddings: Vec<Vec<f32>>, config: HNSWConfig) -> Self {
let embedding_dim = embeddings[0].len();
let n = embeddings.len();
let seed = config.seed.unwrap_or_else(rand::random);
let mut rng = StdRng::seed_from_u64(seed);
let ml = config.ml;
let levels: Vec<usize> = (0..n)
.map(|_| {
let r: f32 = rng.random::<f32>().max(f32::EPSILON);
(-r.ln() * ml).floor() as usize
})
.collect();
let _max_level = *levels.iter().max().unwrap_or(&0);
let mut sorted_indices: Vec<usize> = (0..n).collect();
sorted_indices.sort_by(|&a, &b| levels[b].cmp(&levels[a]));
let mut index = Self::new(embedding_dim, config);
let drop_full_after_build =
index.config.storage != Storage::F32 && index.config.rerank_candidates == 0;
if drop_full_after_build {
index.config.rerank_candidates = 1; }
index
.nodes
.reserve(n * node_stride(index.config.m0, embedding_dim, index.config.storage));
index.fit_codebook(&embeddings);
index.connections.reserve(n);
index.ids.reserve(n);
index.contents.reserve(n);
index.metadata.reserve(n);
for &i in &sorted_indices {
let level = levels[i];
let node_id = index.len();
let mut node_connections: Vec<Vec<u32>> = Vec::with_capacity(level + 1);
for _ in 0..=level {
node_connections.push(Vec::new());
}
index.push_node(&embeddings[i]);
index.connections.push(node_connections);
index.ids.push(i.to_string());
index.contents.push(String::new());
index.metadata.push(None);
if index.entry_point.is_none() {
index.entry_point = Some(node_id);
index.max_layer = level;
continue;
}
index.insert_node(node_id, level);
if level > index.max_layer {
index.max_layer = level;
index.entry_point = Some(node_id);
}
}
if drop_full_after_build {
index.config.rerank_candidates = 0;
index.full = Vec::new();
}
index.shrink_to_fit();
index
}
#[inline]
pub fn len(&self) -> usize {
self.ids.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.ids.is_empty()
}
#[inline(always)]
fn get_embedding(&self, node_id: usize) -> &[f32] {
match self.config.storage {
Storage::F32 => {
let start = node_id * self.stride + self.hdr;
bytemuck::cast_slice(&self.nodes[start..start + self.embedding_dim])
}
Storage::SQ8 | Storage::RaBitQ => {
debug_assert!(
!self.full.is_empty(),
"full-precision vectors were dropped (rerank_candidates = 0); \
no exact embedding exists to return"
);
let start = node_id * self.embedding_dim;
&self.full[start..start + self.embedding_dim]
}
}
}
#[inline(always)]
fn get_codes(&self, node_id: usize) -> &[u8] {
let start = node_id * self.stride + self.hdr;
let words = vec_words(Storage::SQ8, self.embedding_dim);
let bytes: &[u8] = bytemuck::cast_slice(&self.nodes[start..start + words]);
&bytes[..self.embedding_dim]
}
#[inline(always)]
fn get_rabitq_code(&self, node_id: usize) -> (f32, f32, &[u8]) {
let base = node_id * self.stride + self.hdr;
let dtc_sq = f32::from_bits(self.nodes[base]);
let est_factor = f32::from_bits(self.nodes[base + 1]);
let bit_words = rabitq_bit_words(self.embedding_dim);
let bytes: &[u8] = bytemuck::cast_slice(&self.nodes[base + 2..base + 2 + bit_words]);
let n_bytes = self.embedding_dim.div_ceil(8);
(dtc_sq, est_factor, &bytes[..n_bytes])
}
#[inline(always)]
fn get_norm(&self, node_id: usize) -> f32 {
f32::from_bits(self.nodes[node_id * self.stride + 1])
}
#[inline(always)]
fn get_neighbors_l0(&self, node_id: usize) -> &[u32] {
let base = node_id * self.stride;
let count = self.nodes[base] as usize;
&self.nodes[base + 2..base + 2 + count]
}
#[inline]
fn l0_contains(&self, node_id: usize, neighbor: u32) -> bool {
self.get_neighbors_l0(node_id).contains(&neighbor)
}
#[inline]
fn l0_push(&mut self, node_id: usize, neighbor: u32) {
let base = node_id * self.stride;
let count = self.nodes[base] as usize;
debug_assert!(
count < self.config.m0 + 1,
"layer-0 overflow: prune before pushing again"
);
self.nodes[base + 2 + count] = neighbor;
self.nodes[base] = (count + 1) as u32;
}
#[inline]
fn l0_replace(&mut self, node_id: usize, neighbors: &[u32]) {
let base = node_id * self.stride;
let count = neighbors.len().min(self.config.m0);
self.nodes[base + 2..base + 2 + count].copy_from_slice(&neighbors[..count]);
self.nodes[base] = count as u32;
}
fn push_node(&mut self, embedding: &[f32]) {
debug_assert_eq!(embedding.len(), self.embedding_dim);
let base = self.nodes.len();
self.nodes.resize(base + self.stride, 0);
self.nodes[base + 1] = crate::vector::simd::norm_simd(embedding).to_bits();
let v = base + self.hdr;
match self.config.storage {
Storage::F32 => {
self.nodes[v..v + self.embedding_dim]
.copy_from_slice(bytemuck::cast_slice(embedding));
}
Storage::SQ8 => {
let words = vec_words(Storage::SQ8, self.embedding_dim);
let bytes: &mut [u8] = bytemuck::cast_slice_mut(&mut self.nodes[v..v + words]);
for (d, &x) in embedding.iter().enumerate() {
let s = self.q_scale[d];
bytes[d] = if s <= 0.0 {
0
} else {
(((x - self.q_min[d]) / s).round().clamp(0.0, 255.0)) as u8
};
}
if self.config.rerank_candidates > 0 {
self.full.extend_from_slice(embedding);
}
}
Storage::RaBitQ => {
let rq = self
.rabitq
.as_ref()
.expect("RaBitQ storage requires fit_codebook to run before push_node");
let encode_input = self.rabitq_cosine_input(embedding);
let code = rq.encode(&encode_input);
self.nodes[v] = code.dtc_sq.to_bits();
self.nodes[v + 1] = code.est_factor.to_bits();
let bit_words = rabitq_bit_words(self.embedding_dim);
let bytes: &mut [u8] =
bytemuck::cast_slice_mut(&mut self.nodes[v + 2..v + 2 + bit_words]);
bytes[..code.bits.len()].copy_from_slice(&code.bits);
if self.config.rerank_candidates > 0 {
self.full.extend_from_slice(embedding);
}
}
}
}
fn migrate_l0_into_arena(&mut self) {
for i in 0..self.len() {
if !self.connections[i].is_empty() {
let neighbors = std::mem::take(&mut self.connections[i][0]);
self.l0_replace(i, &neighbors);
}
}
}
pub fn add(&mut self, document: Document) -> Result<()> {
if document.embedding.len() != self.embedding_dim {
return Err(crate::RagError::DimensionMismatch {
expected: self.embedding_dim,
actual: document.embedding.len(),
});
}
if !self.is_trained() {
return Err(crate::RagError::NotTrained(format!(
"Storage::{:?} requires a fitted codebook before add() — call \
`index.train(&sample)` first, or build the index via `HNSWIndex::build`/ \
`build_parallel`, which trains internally from the full corpus",
self.config.storage
)));
}
if document.embedding.iter().any(|v| !v.is_finite()) {
return Err(crate::RagError::InvalidInput(
"embedding contains non-finite values (NaN or Inf)".into(),
));
}
let node_id = self.len();
let node_level = self.random_level();
let mut node_connections: Vec<Vec<u32>> = Vec::with_capacity(node_level + 1);
for _ in 0..=node_level {
node_connections.push(Vec::new());
}
self.push_node(&document.embedding);
self.connections.push(node_connections);
self.ids.push(document.id);
self.contents.push(document.content);
self.metadata.push(document.metadata);
if self.entry_point.is_none() {
self.entry_point = Some(node_id);
self.max_layer = node_level;
return Ok(());
}
self.insert_node(node_id, node_level);
if node_level > self.max_layer {
self.max_layer = node_level;
self.entry_point = Some(node_id);
}
Ok(())
}
pub fn add_embedding(&mut self, id: String, embedding: Vec<f32>) -> Result<()> {
if embedding.len() != self.embedding_dim {
return Err(crate::RagError::DimensionMismatch {
expected: self.embedding_dim,
actual: embedding.len(),
});
}
if !self.is_trained() {
return Err(crate::RagError::NotTrained(format!(
"Storage::{:?} requires a fitted codebook before add_embedding() — call \
`index.train(&sample)` first, or build the index via `HNSWIndex::build`/ \
`build_parallel`, which trains internally from the full corpus",
self.config.storage
)));
}
if embedding.iter().any(|v| !v.is_finite()) {
return Err(crate::RagError::InvalidInput(
"embedding contains non-finite values (NaN or Inf)".into(),
));
}
let node_id = self.len();
let node_level = self.random_level();
let mut node_connections: Vec<Vec<u32>> = Vec::with_capacity(node_level + 1);
for _ in 0..=node_level {
node_connections.push(Vec::new());
}
self.push_node(&embedding);
self.connections.push(node_connections);
self.ids.push(id);
self.contents.push(String::new());
self.metadata.push(None);
if self.entry_point.is_none() {
self.entry_point = Some(node_id);
self.max_layer = node_level;
return Ok(());
}
self.insert_node(node_id, node_level);
if node_level > self.max_layer {
self.max_layer = node_level;
self.entry_point = Some(node_id);
}
Ok(())
}
pub fn search(&self, query: &[f32], k: usize) -> Result<Vec<SearchResult>> {
let mut ctx = SearchContext::new(self.len());
self.search_inner(query, k, &mut ctx)
}
pub fn search_batch(&self, queries: &[Vec<f32>], k: usize) -> Result<Vec<Vec<SearchResult>>> {
use rayon::prelude::*;
queries
.par_iter()
.map_init(
|| SearchContext::new(self.len()),
|ctx, query| self.search_inner(query, k, ctx),
)
.collect()
}
pub fn searcher(&self) -> Searcher<'_> {
Searcher {
index: self,
ctx: SearchContext::new(self.len()),
}
}
fn search_inner(
&self,
query: &[f32],
k: usize,
ctx: &mut SearchContext,
) -> Result<Vec<SearchResult>> {
if query.len() != self.embedding_dim {
return Err(crate::RagError::DimensionMismatch {
expected: self.embedding_dim,
actual: query.len(),
});
}
if self.is_empty() {
return Ok(Vec::new());
}
if ctx.capacity < self.len() {
*ctx = SearchContext::new(self.len());
}
let query_norm = crate::vector::simd::norm_simd(query);
let rq_prepared = self.prepare_rabitq_query(query);
let qprep = QueryPrep {
norm: query_norm,
rabitq: rq_prepared.as_ref(),
};
let entry_point = self.entry_point.unwrap();
let mut current_nearest = vec![entry_point];
for layer in (1..=self.max_layer).rev() {
current_nearest = self
.search_layer(query, ¤t_nearest, 1, layer, ctx, &qprep)
.into_iter()
.map(|(_, id)| id)
.collect();
}
let ef = self.config.ef_search.max(k);
let mut found = self.search_layer(query, ¤t_nearest, ef, 0, ctx, &qprep);
if self.config.storage != Storage::F32 && self.config.rerank_candidates > 0 {
let pool = self.config.rerank_candidates.max(k).min(found.len());
found.truncate(pool);
for entry in found.iter_mut() {
entry.0 = self.exact_distance(query, entry.1);
}
found.sort_unstable_by(|a, b| a.0.total_cmp(&b.0));
}
Ok(found
.into_iter()
.take(k)
.map(|(dist, node_id)| SearchResult {
id: self.ids[node_id].clone(),
content: self.contents[node_id].clone(),
score: self.score_from_distance(dist),
metadata: self.metadata[node_id].clone(),
})
.collect())
}
pub fn clear(&mut self) {
self.nodes.clear();
self.full.clear();
self.connections.clear();
self.ids.clear();
self.contents.clear();
self.metadata.clear();
self.entry_point = None;
self.max_layer = 0;
}
pub fn get_all_documents(&self) -> Vec<Document> {
(0..self.len())
.map(|i| Document {
id: self.ids[i].clone(),
content: self.contents[i].clone(),
embedding: self.get_embedding(i).to_vec(),
metadata: self.metadata[i].clone(),
})
.collect()
}
pub fn config(&self) -> &HNSWConfig {
&self.config
}
pub fn entry_point(&self) -> Option<usize> {
self.entry_point
}
pub fn max_layer(&self) -> usize {
self.max_layer
}
pub fn embedding_dim(&self) -> usize {
self.embedding_dim
}
pub fn shrink_to_fit(&mut self) {
self.nodes.shrink_to_fit();
self.full.shrink_to_fit();
self.ids.shrink_to_fit();
self.contents.shrink_to_fit();
self.metadata.shrink_to_fit();
for layers in &mut self.connections {
for l in layers.iter_mut() {
l.shrink_to_fit();
}
layers.shrink_to_fit();
}
self.connections.shrink_to_fit();
}
pub fn avg_degree_l0(&self) -> f32 {
if self.is_empty() {
return 0.0;
}
let total: usize = (0..self.len())
.map(|i| self.get_neighbors_l0(i).len())
.sum();
total as f32 / self.len() as f32
}
pub fn memory_breakdown(&self) -> MemoryBreakdown {
let vec_header = std::mem::size_of::<Vec<u32>>();
let nested: usize = self
.connections
.iter()
.map(|layers| {
vec_header
+ layers
.iter()
.map(|l| vec_header + l.capacity() * std::mem::size_of::<u32>())
.sum::<usize>()
})
.sum();
let n = self.len();
let arena = self.nodes.capacity() * std::mem::size_of::<u32>();
let hot_vectors = n * vec_words(self.config.storage, self.embedding_dim) * 4;
let cold_vectors = self.full.capacity() * std::mem::size_of::<f32>();
MemoryBreakdown {
embeddings: hot_vectors + cold_vectors,
norms: n * std::mem::size_of::<f32>(),
layer0_links: arena
.saturating_sub(hot_vectors)
.saturating_sub(n * std::mem::size_of::<f32>()),
upper_layer_links: nested,
payload: self
.ids
.iter()
.map(|s| s.capacity() + std::mem::size_of::<String>())
.sum::<usize>()
+ self
.contents
.iter()
.map(|s| s.capacity() + std::mem::size_of::<String>())
.sum::<usize>(),
}
}
fn random_level(&mut self) -> usize {
let uniform: f32 = self.level_rng.random::<f32>().max(f32::EPSILON);
(-uniform.ln() * self.config.ml).floor() as usize
}
fn insert_node(&mut self, node_id: usize, node_level: usize) {
let entry_point = self.entry_point.unwrap();
let mut current_nearest = vec![entry_point];
let node_embedding = self.get_embedding(node_id).to_vec();
let query_norm = crate::vector::simd::norm_simd(&node_embedding);
let rq_prepared = self.prepare_rabitq_query(&node_embedding);
let qprep = QueryPrep {
norm: query_norm,
rabitq: rq_prepared.as_ref(),
};
let mut ctx = SearchContext::new(self.len());
for layer in (node_level + 1..=self.max_layer).rev() {
current_nearest = self
.search_layer(
&node_embedding,
¤t_nearest,
1,
layer,
&mut ctx,
&qprep,
)
.into_iter()
.map(|(_, id)| id)
.collect();
}
for layer in (0..=node_level).rev() {
current_nearest = self
.search_layer(
&node_embedding,
¤t_nearest,
self.config.ef_construction,
layer,
&mut ctx,
&qprep,
)
.into_iter()
.map(|(_, id)| id)
.collect();
let m = if layer == 0 {
self.config.m0
} else {
self.config.m
};
let neighbors = self.select_neighbors(¤t_nearest, &node_embedding, m, layer);
for &neighbor_id in &neighbors {
let neighbor_u32 = neighbor_id as u32;
let node_u32 = node_id as u32;
if layer == 0 {
if !self.l0_contains(node_id, neighbor_u32) {
self.l0_push(node_id, neighbor_u32);
}
if !self.l0_contains(neighbor_id, node_u32) {
self.l0_push(neighbor_id, node_u32);
}
let m0 = self.config.m0;
if self.get_neighbors_l0(neighbor_id).len() > m0 {
let neighbor_embedding = self.get_embedding(neighbor_id).to_vec();
let current: Vec<usize> = self
.get_neighbors_l0(neighbor_id)
.iter()
.map(|&x| x as usize)
.collect();
let pruned =
self.select_neighbors(¤t, &neighbor_embedding, m0, layer);
let pruned: Vec<u32> = pruned.into_iter().map(|x| x as u32).collect();
self.l0_replace(neighbor_id, &pruned);
}
continue;
}
if !self.connections[node_id][layer].contains(&neighbor_u32) {
self.connections[node_id][layer].push(neighbor_u32);
}
if layer < self.connections[neighbor_id].len() {
if !self.connections[neighbor_id][layer].contains(&node_u32) {
self.connections[neighbor_id][layer].push(node_u32);
}
let neighbor_m = self.config.m;
if self.connections[neighbor_id][layer].len() > neighbor_m {
let neighbor_embedding = self.get_embedding(neighbor_id).to_vec();
let neighbor_connections: Vec<usize> = self.connections[neighbor_id][layer]
.iter()
.map(|&x| x as usize)
.collect();
let pruned = self.select_neighbors(
&neighbor_connections,
&neighbor_embedding,
neighbor_m,
layer,
);
self.connections[neighbor_id][layer] =
pruned.into_iter().map(|x| x as u32).collect();
}
}
}
}
}
#[inline]
fn search_layer(
&self,
query: &[f32],
entry_points: &[usize],
ef: usize,
layer: usize,
ctx: &mut SearchContext,
qprep: &QueryPrep,
) -> Vec<(f32, usize)> {
ctx.reset();
for &ep in entry_points {
let dist = self.distance_to_node(query, ep, qprep);
ctx.distance_calls += 1;
ctx.candidates.push(Reverse((OrderedFloat(dist), ep)));
ctx.best.push((OrderedFloat(dist), ep));
ctx.mark_visited(ep);
}
while let Some(Reverse((current_dist, current_id))) = ctx.candidates.pop() {
if ctx.best.len() >= ef {
if let Some(&(furthest_dist, _)) = ctx.best.peek() {
if current_dist > furthest_dist {
break;
}
}
}
let neighbors_l0_slice;
let neighbors: &[u32] = if layer == 0 && !self.nodes.is_empty() {
neighbors_l0_slice = self.get_neighbors_l0(current_id);
neighbors_l0_slice
} else if layer < self.connections[current_id].len() {
&self.connections[current_id][layer]
} else {
&[]
};
if !neighbors.is_empty() {
let n_neighbors = neighbors.len();
const PREFETCH_AHEAD: usize = 2;
const VECTOR_LINES: usize = 3;
let stride = self.stride;
let vec_byte_offset = self.hdr * std::mem::size_of::<u32>();
let mut batch_buf: [(f32, usize); 64] = [(0.0, 0); 64];
let mut batch_count = 0usize;
let mut overflow = Vec::new();
for (i, &neighbor_u32) in neighbors.iter().enumerate() {
let neighbor_id = neighbor_u32 as usize;
unsafe {
let lookahead = i + PREFETCH_AHEAD;
if lookahead < n_neighbors {
let ahead_id = neighbors[lookahead] as usize;
let block =
self.nodes.as_ptr().wrapping_add(ahead_id * stride) as *const u8;
prefetch_read(block);
prefetch_embedding(block.wrapping_add(vec_byte_offset), VECTOR_LINES);
let bitset_ptr =
ctx.visited.bits.as_ptr().wrapping_add(ahead_id >> 6) as *const u8;
prefetch_read(bitset_ptr);
}
}
if !ctx.is_visited(neighbor_id) {
ctx.mark_visited(neighbor_id);
let dist = self.distance_to_node(query, neighbor_id, qprep);
ctx.distance_calls += 1;
if batch_count < batch_buf.len() {
batch_buf[batch_count] = (dist, neighbor_id);
batch_count += 1;
} else {
overflow.push((dist, neighbor_id));
}
}
}
let mut consider = |dist: f32, neighbor_id: usize| {
let dist_ord = OrderedFloat(dist);
if ctx.best.len() < ef {
ctx.candidates.push(Reverse((dist_ord, neighbor_id)));
ctx.best.push((dist_ord, neighbor_id));
} else if let Some(&(furthest_dist, _)) = ctx.best.peek() {
if dist_ord < furthest_dist {
ctx.candidates.push(Reverse((dist_ord, neighbor_id)));
ctx.best.push((dist_ord, neighbor_id));
if ctx.best.len() > ef {
ctx.best.pop();
}
}
}
};
for &(dist, neighbor_id) in &batch_buf[..batch_count] {
consider(dist, neighbor_id);
}
for &(dist, neighbor_id) in &overflow {
consider(dist, neighbor_id);
}
}
}
let mut results: Vec<(f32, usize)> = ctx
.best
.drain()
.map(|(OrderedFloat(dist), id)| (dist, id))
.collect();
results.sort_by(|a, b| a.0.total_cmp(&b.0));
results
}
fn select_neighbors(
&self,
candidates: &[usize],
query: &[f32],
m: usize,
layer: usize,
) -> Vec<usize> {
let mut scored: Vec<(f32, usize)> = candidates
.iter()
.map(|&id| (self.distance(query, self.get_embedding(id)), id))
.collect();
scored.sort_by(|a, b| a.0.total_cmp(&b.0));
select_neighbors_core(
&scored,
m,
self.config.use_heuristic,
self.config.extend_candidates,
self.config.keep_pruned_connections,
|id| {
if layer == 0 {
self.get_neighbors_l0(id).iter().map(|&n| n as usize).collect()
} else if layer < self.connections[id].len() {
self.connections[id][layer].iter().map(|&n| n as usize).collect()
} else {
Vec::new()
}
},
|id| self.distance(query, self.get_embedding(id)),
|a, b| self.distance(self.get_embedding(a), self.get_embedding(b)),
)
.into_iter()
.map(|(_, id)| id)
.collect()
}
#[inline]
fn distance(&self, a: &[f32], b: &[f32]) -> f32 {
Self::metric_distance(self.config.metric, a, b)
}
#[inline]
fn metric_distance(metric: DistanceMetric, a: &[f32], b: &[f32]) -> f32 {
match metric {
DistanceMetric::Cosine => 1.0 - crate::vector::simd::cosine_similarity_simd(a, b),
DistanceMetric::L2 => crate::vector::simd::l2_squared_distance_simd(a, b),
}
}
#[inline]
fn prepare_rabitq_query(&self, query: &[f32]) -> Option<crate::vector::rabitq::PreparedQuery> {
if self.config.storage != Storage::RaBitQ {
return None;
}
let rq = self
.rabitq
.as_ref()
.expect("RaBitQ storage requires fit_codebook to run before any query");
let query = self.rabitq_cosine_input(query);
Some(rq.prepare_query(&query))
}
#[inline]
fn distance_to_node(&self, query: &[f32], node_id: usize, qprep: &QueryPrep) -> f32 {
if self.config.storage == Storage::SQ8 {
return match self.config.metric {
DistanceMetric::L2 => crate::vector::simd::sq8_asymmetric_l2_simd(
query,
self.get_codes(node_id),
&self.q_min,
&self.q_scale,
),
DistanceMetric::Cosine => {
let norm_b = self.get_norm(node_id);
if qprep.norm == 0.0 || norm_b == 0.0 {
return 1.0;
}
let dot = crate::vector::simd::sq8_asymmetric_dot_simd(
query,
self.get_codes(node_id),
&self.q_min,
&self.q_scale,
);
let similarity = (dot / (qprep.norm * norm_b)).clamp(-1.0, 1.0);
1.0 - similarity
}
};
}
if self.config.storage == Storage::RaBitQ {
let prepared = qprep.rabitq.expect(
"Storage::RaBitQ traversal requires a query prepared via prepare_rabitq_query",
);
let (dtc_sq, est_factor, bits) = self.get_rabitq_code(node_id);
let raw = crate::vector::simd::rabitq_asymmetric_l2_simd(
prepared.rq(),
bits,
dtc_sq,
est_factor,
prepared.qn_sq(),
);
return match self.config.metric {
DistanceMetric::L2 => raw,
DistanceMetric::Cosine => (raw * 0.5).clamp(0.0, 2.0),
};
}
let embedding = self.get_embedding(node_id);
match self.config.metric {
DistanceMetric::Cosine => {
let norm_b = self.get_norm(node_id);
if qprep.norm == 0.0 || norm_b == 0.0 {
return 1.0;
}
crate::vector::simd::cosine_distance_prenorm(query, embedding, norm_b)
}
DistanceMetric::L2 => crate::vector::simd::l2_squared_distance_simd(query, embedding),
}
}
#[inline]
fn exact_distance(&self, query: &[f32], node_id: usize) -> f32 {
let embedding = self.get_embedding(node_id);
match self.config.metric {
DistanceMetric::L2 => crate::vector::simd::l2_squared_distance_simd(query, embedding),
DistanceMetric::Cosine => {
crate::vector::simd::cosine_distance_prenorm(
query,
embedding,
self.get_norm(node_id),
)
}
}
}
#[inline]
fn score_from_distance(&self, dist: f32) -> f32 {
match self.config.metric {
DistanceMetric::Cosine => 1.0 - dist,
DistanceMetric::L2 => 1.0 / (1.0 + dist.max(0.0).sqrt()),
}
}
pub fn build_parallel(embeddings: Vec<Vec<f32>>, config: HNSWConfig) -> Self {
assert!(!embeddings.is_empty(), "Cannot build from empty embeddings");
let embedding_dim = embeddings[0].len();
let n = embeddings.len();
if n == 1 {
return Self::build_single(embeddings, config);
}
let m0 = config.m0.min(M0_MAX);
let m = config.m.min(M_MAX);
assert!(
config.m0 <= M0_MAX && config.m <= M_MAX,
"parallel build supports m <= {M_MAX} and m0 <= {M0_MAX} (the node arrays are fixed-size); got m={} m0={}. Use BuildStrategy::Sequential for a larger degree.",
config.m,
config.m0
);
let ml = config.ml;
let ef_construction = config.ef_construction;
let seed = config.seed.unwrap_or_else(rand::random);
let mut rng = StdRng::seed_from_u64(seed);
let mut sizes = Vec::new();
let mut num = n;
loop {
let next = (num as f32 * ml) as usize;
if next < M_MAX {
break;
}
sizes.push((num - next, num));
num = next;
}
sizes.push((num, num));
sizes.reverse();
let num_batches = sizes.len();
let top = LayerId(num_batches - 1);
assert!(n < u32::MAX as usize);
let mut shuffled: Vec<(u32, usize)> = (0..n).map(|i| (rng.random::<u32>(), i)).collect();
shuffled.sort_unstable_by_key(|&(r, _)| r);
let points: Vec<Vec<f32>> = shuffled
.iter()
.map(|&(_, idx)| embeddings[idx].clone())
.collect();
let mut ranges = Vec::with_capacity(num_batches);
for (i, (size, cumulative)) in sizes.into_iter().enumerate() {
let start = cumulative - size;
let batch_id = LayerId(num_batches - i - 1);
ranges.push((batch_id, max(start, 1)..cumulative));
}
let zero: Vec<RwLock<ZeroNode>> =
(0..n).map(|_| RwLock::new(ZeroNode::default())).collect();
let mut layers: Vec<Vec<UpperNode>> = vec![Vec::new(); top.0];
let pool = SearchPool::new(n, config.metric, config.keep_pruned_connections);
let use_heuristic = config.use_heuristic;
let extend_candidates = config.extend_candidates;
for (batch, range) in ranges {
let end = range.end;
if batch.0 == top.0 {
for i in range {
Self::par_insert(
PointId(i as u32),
batch,
&zero,
&layers,
&points,
&pool,
ef_construction,
top,
m,
m0,
use_heuristic,
extend_candidates,
);
}
} else {
range.into_par_iter().for_each(|i| {
Self::par_insert(
PointId(i as u32),
batch,
&zero,
&layers,
&points,
&pool,
ef_construction,
top,
m,
m0,
use_heuristic,
extend_candidates,
);
});
}
if !batch.is_zero() {
zero[..end]
.par_iter()
.map(|z| UpperNode::from_zero(&z.read(), m))
.collect_into_vec(&mut layers[batch.0 - 1]);
}
}
Self::convert_parallel_to_index(zero, layers, points, shuffled, embedding_dim, config, top)
}
#[allow(clippy::too_many_arguments)]
fn par_insert(
new: PointId,
target_layer: LayerId, zero: &[RwLock<ZeroNode>],
layers: &[Vec<UpperNode>],
points: &[Vec<f32>],
pool: &SearchPool,
ef_construction: usize,
top: LayerId,
m: usize,
m0: usize,
use_heuristic: bool,
extend_candidates: bool,
) {
let metric = pool.metric;
let keep_pruned = pool.keep_pruned;
let mut search = pool.pop();
search.visited.reserve(points.len());
let point = &points[new.as_usize()];
search.reset();
search.push(PointId(0), point, points);
for cur in top.descend() {
search.ef = if cur.0 <= target_layer.0 {
ef_construction
} else {
1
};
if cur.0 > target_layer.0 {
if cur.0 <= layers.len() && !layers[cur.0 - 1].is_empty() {
search.search_upper(point, &layers[cur.0 - 1], points, m);
search.cull();
}
} else {
search.search_zero(point, zero, points, ef_construction);
break; }
}
let found = Self::par_select_heuristic(
metric,
point,
search.select_simple(),
points,
m0,
keep_pruned,
use_heuristic,
extend_candidates,
zero,
);
{
let mut node = zero[new.as_usize()].write();
for (i, candidate) in found.iter().take(m0).enumerate() {
node.nearest[i] = candidate.pid;
}
}
for candidate in found.iter().take(m0) {
Self::add_reverse_connection(
metric,
zero,
points,
new,
candidate.pid,
keep_pruned,
m0,
use_heuristic,
);
}
pool.push(search);
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
fn par_select_heuristic(
metric: DistanceMetric,
query: &[f32],
sorted: &[Candidate],
points: &[Vec<f32>],
m: usize,
keep_pruned: bool,
use_heuristic: bool,
extend_candidates: bool,
zero: &[RwLock<ZeroNode>],
) -> Vec<Candidate> {
let scored: Vec<(f32, usize)> =
sorted.iter().map(|c| (c.distance, c.pid.as_usize())).collect();
select_neighbors_core(
&scored,
m,
use_heuristic,
extend_candidates,
keep_pruned,
|id| zero[id].read().iter().map(|pid| pid.as_usize()).collect(),
|id| Self::parallel_distance(metric, query, &points[id]),
|a, b| Self::parallel_distance(metric, &points[a], &points[b]),
)
.into_iter()
.map(|(distance, id)| Candidate { distance, pid: PointId(id as u32) })
.collect()
}
#[allow(clippy::too_many_arguments)]
fn add_reverse_connection(
metric: DistanceMetric,
zero: &[RwLock<ZeroNode>],
points: &[Vec<f32>],
new: PointId,
neighbor: PointId,
keep_pruned: bool,
m0: usize,
use_heuristic: bool,
) {
let mut node = zero[neighbor.as_usize()].write();
let neighbor_point = &points[neighbor.as_usize()];
let count = node.count();
if node.nearest[..count].contains(&new) {
return;
}
if count < m0 {
let new_dist = Self::parallel_distance(metric, neighbor_point, &points[new.as_usize()]);
let pos = {
let mut left = 0;
let mut right = count;
while left < right {
let mid = (left + right) / 2;
let mid_dist = Self::parallel_distance(
metric,
neighbor_point,
&points[node.nearest[mid].as_usize()],
);
if mid_dist < new_dist {
left = mid + 1;
} else {
right = mid;
}
}
left
};
for i in (pos..count).rev() {
node.nearest[i + 1] = node.nearest[i];
}
node.nearest[pos] = new;
return;
}
let mut cands: Vec<Candidate> = node
.iter()
.chain(std::iter::once(new))
.map(|pid| Candidate {
distance: Self::parallel_distance(metric, neighbor_point, &points[pid.as_usize()]),
pid,
})
.collect();
cands.sort_unstable();
let selected = Self::par_select_heuristic(
metric,
neighbor_point,
&cands,
points,
m0,
keep_pruned,
use_heuristic,
false,
zero,
);
for (i, slot) in node.nearest.iter_mut().enumerate() {
*slot = selected.get(i).map_or(PointId(INVALID), |c| c.pid);
}
}
fn convert_parallel_to_index(
zero: Vec<RwLock<ZeroNode>>,
layers: Vec<Vec<UpperNode>>,
points: Vec<Vec<f32>>,
shuffled: Vec<(u32, usize)>,
embedding_dim: usize,
config: HNSWConfig,
top: LayerId,
) -> Self {
let n = points.len();
let num_layers = top.0 + 1; let zero_final: Vec<ZeroNode> = zero.into_iter().map(|n| n.into_inner()).collect();
let mut connections: Vec<Vec<Vec<u32>>> = Vec::with_capacity(n);
for i in 0..n {
let mut node_connections: Vec<Vec<u32>> = Vec::with_capacity(num_layers);
let layer0: Vec<u32> = zero_final[i].iter().map(|p| p.as_usize() as u32).collect();
node_connections.push(layer0);
for layer in &layers {
if i < layer.len() {
let layer_conns: Vec<u32> =
layer[i].iter().map(|p| p.as_usize() as u32).collect();
node_connections.push(layer_conns);
}
}
connections.push(node_connections);
}
let ids: Vec<String> = shuffled.iter().map(|&(_, orig)| orig.to_string()).collect();
let mut index = Self {
level_rng: StdRng::seed_from_u64(config.seed.unwrap_or_else(rand::random)),
embedding_dim,
stride: node_stride(config.m0, embedding_dim, config.storage),
hdr: node_hdr_len(config.m0),
config,
nodes: Vec::new(),
connections,
q_min: Vec::new(),
q_scale: Vec::new(),
full: Vec::new(),
rabitq: None,
ids,
contents: vec![String::new(); n],
metadata: vec![None; n],
entry_point: Some(0),
max_layer: top.0,
};
index.fit_codebook(&points);
index.nodes.reserve(n * index.stride);
if index.config.storage != Storage::F32 && index.config.rerank_candidates > 0 {
index.full.reserve(n * embedding_dim);
}
for p in &points {
index.push_node(p);
}
index.migrate_l0_into_arena();
index.shrink_to_fit();
index
}
fn build_single(embeddings: Vec<Vec<f32>>, config: HNSWConfig) -> Self {
let embedding_dim = embeddings[0].len();
let mut index = Self::new(embedding_dim, config);
index.fit_codebook(&embeddings);
index.push_node(&embeddings[0]);
index.connections.push(vec![Vec::new()]);
index.ids.push("0".to_string());
index.contents.push(String::new());
index.metadata.push(None);
index.entry_point = Some(0);
index
}
#[inline]
fn parallel_distance(metric: DistanceMetric, a: &[f32], b: &[f32]) -> f32 {
Self::metric_distance(metric, a, b)
}
}
impl crate::index::VectorIndex for HNSWIndex {
fn add(&mut self, document: Document) -> Result<()> {
self.add(document)
}
fn search(&self, query: &[f32], k: usize) -> Result<Vec<SearchResult>> {
self.search(query, k)
}
fn len(&self) -> usize {
self.len()
}
fn clear(&mut self) {
self.clear()
}
fn embedding_dim(&self) -> usize {
self.embedding_dim()
}
}
impl crate::index::VectorIndexSnapshot for HNSWIndex {
fn get_all_documents(&self) -> Vec<Document> {
self.get_all_documents()
}
}
const M0_MAX: usize = 64;
const M_MAX: usize = 32;
const INVALID: u32 = u32::MAX;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
struct PointId(u32);
impl PointId {
fn as_usize(self) -> usize {
self.0 as usize
}
fn is_valid(self) -> bool {
self.0 != INVALID
}
}
#[derive(Clone)]
struct ZeroNode {
nearest: [PointId; M0_MAX],
}
impl Default for ZeroNode {
fn default() -> Self {
Self {
nearest: [PointId(INVALID); M0_MAX],
}
}
}
impl ZeroNode {
fn count(&self) -> usize {
self.nearest.iter().take_while(|p| p.is_valid()).count()
}
fn iter(&self) -> impl Iterator<Item = PointId> + '_ {
self.nearest.iter().copied().take_while(|p| p.is_valid())
}
}
#[derive(Clone)]
struct UpperNode {
nearest: [PointId; M_MAX],
}
impl Default for UpperNode {
fn default() -> Self {
Self {
nearest: [PointId(INVALID); M_MAX],
}
}
}
impl UpperNode {
fn from_zero(zero: &ZeroNode, m: usize) -> Self {
let mut node = Self::default();
for (i, &pid) in zero.nearest.iter().take(m.min(M_MAX)).enumerate() {
node.nearest[i] = pid;
}
node
}
fn iter(&self) -> impl Iterator<Item = PointId> + '_ {
self.nearest.iter().copied().take_while(|p| p.is_valid())
}
}
struct Visited {
store: Vec<u8>,
generation: u8,
}
impl Visited {
fn new(capacity: usize) -> Self {
Self {
store: vec![0; capacity],
generation: 1,
}
}
fn clear(&mut self) {
if self.generation == 255 {
self.store.fill(0);
self.generation = 1;
} else {
self.generation += 1;
}
}
fn insert(&mut self, pid: PointId) -> bool {
let idx = pid.as_usize();
if self.store[idx] == self.generation {
false
} else {
self.store[idx] = self.generation;
true
}
}
fn reserve(&mut self, capacity: usize) {
if self.store.len() < capacity {
self.store.resize(capacity, 0);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
struct Candidate {
distance: f32,
pid: PointId,
}
impl Eq for Candidate {}
impl PartialOrd for Candidate {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for Candidate {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.distance
.total_cmp(&other.distance)
.then_with(|| self.pid.cmp(&other.pid))
}
}
struct Search {
metric: DistanceMetric,
candidates: BinaryHeap<Reverse<Candidate>>,
nearest: Vec<Candidate>,
visited: Visited,
ef: usize,
}
impl Search {
fn new(capacity: usize, metric: DistanceMetric) -> Self {
Self {
metric,
candidates: BinaryHeap::new(),
nearest: Vec::new(),
visited: Visited::new(capacity),
ef: 1,
}
}
fn reset(&mut self) {
self.candidates.clear();
self.nearest.clear();
self.visited.clear();
}
fn push(&mut self, pid: PointId, point: &[f32], points: &[Vec<f32>]) {
let distance = HNSWIndex::parallel_distance(self.metric, point, &points[pid.as_usize()]);
let candidate = Candidate { distance, pid };
self.candidates.push(Reverse(candidate));
self.nearest.push(candidate);
self.visited.insert(pid);
}
fn cull(&mut self) {
self.candidates.clear();
for &candidate in &self.nearest {
self.candidates.push(Reverse(candidate));
}
self.visited.clear();
for c in &self.nearest {
self.visited.insert(c.pid);
}
}
fn search_zero(
&mut self,
point: &[f32],
layer: &[RwLock<ZeroNode>],
points: &[Vec<f32>],
num: usize,
) {
while let Some(Reverse(candidate)) = self.candidates.pop() {
if let Some(furthest) = self.nearest.last() {
if candidate.distance > furthest.distance && self.nearest.len() >= self.ef {
break;
}
}
let node = layer[candidate.pid.as_usize()].read();
for neighbor_pid in node.iter() {
if self.visited.insert(neighbor_pid) {
let distance = HNSWIndex::parallel_distance(
self.metric,
point,
&points[neighbor_pid.as_usize()],
);
let new_candidate = Candidate {
distance,
pid: neighbor_pid,
};
let dominated = self.nearest.len() >= self.ef
&& self
.nearest
.last()
.map(|f| distance > f.distance)
.unwrap_or(false);
if !dominated {
self.candidates.push(Reverse(new_candidate));
let pos = self
.nearest
.binary_search(&new_candidate)
.unwrap_or_else(|i| i);
if pos < self.ef {
self.nearest.insert(pos, new_candidate);
if self.nearest.len() > self.ef {
self.nearest.pop();
}
}
}
}
}
}
self.nearest.truncate(num);
}
fn search_upper(
&mut self,
point: &[f32],
layer: &[UpperNode],
points: &[Vec<f32>],
num: usize,
) {
if layer.is_empty() {
return;
}
while let Some(Reverse(candidate)) = self.candidates.pop() {
if let Some(furthest) = self.nearest.last() {
if candidate.distance > furthest.distance && self.nearest.len() >= self.ef {
break;
}
}
if candidate.pid.as_usize() >= layer.len() {
continue;
}
let node = &layer[candidate.pid.as_usize()];
for neighbor_pid in node.iter() {
if self.visited.insert(neighbor_pid) {
let distance = HNSWIndex::parallel_distance(
self.metric,
point,
&points[neighbor_pid.as_usize()],
);
let new_candidate = Candidate {
distance,
pid: neighbor_pid,
};
let dominated = self.nearest.len() >= self.ef
&& self
.nearest
.last()
.map(|f| distance > f.distance)
.unwrap_or(false);
if !dominated {
self.candidates.push(Reverse(new_candidate));
let pos = self
.nearest
.binary_search(&new_candidate)
.unwrap_or_else(|i| i);
if pos < self.ef {
self.nearest.insert(pos, new_candidate);
if self.nearest.len() > self.ef {
self.nearest.pop();
}
}
}
}
}
}
self.nearest.truncate(num);
}
fn select_simple(&self) -> &[Candidate] {
&self.nearest
}
}
struct SearchPool {
pool: Mutex<Vec<Search>>,
capacity: usize,
metric: DistanceMetric,
keep_pruned: bool,
}
impl SearchPool {
fn new(capacity: usize, metric: DistanceMetric, keep_pruned: bool) -> Self {
Self {
pool: Mutex::new(Vec::new()),
capacity,
keep_pruned,
metric,
}
}
fn pop(&self) -> Search {
self.pool
.lock()
.pop()
.unwrap_or_else(|| Search::new(self.capacity, self.metric))
}
fn push(&self, mut search: Search) {
search.reset();
self.pool.lock().push(search);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
struct LayerId(usize);
impl LayerId {
fn is_zero(self) -> bool {
self.0 == 0
}
fn descend(self) -> impl Iterator<Item = LayerId> {
(0..=self.0).rev().map(LayerId)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_document(id: &str, embedding: Vec<f32>) -> Document {
Document {
id: id.to_string(),
content: format!("Content for {}", id),
embedding,
metadata: None,
}
}
fn generate_random_vector(dim: usize, seed: u64) -> Vec<f32> {
use rand::SeedableRng;
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
(0..dim).map(|_| rng.random::<f32>() * 2.0 - 1.0).collect()
}
#[test]
fn test_hnsw_config_default() {
let config = HNSWConfig::default();
assert_eq!(config.m, 32); assert_eq!(config.m0, 64); assert_eq!(config.ef_construction, 100);
assert_eq!(config.ef_search, 100);
assert!((config.ml - (1.0 / 32_f32.ln())).abs() < 0.01);
assert!(config.use_heuristic); assert!(!config.extend_candidates);
assert!(config.keep_pruned_connections);
}
#[test]
fn test_hnsw_config_builders() {
let config = HNSWConfig::default()
.with_m(32)
.with_ef_search(100)
.with_ef_construction(400)
.with_simple_selection()
.with_extended_candidates();
assert_eq!(config.m, 32);
assert_eq!(config.m0, 64);
assert_eq!(config.ef_search, 100);
assert_eq!(config.ef_construction, 400);
assert!(!config.use_heuristic);
assert!(config.extend_candidates);
}
#[test]
fn test_hnsw_new() {
let index = HNSWIndex::with_defaults(128);
assert_eq!(index.embedding_dim, 128);
assert_eq!(index.len(), 0);
assert!(index.is_empty());
}
#[test]
fn search_layer_considers_neighbors_beyond_fixed_stack_batch() {
let mut config = HNSWConfig::default().with_m(64);
config.m0 = 128;
let mut index = HNSWIndex::new(2, config);
let total_nodes = 66usize; index.connections = vec![vec![Vec::new()]; total_nodes];
index.metadata = vec![None; total_nodes];
index.entry_point = Some(0);
index.max_layer = 0;
for node in 0..total_nodes {
let v: [f32; 2] = if node == 65 { [1.0, 0.0] } else { [0.0, 1.0] };
index.push_node(&v);
index.ids.push(format!("doc-{node}"));
index.contents.push(String::new());
}
let neighbors: Vec<u32> = (1..=65).map(|n| n as u32).collect();
index.l0_replace(0, &neighbors);
let mut ctx = SearchContext::new(index.len());
let qprep = QueryPrep {
norm: 1.0,
rabitq: None,
};
let candidates = index.search_layer(&[1.0, 0.0], &[0], 66, 0, &mut ctx, &qprep);
assert!(
candidates.iter().any(|&(_, id)| id == 65),
"best neighbor from position >64 should be considered"
);
}
#[test]
fn test_add_single_document() {
let mut index = HNSWIndex::with_defaults(3);
let doc = create_test_document("doc1", vec![1.0, 0.0, 0.0]);
assert!(index.add(doc).is_ok());
assert_eq!(index.len(), 1);
assert!(!index.is_empty());
}
#[test]
fn test_add_dimension_mismatch() {
let mut index = HNSWIndex::with_defaults(3);
let doc = create_test_document("doc1", vec![1.0, 0.0]);
assert!(index.add(doc).is_err());
}
#[test]
fn sq8_add_without_train_errs_not_panics() {
let mut index = HNSWIndex::new(
4,
HNSWConfig {
storage: Storage::SQ8,
rerank_candidates: 100,
metric: DistanceMetric::L2,
..Default::default()
},
);
let err = index
.add(create_test_document("doc1", vec![0.5, -0.3, 0.8, 0.1]))
.expect_err("add() before train() must error, not panic");
assert!(
matches!(err, crate::RagError::NotTrained(_)),
"expected NotTrained, got {err:?}"
);
let msg = err.to_string();
assert!(
msg.contains("train("),
"error message should name the method to call: {msg}"
);
}
#[test]
fn rabitq_add_without_train_errs_not_panics() {
let mut index = HNSWIndex::new(
4,
HNSWConfig {
storage: Storage::RaBitQ,
rerank_candidates: 100,
metric: DistanceMetric::L2,
..Default::default()
},
);
let err = index
.add(create_test_document("doc1", vec![0.5, -0.3, 0.8, 0.1]))
.expect_err("add() before train() must error, not panic");
assert!(
matches!(err, crate::RagError::NotTrained(_)),
"expected NotTrained, got {err:?}"
);
}
#[test]
fn f32_add_works_without_training() {
let mut index = HNSWIndex::new(4, HNSWConfig::default());
assert!(index
.add(create_test_document("doc1", vec![0.5, -0.3, 0.8, 0.1]))
.is_ok());
assert_eq!(index.len(), 1);
assert!(index.search(&[0.5, -0.3, 0.8, 0.1], 1).unwrap()[0].id == "doc1");
}
fn train_then_add_recall(storage: Storage) -> f32 {
let mut rng = StdRng::seed_from_u64(303);
let dim = 16;
let n_clusters = 8;
let per_cluster = 40;
let centers: Vec<Vec<f32>> = (0..n_clusters)
.map(|_| (0..dim).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..n_clusters * per_cluster)
.map(|i| {
let c = ¢ers[i % n_clusters];
c.iter().map(|x| x + rng.random::<f32>() * 0.4).collect()
})
.collect();
let queries: Vec<Vec<f32>> = (0..40)
.map(|i| {
let c = ¢ers[i % n_clusters];
c.iter().map(|x| x + rng.random::<f32>() * 0.4).collect()
})
.collect();
let mut index = HNSWIndex::new(
dim,
HNSWConfig {
metric: DistanceMetric::L2,
m: 16,
m0: 32,
ef_construction: 150,
ef_search: 150,
storage,
rerank_candidates: 50,
..Default::default()
},
);
index.train(&base[..base.len() / 2]).expect("train");
for (i, v) in base.iter().enumerate() {
index
.add_embedding(i.to_string(), v.clone())
.expect("add_embedding");
}
let k = 10;
let mut total_recall = 0.0f32;
for q in &queries {
let mut exact: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| {
let d: f32 = v.iter().zip(q).map(|(a, b)| (a - b) * (a - b)).sum();
(d, i)
})
.collect();
exact.sort_by(|a, b| a.0.total_cmp(&b.0));
let truth: HashSet<usize> = exact.iter().take(k).map(|(_, i)| *i).collect();
let got: HashSet<usize> = index
.search(q, k)
.expect("search")
.into_iter()
.filter_map(|r| r.id.parse::<usize>().ok())
.collect();
total_recall += truth.intersection(&got).count() as f32 / k as f32;
}
total_recall / queries.len() as f32
}
#[test]
fn sq8_train_then_add_retrieves_correctly() {
let recall = train_then_add_recall(Storage::SQ8);
assert!(
recall > 0.6,
"Storage::SQ8 via train()+add_embedding(): recall@10 = {:.1}%, held-out queries, \
brute-force ground truth",
recall * 100.0
);
}
#[test]
fn rabitq_train_then_add_retrieves_correctly() {
let recall = train_then_add_recall(Storage::RaBitQ);
assert!(
recall > 0.6,
"Storage::RaBitQ via train()+add_embedding(): recall@10 = {:.1}%, held-out \
queries, brute-force ground truth",
recall * 100.0
);
}
#[test]
fn train_on_nonempty_index_errs() {
let mut index = HNSWIndex::new(
3,
HNSWConfig {
storage: Storage::SQ8,
..Default::default()
},
);
index.train(&[vec![1.0, 2.0, 3.0]]).expect("first train");
index
.add_embedding("0".into(), vec![1.0, 2.0, 3.0])
.expect("add after train");
assert!(
index.train(&[vec![4.0, 5.0, 6.0]]).is_err(),
"retraining a non-empty index must be rejected"
);
}
#[test]
fn rejects_inf_embedding() {
let mut index = HNSWIndex::new(3, HNSWConfig::default());
let doc = Document {
id: "inf".to_string(),
content: "test".to_string(),
embedding: vec![f32::INFINITY, 0.0, 0.0],
metadata: None,
};
assert!(index.add(doc).is_err());
let doc_neg = Document {
id: "neg_inf".to_string(),
content: "test".to_string(),
embedding: vec![0.0, f32::NEG_INFINITY, 0.0],
metadata: None,
};
assert!(index.add(doc_neg).is_err());
}
#[test]
fn an_undersized_scratch_context_is_regrown_before_use() {
let mut index = HNSWIndex::new(3, HNSWConfig::default());
for i in 0..64 {
index
.add(Document {
id: format!("doc-{i}"),
content: String::new(),
embedding: vec![(i as f32) * 0.01, 1.0 - (i as f32) * 0.01, 0.0],
metadata: None,
})
.unwrap();
}
let mut stale = SearchContext::new(1);
assert!(stale.capacity < index.len());
let results = index.search_inner(&[1.0, 0.0, 0.0], 5, &mut stale).unwrap();
assert_eq!(results.len(), 5);
assert!(
stale.capacity >= index.len(),
"search_inner must regrow an undersized context, not index past the end of its bitset"
);
}
#[test]
fn test_search_empty_index() {
let index = HNSWIndex::with_defaults(3);
let query = vec![1.0, 0.0, 0.0];
let results = index.search(&query, 5).unwrap();
assert!(results.is_empty());
}
#[test]
fn test_search_single_document() {
let mut index = HNSWIndex::with_defaults(3);
let doc = create_test_document("doc1", vec![1.0, 0.0, 0.0]);
index.add(doc).unwrap();
let query = vec![1.0, 0.0, 0.0];
let results = index.search(&query, 1).unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].id, "doc1");
assert!((results[0].score - 1.0).abs() < 1e-6);
}
#[test]
fn test_search_multiple_documents() {
let mut index = HNSWIndex::with_defaults(3);
let docs = vec![
create_test_document("doc1", vec![1.0, 0.0, 0.0]),
create_test_document("doc2", vec![0.0, 1.0, 0.0]),
create_test_document("doc3", vec![0.0, 0.0, 1.0]),
create_test_document("doc4", vec![1.0, 1.0, 0.0]),
];
for doc in docs {
index.add(doc).unwrap();
}
let query = vec![1.0, 0.0, 0.0];
let results = index.search(&query, 2).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].id, "doc1");
assert!(results[0].score > 0.9);
}
#[test]
fn test_search_exact_match() {
let mut index = HNSWIndex::with_defaults(3);
let embedding = vec![0.5, 0.5, 0.7072];
let doc = create_test_document("doc1", embedding.clone());
index.add(doc).unwrap();
let results = index.search(&embedding, 1).unwrap();
assert_eq!(results.len(), 1);
assert!((results[0].score - 1.0).abs() < 1e-5);
}
#[test]
fn test_clear() {
let mut index = HNSWIndex::with_defaults(3);
for i in 0..5 {
let doc = create_test_document(&format!("doc{}", i), vec![i as f32, 0.0, 0.0]);
index.add(doc).unwrap();
}
assert_eq!(index.len(), 5);
index.clear();
assert_eq!(index.len(), 0);
assert!(index.is_empty());
}
#[test]
fn test_random_dataset_100_vectors() {
let dim = 128;
let mut index = HNSWIndex::with_defaults(dim);
for i in 0..100 {
let embedding = generate_random_vector(dim, i);
let doc = create_test_document(&format!("doc{}", i), embedding);
index.add(doc).unwrap();
}
assert_eq!(index.len(), 100);
let query = generate_random_vector(dim, 9999);
let results = index.search(&query, 10).unwrap();
assert_eq!(results.len(), 10);
for i in 0..results.len() - 1 {
assert!(results[i].score >= results[i + 1].score);
}
}
#[test]
fn test_random_dataset_1000_vectors() {
let dim = 64;
let mut index = HNSWIndex::with_defaults(dim);
for i in 0..1000 {
let embedding = generate_random_vector(dim, i);
let doc = create_test_document(&format!("doc{}", i), embedding);
index.add(doc).unwrap();
}
assert_eq!(index.len(), 1000);
for seed in [111, 222, 333, 444, 555] {
let query = generate_random_vector(dim, seed);
let results = index.search(&query, 20).unwrap();
assert_eq!(results.len(), 20);
for i in 0..results.len() - 1 {
assert!(results[i].score >= results[i + 1].score);
}
for result in &results {
assert!(result.score >= -1.0 && result.score <= 1.0);
}
}
}
#[test]
fn test_recall_with_known_neighbors() {
let dim = 32;
let mut index = HNSWIndex::with_defaults(dim);
let query = generate_random_vector(dim, 0);
for i in 0..100 {
let mut embedding = generate_random_vector(dim, i + 1);
if i < 10 {
for j in 0..dim {
embedding[j] = query[j] * 0.9 + embedding[j] * 0.1;
}
}
let doc = create_test_document(&format!("doc{}", i), embedding);
index.add(doc).unwrap();
}
let results = index.search(&query, 10).unwrap();
let mut recall_count = 0;
for result in &results {
let doc_num: usize = result.id.strip_prefix("doc").unwrap().parse().unwrap();
if doc_num < 10 {
recall_count += 1;
}
}
assert!(recall_count >= 7, "Recall too low: {}/10", recall_count);
}
#[test]
fn test_search_dimension_mismatch() {
let mut index = HNSWIndex::with_defaults(3);
let doc = create_test_document("doc1", vec![1.0, 0.0, 0.0]);
index.add(doc).unwrap();
let query = vec![1.0, 0.0]; assert!(index.search(&query, 1).is_err());
}
#[test]
fn test_metadata_preservation() {
let mut index = HNSWIndex::with_defaults(3);
let mut doc = create_test_document("doc1", vec![1.0, 0.0, 0.0]);
doc.metadata = Some(serde_json::json!({"category": "test", "priority": 5}));
index.add(doc).unwrap();
let query = vec![1.0, 0.0, 0.0];
let results = index.search(&query, 1).unwrap();
assert_eq!(results.len(), 1);
assert!(results[0].metadata.is_some());
let metadata = results[0].metadata.as_ref().unwrap();
assert_eq!(metadata["category"], "test");
assert_eq!(metadata["priority"], 5);
}
#[test]
fn test_search_with_nan_query_does_not_panic() {
let mut index = HNSWIndex::with_defaults(3);
index
.add(create_test_document("doc1", vec![1.0, 0.0, 0.0]))
.unwrap();
index
.add(create_test_document("doc2", vec![0.0, 1.0, 0.0]))
.unwrap();
let query = vec![f32::NAN, 0.0, 0.0];
let outcome = std::panic::catch_unwind(|| index.search(&query, 2));
assert!(outcome.is_ok(), "search panicked when query contains NaN");
}
#[test]
#[should_panic(expected = "All embeddings must have the same dimension")]
fn test_build_rejects_mismatched_dimensions() {
let _ = HNSWIndex::build(
vec![vec![1.0, 0.0, 0.0], vec![1.0, 0.0]],
HNSWConfig::default(),
);
}
#[test]
fn test_add_rejects_nan_embedding() {
let mut index = HNSWIndex::with_defaults(3);
let doc = create_test_document("nan_doc", vec![1.0, f32::NAN, 0.0]);
let result = index.add(doc);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("NaN"),
"Error should mention NaN, got: {}",
err_msg
);
}
#[test]
fn test_add_embedding_rejects_nan() {
let mut index = HNSWIndex::with_defaults(3);
let result = index.add_embedding("nan_vec".into(), vec![f32::NAN, 0.0, 0.0]);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("NaN"),
"Error should mention NaN, got: {}",
err_msg
);
}
#[test]
fn test_add_rejects_all_nan_embedding() {
let mut index = HNSWIndex::with_defaults(3);
let doc = create_test_document("all_nan", vec![f32::NAN, f32::NAN, f32::NAN]);
assert!(index.add(doc).is_err());
}
#[test]
fn test_zero_vector_accepted_and_searchable() {
let mut index = HNSWIndex::with_defaults(3);
let doc_zero = create_test_document("zero", vec![0.0, 0.0, 0.0]);
assert!(index.add(doc_zero).is_ok());
let doc_normal = create_test_document("normal", vec![1.0, 0.0, 0.0]);
assert!(index.add(doc_normal).is_ok());
let query = vec![1.0, 0.0, 0.0];
let results = index.search(&query, 2).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].id, "normal");
}
#[test]
fn test_zero_vector_query_does_not_panic() {
let mut index = HNSWIndex::with_defaults(3);
let doc = create_test_document("doc1", vec![1.0, 0.0, 0.0]);
index.add(doc).unwrap();
let query = vec![0.0, 0.0, 0.0];
let results = index.search(&query, 1).unwrap();
assert_eq!(results.len(), 1);
}
#[test]
fn random_level_never_panics_and_comes_from_the_seeded_stream() {
let draws = |seed: Option<u64>| -> Vec<usize> {
let mut index = HNSWIndex::new(3, HNSWConfig { seed, ..Default::default() });
(0..10_000).map(|_| index.random_level()).collect()
};
let a = draws(Some(7));
assert!(a.iter().all(|&l| l < 64), "exponential decay must not produce absurd levels");
assert!(a.iter().any(|&l| l > 0), "every node landing on layer 0 means ml is not applied");
assert_eq!(a, draws(Some(7)), "a fixed seed must give a reproducible level sequence");
assert_ne!(a, draws(Some(8)), "a different seed must give a different level sequence");
}
#[test]
fn seed_reaches_the_incremental_add_path() {
let vecs: Vec<Vec<f32>> = (0..200)
.map(|i| (0..8).map(|d| ((i * 7 + d * 13) % 40) as f32 * 0.1).collect())
.collect();
let grown = |seed: u64| -> Vec<Vec<u32>> {
let mut ix = HNSWIndex::new(
8,
HNSWConfig { seed: Some(seed), m: 4, m0: 8, ..Default::default() },
);
for (i, v) in vecs.iter().enumerate() {
ix.add_embedding(i.to_string(), v.clone()).unwrap();
}
(0..ix.len())
.map(|i| {
let mut n = ix.get_neighbors_l0(i).to_vec();
n.sort_unstable();
n
})
.collect()
};
assert_eq!(
grown(7),
grown(7),
"an index grown by add() must be reproducible at a fixed seed -- `add` is inherently \
sequential, so unlike the parallel bulk builder it has no thread race to blame"
);
assert_ne!(grown(7), grown(9), "a different seed must give a different graph");
}
#[test]
fn cosine_and_l2_pick_different_neighbors() {
let query = vec![10.0, 0.0];
let docs = [
("far_same_direction", vec![100.0, 0.0]),
("near_off_axis", vec![9.0, 3.0]),
];
let winner = |metric: DistanceMetric| {
let mut index = HNSWIndex::new(
2,
HNSWConfig {
metric,
..Default::default()
},
);
for (id, v) in &docs {
index
.add(Document {
id: (*id).to_string(),
content: String::new(),
embedding: v.clone(),
metadata: None,
})
.unwrap();
}
index.search(&query, 1).unwrap()[0].id.clone()
};
assert_eq!(winner(DistanceMetric::Cosine), "far_same_direction");
assert_eq!(winner(DistanceMetric::L2), "near_off_axis");
}
#[test]
fn l2_scores_are_bounded_and_monotonic() {
let mut index = HNSWIndex::new(
2,
HNSWConfig {
metric: DistanceMetric::L2,
..Default::default()
},
);
for (i, v) in [vec![0.0, 0.0], vec![50.0, 0.0], vec![500.0, 0.0]]
.into_iter()
.enumerate()
{
index
.add(Document {
id: i.to_string(),
content: String::new(),
embedding: v,
metadata: None,
})
.unwrap();
}
let results = index.search(&[0.0, 0.0], 3).unwrap();
assert_eq!(results[0].id, "0", "nearest must come first");
for r in &results {
assert!(
r.score > 0.0 && r.score <= 1.0,
"L2 score {} outside (0, 1]",
r.score
);
}
for w in results.windows(2) {
assert!(w[0].score >= w[1].score, "scores must be descending");
}
}
#[test]
fn cosine_is_still_the_default() {
assert_eq!(HNSWConfig::default().metric, DistanceMetric::Cosine);
assert_eq!(DistanceMetric::default(), DistanceMetric::Cosine);
}
#[test]
fn keep_pruned_connections_controls_graph_density_in_both_builders() {
let mut rng = StdRng::seed_from_u64(11);
let centers: Vec<Vec<f32>> = (0..12)
.map(|_| (0..24).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let embeddings: Vec<Vec<f32>> = (0..600)
.map(|i| {
let c = ¢ers[i % 12];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
for strategy in [BuildStrategy::Sequential, BuildStrategy::Parallel] {
let build = |keep: bool| {
HNSWIndex::build(
embeddings.clone(),
HNSWConfig {
m: 16,
m0: 32,
ef_construction: 100,
keep_pruned_connections: keep,
build_strategy: strategy,
seed: Some(3),
..Default::default()
},
)
};
let dense = build(true).avg_degree_l0();
let sparse = build(false).avg_degree_l0();
assert!(
sparse < dense,
"{strategy:?}: keep_pruned_connections has no effect \
(degree {sparse:.1} with it off vs {dense:.1} with it on) — \
the flag is being ignored and the diversity heuristic's pruning is discarded"
);
}
}
#[test]
fn every_build_strategy_produces_a_searchable_graph() {
let mut rng = StdRng::seed_from_u64(7);
let centers: Vec<Vec<f32>> = (0..8)
.map(|_| (0..16).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let embeddings: Vec<Vec<f32>> = (0..400)
.map(|i| {
let c = ¢ers[i % 8];
c.iter().map(|x| x + rng.random::<f32>() * 0.3).collect()
})
.collect();
for strategy in [BuildStrategy::Sequential, BuildStrategy::Parallel] {
let config = HNSWConfig::default()
.with_build_strategy(strategy)
.with_seed(42)
.with_ef_search(100);
let index = HNSWIndex::build(embeddings.clone(), config);
assert_eq!(
index.len(),
embeddings.len(),
"{strategy:?}: wrong node count"
);
let hits = embeddings
.iter()
.enumerate()
.filter(|(i, e)| {
index
.search(e, 1)
.expect("search")
.first()
.is_some_and(|r| r.id == i.to_string())
})
.count();
let recall = hits as f32 / embeddings.len() as f32;
assert!(
recall > 0.95,
"{strategy:?}: self-retrieval recall {:.1}%, graph is broken",
recall * 100.0,
);
}
}
#[test]
fn rabitq_vec_words_matches_documented_layout() {
assert_eq!(rabitq_bit_words(100), 4);
assert_eq!(vec_words(Storage::RaBitQ, 100), 2 + 4);
assert_eq!(vec_words(Storage::RaBitQ, 128), 2 + 4);
assert_eq!(vec_words(Storage::RaBitQ, 128) * 4, 24);
}
#[test]
fn rabitq_recall_on_clustered_data_with_held_out_queries() {
let mut rng = StdRng::seed_from_u64(2024);
let dim = 32;
let n_clusters = 16;
let per_cluster = 50;
let centers: Vec<Vec<f32>> = (0..n_clusters)
.map(|_| (0..dim).map(|_| rng.random::<f32>() * 20.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..n_clusters * per_cluster)
.map(|i| {
let c = ¢ers[i % n_clusters];
c.iter().map(|x| x + rng.random::<f32>() * 0.8).collect()
})
.collect();
let n_queries = 60;
let queries: Vec<Vec<f32>> = (0..n_queries)
.map(|i| {
let c = ¢ers[i % n_clusters];
c.iter().map(|x| x + rng.random::<f32>() * 0.8).collect()
})
.collect();
let config = HNSWConfig {
metric: DistanceMetric::L2,
m: 16,
m0: 32,
ef_construction: 150,
ef_search: 150,
storage: Storage::RaBitQ,
rerank_candidates: 50,
seed: Some(7),
..Default::default()
};
let index = HNSWIndex::build_parallel(base.clone(), config);
let k = 10;
let mut total_recall = 0.0f32;
for q in &queries {
let mut exact: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| {
let d: f32 = v.iter().zip(q).map(|(a, b)| (a - b) * (a - b)).sum();
(d, i)
})
.collect();
exact.sort_by(|a, b| a.0.total_cmp(&b.0));
let truth: HashSet<usize> = exact.iter().take(k).map(|(_, i)| *i).collect();
let got: HashSet<usize> = index
.search(q, k)
.expect("search")
.into_iter()
.filter_map(|r| r.id.parse::<usize>().ok())
.collect();
total_recall += truth.intersection(&got).count() as f32 / k as f32;
}
let recall = total_recall / n_queries as f32;
assert!(
recall > 0.75,
"Storage::RaBitQ recall@{k} on clustered data = {:.1}% (held-out queries, \
brute-force ground truth) — below floor, traversal metric likely broken",
recall * 100.0,
);
}
#[test]
fn zero_rerank_quantized_builds_on_the_default_strategy_and_still_drops_its_vectors() {
let mut rng = StdRng::seed_from_u64(11);
let dim = 24;
let centers: Vec<Vec<f32>> = (0..8)
.map(|_| (0..dim).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..320)
.map(|i| {
let c = ¢ers[i % 8];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
for storage in [Storage::SQ8, Storage::RaBitQ] {
for build_strategy in [BuildStrategy::Sequential, BuildStrategy::Parallel] {
let config = HNSWConfig {
metric: DistanceMetric::L2,
storage,
rerank_candidates: 0,
seed: Some(4),
build_strategy,
..Default::default()
};
let index = HNSWIndex::build(base.clone(), config);
assert!(
index.full.is_empty(),
"{storage:?}/{build_strategy:?}: rerank_candidates = 0 must still DROP the f32 \
vectors — retaining them would 'fix' the panic by silently ignoring the \
caller's memory request"
);
assert_eq!(
index.rerank_candidates(),
0,
"{storage:?}/{build_strategy:?}: the caller's rerank_candidates must be \
restored after the build"
);
let mut hits = 0;
for (i, q) in base.iter().enumerate().step_by(17) {
let got = index.search(q, 5).expect("search must not panic");
assert_eq!(got.len(), 5);
if got.iter().any(|r| r.id == i.to_string()) {
hits += 1;
}
}
assert!(
hits > 0,
"{storage:?}/{build_strategy:?}: index returns results but finds nothing — \
graph is broken"
);
}
}
}
#[test]
fn raising_the_rerank_pool_on_an_index_that_dropped_its_vectors_is_an_error() {
let mut rng = StdRng::seed_from_u64(77);
let dim = 24;
let centers: Vec<Vec<f32>> = (0..8)
.map(|_| (0..dim).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..320)
.map(|i| {
let c = ¢ers[i % 8];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
let cfg = |rerank| HNSWConfig {
metric: DistanceMetric::L2,
storage: Storage::RaBitQ,
rerank_candidates: rerank,
seed: Some(9),
..Default::default()
};
let mut dropped = HNSWIndex::build(base.clone(), cfg(0));
assert!(
matches!(
dropped.set_rerank_candidates(64),
Err(crate::RagError::FullPrecisionDropped)
),
"raising the rerank pool on a vectors-dropped index must be an error"
);
assert_eq!(dropped.rerank_candidates(), 0, "the refused set must not take effect");
assert!(dropped.set_rerank_candidates(0).is_ok());
let mut kept = HNSWIndex::build(base, cfg(100));
assert!(kept.set_rerank_candidates(64).is_ok());
assert_eq!(kept.rerank_candidates(), 64);
assert!(kept.set_rerank_candidates(0).is_ok());
assert_eq!(kept.rerank_candidates(), 0);
assert!(kept.set_rerank_candidates(200).is_ok());
assert_eq!(kept.rerank_candidates(), 200);
}
#[test]
fn rabitq_zero_rerank_drops_full_precision_vectors_and_does_not_panic() {
let mut rng = StdRng::seed_from_u64(55);
let dim = 24;
let centers: Vec<Vec<f32>> = (0..8)
.map(|_| (0..dim).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..320)
.map(|i| {
let c = ¢ers[i % 8];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
let config = HNSWConfig {
metric: DistanceMetric::L2,
storage: Storage::RaBitQ,
rerank_candidates: 0,
seed: Some(3),
..Default::default()
};
let index = HNSWIndex::build_parallel(base.clone(), config);
assert!(
index.full.is_empty(),
"rerank_candidates = 0 must drop the full-precision side array"
);
for q in base.iter().step_by(37) {
let results = index.search(q, 5).expect("search must not panic");
assert_eq!(results.len(), 5);
}
}
fn nonuniform_norm_clusters(
seed: u64,
dim: usize,
n_clusters: usize,
per_cluster: usize,
n_queries: usize,
) -> (Vec<Vec<f32>>, Vec<Vec<f32>>) {
let mut rng = StdRng::seed_from_u64(seed);
let directions: Vec<Vec<f32>> = (0..n_clusters)
.map(|_| {
let mut v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>() * 2.0 - 1.0).collect();
crate::vector::ops::normalize(&mut v);
v
})
.collect();
let make = |n: usize, rng: &mut StdRng| -> Vec<Vec<f32>> {
(0..n)
.map(|i| {
let dir = &directions[i % n_clusters];
let jittered: Vec<f32> =
dir.iter().map(|x| x + rng.random::<f32>() * 0.05).collect();
let scale = 0.5 + rng.random::<f32>() * 49.5;
jittered.into_iter().map(|x| x * scale).collect()
})
.collect()
};
let base = make(n_clusters * per_cluster, &mut rng);
let queries = make(n_queries, &mut rng);
(base, queries)
}
#[test]
fn nonuniform_norm_fixture_discriminates_cosine_from_l2() {
let (base, queries) = nonuniform_norm_clusters(11, 16, 10, 30, 20);
let build = |metric: DistanceMetric| {
HNSWIndex::build_parallel(
base.clone(),
HNSWConfig {
metric,
ef_construction: 150,
ef_search: 150,
seed: Some(1),
..Default::default()
},
)
};
let cosine_idx = build(DistanceMetric::Cosine);
let l2_idx = build(DistanceMetric::L2);
let disagreements = queries
.iter()
.filter(|q| {
let c = cosine_idx.search(q, 1).unwrap()[0].id.clone();
let l = l2_idx.search(q, 1).unwrap()[0].id.clone();
c != l
})
.count();
assert!(
disagreements * 2 >= queries.len(),
"fixture is not discriminating: cosine and L2 only disagreed on {disagreements}/{} \
queries — a metric mix-up test built on this fixture could pass by accident",
queries.len()
);
}
#[test]
fn sq8_default_metric_is_cosine_not_l2() {
let (base, queries) = nonuniform_norm_clusters(23, 24, 12, 30, 40);
let config = HNSWConfig {
storage: Storage::SQ8,
rerank_candidates: 50,
ef_construction: 150,
ef_search: 150,
seed: Some(2),
..Default::default() };
assert_eq!(
config.metric,
DistanceMetric::Cosine,
"test setup sanity check"
);
let index = HNSWIndex::build_parallel(base.clone(), config);
let k = 10;
let mut cosine_recall = 0.0f32;
let mut l2_recall = 0.0f32;
for q in &queries {
let mut by_cosine: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| (1.0 - crate::vector::simd::cosine_similarity_simd(v, q), i))
.collect();
by_cosine.sort_by(|a, b| a.0.total_cmp(&b.0));
let cosine_truth: HashSet<usize> = by_cosine.iter().take(k).map(|(_, i)| *i).collect();
let mut by_l2: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| (crate::vector::simd::l2_squared_distance_simd(v, q), i))
.collect();
by_l2.sort_by(|a, b| a.0.total_cmp(&b.0));
let l2_truth: HashSet<usize> = by_l2.iter().take(k).map(|(_, i)| *i).collect();
let got: HashSet<usize> = index
.search(q, k)
.expect("search")
.into_iter()
.filter_map(|r| r.id.parse::<usize>().ok())
.collect();
cosine_recall += cosine_truth.intersection(&got).count() as f32 / k as f32;
l2_recall += l2_truth.intersection(&got).count() as f32 / k as f32;
}
cosine_recall /= queries.len() as f32;
l2_recall /= queries.len() as f32;
assert!(
cosine_recall > 0.6,
"Storage::SQ8 with default metric: recall@{k} against COSINE ground truth = \
{:.1}% — the default metric is meant to be cosine",
cosine_recall * 100.0
);
assert!(
cosine_recall > l2_recall + 0.2,
"Storage::SQ8 with default metric answers cosine ({:.1}% recall) no better than \
L2 ({:.1}% recall) — this is the exact shape of the metric-ignoring bug",
cosine_recall * 100.0,
l2_recall * 100.0
);
}
#[test]
fn rabitq_default_metric_is_cosine_not_l2() {
let (base, queries) = nonuniform_norm_clusters(29, 24, 12, 30, 40);
let config = HNSWConfig {
storage: Storage::RaBitQ,
rerank_candidates: 50,
ef_construction: 150,
ef_search: 150,
seed: Some(4),
..Default::default() };
assert_eq!(
config.metric,
DistanceMetric::Cosine,
"test setup sanity check"
);
let index = HNSWIndex::build_parallel(base.clone(), config);
let k = 10;
let mut cosine_recall = 0.0f32;
let mut l2_recall = 0.0f32;
for q in &queries {
let mut by_cosine: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| (1.0 - crate::vector::simd::cosine_similarity_simd(v, q), i))
.collect();
by_cosine.sort_by(|a, b| a.0.total_cmp(&b.0));
let cosine_truth: HashSet<usize> = by_cosine.iter().take(k).map(|(_, i)| *i).collect();
let mut by_l2: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| (crate::vector::simd::l2_squared_distance_simd(v, q), i))
.collect();
by_l2.sort_by(|a, b| a.0.total_cmp(&b.0));
let l2_truth: HashSet<usize> = by_l2.iter().take(k).map(|(_, i)| *i).collect();
let got: HashSet<usize> = index
.search(q, k)
.expect("search")
.into_iter()
.filter_map(|r| r.id.parse::<usize>().ok())
.collect();
cosine_recall += cosine_truth.intersection(&got).count() as f32 / k as f32;
l2_recall += l2_truth.intersection(&got).count() as f32 / k as f32;
}
cosine_recall /= queries.len() as f32;
l2_recall /= queries.len() as f32;
assert!(
cosine_recall > 0.5,
"Storage::RaBitQ with default metric: recall@{k} against COSINE ground truth = \
{:.1}%",
cosine_recall * 100.0
);
assert!(
cosine_recall > l2_recall + 0.2,
"Storage::RaBitQ with default metric answers cosine ({:.1}% recall) no better \
than L2 ({:.1}% recall) — this is the exact shape of the metric-ignoring bug",
cosine_recall * 100.0,
l2_recall * 100.0
);
}
#[test]
fn sq8_cosine_scores_are_bounded_and_monotonic() {
let query = vec![10.0, 0.0, 0.0, 0.0];
let vectors = [
("same_dir", vec![500.0, 0.0, 0.0, 0.0]), ("close_dir", vec![8.0, 2.0, 0.0, 0.0]),
("far_dir", vec![2.0, 8.0, 0.0, 0.0]),
("opposite_dir", vec![-30.0, 0.0, 0.0, 0.0]),
];
let mut index = HNSWIndex::new(
4,
HNSWConfig {
storage: Storage::SQ8,
rerank_candidates: 100,
..Default::default() },
);
index
.train(&vectors.iter().map(|(_, v)| v.clone()).collect::<Vec<_>>())
.unwrap();
for (id, v) in &vectors {
index.add_embedding((*id).to_string(), v.clone()).unwrap();
}
let results = index.search(&query, vectors.len()).unwrap();
assert_eq!(results.len(), vectors.len());
for r in &results {
assert!(
(-1.0..=1.0).contains(&r.score),
"SQ8 cosine score {} for {} outside [-1, 1] — the metric-ignoring bug fed a \
squared-L2 value into the cosine score formula",
r.score,
r.id
);
}
for w in results.windows(2) {
assert!(
w[0].score >= w[1].score,
"scores must be descending: {:?}",
results
);
}
assert_eq!(
results[0].id, "same_dir",
"identical direction must rank first under cosine"
);
assert_eq!(
results.last().unwrap().id,
"opposite_dir",
"opposite direction must rank last under cosine"
);
}
#[test]
fn ef_search_controls_distance_calls() {
let mut rng = StdRng::seed_from_u64(9001);
let centers: Vec<Vec<f32>> = (0..20)
.map(|_| (0..32).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let embeddings: Vec<Vec<f32>> = (0..2000)
.map(|i| {
let c = ¢ers[i % 20];
c.iter().map(|x| x + rng.random::<f32>() * 0.4).collect()
})
.collect();
let mut index = HNSWIndex::build(embeddings.clone(), HNSWConfig::default().with_seed(7));
let query = embeddings[0].clone();
let calls_at = |index: &mut HNSWIndex, ef: usize| -> u64 {
index.set_ef_search(ef);
let mut searcher = index.searcher();
searcher.search(&query, 10).unwrap();
searcher.distance_calls()
};
let low = calls_at(&mut index, 10);
let high = calls_at(&mut index, 800);
assert!(
high > low,
"ef_search has no effect on work done: {high} distance calls at ef=800 vs {low} at \
ef=10 — ef_search is being ignored"
);
}
#[test]
fn ef_construction_controls_graph_quality() {
let mut rng = StdRng::seed_from_u64(3113);
let centers: Vec<Vec<f32>> = (0..16)
.map(|_| (0..24).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..800)
.map(|i| {
let c = ¢ers[i % 16];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
let queries: Vec<Vec<f32>> = (0..60)
.map(|i| {
let c = ¢ers[i % 16];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
let recall_for = |ef_construction: usize| -> f32 {
let index = HNSWIndex::build(
base.clone(),
HNSWConfig {
m: 8,
m0: 16,
ef_construction,
ef_search: 12,
seed: Some(11),
build_strategy: BuildStrategy::Sequential,
..Default::default()
},
);
let k = 10;
let mut total = 0.0f32;
for q in &queries {
let mut exact: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| (1.0 - crate::vector::simd::cosine_similarity_simd(v, q), i))
.collect();
exact.sort_by(|a, b| a.0.total_cmp(&b.0));
let truth: HashSet<usize> = exact.iter().take(k).map(|(_, i)| *i).collect();
let got: HashSet<usize> = index
.search(q, k)
.expect("search")
.into_iter()
.filter_map(|r| r.id.parse::<usize>().ok())
.collect();
total += truth.intersection(&got).count() as f32 / k as f32;
}
total / queries.len() as f32
};
let starved = recall_for(1);
let generous = recall_for(200);
assert!(
generous > starved + 0.1,
"ef_construction has no measurable effect on graph quality: recall@10 = {:.3} at \
ef_construction=1 vs {:.3} at ef_construction=200 — ef_construction is being \
ignored at build time",
starved,
generous
);
}
#[test]
fn use_heuristic_selects_a_different_neighbor_set_than_simple() {
let a = [1.0f32, 0.0];
let b = [1.05f32, 0.0];
let c = [0.0f32, 1.2];
let query = [0.0f32, 0.0];
let build_index = |use_heuristic: bool| -> HNSWIndex {
let config = HNSWConfig {
metric: DistanceMetric::L2,
use_heuristic,
extend_candidates: false,
..Default::default()
};
let mut index = HNSWIndex::new(2, config);
index.push_node(&a); index.push_node(&b); index.push_node(&c); index
};
let heuristic_selected: HashSet<usize> = build_index(true)
.select_neighbors(&[0, 1, 2], &query, 2, 0)
.into_iter()
.collect();
let simple_selected: HashSet<usize> = build_index(false)
.select_neighbors(&[0, 1, 2], &query, 2, 0)
.into_iter()
.collect();
assert_eq!(
heuristic_selected,
HashSet::from([0, 2]),
"Algorithm-4 heuristic should pick the diverse pair {{A, C}}, got \
{heuristic_selected:?}"
);
assert_eq!(
simple_selected,
HashSet::from([0, 1]),
"simple selection should pick the two nearest {{A, B}}, got {simple_selected:?}"
);
assert_ne!(
heuristic_selected, simple_selected,
"use_heuristic has no effect: both configs picked the same neighbours — \
use_heuristic is being ignored"
);
}
#[test]
fn extend_candidates_pulls_in_neighbors_of_candidates() {
let d = [1.0f32, 0.0];
let e = [2.0f32, 0.0];
let query = [0.0f32, 0.0];
let build_index = |extend_candidates: bool| -> HNSWIndex {
let config = HNSWConfig {
metric: DistanceMetric::L2,
use_heuristic: true,
extend_candidates,
keep_pruned_connections: true,
m0: 4,
..Default::default()
};
let mut index = HNSWIndex::new(2, config);
index.push_node(&d); index.push_node(&e); index.l0_push(0, 1);
index
};
let extended = build_index(true).select_neighbors(&[0], &query, 2, 0);
let not_extended = build_index(false).select_neighbors(&[0], &query, 2, 0);
assert_eq!(
not_extended.len(),
1,
"without extend_candidates, only the directly-passed candidate D can be selected, \
got {not_extended:?}"
);
assert_eq!(
extended.len(),
2,
"extend_candidates has no effect: expected D's neighbour E to be pulled into the \
pool and selected alongside D, got {extended:?} — extend_candidates is being \
ignored"
);
}
#[test]
fn build_parallel_returns_original_row_indices_despite_its_shuffle() {
let n = 300;
let base: Vec<Vec<f32>> = (0..n)
.map(|i| {
let mut v = vec![0.0f32; 32];
v[i % 32] = 1.0 + i as f32;
v[(i * 7 + 3) % 32] = 0.5 + (i % 13) as f32;
v
})
.collect();
let ix = HNSWIndex::build(
base.clone(),
HNSWConfig {
metric: DistanceMetric::L2,
seed: Some(4),
build_strategy: BuildStrategy::Parallel,
..Default::default()
},
);
assert_eq!(ix.len(), n, "the build dropped or duplicated rows");
for j in 0..ix.len() {
let claimed: usize = ix.ids[j]
.parse()
.expect("build_parallel labels every node with its original row index");
assert_eq!(
ix.get_embedding(j),
base[claimed].as_slice(),
"node {j} claims to be row {claimed}, but the vector it stores is not row \
{claimed}'s. build_parallel's insertion shuffle has leaked into the ids it \
returns, so every id this index reports is a permutation of the caller's rows."
);
}
let seq = HNSWIndex::build(
base.clone(),
HNSWConfig {
metric: DistanceMetric::L2,
seed: Some(4),
ef_search: 300,
build_strategy: BuildStrategy::Sequential,
..Default::default()
},
);
for i in [0, 7, 42, 199, 292, n - 1] {
assert_eq!(
seq.search(&base[i], 1).unwrap()[0].id,
i.to_string(),
"row {i} queried with its own vector did not come back as itself"
);
}
}
#[test]
fn both_builders_produce_the_same_graph() {
let mut rng = StdRng::seed_from_u64(31337);
let centers: Vec<Vec<f32>> = (0..8)
.map(|_| (0..16).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..600)
.map(|i| {
let c: &Vec<f32> = ¢ers[i % 8];
c.iter().map(|x| x + rng.random::<f32>() * 0.4).collect()
})
.collect();
let avg_degree = |strategy: BuildStrategy, heuristic: bool, extend: bool| -> f64 {
let reps = 3;
(0..reps)
.map(|r| {
let ix = HNSWIndex::build(
base.clone(),
HNSWConfig {
metric: DistanceMetric::L2,
m: 8,
m0: 16,
ef_construction: 16,
seed: Some(21 + r),
use_heuristic: heuristic,
extend_candidates: extend,
keep_pruned_connections: false,
build_strategy: strategy,
..Default::default()
},
);
let total: usize = (0..ix.len()).map(|i| ix.get_neighbors_l0(i).len()).sum();
total as f64 / ix.len() as f64
})
.sum::<f64>()
/ reps as f64
};
for (heuristic, extend) in [(false, false), (true, false), (true, true)] {
let par = avg_degree(BuildStrategy::Parallel, heuristic, extend);
let seq = avg_degree(BuildStrategy::Sequential, heuristic, extend);
assert!(
(par - seq).abs() < 0.5,
"use_heuristic={heuristic} extend_candidates={extend}: the parallel builder \
produced average degree {par:.2} and the sequential one {seq:.2}. They run the \
same algorithm on the same seed, so a gap this size means an adapter is dropping \
an option or computing a distance against the wrong point."
);
}
}
#[test]
fn both_builders_reach_similar_recall() {
let mut rng = StdRng::seed_from_u64(0xEF54);
let dim = 128;
let norm = |v: &mut Vec<f32>| {
let n = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if n > 0.0 {
v.iter_mut().for_each(|x| *x /= n);
}
};
let mut gauss = |rng: &mut StdRng| -> f32 {
let u1 = (rng.random::<f32>()).max(1e-7);
let u2 = rng.random::<f32>();
(-2.0 * u1.ln()).sqrt() * (std::f32::consts::TAU * u2).cos()
};
let unit = |rng: &mut StdRng| -> Vec<f32> {
let mut v: Vec<f32> = (0..dim).map(|_| gauss(rng)).collect();
norm(&mut v);
v
};
let sample_from = |rng: &mut StdRng, centers: &[Vec<f32>], gauss: &mut dyn FnMut(&mut StdRng) -> f32| -> Vec<f32> {
let c = ¢ers[rng.random::<u64>() as usize % centers.len()];
let mut v: Vec<f32> = c.iter().map(|x| x + 0.05 * gauss(rng)).collect();
norm(&mut v);
v
};
let base_centers: Vec<Vec<f32>> = (0..100).map(|_| unit(&mut rng)).collect();
let query_centers: Vec<Vec<f32>> = (0..100).map(|_| unit(&mut rng)).collect();
let base: Vec<Vec<f32>> =
(0..10_000).map(|_| sample_from(&mut rng, &base_centers, &mut gauss)).collect();
let queries: Vec<Vec<f32>> =
(0..200).map(|_| sample_from(&mut rng, &query_centers, &mut gauss)).collect();
const K: usize = 10;
let truth: Vec<Vec<usize>> = queries
.iter()
.map(|q| {
let mut d: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(j, v)| {
(q.iter().zip(v).map(|(a, b)| (a - b).powi(2)).sum::<f32>(), j)
})
.collect();
d.sort_by(|a, b| a.0.total_cmp(&b.0));
d.iter().take(K).map(|(_, j)| *j).collect()
})
.collect();
let recall_of = |strategy: BuildStrategy| -> f32 {
let mut ix = HNSWIndex::build(
base.clone(),
HNSWConfig {
metric: DistanceMetric::L2,
m: 32,
m0: 64,
ef_construction: 200,
ef_search: 40,
seed: Some(7),
build_strategy: strategy,
..Default::default()
},
);
ix.set_ef_search(40);
let mut hit = 0.0f32;
for (qi, q) in queries.iter().enumerate() {
let got: std::collections::HashSet<usize> = ix
.search(q, K)
.unwrap()
.iter()
.filter_map(|r| r.id.parse::<usize>().ok())
.collect();
hit += truth[qi].iter().filter(|t| got.contains(t)).count() as f32 / K as f32;
}
hit / queries.len() as f32
};
let seq = recall_of(BuildStrategy::Sequential);
let par = recall_of(BuildStrategy::Parallel);
assert!(
seq > 0.65 && par > 0.55,
"recall floor not met (seq={seq:.3}, par={par:.3}) — the test is comparing broken \
indexes and would pass vacuously"
);
assert!(
seq - par < 0.05,
"parallel recall {par:.3} trails sequential {seq:.3} by {:.3} at equal ef_search — the \
parallel builder is producing a worse graph (last time: it truncated the candidate \
pool to m0 before the diversity heuristic).",
seq - par
);
}
#[test]
fn use_heuristic_and_extend_candidates_are_honoured_by_both_builders() {
let mut rng = StdRng::seed_from_u64(31337);
let centers: Vec<Vec<f32>> = (0..8)
.map(|_| (0..16).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..600)
.map(|i| {
let c: &Vec<f32> = ¢ers[i % 8];
c.iter().map(|x| x + rng.random::<f32>() * 0.4).collect()
})
.collect();
const M0: usize = 16;
let avg_degree = |heuristic: bool, extend: bool, strategy: BuildStrategy| -> f64 {
let reps = 3;
(0..reps)
.map(|r| {
let ix = HNSWIndex::build(
base.clone(),
HNSWConfig {
metric: DistanceMetric::L2,
m: 8,
m0: M0,
ef_construction: 16,
seed: Some(21 + r),
use_heuristic: heuristic,
extend_candidates: extend,
keep_pruned_connections: false,
build_strategy: strategy,
..Default::default()
},
);
let total: usize = (0..ix.len()).map(|i| ix.get_neighbors_l0(i).len()).sum();
total as f64 / ix.len() as f64
})
.sum::<f64>()
/ reps as f64
};
for strategy in [BuildStrategy::Parallel, BuildStrategy::Sequential] {
let greedy = avg_degree(false, false, strategy);
let heuristic = avg_degree(true, false, strategy);
let extended = avg_degree(true, true, strategy);
let noise = (avg_degree(true, false, strategy) - heuristic).abs();
assert!(
(greedy - M0 as f64).abs() < 0.01,
"{strategy:?}: use_heuristic=false must keep the m0 nearest candidates, filling \
every slot, but average degree was {greedy:.2} of a possible {M0}"
);
assert!(
heuristic < 0.75 * M0 as f64 && (greedy - heuristic) > noise * 10.0,
"{strategy:?}: use_heuristic=true must prune candidates that hide behind an \
accepted neighbour, dropping degree below m0={M0}, but degree was \
{heuristic:.2} vs greedy's {greedy:.2} (build-to-build noise {noise:.3}) — the \
flag is not reaching this builder"
);
assert!(
extended > heuristic * 1.05 && (extended - heuristic) > noise * 10.0,
"{strategy:?}: extend_candidates=true must widen the candidate pool and let the \
heuristic accept more of it, but degree was {extended:.2} vs {heuristic:.2} \
without it (build-to-build noise {noise:.3}) — the flag is not reaching this \
builder"
);
}
}
#[test]
fn seed_gives_reproducible_builds_only_on_the_sequential_builder() {
let base: Vec<Vec<f32>> = (0..600)
.map(|i| {
(0..16)
.map(|d| (((i * 7 + d * 13) % 50) as f32) * 0.1 + (i % 8) as f32)
.collect()
})
.collect();
let graph_of = |strategy: BuildStrategy| -> Vec<Vec<u32>> {
let ix = HNSWIndex::build(
base.clone(),
HNSWConfig {
metric: DistanceMetric::L2,
m: 8,
m0: 16,
ef_construction: 100,
seed: Some(21),
build_strategy: strategy,
..Default::default()
},
);
(0..ix.len())
.map(|i| {
let mut n: Vec<u32> = ix.get_neighbors_l0(i).to_vec();
n.sort_unstable();
n
})
.collect()
};
let seq_a = graph_of(BuildStrategy::Sequential);
let seq_b = graph_of(BuildStrategy::Sequential);
assert_eq!(
seq_a, seq_b,
"Sequential + a fixed seed must be bit-reproducible — that is the whole contract of \
`seed`, and it is the only builder that honours it"
);
let par_a = graph_of(BuildStrategy::Parallel);
let par_b = graph_of(BuildStrategy::Parallel);
let differing = par_a
.iter()
.zip(&par_b)
.filter(|(x, y)| x != y)
.count();
assert!(
differing > 0,
"the parallel builder has become reproducible under a fixed seed ({differing} nodes \
differ). That is GOOD — but `HNSWConfig::seed`'s documentation says it is not, and \
this test exists to catch that contract changing. Update the docs and this assertion."
);
}
#[test]
fn m_caps_upper_layer_degree_in_the_parallel_builder_too() {
let mut rng = StdRng::seed_from_u64(4242);
let centers: Vec<Vec<f32>> = (0..10)
.map(|_| (0..16).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let embeddings: Vec<Vec<f32>> = (0..2000)
.map(|i| {
let c = ¢ers[i % 10];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
let build = |m: usize| {
HNSWIndex::build(
embeddings.clone(),
HNSWConfig {
m,
m0: 64,
ef_construction: 100,
seed: Some(5),
keep_pruned_connections: true, build_strategy: BuildStrategy::Parallel, ..Default::default()
},
)
};
let peak_l1 = |ix: &HNSWIndex| -> (usize, usize) {
let mut peak = 0;
let mut count = 0;
for id in 0..ix.len() {
if ix.connections[id].len() > 1 {
peak = peak.max(ix.connections[id][1].len());
count += 1;
}
}
(peak, count)
};
let (narrow_peak, narrow_count) = peak_l1(&build(4));
let (wide_peak, _) = peak_l1(&build(32));
assert!(
narrow_count > 10,
"fixture put only {narrow_count} nodes on layer 1 — the assertions below would be \
vacuous"
);
assert_eq!(
narrow_peak, 4,
"the m=4 cap never bound (peak layer-1 degree {narrow_peak}) — nothing is pressing \
against the limit, so the invariant below proves nothing"
);
assert!(
narrow_peak <= 4,
"m = 4 but a layer-1 node holds {narrow_peak} neighbours — the parallel builder is \
ignoring config.m (it used to hardcode M_MAX = 32)"
);
assert!(
wide_peak > 4,
"m = 32 produced a peak layer-1 degree of only {wide_peak}, no better than m=4 — \
config.m is not reaching the parallel builder"
);
}
#[test]
fn m_caps_upper_layer_degree_independent_of_m0() {
let mut rng = StdRng::seed_from_u64(5150);
let centers: Vec<Vec<f32>> = (0..12)
.map(|_| (0..16).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let embeddings: Vec<Vec<f32>> = (0..800)
.map(|i| {
let c = ¢ers[i % 12];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
let peak_degree_at = |index: &HNSWIndex, layer: usize| -> (usize, usize) {
let mut peak = 0usize;
let mut count = 0usize;
for node_id in 0..index.len() {
if layer < index.connections[node_id].len() {
peak = peak.max(index.connections[node_id][layer].len());
count += 1;
}
}
(peak, count)
};
let build_with_m = |m: usize| -> HNSWIndex {
HNSWIndex::build(
embeddings.clone(),
HNSWConfig {
m,
m0: 64,
ml: 3.0,
ef_construction: 150,
seed: Some(99),
build_strategy: BuildStrategy::Sequential,
..Default::default()
},
)
};
let narrow = build_with_m(4);
let wide = build_with_m(32);
let (narrow_peak, narrow_count) = peak_degree_at(&narrow, 1);
let (wide_peak, wide_count) = peak_degree_at(&wide, 1);
assert!(
narrow_count > 20 && wide_count > 20,
"fixture put almost nothing on layer 1 ({narrow_count} / {wide_count} nodes) — a \
degree cap cannot bind on an empty layer, so the assertions below would be vacuous"
);
assert_eq!(
narrow_peak, 4,
"the m=4 cap never bound (peak layer-1 degree was {narrow_peak}) — with nothing
pressing against the limit, the over-degree assertion below proves nothing"
);
assert!(
narrow_peak <= 4,
"m = 4 but a layer-1 node holds {narrow_peak} neighbours — config.m is not capping \
upper-layer degree"
);
assert!(
wide_peak > 4,
"m = 32 produced a peak layer-1 degree of only {wide_peak}, no better than m=4 — \
config.m is being ignored during insertion"
);
}
#[test]
fn seed_makes_builds_reproducible_and_distinguishable() {
let mut rng = StdRng::seed_from_u64(2718);
let centers: Vec<Vec<f32>> = (0..10)
.map(|_| (0..12).map(|_| rng.random::<f32>() * 10.0).collect())
.collect();
let embeddings: Vec<Vec<f32>> = (0..400)
.map(|i| {
let c = ¢ers[i % 10];
c.iter().map(|x| x + rng.random::<f32>() * 0.5).collect()
})
.collect();
let build_with_seed = |seed: u64| -> HNSWIndex {
HNSWIndex::build(
embeddings.clone(),
HNSWConfig {
seed: Some(seed),
build_strategy: BuildStrategy::Sequential,
..Default::default()
},
)
};
#[allow(clippy::type_complexity)]
let fingerprint =
|index: &HNSWIndex| -> (Vec<u32>, Vec<Vec<Vec<u32>>>, Option<usize>, usize) {
(
index.nodes.clone(),
index.connections.clone(),
index.entry_point,
index.max_layer,
)
};
let a1 = build_with_seed(42);
let a2 = build_with_seed(42);
let b = build_with_seed(43);
assert_eq!(
fingerprint(&a1),
fingerprint(&a2),
"two builds with the same seed produced different graphs — seed is not being used \
deterministically (or is being ignored in favor of a fresh random seed each time)"
);
assert_ne!(
fingerprint(&a1),
fingerprint(&b),
"two builds with DIFFERENT seeds produced the identical graph — seed is being \
ignored in favor of some fixed internal value"
);
}
#[test]
fn rerank_candidates_nonzero_beats_zero_on_the_same_index() {
let mut rng = StdRng::seed_from_u64(707);
let dim = 16;
let n_clusters = 10;
let per_cluster = 60;
let centers: Vec<Vec<f32>> = (0..n_clusters)
.map(|_| (0..dim).map(|_| rng.random::<f32>() * 20.0).collect())
.collect();
let base: Vec<Vec<f32>> = (0..n_clusters * per_cluster)
.map(|i| {
let c = ¢ers[i % n_clusters];
c.iter().map(|x| x + rng.random::<f32>() * 0.9).collect()
})
.collect();
let n_queries = 50;
let queries: Vec<Vec<f32>> = (0..n_queries)
.map(|i| {
let c = ¢ers[i % n_clusters];
c.iter().map(|x| x + rng.random::<f32>() * 0.9).collect()
})
.collect();
let recall_for = |rerank_candidates: usize| -> f32 {
let config = HNSWConfig {
metric: DistanceMetric::L2,
m: 16,
m0: 32,
ef_construction: 150,
ef_search: 150,
storage: Storage::RaBitQ,
rerank_candidates,
seed: Some(21),
..Default::default()
};
let index = HNSWIndex::build_parallel(base.clone(), config);
let k = 10;
let mut total = 0.0f32;
for q in &queries {
let mut exact: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| {
let d: f32 = v.iter().zip(q).map(|(a, b)| (a - b) * (a - b)).sum();
(d, i)
})
.collect();
exact.sort_by(|a, b| a.0.total_cmp(&b.0));
let truth: HashSet<usize> = exact.iter().take(k).map(|(_, i)| *i).collect();
let got: HashSet<usize> = index
.search(q, k)
.expect("search")
.into_iter()
.filter_map(|r| r.id.parse::<usize>().ok())
.collect();
total += truth.intersection(&got).count() as f32 / k as f32;
}
total / queries.len() as f32
};
let coarse_only = recall_for(0);
let reranked = recall_for(50);
assert!(
reranked > coarse_only,
"rerank_candidates=50 must beat rerank_candidates=0 on the same index — got \
{reranked:.3} vs {coarse_only:.3}. Equal recall means the rerank pool is being \
silently ignored, exactly the shape of the historical PQHNSWConfig bug."
);
}
}