use std::cmp::Ordering;
use std::collections::{BinaryHeap, HashMap, HashSet};
use rand::rngs::SmallRng;
use rand::{Rng, SeedableRng};
use crate::error::Result;
#[inline]
pub(crate) fn dot_f64(a: &[f32], b: &[f32]) -> f64 {
crate::vector::kernels::dot_f64(a, b)
}
#[inline]
pub(crate) fn norm_f64(v: &[f32]) -> f64 {
v.iter()
.map(|&x| (x as f64) * (x as f64))
.sum::<f64>()
.sqrt()
}
#[inline]
pub(crate) fn dist_from(dot: f64, norm_a: f64, norm_b: f64) -> f64 {
if !(norm_a > 0.0) || !(norm_b > 0.0) {
return 1.0;
}
(1.0 - (dot / (norm_a * norm_b))).clamp(0.0, 2.0)
}
#[inline]
fn cosine_distance(a: &[f32], b: &[f32]) -> f64 {
dist_from(dot_f64(a, b), norm_f64(a), norm_f64(b))
}
#[derive(Clone)]
struct Candidate {
idx: usize,
distance: f64,
}
impl PartialEq for Candidate {
fn eq(&self, other: &Self) -> bool {
self.distance == other.distance
}
}
impl Eq for Candidate {}
impl Ord for Candidate {
fn cmp(&self, other: &Self) -> Ordering {
other.distance.total_cmp(&self.distance)
}
}
impl PartialOrd for Candidate {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
#[derive(Clone)]
struct FarCandidate {
idx: usize,
distance: f64,
}
impl PartialEq for FarCandidate {
fn eq(&self, other: &Self) -> bool {
self.distance == other.distance
}
}
impl Eq for FarCandidate {}
impl Ord for FarCandidate {
fn cmp(&self, other: &Self) -> Ordering {
self.distance.total_cmp(&other.distance)
}
}
impl PartialOrd for FarCandidate {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
#[derive(Clone)]
struct HnswNode {
embedding: Vec<f32>,
norm: f64,
neighbors: Vec<Vec<usize>>,
tombstoned: bool,
}
const HNSW_LEVEL_SEED: u64 = 0x11DB_5EED;
#[derive(Clone)]
pub struct HnswIndex {
dim: usize,
m: usize,
m_max0: usize,
ef_construction: usize,
ef_search: usize,
ml: f64,
entry_point: Option<usize>,
max_layer: usize,
nodes: Vec<HnswNode>,
rid_to_idx: HashMap<String, usize>,
idx_to_rid: Vec<String>,
free_list: Vec<usize>,
active_count: usize,
rng: SmallRng,
incoming0: Vec<usize>,
}
impl HnswIndex {
pub fn new(dim: usize) -> Self {
Self::with_params(dim, 16, 200, 200)
}
pub fn with_params(dim: usize, m: usize, ef_construction: usize, ef_search: usize) -> Self {
#[cfg(feature = "testing")]
let rng = match std::env::var("YANTRIKDB_HNSW_SEED")
.ok()
.and_then(|s| s.parse::<u64>().ok())
{
Some(seed) => SmallRng::seed_from_u64(seed),
None => SmallRng::seed_from_u64(HNSW_LEVEL_SEED),
};
#[cfg(not(feature = "testing"))]
let rng = SmallRng::seed_from_u64(HNSW_LEVEL_SEED);
Self {
dim,
m,
m_max0: m * 2,
ef_construction,
ef_search,
ml: 1.0 / (m as f64).ln(),
entry_point: None,
max_layer: 0,
nodes: Vec::new(),
rid_to_idx: HashMap::new(),
idx_to_rid: Vec::new(),
free_list: Vec::new(),
active_count: 0,
rng,
incoming0: Vec::new(),
}
}
pub fn with_params_seeded(
dim: usize,
m: usize,
ef_construction: usize,
ef_search: usize,
seed: u64,
) -> Self {
let mut idx = Self::with_params(dim, m, ef_construction, ef_search);
idx.rng = SmallRng::seed_from_u64(seed);
idx
}
pub fn dim(&self) -> usize {
self.dim
}
pub fn len(&self) -> usize {
self.active_count
}
pub fn is_empty(&self) -> bool {
self.active_count == 0
}
pub fn insert(&mut self, rid: &str, embedding: &[f32]) -> Result<()> {
assert_eq!(embedding.len(), self.dim, "embedding dimension mismatch");
if let Some(&existing_idx) = self.rid_to_idx.get(rid) {
let node = &mut self.nodes[existing_idx];
if node.tombstoned {
node.embedding = embedding.to_vec();
node.norm = norm_f64(embedding);
node.tombstoned = false;
self.active_count += 1;
let level = node.neighbors.len().saturating_sub(1);
self.connect_node(existing_idx, level);
return Ok(());
}
node.embedding = embedding.to_vec();
node.norm = norm_f64(embedding);
return Ok(());
}
let level = self.random_level();
let idx = if let Some(free_idx) = self.free_list.pop() {
let neighbors = (0..=level).map(|_| Vec::new()).collect();
self.nodes[free_idx] = HnswNode {
embedding: embedding.to_vec(),
norm: norm_f64(embedding),
neighbors,
tombstoned: false,
};
self.idx_to_rid[free_idx] = rid.to_string();
free_idx
} else {
let neighbors = (0..=level).map(|_| Vec::new()).collect();
self.nodes.push(HnswNode {
embedding: embedding.to_vec(),
norm: norm_f64(embedding),
neighbors,
tombstoned: false,
});
self.idx_to_rid.push(rid.to_string());
self.incoming0.push(0);
self.nodes.len() - 1
};
self.rid_to_idx.insert(rid.to_string(), idx);
self.active_count += 1;
if self.entry_point.is_none() {
self.entry_point = Some(idx);
self.max_layer = level;
return Ok(());
}
self.connect_node(idx, level);
if level > self.max_layer {
self.entry_point = Some(idx);
self.max_layer = level;
}
Ok(())
}
#[inline]
fn dist_to(&self, query: &[f32], qnorm: f64, idx: usize) -> f64 {
let node = &self.nodes[idx];
dist_from(dot_f64(query, &node.embedding), qnorm, node.norm)
}
fn connect_node(&mut self, idx: usize, level: usize) {
let ep = match self.entry_point {
Some(ep) => ep,
None => return,
};
let query = self.nodes[idx].embedding.clone();
let qnorm = self.nodes[idx].norm;
let mut current_ep = ep;
for lc in (level + 1..=self.max_layer).rev() {
current_ep = self.greedy_closest(&query, qnorm, current_ep, lc);
}
let insert_top = level.min(self.max_layer);
let mut ep_candidates = vec![current_ep];
for lc in (0..=insert_top).rev() {
let max_m = if lc == 0 { self.m_max0 } else { self.m };
let ef = self.ef_construction;
let nearest = self.search_layer(&query, qnorm, &ep_candidates, ef, lc, Some(idx));
let selected: Vec<usize> = nearest.iter().take(max_m).map(|c| c.idx).collect();
if lc == 0 {
for old_i in 0..self.nodes[idx].neighbors[0].len() {
let old = self.nodes[idx].neighbors[0][old_i];
self.incoming0[old] = self.incoming0[old].saturating_sub(1);
}
}
self.nodes[idx].neighbors[lc] = selected.clone();
if lc == 0 {
for &t in &selected {
self.incoming0[t] += 1;
}
}
for &neighbor_idx in &selected {
if self.nodes[neighbor_idx].neighbors.len() > lc {
self.nodes[neighbor_idx].neighbors[lc].push(idx);
if lc == 0 {
self.incoming0[idx] += 1;
}
if self.nodes[neighbor_idx].neighbors[lc].len() > max_m {
self.prune_neighbors(neighbor_idx, lc, max_m);
}
}
}
ep_candidates = selected;
if ep_candidates.is_empty() {
ep_candidates = vec![current_ep];
}
}
}
fn prune_neighbors(&mut self, node_idx: usize, layer: usize, max_m: usize) {
let node_emb = self.nodes[node_idx].embedding.clone();
let node_norm = self.nodes[node_idx].norm;
let mut neighbors_with_dist: Vec<(usize, f64)> = self.nodes[node_idx].neighbors[layer]
.iter()
.filter(|&&n| !self.nodes[n].tombstoned)
.map(|&n| (n, self.dist_to(&node_emb, node_norm, n)))
.collect();
neighbors_with_dist.sort_by(|a, b| a.1.total_cmp(&b.1));
if layer == 0 {
for old_i in 0..self.nodes[node_idx].neighbors[0].len() {
let old = self.nodes[node_idx].neighbors[0][old_i];
if self.nodes[old].tombstoned {
self.incoming0[old] = self.incoming0[old].saturating_sub(1);
}
}
let mut kept: Vec<usize> = Vec::with_capacity(max_m + 2);
for (i, &(n, _)) in neighbors_with_dist.iter().enumerate() {
if i < max_m || self.incoming0[n] <= 1 {
kept.push(n);
} else {
self.incoming0[n] = self.incoming0[n].saturating_sub(1);
}
}
self.nodes[node_idx].neighbors[0] = kept;
return;
}
neighbors_with_dist.truncate(max_m);
self.nodes[node_idx].neighbors[layer] =
neighbors_with_dist.iter().map(|&(n, _)| n).collect();
}
pub fn search(&self, query: &[f32], k: usize) -> Result<Vec<(String, f64)>> {
if self.active_count == 0 || self.entry_point.is_none() {
return Ok(vec![]);
}
let ep = self.entry_point.unwrap();
let mut current_ep = ep;
let qnorm = norm_f64(query);
for lc in (1..=self.max_layer).rev() {
current_ep = self.greedy_closest(query, qnorm, current_ep, lc);
}
let ef = self.ef_search.max(k * 2);
let nearest = self.search_layer(query, qnorm, &[current_ep], ef, 0, None);
let mut results: Vec<(String, f64)> = Vec::with_capacity(k);
for c in &nearest {
if !self.nodes[c.idx].tombstoned {
results.push((self.idx_to_rid[c.idx].clone(), c.distance));
if results.len() >= k {
break;
}
}
}
Ok(results)
}
pub fn ensure_all_reachable(&mut self) -> usize {
let Some(ep) = self.entry_point else {
return 0;
};
let n = self.nodes.len();
let mut restore: Vec<(usize, usize)> = Vec::new();
for u in 0..n {
if self.nodes[u].neighbors.is_empty() {
continue;
}
for &v in &self.nodes[u].neighbors[0] {
if v < n && !self.nodes[v].neighbors[0].contains(&u) {
restore.push((v, u));
}
}
}
for (v, u) in restore {
if !self.nodes[v].neighbors[0].contains(&u) {
self.nodes[v].neighbors[0].push(u);
self.incoming0[u] += 1;
}
}
let mut seen = vec![false; n];
let mut stack = vec![ep];
seen[ep] = true;
while let Some(cur) = stack.pop() {
if let Some(nbrs) = self.nodes[cur].neighbors.first() {
for &n in nbrs {
if n < self.nodes.len() && !seen[n] {
seen[n] = true;
stack.push(n);
}
}
}
}
let mut rescued = 0;
for i in 0..self.nodes.len() {
if seen[i] || self.nodes[i].tombstoned {
continue;
}
let emb = self.nodes[i].embedding.clone();
let norm = self.nodes[i].norm;
let mut best: Option<(usize, f64)> = None;
for (j, seen_j) in seen.iter().enumerate() {
if !seen_j {
continue;
}
let d = self.dist_to(&emb, norm, j);
if best.is_none_or(|(_, bd)| d < bd) {
best = Some((j, d));
}
}
let Some((j, _)) = best else {
continue;
};
self.nodes[j].neighbors[0].push(i);
self.incoming0[i] += 1;
self.nodes[i].neighbors[0].push(j);
self.incoming0[j] += 1;
rescued += 1;
seen[i] = true;
let mut st = vec![i];
while let Some(c) = st.pop() {
if let Some(nbrs) = self.nodes[c].neighbors.first() {
for &n in nbrs {
if n < self.nodes.len() && !seen[n] {
seen[n] = true;
st.push(n);
}
}
}
}
}
rescued
}
pub fn remove(&mut self, rid: &str) -> bool {
if let Some(&idx) = self.rid_to_idx.get(rid) {
if !self.nodes[idx].tombstoned {
self.nodes[idx].tombstoned = true;
self.active_count -= 1;
self.free_list.push(idx);
return true;
}
}
false
}
pub fn clear(&mut self) {
self.nodes.clear();
self.rid_to_idx.clear();
self.idx_to_rid.clear();
self.free_list.clear();
self.entry_point = None;
self.max_layer = 0;
self.active_count = 0;
self.incoming0.clear();
}
fn random_level(&mut self) -> usize {
let r: f64 = self.rng.gen();
let level = (-r.ln() * self.ml).floor() as usize;
level.min(32) }
fn greedy_closest(&self, query: &[f32], qnorm: f64, entry: usize, layer: usize) -> usize {
let mut current = entry;
let mut current_dist = self.dist_to(query, qnorm, current);
loop {
let mut changed = false;
if layer < self.nodes[current].neighbors.len() {
for &neighbor in &self.nodes[current].neighbors[layer] {
if neighbor >= self.nodes.len() || self.nodes[neighbor].tombstoned {
continue;
}
let dist = self.dist_to(query, qnorm, neighbor);
if dist < current_dist {
current = neighbor;
current_dist = dist;
changed = true;
}
}
}
if !changed {
break;
}
}
current
}
fn search_layer(
&self,
query: &[f32],
qnorm: f64,
entry_points: &[usize],
ef: usize,
layer: usize,
exclude_idx: Option<usize>,
) -> Vec<Candidate> {
let mut visited = HashSet::new();
let mut candidates = BinaryHeap::new();
let mut results = BinaryHeap::<FarCandidate>::new();
for &ep in entry_points {
if ep >= self.nodes.len() || visited.contains(&ep) {
continue;
}
visited.insert(ep);
let dist = self.dist_to(query, qnorm, ep);
if exclude_idx != Some(ep) && !self.nodes[ep].tombstoned {
candidates.push(Candidate {
idx: ep,
distance: dist,
});
results.push(FarCandidate {
idx: ep,
distance: dist,
});
} else {
candidates.push(Candidate {
idx: ep,
distance: dist,
});
}
}
while let Some(closest) = candidates.pop() {
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f64::MAX);
if closest.distance > worst_dist && results.len() >= ef {
break;
}
let node = &self.nodes[closest.idx];
if layer < node.neighbors.len() {
for &neighbor in &node.neighbors[layer] {
if neighbor >= self.nodes.len() || visited.contains(&neighbor) {
continue;
}
visited.insert(neighbor);
let dist = self.dist_to(query, qnorm, neighbor);
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f64::MAX);
if dist < worst_dist || results.len() < ef {
candidates.push(Candidate {
idx: neighbor,
distance: dist,
});
if exclude_idx != Some(neighbor) && !self.nodes[neighbor].tombstoned {
results.push(FarCandidate {
idx: neighbor,
distance: dist,
});
if results.len() > ef {
results.pop(); }
}
}
}
}
}
let mut sorted: Vec<Candidate> = results
.into_iter()
.map(|fc| Candidate {
idx: fc.idx,
distance: fc.distance,
})
.collect();
sorted.sort_by(|a, b| a.distance.total_cmp(&b.distance));
sorted
}
}
pub struct BruteForceIndex {
dim: usize,
entries: Vec<(String, Vec<f32>, bool)>, rid_to_idx: HashMap<String, usize>,
}
impl BruteForceIndex {
pub fn new(dim: usize) -> Self {
Self {
dim,
entries: Vec::new(),
rid_to_idx: HashMap::new(),
}
}
pub fn insert(&mut self, rid: &str, embedding: &[f32]) {
assert_eq!(embedding.len(), self.dim);
if let Some(&idx) = self.rid_to_idx.get(rid) {
self.entries[idx].1 = embedding.to_vec();
self.entries[idx].2 = false;
} else {
let idx = self.entries.len();
self.entries
.push((rid.to_string(), embedding.to_vec(), false));
self.rid_to_idx.insert(rid.to_string(), idx);
}
}
pub fn remove(&mut self, rid: &str) -> bool {
if let Some(&idx) = self.rid_to_idx.get(rid) {
if !self.entries[idx].2 {
self.entries[idx].2 = true;
return true;
}
}
false
}
pub fn search(&self, query: &[f32], k: usize) -> Vec<(String, f64)> {
let mut scored: Vec<(String, f64)> = self
.entries
.iter()
.filter(|(_, _, tombstoned)| !tombstoned)
.map(|(rid, emb, _)| (rid.clone(), cosine_distance(query, emb)))
.collect();
scored.sort_by(|a, b| a.1.total_cmp(&b.1));
scored.truncate(k);
scored
}
pub fn dim(&self) -> usize {
self.dim
}
pub fn len(&self) -> usize {
self.entries.iter().filter(|(_, _, t)| !t).count()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn seeded_construction_is_deterministic() {
let build = |seed: u64| {
let mut idx = HnswIndex::with_params_seeded(16, 16, 200, 200, seed);
for i in 0..200 {
idx.insert(&format!("rid-{i}"), &vec_seed(i as f32, 16))
.unwrap();
}
idx.search(&vec_seed(7.0, 16), 10).unwrap()
};
let a = build(42);
let b = build(42);
assert_eq!(a, b, "same seed must reproduce identical results");
}
#[test]
fn default_construction_is_deterministic_across_opens() {
let build = || {
let mut idx = HnswIndex::new(16);
for i in 0..200 {
idx.insert(&format!("rid-{i:03}"), &vec_seed(i as f32, 16))
.unwrap();
}
idx
};
let (a, b) = (build(), build());
for q in 0..8 {
let qa = a.search(&vec_seed(q as f32 * 3.7, 16), 100).unwrap();
let qb = b.search(&vec_seed(q as f32 * 3.7, 16), 100).unwrap();
assert_eq!(
qa, qb,
"query {q}: fresh default constructions disagree — the \
level RNG is drawing per-instance entropy again"
);
}
}
#[test]
fn every_insert_is_reachable_by_its_own_vector() {
for round in 0..5u64 {
let mut idx = HnswIndex::with_params_seeded(16, 16, 200, 200, round * 7919 + 1);
for i in 0..120 {
let mut v = vec_seed(1.0, 16);
v[i % 16] += 0.001 * ((i as f32) + 1.0);
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
let v: Vec<f32> = v.iter().map(|x| x / norm).collect();
idx.insert(&format!("dense-{round}-{i}"), &v).unwrap();
}
for i in 0..30 {
idx.insert(&format!("far-{round}-{i}"), &vec_seed(100.0 + i as f32, 16))
.unwrap();
}
idx.ensure_all_reachable();
assert_eq!(
idx.ensure_all_reachable(),
0,
"round {round}: repair must be idempotent"
);
let rids: Vec<String> = idx.idx_to_rid.clone();
for (i, rid) in rids.iter().enumerate() {
let emb = idx.nodes[i].embedding.clone();
let hits = idx.search(&emb, 5).unwrap();
assert!(
hits.iter().any(|(r, _)| r == rid),
"round {round}: {rid} is stored but unreachable by its \
own vector — orphaned by pruning"
);
}
}
}
#[test]
fn cosine_distance_guards_nan_and_zero_norms() {
let nan_vec = vec![f32::NAN, 1.0, 0.0];
let finite = vec![1.0f32, 0.0, 0.0];
let zero = vec![0.0f32, 0.0, 0.0];
let d_nan = cosine_distance(&nan_vec, &finite);
assert!(
d_nan.is_finite(),
"NaN embedding must not yield NaN distance"
);
assert_eq!(d_nan, 1.0);
assert_eq!(cosine_distance(&finite, &nan_vec), 1.0);
assert_eq!(cosine_distance(&zero, &finite), 1.0);
let d = cosine_distance(&finite, &finite);
assert!(d.is_finite() && (0.0..=2.0).contains(&d));
}
fn vec_seed(seed: f32, dim: usize) -> Vec<f32> {
let raw: Vec<f32> = (0..dim)
.map(|i| ((seed + i as f32) * 0.7123 + (i as f32) * 0.3171).sin())
.collect();
let norm: f32 = raw.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm == 0.0 {
return vec![1.0 / (dim as f32).sqrt(); dim];
}
raw.iter().map(|x| x / norm).collect()
}
#[test]
fn test_cosine_distance_identical() {
let v = vec_seed(1.0, 8);
let d = cosine_distance(&v, &v);
assert!(d.abs() < 1e-6, "distance to self should be ~0, got {d}");
}
#[test]
fn test_cosine_distance_orthogonal() {
let a = vec![1.0f32, 0.0, 0.0, 0.0];
let b = vec![0.0f32, 1.0, 0.0, 0.0];
let d = cosine_distance(&a, &b);
assert!(
(d - 1.0).abs() < 1e-6,
"orthogonal distance should be ~1, got {d}"
);
}
#[test]
fn test_empty_index() {
let index = HnswIndex::new(8);
let results = index.search(&vec_seed(1.0, 8), 10).unwrap();
assert!(results.is_empty());
assert_eq!(index.len(), 0);
assert!(index.is_empty());
}
#[test]
fn test_single_insert_search() {
let mut index = HnswIndex::new(8);
index.insert("a", &vec_seed(1.0, 8)).unwrap();
assert_eq!(index.len(), 1);
let results = index.search(&vec_seed(1.0, 8), 10).unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, "a");
assert!(results[0].1 < 1e-6); }
#[test]
fn test_insert_search_nearest() {
let dim = 64;
let mut index = HnswIndex::new(dim);
for i in 0..100 {
index
.insert(&format!("v{i}"), &vec_seed(i as f32 * 0.37, dim))
.unwrap();
}
assert_eq!(index.len(), 100);
let query = vec_seed(0.0, dim);
let results = index.search(&query, 1).unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, "v0");
}
#[test]
fn test_tombstone_excludes_from_search() {
let dim = 8;
let mut index = HnswIndex::new(dim);
index.insert("a", &vec_seed(1.0, dim)).unwrap();
index.insert("b", &vec_seed(2.0, dim)).unwrap();
assert_eq!(index.len(), 2);
assert!(index.remove("a"));
assert_eq!(index.len(), 1);
let results = index.search(&vec_seed(1.0, dim), 10).unwrap();
assert!(!results.iter().any(|(rid, _)| rid == "a"));
}
#[test]
fn test_remove_nonexistent() {
let mut index = HnswIndex::new(8);
assert!(!index.remove("nonexistent"));
}
#[test]
fn test_free_list_reuse() {
let dim = 8;
let mut index = HnswIndex::new(dim);
index.insert("a", &vec_seed(1.0, dim)).unwrap();
let initial_nodes = index.nodes.len();
index.remove("a");
index.insert("b", &vec_seed(2.0, dim)).unwrap();
assert_eq!(index.nodes.len(), initial_nodes);
assert_eq!(index.len(), 1);
let results = index.search(&vec_seed(2.0, dim), 10).unwrap();
assert_eq!(results[0].0, "b");
}
#[test]
fn test_clear() {
let dim = 8;
let mut index = HnswIndex::new(dim);
for i in 0..50 {
index
.insert(&format!("v{i}"), &vec_seed(i as f32, dim))
.unwrap();
}
assert_eq!(index.len(), 50);
index.clear();
assert_eq!(index.len(), 0);
assert!(index.is_empty());
assert!(index.search(&vec_seed(1.0, dim), 10).unwrap().is_empty());
}
#[test]
fn test_duplicate_insert_updates() {
let dim = 8;
let mut index = HnswIndex::new(dim);
index.insert("a", &vec_seed(1.0, dim)).unwrap();
index.insert("a", &vec_seed(2.0, dim)).unwrap();
assert_eq!(index.len(), 1);
let results = index.search(&vec_seed(2.0, dim), 1).unwrap();
assert_eq!(results[0].0, "a");
assert!(results[0].1 < 0.01);
}
#[test]
fn test_resurrect_tombstoned() {
let dim = 8;
let mut index = HnswIndex::new(dim);
index.insert("a", &vec_seed(1.0, dim)).unwrap();
index.remove("a");
assert_eq!(index.len(), 0);
index.insert("a", &vec_seed(2.0, dim)).unwrap();
assert_eq!(index.len(), 1);
let results = index.search(&vec_seed(2.0, dim), 1).unwrap();
assert_eq!(results[0].0, "a");
}
#[test]
fn test_recall_quality_dim64() {
let dim = 64;
let n = 1000;
let k = 10;
let mut hnsw = HnswIndex::with_params(dim, 16, 200, 50);
let mut brute = BruteForceIndex::new(dim);
for i in 0..n {
let emb = vec_seed(i as f32 * 0.37, dim);
hnsw.insert(&format!("v{i}"), &emb).unwrap();
brute.insert(&format!("v{i}"), &emb);
}
let mut total_recall = 0.0;
let num_queries = 20;
for q in 0..num_queries {
let query = vec_seed(q as f32 * 7.13 + 100.0, dim);
let hnsw_results: HashSet<String> = hnsw
.search(&query, k)
.unwrap()
.into_iter()
.map(|(rid, _)| rid)
.collect();
let brute_results: HashSet<String> = brute
.search(&query, k)
.into_iter()
.map(|(rid, _)| rid)
.collect();
let intersection = hnsw_results.intersection(&brute_results).count();
total_recall += intersection as f64 / k as f64;
}
let avg_recall = total_recall / num_queries as f64;
assert!(
avg_recall > 0.90,
"recall@{k} should be > 0.90, got {avg_recall:.3}"
);
}
#[test]
fn test_recall_quality_dim384() {
let dim = 384;
let n = 500;
let k = 10;
let mut hnsw = HnswIndex::with_params(dim, 16, 200, 50);
let mut brute = BruteForceIndex::new(dim);
for i in 0..n {
let emb = vec_seed(i as f32 * 0.37, dim);
hnsw.insert(&format!("v{i}"), &emb).unwrap();
brute.insert(&format!("v{i}"), &emb);
}
let mut total_recall = 0.0;
let num_queries = 10;
for q in 0..num_queries {
let query = vec_seed(q as f32 * 7.13 + 100.0, dim);
let hnsw_results: HashSet<String> = hnsw
.search(&query, k)
.unwrap()
.into_iter()
.map(|(rid, _)| rid)
.collect();
let brute_results: HashSet<String> = brute
.search(&query, k)
.into_iter()
.map(|(rid, _)| rid)
.collect();
let intersection = hnsw_results.intersection(&brute_results).count();
total_recall += intersection as f64 / k as f64;
}
let avg_recall = total_recall / num_queries as f64;
assert!(
avg_recall > 0.85,
"recall@{k} at dim=384 should be > 0.85, got {avg_recall:.3}"
);
}
#[test]
fn test_search_results_sorted_by_distance() {
let dim = 64;
let mut index = HnswIndex::new(dim);
for i in 0..200 {
index
.insert(&format!("v{i}"), &vec_seed(i as f32 * 0.37, dim))
.unwrap();
}
let query = vec_seed(999.0, dim);
let results = index.search(&query, 20).unwrap();
for i in 1..results.len() {
assert!(
results[i - 1].1 <= results[i].1 + 1e-10,
"results not sorted: {} > {}",
results[i - 1].1,
results[i].1
);
}
}
#[test]
fn test_large_insert_search() {
let dim = 64;
let n = 5000;
let mut index = HnswIndex::new(dim);
for i in 0..n {
index
.insert(&format!("v{i}"), &vec_seed(i as f32 * 0.37, dim))
.unwrap();
}
assert_eq!(index.len(), n);
let results = index.search(&vec_seed(999.0, dim), 10).unwrap();
assert_eq!(results.len(), 10);
}
#[test]
fn test_search_with_many_tombstones() {
let dim = 32;
let mut index = HnswIndex::new(dim);
for i in 0..100 {
index
.insert(&format!("v{i}"), &vec_seed(i as f32, dim))
.unwrap();
}
for i in 0..90 {
index.remove(&format!("v{i}"));
}
assert_eq!(index.len(), 10);
let results = index.search(&vec_seed(95.0, dim), 5).unwrap();
for (rid, _) in &results {
let num: usize = rid[1..].parse().unwrap();
assert!(num >= 90, "got tombstoned result {rid}");
}
}
#[test]
fn test_brute_force_index() {
let dim = 8;
let mut bf = BruteForceIndex::new(dim);
bf.insert("a", &vec_seed(1.0, dim));
bf.insert("b", &vec_seed(2.0, dim));
bf.insert("c", &vec_seed(3.0, dim));
assert_eq!(bf.len(), 3);
bf.remove("b");
assert_eq!(bf.len(), 2);
let results = bf.search(&vec_seed(1.0, dim), 2);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, "a"); }
}