#![cfg(feature = "diskann")]
#![allow(clippy::unwrap_used, clippy::expect_used, clippy::needless_update)]
use vicinity::diskann::{DiskANNIndex, DiskANNParams, DiskANNSearcher};
fn generate_vectors(n: usize, d: usize, seed: u64) -> Vec<Vec<f32>> {
let mut state = seed;
let mut next = || {
state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
((state >> 33) as f32) / (u32::MAX as f32) - 0.5
};
(0..n).map(|_| (0..d).map(|_| next()).collect()).collect()
}
fn l2_distance(a: &[f32], b: &[f32]) -> f32 {
a.iter()
.zip(b.iter())
.map(|(x, y)| (x - y).powi(2))
.sum::<f32>()
.sqrt()
}
fn brute_force_knn(vectors: &[Vec<f32>], query: &[f32], k: usize) -> Vec<(u32, f32)> {
let mut dists: Vec<(u32, f32)> = vectors
.iter()
.enumerate()
.map(|(i, v)| (i as u32, l2_distance(query, v)))
.collect();
dists.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
dists.truncate(k);
dists
}
fn compute_recall(results: &[(u32, f32)], ground_truth: &[(u32, f32)], k: usize) -> f64 {
let gt_ids: std::collections::HashSet<_> =
ground_truth.iter().take(k).map(|(id, _)| *id).collect();
let result_ids: std::collections::HashSet<_> =
results.iter().take(k).map(|(id, _)| *id).collect();
let intersection = gt_ids.intersection(&result_ids).count();
intersection as f64 / k as f64
}
#[test]
fn test_diskann_save_load_roundtrip() {
let n = 1000;
let d = 32;
let k = 10;
let vectors = generate_vectors(n, d, 42);
let params = DiskANNParams {
m: 16,
ef_construction: 50,
alpha: 1.2,
ef_search: 50,
seed: None,
..DiskANNParams::default()
};
let mut index = DiskANNIndex::new(d, params.clone()).expect("Failed to create index");
for (i, vec) in vectors.iter().enumerate() {
index
.add(i as u32, vec.clone())
.expect("Failed to add vector");
}
index.build().expect("Failed to build index");
let temp_dir = tempfile::tempdir().expect("Failed to create temp dir");
let index_path = temp_dir.path().join("diskann_test");
index.save(&index_path).expect("Failed to save index");
assert!(
index_path.join("vectors.bin").exists(),
"vectors.bin missing"
);
assert!(
index_path.join("graph.index").exists(),
"graph.index missing"
);
assert!(
index_path.join("metadata.json").exists(),
"metadata.json missing"
);
assert!(
index_path.join("doc_ids.bin").exists(),
"doc_ids.bin missing"
);
let mut searcher = DiskANNSearcher::load(&index_path).expect("Failed to load index");
let queries = generate_vectors(10, d, 123);
let mut total_recall = 0.0;
for (query_idx, query) in queries.iter().enumerate() {
let ground_truth = brute_force_knn(&vectors, query, k);
let (results, diagnostics) = searcher
.search_with_diagnostics(query, k, params.ef_search)
.expect("Search failed");
let plain_results = searcher
.search(query, k, params.ef_search)
.expect("plain search failed");
assert_eq!(
results, plain_results,
"diagnostic search must preserve result ordering"
);
assert!(
diagnostics.graph_reads > 0,
"query {query_idx} should read at least one graph record"
);
assert!(
diagnostics.vector_reads >= results.len(),
"query {query_idx} should read vectors for at least the returned hits"
);
assert_eq!(
diagnostics.vector_bytes,
diagnostics.vector_reads * d * std::mem::size_of::<f32>()
);
assert!(
diagnostics.graph_bytes >= diagnostics.graph_reads * std::mem::size_of::<u32>(),
"query {query_idx} graph bytes should include one degree per read"
);
assert_eq!(
diagnostics.visited_nodes, diagnostics.vector_reads,
"file-backed search reads one vector for each first visit"
);
let recall = compute_recall(&results, &ground_truth, k);
total_recall += recall;
}
let avg_recall = total_recall / queries.len() as f64;
println!("DiskANN persistence test:");
println!(" Vectors: {}", n);
println!(" Dimension: {}", d);
println!(" Average Recall@{}: {:.2}%", k, avg_recall * 100.0);
assert!(
avg_recall > 0.5,
"Recall too low: {:.2}%",
avg_recall * 100.0
);
}
#[test]
fn test_diskann_file_search_matches_in_memory_search() {
let n = 500;
let d = 24;
let k = 10;
let ef = 40;
let vectors = generate_vectors(n, d, 99);
let queries = generate_vectors(12, d, 199);
let params = DiskANNParams {
m: 16,
ef_construction: 50,
alpha: 1.2,
ef_search: ef,
seed: Some(7),
..DiskANNParams::default()
};
let mut index = DiskANNIndex::new(d, params).expect("create index");
for (i, vec) in vectors.iter().enumerate() {
index.add(i as u32, vec.clone()).expect("add vector");
}
index.build().expect("build index");
let temp_dir = tempfile::tempdir().expect("temp dir");
let index_path = temp_dir.path().join("diskann_parity");
index.save(&index_path).expect("save index");
let mut searcher = DiskANNSearcher::load(&index_path).expect("load searcher");
let mut mmap_searcher = DiskANNSearcher::load_mmap(&index_path).expect("load mmap searcher");
for query in &queries {
let in_memory = index.search(query, k, ef).expect("in-memory search");
let file_backed = searcher.search(query, k, ef).expect("file-backed search");
let mmap_backed = mmap_searcher
.search(query, k, ef)
.expect("mmap-backed search");
assert_eq!(
file_backed, in_memory,
"file-backed search must preserve in-memory DiskANN ranking"
);
assert_eq!(
mmap_backed, in_memory,
"mmap-backed search must preserve in-memory DiskANN ranking"
);
}
}
#[test]
fn test_diskann_metadata_roundtrip() {
let n = 100;
let d = 16;
let vectors = generate_vectors(n, d, 42);
let params = DiskANNParams {
m: 8,
ef_construction: 20,
alpha: 1.1,
ef_search: 20,
seed: None,
..DiskANNParams::default()
};
let mut index = DiskANNIndex::new(d, params.clone()).expect("Failed to create index");
for (i, vec) in vectors.iter().enumerate() {
index
.add(i as u32, vec.clone())
.expect("Failed to add vector");
}
index.build().expect("Failed to build index");
let temp_dir = tempfile::tempdir().expect("Failed to create temp dir");
let index_path = temp_dir.path().join("diskann_meta_test");
index.save(&index_path).expect("Failed to save index");
let metadata_path = index_path.join("metadata.json");
let metadata: serde_json::Value = serde_json::from_reader(
std::fs::File::open(&metadata_path).expect("Failed to open metadata"),
)
.expect("Failed to parse metadata");
assert_eq!(metadata["dimension"], d);
assert_eq!(metadata["num_vectors"], n);
assert_eq!(metadata["params"]["m"], params.m);
assert_eq!(
metadata["params"]["ef_construction"],
params.ef_construction
);
}