#![allow(clippy::doc_markdown)]
use std::io::{BufReader, BufWriter, Read, Write};
use std::path::Path;
use std::sync::Arc;
use instant_distance::{Builder, HnswMap, Point as IdPoint, Search};
use mnem_core::codec::from_canonical_bytes;
use mnem_core::error::{Error, RepoError, StoreError};
use mnem_core::id::{Cid, NodeId};
use mnem_core::index::vector::{VectorHit, VectorIndex};
use mnem_core::objects::{Dtype, Embedding, Node};
use mnem_core::prolly::Cursor;
use mnem_core::repo::ReadonlyRepo;
use mnem_core::store::Blockstore;
#[derive(Clone, Debug)]
pub struct HnswConfig {
pub ef_construction: usize,
pub ef_search: usize,
pub seed: u64,
}
impl Default for HnswConfig {
fn default() -> Self {
Self {
ef_construction: 200,
ef_search: 100,
seed: 0x6DEF_1EE7_5CE8_7D55,
}
}
}
#[derive(Clone, Debug)]
pub(crate) struct Point {
pub(crate) vec: Vec<f32>,
}
impl IdPoint for Point {
fn distance(&self, other: &Self) -> f32 {
debug_assert_eq!(self.vec.len(), other.vec.len());
let mut acc = 0.0_f32;
for (x, y) in self.vec.iter().zip(other.vec.iter()) {
let d = x - y;
acc += d * d;
}
acc
}
}
pub struct HnswVectorIndex {
model: String,
dim: u32,
pub(crate) ids: Vec<NodeId>,
pub(crate) points: Vec<Point>,
inner: HnswMap<Point, usize>,
ef_search: usize,
}
impl HnswVectorIndex {
pub fn points_iter(&self) -> impl Iterator<Item = (NodeId, &[f32])> + '_ {
self.ids
.iter()
.zip(self.points.iter())
.map(|(id, p)| (*id, p.vec.as_slice()))
}
#[must_use]
pub fn points_len(&self) -> usize {
self.ids.len()
}
}
impl std::fmt::Debug for HnswVectorIndex {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HnswVectorIndex")
.field("model", &self.model)
.field("dim", &self.dim)
.field("len", &self.ids.len())
.finish()
}
}
impl HnswVectorIndex {
pub fn build_from_repo(repo: &ReadonlyRepo, model: &str) -> Result<Self, Error> {
Self::build_from_repo_with(repo, model, HnswConfig::default())
}
pub fn build_from_repo_with(
repo: &ReadonlyRepo,
model: &str,
cfg: HnswConfig,
) -> Result<Self, Error> {
let bs: Arc<dyn Blockstore> = repo.blockstore().clone();
let Some(commit) = repo.head_commit() else {
return Err(RepoError::Uninitialized.into());
};
let mut ids: Vec<NodeId> = Vec::new();
let mut points: Vec<Point> = Vec::new();
let mut dim: Option<u32> = None;
let cursor = Cursor::new(&*bs, &commit.nodes)?;
for entry in cursor {
let (_k, node_cid) = entry?;
let bytes = bs
.get(&node_cid)
.map_err(Error::from)?
.ok_or_else(|| Error::from(RepoError::NotFound))?;
let node: Node = from_canonical_bytes(&bytes).map_err(Error::from)?;
if repo.is_tombstoned(&node.id) {
continue;
}
if mnem_core::anchor::is_system_node(&node) {
continue;
}
let Some(embed) = repo.embedding_for(&node_cid, model)? else {
continue;
};
embed.validate()?;
if let Some(d) = dim {
if embed.dim != d {
continue;
}
} else {
dim = Some(embed.dim);
}
let Some(vec_f32) = decode_to_f32(&embed) else {
continue;
};
let normalised = normalise(vec_f32);
ids.push(node.id);
points.push(Point { vec: normalised });
}
let dim = dim.unwrap_or(0);
if points.is_empty() {
return Ok(Self {
model: model.into(),
dim,
ids: Vec::new(),
points: Vec::new(),
inner: Builder::default().build(Vec::<Point>::new(), Vec::<usize>::new()),
ef_search: cfg.ef_search,
});
}
let values: Vec<usize> = (0..points.len()).collect();
let points_retained = points.clone();
let inner = Builder::default()
.ef_construction(cfg.ef_construction)
.seed(cfg.seed)
.build(points, values);
Ok(Self {
model: model.into(),
dim,
ids,
points: points_retained,
inner,
ef_search: cfg.ef_search,
})
}
#[doc(hidden)]
#[must_use]
pub fn from_parts_for_test(
model: &str,
dim: u32,
ids: Vec<NodeId>,
normalised_vecs: Vec<Vec<f32>>,
cfg: &HnswConfig,
) -> Self {
assert_eq!(ids.len(), normalised_vecs.len(), "ids/vecs length mismatch");
let points: Vec<Point> = normalised_vecs
.into_iter()
.map(|v| Point { vec: v })
.collect();
if points.is_empty() {
return Self {
model: model.into(),
dim,
ids,
points,
inner: Builder::default().build(Vec::<Point>::new(), Vec::<usize>::new()),
ef_search: cfg.ef_search,
};
}
let values: Vec<usize> = (0..points.len()).collect();
let points_retained = points.clone();
let inner = Builder::default()
.ef_construction(cfg.ef_construction)
.seed(cfg.seed)
.build(points, values);
Self {
model: model.into(),
dim,
ids,
points: points_retained,
inner,
ef_search: cfg.ef_search,
}
}
pub fn save_to_path(&self, path: &Path, op_id: &Cid) -> Result<(), Error> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
}
let file =
std::fs::File::create(path).map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let mut w = BufWriter::new(file);
w.write_all(b"MNEMHNSW")
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
w.write_all(&1u32.to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let op_id_bytes = op_id.to_bytes();
let op_id_len = u32::try_from(op_id_bytes.len()).expect("op_id too large");
w.write_all(&op_id_len.to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
w.write_all(&op_id_bytes)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let model_bytes = self.model.as_bytes();
let model_len = u32::try_from(model_bytes.len()).expect("model too large");
w.write_all(&model_len.to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
w.write_all(model_bytes)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
w.write_all(&self.dim.to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
w.write_all(&(self.ef_search as u64).to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
w.write_all(&(self.ef_search as u64).to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
w.write_all(&0u64.to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let n_points = u64::try_from(self.ids.len()).expect("too many points");
w.write_all(&n_points.to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
for id in &self.ids {
w.write_all(id.as_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
}
for point in &self.points {
for &val in &point.vec {
w.write_all(&val.to_le_bytes())
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
}
}
w.flush()
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
Ok(())
}
pub fn load_from_path(
path: &Path,
expected_op_id: &Cid,
cfg: &HnswConfig,
) -> Result<Option<Self>, Error> {
if !path.exists() {
return Ok(None);
}
let file =
std::fs::File::open(path).map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let mut r = BufReader::new(file);
let mut magic = [0u8; 8];
r.read_exact(&mut magic)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
if &magic != b"MNEMHNSW" {
tracing::debug!("ann cache: bad magic, ignoring {:?}", path);
return Ok(None);
}
let mut ver_buf = [0u8; 4];
r.read_exact(&mut ver_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let version = u32::from_le_bytes(ver_buf);
if version != 1 {
tracing::debug!(
"ann cache: unsupported version {}, ignoring {:?}",
version,
path
);
return Ok(None);
}
let mut len_buf = [0u8; 4];
r.read_exact(&mut len_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let op_id_len = u32::from_le_bytes(len_buf) as usize;
let mut op_id_bytes = vec![0u8; op_id_len];
r.read_exact(&mut op_id_bytes)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let stored_op_id = Cid::from_bytes(&op_id_bytes).map_err(|e| Error::from(e))?;
if &stored_op_id != expected_op_id {
tracing::debug!("ann cache: stale op_id, ignoring {:?}", path);
return Ok(None);
}
r.read_exact(&mut len_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let model_len = u32::from_le_bytes(len_buf) as usize;
let mut model_bytes = vec![0u8; model_len];
r.read_exact(&mut model_bytes)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let model = String::from_utf8(model_bytes)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let mut dim_buf = [0u8; 4];
r.read_exact(&mut dim_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let dim = u32::from_le_bytes(dim_buf);
let mut u64_buf = [0u8; 8];
r.read_exact(&mut u64_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
r.read_exact(&mut u64_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
r.read_exact(&mut u64_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
r.read_exact(&mut u64_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
let n_points = u64::from_le_bytes(u64_buf) as usize;
let mut ids: Vec<NodeId> = Vec::with_capacity(n_points);
for _ in 0..n_points {
let mut id_buf = [0u8; 16];
r.read_exact(&mut id_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
ids.push(NodeId::from_bytes_raw(id_buf));
}
let mut normalised_vecs: Vec<Vec<f32>> = Vec::with_capacity(n_points);
let dim_usize = dim as usize;
for _ in 0..n_points {
let mut vec = Vec::with_capacity(dim_usize);
for _ in 0..dim_usize {
let mut f_buf = [0u8; 4];
r.read_exact(&mut f_buf)
.map_err(|e| Error::from(StoreError::Io(e.to_string())))?;
vec.push(f32::from_le_bytes(f_buf));
}
normalised_vecs.push(vec);
}
Ok(Some(Self::from_parts_for_test(
&model,
dim,
ids,
normalised_vecs,
cfg,
)))
}
}
impl VectorIndex for HnswVectorIndex {
fn model(&self) -> &str {
&self.model
}
fn dim(&self) -> u32 {
self.dim
}
fn search(&self, query: &[f32], k: usize) -> Result<Vec<VectorHit>, Error> {
if self.dim == 0 && self.ids.is_empty() {
return Ok(Vec::new());
}
if query.len() != self.dim as usize {
return Err(RepoError::VectorDimMismatch {
index_dim: self.dim,
query_dim: query.len(),
}
.into());
}
if k == 0 {
return Ok(Vec::new());
}
let q = Point {
vec: normalise(query.to_vec()),
};
let mut searcher = Search::default();
let fetch = std::cmp::max(k, self.ef_search);
let mut hits: Vec<VectorHit> = Vec::with_capacity(k);
for item in self.inner.search(&q, &mut searcher).take(fetch) {
let ord = *item.value;
let node_id = self.ids[ord];
let score = 1.0 - item.distance * 0.5;
hits.push(VectorHit::new(node_id, score));
}
hits.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.node_id.cmp(&b.node_id))
});
hits.truncate(k);
Ok(hits)
}
fn len(&self) -> usize {
self.ids.len()
}
}
fn decode_to_f32(embed: &Embedding) -> Option<Vec<f32>> {
let dim = embed.dim as usize;
let bytes = &embed.vector;
if bytes.len() != dim * embed.dtype.byte_width() {
return None;
}
match embed.dtype {
Dtype::F32 => {
let mut out = Vec::with_capacity(dim);
for chunk in bytes.chunks_exact(4) {
out.push(f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]));
}
Some(out)
}
Dtype::F64 => {
let mut out = Vec::with_capacity(dim);
for chunk in bytes.chunks_exact(8) {
out.push(f64::from_le_bytes([
chunk[0], chunk[1], chunk[2], chunk[3], chunk[4], chunk[5], chunk[6], chunk[7],
]) as f32);
}
Some(out)
}
_ => None,
}
}
fn normalise(mut v: Vec<f32>) -> Vec<f32> {
let mut sq = 0.0_f32;
for x in &v {
sq += x * x;
}
if sq > 0.0 {
let inv = sq.sqrt().recip();
for x in &mut v {
*x *= inv;
}
}
v
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_build_returns_len_zero_index() {
let cfg = HnswConfig::default();
let built = Builder::default()
.ef_construction(cfg.ef_construction)
.seed(cfg.seed)
.build(Vec::<Point>::new(), Vec::<usize>::new());
let idx = HnswVectorIndex {
model: "m".into(),
dim: 0,
ids: Vec::new(),
points: Vec::new(),
inner: built,
ef_search: cfg.ef_search,
};
assert!(idx.is_empty());
let hits = idx.search(&[0.0_f32; 3], 5).unwrap();
assert!(hits.is_empty());
}
#[test]
fn dim_mismatch_errors() {
use mnem_core::error::RepoError;
let points = vec![
Point {
vec: normalise(vec![1.0, 0.0, 0.0]),
},
Point {
vec: normalise(vec![0.0, 1.0, 0.0]),
},
];
let values = vec![0_usize, 1];
let points_retained = points.clone();
let inner = Builder::default().build(points, values);
let idx = HnswVectorIndex {
model: "m".into(),
dim: 3,
ids: vec![NodeId::new_v7(), NodeId::new_v7()],
points: points_retained,
inner,
ef_search: 10,
};
let err = idx.search(&[1.0, 0.0], 1).unwrap_err();
assert!(matches!(
err,
Error::Repo(RepoError::VectorDimMismatch {
index_dim: 3,
query_dim: 2,
})
));
}
#[test]
fn identical_query_is_top_hit() {
let id_a = NodeId::new_v7();
let id_b = NodeId::new_v7();
let points = vec![
Point {
vec: normalise(vec![1.0, 0.0, 0.0]),
},
Point {
vec: normalise(vec![0.0, 1.0, 0.0]),
},
];
let points_retained = points.clone();
let inner = Builder::default().build(points, vec![0_usize, 1]);
let idx = HnswVectorIndex {
model: "m".into(),
dim: 3,
ids: vec![id_a, id_b],
points: points_retained,
inner,
ef_search: 10,
};
let hits = idx.search(&[1.0, 0.0, 0.0], 2).unwrap();
assert_eq!(hits[0].node_id, id_a, "exact match should rank #1");
assert!(
(hits[0].score - 1.0).abs() < 1e-5,
"expected cos == 1, got {}",
hits[0].score
);
}
#[test]
fn score_is_cosine_not_euclidean() {
let id_a = NodeId::new_v7();
let id_b = NodeId::new_v7();
let points = vec![
Point {
vec: normalise(vec![1.0, 0.0]),
},
Point {
vec: normalise(vec![0.0, 1.0]),
},
];
let points_retained = points.clone();
let inner = Builder::default().build(points, vec![0_usize, 1]);
let idx = HnswVectorIndex {
model: "m".into(),
dim: 2,
ids: vec![id_a, id_b],
points: points_retained,
inner,
ef_search: 10,
};
let hits = idx.search(&[1.0, 0.0], 2).unwrap();
let orth = hits.iter().find(|h| h.node_id == id_b).unwrap();
assert!(
orth.score.abs() < 1e-5,
"expected orthogonal cos ~= 0; got {}",
orth.score
);
}
fn f32_embed(model: &str, v: &[f32]) -> Embedding {
let mut bytes = Vec::with_capacity(v.len() * 4);
for x in v {
bytes.extend_from_slice(&x.to_le_bytes());
}
Embedding {
model: model.to_string(),
dtype: Dtype::F32,
dim: u32::try_from(v.len()).expect("test vec fits in u32"),
vector: bytes::Bytes::from(bytes),
}
}
fn stores() -> (
Arc<dyn mnem_core::store::Blockstore>,
Arc<dyn mnem_core::store::OpHeadsStore>,
) {
(
Arc::new(mnem_core::store::MemoryBlockstore::new()),
Arc::new(mnem_core::store::MemoryOpHeadsStore::new()),
)
}
#[test]
fn build_from_repo_reads_sidecar_embeddings() {
let (bs, ohs) = stores();
let repo = ReadonlyRepo::init(bs, ohs).unwrap();
let mut tx = repo.start_transaction();
let id_a = NodeId::from_bytes_raw([1u8; 16]);
let id_b = NodeId::from_bytes_raw([2u8; 16]);
let cid_a = tx.add_node(&Node::new(id_a, "Doc")).unwrap();
let cid_b = tx.add_node(&Node::new(id_b, "Doc")).unwrap();
tx.set_embedding(cid_a, "mA".into(), f32_embed("mA", &[1.0, 0.0]))
.unwrap();
tx.set_embedding(cid_b, "mA".into(), f32_embed("mA", &[0.0, 1.0]))
.unwrap();
let id_c = NodeId::from_bytes_raw([3u8; 16]);
let cid_c = tx.add_node(&Node::new(id_c, "Doc")).unwrap();
tx.set_embedding(cid_c, "mB".into(), f32_embed("mB", &[1.0, 0.0]))
.unwrap();
tx.add_node(&Node::new(NodeId::from_bytes_raw([4u8; 16]), "Doc"))
.unwrap();
let repo = tx.commit("t", "seed").unwrap();
let idx = HnswVectorIndex::build_from_repo(&repo, "mA").unwrap();
assert_eq!(idx.len(), 2, "only the two mA nodes should index");
assert_eq!(idx.dim(), 2);
let hits = idx.search(&[1.0, 0.0], 2).unwrap();
assert_eq!(hits[0].node_id, id_a, "exact-match node should rank #1");
assert!(
(hits[0].score - 1.0).abs() < 1e-5,
"expected cos == 1, got {}",
hits[0].score
);
}
fn small_index() -> (HnswVectorIndex, Vec<NodeId>) {
let id_a = NodeId::from_bytes_raw([10u8; 16]);
let id_b = NodeId::from_bytes_raw([20u8; 16]);
let id_c = NodeId::from_bytes_raw([30u8; 16]);
let ids = vec![id_a, id_b, id_c];
let vecs = vec![
normalise(vec![1.0, 0.0, 0.0]),
normalise(vec![0.0, 1.0, 0.0]),
normalise(vec![0.0, 0.0, 1.0]),
];
let cfg = HnswConfig::default();
let idx = HnswVectorIndex::from_parts_for_test("test-model", 3, ids.clone(), vecs, &cfg);
(idx, ids)
}
fn make_op_id(seed: &[u8]) -> mnem_core::id::Cid {
use mnem_core::id::{CODEC_RAW, Multihash};
mnem_core::id::Cid::new(CODEC_RAW, Multihash::sha2_256(seed))
}
#[test]
fn ann_cache_round_trip() {
let (idx, ids) = small_index();
let op_id = make_op_id(b"test-op-1");
let cfg = HnswConfig::default();
let path = std::env::temp_dir().join("mnem_ann_cache_round_trip.bin");
idx.save_to_path(&path, &op_id).expect("save_to_path");
let loaded = HnswVectorIndex::load_from_path(&path, &op_id, &cfg)
.expect("load_from_path ok")
.expect("Some(index)");
assert_eq!(loaded.len(), idx.len(), "same number of points");
assert_eq!(loaded.dim(), idx.dim(), "same dim");
let query = [1.0_f32, 0.0, 0.0];
let orig_hits = idx.search(&query, 3).unwrap();
let load_hits = loaded.search(&query, 3).unwrap();
assert_eq!(orig_hits.len(), load_hits.len(), "hit count matches");
assert_eq!(orig_hits[0].node_id, ids[0], "top hit is the x-axis vector");
assert_eq!(
load_hits[0].node_id, orig_hits[0].node_id,
"same top hit after round-trip"
);
assert!(
(load_hits[0].score - orig_hits[0].score).abs() < 1e-5,
"scores match after round-trip"
);
let _ = std::fs::remove_file(&path);
}
#[test]
fn ann_cache_stale_op_id_returns_none() {
let (idx, _ids) = small_index();
let op_id_a = make_op_id(b"test-op-A");
let op_id_b = make_op_id(b"test-op-B");
let cfg = HnswConfig::default();
let path = std::env::temp_dir().join("mnem_ann_cache_stale_op_id.bin");
idx.save_to_path(&path, &op_id_a).expect("save");
let result = HnswVectorIndex::load_from_path(&path, &op_id_b, &cfg).expect("no I/O error");
assert!(result.is_none(), "stale op_id should return None");
let _ = std::fs::remove_file(&path);
}
#[test]
fn ann_cache_missing_file_returns_none() {
let op_id = make_op_id(b"test-op-missing");
let cfg = HnswConfig::default();
let path = std::env::temp_dir().join("mnem_ann_cache_does_not_exist_xyz.bin");
let _ = std::fs::remove_file(&path);
let result = HnswVectorIndex::load_from_path(&path, &op_id, &cfg)
.expect("missing file is Ok(None), not Err");
assert!(result.is_none(), "missing file should return None");
}
}