use crate::error::{Error, Result};
use crate::types::MemoryId;
use serde::{Deserialize, Serialize};
use std::cmp::Ordering;
use std::collections::{BinaryHeap, HashMap, HashSet};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HnswConfig {
pub m: usize,
pub ef_construction: usize,
pub ef_search: usize,
}
impl Default for HnswConfig {
fn default() -> Self {
Self {
m: 16,
ef_construction: 100,
ef_search: 50,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HnswStats {
pub point_count: usize,
pub deleted_count: usize,
pub dimensions: usize,
pub max_layer: usize,
pub memory_bytes: usize,
}
#[derive(Clone)]
struct HnswPoint {
id: MemoryId,
embedding: Vec<f32>,
neighbors: Vec<Vec<usize>>, deleted: bool,
}
#[derive(Clone)]
struct Candidate {
index: usize,
distance: f32,
}
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
.partial_cmp(&self.distance)
.unwrap_or(Ordering::Equal)
}
}
impl PartialOrd for Candidate {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
#[derive(Clone)]
struct MaxCandidate {
index: usize,
distance: f32,
}
impl PartialEq for MaxCandidate {
fn eq(&self, other: &Self) -> bool {
self.distance == other.distance
}
}
impl Eq for MaxCandidate {}
impl Ord for MaxCandidate {
fn cmp(&self, other: &Self) -> Ordering {
self.distance
.partial_cmp(&other.distance)
.unwrap_or(Ordering::Equal)
}
}
impl PartialOrd for MaxCandidate {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
pub struct HnswIndex {
points: Vec<HnswPoint>,
id_to_index: HashMap<MemoryId, usize>,
config: HnswConfig,
max_layer: usize,
entry_point: Option<usize>,
dirty: bool,
deleted_count: usize,
dimensions: usize,
}
impl HnswIndex {
pub fn new(config: HnswConfig) -> Self {
Self {
points: Vec::new(),
id_to_index: HashMap::new(),
config,
max_layer: 0,
entry_point: None,
dirty: false,
deleted_count: 0,
dimensions: 0,
}
}
pub fn insert(&mut self, id: MemoryId, embedding: Vec<f32>) {
if embedding.is_empty() {
return;
}
if self.points.is_empty() {
self.dimensions = embedding.len();
}
if self.id_to_index.contains_key(&id) {
self.remove(&id);
}
let index = self.points.len();
let level = self.random_level();
let mut neighbors = Vec::with_capacity(level + 1);
for _ in 0..=level {
neighbors.push(Vec::new());
}
let point = HnswPoint {
id: id.clone(),
embedding,
neighbors,
deleted: false,
};
self.points.push(point);
self.id_to_index.insert(id, index);
if let Some(ep) = self.entry_point {
if level > self.max_layer {
self.entry_point = Some(index);
self.max_layer = level;
}
self.connect_new_point(index, ep);
} else {
self.entry_point = Some(index);
self.max_layer = level;
}
self.dirty = true;
}
pub fn search(&self, query: &[f32], k: usize) -> Vec<(MemoryId, f32)> {
if self.points.is_empty() || query.is_empty() {
return Vec::new();
}
let ep = match self.entry_point {
Some(ep) if !self.points[ep].deleted => ep,
_ => return Vec::new(),
};
let mut current = ep;
for layer in (1..=self.max_layer).rev() {
current = self.greedy_search(current, query, layer);
}
let candidates = self.search_layer(current, query, self.config.ef_search, 0);
candidates
.into_iter()
.filter(|c| !self.points[c.index].deleted)
.take(k)
.map(|c| {
let similarity = 1.0 - c.distance; (self.points[c.index].id.clone(), similarity)
})
.collect()
}
pub fn remove(&mut self, id: &MemoryId) {
if let Some(&point_index) = self.id_to_index.get(id) {
if !self.points[point_index].deleted {
self.points[point_index].deleted = true;
self.deleted_count += 1;
self.dirty = true;
for i in 0..self.points.len() {
if i == point_index {
continue;
}
for layer in &mut self.points[i].neighbors {
layer.retain(|&n| n != point_index);
}
}
}
}
}
pub fn rebuild(&mut self) {
let active_points: Vec<(MemoryId, Vec<f32>)> = self
.points
.iter()
.filter(|p| !p.deleted)
.map(|p| (p.id.clone(), p.embedding.clone()))
.collect();
self.points.clear();
self.id_to_index.clear();
self.entry_point = None;
self.max_layer = 0;
self.deleted_count = 0;
for (id, embedding) in active_points {
self.insert(id, embedding);
}
self.dirty = true;
}
pub fn serialize(&self) -> Result<Vec<u8>> {
let points: Vec<HnswPointSnapshot> = self
.points
.iter()
.map(|p| HnswPointSnapshot {
id: p.id.clone(),
embedding: p.embedding.clone(),
neighbors: p.neighbors.clone(),
deleted: p.deleted,
})
.collect();
let snapshot = HnswSnapshotV2 {
version: 2,
max_layer: self.max_layer,
entry_point: self.entry_point,
ef_construction: self.config.ef_construction,
ef_search: self.config.ef_search,
m: self.config.m,
points,
};
serde_json::to_vec(&snapshot).map_err(|e| Error::internal(format!("HNSW serialize: {e}")))
}
pub fn deserialize(data: &[u8]) -> Result<Self> {
if let Ok(snapshot) = serde_json::from_slice::<HnswSnapshotV2>(data) {
if snapshot.version >= 2 {
let config = HnswConfig {
m: snapshot.m,
ef_construction: snapshot.ef_construction,
ef_search: snapshot.ef_search,
};
let mut index = Self::new(config);
index.max_layer = snapshot.max_layer;
index.entry_point = snapshot.entry_point;
let dimensions = snapshot
.points
.first()
.map(|p| p.embedding.len())
.unwrap_or(0);
index.dimensions = dimensions;
for pt in &snapshot.points {
let hnsw_point = HnswPoint {
id: pt.id.clone(),
embedding: pt.embedding.clone(),
neighbors: pt.neighbors.clone(),
deleted: pt.deleted,
};
let idx = index.points.len();
index.id_to_index.insert(pt.id.clone(), idx);
if pt.deleted {
index.deleted_count += 1;
}
index.points.push(hnsw_point);
}
return Ok(index);
}
}
let snapshot: HnswSnapshotLegacy = serde_json::from_slice(data)
.map_err(|e| Error::internal(format!("HNSW deserialize: {e}")))?;
let mut index = Self::new(snapshot.config);
for (id, embedding) in snapshot.points {
index.insert(id, embedding);
}
Ok(index)
}
pub fn stats(&self) -> HnswStats {
let mut memory_bytes = std::mem::size_of::<Self>();
memory_bytes += 24; for p in &self.points {
memory_bytes += 32 + 24 + p.embedding.len() * 4;
memory_bytes += 24;
for layer in &p.neighbors {
memory_bytes += 24 + layer.len() * std::mem::size_of::<usize>();
}
memory_bytes += 1;
}
memory_bytes += self.id_to_index.len() * (32 + std::mem::size_of::<usize>() + 32);
HnswStats {
point_count: self.points.len() - self.deleted_count,
deleted_count: self.deleted_count,
dimensions: self.dimensions,
max_layer: self.max_layer,
memory_bytes,
}
}
pub fn is_dirty(&self) -> bool {
self.dirty
}
pub fn mark_clean(&mut self) {
self.dirty = false;
}
pub fn len(&self) -> usize {
self.points.len() - self.deleted_count
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
fn random_level(&self) -> usize {
let ml = 1.0 / (self.config.m as f64).ln();
let r: f64 = rand_f64();
(-r.ln() * ml).floor() as usize
}
fn cosine_distance(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return 1.0;
}
let mut dot = 0.0f32;
let mut norm_a = 0.0f32;
let mut norm_b = 0.0f32;
for i in 0..a.len() {
dot += a[i] * b[i];
norm_a += a[i] * a[i];
norm_b += b[i] * b[i];
}
let denom = norm_a.sqrt() * norm_b.sqrt();
if denom < f32::EPSILON {
return 1.0;
}
1.0 - (dot / denom)
}
fn greedy_search(&self, start: usize, query: &[f32], layer: usize) -> usize {
let mut current = start;
let mut best_dist = Self::cosine_distance(&self.points[current].embedding, query);
loop {
let mut changed = false;
let neighbors = if layer < self.points[current].neighbors.len() {
&self.points[current].neighbors[layer]
} else {
break;
};
for &neighbor_idx in neighbors {
if neighbor_idx >= self.points.len() || self.points[neighbor_idx].deleted {
continue;
}
let dist = Self::cosine_distance(&self.points[neighbor_idx].embedding, query);
if dist < best_dist {
best_dist = dist;
current = neighbor_idx;
changed = true;
}
}
if !changed {
break;
}
}
current
}
fn search_layer(
&self,
start: usize,
query: &[f32],
ef: usize,
_layer: usize,
) -> Vec<Candidate> {
let mut visited = HashSet::new();
let start_dist = Self::cosine_distance(&self.points[start].embedding, query);
let mut candidates = BinaryHeap::new(); let mut result = BinaryHeap::<MaxCandidate>::new();
candidates.push(Candidate {
index: start,
distance: start_dist,
});
result.push(MaxCandidate {
index: start,
distance: start_dist,
});
visited.insert(start);
while let Some(current) = candidates.pop() {
if let Some(worst) = result.peek() {
if current.distance > worst.distance && result.len() >= ef {
break;
}
}
let neighbors = if !self.points[current.index].neighbors.is_empty() {
&self.points[current.index].neighbors[0]
} else {
continue;
};
for &neighbor_idx in neighbors {
if neighbor_idx >= self.points.len() || visited.contains(&neighbor_idx) {
continue;
}
visited.insert(neighbor_idx);
if self.points[neighbor_idx].deleted {
continue;
}
let dist = Self::cosine_distance(&self.points[neighbor_idx].embedding, query);
let should_add = result.len() < ef || {
if let Some(worst) = result.peek() {
dist < worst.distance
} else {
true
}
};
if should_add {
candidates.push(Candidate {
index: neighbor_idx,
distance: dist,
});
result.push(MaxCandidate {
index: neighbor_idx,
distance: dist,
});
if result.len() > ef {
result.pop(); }
}
}
}
let mut results: Vec<Candidate> = result
.into_iter()
.map(|mc| Candidate {
index: mc.index,
distance: mc.distance,
})
.collect();
results.sort_by(|a, b| {
a.distance
.partial_cmp(&b.distance)
.unwrap_or(Ordering::Equal)
});
results
}
fn connect_new_point(&mut self, new_idx: usize, entry_point: usize) {
let query = self.points[new_idx].embedding.clone();
let new_level = self.points[new_idx].neighbors.len().saturating_sub(1);
let m = self.config.m;
let mut current = entry_point;
for layer in (new_level + 1..=self.max_layer).rev() {
current = self.greedy_search(current, &query, layer);
}
for layer in (0..=new_level.min(self.max_layer)).rev() {
let candidates = self.search_layer(current, &query, self.config.ef_construction, layer);
let max_neighbors = if layer == 0 { m * 2 } else { m };
let selected: Vec<usize> = candidates
.iter()
.filter(|c| c.index != new_idx && !self.points[c.index].deleted)
.take(max_neighbors)
.map(|c| c.index)
.collect();
if layer < self.points[new_idx].neighbors.len() {
self.points[new_idx].neighbors[layer] = selected.clone();
}
for &neighbor_idx in &selected {
if layer < self.points[neighbor_idx].neighbors.len() {
let already_linked =
self.points[neighbor_idx].neighbors[layer].contains(&new_idx);
if !already_linked {
self.points[neighbor_idx].neighbors[layer].push(new_idx);
if self.points[neighbor_idx].neighbors[layer].len() > max_neighbors {
let emb = self.points[neighbor_idx].embedding.clone();
let mut scored: Vec<(usize, f32)> = self.points[neighbor_idx].neighbors
[layer]
.iter()
.map(|&n| {
(n, Self::cosine_distance(&self.points[n].embedding, &emb))
})
.collect();
scored.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(Ordering::Equal));
scored.truncate(max_neighbors);
self.points[neighbor_idx].neighbors[layer] =
scored.into_iter().map(|(idx, _)| idx).collect();
}
}
}
}
if let Some(c) = candidates.first() {
current = c.index;
}
}
}
}
#[derive(Serialize, Deserialize)]
struct HnswPointSnapshot {
id: MemoryId,
embedding: Vec<f32>,
neighbors: Vec<Vec<usize>>,
deleted: bool,
}
#[derive(Serialize, Deserialize)]
struct HnswSnapshotV2 {
version: u8,
max_layer: usize,
entry_point: Option<usize>,
ef_construction: usize,
ef_search: usize,
m: usize,
points: Vec<HnswPointSnapshot>,
}
#[derive(Serialize, Deserialize)]
struct HnswSnapshotLegacy {
points: Vec<(MemoryId, Vec<f32>)>,
config: HnswConfig,
}
fn rand_f64() -> f64 {
use std::cell::Cell;
thread_local! {
static SEED: Cell<u64> = const { Cell::new(0x12345678_9abcdef0) };
}
SEED.with(|s| {
let mut x = s.get();
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
s.set(x);
(x as f64) / (u64::MAX as f64)
})
}
#[cfg(test)]
mod tests {
use super::*;
fn make_id(n: u8) -> MemoryId {
MemoryId::from_bytes([n; 32])
}
fn make_embedding(values: &[f32]) -> Vec<f32> {
values.to_vec()
}
#[test]
fn test_new_index() {
let index = HnswIndex::new(HnswConfig::default());
assert!(index.is_empty());
assert_eq!(index.len(), 0);
}
#[test]
fn test_insert_and_search() {
let mut index = HnswIndex::new(HnswConfig::default());
index.insert(make_id(1), make_embedding(&[1.0, 0.0, 0.0]));
index.insert(make_id(2), make_embedding(&[0.0, 1.0, 0.0]));
index.insert(make_id(3), make_embedding(&[1.0, 1.0, 0.0]));
assert_eq!(index.len(), 3);
let results = index.search(&[1.0, 0.0, 0.0], 2);
assert!(!results.is_empty());
assert_eq!(results[0].0, make_id(1));
assert!(results[0].1 > 0.9); }
#[test]
fn test_cosine_distance() {
let a = &[1.0, 0.0, 0.0];
let b = &[1.0, 0.0, 0.0];
let dist = HnswIndex::cosine_distance(a, b);
assert!(dist.abs() < 0.01);
let c = &[0.0, 1.0, 0.0];
let dist2 = HnswIndex::cosine_distance(a, c);
assert!((dist2 - 1.0).abs() < 0.01); }
#[test]
fn test_remove() {
let mut index = HnswIndex::new(HnswConfig::default());
index.insert(make_id(1), make_embedding(&[1.0, 0.0]));
index.insert(make_id(2), make_embedding(&[0.0, 1.0]));
assert_eq!(index.len(), 2);
index.remove(&make_id(1));
assert_eq!(index.len(), 1);
let results = index.search(&[1.0, 0.0], 5);
for (id, _) in &results {
assert_ne!(id, &make_id(1));
}
}
#[test]
fn test_rebuild() {
let mut index = HnswIndex::new(HnswConfig::default());
for i in 0..10u8 {
let v = vec![i as f32, (10 - i) as f32];
index.insert(make_id(i), v);
}
index.remove(&make_id(3));
index.remove(&make_id(7));
assert_eq!(index.len(), 8);
index.rebuild();
assert_eq!(index.len(), 8);
assert_eq!(index.deleted_count, 0);
}
#[test]
fn test_serialize_deserialize() {
let mut index = HnswIndex::new(HnswConfig::default());
index.insert(make_id(1), make_embedding(&[1.0, 0.0, 0.0]));
index.insert(make_id(2), make_embedding(&[0.0, 1.0, 0.0]));
let data = index.serialize().unwrap();
let restored = HnswIndex::deserialize(&data).unwrap();
assert_eq!(restored.len(), 2);
let results = restored.search(&[1.0, 0.0, 0.0], 1);
assert_eq!(results[0].0, make_id(1));
}
#[test]
fn test_stats() {
let mut index = HnswIndex::new(HnswConfig::default());
index.insert(make_id(1), make_embedding(&[1.0, 0.0, 0.0]));
index.insert(make_id(2), make_embedding(&[0.0, 1.0, 0.0]));
let stats = index.stats();
assert_eq!(stats.point_count, 2);
assert_eq!(stats.dimensions, 3);
assert!(stats.memory_bytes > 0);
}
#[test]
fn test_empty_search() {
let index = HnswIndex::new(HnswConfig::default());
let results = index.search(&[1.0, 0.0], 5);
assert!(results.is_empty());
}
#[test]
fn test_empty_embedding() {
let mut index = HnswIndex::new(HnswConfig::default());
index.insert(make_id(1), vec![]); assert_eq!(index.len(), 0);
}
#[test]
fn test_serialize_roundtrip_preserves_topology() {
let mut index = HnswIndex::new(HnswConfig::default());
for i in 0..20u8 {
let mut emb = vec![0.0f32; 8];
emb[i as usize % 8] = 1.0;
emb[(i as usize + 1) % 8] = 0.5;
index.insert(make_id(i), emb);
}
let query = vec![1.0, 0.5, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let before = index.search(&query, 5);
assert!(!before.is_empty(), "should have results before serialize");
let data = index.serialize().unwrap();
let restored = HnswIndex::deserialize(&data).unwrap();
assert_eq!(restored.len(), index.len());
assert_eq!(restored.max_layer, index.max_layer);
assert_eq!(restored.entry_point, index.entry_point);
let after = restored.search(&query, 5);
assert_eq!(
before.iter().map(|(id, _)| id.clone()).collect::<Vec<_>>(),
after.iter().map(|(id, _)| id.clone()).collect::<Vec<_>>(),
"search results must be identical after round-trip"
);
for (b, a) in before.iter().zip(after.iter()) {
assert!(
(b.1 - a.1).abs() < 1e-6,
"scores diverged: {} vs {}",
b.1,
a.1
);
}
}
#[test]
fn test_legacy_format_backward_compat() {
let legacy = HnswSnapshotLegacy {
points: vec![
(make_id(1), vec![1.0, 0.0, 0.0]),
(make_id(2), vec![0.0, 1.0, 0.0]),
],
config: HnswConfig::default(),
};
let data = serde_json::to_vec(&legacy).unwrap();
let restored = HnswIndex::deserialize(&data).unwrap();
assert_eq!(restored.len(), 2);
let results = restored.search(&[1.0, 0.0, 0.0], 1);
assert_eq!(results[0].0, make_id(1));
}
#[test]
fn test_entry_point_updates_on_higher_level() {
let config = HnswConfig {
m: 2,
ef_construction: 10,
ef_search: 10,
};
let mut index = HnswIndex::new(config);
for i in 0..50u8 {
let emb = vec![i as f32, (50 - i) as f32];
index.insert(make_id(i), emb);
}
assert!(
index.max_layer >= 1,
"expected max_layer >= 1 after 50 inserts with m=2, got {}",
index.max_layer
);
if let Some(ep) = index.entry_point {
let ep_level = index.points[ep].neighbors.len().saturating_sub(1);
assert_eq!(
ep_level, index.max_layer,
"entry point level ({}) should equal max_layer ({})",
ep_level, index.max_layer
);
} else {
panic!("entry_point should be Some after inserts");
}
}
#[test]
fn test_deletion_prunes_neighbor_lists() {
let config = HnswConfig {
m: 4,
ef_construction: 20,
ef_search: 10,
};
let mut index = HnswIndex::new(config);
for i in 0..10u8 {
let emb = vec![i as f32, (10 - i) as f32, 1.0];
index.insert(make_id(i), emb);
}
let victim_id = make_id(5);
let &victim_idx = index.id_to_index.get(&victim_id).unwrap();
index.remove(&victim_id);
for (idx, point) in index.points.iter().enumerate() {
if idx == victim_idx {
continue; }
for (layer, neighbors) in point.neighbors.iter().enumerate() {
assert!(
!neighbors.contains(&victim_idx),
"point {} layer {} still references deleted point {}",
idx,
layer,
victim_idx
);
}
}
}
#[test]
fn test_memory_stats_lower_bound() {
let mut index = HnswIndex::new(HnswConfig::default());
let dim = 128;
for i in 0..100u8 {
let mut emb = vec![0.0f32; dim];
emb[i as usize % dim] = 1.0;
index.insert(make_id(i), emb);
}
let stats = index.stats();
assert_eq!(stats.point_count, 100);
assert_eq!(stats.dimensions, dim);
let lower_bound = 100 * (32 + dim * 4);
assert!(
stats.memory_bytes >= lower_bound,
"memory_bytes ({}) should be >= conservative lower bound ({})",
stats.memory_bytes,
lower_bound
);
}
}