use crate::index::IndexConfig;
use crate::retrieval::{RerankedResult, SearchResult};
use crate::similarity::compute_similarity;
use embeddenator_vsa::SparseVec;
use serde::{Deserialize, Serialize};
use std::cmp::Reverse;
use std::collections::{BinaryHeap, HashMap, HashSet};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HNSWConfig {
pub m: usize,
pub m_max0: usize,
pub ef_construction: usize,
pub ef_search: usize,
pub ml: f64,
}
impl Default for HNSWConfig {
fn default() -> Self {
let m = 16;
Self {
m,
m_max0: m * 2,
ef_construction: 200,
ef_search: 50,
ml: 1.0 / (m as f64).ln(),
}
}
}
impl HNSWConfig {
pub fn fast() -> Self {
Self {
m: 12,
m_max0: 24,
ef_construction: 100,
ef_search: 20,
ml: 1.0 / 12_f64.ln(),
}
}
pub fn accurate() -> Self {
Self {
m: 32,
m_max0: 64,
ef_construction: 400,
ef_search: 200,
ml: 1.0 / 32_f64.ln(),
}
}
pub fn with_m(mut self, m: usize) -> Self {
self.m = m;
self.m_max0 = m * 2;
self.ml = 1.0 / (m as f64).ln();
self
}
pub fn with_ef_search(mut self, ef: usize) -> Self {
self.ef_search = ef;
self
}
pub fn with_ef_construction(mut self, ef: usize) -> Self {
self.ef_construction = ef;
self
}
}
#[derive(Clone, Debug)]
struct HNSWNode {
#[allow(dead_code)]
id: usize,
vec: SparseVec,
#[allow(dead_code)]
level: usize,
neighbors: Vec<Vec<usize>>,
}
impl HNSWNode {
fn new(id: usize, vec: SparseVec, level: usize) -> Self {
Self {
id,
vec,
level,
neighbors: vec![Vec::new(); level + 1],
}
}
}
#[derive(Clone, Copy, Debug)]
struct Candidate {
id: usize,
distance: f64,
}
impl PartialEq for Candidate {
fn eq(&self, other: &Self) -> bool {
self.id == other.id
}
}
impl Eq for Candidate {}
impl PartialOrd for Candidate {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for Candidate {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other
.distance
.partial_cmp(&self.distance)
.unwrap_or(std::cmp::Ordering::Equal)
}
}
#[derive(Clone, Debug)]
pub struct HNSWIndex {
hnsw_config: HNSWConfig,
index_config: IndexConfig,
nodes: HashMap<usize, HNSWNode>,
entry_point: Option<usize>,
max_level: usize,
rng_state: u64,
}
impl HNSWIndex {
pub fn new(hnsw_config: HNSWConfig, index_config: IndexConfig) -> Self {
Self {
hnsw_config,
index_config,
nodes: HashMap::new(),
entry_point: None,
max_level: 0,
rng_state: 0x5DEECE66D, }
}
pub fn with_index_config(index_config: IndexConfig) -> Self {
Self::new(HNSWConfig::default(), index_config)
}
pub fn len(&self) -> usize {
self.nodes.len()
}
pub fn is_empty(&self) -> bool {
self.nodes.is_empty()
}
pub fn max_level(&self) -> usize {
self.max_level
}
pub fn set_ef_search(&mut self, ef: usize) {
self.hnsw_config.ef_search = ef;
}
fn generate_level(&mut self) -> usize {
self.rng_state = self.rng_state.wrapping_mul(0x5DEECE66D).wrapping_add(0xB);
let r = (self.rng_state >> 17) as f64 / (1u64 << 47) as f64;
(-(r.max(f64::MIN_POSITIVE).ln()) * self.hnsw_config.ml) as usize
}
fn distance(&self, a: &SparseVec, b: &SparseVec) -> f64 {
1.0 - compute_similarity(a, b, self.index_config.metric)
}
fn search_layer(
&self,
query: &SparseVec,
entry_points: &[usize],
ef: usize,
layer: usize,
) -> Vec<(usize, f64)> {
let mut visited: HashSet<usize> = entry_points.iter().copied().collect();
let mut candidates: BinaryHeap<Candidate> = BinaryHeap::new();
let mut results: BinaryHeap<Reverse<Candidate>> = BinaryHeap::new();
for &ep in entry_points {
if let Some(node) = self.nodes.get(&ep) {
let dist = self.distance(query, &node.vec);
candidates.push(Candidate {
id: ep,
distance: dist,
});
results.push(Reverse(Candidate {
id: ep,
distance: dist,
}));
}
}
while let Some(Candidate {
id: current_id,
distance: current_dist,
}) = candidates.pop()
{
let furthest_dist = results
.peek()
.map(|Reverse(c)| c.distance)
.unwrap_or(f64::INFINITY);
if current_dist > furthest_dist && results.len() >= ef {
break;
}
if let Some(node) = self.nodes.get(¤t_id) {
if layer < node.neighbors.len() {
for &neighbor_id in &node.neighbors[layer] {
if visited.insert(neighbor_id) {
if let Some(neighbor) = self.nodes.get(&neighbor_id) {
let dist = self.distance(query, &neighbor.vec);
let should_add = results.len() < ef || {
let worst = results
.peek()
.map(|Reverse(c)| c.distance)
.unwrap_or(f64::INFINITY);
dist < worst
};
if should_add {
candidates.push(Candidate {
id: neighbor_id,
distance: dist,
});
results.push(Reverse(Candidate {
id: neighbor_id,
distance: dist,
}));
if results.len() > ef {
results.pop();
}
}
}
}
}
}
}
}
let mut result_vec: Vec<(usize, f64)> = results
.into_iter()
.map(|Reverse(c)| (c.id, c.distance))
.collect();
result_vec.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
result_vec
}
fn select_neighbors_heuristic(
&self,
_query: &SparseVec,
candidates: &[(usize, f64)],
m: usize,
) -> Vec<usize> {
if candidates.len() <= m {
return candidates.iter().map(|(id, _)| *id).collect();
}
let mut selected: Vec<usize> = Vec::with_capacity(m);
let mut working: Vec<(usize, f64)> = candidates.to_vec();
while selected.len() < m && !working.is_empty() {
working.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
let (best_id, best_dist) = working.remove(0);
let is_diverse = selected.iter().all(|&sel_id| {
if let (Some(best_node), Some(sel_node)) =
(self.nodes.get(&best_id), self.nodes.get(&sel_id))
{
let inter_dist = self.distance(&best_node.vec, &sel_node.vec);
best_dist <= inter_dist
} else {
true
}
});
if is_diverse || selected.len() < m / 2 {
selected.push(best_id);
}
if selected.len() < m && working.is_empty() && !is_diverse {
working.push((best_id, best_dist));
for (id, dist) in candidates {
if !selected.contains(id) && !working.iter().any(|(wid, _)| wid == id) {
working.push((*id, *dist));
}
}
}
}
if selected.len() < m {
let remaining: Vec<usize> = candidates
.iter()
.filter(|(id, _)| !selected.contains(id))
.map(|(id, _)| *id)
.take(m - selected.len())
.collect();
selected.extend(remaining);
}
selected
}
fn connect_nodes(&mut self, id1: usize, id2: usize, layer: usize) {
let m_max = if layer == 0 {
self.hnsw_config.m_max0
} else {
self.hnsw_config.m
};
if let Some(node1) = self.nodes.get_mut(&id1) {
if layer < node1.neighbors.len() && !node1.neighbors[layer].contains(&id2) {
node1.neighbors[layer].push(id2);
}
}
if let Some(node2) = self.nodes.get_mut(&id2) {
if layer < node2.neighbors.len() && !node2.neighbors[layer].contains(&id1) {
node2.neighbors[layer].push(id1);
}
}
for id in [id1, id2] {
if let Some(node) = self.nodes.get(&id) {
if layer < node.neighbors.len() && node.neighbors[layer].len() > m_max {
let query = node.vec.clone();
let neighbors: Vec<usize> = node.neighbors[layer].clone();
let candidates: Vec<(usize, f64)> = neighbors
.iter()
.filter_map(|&nid| {
self.nodes
.get(&nid)
.map(|n| (nid, self.distance(&query, &n.vec)))
})
.collect();
let selected = self.select_neighbors_heuristic(&query, &candidates, m_max);
if let Some(node) = self.nodes.get_mut(&id) {
if layer < node.neighbors.len() {
node.neighbors[layer] = selected;
}
}
}
}
}
}
pub fn insert(&mut self, id: usize, vec: &SparseVec) {
let level = self.generate_level();
let node = HNSWNode::new(id, vec.clone(), level);
if self.nodes.is_empty() {
self.nodes.insert(id, node);
self.entry_point = Some(id);
self.max_level = level;
return;
}
let entry_point = self.entry_point.expect("Entry point should exist");
let mut curr_ep = vec![entry_point];
for lc in (level + 1..=self.max_level).rev() {
let nearest = self.search_layer(vec, &curr_ep, 1, lc);
if let Some((nearest_id, _)) = nearest.first() {
curr_ep = vec![*nearest_id];
}
}
self.nodes.insert(id, node);
let top_layer = level.min(self.max_level);
for lc in (0..=top_layer).rev() {
let candidates = self.search_layer(vec, &curr_ep, self.hnsw_config.ef_construction, lc);
let m = if lc == 0 {
self.hnsw_config.m_max0
} else {
self.hnsw_config.m
};
let neighbors = self.select_neighbors_heuristic(vec, &candidates, m);
for &neighbor_id in &neighbors {
self.connect_nodes(id, neighbor_id, lc);
}
curr_ep = candidates.iter().map(|(cid, _)| *cid).collect();
}
if level > self.max_level {
self.entry_point = Some(id);
self.max_level = level;
}
}
fn search(&self, query: &SparseVec, k: usize) -> Vec<(usize, f64)> {
if self.nodes.is_empty() {
return Vec::new();
}
let entry_point = match self.entry_point {
Some(ep) => ep,
None => return Vec::new(),
};
let mut curr_ep = vec![entry_point];
for lc in (1..=self.max_level).rev() {
let nearest = self.search_layer(query, &curr_ep, 1, lc);
if let Some((nearest_id, _)) = nearest.first() {
curr_ep = vec![*nearest_id];
}
}
let candidates = self.search_layer(query, &curr_ep, self.hnsw_config.ef_search.max(k), 0);
candidates.into_iter().take(k).collect()
}
pub fn stats(&self) -> HNSWStats {
let mut nodes_per_level = vec![0usize; self.max_level + 1];
let mut total_edges = 0usize;
for node in self.nodes.values() {
for (layer, neighbors) in node.neighbors.iter().enumerate() {
if layer <= self.max_level {
nodes_per_level[layer] += 1;
total_edges += neighbors.len();
}
}
}
HNSWStats {
num_vectors: self.nodes.len(),
max_level: self.max_level,
nodes_per_level,
total_edges: total_edges / 2, avg_degree: if self.nodes.is_empty() {
0.0
} else {
total_edges as f64 / self.nodes.len() as f64
},
}
}
}
#[derive(Debug, Clone)]
pub struct HNSWStats {
pub num_vectors: usize,
pub max_level: usize,
pub nodes_per_level: Vec<usize>,
pub total_edges: usize,
pub avg_degree: f64,
}
impl crate::index::RetrievalIndex for HNSWIndex {
fn add(&mut self, id: usize, vec: &SparseVec) {
self.insert(id, vec);
}
fn finalize(&mut self) {
}
fn query_top_k(&self, query: &SparseVec, k: usize) -> Vec<SearchResult> {
if k == 0 {
return Vec::new();
}
self.search(query, k)
.into_iter()
.map(|(id, distance)| {
let similarity = 1.0 - distance;
let score = (similarity * 1000.0) as i32;
SearchResult { id, score }
})
.collect()
}
fn query_top_k_reranked(
&self,
query: &SparseVec,
vectors: &HashMap<usize, SparseVec>,
candidate_k: usize,
k: usize,
) -> Vec<RerankedResult> {
if k == 0 {
return Vec::new();
}
let candidates = self.search(query, candidate_k);
let mut results: Vec<RerankedResult> = candidates
.into_iter()
.filter_map(|(id, _dist)| {
let vec = vectors
.get(&id)
.or_else(|| self.nodes.get(&id).map(|n| &n.vec))?;
let cosine = query.cosine(vec);
let approx_score = (cosine * 1000.0) as i32;
Some(RerankedResult {
id,
approx_score,
cosine,
})
})
.collect();
results.sort_by(|a, b| {
b.cosine
.partial_cmp(&a.cosine)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.id.cmp(&b.id))
});
results.truncate(k);
results
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::index::RetrievalIndex;
#[allow(deprecated)]
fn create_test_vector(data: &[u8]) -> SparseVec {
SparseVec::from_data(data)
}
#[test]
fn test_hnsw_basic() {
let config = HNSWConfig::default();
let index_config = IndexConfig::default();
let mut index = HNSWIndex::new(config, index_config);
let vec1 = create_test_vector(b"apple");
let vec2 = create_test_vector(b"banana");
let vec3 = create_test_vector(b"cherry");
index.add(1, &vec1);
index.add(2, &vec2);
index.add(3, &vec3);
assert_eq!(index.len(), 3);
let results = index.query_top_k(&vec1, 2);
assert!(!results.is_empty());
assert_eq!(results[0].id, 1);
}
#[test]
fn test_hnsw_empty() {
let config = HNSWConfig::default();
let index_config = IndexConfig::default();
let index = HNSWIndex::new(config, index_config);
let query = create_test_vector(b"test");
let results = index.query_top_k(&query, 5);
assert!(results.is_empty());
}
#[test]
fn test_hnsw_single_element() {
let config = HNSWConfig::default();
let index_config = IndexConfig::default();
let mut index = HNSWIndex::new(config, index_config);
let vec = create_test_vector(b"single");
index.add(42, &vec);
let results = index.query_top_k(&vec, 5);
assert_eq!(results.len(), 1);
assert_eq!(results[0].id, 42);
}
#[test]
fn test_hnsw_many_vectors() {
let config = HNSWConfig::fast(); let index_config = IndexConfig::default();
let mut index = HNSWIndex::new(config, index_config);
for i in 0..10 {
let data = format!("document-{}", i);
let vec = create_test_vector(data.as_bytes());
index.add(i, &vec);
}
assert_eq!(index.len(), 10);
let stats = index.stats();
assert_eq!(stats.num_vectors, 10);
let query = create_test_vector(b"document-5");
let results = index.query_top_k(&query, 5);
assert!(!results.is_empty());
assert!(results.len() <= 5);
let top_ids: Vec<usize> = results.iter().map(|r| r.id).collect();
assert!(
top_ids.contains(&5),
"Query vector should be in top results"
);
}
#[test]
fn test_hnsw_reranking() {
let config = HNSWConfig::default();
let index_config = IndexConfig::default();
let mut index = HNSWIndex::new(config, index_config);
let mut vectors = HashMap::new();
for i in 0..20 {
let data = format!("item-{}", i);
let vec = create_test_vector(data.as_bytes());
index.add(i, &vec);
vectors.insert(i, vec);
}
let query = create_test_vector(b"item-5");
let results = index.query_top_k_reranked(&query, &vectors, 10, 5);
assert!(!results.is_empty());
assert!(results.len() <= 5);
for i in 1..results.len() {
assert!(results[i - 1].cosine >= results[i].cosine);
}
}
#[test]
fn test_hnsw_config_builders() {
let fast = HNSWConfig::fast();
assert!(fast.m < HNSWConfig::default().m);
assert!(fast.ef_search < HNSWConfig::default().ef_search);
let accurate = HNSWConfig::accurate();
assert!(accurate.m > HNSWConfig::default().m);
assert!(accurate.ef_search > HNSWConfig::default().ef_search);
let custom = HNSWConfig::default()
.with_m(24)
.with_ef_search(100)
.with_ef_construction(300);
assert_eq!(custom.m, 24);
assert_eq!(custom.ef_search, 100);
assert_eq!(custom.ef_construction, 300);
}
#[test]
fn test_hnsw_stats() {
let config = HNSWConfig::fast(); let index_config = IndexConfig::default();
let mut index = HNSWIndex::new(config, index_config);
for i in 0..20 {
let data = format!("vec-{}", i);
let vec = create_test_vector(data.as_bytes());
index.add(i, &vec);
}
let stats = index.stats();
assert_eq!(stats.num_vectors, 20);
assert!(stats.total_edges > 0);
assert!(stats.avg_degree > 0.0);
}
}