use std::path::Path;
use parking_lot::RwLock;
use astraea_core::error::Result;
use astraea_core::traits::VectorIndex;
use astraea_core::types::{DistanceMetric, NodeId, SimilarityResult};
use crate::hnsw::HnswIndex;
const DEFAULT_M: usize = 16;
const DEFAULT_EF_CONSTRUCTION: usize = 200;
const DEFAULT_EF_SEARCH: usize = 50;
pub struct HnswVectorIndex {
inner: RwLock<HnswIndex>,
ef_search: usize,
}
impl HnswVectorIndex {
pub fn new(dimension: usize, metric: DistanceMetric) -> Self {
Self {
inner: RwLock::new(HnswIndex::new(
dimension,
metric,
DEFAULT_M,
DEFAULT_EF_CONSTRUCTION,
)),
ef_search: DEFAULT_EF_SEARCH,
}
}
pub fn with_params(
dimension: usize,
metric: DistanceMetric,
m: usize,
ef_construction: usize,
ef_search: usize,
) -> Self {
Self {
inner: RwLock::new(HnswIndex::new(dimension, metric, m, ef_construction)),
ef_search,
}
}
pub fn with_seed(dimension: usize, metric: DistanceMetric, seed: u64) -> Self {
Self {
inner: RwLock::new(HnswIndex::with_seed(
dimension,
metric,
DEFAULT_M,
DEFAULT_EF_CONSTRUCTION,
seed,
)),
ef_search: DEFAULT_EF_SEARCH,
}
}
pub fn save_to_file(&self, path: &Path) -> Result<()> {
let idx = self.inner.read();
idx.save(path)
}
pub fn load_from_file(path: &Path) -> Result<Self> {
let idx = HnswIndex::load(path)?;
Ok(Self {
inner: RwLock::new(idx),
ef_search: DEFAULT_EF_SEARCH,
})
}
}
impl VectorIndex for HnswVectorIndex {
fn insert(&self, node_id: NodeId, embedding: &[f32]) -> Result<()> {
let mut idx = self.inner.write();
idx.insert(node_id, embedding)
}
fn remove(&self, node_id: NodeId) -> Result<bool> {
let mut idx = self.inner.write();
idx.remove(node_id)
}
fn search(&self, query: &[f32], k: usize) -> Result<Vec<SimilarityResult>> {
let idx = self.inner.read();
let raw_results = idx.search(query, k, self.ef_search)?;
Ok(raw_results
.into_iter()
.map(|(node_id, distance)| SimilarityResult { node_id, distance })
.collect())
}
fn dimension(&self) -> usize {
let idx = self.inner.read();
idx.dimension()
}
fn metric(&self) -> DistanceMetric {
let idx = self.inner.read();
idx.metric()
}
fn len(&self) -> usize {
let idx = self.inner.read();
idx.len()
}
fn node_ids(&self) -> Vec<NodeId> {
let idx = self.inner.read();
idx.node_ids()
}
fn save_to_path(&self, path: &Path) -> Result<()> {
self.save_to_file(path)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_vector_index_trait_basic() {
let idx = HnswVectorIndex::new(3, DistanceMetric::Euclidean);
assert_eq!(idx.dimension(), 3);
assert_eq!(idx.metric(), DistanceMetric::Euclidean);
assert!(idx.is_empty());
assert_eq!(idx.len(), 0);
idx.insert(NodeId(1), &[1.0, 0.0, 0.0]).unwrap();
idx.insert(NodeId(2), &[0.0, 1.0, 0.0]).unwrap();
idx.insert(NodeId(3), &[0.0, 0.0, 1.0]).unwrap();
assert_eq!(idx.len(), 3);
assert!(!idx.is_empty());
let results = idx.search(&[1.0, 0.0, 0.0], 2).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].node_id, NodeId(1));
assert!(results[0].distance < 1e-6);
}
#[test]
fn test_vector_index_trait_remove() {
let idx = HnswVectorIndex::new(2, DistanceMetric::Cosine);
idx.insert(NodeId(1), &[1.0, 0.0]).unwrap();
idx.insert(NodeId(2), &[0.0, 1.0]).unwrap();
assert!(idx.remove(NodeId(1)).unwrap());
assert_eq!(idx.len(), 1);
assert!(!idx.remove(NodeId(99)).unwrap());
}
#[test]
fn test_vector_index_custom_params() {
let idx = HnswVectorIndex::with_params(4, DistanceMetric::DotProduct, 8, 100, 30);
assert_eq!(idx.dimension(), 4);
assert_eq!(idx.metric(), DistanceMetric::DotProduct);
}
#[test]
fn test_node_ids_via_trait() {
let idx: Box<dyn VectorIndex> =
Box::new(HnswVectorIndex::new(2, DistanceMetric::Euclidean));
assert!(idx.node_ids().is_empty());
idx.insert(NodeId(10), &[1.0, 0.0]).unwrap();
idx.insert(NodeId(20), &[0.0, 1.0]).unwrap();
let mut ids = idx.node_ids();
ids.sort();
assert_eq!(ids, vec![NodeId(10), NodeId(20)]);
idx.remove(NodeId(10)).unwrap();
let ids = idx.node_ids();
assert_eq!(ids, vec![NodeId(20)]);
}
#[test]
fn test_save_to_path_via_trait() {
use astraea_core::traits::VectorIndex as VTrait;
let tmp = tempfile::NamedTempFile::new().unwrap();
let path = tmp.path().to_owned();
drop(tmp);
let original: Box<dyn VTrait> =
Box::new(HnswVectorIndex::new(3, DistanceMetric::Euclidean));
original.insert(NodeId(1), &[1.0, 0.0, 0.0]).unwrap();
original.insert(NodeId(2), &[0.0, 1.0, 0.0]).unwrap();
original.save_to_path(&path).unwrap();
let loaded = HnswVectorIndex::load_from_file(&path).unwrap();
let mut ids = loaded.node_ids();
ids.sort();
assert_eq!(ids, vec![NodeId(1), NodeId(2)]);
assert_eq!(loaded.dimension(), 3);
}
}