use std::fs::File;
use std::io::{BufReader, BufWriter, Read, Write};
use std::path::Path;
use astraea_core::error::{AstraeaError, Result};
use astraea_core::types::DistanceMetric;
use bincode::Options as _;
use crate::hnsw::HnswIndex;
const MAGIC: u32 = 0x48_4E_53_57;
const FORMAT_VERSION: u32 = 1;
const MAX_HNSW_BYTES: u64 = 4 * 1024 * 1024 * 1024;
const HEADER_SIZE: u64 = 37;
#[derive(Debug, Clone, Copy)]
#[repr(C)]
struct HnswFileHeader {
magic: u32,
version: u32,
dimension: u32,
metric: u8,
m: u32,
m_max0: u32,
ef_construction: u32,
num_vectors: u64,
num_layers: u32,
}
fn metric_to_byte(metric: DistanceMetric) -> u8 {
match metric {
DistanceMetric::Cosine => 0,
DistanceMetric::Euclidean => 1,
DistanceMetric::DotProduct => 2,
}
}
fn byte_to_metric(b: u8) -> Result<DistanceMetric> {
match b {
0 => Ok(DistanceMetric::Cosine),
1 => Ok(DistanceMetric::Euclidean),
2 => Ok(DistanceMetric::DotProduct),
_ => Err(AstraeaError::Deserialization(format!(
"unknown distance metric byte: {b}"
))),
}
}
fn write_header<W: Write>(writer: &mut W, header: &HnswFileHeader) -> Result<()> {
writer.write_all(&header.magic.to_le_bytes())?;
writer.write_all(&header.version.to_le_bytes())?;
writer.write_all(&header.dimension.to_le_bytes())?;
writer.write_all(&[header.metric])?;
writer.write_all(&header.m.to_le_bytes())?;
writer.write_all(&header.m_max0.to_le_bytes())?;
writer.write_all(&header.ef_construction.to_le_bytes())?;
writer.write_all(&header.num_vectors.to_le_bytes())?;
writer.write_all(&header.num_layers.to_le_bytes())?;
Ok(())
}
fn read_header<R: Read>(reader: &mut R) -> Result<HnswFileHeader> {
let mut buf4 = [0u8; 4];
let mut buf8 = [0u8; 8];
let mut buf1 = [0u8; 1];
reader.read_exact(&mut buf4)?;
let magic = u32::from_le_bytes(buf4);
if magic != MAGIC {
return Err(AstraeaError::Deserialization(format!(
"invalid HNSW file magic: expected 0x{MAGIC:08X}, got 0x{magic:08X}"
)));
}
reader.read_exact(&mut buf4)?;
let version = u32::from_le_bytes(buf4);
if version != FORMAT_VERSION {
return Err(AstraeaError::Deserialization(format!(
"unsupported HNSW file version: expected {FORMAT_VERSION}, got {version}"
)));
}
reader.read_exact(&mut buf4)?;
let dimension = u32::from_le_bytes(buf4);
reader.read_exact(&mut buf1)?;
let metric = buf1[0];
reader.read_exact(&mut buf4)?;
let m = u32::from_le_bytes(buf4);
reader.read_exact(&mut buf4)?;
let m_max0 = u32::from_le_bytes(buf4);
reader.read_exact(&mut buf4)?;
let ef_construction = u32::from_le_bytes(buf4);
reader.read_exact(&mut buf8)?;
let num_vectors = u64::from_le_bytes(buf8);
reader.read_exact(&mut buf4)?;
let num_layers = u32::from_le_bytes(buf4);
Ok(HnswFileHeader {
magic,
version,
dimension,
metric,
m,
m_max0,
ef_construction,
num_vectors,
num_layers,
})
}
pub fn save_to_file(index: &HnswIndex, path: &Path) -> Result<()> {
let file = File::create(path)?;
let mut writer = BufWriter::new(file);
let dimension_u32 = u32::try_from(index.dimension()).map_err(|_| {
AstraeaError::Serialization(format!(
"index dimension {} exceeds u32::MAX and cannot be written to the HNSW file header",
index.dimension()
))
})?;
let header = HnswFileHeader {
magic: MAGIC,
version: FORMAT_VERSION,
dimension: dimension_u32,
metric: metric_to_byte(index.metric()),
m: index.m() as u32,
m_max0: index.m_max0() as u32,
ef_construction: index.ef_construction() as u32,
num_vectors: index.len() as u64,
num_layers: index.num_layers() as u32,
};
write_header(&mut writer, &header)?;
bincode::serialize_into(&mut writer, index)
.map_err(|e| AstraeaError::Serialization(format!("bincode serialization failed: {e}")))?;
writer.flush()?;
Ok(())
}
pub fn load_from_file(path: &Path) -> Result<HnswIndex> {
let file = File::open(path)?;
let file_size = file.metadata()?.len();
if file_size > MAX_HNSW_BYTES {
return Err(AstraeaError::Deserialization(format!(
"HNSW file is too large ({file_size} bytes > {MAX_HNSW_BYTES} byte cap): \
refusing to load"
)));
}
let mut reader = BufReader::new(file);
let header = read_header(&mut reader)?;
let _metric = byte_to_metric(header.metric)?;
let body_limit = file_size.saturating_sub(HEADER_SIZE).max(1);
let index: HnswIndex = bincode::DefaultOptions::new()
.with_fixint_encoding()
.allow_trailing_bytes()
.with_limit(body_limit)
.deserialize_from(&mut reader)
.map_err(|e| {
AstraeaError::Deserialization(format!("bincode deserialization failed: {e}"))
})?;
if index.dimension() != header.dimension as usize {
return Err(AstraeaError::Deserialization(format!(
"header/body dimension mismatch: header says {}, body has {}",
header.dimension,
index.dimension()
)));
}
Ok(index)
}
pub fn load_from_file_with_dimension(path: &Path, expected_dimension: usize) -> Result<HnswIndex> {
let index = load_from_file(path)?;
let got = index.dimension();
if got != expected_dimension {
return Err(AstraeaError::DimensionMismatch {
expected: expected_dimension,
got,
});
}
Ok(index)
}
impl HnswIndex {
pub fn save(&self, path: &Path) -> Result<()> {
save_to_file(self, path)
}
pub fn load(path: &Path) -> Result<Self> {
load_from_file(path)
}
pub fn load_expecting_dimension(path: &Path, expected_dimension: usize) -> Result<Self> {
load_from_file_with_dimension(path, expected_dimension)
}
}
#[cfg(test)]
mod tests {
use super::*;
use astraea_core::types::NodeId;
use rand::Rng;
use tempfile::NamedTempFile;
fn build_test_index(dim: usize, n: usize) -> HnswIndex {
let mut idx = HnswIndex::new(dim, DistanceMetric::Euclidean, 16, 200);
let mut rng = rand::thread_rng();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>()).collect();
idx.insert(NodeId(i as u64), &v).unwrap();
}
idx
}
#[test]
fn test_round_trip_100_vectors() {
let dim = 32;
let n = 100;
let original = build_test_index(dim, n);
let tmp = NamedTempFile::new().unwrap();
original.save(tmp.path()).unwrap();
let loaded = HnswIndex::load(tmp.path()).unwrap();
assert_eq!(loaded.dimension(), original.dimension());
assert_eq!(loaded.metric(), original.metric());
assert_eq!(loaded.m(), original.m());
assert_eq!(loaded.m_max0(), original.m_max0());
assert_eq!(loaded.ef_construction(), original.ef_construction());
assert_eq!(loaded.len(), original.len());
let mut rng = rand::thread_rng();
let query: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>()).collect();
let k = 5;
let ef_search = 100;
let orig_results = original.search(&query, k, ef_search).unwrap();
let loaded_results = loaded.search(&query, k, ef_search).unwrap();
assert_eq!(orig_results.len(), loaded_results.len());
assert_eq!(orig_results[0].0, loaded_results[0].0);
assert!((orig_results[0].1 - loaded_results[0].1).abs() < 1e-6);
}
#[test]
fn test_round_trip_empty_index() {
let dim = 8;
let original = HnswIndex::new(dim, DistanceMetric::Cosine, 16, 200);
assert!(original.is_empty());
let tmp = NamedTempFile::new().unwrap();
original.save(tmp.path()).unwrap();
let loaded = HnswIndex::load(tmp.path()).unwrap();
assert_eq!(loaded.dimension(), dim);
assert_eq!(loaded.metric(), DistanceMetric::Cosine);
assert!(loaded.is_empty());
assert_eq!(loaded.len(), 0);
let results = loaded.search(&vec![0.0; dim], 5, 50).unwrap();
assert!(results.is_empty());
}
#[test]
fn test_invalid_magic_bytes() {
let dim = 4;
let original = build_test_index(dim, 5);
let tmp = NamedTempFile::new().unwrap();
original.save(tmp.path()).unwrap();
let mut data = std::fs::read(tmp.path()).unwrap();
data[0] = 0xFF;
data[1] = 0xFF;
data[2] = 0xFF;
data[3] = 0xFF;
std::fs::write(tmp.path(), &data).unwrap();
let result = HnswIndex::load(tmp.path());
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(
err_msg.contains("invalid HNSW file magic"),
"expected magic error, got: {err_msg}"
);
}
#[test]
fn test_invalid_version() {
let dim = 4;
let original = build_test_index(dim, 5);
let tmp = NamedTempFile::new().unwrap();
original.save(tmp.path()).unwrap();
let mut data = std::fs::read(tmp.path()).unwrap();
let bad_version: u32 = 99;
data[4..8].copy_from_slice(&bad_version.to_le_bytes());
std::fs::write(tmp.path(), &data).unwrap();
let result = HnswIndex::load(tmp.path());
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(
err_msg.contains("unsupported HNSW file version"),
"expected version error, got: {err_msg}"
);
}
#[test]
fn test_round_trip_cosine_metric() {
let dim = 16;
let n = 50;
let mut idx = HnswIndex::new(dim, DistanceMetric::Cosine, 8, 100);
let mut rng = rand::thread_rng();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>() + 0.01).collect();
idx.insert(NodeId(i as u64), &v).unwrap();
}
let tmp = NamedTempFile::new().unwrap();
idx.save(tmp.path()).unwrap();
let loaded = HnswIndex::load(tmp.path()).unwrap();
assert_eq!(loaded.metric(), DistanceMetric::Cosine);
assert_eq!(loaded.len(), n);
assert_eq!(loaded.m(), 8);
assert_eq!(loaded.ef_construction(), 100);
}
#[test]
fn test_round_trip_dot_product_metric() {
let dim = 8;
let n = 20;
let mut idx = HnswIndex::new(dim, DistanceMetric::DotProduct, 12, 150);
let mut rng = rand::thread_rng();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>()).collect();
idx.insert(NodeId(i as u64), &v).unwrap();
}
let tmp = NamedTempFile::new().unwrap();
idx.save(tmp.path()).unwrap();
let loaded = HnswIndex::load(tmp.path()).unwrap();
assert_eq!(loaded.metric(), DistanceMetric::DotProduct);
assert_eq!(loaded.len(), n);
}
#[test]
fn test_search_consistency_after_load() {
let dim = 16;
let n = 80;
let original = build_test_index(dim, n);
let tmp = NamedTempFile::new().unwrap();
original.save(tmp.path()).unwrap();
let loaded = HnswIndex::load(tmp.path()).unwrap();
let mut rng = rand::thread_rng();
for _ in 0..10 {
let query: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>()).collect();
let orig_results = original.search(&query, 3, 100).unwrap();
let loaded_results = loaded.search(&query, 3, 100).unwrap();
assert_eq!(orig_results.len(), loaded_results.len());
for (o, l) in orig_results.iter().zip(loaded_results.iter()) {
assert_eq!(o.0, l.0, "node IDs should match");
assert!((o.1 - l.1).abs() < 1e-6, "distances should match");
}
}
}
#[test]
fn test_load_with_dimension_mismatch_returns_error() {
let dim = 128;
let original = build_test_index(dim, 10);
let tmp = NamedTempFile::new().unwrap();
original.save(tmp.path()).unwrap();
let result = HnswIndex::load_expecting_dimension(tmp.path(), 768);
assert!(
result.is_err(),
"expected DimensionMismatch error when loading 128-dim index expecting 768"
);
match result.unwrap_err() {
astraea_core::error::AstraeaError::DimensionMismatch { expected, got } => {
assert_eq!(expected, 768);
assert_eq!(got, 128);
}
other => panic!("expected DimensionMismatch, got: {other:?}"),
}
}
#[test]
fn test_load_with_dimension_matching_succeeds() {
let dim = 128;
let original = build_test_index(dim, 10);
let tmp = NamedTempFile::new().unwrap();
original.save(tmp.path()).unwrap();
let loaded = HnswIndex::load_expecting_dimension(tmp.path(), dim);
assert!(
loaded.is_ok(),
"loading at the matching dimension should succeed"
);
assert_eq!(loaded.unwrap().dimension(), dim);
}
#[test]
fn test_save_dimension_exceeding_u32_max_returns_error() {
let huge_dim: usize = (u32::MAX as usize) + 1;
let idx = HnswIndex::new(huge_dim, DistanceMetric::Euclidean, 16, 200);
let tmp = NamedTempFile::new().unwrap();
let result = idx.save(tmp.path());
assert!(
result.is_err(),
"saving an index with dimension > u32::MAX must fail"
);
match result.unwrap_err() {
astraea_core::error::AstraeaError::Serialization(msg) => {
assert!(
msg.contains("u32::MAX"),
"error message should mention u32::MAX, got: {msg}"
);
}
other => panic!("expected Serialization error, got: {other:?}"),
}
}
#[test]
fn test_corrupt_body_garbage_returns_err_not_abort() {
let dim: u32 = 4;
let mut data: Vec<u8> = Vec::new();
data.extend_from_slice(&MAGIC.to_le_bytes());
data.extend_from_slice(&FORMAT_VERSION.to_le_bytes());
data.extend_from_slice(&dim.to_le_bytes());
data.push(0u8); data.extend_from_slice(&16u32.to_le_bytes()); data.extend_from_slice(&32u32.to_le_bytes()); data.extend_from_slice(&200u32.to_le_bytes()); data.extend_from_slice(&0u64.to_le_bytes()); data.extend_from_slice(&0u32.to_le_bytes());
data.extend(std::iter::repeat_n(0xFFu8, 200));
let tmp = NamedTempFile::new().unwrap();
std::fs::write(tmp.path(), &data).unwrap();
let result = HnswIndex::load(tmp.path());
assert!(
result.is_err(),
"loading a file with a garbage body must return Err, not panic or abort"
);
}
#[test]
fn test_corrupt_body_huge_vector_count_returns_err_not_abort() {
let dim: u32 = 4;
let m: u32 = 16;
let m_max0: u32 = 32;
let ef_construction: u32 = 200;
let ml: f64 = 1.0_f64 / (m as f64).ln();
let mut data: Vec<u8> = Vec::new();
data.extend_from_slice(&MAGIC.to_le_bytes());
data.extend_from_slice(&FORMAT_VERSION.to_le_bytes());
data.extend_from_slice(&dim.to_le_bytes());
data.push(0u8); data.extend_from_slice(&m.to_le_bytes());
data.extend_from_slice(&m_max0.to_le_bytes());
data.extend_from_slice(&ef_construction.to_le_bytes());
data.extend_from_slice(&0u64.to_le_bytes()); data.extend_from_slice(&0u32.to_le_bytes());
data.extend_from_slice(&(dim as u64).to_le_bytes()); data.extend_from_slice(&0u32.to_le_bytes()); data.extend_from_slice(&(m as u64).to_le_bytes()); data.extend_from_slice(&(m_max0 as u64).to_le_bytes()); data.extend_from_slice(&(ef_construction as u64).to_le_bytes()); data.extend_from_slice(&ml.to_le_bytes()); data.extend_from_slice(&u64::MAX.to_le_bytes());
let tmp = NamedTempFile::new().unwrap();
std::fs::write(tmp.path(), &data).unwrap();
let result = HnswIndex::load(tmp.path());
assert!(
result.is_err(),
"loading a file claiming u64::MAX vectors must return Err, not abort"
);
match result.unwrap_err() {
AstraeaError::Deserialization(_) => {} other => panic!("expected Deserialization error, got: {other:?}"),
}
}
#[test]
fn test_round_trip_preserves_non_128_dimension_768() {
const DIM: usize = 768;
let mut idx = HnswIndex::new(DIM, DistanceMetric::Cosine, 16, 200);
let mut rng = rand::thread_rng();
for i in 0..5u64 {
let v: Vec<f32> = (0..DIM).map(|_| rng.r#gen::<f32>()).collect();
idx.insert(NodeId(i), &v).unwrap();
}
assert_eq!(idx.dimension(), DIM);
let tmp = NamedTempFile::new().unwrap();
idx.save(tmp.path()).unwrap();
let loaded = HnswIndex::load(tmp.path()).unwrap();
assert_eq!(
loaded.dimension(),
DIM,
"loaded index dimension must equal the saved 768, not be truncated or defaulted"
);
assert_eq!(loaded.metric(), DistanceMetric::Cosine);
assert_eq!(loaded.len(), 5);
let loaded2 = HnswIndex::load_expecting_dimension(tmp.path(), DIM).unwrap();
assert_eq!(loaded2.dimension(), DIM);
let wrong = HnswIndex::load_expecting_dimension(tmp.path(), 128);
match wrong {
Err(astraea_core::error::AstraeaError::DimensionMismatch { expected, got }) => {
assert_eq!(expected, 128);
assert_eq!(got, DIM);
}
other => panic!("expected DimensionMismatch(128, 768), got: {other:?}"),
}
}
}