use alloc::vec::Vec;
#[cfg(feature = "counters")]
use core::cell::Cell;
use plugmem_arena::{Arena, ArenaCfg, ChunkPool, ChunkPoolCfg, ListHandle, ShardMode, Slot, key};
use xxhash_rust::xxh3::xxh3_64;
use crate::error::Error;
use crate::id::NONE_U32;
use crate::index::vecpool::{VecPool, dot_i8};
const MAX_LEVEL: usize = 16;
const UPPER_SHARDS: usize = 64;
const NEIGHBOR_BYTES: usize = core::mem::size_of::<u32>();
const META_BYTES: usize = 2 * core::mem::size_of::<u32>();
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct UpperSlot {
slot: u32,
level: u32,
handle: ListHandle,
}
impl Slot for UpperSlot {
const SIZE: usize = 20;
const KEY_LEN: usize = 8;
fn write(&self, out: &mut [u8]) {
key::write_u32(out, self.slot);
key::write_u32(&mut out[4..], self.level);
out[8..20].copy_from_slice(&self.handle.to_bytes());
}
fn read(bytes: &[u8]) -> Self {
Self {
slot: key::read_u32(bytes),
level: key::read_u32(&bytes[4..]),
handle: ListHandle::from_bytes(bytes[8..20].try_into().unwrap()),
}
}
}
#[derive(Debug, Default)]
pub struct HnswScratch {
visited: Vec<u32>,
epoch: u32,
cand: Vec<(f32, u32)>,
found: Vec<(f32, u32)>,
nbrs: Vec<u32>,
sel: Vec<u32>,
pruned: Vec<u32>,
relink: Vec<(f32, u32)>,
}
#[inline]
fn better(a: (f32, u32), b: (f32, u32)) -> core::cmp::Ordering {
a.0.total_cmp(&b.0).then(b.1.cmp(&a.1))
}
pub struct HnswGraph<'a> {
m: usize,
m0: usize,
level0: Vec<u32>,
upper: Arena<'a, UpperSlot>,
lists: ChunkPool<'a>,
entry: u32,
indexed: u32,
thresholds: [u64; MAX_LEVEL],
#[cfg(feature = "counters")]
dist_evals: Cell<u64>,
}
impl<'a> HnswGraph<'a> {
pub fn new(m: usize, m0: usize, max_bytes: usize) -> Result<Self, Error> {
let mut thresholds = [0u64; MAX_LEVEL];
let mut t = u64::MAX;
for slot in &mut thresholds {
t /= m as u64;
*slot = t;
}
Ok(Self {
m,
m0,
level0: Vec::new(),
upper: Arena::new(
ArenaCfg::new(UPPER_SHARDS, ShardMode::Uniform).with_max_bytes(max_bytes),
)?,
lists: ChunkPool::new(ChunkPoolCfg::new().with_max_bytes(max_bytes)),
entry: NONE_U32,
indexed: 0,
thresholds,
#[cfg(feature = "counters")]
dist_evals: Cell::new(0),
})
}
pub fn indexed(&self) -> u32 {
self.indexed
}
fn level_of(&self, fact: u32) -> usize {
let h = xxh3_64(&fact.to_le_bytes());
self.thresholds.iter().take_while(|&&t| h < t).count()
}
#[inline]
fn sim_q(&self, pool: &VecPool<'_>, q: (f32, &[u8]), slot: u32) -> f32 {
#[cfg(feature = "counters")]
self.dist_evals.set(self.dist_evals.get() + 1);
let (s, qb) = pool.quant(slot as usize);
q.0 * s * dot_i8(q.1, qb) as f32
}
#[inline]
fn block(&self, slot: u32) -> &[u32] {
let at = slot as usize * self.m0;
&self.level0[at..at + self.m0]
}
fn neighbors_into(&self, slot: u32, level: usize, out: &mut Vec<u32>) {
out.clear();
if level == 0 {
out.extend(
self.block(slot)
.iter()
.copied()
.take_while(|&n| n != NONE_U32),
);
return;
}
let mut kb = [0u8; 8];
key::write_u32(&mut kb, slot);
key::write_u32(&mut kb[4..], level as u32);
let Some(entry) = self.upper.get(&kb) else {
return;
};
for chunk in self.lists.iter(&entry.handle) {
for raw in chunk.chunks_exact(4) {
out.push(u32::from_le_bytes(raw.try_into().unwrap()));
}
}
}
#[inline]
fn visit(scratch: &mut HnswScratch, slot: u32) -> bool {
let at = slot as usize;
if scratch.visited[at] == scratch.epoch {
return true;
}
scratch.visited[at] = scratch.epoch;
false
}
fn search_layer(
&self,
pool: &VecPool<'_>,
q: (f32, &[u8]),
level: usize,
ep: u32,
ef: usize,
scratch: &mut HnswScratch,
) {
scratch.epoch = scratch.epoch.wrapping_add(1);
if scratch.epoch == 0 {
scratch.visited.fill(u32::MAX);
scratch.epoch = 1;
}
scratch
.visited
.resize(self.indexed as usize, scratch.epoch.wrapping_sub(1));
scratch.cand.clear();
scratch.found.clear();
Self::visit(scratch, ep);
let s = self.sim_q(pool, q, ep);
scratch.cand.push((s, ep));
scratch.found.push((s, ep));
while let Some(best) = scratch.cand.pop() {
if scratch.found.len() >= ef && better(best, scratch.found[0]).is_lt() {
break;
}
let nbrs = core::mem::take(&mut scratch.nbrs);
let mut nbrs = nbrs;
self.neighbors_into(best.1, level, &mut nbrs);
for &nb in &nbrs {
if Self::visit(scratch, nb) {
continue;
}
let s = self.sim_q(pool, q, nb);
let entry = (s, nb);
if scratch.found.len() < ef || better(entry, scratch.found[0]).is_gt() {
let at = scratch.found.partition_point(|&e| better(e, entry).is_lt());
scratch.found.insert(at, entry);
if scratch.found.len() > ef {
scratch.found.remove(0);
}
let at = scratch.cand.partition_point(|&e| better(e, entry).is_lt());
scratch.cand.insert(at, entry);
}
}
scratch.nbrs = nbrs;
}
}
fn select_neighbors(&self, pool: &VecPool<'_>, cap: usize, scratch: &mut HnswScratch) {
scratch.sel.clear();
scratch.pruned.clear();
for i in (0..scratch.found.len()).rev() {
let (sim, cand) = scratch.found[i];
if scratch.sel.len() >= cap {
break;
}
let dominated = scratch.sel.iter().any(|&kept| {
#[cfg(feature = "counters")]
self.dist_evals.set(self.dist_evals.get() + 1);
pool.sim(cand, kept) > sim
});
if dominated {
scratch.pruned.push(cand);
} else {
scratch.sel.push(cand);
}
}
for &p in scratch.pruned.iter() {
if scratch.sel.len() >= cap {
break;
}
scratch.sel.push(p);
}
}
fn write_list(&mut self, slot: u32, level: usize, sel: &[u32]) -> Result<(), Error> {
if level == 0 {
let at = slot as usize * self.m0;
let block = &mut self.level0[at..at + self.m0];
block.fill(NONE_U32);
block[..sel.len()].copy_from_slice(sel);
return Ok(());
}
let mut kb = [0u8; 8];
key::write_u32(&mut kb, slot);
key::write_u32(&mut kb[4..], level as u32);
let mut handle = match self.upper.get(&kb) {
Some(entry) => {
let mut h = entry.handle;
self.lists.free(&mut h);
h
}
None => ListHandle::EMPTY,
};
for &n in sel {
self.lists.push(&mut handle, &n.to_le_bytes())?;
}
let updated = UpperSlot {
slot,
level: level as u32,
handle,
};
if self.upper.contains(&kb) {
let payload = self.upper.payload_mut(&kb).expect("checked above");
let mut full = [0u8; UpperSlot::SIZE];
updated.write(&mut full);
payload.copy_from_slice(&full[UpperSlot::KEY_LEN..]);
} else {
self.upper.insert(&updated)?;
}
Ok(())
}
fn add_link(
&mut self,
pool: &VecPool<'_>,
v: u32,
new: u32,
level: usize,
scratch: &mut HnswScratch,
) -> Result<(), Error> {
let cap = if level == 0 { self.m0 } else { self.m };
let mut nbrs = core::mem::take(&mut scratch.nbrs);
self.neighbors_into(v, level, &mut nbrs);
if nbrs.len() < cap {
nbrs.push(new);
let sel = core::mem::take(&mut scratch.sel);
let mut sel = sel;
sel.clear();
sel.extend_from_slice(&nbrs);
let res = self.write_list(v, level, &sel);
scratch.sel = sel;
scratch.nbrs = nbrs;
return res;
}
let vq = pool.quant(v as usize);
scratch.relink.clear();
for &n in nbrs.iter().chain(core::iter::once(&new)) {
let s = self.sim_q(pool, (vq.0, vq.1), n);
scratch.relink.push((s, n));
}
scratch.nbrs = nbrs;
scratch.relink.sort_unstable_by(|a, b| better(*a, *b));
scratch.found.clear();
scratch.found.extend_from_slice(&scratch.relink);
self.select_neighbors(pool, cap, scratch);
let sel = core::mem::take(&mut scratch.sel);
let res = self.write_list(v, level, &sel);
scratch.sel = sel;
res
}
pub fn insert_bulk(
&mut self,
pool: &VecPool<'_>,
upto: u32,
ef_construction: usize,
scratch: &mut HnswScratch,
) -> Result<(), Error> {
debug_assert!(upto as usize <= pool.len());
self.level0.resize(upto as usize * self.m0, NONE_U32);
for slot in self.indexed..upto {
self.insert_one(pool, slot, ef_construction, scratch)?;
self.indexed = slot + 1;
}
Ok(())
}
fn insert_one(
&mut self,
pool: &VecPool<'_>,
slot: u32,
ef_construction: usize,
scratch: &mut HnswScratch,
) -> Result<(), Error> {
let level = self.level_of(pool.slot_fact(slot as usize));
if self.entry == NONE_U32 {
self.entry = slot;
return Ok(());
}
let q = pool.quant(slot as usize);
let q = (q.0, q.1);
let top = self.level_of(pool.slot_fact(self.entry as usize));
let mut ep = self.entry;
let mut lev = top;
while lev > level {
self.search_layer(pool, q, lev, ep, 1, scratch);
ep = scratch.found.last().expect("entry is always found").1;
lev -= 1;
}
let mut lev = level.min(top);
loop {
self.search_layer(pool, q, lev, ep, ef_construction, scratch);
ep = scratch.found.last().expect("entry is always found").1;
let cap = if lev == 0 { self.m0 } else { self.m };
self.select_neighbors(pool, cap, scratch);
let sel = core::mem::take(&mut scratch.sel);
self.write_list(slot, lev, &sel)?;
for &nb in &sel {
self.add_link(pool, nb, slot, lev, scratch)?;
}
scratch.sel = sel;
if lev == 0 {
break;
}
lev -= 1;
}
if level > top {
self.entry = slot;
}
Ok(())
}
pub fn search(
&self,
pool: &VecPool<'_>,
query: &[f32],
ef: usize,
vec_scratch: &mut crate::index::vecpool::VecScratch,
scratch: &mut HnswScratch,
out: &mut Vec<(u32, f32)>,
) -> Result<(), Error> {
pool.quantize_query(query, vec_scratch)?;
let q = pool.quantized(vec_scratch);
self.search_quantized(pool, q, ef, scratch, out);
Ok(())
}
pub(crate) fn search_quantized(
&self,
pool: &VecPool<'_>,
q: (f32, &[u8]),
ef: usize,
scratch: &mut HnswScratch,
out: &mut Vec<(u32, f32)>,
) {
out.clear();
if self.entry == NONE_U32 {
return;
}
let mut ep = self.entry;
let top = self.level_of(pool.slot_fact(self.entry as usize));
for lev in (1..=top).rev() {
self.search_layer(pool, q, lev, ep, 1, scratch);
ep = scratch.found.last().expect("entry is always found").1;
}
self.search_layer(pool, q, 0, ep, ef.max(1), scratch);
for &(sim, slot) in scratch.found.iter().rev() {
out.push((slot, sim));
}
}
pub(crate) fn pool_bytes(&self) -> usize {
self.level0.len() * NEIGHBOR_BYTES + self.upper.pool_bytes() + self.lists.pool_bytes()
}
pub(crate) fn remapped(
&self,
map: &[u32],
new_pool: &VecPool<'_>,
max_bytes: usize,
) -> Result<HnswGraph<'static>, Error> {
let mut g: HnswGraph<'static> = HnswGraph::new(self.m, self.m0, max_bytes)?;
let old_indexed = self.indexed as usize;
let new_indexed = map[..old_indexed]
.iter()
.filter(|&&m| m != NONE_U32)
.count() as u32;
g.level0 = alloc::vec![NONE_U32; new_indexed as usize * self.m0];
g.indexed = new_indexed;
let mut nbrs = Vec::new();
let mut sel = Vec::new();
for old in 0..old_indexed as u32 {
let new = map[old as usize];
if new == NONE_U32 {
continue;
}
self.neighbors_into(old, 0, &mut nbrs);
sel.clear();
sel.extend(
nbrs.iter()
.map(|&n| map[n as usize])
.filter(|&n| n != NONE_U32),
);
g.write_list(new, 0, &sel)?;
let levels = g.level_of(new_pool.slot_fact(new as usize));
for level in 1..=levels {
self.neighbors_into(old, level, &mut nbrs);
if nbrs.is_empty() {
continue;
}
sel.clear();
sel.extend(
nbrs.iter()
.map(|&n| map[n as usize])
.filter(|&n| n != NONE_U32),
);
g.write_list(new, level, &sel)?;
}
}
g.entry = if self.entry != NONE_U32 && map[self.entry as usize] != NONE_U32 {
map[self.entry as usize]
} else {
let mut best = NONE_U32;
let mut best_level = 0usize;
for slot in 0..new_indexed {
let level = g.level_of(new_pool.slot_fact(slot as usize));
if best == NONE_U32 || level > best_level {
best = slot;
best_level = level;
}
}
best
};
Ok(g)
}
pub(crate) fn dump_meta(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(META_BYTES);
out.extend_from_slice(&self.entry.to_le_bytes());
out.extend_from_slice(&self.indexed.to_le_bytes());
out
}
pub(crate) fn dump_level0(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(self.level0.len() * NEIGHBOR_BYTES);
for &n in &self.level0 {
out.extend_from_slice(&n.to_le_bytes());
}
out
}
pub(crate) fn dump_upper(&self) -> [Vec<u8>; 4] {
let (mut am, mut ap) = (Vec::new(), Vec::new());
self.upper.dump_meta(&mut am);
self.upper.dump_pool(&mut ap);
let (mut cm, mut cp) = (Vec::new(), Vec::new());
self.lists.dump_meta(&mut cm);
self.lists.dump_pool(&mut cp);
[am, ap, cm, cp]
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn from_parts(
m: usize,
m0: usize,
max_bytes: usize,
meta: &[u8],
level0: &[u8],
upper_meta: &[u8],
upper_pool: &[u8],
lists_meta: &[u8],
lists_pool: &[u8],
) -> Result<Self, Error> {
let mut g = Self::new(m, m0, max_bytes)?;
if meta.len() != META_BYTES {
return Err(Error::Corrupt("hnsw meta section has a wrong length"));
}
g.entry = u32::from_le_bytes(meta[0..4].try_into().unwrap());
g.indexed = u32::from_le_bytes(meta[4..8].try_into().unwrap());
if level0.len() as u64 != u64::from(g.indexed) * m0 as u64 * NEIGHBOR_BYTES as u64 {
return Err(Error::Corrupt("hnsw level0 length mismatch"));
}
g.level0 = level0
.chunks_exact(NEIGHBOR_BYTES)
.map(|b| u32::from_le_bytes(b.try_into().unwrap()))
.collect();
g.upper = Arena::load(
ArenaCfg::new(UPPER_SHARDS, ShardMode::Uniform).with_max_bytes(max_bytes),
upper_meta,
upper_pool,
)?;
g.lists = ChunkPool::load(
ChunkPoolCfg::new().with_max_bytes(max_bytes),
lists_meta,
lists_pool,
)?;
Ok(g)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn from_parts_borrowed(
m: usize,
m0: usize,
max_bytes: usize,
meta: &[u8],
level0: &[u8],
upper_meta: &[u8],
upper_pool: &'a [u8],
lists_meta: &[u8],
lists_pool: &'a [u8],
) -> Result<Self, Error> {
let mut g = Self::new(m, m0, max_bytes)?;
if meta.len() != META_BYTES {
return Err(Error::Corrupt("hnsw meta section has a wrong length"));
}
g.entry = u32::from_le_bytes(meta[0..4].try_into().unwrap());
g.indexed = u32::from_le_bytes(meta[4..8].try_into().unwrap());
if level0.len() as u64 != u64::from(g.indexed) * m0 as u64 * NEIGHBOR_BYTES as u64 {
return Err(Error::Corrupt("hnsw level0 length mismatch"));
}
g.level0 = level0
.chunks_exact(NEIGHBOR_BYTES)
.map(|b| u32::from_le_bytes(b.try_into().unwrap()))
.collect();
g.upper = Arena::load_borrowed(
ArenaCfg::new(UPPER_SHARDS, ShardMode::Uniform).with_max_bytes(max_bytes),
upper_meta,
upper_pool,
)?;
g.lists = ChunkPool::load_borrowed(
ChunkPoolCfg::new().with_max_bytes(max_bytes),
lists_meta,
lists_pool,
)?;
Ok(g)
}
pub(crate) fn validate(&self, pool: &VecPool<'_>) -> Result<(), Error> {
if self.indexed as usize > pool.len() {
return Err(Error::Corrupt("hnsw indexes more slots than the pool"));
}
if self.level0.len() != self.indexed as usize * self.m0 {
return Err(Error::Corrupt("hnsw level0 disagrees with indexed"));
}
if self.indexed == 0 {
if self.entry != NONE_U32 || !self.upper.is_empty() || self.lists.chunks() != 0 {
return Err(Error::Corrupt("hnsw empty graph carries state"));
}
return Ok(());
}
if self.entry >= self.indexed {
return Err(Error::Corrupt("hnsw entry out of range"));
}
for slot in 0..self.indexed {
let block = self.block(slot);
let mut ended = false;
for &n in block {
if n == NONE_U32 {
ended = true;
continue;
}
if ended {
return Err(Error::Corrupt("hnsw level0 padding is not canonical"));
}
if n >= self.indexed || n == slot {
return Err(Error::Corrupt("hnsw level0 neighbor out of range"));
}
}
}
let mut visited = alloc::vec![false; self.lists.chunks()];
for entry in self.upper.iter() {
if entry.slot >= self.indexed {
return Err(Error::Corrupt("hnsw upper handle out of range"));
}
let max_level = self.level_of(pool.slot_fact(entry.slot as usize));
if entry.level == 0 || entry.level as usize > max_level {
return Err(Error::Corrupt("hnsw upper level disagrees with the hash"));
}
self.lists.validate_chain(&entry.handle, &mut visited)?;
let mut count = 0u32;
for chunk in self.lists.iter(&entry.handle) {
if !chunk.len().is_multiple_of(4) {
return Err(Error::Corrupt("hnsw upper list is not a slot sequence"));
}
for raw in chunk.chunks_exact(4) {
let n = u32::from_le_bytes(raw.try_into().unwrap());
if n >= self.indexed || n == entry.slot {
return Err(Error::Corrupt("hnsw upper neighbor out of range"));
}
count += 1;
}
}
if count != entry.handle.len() || count as usize > self.m {
return Err(Error::Corrupt("hnsw upper list disagrees with its handle"));
}
}
if self.lists.orphan_count(&visited) != 0 {
return Err(Error::Corrupt("hnsw list pool has orphan chunks"));
}
Ok(())
}
#[cfg(feature = "counters")]
pub fn dist_evals(&self) -> u64 {
self.dist_evals.get()
}
#[cfg(feature = "counters")]
pub fn reset_dist_evals(&self) {
self.dist_evals.set(0);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::id::FactId;
use alloc::vec;
struct Lcg(u64);
impl Lcg {
fn next(&mut self) -> f32 {
self.0 = self
.0
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((self.0 >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
}
}
fn cluster_pool(n: usize, dim: usize, clusters: usize, seed: u64) -> VecPool<'static> {
let mut rng = Lcg(seed);
let centers: Vec<Vec<f32>> = (0..clusters)
.map(|_| (0..dim).map(|_| rng.next()).collect())
.collect();
let mut pool = VecPool::new(dim, usize::MAX);
for i in 0..n {
let c = ¢ers[i % clusters];
let v: Vec<f32> = c.iter().map(|&x| x + rng.next() * 0.3).collect();
pool.push(FactId(i as u32), &v).unwrap();
}
pool
}
fn build(pool: &VecPool<'_>, m: usize, m0: usize) -> HnswGraph<'static> {
let mut g = HnswGraph::new(m, m0, usize::MAX).unwrap();
let mut scratch = HnswScratch::default();
g.insert_bulk(pool, pool.len() as u32, 200, &mut scratch)
.unwrap();
g
}
fn brute_force(pool: &VecPool<'_>, q: u32, k: usize) -> Vec<u32> {
let mut all: Vec<(f32, u32)> = (0..pool.len() as u32)
.map(|i| (pool.sim(q, i), i))
.collect();
all.sort_unstable_by(|a, b| better(*b, *a));
all.into_iter().take(k).map(|(_, s)| s).collect()
}
#[test]
#[cfg_attr(miri, ignore)] fn levels_are_geometric_and_pure() {
let g = HnswGraph::new(16, 32, usize::MAX).unwrap();
let n = 100_000u32;
let mut per_level = [0usize; 4];
for fact in 0..n {
let l = g.level_of(fact).min(3);
per_level[l] += 1;
}
let at_least_1: usize = per_level[1..].iter().sum();
assert!(
(4_000..9_000).contains(&at_least_1),
"level>=1 count {at_least_1} is out of band"
);
let at_least_2: usize = per_level[2..].iter().sum();
assert!(
(150..800).contains(&at_least_2),
"level>=2 count {at_least_2} is out of band"
);
assert_eq!(g.level_of(42), g.level_of(42));
}
#[test]
#[cfg_attr(miri, ignore)] fn recall_against_brute_force() {
let dim = 32;
let pool = cluster_pool(2_000, dim, 64, 0xA11CE);
let g = build(&pool, 16, 32);
let mut scratch = HnswScratch::default();
let mut out = Vec::new();
let mut hits = 0usize;
let mut total = 0usize;
for q in (0..2_000u32).step_by(97) {
let truth = brute_force(&pool, q, 10);
let (scale, qb) = pool.quant(q as usize);
g.search_quantized(&pool, (scale, qb), 64, &mut scratch, &mut out);
let got: Vec<u32> = out.iter().take(10).map(|&(s, _)| s).collect();
hits += truth.iter().filter(|t| got.contains(t)).count();
total += truth.len();
}
let recall = hits as f64 / total as f64;
assert!(recall >= 0.9, "recall@10 {recall} below the 0.9 gate");
}
#[test]
#[cfg_attr(miri, ignore)] fn build_is_deterministic() {
let pool = cluster_pool(600, 24, 16, 7);
let a = build(&pool, 8, 16);
let b = build(&pool, 8, 16);
assert_eq!(a.level0, b.level0);
assert_eq!(a.entry, b.entry);
assert_eq!(a.indexed, b.indexed);
let (mut am, mut bm) = (Vec::new(), Vec::new());
a.upper.dump_meta(&mut am);
b.upper.dump_meta(&mut bm);
assert_eq!(am, bm);
let (mut ap, mut bp) = (Vec::new(), Vec::new());
a.lists.dump_pool(&mut ap);
b.lists.dump_pool(&mut bp);
assert_eq!(ap, bp);
}
#[test]
#[cfg_attr(miri, ignore)] fn degree_caps_hold() {
let pool = cluster_pool(800, 16, 8, 3);
let g = build(&pool, 6, 12);
let mut nbrs = Vec::new();
for slot in 0..g.indexed() {
g.neighbors_into(slot, 0, &mut nbrs);
assert!(nbrs.len() <= 12);
assert!(!nbrs.contains(&slot));
let mut sorted = nbrs.clone();
sorted.sort_unstable();
sorted.dedup();
assert_eq!(sorted.len(), nbrs.len());
assert!(nbrs.iter().all(|&n| n < g.indexed()));
for level in 1..=g.level_of(pool.slot_fact(slot as usize)) {
g.neighbors_into(slot, level, &mut nbrs);
assert!(nbrs.len() <= 6, "level {level} degree overflow");
}
}
}
#[test]
#[cfg_attr(miri, ignore)] fn remap_survives_a_dead_entry() {
let pool = cluster_pool(300, 16, 8, 21);
let g = build(&pool, 6, 12);
let entry = g.entry;
let mut map = alloc::vec![NONE_U32; 300];
let mut new_pool = VecPool::new(16, usize::MAX);
let mut next = 0u32;
for old in 0..300u32 {
if old == entry || old % 7 == 0 {
continue;
}
map[old as usize] = next;
new_pool.copy_slot(&pool, old);
next += 1;
}
let remapped = g.remapped(&map, &new_pool, usize::MAX).unwrap();
assert_eq!(remapped.indexed(), next);
assert_ne!(remapped.entry, NONE_U32, "a survivor takes the entry");
remapped.validate(&new_pool).unwrap();
let probe_old = 1u32; let probe_old = if probe_old == entry { 2 } else { probe_old };
let probe_new = map[probe_old as usize];
let (scale, qb) = new_pool.quant(probe_new as usize);
let mut scratch = HnswScratch::default();
let mut out = Vec::new();
remapped.search_quantized(&new_pool, (scale, qb), 32, &mut scratch, &mut out);
assert_eq!(out[0].0, probe_new);
}
#[test]
fn validate_rejects_malformed_graphs() {
let pool = cluster_pool(50, 8, 4, 5);
let g = build(&pool, 4, 8);
let dump = (g.dump_meta(), g.dump_level0(), g.dump_upper());
let load = |meta: &[u8], level0: &[u8]| {
HnswGraph::from_parts(
4,
8,
usize::MAX,
meta,
level0,
&dump.2[0],
&dump.2[1],
&dump.2[2],
&dump.2[3],
)
};
load(&dump.0, &dump.1).unwrap().validate(&pool).unwrap();
let mut meta = dump.0.clone();
meta[0..4].copy_from_slice(&999u32.to_le_bytes());
assert!(load(&meta, &dump.1).unwrap().validate(&pool).is_err());
let mut level0 = dump.1.clone();
level0[0..4].copy_from_slice(&500u32.to_le_bytes());
assert!(load(&dump.0, &level0).unwrap().validate(&pool).is_err());
let mut level0 = dump.1.clone();
level0[0..4].copy_from_slice(&NONE_U32.to_le_bytes());
level0[4..8].copy_from_slice(&1u32.to_le_bytes());
assert!(load(&dump.0, &level0).unwrap().validate(&pool).is_err());
assert!(load(&dump.0, &dump.1[..dump.1.len() - 4]).is_err());
let mut meta = dump.0.clone();
meta[4..8].copy_from_slice(&0u32.to_le_bytes());
assert!(load(&meta, &[]).unwrap().validate(&pool).is_err());
}
#[test]
fn tiny_graphs() {
let pool = cluster_pool(1, 8, 1, 1);
let mut g = HnswGraph::new(4, 8, usize::MAX).unwrap();
let mut scratch = HnswScratch::default();
let mut out = vec![(0u32, 0.0f32)];
let (scale, qb) = pool.quant(0);
g.search_quantized(&pool, (scale, qb), 8, &mut scratch, &mut out);
assert!(out.is_empty(), "an empty graph must answer empty");
g.insert_bulk(&pool, 1, 50, &mut scratch).unwrap();
g.search_quantized(&pool, (scale, qb), 8, &mut scratch, &mut out);
assert_eq!(out.len(), 1);
assert_eq!(out[0].0, 0);
}
}