use crate::error::NopalError;
use crate::types::NodeId;
use hnsw_rs::prelude::*;
use std::collections::HashMap;
const DEFAULT_MAX_NB_CONNECTION: usize = 24;
const DEFAULT_EF_CONSTRUCTION: usize = 400;
const DEFAULT_MAX_LAYER: usize = 16;
const DEFAULT_EF_SEARCH: usize = 30;
pub struct HnswIndex {
inner: Hnsw<'static, f32, DistCosine>,
id_map: HashMap<usize, NodeId>,
reverse_map: HashMap<NodeId, usize>,
model: String,
dimension: usize,
next_data_id: usize,
}
impl HnswIndex {
pub fn new(
model: impl Into<String>,
dimension: usize,
max_elements: usize,
) -> Self {
let inner = Hnsw::<f32, DistCosine>::new(
DEFAULT_MAX_NB_CONNECTION,
max_elements,
DEFAULT_MAX_LAYER,
DEFAULT_EF_CONSTRUCTION,
DistCosine {},
);
Self {
inner,
id_map: HashMap::new(),
reverse_map: HashMap::new(),
model: model.into(),
dimension,
next_data_id: 0,
}
}
pub fn with_params(
model: impl Into<String>,
dimension: usize,
max_elements: usize,
max_nb_connection: usize,
ef_construction: usize,
max_layer: usize,
) -> Self {
let inner = Hnsw::<f32, DistCosine>::new(
max_nb_connection,
max_elements,
max_layer,
ef_construction,
DistCosine {},
);
Self {
inner,
id_map: HashMap::new(),
reverse_map: HashMap::new(),
model: model.into(),
dimension,
next_data_id: 0,
}
}
pub fn build_batch(
vectors: Vec<(NodeId, Vec<f32>)>,
model: impl Into<String>,
dimension: usize,
) -> Result<Self, NopalError> {
if vectors.is_empty() {
return Err(NopalError::custom("HnswIndex::build_batch: no vectors provided"));
}
for (node_id, vec) in &vectors {
if vec.len() != dimension {
return Err(NopalError::custom(format!(
"HnswIndex::build_batch: node {} has dimension {}, expected {}",
node_id, vec.len(), dimension
)));
}
}
let model_str = model.into();
let nb_elements = vectors.len();
let mut index = Self::new(&model_str, dimension, nb_elements);
let mut owned_vectors: Vec<Vec<f32>> = Vec::with_capacity(nb_elements);
let mut data_ids: Vec<usize> = Vec::with_capacity(nb_elements);
for (node_id, vec) in vectors {
let data_id = index.next_data_id;
index.next_data_id += 1;
index.id_map.insert(data_id, node_id);
index.reverse_map.insert(node_id, data_id);
owned_vectors.push(vec);
data_ids.push(data_id);
}
let insert_data: Vec<(&Vec<f32>, usize)> = owned_vectors
.iter()
.zip(data_ids.iter())
.map(|(v, &id)| (v, id))
.collect();
index.inner.parallel_insert(&insert_data);
index.inner.set_searching_mode(true);
Ok(index)
}
pub fn insert(&mut self, node_id: NodeId, vector: Vec<f32>) -> Result<(), NopalError> {
if vector.len() != self.dimension {
return Err(NopalError::custom(format!(
"HnswIndex({}): expected dimension {}, got {}",
self.model, self.dimension, vector.len()
)));
}
if self.reverse_map.contains_key(&node_id) {
return Err(NopalError::custom(format!(
"HnswIndex({}): node {} already indexed — remove first to update",
self.model, node_id
)));
}
let data_id = self.next_data_id;
self.next_data_id += 1;
self.inner.insert((&vector, data_id));
self.id_map.insert(data_id, node_id);
self.reverse_map.insert(node_id, data_id);
Ok(())
}
pub fn search_knn(
&self,
query: &[f32],
k: usize,
) -> Result<Vec<(NodeId, f32)>, NopalError> {
self.search_knn_with_ef(query, k, DEFAULT_EF_SEARCH)
}
pub fn search_knn_with_ef(
&self,
query: &[f32],
k: usize,
ef_search: usize,
) -> Result<Vec<(NodeId, f32)>, NopalError> {
if query.len() != self.dimension {
return Err(NopalError::custom(format!(
"HnswIndex({}): query dimension {} != index dimension {}",
self.model, query.len(), self.dimension
)));
}
if self.id_map.is_empty() {
return Ok(Vec::new());
}
let neighbors = self.inner.search(query, k, ef_search);
let mut results = Vec::with_capacity(neighbors.len());
for neighbor in neighbors {
let data_id = neighbor.d_id;
if let Some(&node_id) = self.id_map.get(&data_id) {
results.push((node_id, neighbor.distance));
}
}
Ok(results)
}
pub fn search_knn_filtered<F>(
&self,
query: &[f32],
k: usize,
ef_search: usize,
filter: F,
) -> Result<Vec<(NodeId, f32)>, NopalError>
where
F: Fn(&NodeId) -> bool,
{
let over_fetch = k * 4;
let mut results = self.search_knn_with_ef(query, over_fetch, ef_search)?;
results.retain(|(node_id, _)| filter(node_id));
results.truncate(k);
Ok(results)
}
pub fn len(&self) -> usize {
self.id_map.len()
}
pub fn is_empty(&self) -> bool {
self.id_map.is_empty()
}
pub fn model(&self) -> &str {
&self.model
}
pub fn dimension(&self) -> usize {
self.dimension
}
#[allow(dead_code)] pub(crate) fn inner(&self) -> &Hnsw<'static, f32, DistCosine> {
&self.inner
}
#[allow(dead_code)] pub(crate) fn id_map(&self) -> &HashMap<usize, NodeId> {
&self.id_map
}
#[allow(dead_code)] pub(crate) fn reverse_map(&self) -> &HashMap<NodeId, usize> {
&self.reverse_map
}
#[allow(dead_code)] pub(crate) fn next_data_id(&self) -> usize {
self.next_data_id
}
}
pub type EmbeddingIndex = HnswIndex;
#[cfg(test)]
mod tests {
use super::*;
use uuid::Uuid;
#[test]
fn test_build_batch_and_search() {
let id_a = Uuid::new_v4();
let id_b = Uuid::new_v4();
let id_c = Uuid::new_v4();
let vectors = vec![
(id_a, vec![1.0, 0.0, 0.0]),
(id_b, vec![0.0, 1.0, 0.0]),
(id_c, vec![0.9, 0.1, 0.0]),
];
let index = HnswIndex::build_batch(vectors, "test", 3).unwrap();
let results = index.search_knn(&[1.0, 0.0, 0.0], 1).unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, id_a);
assert!(results[0].1 < 0.01, "self-distance should be ~0, got {}", results[0].1);
}
#[test]
fn test_incremental_insert() {
let mut index = HnswIndex::new("test", 2, 10);
let id_a = Uuid::new_v4();
let id_b = Uuid::new_v4();
index.insert(id_a, vec![1.0, 0.0]).unwrap();
index.insert(id_b, vec![0.0, 1.0]).unwrap();
assert_eq!(index.len(), 2);
let results = index.search_knn(&[0.9, 0.1], 1).unwrap();
assert_eq!(results[0].0, id_a);
}
#[test]
fn test_duplicate_insert_returns_error() {
let mut index = HnswIndex::new("test", 2, 10);
let id = Uuid::new_v4();
index.insert(id, vec![1.0, 0.0]).unwrap();
let result = index.insert(id, vec![0.0, 1.0]);
assert!(result.is_err());
}
#[test]
fn test_search_top_k() {
let ids: Vec<Uuid> = (0..10).map(|_| Uuid::new_v4()).collect();
let vectors: Vec<(Uuid, Vec<f32>)> = ids
.iter()
.enumerate()
.map(|(i, &id)| {
let mut v = vec![0.0; 4];
v[0] = 1.0 - (i as f32 * 0.1);
v[1] = i as f32 * 0.1;
(id, v)
})
.collect();
let index = HnswIndex::build_batch(vectors, "test", 4).unwrap();
let results = index.search_knn(&[1.0, 0.0, 0.0, 0.0], 3).unwrap();
assert_eq!(results.len(), 3);
assert_eq!(results[0].0, ids[0]);
}
#[test]
fn test_dimension_mismatch_on_insert() {
let mut index = HnswIndex::new("test", 3, 10);
let result = index.insert(Uuid::new_v4(), vec![1.0, 2.0]); assert!(result.is_err());
}
#[test]
fn test_dimension_mismatch_on_search() {
let index = HnswIndex::build_batch(
vec![(Uuid::new_v4(), vec![1.0, 0.0])],
"test",
2,
)
.unwrap();
let result = index.search_knn(&[1.0, 0.0, 0.0], 1); assert!(result.is_err());
}
#[test]
fn test_build_batch_empty_returns_error() {
let result = HnswIndex::build_batch(Vec::new(), "test", 2);
assert!(result.is_err());
}
#[test]
fn test_len_and_is_empty() {
let index = HnswIndex::new("test", 2, 10);
assert_eq!(index.len(), 0);
assert!(index.is_empty());
let index = HnswIndex::build_batch(
vec![(Uuid::new_v4(), vec![1.0, 0.0])],
"test",
2,
)
.unwrap();
assert_eq!(index.len(), 1);
assert!(!index.is_empty());
}
#[test]
fn test_search_empty_index_returns_empty() {
let index = HnswIndex::new("test", 2, 10);
let results = index.search_knn(&[1.0, 0.0], 5).unwrap();
assert!(results.is_empty());
}
#[test]
fn test_filtered_search() {
let id_a = Uuid::new_v4();
let id_b = Uuid::new_v4();
let id_c = Uuid::new_v4();
let vectors = vec![
(id_a, vec![1.0, 0.0, 0.0]),
(id_b, vec![0.9, 0.1, 0.0]),
(id_c, vec![0.0, 1.0, 0.0]),
];
let index = HnswIndex::build_batch(vectors, "test", 3).unwrap();
let allowed = vec![id_b, id_c];
let results = index
.search_knn_filtered(&[1.0, 0.0, 0.0], 1, DEFAULT_EF_SEARCH, |nid| {
allowed.contains(nid)
})
.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, id_b);
}
#[test]
fn test_model_and_dimension_accessors() {
let index = HnswIndex::new("minilm", 384, 100);
assert_eq!(index.model(), "minilm");
assert_eq!(index.dimension(), 384);
}
}