use super::*;
#[test]
fn test_create_index() {
let index = HNSWIndex::new(128, 16, 16).unwrap();
assert_eq!(index.dimension, 128);
assert_eq!(index.num_vectors, 0);
}
#[test]
fn neighbor_list_holds_default_base_degree_inline() {
let neighbors: NeighborList = (0..32).collect();
assert_eq!(neighbors.len(), 32);
assert!(!neighbors.spilled());
}
#[test]
fn test_add_vectors() {
let mut index = HNSWIndex::new(3, 16, 16).unwrap();
index.add(0, vec![1.0, 0.0, 0.0]).unwrap();
index.add(1, vec![0.0, 1.0, 0.0]).unwrap();
assert_eq!(index.num_vectors, 2);
}
#[test]
fn test_dimension_mismatch() {
let mut index = HNSWIndex::new(3, 16, 16).unwrap();
let result = index.add(0, vec![1.0, 0.0]); assert!(result.is_err());
}
fn build_test_index() -> (HNSWIndex, Vec<f32>) {
let dim = 32;
let n = 200;
let mut index = HNSWIndex::new(dim, 16, 32).unwrap();
let mut seed: u64 = 42;
let mut next = || -> f32 {
seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1);
((seed >> 33) as f32) / (u32::MAX as f32) - 0.5
};
for i in 0..n {
let mut v: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
v.iter_mut().for_each(|x| *x /= norm);
}
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
let mut q: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = q.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
q.iter_mut().for_each(|x| *x /= norm);
}
(index, q)
}
#[test]
fn test_search_adaptive_conservative_matches_search() {
let (index, q) = build_test_index();
let k = 5;
let ef = 64;
let baseline = index.search(&q, k, ef).unwrap();
let config = crate::adaptive::AdaptiveConfig::conservative();
let (adaptive, _evaluated) = index.search_adaptive(&q, k, ef, &config).unwrap();
assert_eq!(adaptive.len(), baseline.len());
assert!(
adaptive.windows(2).all(|w| w[0].1 <= w[1].1),
"adaptive results should stay sorted by ascending distance",
);
assert_eq!(adaptive[0].0, baseline[0].0);
}
#[test]
fn test_search_adaptive_aggressive_fewer_evaluations() {
let (index, q) = build_test_index();
let k = 5;
let ef = 64;
let conservative = crate::adaptive::AdaptiveConfig::conservative();
let aggressive = crate::adaptive::AdaptiveConfig::aggressive();
let (_results_c, evaluated_c) = index.search_adaptive(&q, k, ef, &conservative).unwrap();
let (_results_a, evaluated_a) = index.search_adaptive(&q, k, ef, &aggressive).unwrap();
assert!(
evaluated_a <= evaluated_c,
"aggressive ({}) should evaluate <= conservative ({})",
evaluated_a,
evaluated_c,
);
}
#[test]
fn test_search_acorn_filters_by_metadata() {
let dim = 8;
let mut index = HNSWIndex::with_filtering(dim, 8, 16, "group").unwrap();
let mut doc_ids = Vec::new();
let mut vectors = Vec::new();
for i in 0..40_u32 {
let doc_id = 10_000 + i * 7;
let mut v = vec![0.0; dim];
v[(i as usize) % dim] = 1.0;
v[((i as usize) * 3 + 1) % dim] += 0.1;
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
for x in &mut v {
*x /= norm;
}
index.add(doc_id, v.clone()).unwrap();
doc_ids.push(doc_id);
vectors.push(v);
}
index.build().unwrap();
for (i, &doc_id) in doc_ids.iter().enumerate() {
let mut metadata = crate::filtering::DocumentMetadata::new();
metadata.insert(
"group".to_string(),
crate::filtering::MetadataValue::from((i as u32) % 2),
);
index.add_metadata(doc_id, metadata).unwrap();
}
let filter = crate::filtering::MetadataFilter::equals("group", 0_u32);
let config = crate::hnsw::AcornConfig {
ef_search: 32,
..Default::default()
};
let (results, _stats) = index
.search_acorn_with_stats(&vectors[2], 5, &config, &filter)
.unwrap();
assert!(!results.is_empty());
assert!(results.iter().all(|(doc_id, _)| doc_ids.contains(doc_id)));
assert!(
results
.iter()
.all(|(doc_id, _)| ((doc_id - 10_000) / 7) % 2 == 0),
"all ACORN results should satisfy the metadata filter: {results:?}"
);
}
#[test]
fn test_search_acorn_requires_filtering_metadata() {
let (index, q) = build_test_index();
let filter = crate::filtering::MetadataFilter::equals("group", 0_u32);
let err = index
.search_acorn(&q, 5, &crate::hnsw::AcornConfig::default(), &filter)
.unwrap_err();
assert!(matches!(err, RetrieveError::InvalidParameter(_)));
}
#[cfg(feature = "serde")]
#[test]
fn test_hnsw_save_load_roundtrip() {
let (index, q) = build_test_index();
let k = 5;
let ef = 64;
let original_results = index.search(&q, k, ef).unwrap();
let mut buf = Vec::new();
index.save_to_writer(&mut buf).unwrap();
assert!(!buf.is_empty(), "serialized output should not be empty");
let loaded = HNSWIndex::load_from_reader(buf.as_slice()).unwrap();
assert_eq!(loaded.dimension, index.dimension);
assert_eq!(loaded.num_vectors, index.num_vectors);
assert!(loaded.is_built());
assert_eq!(loaded.doc_ids, index.doc_ids);
assert_eq!(
loaded.doc_id_to_internal.len(),
index.doc_id_to_internal.len()
);
assert_eq!(loaded.entry_point(), index.entry_point());
let loaded_results = loaded.search(&q, k, ef).unwrap();
assert_eq!(
loaded_results, original_results,
"search results should be identical after save/load roundtrip"
);
#[cfg(any(feature = "persistence", feature = "store"))]
{
let postcard = index.to_postcard().unwrap();
let loaded = HNSWIndex::from_postcard(&postcard).unwrap();
assert_eq!(loaded.entry_point(), index.entry_point());
assert_eq!(loaded.search(&q, k, ef).unwrap(), original_results);
}
}
#[cfg(feature = "serde")]
#[test]
fn test_hnsw_save_load_preserves_tombstones() {
let (mut index, q) = build_test_index();
let k = 20;
let ef = 64;
let nearest_id = index.search(&q, 1, ef).unwrap()[0].0;
index.delete(nearest_id).unwrap();
assert!(index.is_deleted(nearest_id));
let mut buf = Vec::new();
index.save_to_writer(&mut buf).unwrap();
let loaded = HNSWIndex::load_from_reader(buf.as_slice()).unwrap();
assert!(loaded.is_deleted(nearest_id));
assert_eq!(loaded.num_active(), index.num_active());
let loaded_ids: Vec<_> = loaded
.search(&q, k, ef)
.unwrap()
.into_iter()
.map(|(id, _)| id)
.collect();
assert!(
!loaded_ids.contains(&nearest_id),
"serde roundtrip resurrected deleted doc_id {}",
nearest_id
);
#[cfg(any(feature = "persistence", feature = "store"))]
{
let postcard = index.to_postcard().unwrap();
let loaded = HNSWIndex::from_postcard(&postcard).unwrap();
assert!(loaded.is_deleted(nearest_id));
assert_eq!(loaded.num_active(), index.num_active());
let loaded_ids: Vec<_> = loaded
.search(&q, k, ef)
.unwrap()
.into_iter()
.map(|(id, _)| id)
.collect();
assert!(
!loaded_ids.contains(&nearest_id),
"postcard roundtrip resurrected deleted doc_id {}",
nearest_id
);
}
}
#[cfg(feature = "serde")]
#[test]
fn test_hnsw_save_to_file_roundtrip_and_no_temp_leftover() {
let (index, q) = build_test_index();
let k = 5;
let ef = 64;
let original_results = index.search(&q, k, ef).unwrap();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("index.json");
index.save_to_file(&path).unwrap();
assert!(path.exists(), "saved file should exist at target path");
let tmp_path = dir.path().join("index.json.tmp");
assert!(
!tmp_path.exists(),
"save_to_file left a .tmp sibling: {}",
tmp_path.display()
);
let loaded = HNSWIndex::load_from_file(&path).unwrap();
let loaded_results = loaded.search(&q, k, ef).unwrap();
assert_eq!(loaded_results, original_results);
}
#[cfg(feature = "serde")]
#[test]
fn test_hnsw_save_to_file_overwrites_existing_atomically() {
let (index, _q) = build_test_index();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("index.json");
std::fs::write(&path, b"{\"sentinel\": true}").unwrap();
let pre_size = std::fs::metadata(&path).unwrap().len();
index.save_to_file(&path).unwrap();
let post_size = std::fs::metadata(&path).unwrap().len();
assert!(
post_size > pre_size,
"post-save size {} should exceed sentinel size {}",
post_size,
pre_size
);
let _ = HNSWIndex::load_from_file(&path).expect("post-overwrite file must load");
assert!(!dir.path().join("index.json.tmp").exists());
}
#[cfg(feature = "serde")]
#[test]
fn test_hnsw_save_to_file_cleans_up_tmp_on_rename_failure() {
let (index, _q) = build_test_index();
let dir = tempfile::tempdir().unwrap();
let target = dir.path().join("save.json");
std::fs::create_dir(&target).unwrap();
std::fs::write(target.join("blocker"), b"x").unwrap();
let result = index.save_to_file(&target);
assert!(
result.is_err(),
"rename into a non-empty directory should fail"
);
let tmp_path = dir.path().join("save.json.tmp");
assert!(
!tmp_path.exists(),
"save_to_file failed but left a temp file at {} \
(rename-failure cleanup branch did not fire)",
tmp_path.display(),
);
}
#[cfg(feature = "serde")]
#[test]
#[ignore = "measurement only; run with --release --ignored --nocapture"]
fn save_to_file_overhead_probe() {
use std::time::Instant;
let (index, _q) = build_test_index();
let dir = tempfile::tempdir().unwrap();
let iters = 30;
let path = dir.path().join("atomic.json");
let t = Instant::now();
for _ in 0..iters {
index.save_to_file(&path).unwrap();
}
let atomic_ms_per = t.elapsed().as_secs_f64() * 1000.0 / iters as f64;
let path2 = dir.path().join("naive.json");
let t = Instant::now();
for _ in 0..iters {
let f = std::fs::File::create(&path2).unwrap();
index.save_to_writer(std::io::BufWriter::new(f)).unwrap();
}
let naive_ms_per = t.elapsed().as_secs_f64() * 1000.0 / iters as f64;
let file_kb = std::fs::metadata(&path).unwrap().len() / 1024;
println!("# save_to_file_overhead_probe");
println!(
"# build_test_index ({} vectors, {} dim, ~{} KB JSON)",
index.num_vectors, index.dimension, file_kb
);
println!("# atomic {:.2} ms/op", atomic_ms_per);
println!("# naive {:.2} ms/op", naive_ms_per);
println!(
"# overhead {:.2} ms/op ({:+.0}%)",
atomic_ms_per - naive_ms_per,
100.0 * (atomic_ms_per - naive_ms_per) / naive_ms_per
);
}
#[cfg(feature = "serde")]
#[test]
fn test_hnsw_build_save_reload_search_equality() {
let (index, q) = build_test_index();
let k = 10;
let ef = 64;
let original_results = index.search(&q, k, ef).unwrap();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("index.json");
index.save_to_file(&path).unwrap();
drop(index);
let loaded = HNSWIndex::load_from_file(&path).expect("load_from_file roundtrip");
let loaded_results = loaded.search(&q, k, ef).unwrap();
assert_eq!(
loaded_results, original_results,
"search results diverged across path-based save/load roundtrip. \
likely cause: writer and reader formats drifted (one side \
changed without the other), or a transient field was \
dropped during serialization that affects search behavior."
);
}
#[cfg(all(feature = "serde", unix))]
#[test]
fn test_hnsw_save_to_file_uses_rename_not_truncate() {
use std::os::unix::fs::MetadataExt;
let (index, _q) = build_test_index();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("index.json");
index.save_to_file(&path).unwrap();
let ino_before = std::fs::metadata(&path).unwrap().ino();
index.save_to_file(&path).unwrap();
let ino_after = std::fs::metadata(&path).unwrap().ino();
assert_ne!(
ino_before, ino_after,
"save_to_file must replace via rename (inode changes), not truncate \
in place (inode preserved). truncate-overwrite would leave the file \
in a partially-written state if a crash interrupts the write -- this \
is exactly the partial-write corruption mode the atomic-rename refactor \
is meant to prevent."
);
}
#[test]
fn test_delete_excludes_from_results() {
let (mut index, q) = build_test_index();
let k = 10;
let ef = 64;
let before = index.search(&q, k, ef).unwrap();
assert!(!before.is_empty());
let nearest_id = before[0].0;
index.delete(nearest_id).unwrap();
let after = index.search(&q, k, ef).unwrap();
let after_ids: Vec<u32> = after.iter().map(|(id, _)| *id).collect();
assert!(
!after_ids.contains(&nearest_id),
"deleted doc_id {} should not appear in results",
nearest_id
);
}
#[test]
fn test_delete_all_returns_empty() {
let dim = 4;
let mut index = HNSWIndex::new(dim, 4, 4).unwrap();
for i in 0..5u32 {
let mut v = vec![0.0f32; dim];
v[i as usize % dim] = 1.0;
index.add(i, v).unwrap();
}
index.build().unwrap();
for i in 0..5u32 {
index.delete(i).unwrap();
}
let results = index.search(&[1.0, 0.0, 0.0, 0.0], 5, 50).unwrap();
assert!(
results.is_empty(),
"all vectors deleted, results should be empty"
);
}
#[test]
fn test_delete_nonexistent_returns_error() {
let mut index = HNSWIndex::new(4, 4, 4).unwrap();
index.add(0, vec![1.0, 0.0, 0.0, 0.0]).unwrap();
index.build().unwrap();
let result = index.delete(999);
assert!(result.is_err(), "deleting nonexistent doc_id should error");
}
#[test]
fn test_delete_idempotent() {
let (mut index, _q) = build_test_index();
index.delete(0).unwrap();
index.delete(0).unwrap();
}
#[test]
fn test_is_deleted() {
let (mut index, _q) = build_test_index();
assert!(!index.is_deleted(0));
index.delete(0).unwrap();
assert!(index.is_deleted(0));
assert!(!index.is_deleted(1));
}
#[test]
fn test_num_active() {
let (mut index, _q) = build_test_index();
let total = index.num_vectors;
assert_eq!(index.num_active(), total);
index.delete(0).unwrap();
index.delete(1).unwrap();
assert_eq!(index.num_active(), total - 2);
}
#[test]
fn test_builder_produces_working_index() {
let mut index = HNSWIndex::builder(4)
.m(8)
.ef_search(32)
.auto_normalize(true)
.build()
.unwrap();
index.add_slice(0, &[3.0, 4.0, 0.0, 0.0]).unwrap();
index.add_slice(1, &[0.0, 0.0, 3.0, 4.0]).unwrap();
index.build().unwrap();
let results = index.search(&[3.0, 4.0, 0.0, 0.0], 1, 32).unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, 0, "nearest neighbor should be doc 0");
assert!(
results[0].1 < 0.01,
"self-distance after normalization should be ~0, got {}",
results[0].1
);
}
#[cfg(feature = "id-compression")]
#[test]
fn unsupported_id_compression_method_fails_build() {
let params = HNSWParams {
m: 4,
m_max: 8,
auto_normalize: true,
seed: Some(42),
id_compression: Some(crate::compression::IdCompressionMethod::EliasFano),
compression_threshold: 0,
..Default::default()
};
let mut index = HNSWIndex::with_params(4, params).unwrap();
for i in 0..16_u32 {
let mut vector = vec![0.0; 4];
vector[i as usize % 4] = 1.0;
vector[((i as usize) + 1) % 4] = 0.5;
index.add(i, vector).unwrap();
}
let err = index.build().unwrap_err();
assert!(
err.to_string()
.contains("unsupported HNSW ID compression method"),
"unexpected error: {err}"
);
}
#[cfg(feature = "id-compression")]
#[test]
fn compression_threshold_preserves_small_neighbor_lists() {
let neighbors: Vec<NeighborList> = vec![
[1, 2].into_iter().collect(),
[0, 2, 3, 4].into_iter().collect(),
];
let mut layer = Layer::new_uncompressed(neighbors.clone());
let compressor = crate::compression::DeltaVarintCompressor::new();
layer.compress(&compressor, 8, 3).unwrap();
assert_eq!(layer.get_neighbors(0).as_ref(), neighbors[0].as_slice());
assert_eq!(layer.get_neighbors(1).as_ref(), neighbors[1].as_slice());
}
#[test]
fn test_auto_normalize_symmetric_for_angular() {
let mut index = HNSWIndex::builder(4)
.m(8)
.ef_search(32)
.metric(DistanceMetric::Angular)
.auto_normalize(true)
.build()
.unwrap();
index.add_slice(0, &[3.0, 4.0, 0.0, 0.0]).unwrap();
index.add_slice(1, &[0.0, 0.0, 3.0, 4.0]).unwrap();
index.build().unwrap();
let results = index.search(&[3.0, 4.0, 0.0, 0.0], 1, 32).unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, 0, "nearest neighbor should be doc 0");
assert!(
results[0].1 < 0.05,
"angular self-distance after normalization should be ~0, got {}",
results[0].1
);
}
#[cfg(feature = "parallel")]
#[test]
fn test_search_batch_matches_sequential() {
let (index, q) = build_test_index();
let k = 5;
let ef = 64;
let dim = index.dimension;
let mut queries_flat: Vec<f32> = Vec::new();
let num_queries = 8;
let mut seed: u64 = 99;
let mut next = || -> f32 {
seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1);
((seed >> 33) as f32) / (u32::MAX as f32) - 0.5
};
for _ in 0..num_queries {
let mut v: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
v.iter_mut().for_each(|x| *x /= norm);
}
queries_flat.extend_from_slice(&v);
}
let sequential: Vec<Vec<(u32, f32)>> = (0..num_queries)
.map(|i| {
let query = &queries_flat[i * dim..(i + 1) * dim];
index.search(query, k, ef).unwrap()
})
.collect();
let query_slices: Vec<&[f32]> = (0..num_queries)
.map(|i| &queries_flat[i * dim..(i + 1) * dim])
.collect();
let batch = index.search_batch(&query_slices, k, ef).unwrap();
let batch_flat = index
.search_batch_flat(&queries_flat, num_queries, k, ef)
.unwrap();
let _ = index.search(&q, k, ef).unwrap();
for i in 0..num_queries {
assert_eq!(
sequential[i], batch[i],
"search_batch result {} differs from sequential",
i
);
assert_eq!(
sequential[i], batch_flat[i],
"search_batch_flat result {} differs from sequential",
i
);
}
}
fn build_structural_test_index(n: usize, dim: usize, m: usize) -> HNSWIndex {
let params = HNSWParams {
m,
m_max: 2 * m,
seed: Some(42),
..Default::default()
};
let mut index = HNSWIndex::with_params(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim)
.map(|j| ((i * 7 + j * 3) % 100) as f32 / 100.0)
.collect();
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
let normed: Vec<f32> = v.iter().map(|x| x / norm).collect();
index.add(i as u32, normed).unwrap();
}
index.build().unwrap();
index
}
#[test]
fn test_layer_assignment_distribution() {
let index = build_structural_test_index(500, 32, 16);
let max_layer = *index.layer_assignments.iter().max().unwrap_or(&0) as usize;
let mut layer_counts = vec![0usize; max_layer + 1];
for &l in &index.layer_assignments {
layer_counts[l as usize] += 1;
}
let layer0_frac = layer_counts[0] as f64 / 500.0;
assert!(
layer0_frac > 0.5,
"Layer 0 fraction {:.2} should be > 0.5 (got {} of 500)",
layer0_frac,
layer_counts[0]
);
for l in 1..layer_counts.len() {
assert!(
layer_counts[l] < layer_counts[0],
"Layer {} ({}) should have fewer vectors than layer 0 ({})",
l,
layer_counts[l],
layer_counts[0]
);
}
}
#[test]
fn test_m_max_enforced() {
let m = 8;
let m_max = 2 * m;
let index = build_structural_test_index(200, 16, m);
for (layer_idx, layer) in index.layers.iter().enumerate() {
let limit = if layer_idx == 0 { m_max } else { m };
if let Some(neighbors) = layer.get_all_neighbors() {
for (node_id, nbrs) in neighbors.iter().enumerate() {
assert!(
nbrs.len() <= limit,
"Node {} at layer {} has {} neighbors, limit is {}",
node_id,
layer_idx,
nbrs.len(),
limit
);
}
}
}
}
#[test]
fn test_neighbor_ids_in_bounds() {
let index = build_structural_test_index(100, 16, 8);
for (layer_idx, layer) in index.layers.iter().enumerate() {
if let Some(neighbors) = layer.get_all_neighbors() {
for (node_id, nbrs) in neighbors.iter().enumerate() {
for &nbr in nbrs.iter() {
assert!(
(nbr as usize) < index.num_vectors,
"Node {} at layer {} has out-of-bounds neighbor {}",
node_id,
layer_idx,
nbr
);
}
}
}
}
}
#[test]
fn test_layer_assignment_matches_layers() {
let index = build_structural_test_index(100, 16, 8);
for (node_id, &assigned_layer) in index.layer_assignments.iter().enumerate() {
for l in 0..=assigned_layer as usize {
assert!(
l < index.layers.len(),
"Node {} assigned to layer {} but only {} layers exist",
node_id,
assigned_layer,
index.layers.len()
);
}
}
}
#[test]
fn test_batch_search_mqo_order_and_recall() {
let (index, _) = build_test_index();
let k = 5;
let ef = 64;
let dim = index.dimension;
let mut seed: u64 = 99;
let mut next = || -> f32 {
seed = seed
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((seed >> 33) as f32) / (u32::MAX as f32) - 0.5
};
let num_queries = 10;
let queries_owned: Vec<Vec<f32>> = (0..num_queries)
.map(|_| {
let mut v: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
v.iter_mut().for_each(|x| *x /= norm);
}
v
})
.collect();
let queries: Vec<&[f32]> = queries_owned.iter().map(|v| v.as_slice()).collect();
let individual: Vec<Vec<(u32, f32)>> = queries
.iter()
.map(|q| index.search(q, k, ef).unwrap())
.collect();
let batch = index.batch_search_mqo(&queries, k, ef).unwrap();
assert_eq!(batch.len(), num_queries);
for (i, (ind, mqo)) in individual.iter().zip(batch.iter()).enumerate() {
assert!(
!mqo.is_empty(),
"query {}: batch_search_mqo returned no results",
i
);
let ind_best_dist = ind.first().map(|(_, d)| *d).unwrap_or(f32::INFINITY);
let mqo_best_dist = mqo.first().map(|(_, d)| *d).unwrap_or(f32::INFINITY);
assert!(
mqo_best_dist <= ind_best_dist + 1e-5,
"query {}: MQO nearest dist {} > individual nearest dist {} (regression)",
i,
mqo_best_dist,
ind_best_dist,
);
}
}
#[test]
fn test_batch_search_mqo_single_query() {
let (index, q) = build_test_index();
let k = 5;
let ef = 64;
let single = index.search(&q, k, ef).unwrap();
let batch = index.batch_search_mqo(&[q.as_slice()], k, ef).unwrap();
assert_eq!(batch.len(), 1);
assert_eq!(batch[0], single, "single-query MQO should match search");
}
#[test]
fn test_search_with_distance_matches_standard() {
let (index, q) = build_test_index();
let k = 5;
let ef = 64;
let standard = index.search(&q, k, ef).unwrap();
let vectors = &index.vectors;
let dim = index.dimension;
let dist_fn = |query: &[f32], internal_id: u32| -> f32 {
let start = internal_id as usize * dim;
let vec = &vectors[start..start + dim];
crate::distance::cosine_distance_normalized(query, vec)
};
let custom = index.search_with_distance(&q, k, ef, &dist_fn).unwrap();
assert_eq!(standard.len(), custom.len());
for (s, c) in standard.iter().zip(custom.iter()) {
assert_eq!(s.0, c.0, "doc_ids should match");
assert!((s.1 - c.1).abs() < 1e-6, "distances should match");
}
}
#[test]
fn test_search_with_distance_custom_metric() {
let (index, q) = build_test_index();
let k = 5;
let ef = 64;
let vectors = &index.vectors;
let dim = index.dimension;
let l2_dist = |query: &[f32], internal_id: u32| -> f32 {
let start = internal_id as usize * dim;
let vec = &vectors[start..start + dim];
crate::distance::l2_distance(query, vec)
};
let results = index.search_with_distance(&q, k, ef, &l2_dist).unwrap();
for w in results.windows(2) {
assert!(w[0].1 <= w[1].1, "results not sorted by custom distance");
}
}
#[test]
fn test_delete_with_repair_excludes_from_results() {
let (mut index, q) = build_test_index();
let k = 10;
let ef = 64;
let before = index.search(&q, k, ef).unwrap();
let nearest_id = before[0].0;
let repairs = index.delete_with_repair(nearest_id).unwrap();
let after = index.search(&q, k, ef).unwrap();
let after_ids: Vec<u32> = after.iter().map(|(id, _)| *id).collect();
assert!(
!after_ids.contains(&nearest_id),
"deleted doc_id {} should not appear in results (repairs={})",
nearest_id,
repairs,
);
}
#[test]
fn test_delete_with_repair_maintains_recall() {
let (mut index, q) = build_test_index();
let k = 10;
let ef = 100;
let before = index.search(&q, k, ef).unwrap();
assert_eq!(before.len(), k);
for i in 0..40u32 {
index.delete_with_repair(i).unwrap();
}
let after = index.search(&q, k, ef).unwrap();
assert_eq!(
after.len(),
k,
"should still return k={} results after deleting 20%",
k
);
for (id, _) in &after {
assert!(*id >= 40, "deleted id {} appeared in results", id);
}
}
#[test]
fn test_delete_with_repair_recall_floor_against_ground_truth() {
let dim = 32usize;
let n = 200u32;
let deleted_count = 40u32;
let k = 10usize;
let ef = 100usize;
let recall_floor = 0.7f32;
let mut seed: u64 = 42;
let mut next = || -> f32 {
seed = seed.wrapping_mul(6_364_136_223_846_793_005).wrapping_add(1);
((seed >> 33) as f32) / (u32::MAX as f32) - 0.5
};
let mut vectors: Vec<Vec<f32>> = Vec::with_capacity(n as usize);
for _ in 0..n {
let mut v: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
v.iter_mut().for_each(|x| *x /= norm);
}
vectors.push(v);
}
let mut q: Vec<f32> = (0..dim).map(|_| next()).collect();
let qnorm = q.iter().map(|x| x * x).sum::<f32>().sqrt();
if qnorm > 0.0 {
q.iter_mut().for_each(|x| *x /= qnorm);
}
let mut ground: Vec<(u32, f32)> = (deleted_count..n)
.map(|i| {
let v = &vectors[i as usize];
let dot: f32 = q.iter().zip(v.iter()).map(|(a, b)| a * b).sum();
(i, 1.0 - dot)
})
.collect();
ground.sort_by(|a, b| a.1.total_cmp(&b.1));
let truth_ids: std::collections::HashSet<u32> =
ground.iter().take(k).map(|(id, _)| *id).collect();
let mut index = HNSWIndex::new(dim, 16, 32).unwrap();
for (i, v) in vectors.iter().enumerate() {
index.add(i as u32, v.clone()).unwrap();
}
index.build().unwrap();
for i in 0..deleted_count {
index.delete_with_repair(i).unwrap();
}
let results = index.search(&q, k, ef).unwrap();
assert_eq!(results.len(), k, "search should return k results");
let hits = results
.iter()
.filter(|(id, _)| truth_ids.contains(id))
.count();
let recall = hits as f32 / k as f32;
assert!(
recall >= recall_floor,
"post-delete recall@{}={:.2} fell below floor {:.2}; deleted_count={}, n={}",
k,
recall,
recall_floor,
deleted_count,
n
);
}
#[test]
fn test_delete_with_repair_preserves_self_search_reachability() {
let dim = 32usize;
let n = 300u32;
let delete_ratio = 0.6f32;
let deleted_count = (n as f32 * delete_ratio) as u32;
let k = 5usize;
let ef = 200usize;
let mut seed: u64 = 42;
let mut next = || -> f32 {
seed = seed.wrapping_mul(6_364_136_223_846_793_005).wrapping_add(1);
((seed >> 33) as f32) / (u32::MAX as f32) - 0.5
};
let mut vectors: Vec<Vec<f32>> = Vec::with_capacity(n as usize);
for _ in 0..n {
let mut v: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
v.iter_mut().for_each(|x| *x /= norm);
}
vectors.push(v);
}
let mut index = HNSWIndex::new(dim, 16, 32).unwrap();
for (i, v) in vectors.iter().enumerate() {
index.add(i as u32, v.clone()).unwrap();
}
index.build().unwrap();
for i in 0..deleted_count {
index.delete_with_repair(i).unwrap();
}
let mut unreachable: Vec<u32> = Vec::new();
for id in deleted_count..n {
let results = index.search(&vectors[id as usize], k, ef).unwrap();
if !results.iter().any(|(rid, _)| *rid == id) {
unreachable.push(id);
}
}
assert!(
unreachable.is_empty(),
"{} of {} live ids became unreachable after deleting {} (ratio {:.0}%): {:?}. \
likely cause: cumulative delete_with_repair calls left live nodes orphaned -- \
crescent-locus replacement may have failed to add in-edges, or the entry-point \
repromotion in find_entry_point_excluding picked an anchor whose neighborhood \
has drifted away from these ids.",
unreachable.len(),
n - deleted_count,
deleted_count,
delete_ratio * 100.0,
&unreachable[..unreachable.len().min(10)]
);
}
#[test]
fn test_delete_with_repair_entry_point() {
let (mut index, q) = build_test_index();
let ep = index.cached_entry_point.unwrap();
let ep_doc_id = index.doc_ids[ep as usize];
index.delete_with_repair(ep_doc_id).unwrap();
assert_ne!(
index.cached_entry_point,
Some(ep),
"entry point should change after deleting it"
);
assert!(
index.cached_entry_point.is_some(),
"should find a new entry point"
);
let results = index.search(&q, 5, 64).unwrap();
assert!(!results.is_empty(), "search should work after EP deletion");
}
#[test]
fn test_delete_with_repair_graph_edges_cleaned() {
let (mut index, _q) = build_test_index();
let target_doc_id = 50u32;
let internal_id = index.doc_id_to_internal[&target_doc_id];
index.delete_with_repair(target_doc_id).unwrap();
let layer = &index.layers[0];
for node_id in 0..layer.len() as u32 {
let neighbors = layer.get_neighbors(node_id);
assert!(
!neighbors.contains(&internal_id),
"node {} still points to deleted node {} after repair",
node_id,
internal_id
);
}
assert!(
layer.get_neighbors(internal_id).is_empty(),
"deleted node should have empty neighbor list"
);
}
#[test]
fn test_delete_batch_with_repair() {
let (mut index, q) = build_test_index();
let k = 10;
let ef = 100;
let ids_to_delete: Vec<u32> = (10..40).collect();
let repairs = index.delete_batch_with_repair(&ids_to_delete).unwrap();
let after = index.search(&q, k, ef).unwrap();
assert_eq!(after.len(), k);
let deleted_set: std::collections::HashSet<u32> = ids_to_delete.iter().copied().collect();
for (id, _) in &after {
assert!(
!deleted_set.contains(id),
"deleted id {} in results (repairs={})",
id,
repairs,
);
}
}
#[test]
fn test_delete_with_repair_nonexistent_errors() {
let (mut index, _q) = build_test_index();
assert!(index.delete_with_repair(9999).is_err());
}
#[test]
fn test_delete_with_repair_recall_vs_tombstone() {
let dim = 32;
let n = 300;
let mut seed: u64 = 77;
let mut next = || -> f32 {
seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1);
((seed >> 33) as f32) / (u32::MAX as f32) - 0.5
};
let mut vectors: Vec<Vec<f32>> = Vec::new();
for _ in 0..n {
let mut v: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
v.iter_mut().for_each(|x| *x /= norm);
}
vectors.push(v);
}
let mut query: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = query.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
query.iter_mut().for_each(|x| *x /= norm);
}
let build = |seed_val: u64| -> HNSWIndex {
let params = HNSWParams {
seed: Some(seed_val),
..HNSWParams::default()
};
let mut idx = HNSWIndex::with_params(dim, params).unwrap();
for (i, v) in vectors.iter().enumerate() {
idx.add(i as u32, v.clone()).unwrap();
}
idx.build().unwrap();
idx
};
let mut repair_idx = build(42);
let mut tombstone_idx = build(42);
let delete_ids: Vec<u32> = (0..120).collect();
for &id in &delete_ids {
repair_idx.delete_with_repair(id).unwrap();
tombstone_idx.delete(id).unwrap();
}
let k = 10;
let ef = 100;
let repair_results = repair_idx.search(&query, k, ef).unwrap();
let tombstone_results = tombstone_idx.search(&query, k, ef).unwrap();
assert!(
!repair_results.is_empty(),
"repair search should return results"
);
assert!(
!tombstone_results.is_empty(),
"tombstone search should return results"
);
assert!(
repair_results.len() >= tombstone_results.len(),
"repair ({}) should return >= tombstone ({}) results",
repair_results.len(),
tombstone_results.len(),
);
}