use std::sync::{Mutex, OnceLock};
use super::hnsw::{AnnIndex, FlatEmbeddings, brute_force_topk};
pub const ANN_MIN_VECTORS: usize = 2_500;
struct Cached {
fingerprint: u64,
index: AnnIndex,
}
fn cache() -> &'static Mutex<Option<Cached>> {
static CACHE: OnceLock<Mutex<Option<Cached>>> = OnceLock::new();
CACHE.get_or_init(|| Mutex::new(None))
}
#[must_use]
pub fn topk(embeddings: &FlatEmbeddings, query: &[f32], top_k: usize) -> Vec<(usize, f32)> {
topk_gated(embeddings, query, top_k, ANN_MIN_VECTORS)
}
fn topk_gated(
embeddings: &FlatEmbeddings,
query: &[f32],
top_k: usize,
min_vectors: usize,
) -> Vec<(usize, f32)> {
if embeddings.n_vectors() < min_vectors {
return brute_force_topk(embeddings, query, top_k);
}
let fp = fingerprint(embeddings);
let Ok(mut guard) = cache().lock() else {
return brute_force_topk(embeddings, query, top_k);
};
let needs_build = match guard.as_ref() {
Some(c) => c.fingerprint != fp,
None => true,
};
if needs_build {
*guard = Some(Cached {
fingerprint: fp,
index: AnnIndex::build(embeddings.clone()),
});
}
match guard.as_ref() {
Some(c) => c.index.search(query, top_k),
None => brute_force_topk(embeddings, query, top_k),
}
}
fn fingerprint(embeddings: &FlatEmbeddings) -> u64 {
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
macro_rules! mix {
($x:expr_2021) => {{
h ^= $x;
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}};
}
let n = embeddings.n_vectors();
mix!(n as u64);
mix!(embeddings.dim as u64);
for i in 0..n {
let v = embeddings.get(i);
mix!(i as u64);
if let Some(&f) = v.first() {
mix!(u64::from(f.to_bits()));
}
if let Some(&f) = v.get(v.len() / 2) {
mix!(u64::from(f.to_bits()));
}
if let Some(&f) = v.last() {
mix!(u64::from(f.to_bits()));
}
}
h
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashSet;
const TEST_GATE: usize = 1000;
static TEST_LOCK: Mutex<()> = Mutex::new(());
fn serial() -> std::sync::MutexGuard<'static, ()> {
TEST_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn cached_fingerprint() -> Option<u64> {
cache()
.lock()
.ok()
.and_then(|g| g.as_ref().map(|c| c.fingerprint))
}
fn random_vec(dim: usize, seed: u64) -> Vec<f32> {
let mut v = Vec::with_capacity(dim);
let mut s = seed;
for _ in 0..dim {
s = s.wrapping_mul(6_364_136_223_846_793_005).wrapping_add(1);
v.push((s as f32 / u64::MAX as f32) * 2.0 - 1.0);
}
v
}
fn flat_from(vecs: Vec<Vec<f32>>) -> FlatEmbeddings {
FlatEmbeddings::from_vecs(vecs)
}
fn jitter(base: &[f32], seed: u64, scale: f32) -> Vec<f32> {
base.iter()
.enumerate()
.map(|(i, &b)| {
let s = seed
.wrapping_add(i as u64)
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
b + ((s as f32 / u64::MAX as f32) * 2.0 - 1.0) * scale
})
.collect()
}
fn clustered(
n_clusters: usize,
per_cluster: usize,
dim: usize,
) -> (FlatEmbeddings, Vec<Vec<f32>>) {
let centers: Vec<Vec<f32>> = (0..n_clusters)
.map(|c| random_vec(dim, (c as u64 + 1) * 1_000))
.collect();
let mut vectors = Vec::with_capacity(n_clusters * per_cluster);
for (c, center) in centers.iter().enumerate() {
for j in 0..per_cluster {
vectors.push(jitter(center, (c * per_cluster + j) as u64 + 7, 0.02));
}
}
(flat_from(vectors), centers)
}
#[test]
fn small_corpus_matches_brute_force_exactly() {
let flat = flat_from((0..200).map(|i| random_vec(32, i)).collect());
let query = random_vec(32, 9_999);
let via_cache = topk(&flat, &query, 8);
let exact = brute_force_topk(&flat, &query, 8);
assert_eq!(via_cache.len(), exact.len());
for (a, b) in via_cache.iter().zip(exact.iter()) {
assert_eq!(a.0, b.0, "below threshold must be exact brute force");
}
}
#[test]
fn hnsw_path_recall_matches_brute_force_on_clusters() {
let _serial = serial();
let (flat, centers) = clustered(24, 60, 32); let query = ¢ers[5];
let k = 20;
let ann = topk_gated(&flat, query, k, TEST_GATE); let exact = brute_force_topk(&flat, query, k);
assert_eq!(ann.len(), k);
let exact_set: HashSet<usize> = exact.iter().map(|(i, _)| *i).collect();
let overlap = ann.iter().filter(|(i, _)| exact_set.contains(i)).count();
assert!(
overlap * 100 >= k * 50,
"HNSW recall@{k} too low: {overlap}/{k}"
);
}
#[test]
fn hnsw_path_results_are_descending() {
let _serial = serial();
let (flat, centers) = clustered(20, 60, 24); let results = topk_gated(&flat, ¢ers[3], 10, TEST_GATE);
for w in results.windows(2) {
assert!(
w[0].1 >= w[1].1,
"results must be sorted by descending similarity"
);
}
}
#[test]
fn rebuilds_when_corpus_changes() {
let _serial = serial();
let (a, ca) = clustered(20, 55, 32); let (b, cb) = clustered(18, 60, 32);
let _ = topk_gated(&a, &ca[7], 5, TEST_GATE);
assert_eq!(
cached_fingerprint(),
Some(fingerprint(&a)),
"first query caches corpus A's index"
);
let _ = topk_gated(&b, &cb[4], 5, TEST_GATE);
assert_eq!(
cached_fingerprint(),
Some(fingerprint(&b)),
"a different corpus must force a rebuild to B"
);
let _ = topk_gated(&a, &ca[7], 5, TEST_GATE);
assert_eq!(
cached_fingerprint(),
Some(fingerprint(&a)),
"re-querying A must rebuild A — never serve stale B"
);
}
#[test]
fn fingerprint_differs_on_content_change() {
let a = flat_from((0..10).map(|i| random_vec(8, i)).collect());
let mut b_vecs: Vec<Vec<f32>> = (0..10).map(|i| random_vec(8, i)).collect();
b_vecs[3][0] += 0.5;
let b = flat_from(b_vecs);
assert_ne!(fingerprint(&a), fingerprint(&b));
}
}