use serde::{Deserialize, Serialize};
use std::cell::RefCell;
use std::cmp::Reverse;
use std::collections::{BTreeMap, BTreeSet, BinaryHeap};
const MAX_LEVEL: usize = 16;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct HnswParams {
pub m: usize,
pub m0: usize,
pub ef_construction: usize,
pub ef_search: usize,
pub prune: Prune,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum Prune {
Own,
#[default]
Both,
}
impl Prune {
fn parse(s: &str) -> Option<Self> {
match s.trim() {
"own" => Some(Self::Own),
"both" => Some(Self::Both),
_ => None,
}
}
}
impl Default for HnswParams {
fn default() -> Self {
Self {
m: 16,
m0: 64,
ef_construction: 200,
ef_search: 400,
prune: Prune::Both,
}
}
}
impl HnswParams {
fn parse(s: &str) -> Option<Self> {
let fields: Vec<&str> = s.split(',').collect();
if fields.len() < 4 || fields.len() > 5 {
return None;
}
let mut num = fields
.iter()
.take(4)
.map(|f| f.trim().parse::<usize>().ok().filter(|&v| v > 0));
let mut next = || num.next().flatten();
let (m, m0, ef_construction, ef_search) = (next()?, next()?, next()?, next()?);
let prune = match fields.get(4) {
Some(f) => Prune::parse(f)?,
None => Prune::default(),
};
Some(Self {
m,
m0,
ef_construction,
ef_search,
prune,
})
}
}
pub fn hnsw_params() -> HnswParams {
static PARAMS: std::sync::OnceLock<HnswParams> = std::sync::OnceLock::new();
*PARAMS.get_or_init(|| {
let Ok(raw) = std::env::var("MUSHROOMDB_HNSW_PARAMS") else {
return HnswParams::default();
};
match HnswParams::parse(&raw) {
Some(p) => p,
None => {
eprintln!(
"mushroomdb: MUSHROOMDB_HNSW_PARAMS={raw:?} is not \
`m,m0,ef_construction,ef_search[,own|both]` with non-zero numbers; \
using the defaults {:?}",
HnswParams::default()
);
HnswParams::default()
}
}
})
}
#[deprecated(note = "read hnsw_params() instead")]
pub const M: usize = 16;
#[deprecated(note = "read hnsw_params() instead")]
pub const M0: usize = 64;
#[deprecated(note = "read hnsw_params() instead")]
pub const EF_CONSTRUCTION: usize = 200;
#[deprecated(note = "read hnsw_params() instead")]
pub const EF_SEARCH: usize = 400;
#[cfg(any(test, feature = "test-hooks"))]
thread_local! {
static HNSW_INSERT_COUNT: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
static HNSW_REMOVE_SCANNED: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
static HNSW_SEARCH_COUNT: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
static HNSW_DIST_EVALS: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
static HNSW_DIST_EVALS_PAIRWISE: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
static HNSW_BEAM_SCRATCH_GROWS: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
}
#[inline]
fn note_insert() {
#[cfg(any(test, feature = "test-hooks"))]
HNSW_INSERT_COUNT.with(|c| c.set(c.get().saturating_add(1)));
}
#[inline]
fn note_remove_scanned(n: usize) {
#[cfg(any(test, feature = "test-hooks"))]
HNSW_REMOVE_SCANNED.with(|c| c.set(c.get().saturating_add(n as u64)));
#[cfg(not(any(test, feature = "test-hooks")))]
let _ = n;
}
#[inline]
fn note_dist() {
#[cfg(any(test, feature = "test-hooks"))]
HNSW_DIST_EVALS.with(|c| c.set(c.get().saturating_add(1)));
}
#[inline]
fn note_dist_pairwise() {
#[cfg(any(test, feature = "test-hooks"))]
HNSW_DIST_EVALS_PAIRWISE.with(|c| c.set(c.get().saturating_add(1)));
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_dist_evals() -> u64 {
HNSW_DIST_EVALS.with(|c| c.get())
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_dist_evals_pairwise() -> u64 {
HNSW_DIST_EVALS_PAIRWISE.with(|c| c.get())
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_dist_evals_reset() {
HNSW_DIST_EVALS.with(|c| c.set(0));
HNSW_DIST_EVALS_PAIRWISE.with(|c| c.set(0));
}
#[inline]
fn note_beam_scratch_grow() {
#[cfg(any(test, feature = "test-hooks"))]
HNSW_BEAM_SCRATCH_GROWS.with(|c| c.set(c.get().saturating_add(1)));
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_beam_scratch_grows() -> u64 {
HNSW_BEAM_SCRATCH_GROWS.with(|c| c.get())
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_beam_scratch_grows_reset() {
HNSW_BEAM_SCRATCH_GROWS.with(|c| c.set(0));
}
#[inline]
fn note_search() {
#[cfg(any(test, feature = "test-hooks"))]
HNSW_SEARCH_COUNT.with(|c| c.set(c.get().saturating_add(1)));
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_remove_scanned() -> u64 {
HNSW_REMOVE_SCANNED.with(|c| c.get())
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_remove_scanned_reset() {
HNSW_REMOVE_SCANNED.with(|c| c.set(0));
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_insert_count() -> u64 {
HNSW_INSERT_COUNT.with(|c| c.get())
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_insert_count_reset() {
HNSW_INSERT_COUNT.with(|c| c.set(0));
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_search_count() -> u64 {
HNSW_SEARCH_COUNT.with(|c| c.get())
}
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
pub fn hnsw_search_count_reset() {
HNSW_SEARCH_COUNT.with(|c| c.set(0));
}
#[inline]
fn splitmix64(x: u64) -> u64 {
let x = x.wrapping_add(0x9E3779B97F4A7C15);
let x = (x ^ (x >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
let x = (x ^ (x >> 27)).wrapping_mul(0x94D049BB133111EB);
x ^ (x >> 31)
}
fn gen_level(base_seed: u64, node_id: u32) -> usize {
let mixed = base_seed ^ (node_id as u64).wrapping_mul(0x9E3779B97F4A7C15);
let rng = splitmix64(mixed);
let bits = (rng >> 11) | 1; let uniform = bits as f64 / (1u64 << 53) as f64;
let ml = 1.0 / (hnsw_params().m as f64).ln();
let level = (-uniform.ln() * ml).floor() as usize;
level.min(MAX_LEVEL)
}
#[derive(Debug, Clone, Copy, PartialEq)]
struct OrdF64(f64);
impl Eq for OrdF64 {}
impl PartialOrd for OrdF64 {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for OrdF64 {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.0
.partial_cmp(&other.0)
.unwrap_or(std::cmp::Ordering::Greater)
}
}
#[derive(Serialize, Deserialize, Clone, Debug, Default)]
struct VecSlab {
dim: usize,
data: Vec<f32>,
}
impl VecSlab {
#[inline]
fn get(&self, slot: u32) -> &[f32] {
if self.dim == 0 {
return &[];
}
let start = slot as usize * self.dim;
self.data.get(start..start + self.dim).unwrap_or(&[])
}
fn put(&mut self, slot: u32, v: &[f64]) -> bool {
if self.dim == 0 {
self.dim = v.len();
}
if self.dim == 0 || v.len() != self.dim {
return false;
}
let start = slot as usize * self.dim;
let end = start + self.dim;
if self.data.len() < end {
self.data.resize(end, 0.0);
}
for (d, s) in self.data[start..end].iter_mut().zip(v) {
*d = *s as f32;
}
true
}
#[inline]
fn floats_for(&self, live: usize) -> usize {
live * self.dim
}
}
#[inline]
fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len(), "a dot product needs equal lengths");
let mut acc = [0.0f32; 8];
let mut ca = a.chunks_exact(8);
let mut cb = b.chunks_exact(8);
for (x, y) in ca.by_ref().zip(cb.by_ref()) {
for i in 0..8 {
acc[i] += x[i] * y[i];
}
}
let tail: f32 = ca
.remainder()
.iter()
.zip(cb.remainder())
.map(|(x, y)| x * y)
.sum();
acc.iter().sum::<f32>() + tail
}
#[inline]
fn dist_slots(slab: &VecSlab, a: u32, b: u32) -> f64 {
note_dist();
note_dist_pairwise();
let dot = dot_f32(slab.get(a), slab.get(b)) as f64;
(1.0 - dot.clamp(-1.0, 1.0)).max(0.0)
}
#[inline]
fn dist_to(slab: &VecSlab, slot: u32, q: &[f32]) -> f64 {
note_dist();
let dot = dot_f32(slab.get(slot), q) as f64;
(1.0 - dot.clamp(-1.0, 1.0)).max(0.0)
}
fn slab_of_decoded(vectors: &[Vec<f64>]) -> (VecSlab, u64) {
let dim = vectors
.iter()
.map(|v| v.len())
.find(|&l| l != 0)
.unwrap_or(0);
let mut slab = VecSlab {
dim,
data: vec![0.0; vectors.len() * dim],
};
if dim == 0 {
return (slab, 0);
}
let mut odd = 0usize;
for (slot, v) in vectors.iter().enumerate() {
if v.len() == dim {
slab.put(slot as u32, v);
continue;
}
if v.is_empty() {
continue; }
odd += 1;
let start = slot * dim;
for (d, s) in slab.data[start..start + dim].iter_mut().zip(v) {
*d = *s as f32;
}
}
if odd > 0 {
eprintln!(
"mushroomdb: HNSW loaded {odd} node(s) whose embedding is not {dim} \
dimensions; they were padded to the index's stride. Re-embed the \
collection with one model — their distances were already meaningless. \
This index will not claim the fast path while they are in it."
);
}
(slab, odd as u64)
}
#[cfg(test)]
#[inline]
fn as_f32(v: &[f64]) -> Vec<f32> {
let mut out = Vec::with_capacity(v.len());
as_f32_into(&mut out, v);
out
}
#[inline]
fn as_f32_into(dst: &mut Vec<f32>, v: &[f64]) {
dst.clear();
dst.extend(v.iter().map(|&x| x as f32));
}
struct BeamScratch {
visited: Vec<bool>,
c_heap: BinaryHeap<Reverse<(OrdF64, u32)>>,
w_heap: BinaryHeap<(OrdF64, u32)>,
}
impl BeamScratch {
const fn empty() -> Self {
Self {
visited: Vec::new(),
c_heap: BinaryHeap::new(),
w_heap: BinaryHeap::new(),
}
}
fn reset_beam(&mut self, n_slots: usize) {
if self.visited.len() == n_slots {
self.visited.fill(false);
} else {
if self.visited.capacity() < n_slots {
note_beam_scratch_grow();
}
self.visited.clear();
self.visited.resize(n_slots, false);
}
self.c_heap.clear();
self.w_heap.clear();
}
}
thread_local! {
static BEAM_SCRATCH: RefCell<BeamScratch> = const { RefCell::new(BeamScratch::empty()) };
static QUERY_F32: RefCell<Vec<f32>> = const { RefCell::new(Vec::new()) };
}
fn with_beam_scratch<R>(f: impl FnOnce(&mut BeamScratch) -> R) -> R {
BEAM_SCRATCH.with(|cell| f(&mut cell.borrow_mut()))
}
fn take_query_f32(v: &[f64]) -> Vec<f32> {
QUERY_F32.with(|cell| {
let mut q = std::mem::take(&mut *cell.borrow_mut());
as_f32_into(&mut q, v);
q
})
}
fn stash_query_f32(q: Vec<f32>) {
QUERY_F32.with(|cell| {
*cell.borrow_mut() = q;
});
}
#[derive(Serialize, Deserialize, Clone, Debug, Default)]
struct HnswNode {
level: usize,
layers: Vec<Vec<u32>>,
}
#[derive(Serialize, Deserialize, Clone, Debug, Default)]
pub struct HnswIndex {
base_seed: u64,
slots: Vec<HnswNode>,
slot_of: BTreeMap<u32, u32>,
id_of: Vec<u32>,
free: Vec<u32>,
slab: VecSlab,
#[serde(skip)]
dim_mismatches: u64,
#[serde(skip)]
refused: BTreeSet<u32>,
#[serde(skip)]
parked: Vec<(u32, Vec<f64>)>,
#[serde(skip)]
back_refs: BTreeMap<u32, BTreeSet<u32>>,
#[serde(skip)]
incomplete: bool,
entry_point: Option<u32>,
max_level: usize,
#[doc(hidden)]
#[cfg(any(test, feature = "test-hooks"))]
#[serde(skip)]
pub prune_override_for_test: Option<Prune>,
}
impl HnswIndex {
#[inline]
fn resolved_prune(&self, params: &HnswParams) -> Prune {
#[cfg(any(test, feature = "test-hooks"))]
if let Some(p) = self.prune_override_for_test {
return p;
}
params.prune
}
}
const DEAD: u32 = u32::MAX;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub struct HnswMemoryStats {
pub live_nodes: usize,
pub neighbour_slots: usize,
pub back_ref_entries: usize,
pub vector_floats: usize,
}
impl HnswMemoryStats {
pub fn bytes_per_node(&self) -> f64 {
if self.live_nodes == 0 {
return 0.0;
}
let bytes = (self.neighbour_slots + self.back_ref_entries) * 4
+ self.vector_floats * std::mem::size_of::<f32>();
bytes as f64 / self.live_nodes as f64
}
pub fn adjacency_bytes_per_node(&self) -> f64 {
if self.live_nodes == 0 {
return 0.0;
}
((self.neighbour_slots + self.back_ref_entries) * 4) as f64 / self.live_nodes as f64
}
}
impl HnswIndex {
pub fn new(base_seed: u64) -> Self {
Self {
base_seed,
..Self::default()
}
}
pub fn len(&self) -> usize {
self.slot_of.len()
}
pub fn contains(&self, id: u32) -> bool {
self.slot_of.contains_key(&id)
}
pub fn is_empty(&self) -> bool {
self.slot_of.is_empty()
}
pub fn can_answer(&self, q_len: usize) -> bool {
!self.is_empty()
&& !self.incomplete
&& self.dim_mismatches == 0
&& self.slab.dim == q_len
&& !self.parked.iter().any(|(_, v)| v.len() == q_len)
}
pub fn mark_complete(&mut self) {
self.incomplete = false;
}
#[doc(hidden)]
pub fn is_incomplete(&self) -> bool {
self.incomplete
}
pub fn node_ids(&self) -> BTreeSet<u32> {
self.slot_of.keys().copied().collect()
}
pub fn accounted_ids(&self) -> BTreeSet<u32> {
let mut out = self.node_ids();
out.extend(self.parked.iter().map(|(id, _)| *id));
out.extend(self.refused.iter().copied());
out
}
pub fn accounts_for(&self, id: u32) -> bool {
self.slot_of.contains_key(&id)
|| self.refused.contains(&id)
|| self.parked.iter().any(|(pid, _)| *pid == id)
}
#[doc(hidden)]
pub fn dim_mismatches(&self) -> u64 {
self.dim_mismatches
}
#[doc(hidden)]
pub fn refused_ids(&self) -> BTreeSet<u32> {
self.refused.clone()
}
fn parked_rows_f32(&self) -> Vec<(u32, Vec<f32>)> {
self.parked
.iter()
.map(|(id, v)| (*id, v.iter().map(|&x| x as f32).collect()))
.collect()
}
fn restore_side_state(
&mut self,
dim_mismatches: u64,
refused: BTreeSet<u32>,
parked: Vec<(u32, Vec<f32>)>,
) {
self.dim_mismatches = dim_mismatches;
self.refused = refused
.into_iter()
.filter(|id| !self.slot_of.contains_key(id))
.collect();
self.parked = parked
.into_iter()
.filter(|(id, v)| {
!self.slot_of.contains_key(id) && (self.slab.dim == 0 || v.len() != self.slab.dim)
})
.map(|(id, v)| (id, v.into_iter().map(f64::from).collect()))
.collect();
}
#[inline]
fn is_live(id_of: &[u32], slot: u32) -> bool {
id_of.get(slot as usize).is_some_and(|&i| i != DEAD)
}
fn alloc_slot(&mut self, id: u32, node: HnswNode, unit: &[f64]) -> u32 {
let slot = match self.free.pop() {
Some(s) => {
self.slots[s as usize] = node;
self.id_of[s as usize] = id;
s
}
None => {
self.slots.push(node);
self.id_of.push(id);
(self.slots.len() - 1) as u32
}
};
let written = self.slab.put(slot, unit);
debug_assert!(written, "alloc_slot was handed a vector the slab refused");
self.slot_of.insert(id, slot);
slot
}
fn set_layer(&mut self, slot: u32, lc: usize, next: Vec<u32>) {
let next_set: BTreeSet<u32> = next.iter().copied().collect();
let prev = std::mem::replace(&mut self.slots[slot as usize].layers[lc], next);
let prev_set: BTreeSet<u32> = prev.into_iter().collect();
for &old in prev_set.difference(&next_set) {
if self.slots[slot as usize]
.layers
.iter()
.any(|l| l.contains(&old))
{
continue;
}
if let Some(refs) = self.back_refs.get_mut(&old) {
refs.remove(&slot);
if refs.is_empty() {
self.back_refs.remove(&old);
}
}
}
for &added in next_set.difference(&prev_set) {
self.back_refs.entry(added).or_default().insert(slot);
}
}
fn link_back_ref(&mut self, target: u32, from: u32) {
self.back_refs.entry(target).or_default().insert(from);
}
fn referrers_of(&self, slot: u32) -> Vec<u32> {
self.back_refs
.get(&slot)
.map(|s| s.iter().copied().collect())
.unwrap_or_default()
}
fn clear_back_refs(&mut self, slot: u32) {
self.back_refs.remove(&slot);
}
fn from_v1(v1: HnswIndexV1) -> Self {
let id_of: Vec<u32> = v1.nodes.keys().copied().collect();
let slot_of: BTreeMap<u32, u32> = id_of
.iter()
.enumerate()
.map(|(s, &id)| (id, s as u32))
.collect();
let mut vectors: Vec<Vec<f64>> = Vec::with_capacity(id_of.len());
let slots: Vec<HnswNode> = v1
.nodes
.into_values()
.map(|n| {
vectors.push(n.vector);
HnswNode {
level: n.level,
layers: n
.layers
.into_iter()
.map(|l| l.iter().filter_map(|id| slot_of.get(id).copied()).collect())
.collect(),
}
})
.collect();
let (slab, odd) = slab_of_decoded(&vectors);
let mut out = Self {
base_seed: v1.base_seed,
entry_point: v1.entry_point.and_then(|e| slot_of.get(&e).copied()),
max_level: v1.max_level,
slab,
dim_mismatches: odd,
slots,
slot_of,
id_of,
free: Vec::new(),
..Self::default()
};
if out.entry_point.is_none() && !out.slot_of.is_empty() {
let (_, &ep) = out
.slot_of
.iter()
.max_by_key(|(_, &s)| out.slots[s as usize].level)
.expect("slot_of is non-empty");
out.entry_point = Some(ep);
out.max_level = out.slots[ep as usize].level;
}
out.rebuild_back_refs();
out
}
fn from_v2(v2: HnswIndexV2) -> Self {
let mut vectors: Vec<Vec<f64>> = Vec::with_capacity(v2.slots.len());
let slots: Vec<HnswNode> = v2
.slots
.into_iter()
.map(|n| {
vectors.push(n.vector);
HnswNode {
level: n.level,
layers: n.layers,
}
})
.collect();
let (slab, odd) = slab_of_decoded(&vectors);
let mut out = Self {
base_seed: v2.base_seed,
slab,
dim_mismatches: odd,
slots,
slot_of: v2.slot_of,
id_of: v2.id_of,
free: v2.free,
entry_point: v2.entry_point,
max_level: v2.max_level,
..Self::default()
};
out.rebuild_back_refs();
out
}
fn rebuild_back_refs(&mut self) {
let mut refs: BTreeMap<u32, BTreeSet<u32>> = BTreeMap::new();
for (s, node) in self.slots.iter().enumerate() {
if !Self::is_live(&self.id_of, s as u32) {
continue;
}
for layer in &node.layers {
for &t in layer {
refs.entry(t).or_default().insert(s as u32);
}
}
}
self.back_refs = refs;
}
fn free_slot(&mut self, slot: u32) {
let id = std::mem::replace(&mut self.id_of[slot as usize], DEAD);
self.slot_of.remove(&id);
self.slots[slot as usize] = HnswNode::default();
self.free.push(slot);
}
fn beam_search(
slots: &[HnswNode],
id_of: &[u32],
slab: &VecSlab,
q: &[f32],
ep: u32,
layer: usize,
ef: usize,
) -> Vec<(u32, f64)> {
with_beam_scratch(|scratch| {
scratch.reset_beam(slots.len());
if let Some(v) = scratch.visited.get_mut(ep as usize) {
*v = true;
}
let ep_dist = dist_to(slab, ep, q);
scratch.c_heap.push(Reverse((OrdF64(ep_dist), ep)));
scratch.w_heap.push((OrdF64(ep_dist), ep));
while let Some(&Reverse((OrdF64(c_dist), c))) = scratch.c_heap.peek() {
let f_dist = scratch
.w_heap
.peek()
.map(|(OrdF64(d), _)| *d)
.unwrap_or(f64::MAX);
if c_dist > f_dist {
break; }
scratch.c_heap.pop();
let neighbors: &[u32] = slots
.get(c as usize)
.and_then(|n| n.layers.get(layer))
.map(|l| l.as_slice())
.unwrap_or_default();
for &e in neighbors {
if scratch.visited.get(e as usize).copied().unwrap_or(true) {
continue;
}
if !Self::is_live(id_of, e) {
continue; }
scratch.visited[e as usize] = true;
let e_dist = dist_to(slab, e, q);
let f_dist = scratch
.w_heap
.peek()
.map(|(OrdF64(d), _)| *d)
.unwrap_or(f64::MAX);
if e_dist < f_dist || scratch.w_heap.len() < ef {
scratch.c_heap.push(Reverse((OrdF64(e_dist), e)));
scratch.w_heap.push((OrdF64(e_dist), e));
if scratch.w_heap.len() > ef {
scratch.w_heap.pop(); }
}
}
}
scratch
.w_heap
.drain()
.map(|(OrdF64(d), s)| (s, d))
.collect()
})
}
fn select_neighbors_first_rejection(
slab: &VecSlab,
id_of: &[u32],
candidates: &[(u32, f64)],
m: usize,
) -> Vec<u32> {
let mut kept: Vec<u32> = Vec::with_capacity(m);
for (i, &(cand, d_base)) in candidates.iter().enumerate() {
if kept.len() >= m {
break;
}
if kept.len() + (candidates.len() - i) <= m {
kept.extend(
candidates[i..]
.iter()
.map(|&(c, _)| c)
.filter(|&c| Self::is_live(id_of, c)),
);
break;
}
if !Self::is_live(id_of, cand) {
continue;
}
let diverse = kept.iter().all(|&k| d_base < dist_slots(slab, k, cand));
if diverse {
kept.push(cand);
}
}
debug_assert!(kept.len() <= m, "the prune must respect its allowance");
kept
}
fn greedy_step(
slots: &[HnswNode],
id_of: &[u32],
slab: &VecSlab,
q: &[f32],
ep: u32,
layer: usize,
) -> u32 {
let mut curr = ep;
let mut curr_dist = dist_to(slab, ep, q);
loop {
let mut improved = false;
let neighbors: &[u32] = slots
.get(curr as usize)
.and_then(|n| n.layers.get(layer))
.map(|l| l.as_slice())
.unwrap_or_default();
for &nb in neighbors {
if !Self::is_live(id_of, nb) {
continue;
}
let d = dist_to(slab, nb, q);
if d < curr_dist {
curr_dist = d;
curr = nb;
improved = true;
}
}
if !improved {
break;
}
}
curr
}
pub fn insert(&mut self, id: u32, v: &[f64]) {
self.parked.retain(|(pid, _)| *pid != id);
self.refused.remove(&id);
let Some(unit) = l2_normalize(v) else {
return; };
if self.slab.dim != 0 && unit.len() != self.slab.dim {
if self.len() <= 1 {
let was = self.slab.dim;
if let Some((&evicted, &slot)) = self.slot_of.iter().next() {
let kept: Vec<f64> = self.slab.get(slot).iter().map(|&x| x as f64).collect();
eprintln!(
"mushroomdb: HNSW re-elected its embedding dimension from {was} to {} \
at node {id}, and parked node {evicted}: the first vector indexed \
set the dimension and was the odd one out. Node {evicted} returns \
to the index if {was} dimensions are elected again; until then \
this index answers {was}-dimension queries through the full scan.",
unit.len()
);
self.remove(evicted);
self.parked.push((evicted, kept));
}
self.slab = VecSlab::default();
self.slab.dim = unit.len();
self.refused.clear();
let mut revive: Vec<(u32, Vec<f64>)> = Vec::new();
self.parked.retain(|entry| {
if entry.1.len() == unit.len() {
revive.push(entry.clone());
false
} else {
true
}
});
for (pid, pv) in revive {
self.insert(pid, &pv);
}
} else {
if self.slot_of.contains_key(&id) {
self.remove(id);
}
self.dim_mismatches += 1;
self.refused.insert(id);
if self.dim_mismatches == 1 {
eprintln!(
"mushroomdb: HNSW skipped node {id}: its embedding has {} dimensions \
and this index holds {}. A mixed-dimension index cannot be \
searched, so this index will now answer through the full scan \
instead; re-embed the collection with one model. Further skips \
on this index are silent.",
unit.len(),
self.slab.dim
);
}
return;
}
}
note_insert();
if self.slot_of.contains_key(&id) {
self.remove(id);
}
let level = gen_level(self.base_seed, id);
let slot = self.alloc_slot(
id,
HnswNode {
level,
layers: vec![vec![]; level + 1],
},
&unit,
);
let Some(ep) = self.entry_point else {
self.entry_point = Some(slot);
self.max_level = level;
return;
};
let q = take_query_f32(&unit);
let params = hnsw_params();
let prune = self.resolved_prune(¶ms);
let max_level = self.max_level;
let mut curr_ep = ep;
for lc in ((level + 1)..=max_level).rev() {
curr_ep = Self::greedy_step(&self.slots, &self.id_of, &self.slab, &q, curr_ep, lc);
}
for lc in (0..=level.min(max_level)).rev() {
let m_lc = if lc == 0 { params.m0 } else { params.m };
let mut candidates = Self::beam_search(
&self.slots,
&self.id_of,
&self.slab,
&q,
curr_ep,
lc,
params.ef_construction,
);
sort_by_distance(&mut candidates);
if let Some(&(nearest, _)) = candidates.first() {
curr_ep = nearest;
}
let neighbors =
Self::select_neighbors_first_rejection(&self.slab, &self.id_of, &candidates, m_lc);
self.set_layer(slot, lc, neighbors.clone());
for &nb in &neighbors {
let nb_node = &mut self.slots[nb as usize];
while nb_node.layers.len() <= lc {
nb_node.layers.push(vec![]);
}
if !nb_node.layers[lc].contains(&slot) {
nb_node.layers[lc].push(slot);
self.link_back_ref(slot, nb);
}
if self.slots[nb as usize].layers[lc].len() <= m_lc {
continue;
}
let current: Vec<u32> = self.slots[nb as usize].layers[lc].clone();
let mut scored: Vec<(u32, f64)> = current
.iter()
.filter(|&&s| Self::is_live(&self.id_of, s))
.map(|&s| (s, dist_slots(&self.slab, s, nb)))
.collect();
sort_by_distance(&mut scored);
let kept: Vec<u32> = match prune {
Prune::Own => scored.iter().take(m_lc).map(|(s, _)| *s).collect(),
Prune::Both => Self::select_neighbors_first_rejection(
&self.slab,
&self.id_of,
&scored,
m_lc,
),
};
self.set_layer(nb, lc, kept);
}
}
if level > max_level {
self.entry_point = Some(slot);
self.max_level = level;
}
stash_query_f32(q);
}
pub fn remove(&mut self, id: u32) {
self.parked.retain(|(pid, _)| *pid != id);
self.refused.remove(&id);
let Some(slot) = self.slot_of.get(&id).copied() else {
return;
};
let mut targets: BTreeSet<u32> = BTreeSet::new();
for layer in &self.slots[slot as usize].layers {
targets.extend(layer.iter().copied());
}
for layer in self.slots[slot as usize].layers.iter_mut() {
layer.clear();
}
for t in targets {
if let Some(refs) = self.back_refs.get_mut(&t) {
refs.remove(&slot);
if refs.is_empty() {
self.back_refs.remove(&t);
}
}
}
let referrers = self.referrers_of(slot);
note_remove_scanned(referrers.len());
for referrer in referrers {
for layer in self.slots[referrer as usize].layers.iter_mut() {
layer.retain(|&x| x != slot);
}
}
self.clear_back_refs(slot);
self.free_slot(slot);
if self.entry_point == Some(slot) {
if self.slot_of.is_empty() {
self.entry_point = None;
self.max_level = 0;
} else {
let (_, &new_ep) = self
.slot_of
.iter()
.max_by_key(|(_, &s)| self.slots[s as usize].level)
.expect("slot_of is non-empty");
self.entry_point = Some(new_ep);
self.max_level = self.slots[new_ep as usize].level;
}
}
}
pub fn search(&self, q: &[f64], k: usize) -> Vec<(u32, f64)> {
self.search_with_ef(q, k, self.ef_for(k))
}
pub fn ef_for(&self, k: usize) -> usize {
k.max(hnsw_params().ef_search)
}
pub fn search_with_ef(&self, q: &[f64], k: usize, ef: usize) -> Vec<(u32, f64)> {
let Some(unit_q) = l2_normalize(q) else {
return vec![];
};
if self.slab.dim == 0 || unit_q.len() != self.slab.dim {
return vec![];
}
let Some(ep) = self.entry_point.filter(|&s| Self::is_live(&self.id_of, s)) else {
return vec![];
};
if k == 0 {
return vec![];
}
note_search();
let ef = ef.max(k);
let unit_q = take_query_f32(&unit_q);
let mut curr_ep = ep;
for lc in (1..=self.max_level).rev() {
curr_ep = Self::greedy_step(&self.slots, &self.id_of, &self.slab, &unit_q, curr_ep, lc);
}
let candidates = Self::beam_search(
&self.slots,
&self.id_of,
&self.slab,
&unit_q,
curr_ep,
0,
ef,
);
let mut results: Vec<(u32, f64)> = candidates
.into_iter()
.map(|(s, dist)| (self.id_of[s as usize], (1.0 - dist).clamp(-1.0, 1.0)))
.collect();
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
results.truncate(k);
stash_query_f32(unit_q);
results
}
pub fn memory_stats(&self) -> HnswMemoryStats {
let mut neighbour_slots = 0usize;
for (s, node) in self.slots.iter().enumerate() {
if !Self::is_live(&self.id_of, s as u32) {
continue;
}
neighbour_slots += node.layers.iter().map(|l| l.len()).sum::<usize>();
}
let live_nodes = self.slot_of.len();
HnswMemoryStats {
live_nodes,
neighbour_slots,
back_ref_entries: self.back_refs.values().map(|s| s.len()).sum(),
vector_floats: self.slab.floats_for(live_nodes),
}
}
#[cfg(any(test, feature = "test-hooks"))]
pub fn back_refs_for_test(&self, id: u32) -> BTreeSet<u32> {
let Some(&slot) = self.slot_of.get(&id) else {
return BTreeSet::new();
};
self.back_refs.get(&slot).cloned().unwrap_or_default()
}
#[cfg(any(test, feature = "test-hooks"))]
pub fn scan_back_refs_for_test(&self, id: u32) -> BTreeSet<u32> {
let Some(&slot) = self.slot_of.get(&id) else {
return BTreeSet::new();
};
let mut out = BTreeSet::new();
for (s, node) in self.slots.iter().enumerate() {
if !Self::is_live(&self.id_of, s as u32) {
continue;
}
if node.layers.iter().any(|l| l.contains(&slot)) {
out.insert(s as u32);
}
}
out
}
#[cfg(any(test, feature = "test-hooks"))]
pub fn slot_capacity_for_test(&self) -> usize {
self.slots.len()
}
}
pub const HNSW_BLOB_MAGIC: [u8; 4] = *b"MHNS";
pub const HNSW_BLOB_VERSION: u16 = 4;
const HNSW_BLOB_HEADER_LEN: usize = 6;
#[derive(Serialize, Deserialize, Clone, Debug)]
pub struct HnswBlob {
pub magic: [u8; 4],
pub version: u16,
pub index: HnswIndex,
pub dim_mismatches: u64,
pub refused: BTreeSet<u32>,
pub parked: Vec<(u32, Vec<f32>)>,
pub complete: bool,
}
#[derive(Serialize)]
struct HnswBlobRef<'a> {
magic: [u8; 4],
version: u16,
index: &'a HnswIndex,
dim_mismatches: u64,
refused: BTreeSet<u32>,
parked: Vec<(u32, Vec<f32>)>,
complete: bool,
}
#[derive(Deserialize)]
struct HnswBlobV3 {
#[allow(dead_code)]
magic: [u8; 4],
#[allow(dead_code)]
version: u16,
index: HnswIndex,
complete: bool,
}
#[derive(Deserialize)]
struct HnswIndexV1 {
base_seed: u64,
nodes: BTreeMap<u32, HnswNodeV1>,
entry_point: Option<u32>,
max_level: usize,
}
#[derive(Deserialize)]
struct HnswNodeV1 {
level: usize,
vector: Vec<f64>,
layers: Vec<Vec<u32>>,
}
#[derive(Deserialize)]
struct HnswIndexV2 {
base_seed: u64,
slots: Vec<HnswNodeV2>,
slot_of: BTreeMap<u32, u32>,
id_of: Vec<u32>,
free: Vec<u32>,
entry_point: Option<u32>,
max_level: usize,
}
#[derive(Deserialize)]
struct HnswNodeV2 {
level: usize,
vector: Vec<f64>,
layers: Vec<Vec<u32>>,
}
#[derive(Deserialize)]
struct HnswBlobV2 {
#[allow(dead_code)]
magic: [u8; 4],
#[allow(dead_code)]
version: u16,
index: HnswIndexV2,
}
pub fn encode_hnsw_blob(index: &HnswIndex, complete: bool) -> Option<Vec<u8>> {
bincode::serialize(&HnswBlobRef {
magic: HNSW_BLOB_MAGIC,
version: HNSW_BLOB_VERSION,
index,
dim_mismatches: index.dim_mismatches(),
refused: index.refused_ids(),
parked: index.parked_rows_f32(),
complete,
})
.ok()
}
pub fn hnsw_blob_complete(blob: &[u8]) -> Option<bool> {
if blob.is_empty() {
return None;
}
if blob.len() >= HNSW_BLOB_HEADER_LEN && blob[..4] == HNSW_BLOB_MAGIC {
let version = u16::from_le_bytes([blob[4], blob[5]]);
return match version {
3 | 4 => {
if blob.len() < HNSW_BLOB_HEADER_LEN + 1 {
return None;
}
match blob[blob.len() - 1] {
0 => Some(false),
1 => Some(true),
_ => None,
}
}
1 | 2 => Some(true),
_ => None,
};
}
Some(true)
}
fn finish_decoded_index(
mut index: HnswIndex,
complete: bool,
version: u16,
) -> Result<HnswIndex, String> {
if index.slab.dim != 0 && index.slab.data.len() < index.slots.len() * index.slab.dim {
return Err(format!(
"HNSW v{version} blob is truncated: the slab holds {} floats, {} slots \
of {} dimensions need {}",
index.slab.data.len(),
index.slots.len(),
index.slab.dim,
index.slots.len() * index.slab.dim
));
}
if index.id_of.len() != index.slots.len() {
return Err(format!(
"HNSW v{version} blob is inconsistent: {} slots against {} id entries",
index.slots.len(),
index.id_of.len()
));
}
index.rebuild_back_refs();
index.incomplete = !complete;
Ok(index)
}
pub fn decode_hnsw_blob(blob: &[u8]) -> Result<HnswIndex, String> {
if blob.is_empty() {
return Err("empty blob".to_string());
}
if blob.len() >= HNSW_BLOB_HEADER_LEN && blob[..4] == HNSW_BLOB_MAGIC {
let version = u16::from_le_bytes([blob[4], blob[5]]);
return match version {
4 => bincode::deserialize::<HnswBlob>(blob)
.map_err(|e| format!("HNSW v4 blob did not decode ({e})"))
.and_then(|b| {
let mut index = finish_decoded_index(b.index, b.complete, 4)?;
index.restore_side_state(b.dim_mismatches, b.refused, b.parked);
Ok(index)
}),
3 => bincode::deserialize::<HnswBlobV3>(blob)
.map_err(|e| format!("HNSW v3 blob did not decode ({e})"))
.and_then(|b| finish_decoded_index(b.index, b.complete, 3)),
2 => bincode::deserialize::<HnswBlobV2>(blob)
.map_err(|e| format!("HNSW v2 blob did not decode ({e})"))
.map(|b| HnswIndex::from_v2(b.index)),
v => Err(format!(
"HNSW blob version {v} is not readable by this build (reads up to \
{HNSW_BLOB_VERSION})"
)),
};
}
match bincode::deserialize::<HnswIndexV1>(blob) {
Ok(v1) => Ok(HnswIndex::from_v1(v1)),
Err(v1_err) => Err(format!(
"not a versioned HNSW blob (magic {:?}) and not a v1 one ({v1_err})",
&blob[..4.min(blob.len())]
)),
}
}
pub fn search_hnsw_blob(blob: &[u8], q: &[f64], k: usize) -> Vec<(u32, f64)> {
let Ok(idx) = decode_hnsw_blob(blob) else {
return vec![];
};
idx.search(q, k)
}
fn sort_by_distance(v: &mut [(u32, f64)]) {
v.sort_by(|a, b| {
a.1.partial_cmp(&b.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
}
pub(crate) fn l2_normalize(v: &[f64]) -> Option<Vec<f64>> {
let norm = v.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm == 0.0 {
return None;
}
Some(v.iter().map(|x| x / norm).collect())
}
#[doc(hidden)]
pub fn make_unit_vecs(n: usize, dim: usize, seed: u64) -> Vec<Vec<f64>> {
let mut state = seed;
(0..n)
.map(|_| {
let raw: Vec<f64> = (0..dim)
.map(|_| {
state = splitmix64(state);
(state as i64 as f64) / (i64::MAX as f64)
})
.collect();
l2_normalize(&raw).unwrap_or_else(|| vec![1.0; dim])
})
.collect()
}
#[doc(hidden)]
pub fn make_clustered_unit_vecs(
clusters: usize,
per_cluster: usize,
dim: usize,
seed: u64,
) -> Vec<Vec<f64>> {
let centres = make_unit_vecs(clusters, dim, seed);
let mut state = seed ^ 0xA5A5_5A5A_1234_9876;
let mut out = Vec::with_capacity(clusters * per_cluster);
for centre in ¢res {
for _ in 0..per_cluster {
let raw: Vec<f64> = centre
.iter()
.map(|c| {
state = splitmix64(state);
c + 0.12 * (state as i64 as f64) / (i64::MAX as f64) / (dim as f64).sqrt()
})
.collect();
out.push(l2_normalize(&raw).unwrap_or_else(|| centre.clone()));
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn exact_knn(vecs: &[Vec<f64>], q: &[f64], k: usize) -> Vec<usize> {
let unit_q = l2_normalize(q).unwrap();
let mut scores: Vec<(usize, f64)> = vecs
.iter()
.enumerate()
.map(|(i, v)| {
let dot: f64 = v.iter().zip(unit_q.iter()).map(|(a, b)| a * b).sum();
(i, dot)
})
.collect();
scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scores.into_iter().take(k).map(|(i, _)| i).collect()
}
#[test]
fn hnsw_recalls_near_duplicate() {
let vecs = make_unit_vecs(200, 32, 0xDEAD_BEEF_1234_5678);
let seed = crate::index::fnv1a_u64(b"test-rule");
let mut idx = HnswIndex::new(seed);
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let mut noisy = vecs[7].clone();
noisy[0] += 1e-4;
noisy[1] -= 1e-4;
let results = idx.search(&noisy, 1);
assert!(!results.is_empty(), "HNSW must return at least one result");
assert_eq!(
results[0].0, 7,
"nearest to vec[7]+noise must be 7, got {} (cos={:.6})",
results[0].0, results[0].1
);
}
#[test]
fn a_single_insert_reports_its_distance_evaluations() {
let vecs = make_unit_vecs(201, 32, 0x0D15_7A17);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"dist-evals"));
for (i, v) in vecs.iter().take(200).enumerate() {
idx.insert(i as u32, v);
}
hnsw_dist_evals_reset();
idx.insert(200, &vecs[200]);
let total = hnsw_dist_evals();
let pairwise = hnsw_dist_evals_pairwise();
assert!(
total > 0,
"an insert into a 200-node index must evaluate distances"
);
assert!(
pairwise > 0,
"the prune compares candidates against each other, so some \
evaluations must be pairwise"
);
assert!(
pairwise <= total,
"pairwise ({pairwise}) is a subset of total ({total})"
);
let mut one = HnswIndex::new(7);
one.insert(0, &vecs[0]);
hnsw_dist_evals_reset();
assert_eq!(one.search(&vecs[1], 5).len(), 1);
assert_eq!(
hnsw_dist_evals(),
1 + one.max_level as u64,
"a search of a one-node graph scores the entry point once per \
descended layer and once in the beam"
);
assert_eq!(
hnsw_dist_evals_pairwise(),
0,
"a search compares the query against nodes, never two nodes"
);
}
#[test]
fn beam_search_scratch_does_not_change_hits() {
let vecs = make_unit_vecs(80, 16, 0x5C12_A7C4);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"scratch-hits"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let q = &vecs[13];
let k = 10;
hnsw_dist_evals_reset();
let baseline = idx.search(q, k);
let baseline_evals = hnsw_dist_evals();
assert_eq!(baseline.len(), k, "fixture must return k hits");
hnsw_dist_evals_reset();
hnsw_beam_scratch_grows_reset();
let again = idx.search(q, k);
let again_evals = hnsw_dist_evals();
assert_eq!(
again, baseline,
"scratch reuse must not change hits or order"
);
assert_eq!(
again_evals, baseline_evals,
"scratch reuse must not change dist-eval count"
);
assert_eq!(
hnsw_beam_scratch_grows(),
0,
"a second search on a warm index must reuse the beam visited buffer"
);
let mut exact: Vec<(u32, f64)> = vecs
.iter()
.enumerate()
.map(|(i, v)| {
let dot: f64 = v.iter().zip(q.iter()).map(|(a, b)| a * b).sum();
(i as u32, dot)
})
.collect();
exact.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
let want: Vec<u32> = exact.iter().take(k).map(|(id, _)| *id).collect();
let got: Vec<u32> = baseline.iter().map(|(id, _)| *id).collect();
assert_eq!(
got, want,
"the beam must return the exact top-k for a graph this small; a \
scratch buffer carrying state between queries would show up here \
and nowhere above"
);
}
#[test]
fn hnsw_empty_returns_empty() {
let idx = HnswIndex::new(42);
assert!(idx.search(&[1.0, 0.0], 5).is_empty());
}
#[test]
fn hnsw_zero_vector_skipped() {
let seed = 1;
let mut idx = HnswIndex::new(seed);
idx.insert(0, &[0.0, 0.0]); idx.insert(1, &[1.0, 0.0]);
assert_eq!(idx.len(), 1);
let r = idx.search(&[1.0, 0.0], 5);
assert_eq!(r.len(), 1);
assert_eq!(r[0].0, 1);
}
#[test]
fn hnsw_remove_works() {
let seed = crate::index::fnv1a_u64(b"rm-test");
let mut idx = HnswIndex::new(seed);
idx.insert(0, &[1.0, 0.0]);
idx.insert(1, &[0.0, 1.0]);
idx.insert(2, &[1.0, 0.0]); idx.remove(0);
let r = idx.search(&[1.0, 0.0], 1);
assert!(!r.is_empty());
assert_eq!(r[0].0, 2, "after removing 0, nearest must be 2");
}
#[test]
fn hnsw_cosine_order_preserved() {
let seed = 99;
let mut idx = HnswIndex::new(seed);
idx.insert(0, &[1.0, 0.0]);
idx.insert(1, &[0.6, 0.8]);
idx.insert(2, &[0.0, 1.0]);
let r = idx.search(&[1.0, 0.0], 3);
assert_eq!(r.len(), 3);
assert!(r[0].1 >= r[1].1);
assert!(r[1].1 >= r[2].1);
assert_eq!(r[0].0, 0, "node 0 must be nearest");
}
#[test]
fn the_index_holds_one_f32_per_dimension() {
const DIM: usize = 64;
let vecs = make_unit_vecs(300, DIM, 0x51AB_1234);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"slab"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let mem = idx.memory_stats();
assert_eq!(mem.live_nodes, 300);
assert_eq!(
mem.vector_floats,
300 * DIM,
"the slab holds one float per dimension per live node"
);
let expected = mem.adjacency_bytes_per_node() + (DIM * 4) as f64;
assert!(
(mem.bytes_per_node() - expected).abs() < 1.0,
"bytes per node is {:.1}, expected {expected:.1} = adjacency + {DIM} f32s",
mem.bytes_per_node()
);
let odd = make_unit_vecs(1, DIM - 1, 0x0DD);
hnsw_insert_count_reset();
idx.insert(1_000, &odd[0]);
assert_eq!(idx.len(), 300, "a 63-D vector must not enter a 64-D index");
assert_eq!(
hnsw_insert_count(),
0,
"a skipped vector must not be counted as indexed"
);
assert!(
idx.search(&odd[0], 5).is_empty(),
"a query of the wrong dimension cannot be answered"
);
let more = make_unit_vecs(1, DIM, 0xF00D);
idx.insert(1_001, &more[0]);
assert_eq!(idx.len(), 301, "a 64-D vector is still accepted");
assert_eq!(hnsw_insert_count(), 1);
}
#[test]
fn a_stray_first_vector_re_elects_the_stride() {
const DIM: usize = 64;
let vecs = make_unit_vecs(50, DIM, 0x5712_A140);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"re-elect"));
idx.insert(900, &[1.0, 2.0, 3.0]);
assert_eq!(idx.len(), 1);
assert!(idx.can_answer(3), "a 3-D index can answer a 3-D query");
assert!(!idx.can_answer(DIM), "...and not a 64-D one");
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
assert_eq!(idx.len(), 50, "every real vector must be indexed");
assert!(
!idx.node_ids().contains(&900),
"the stray must have been evicted, not kept"
);
assert!(
idx.can_answer(DIM),
"an index that re-elected its stride has refused nothing and must \
stay the fast path"
);
assert!(!idx.can_answer(3), "the old stride is gone");
assert_eq!(
idx.search(&vecs[7], 1).first().map(|&(id, _)| id),
Some(7),
"and it must still answer correctly"
);
for &id in idx.node_ids().iter() {
assert_eq!(
idx.back_refs_for_test(id),
idx.scan_back_refs_for_test(id),
"back_refs[{id}] disagrees with a full scan after the eviction"
);
}
for id in idx.node_ids() {
idx.remove(id);
}
let sevens = make_unit_vecs(3, 7, 0x5E7E_0007);
for (i, v) in sevens.iter().enumerate() {
idx.insert(i as u32, v);
}
assert_eq!(idx.len(), 3, "an emptied index re-elects its stride");
assert!(idx.can_answer(7) && !idx.can_answer(DIM));
}
#[test]
fn a_real_vector_evicted_by_a_stray_comes_back() {
const DIM: usize = 8;
let vecs = make_unit_vecs(6, DIM, 0x0E01_C7ED);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"parked"));
idx.insert(0, &vecs[0]);
assert!(idx.can_answer(DIM));
idx.insert(900, &[1.0, 2.0, 3.0]);
assert_eq!(idx.node_ids(), [900].into_iter().collect());
assert!(
!idx.can_answer(DIM),
"while an 8-D vector is parked the index must not answer 8-D queries \
— that is the window in which node 0 is missing"
);
idx.insert(1, &vecs[1]);
assert_eq!(
idx.node_ids(),
[0, 1].into_iter().collect(),
"the vector the stray displaced must be re-inserted"
);
assert!(
idx.can_answer(DIM),
"with nothing of this dimension parked the index is complete again"
);
assert!(
!idx.can_answer(3),
"the stray's dimension is not the stride"
);
for (i, v) in vecs.iter().enumerate().skip(2) {
idx.insert(i as u32, v);
}
assert_eq!(idx.len(), 6, "every real vector is indexed");
assert_eq!(
idx.search(&vecs[0], 1).first().map(|&(id, _)| id),
Some(0),
"and the re-inserted one is reachable"
);
for &id in idx.node_ids().iter() {
assert_eq!(
idx.back_refs_for_test(id),
idx.scan_back_refs_for_test(id),
"back_refs[{id}] disagrees with a full scan after a re-insertion"
);
}
}
#[test]
fn a_removed_node_does_not_return_from_the_parked_list() {
const DIM: usize = 8;
let vecs = make_unit_vecs(3, DIM, 0xDE1E_7ED0);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"parked-rm"));
idx.insert(0, &vecs[0]);
idx.insert(900, &[1.0, 2.0, 3.0]); idx.remove(0); idx.insert(1, &vecs[1]); assert_eq!(
idx.node_ids(),
[1].into_iter().collect(),
"a removed node must not be revived by a re-election"
);
assert!(idx.can_answer(DIM));
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"parked-stale"));
idx.insert(0, &vecs[0]); idx.insert(900, &[1.0, 2.0, 3.0]); idx.insert(0, &[7.0, 8.0, 9.0]); assert_eq!(idx.len(), 2, "nodes 900 and 0, both at the 3-D stride");
idx.remove(900); idx.insert(2, &vecs[2]); assert_eq!(
idx.node_ids(),
[2].into_iter().collect(),
"node 0's stale 8-D copy must not be revived — its vector is 3-D now"
);
assert!(
idx.can_answer(8),
"nothing of 8 dimensions is parked, so the index is complete at its \
stride"
);
}
#[test]
fn an_index_that_refused_a_vector_will_not_claim_to_answer() {
const DIM: usize = 64;
let vecs = make_unit_vecs(30, DIM, 0xBADD_14E0);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"refused"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
assert!(idx.can_answer(DIM), "a clean index answers");
hnsw_insert_count_reset();
idx.insert(900, &[1.0, 2.0, 3.0]);
assert_eq!(idx.len(), 30, "the stray must not be indexed");
assert_eq!(hnsw_insert_count(), 0, "nor counted");
assert!(
!idx.can_answer(DIM),
"an index that refused a vector must send the caller to its scan"
);
assert!(
!idx.search(&vecs[3], 5).is_empty(),
"`search` itself still answers — the fallback is the caller's \
decision, taken on `can_answer`"
);
}
#[test]
fn the_dot_kernel_agrees_with_an_f64_reference() {
for dim in [1usize, 7, 8, 15, 64, 1_536] {
let vecs = make_unit_vecs(6, dim, 0x4047_0000 ^ dim as u64);
for i in 0..vecs.len() {
let a32: Vec<f32> = vecs[i].iter().map(|&x| x as f32).collect();
let self_dot = dot_f32(&a32, &a32);
assert!(
(self_dot as f64 - 1.0).abs() < 1e-5,
"dim {dim}: a unit vector dotted with itself is {self_dot}, not 1.0"
);
for j in 0..vecs.len() {
let b32: Vec<f32> = vecs[j].iter().map(|&x| x as f32).collect();
let reference: f64 =
vecs[i].iter().zip(vecs[j].iter()).map(|(a, b)| a * b).sum();
let got = dot_f32(&a32, &b32) as f64;
assert!(
(got - reference).abs() < 2e-5,
"dim {dim}, pair ({i},{j}): kernel {got} against reference \
{reference}"
);
}
}
}
}
#[test]
#[allow(deprecated)]
fn the_deprecated_constants_still_name_the_default_shape() {
let d = HnswParams::default();
assert_eq!(
(M, M0, EF_CONSTRUCTION, EF_SEARCH),
(d.m, d.m0, d.ef_construction, d.ef_search)
);
}
#[test]
fn params_parse_accepts_the_documented_format_and_nothing_else() {
assert_eq!(
HnswParams::parse("16,64,200,400"),
Some(HnswParams::default())
);
assert_eq!(
HnswParams::parse(" 8 , 64 , 300 , 96 "),
Some(HnswParams {
m: 8,
m0: 64,
ef_construction: 300,
ef_search: 96,
prune: Prune::Both
}),
"whitespace around a field must not defeat the override"
);
assert_eq!(
HnswParams::parse("16,64,200,400,both"),
Some(HnswParams::default()),
"the prune field is optional, and `both` is what omitting it means"
);
assert_eq!(
HnswParams::parse("16,64,200,400, own "),
Some(HnswParams {
prune: Prune::Own,
..HnswParams::default()
}),
"`own` opts out of the neighbour-side heuristic"
);
for bad in [
"",
"16",
"16,64,200",
"16,64,200,400,64",
"16,64,200,400,neither",
"16,64,200,400,own,own",
"16,64,200,x",
"0,64,200,400",
"16,0,200,400",
"16,64,0,400",
"16,64,200,0",
"-16,64,200,400",
] {
assert_eq!(HnswParams::parse(bad), None, "{bad:?} must not parse");
}
}
#[test]
fn both_prune_strategies_build_a_searchable_bounded_graph() {
let params = hnsw_params();
let vecs = make_unit_vecs(600, 24, 0x9121_5EED);
for prune in [Prune::Own, Prune::Both] {
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"prune-shapes"));
idx.prune_override_for_test = Some(prune);
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
assert_eq!(idx.len(), 600, "{prune:?} lost nodes");
for (s, node) in idx.slots.iter().enumerate() {
if !HnswIndex::is_live(&idx.id_of, s as u32) {
continue;
}
for (lc, layer) in node.layers.iter().enumerate() {
let allowed = if lc == 0 { params.m0 } else { params.m };
assert!(
layer.len() <= allowed,
"{prune:?} slot {s} layer {lc} holds {} links for an allowance of \
{allowed}",
layer.len()
);
}
}
let hits = idx.search(&vecs[42], 5);
assert_eq!(
hits.first().map(|&(id, _)| id),
Some(42),
"{prune:?} search"
);
for &id in idx.node_ids().iter() {
assert_eq!(
idx.back_refs_for_test(id),
idx.scan_back_refs_for_test(id),
"{prune:?} back_refs[{id}] disagrees with a full scan"
);
}
}
}
fn slab_of(vecs: &[Vec<f64>]) -> (VecSlab, Vec<u32>) {
let mut slab = VecSlab::default();
for (s, v) in vecs.iter().enumerate() {
assert!(slab.put(s as u32, &l2_normalize(v).unwrap()));
}
let id_of = (0..vecs.len() as u32).collect();
(slab, id_of)
}
#[test]
fn the_prune_prefers_a_new_direction_over_a_redundant_neighbour() {
let base = l2_normalize(&[1.0, 0.0, 0.0]).unwrap();
let vecs = vec![
vec![0.80, 0.0, 0.60], vec![0.78, 0.0, 0.63], vec![0.76, 0.0, 0.65], vec![0.74, 0.0, 0.67], vec![0.70, 0.71, 0.0], vec![0.60, 0.0, 0.80], ];
let (slab, id_of) = slab_of(&vecs);
let base = as_f32(&base);
let mut cands: Vec<(u32, f64)> = (0..6u32).map(|s| (s, dist_to(&slab, s, &base))).collect();
sort_by_distance(&mut cands);
assert_eq!(
cands.iter().map(|&(s, _)| s).collect::<Vec<_>>(),
vec![0, 1, 2, 3, 4, 5],
"fixture: candidates must arrive in this nearest-first order"
);
let kept = HnswIndex::select_neighbors_first_rejection(&slab, &id_of, &cands, 2);
assert_eq!(
kept,
vec![0, 4],
"the prune must reject 1, 2 and 3 as reachable through 0 and keep 4, \
which opens a direction 0 does not cover"
);
}
#[test]
fn the_prune_keeps_the_untested_tail_where_algorithm_4_would_backfill() {
let base = l2_normalize(&[1.0, 0.0, 0.0]).unwrap();
let vecs = vec![
vec![0.80, 0.0, 0.60], vec![0.78, 0.0, 0.63], vec![0.55, 0.0, 0.84], ];
let (slab, id_of) = slab_of(&vecs);
let base = as_f32(&base);
let mut cands: Vec<(u32, f64)> = (0..3u32).map(|s| (s, dist_to(&slab, s, &base))).collect();
sort_by_distance(&mut cands);
assert_eq!(
cands.iter().map(|&(s, _)| s).collect::<Vec<_>>(),
vec![0, 1, 2],
"fixture: nearest-first order"
);
let kept = HnswIndex::select_neighbors_first_rejection(&slab, &id_of, &cands, 2);
assert_eq!(
kept,
vec![0, 2],
"the short-cut keeps the untested far candidate; full Algorithm 4 \
would have rejected it too and backfilled with the nearer reject 1, \
answering [0, 1]. If this ever reads [0, 1] the prune has become \
Algorithm 4 and `select_neighbors_first_rejection`'s name, its doc \
comment and docs/site/rules.md are all now wrong."
);
}
#[test]
fn the_prune_never_keeps_a_freed_slot() {
let vecs = make_unit_vecs(4, 8, 0x7EA0_1234);
let (slab, mut id_of) = slab_of(&vecs);
id_of[1] = DEAD;
let base = as_f32(&l2_normalize(&vecs[0]).unwrap());
let mut cands: Vec<(u32, f64)> = (0..4u32).map(|s| (s, dist_to(&slab, s, &base))).collect();
sort_by_distance(&mut cands);
let kept = HnswIndex::select_neighbors_first_rejection(&slab, &id_of, &cands, 4);
assert!(
!kept.contains(&1),
"a freed slot must not become a neighbour"
);
}
#[test]
fn degree_and_adjacency_bytes_stay_within_the_shape() {
let params = hnsw_params();
let vecs = make_unit_vecs(1_200, 32, 0xDE6E_E5EE);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"degree-bound"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let mut worst_layer0 = 0usize;
let mut worst_upper = 0usize;
for (s, node) in idx.slots.iter().enumerate() {
if !HnswIndex::is_live(&idx.id_of, s as u32) {
continue;
}
for (lc, layer) in node.layers.iter().enumerate() {
let allowed = if lc == 0 { params.m0 } else { params.m };
assert!(
layer.len() <= allowed,
"slot {s} layer {lc} holds {} links for an allowance of {allowed}",
layer.len()
);
if lc == 0 {
worst_layer0 = worst_layer0.max(layer.len());
} else {
worst_upper = worst_upper.max(layer.len());
}
}
}
let mem = idx.memory_stats();
let ceiling = 2.0 * (params.m0 + params.m) as f64 * 4.0;
eprintln!(
"degree: layer0 <= {worst_layer0} (allowance {}), upper <= {worst_upper} \
(allowance {}); adjacency {:.1} B/node against a ceiling of {ceiling:.1}; \
vector {:.1} B/node",
params.m0,
params.m,
mem.adjacency_bytes_per_node(),
(mem.vector_floats * 4) as f64 / mem.live_nodes as f64,
);
assert_eq!(mem.live_nodes, 1_200);
assert!(
mem.back_ref_entries <= mem.neighbour_slots,
"the reverse index ({}) cannot hold more pairs than the forward one ({})",
mem.back_ref_entries,
mem.neighbour_slots
);
assert!(
mem.adjacency_bytes_per_node() <= ceiling,
"adjacency is {:.1} B/node against a ceiling of {ceiling:.1} — the \
index shape grew",
mem.adjacency_bytes_per_node()
);
assert_eq!(mem.vector_floats, 1_200 * 32);
}
#[test]
#[ignore = "slow: builds a 4,800-vector clustered index"]
fn clustered_recall_survives_clusters_wider_than_m0() {
const CLUSTERS: usize = 40;
const PER_CLUSTER: usize = 120;
const DIM: usize = 128;
const K: usize = 10;
assert!(
PER_CLUSTER > hnsw_params().m0,
"the fixture only proves anything when a cluster is wider than m0 \
({PER_CLUSTER} vs {})",
hnsw_params().m0
);
let vecs = make_clustered_unit_vecs(CLUSTERS, PER_CLUSTER, DIM, 0xC1_05_7E_12_34_56_78_9A);
let queries: Vec<Vec<f64>> = (0..CLUSTERS)
.map(|c| {
let j = c * PER_CLUSTER + 17;
let mut q = vecs[j].clone();
q[0] += 1e-6;
q
})
.collect();
let exact: Vec<BTreeSet<usize>> = queries
.iter()
.map(|q| exact_knn(&vecs, q, K).into_iter().collect())
.collect();
let mut scored: Vec<(Prune, f64, f64)> = Vec::new();
for prune in [Prune::Own, Prune::Both] {
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"clustered-recall"));
idx.prune_override_for_test = Some(prune);
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let recalls: Vec<f64> = queries
.iter()
.zip(exact.iter())
.map(|(q, truth)| {
let found = idx
.search(q, K)
.into_iter()
.filter(|(id, _)| truth.contains(&(*id as usize)))
.count();
found as f64 / K as f64
})
.collect();
let min = recalls.iter().cloned().fold(f64::MAX, f64::min);
let mean = recalls.iter().sum::<f64>() / recalls.len() as f64;
eprintln!(
"clustered recall@{K} ({CLUSTERS}x{PER_CLUSTER}, dim {DIM}, m0={}, \
prune={prune:?}): min={min:.4} mean={mean:.4}",
hnsw_params().m0
);
scored.push((prune, min, mean));
}
for &(prune, min, mean) in &scored {
let (min_floor, mean_floor) = match prune {
Prune::Own => (0.40, 0.90),
Prune::Both => (0.70, 0.95),
};
assert!(
min >= min_floor,
"{prune:?} min clustered recall@{K} = {min:.4} < {min_floor}"
);
assert!(
mean >= mean_floor,
"{prune:?} mean clustered recall@{K} = {mean:.4} < {mean_floor}"
);
}
let own = scored[0];
let both = scored[1];
assert!(
both.1 >= own.1 && both.2 >= own.2,
"Prune::Both (min {:.4} mean {:.4}) is not better than Prune::Own \
(min {:.4} mean {:.4}) on clusters wider than m0 — the option has no \
justification left",
both.1,
both.2,
own.1,
own.2
);
}
#[test]
#[ignore = "slow: builds a 5k x 1536-D index and churns it"]
fn recall_survives_insert_remove_churn() {
const N: usize = 5_000;
const DIM: usize = 1_536;
const K: usize = 10;
let vecs = make_unit_vecs(N, DIM, 0xC0FF_EE00_5EED_1234);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"churn-recall"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let mut live: BTreeSet<usize> = (0..N).collect();
for i in (0..N).step_by(5) {
idx.remove(i as u32);
live.remove(&i);
}
for i in (0..N).step_by(10) {
idx.insert(i as u32, &vecs[i]);
live.insert(i);
}
for i in (3..N).step_by(7) {
idx.remove(i as u32);
live.remove(&i);
}
assert_eq!(
idx.len(),
live.len(),
"the index and the oracle disagree on size"
);
for &id in idx.node_ids().iter().step_by(53) {
assert_eq!(
idx.back_refs_for_test(id),
idx.scan_back_refs_for_test(id),
"back_refs[{id}] disagrees with a full scan after churn"
);
}
let survivors: Vec<usize> = live.iter().copied().collect();
let queries = make_unit_vecs(40, DIM, 0x9111_0BED);
let mut recalls = Vec::with_capacity(queries.len());
for q in &queries {
let mut scored: Vec<(usize, f64)> = survivors
.iter()
.map(|&i| {
let dot: f64 = vecs[i].iter().zip(q.iter()).map(|(a, b)| a * b).sum();
(i, dot)
})
.collect();
scored.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
let exact: BTreeSet<usize> = scored.into_iter().take(K).map(|(i, _)| i).collect();
let hits = idx.search(q, K);
for (id, _) in &hits {
assert!(
live.contains(&(*id as usize)),
"search returned {id}, which was removed"
);
}
let found = hits
.iter()
.filter(|(id, _)| exact.contains(&(*id as usize)))
.count();
recalls.push(found as f64 / K as f64);
}
let min = recalls.iter().cloned().fold(f64::MAX, f64::min);
let mean = recalls.iter().sum::<f64>() / recalls.len() as f64;
eprintln!(
"churn recall@{K}: min={min:.4} mean={mean:.4} over {} survivors",
live.len()
);
assert!(min >= 0.90, "min recall@{K} after churn = {min:.4} < 0.90");
assert!(
mean >= 0.95,
"mean recall@{K} after churn = {mean:.4} < 0.95"
);
}
#[test]
fn back_refs_match_a_full_scan() {
let vecs = make_unit_vecs(500, 32, 0x5EED_1234);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"backrefs"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
for i in (0..500).step_by(5) {
idx.remove(i as u32);
}
for i in (0..500).step_by(5) {
idx.insert(i as u32, &vecs[i]);
}
assert_eq!(idx.len(), 500, "every node is back in the index");
for &id in idx.node_ids().iter() {
assert_eq!(
idx.back_refs_for_test(id),
idx.scan_back_refs_for_test(id),
"back_refs[{id}] disagrees with a full scan"
);
}
assert_eq!(
idx.slot_capacity_for_test(),
500,
"a remove + re-insert must recycle the slot"
);
}
#[test]
fn insert_then_remove_equals_never_inserted() {
let vecs = make_unit_vecs(200, 16, 0xABCD_0001);
let seed = crate::index::fnv1a_u64(b"rm-equiv");
let skipped = [42usize, 77, 150];
let mut reference = HnswIndex::new(seed);
for (i, v) in vecs.iter().enumerate() {
if skipped.contains(&i) {
continue;
}
reference.insert(i as u32, v);
}
let mut subject = HnswIndex::new(seed);
for (i, v) in vecs.iter().enumerate() {
subject.insert(i as u32, v);
}
for &i in &skipped {
subject.remove(i as u32);
}
assert_eq!(
subject.node_ids(),
reference.node_ids(),
"the removed ids must be gone from the index"
);
for &id in subject.node_ids().iter() {
assert_eq!(
subject.back_refs_for_test(id),
subject.scan_back_refs_for_test(id),
"back_refs[{id}] disagrees with a full scan after the removals"
);
}
for &i in &skipped {
let hits = subject.search(&vecs[i], 20);
assert!(
!hits.iter().any(|&(id, _)| id == i as u32),
"removed id {i} is still reachable through the graph"
);
}
}
#[test]
fn removing_the_entry_point_re_elects_and_still_searches() {
let vecs = make_unit_vecs(300, 16, 0xE47E_9001);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"ep-churn"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let mut seen_levels = Vec::new();
for _ in 0..12 {
let ep_slot = idx
.entry_point
.expect("a non-empty index has an entry point");
let ep_id = idx.id_of[ep_slot as usize];
assert_ne!(ep_id, DEAD, "the entry point must name a live node");
seen_levels.push(idx.max_level);
idx.remove(ep_id);
let new_ep = idx.entry_point.expect("re-election must find a live node");
assert!(
HnswIndex::is_live(&idx.id_of, new_ep),
"the re-elected entry point is a freed slot"
);
assert_ne!(new_ep, ep_slot, "the removed slot is still the entry point");
assert_eq!(
idx.max_level, idx.slots[new_ep as usize].level,
"max_level must follow the re-elected entry point"
);
let highest = idx
.slot_of
.values()
.map(|&s| idx.slots[s as usize].level)
.max()
.unwrap();
assert_eq!(
idx.max_level, highest,
"the entry point must be a highest-level node"
);
assert!(
!idx.node_ids().contains(&ep_id),
"the old entry point lingers"
);
let hits = idx.search(&vecs[200], 5);
assert!(
!hits.is_empty(),
"the graph stopped answering after re-election"
);
for (id, _) in &hits {
assert!(
idx.node_ids().contains(id),
"search returned {id}, which is not in the index"
);
}
for &id in idx.node_ids().iter() {
assert_eq!(
idx.back_refs_for_test(id),
idx.scan_back_refs_for_test(id),
"back_refs[{id}] disagrees with a full scan after an entry-point removal"
);
}
}
assert!(
seen_levels.iter().any(|&l| l > 0),
"the fixture never had a multi-layer entry point, so this proved nothing"
);
for id in idx.node_ids() {
idx.remove(id);
}
assert!(idx.is_empty());
assert_eq!(idx.entry_point, None);
assert_eq!(idx.max_level, 0);
assert!(idx.search(&vecs[0], 5).is_empty());
}
#[test]
fn remove_touches_only_the_nodes_that_list_it() {
let vecs = make_unit_vecs(1_500, 32, 0xD00D_0007);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"rm-cost"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let mut worst = 0u64;
for i in (0..1_500).step_by(100) {
let in_degree = idx.back_refs_for_test(i as u32).len() as u64;
hnsw_remove_scanned_reset();
idx.remove(i as u32);
let scanned = hnsw_remove_scanned();
assert_eq!(
scanned, in_degree,
"removing {i} visited {scanned} nodes for an in-degree of {in_degree}"
);
worst = worst.max(scanned);
}
let live = idx.len() as u64;
assert!(
worst * 2 < live,
"worst removal touched {worst} of {live} live nodes — the O(N·M₀) scan \
is still there"
);
}
#[test]
#[ignore = "slow: builds a 5k index"]
fn remove_is_not_a_full_scan() {
let vecs = make_unit_vecs(5_000, 64, 0xD00D);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"rm-cost"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let t = std::time::Instant::now();
for i in (0..5_000).step_by(100) {
idx.remove(i as u32);
}
let mean = t.elapsed() / 50;
eprintln!("mean removal at n=5000: {mean:?}");
assert!(
mean < std::time::Duration::from_millis(5),
"mean removal {mean:?} at n=5000 — the O(N·M₀) scan is still there"
);
}
#[derive(Serialize)]
struct V1Node {
level: usize,
vector: Vec<f64>,
layers: Vec<Vec<u32>>,
}
#[derive(Serialize)]
struct V1Index {
base_seed: u64,
nodes: BTreeMap<u32, V1Node>,
entry_point: Option<u32>,
max_level: usize,
}
#[derive(Serialize)]
struct V2Node {
level: usize,
vector: Vec<f64>,
layers: Vec<Vec<u32>>,
}
#[derive(Serialize)]
struct V2Index {
base_seed: u64,
slots: Vec<V2Node>,
slot_of: BTreeMap<u32, u32>,
id_of: Vec<u32>,
free: Vec<u32>,
entry_point: Option<u32>,
max_level: usize,
}
#[derive(Serialize)]
struct V2Blob {
magic: [u8; 4],
version: u16,
index: V2Index,
}
fn vector_of(idx: &HnswIndex, slot: u32) -> Vec<f64> {
idx.slab.get(slot).iter().map(|&x| x as f64).collect()
}
fn as_v1_blob(idx: &HnswIndex) -> Vec<u8> {
as_v1_blob_with(idx, &[])
}
fn as_v1_blob_with(idx: &HnswIndex, by_id: &[Vec<f64>]) -> Vec<u8> {
let nodes: BTreeMap<u32, V1Node> = idx
.slot_of
.iter()
.map(|(&id, &s)| {
let n = &idx.slots[s as usize];
(
id,
V1Node {
level: n.level,
vector: by_id
.get(id as usize)
.cloned()
.unwrap_or_else(|| vector_of(idx, s)),
layers: n
.layers
.iter()
.map(|l| l.iter().map(|&t| idx.id_of[t as usize]).collect())
.collect(),
},
)
})
.collect();
bincode::serialize(&V1Index {
base_seed: idx.base_seed,
nodes,
entry_point: idx.entry_point.map(|s| idx.id_of[s as usize]),
max_level: idx.max_level,
})
.unwrap()
}
fn as_v2_blob(idx: &HnswIndex) -> Vec<u8> {
as_v2_blob_with(idx, &[])
}
fn as_v2_blob_with(idx: &HnswIndex, by_id: &[Vec<f64>]) -> Vec<u8> {
let slots: Vec<V2Node> = idx
.slots
.iter()
.enumerate()
.map(|(s, n)| V2Node {
level: n.level,
vector: by_id
.get(*idx.id_of.get(s).unwrap_or(&DEAD) as usize)
.cloned()
.unwrap_or_else(|| vector_of(idx, s as u32)),
layers: n.layers.clone(),
})
.collect();
bincode::serialize(&V2Blob {
magic: HNSW_BLOB_MAGIC,
version: 2,
index: V2Index {
base_seed: idx.base_seed,
slots,
slot_of: idx.slot_of.clone(),
id_of: idx.id_of.clone(),
free: idx.free.clone(),
entry_point: idx.entry_point,
max_level: idx.max_level,
},
})
.unwrap()
}
fn as_v3_blob(idx: &HnswIndex, complete: bool) -> Vec<u8> {
#[derive(Serialize)]
struct V3Ref<'a> {
magic: [u8; 4],
version: u16,
index: &'a HnswIndex,
complete: bool,
}
bincode::serialize(&V3Ref {
magic: HNSW_BLOB_MAGIC,
version: 3,
index: idx,
complete,
})
.expect("v3 encode")
}
fn blob_fixture() -> (Vec<Vec<f64>>, HnswIndex) {
let vecs = make_unit_vecs(120, 24, 0x0B10_B0B0);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"blob-rt"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
(vecs, idx)
}
fn assert_matches(loaded: &HnswIndex, original: &HnswIndex, q: &[f64]) {
assert_eq!(loaded.node_ids(), original.node_ids(), "node ids differ");
for &id in loaded.node_ids().iter() {
assert_eq!(
loaded.back_refs_for_test(id),
loaded.scan_back_refs_for_test(id),
"back_refs[{id}] disagrees with a full scan after the load"
);
}
assert_eq!(
loaded.search(q, 10),
original.search(q, 10),
"the loaded graph answers differently"
);
}
#[test]
fn a_v1_blob_upgrades_in_place() {
let (vecs, idx) = blob_fixture();
let blob = as_v1_blob(&idx);
hnsw_insert_count_reset();
let loaded = decode_hnsw_blob(&blob).expect("a 0.6.5 blob must still load");
assert_eq!(
hnsw_insert_count(),
0,
"up-converting a v1 blob must not re-insert a single vector"
);
assert_matches(&loaded, &idx, &vecs[3]);
}
#[test]
fn a_mixed_dimension_upgrade_declines_the_fast_path() {
for shape in ["v1", "v2"] {
let mut vecs = make_unit_vecs(40, 24, 0x0D1D_0DDD);
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"mixed"));
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
vecs[7].truncate(23);
let blob = match shape {
"v1" => as_v1_blob_with(&idx, &vecs),
_ => as_v2_blob_with(&idx, &vecs),
};
let loaded = decode_hnsw_blob(&blob).expect("a mixed blob still loads");
assert_eq!(
loaded.node_ids(),
idx.node_ids(),
"{shape}: the odd node must stay in the graph"
);
assert!(
!loaded.can_answer(24),
"{shape}: an index holding a padded position must not claim the fast path"
);
let _ = loaded.search(&vecs[3], 5);
}
}
#[test]
fn a_truncated_v3_slab_is_refused() {
let (_vecs, mut idx) = blob_fixture();
let full = idx.slab.data.len();
idx.slab.data.truncate(full - idx.slab.dim);
let blob = encode_hnsw_blob(&idx, true).expect("encode");
let err = decode_hnsw_blob(&blob).expect_err("a truncated slab must be refused");
assert!(
err.contains("truncated"),
"the error must name the problem; got {err:?}"
);
let (_v, mut bad) = blob_fixture();
bad.id_of.pop();
let err = decode_hnsw_blob(&encode_hnsw_blob(&bad, true).expect("encode"))
.expect_err("a slot/id mismatch must be refused");
assert!(
err.contains("inconsistent"),
"the error must name the problem; got {err:?}"
);
}
#[test]
fn an_incomplete_blob_refuses_to_answer() {
let (vecs, idx) = blob_fixture();
let whole = decode_hnsw_blob(&encode_hnsw_blob(&idx, true).expect("encode"))
.expect("a complete blob loads");
assert!(whole.can_answer(24), "a complete blob must answer");
assert!(!whole.is_incomplete());
let partial = decode_hnsw_blob(&encode_hnsw_blob(&idx, false).expect("encode"))
.expect("an incomplete blob still loads");
assert!(
partial.is_incomplete(),
"the flag must survive the round trip"
);
assert!(
!partial.can_answer(24),
"a graph holding a prefix of its corpus must not claim the fast path"
);
assert_eq!(partial.node_ids(), idx.node_ids());
let mut adopted = partial;
adopted.mark_complete();
assert!(adopted.can_answer(24));
assert_eq!(adopted.search(&vecs[3], 5), whole.search(&vecs[3], 5));
}
#[test]
fn hnsw_blob_complete_peeks_the_last_byte() {
let (_vecs, idx) = blob_fixture();
let whole = encode_hnsw_blob(&idx, true).expect("encode");
let partial = encode_hnsw_blob(&idx, false).expect("encode");
assert_eq!(hnsw_blob_complete(&whole), Some(true));
assert_eq!(hnsw_blob_complete(&partial), Some(false));
assert_eq!(&whole[..whole.len() - 1], &partial[..partial.len() - 1]);
assert_eq!(whole[whole.len() - 1], 1);
assert_eq!(partial[partial.len() - 1], 0);
assert_eq!(hnsw_blob_complete(&[]), None);
}
#[test]
fn a_refused_insert_does_not_leave_a_stale_vector() {
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"stale"));
let vecs = make_unit_vecs(4, 8, 0x57A1_E000);
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
assert!(idx.node_ids().contains(&0));
idx.insert(0, &make_unit_vecs(1, 4, 0x57A1_E001)[0]);
assert!(
!idx.node_ids().contains(&0),
"the refused id must not keep its old vector"
);
assert!(
!idx.can_answer(8),
"a refusal means the index is missing a vector it was offered"
);
}
#[test]
fn a_malformed_params_string_yields_the_defaults() {
for bad in [
"",
"16,64,200",
"16,64,200,400,neither",
"a,b,c,d",
"16,0,200,400",
] {
assert_eq!(
HnswParams::parse(bad),
None,
"{bad:?} must not parse to a shape"
);
}
}
#[test]
fn a_v4_blob_round_trips() {
let (vecs, idx) = blob_fixture();
let blob = encode_hnsw_blob(&idx, true).expect("encode");
assert_eq!(
&blob[..4],
&HNSW_BLOB_MAGIC,
"the blob must carry its magic"
);
assert_eq!(
u16::from_le_bytes([blob[4], blob[5]]),
4,
"this build writes blob version 4"
);
hnsw_insert_count_reset();
let loaded = decode_hnsw_blob(&blob).expect("a v4 blob must load");
assert_eq!(hnsw_insert_count(), 0, "a load must not re-insert vectors");
assert_matches(&loaded, &idx, &vecs[3]);
assert_eq!(
loaded.slab.dim, idx.slab.dim,
"the slab's stride must survive the round trip"
);
}
#[test]
fn v3_blob_still_loads_and_reoffers() {
let (vecs, idx) = blob_fixture();
let blob = as_v3_blob(&idx, true);
assert_eq!(
u16::from_le_bytes([blob[4], blob[5]]),
3,
"the fixture must actually be a v3 blob"
);
hnsw_insert_count_reset();
let loaded = decode_hnsw_blob(&blob).expect("a v3 blob must still load");
assert_eq!(hnsw_insert_count(), 0, "a load must not re-insert vectors");
assert_matches(&loaded, &idx, &vecs[3]);
assert_eq!(
loaded.accounted_ids(),
loaded.node_ids(),
"a v3 blob carries no parked or refused state, so the scan's skip \
set is exactly the graph — the pre-v4 behaviour"
);
assert_eq!(loaded.dim_mismatches(), 0);
assert!(loaded.refused_ids().is_empty());
}
#[test]
fn a_v4_blob_restores_parked_and_refused() {
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"v4-state"));
idx.insert(1, &[1.0, 0.0]);
idx.insert(2, &[0.0, 1.0]);
idx.insert(3, &[1.0, 1.0, 1.0]); assert_eq!(idx.dim_mismatches(), 1, "the stray must be refused");
assert_eq!(idx.refused_ids(), BTreeSet::from([3]));
let blob = encode_hnsw_blob(&idx, true).expect("encode");
hnsw_insert_count_reset();
let loaded = decode_hnsw_blob(&blob).expect("decode");
assert_eq!(hnsw_insert_count(), 0, "a load must not re-insert vectors");
assert_eq!(
loaded.dim_mismatches(),
idx.dim_mismatches(),
"the refusal count survives the round trip rather than starting at zero"
);
assert_eq!(
loaded.refused_ids(),
idx.refused_ids(),
"and so do the ids behind it"
);
assert!(
loaded.accounts_for(3),
"the refused id is accounted for, so the open-time scan skips it \
instead of refusing it a second time"
);
assert!(
!loaded.node_ids().contains(&3),
"accounted for is not the same as in the graph"
);
}
#[test]
fn a_v4_blob_restores_a_parked_vector() {
let mut idx = HnswIndex::new(crate::index::fnv1a_u64(b"v4-parked"));
idx.insert(7, &[1.0, 1.0, 1.0]);
idx.insert(1, &[1.0, 0.0]);
idx.insert(2, &[0.0, 1.0]);
assert!(
idx.accounts_for(7) && !idx.node_ids().contains(&7),
"node 7 must be parked, not indexed"
);
let blob = encode_hnsw_blob(&idx, true).expect("encode");
let loaded = decode_hnsw_blob(&blob).expect("decode");
assert!(
loaded.accounts_for(7),
"the parked entry survives the round trip"
);
assert!(
!loaded.node_ids().contains(&7),
"and is still parked rather than indexed"
);
assert!(
loaded.accounted_ids().contains(&7),
"so the open-time scan skips it"
);
}
#[test]
fn a_v2_blob_upgrades_in_place() {
let (vecs, idx) = blob_fixture();
let blob = as_v2_blob(&idx);
hnsw_insert_count_reset();
let loaded = decode_hnsw_blob(&blob).expect("a v2 blob must still load");
assert_eq!(
hnsw_insert_count(),
0,
"up-converting a v2 blob must not re-insert a single vector"
);
assert_eq!(
loaded.slab.dim, idx.slab.dim,
"the slab's stride comes from the decoded vectors"
);
assert_matches(&loaded, &idx, &vecs[3]);
}
#[test]
fn an_unknown_version_is_rejected() {
let (_, idx) = blob_fixture();
let mut blob = encode_hnsw_blob(&idx, true).expect("encode");
blob[4] = HNSW_BLOB_VERSION as u8 + 1; let err = decode_hnsw_blob(&blob).expect_err("a future version must not be read");
assert!(
err.contains("version"),
"the refusal must name the version: {err}"
);
}
#[test]
fn a_v2_reader_refuses_a_v3_blob_and_still_reads_a_v2_one() {
const V2_CEILING: u16 = 2;
fn v2_era_decode(blob: &[u8]) -> Result<usize, String> {
match bincode::deserialize::<HnswBlobV2>(blob) {
Ok(b) if b.magic == HNSW_BLOB_MAGIC => {
if b.version == 0 || b.version > V2_CEILING {
Err(format!("version {} is not readable", b.version))
} else {
Ok(b.index.slots.len())
}
}
Ok(b) => Err(format!("unrecognised magic {:?}", b.magic)),
Err(e) => Err(format!("not a v2 blob ({e})")),
}
}
let (_, idx) = blob_fixture();
let v2 = as_v2_blob(&idx);
assert_eq!(
v2_era_decode(&v2),
Ok(idx.slots.len()),
"the reconstructed v2 reader must read a v2 blob"
);
let v3 = encode_hnsw_blob(&idx, true).expect("encode");
assert_eq!(&v3[..4], &HNSW_BLOB_MAGIC, "same magic, new version");
assert_eq!(u16::from_le_bytes([v3[4], v3[5]]), HNSW_BLOB_VERSION);
let err = v2_era_decode(&v3).expect_err(
"a build that reads up to v2 must refuse a v3 blob rather than \
misread it — its caller then keeps the full scan",
);
eprintln!("a v2-era reader on a v3 blob: {err}");
let mut future = v3.clone();
future[4] = HNSW_BLOB_VERSION as u8 + 1;
let err = decode_hnsw_blob(&future).expect_err("a future version must not be read");
assert!(err.contains("version"), "{err}");
}
#[test]
fn a_foreign_magic_is_rejected() {
let (_, idx) = blob_fixture();
let mut blob = encode_hnsw_blob(&idx, true).expect("encode");
blob[0] = b'X';
assert!(
decode_hnsw_blob(&blob).is_err(),
"a blob with foreign magic must not be read"
);
}
#[test]
fn a_v3_blob_is_not_readable_as_a_v1_index() {
#[derive(Deserialize)]
#[allow(dead_code)]
struct V1ReadNode {
level: usize,
vector: Vec<f64>,
layers: Vec<Vec<u32>>,
}
#[derive(Deserialize)]
#[allow(dead_code)]
struct V1ReadIndex {
base_seed: u64,
nodes: BTreeMap<u32, V1ReadNode>,
entry_point: Option<u32>,
max_level: usize,
}
let (_, idx) = blob_fixture();
let blob = encode_hnsw_blob(&idx, true).expect("encode");
assert!(
bincode::deserialize::<V1ReadIndex>(&blob).is_err(),
"a 0.6.5 reader must reject a 0.6.6 blob, not misread it"
);
}
#[test]
#[ignore]
fn hnsw_5k_1536_recall() {
const N: usize = 5_000;
const DIM: usize = 1_536;
const N_QUERIES: usize = 50;
const K: usize = 10;
let seed = crate::index::fnv1a_u64(b"recall-probe-5k-1536");
let vecs = make_unit_vecs(N, DIM, seed);
let mut idx = HnswIndex::new(seed);
for (i, v) in vecs.iter().enumerate() {
idx.insert(i as u32, v);
}
let q_seed = crate::index::fnv1a_u64(b"recall-queries");
let queries = make_unit_vecs(N_QUERIES, DIM, q_seed);
let mut recalls = Vec::with_capacity(N_QUERIES);
for q in &queries {
let exact_set: std::collections::BTreeSet<usize> =
exact_knn(&vecs, q, K).into_iter().collect();
let approx_ids: Vec<usize> = idx
.search(q, K)
.into_iter()
.map(|(id, _)| id as usize)
.collect();
let hits = approx_ids
.iter()
.filter(|id| exact_set.contains(id))
.count();
recalls.push(hits as f64 / K as f64);
}
let min_recall = recalls.iter().cloned().fold(f64::MAX, f64::min);
let mean_recall = recalls.iter().sum::<f64>() / recalls.len() as f64;
eprintln!("HNSW 5k/1536 recall@{K}: min={min_recall:.4} mean={mean_recall:.4}");
assert!(
min_recall >= 0.90,
"min recall@{K} = {min_recall:.4} < 0.90"
);
assert!(
mean_recall >= 0.95,
"mean recall@{K} = {mean_recall:.4} < 0.95"
);
}
}