use ahash::AHashMap;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HnswGraph {
pub entry_point: Option<u64>,
pub max_level: usize,
id_to_index: AHashMap<u64, usize>,
index_to_id: Vec<u64>,
nodes: Vec<Vec<Vec<u64>>>,
pub m: usize,
pub m_max: usize, pub m_max_0: usize, pub ef_construction: usize,
pub level_mult: f64,
max_doc_id: u64,
}
impl HnswGraph {
#[allow(clippy::too_many_arguments)]
pub fn new(
entry_point: Option<u64>,
max_level: usize,
nodes_map: std::collections::HashMap<u64, Vec<Vec<u64>>>,
m: usize,
m_max: usize,
m_max_0: usize,
ef_construction: usize,
level_mult: f64,
) -> Self {
let mut id_to_index = AHashMap::with_capacity(nodes_map.len());
let mut index_to_id = Vec::with_capacity(nodes_map.len());
let mut nodes = Vec::with_capacity(nodes_map.len());
let mut max_doc_id: u64 = 0;
for (doc_id, layers) in nodes_map {
let index = nodes.len();
id_to_index.insert(doc_id, index);
index_to_id.push(doc_id);
nodes.push(layers);
if doc_id > max_doc_id {
max_doc_id = doc_id;
}
}
Self {
entry_point,
max_level,
id_to_index,
index_to_id,
nodes,
m,
m_max,
m_max_0,
ef_construction,
level_mult,
max_doc_id,
}
}
pub fn max_doc_id(&self) -> u64 {
self.max_doc_id
}
pub fn get_neighbors(&self, doc_id: u64, level: usize) -> Option<&Vec<u64>> {
let &index = self.id_to_index.get(&doc_id)?;
self.nodes.get(index).and_then(|levels| levels.get(level))
}
pub fn set_neighbors(&mut self, doc_id: u64, level: usize, neighbors: Vec<u64>) {
let index = self.get_or_create_index(doc_id);
if level < self.nodes[index].len() {
self.nodes[index][level] = neighbors;
}
}
fn get_or_create_index(&mut self, doc_id: u64) -> usize {
if let Some(&index) = self.id_to_index.get(&doc_id) {
index
} else {
let index = self.nodes.len();
self.id_to_index.insert(doc_id, index);
self.index_to_id.push(doc_id);
self.nodes.push(Vec::new());
if doc_id > self.max_doc_id {
self.max_doc_id = doc_id;
}
index
}
}
pub fn contains_node(&self, doc_id: &u64) -> bool {
self.id_to_index.contains_key(doc_id)
}
pub fn node_count(&self) -> usize {
self.nodes.len()
}
pub fn get_node_layers(&self, doc_id: &u64) -> Option<&Vec<Vec<u64>>> {
let &index = self.id_to_index.get(doc_id)?;
self.nodes.get(index)
}
pub fn iter_nodes(&self) -> impl Iterator<Item = (u64, &Vec<Vec<u64>>)> {
self.index_to_id
.iter()
.zip(self.nodes.iter())
.map(|(&doc_id, layers)| (doc_id, layers))
}
pub fn into_iter_nodes(self) -> impl Iterator<Item = (u64, Vec<Vec<u64>>)> {
self.index_to_id.into_iter().zip(self.nodes)
}
pub fn sorted_nodes(&self) -> Vec<(u64, &Vec<Vec<u64>>)> {
let mut pairs: Vec<_> = self.iter_nodes().collect();
pairs.sort_by_key(|(id, _)| *id);
pairs
}
}
#[derive(Debug, Clone)]
pub struct OrdinalHnswGraph {
entry_point: Option<u32>,
max_level: usize,
doc_ids: std::sync::Arc<[u64]>,
nodes: Vec<Vec<Vec<u32>>>,
}
impl OrdinalHnswGraph {
pub fn from_parts(
entry_point: Option<u32>,
max_level: usize,
doc_ids: std::sync::Arc<[u64]>,
nodes: Vec<Vec<Vec<u32>>>,
) -> crate::error::Result<Self> {
let node_count = doc_ids.len();
if nodes.len() != node_count {
return Err(crate::error::LaurusError::index(format!(
"HNSW ordinal graph corrupt: {} nodes for {} unique record doc ids",
nodes.len(),
node_count
)));
}
if let Some(entry) = entry_point
&& entry as usize >= node_count
{
return Err(crate::error::LaurusError::index(format!(
"HNSW ordinal graph corrupt: entry ordinal {entry} out of range \
(node count {node_count})"
)));
}
for (ord, layers) in nodes.iter().enumerate() {
for neighbors in layers {
for &n in neighbors {
if n as usize >= node_count {
return Err(crate::error::LaurusError::index(format!(
"HNSW ordinal graph corrupt: node {ord} has neighbour \
ordinal {n} out of range (node count {node_count})"
)));
}
}
}
}
Ok(Self {
entry_point,
max_level,
doc_ids,
nodes,
})
}
pub fn entry_point(&self) -> Option<u32> {
self.entry_point
}
pub fn max_level(&self) -> usize {
self.max_level
}
pub fn node_count(&self) -> usize {
self.nodes.len()
}
#[inline]
pub fn doc_id(&self, ord: u32) -> u64 {
self.doc_ids[ord as usize]
}
#[inline]
pub fn neighbors(&self, ord: u32, level: usize) -> Option<&[u32]> {
self.nodes[ord as usize].get(level).map(Vec::as_slice)
}
pub fn doc_ids(&self) -> &std::sync::Arc<[u64]> {
&self.doc_ids
}
pub fn ord_of(&self, doc_id: u64) -> Option<u32> {
self.doc_ids
.binary_search(&doc_id)
.ok()
.map(|ord| ord as u32)
}
pub fn iter_nodes(&self) -> impl Iterator<Item = (u64, &Vec<Vec<u32>>)> {
self.doc_ids
.iter()
.zip(self.nodes.iter())
.map(|(&doc_id, layers)| (doc_id, layers))
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
fn doc_ids(ids: &[u64]) -> Arc<[u64]> {
Arc::from(ids.to_vec())
}
#[test]
fn ordinal_graph_from_parts_roundtrips_accessors() {
let nodes = vec![
vec![vec![1, 2], vec![1]],
vec![vec![0, 2]],
vec![vec![0, 1]],
];
let g = OrdinalHnswGraph::from_parts(Some(0), 1, doc_ids(&[10, 20, 30]), nodes).unwrap();
assert_eq!(g.entry_point(), Some(0));
assert_eq!(g.max_level(), 1);
assert_eq!(g.node_count(), 3);
assert_eq!(g.doc_id(0), 10);
assert_eq!(g.doc_id(2), 30);
assert_eq!(g.neighbors(0, 0), Some(&[1u32, 2][..]));
assert_eq!(g.neighbors(0, 1), Some(&[1u32][..]));
assert_eq!(g.neighbors(1, 1), None);
assert_eq!(g.ord_of(20), Some(1));
assert_eq!(g.ord_of(15), None);
let collected: Vec<u64> = g.iter_nodes().map(|(id, _)| id).collect();
assert_eq!(collected, vec![10, 20, 30]);
}
#[test]
fn ordinal_graph_empty_is_valid() {
let g = OrdinalHnswGraph::from_parts(None, 0, doc_ids(&[]), Vec::new()).unwrap();
assert_eq!(g.entry_point(), None);
assert_eq!(g.node_count(), 0);
assert_eq!(g.ord_of(1), None);
}
#[test]
fn ordinal_graph_rejects_node_count_mismatch() {
let err =
OrdinalHnswGraph::from_parts(None, 0, doc_ids(&[10, 20]), vec![vec![]]).unwrap_err();
assert!(err.to_string().contains("unique record doc ids"));
}
#[test]
fn ordinal_graph_rejects_out_of_range_entry_point() {
let err =
OrdinalHnswGraph::from_parts(Some(2), 0, doc_ids(&[10, 20]), vec![vec![], vec![]])
.unwrap_err();
assert!(err.to_string().contains("entry ordinal"));
}
#[test]
fn ordinal_graph_rejects_out_of_range_neighbor() {
let err = OrdinalHnswGraph::from_parts(
None,
0,
doc_ids(&[10, 20]),
vec![vec![vec![7]], vec![vec![0]]],
)
.unwrap_err();
assert!(err.to_string().contains("neighbour ordinal 7"));
}
}