1use std::fs::File;
8use std::io::{BufReader, BufWriter, Read, Write};
9use std::path::Path;
10
11use astraea_core::error::{AstraeaError, Result};
12use astraea_core::types::DistanceMetric;
13use bincode::Options as _;
14
15use crate::hnsw::HnswIndex;
16
17const MAGIC: u32 = 0x48_4E_53_57;
19
20const FORMAT_VERSION: u32 = 1;
22
23const MAX_HNSW_BYTES: u64 = 4 * 1024 * 1024 * 1024;
31
32const HEADER_SIZE: u64 = 37;
37
38#[derive(Debug, Clone, Copy)]
43#[repr(C)]
44struct HnswFileHeader {
45 magic: u32,
47 version: u32,
49 dimension: u32,
51 metric: u8,
53 m: u32,
55 m_max0: u32,
57 ef_construction: u32,
59 num_vectors: u64,
61 num_layers: u32,
63}
64
65fn metric_to_byte(metric: DistanceMetric) -> u8 {
67 match metric {
68 DistanceMetric::Cosine => 0,
69 DistanceMetric::Euclidean => 1,
70 DistanceMetric::DotProduct => 2,
71 }
72}
73
74fn byte_to_metric(b: u8) -> Result<DistanceMetric> {
76 match b {
77 0 => Ok(DistanceMetric::Cosine),
78 1 => Ok(DistanceMetric::Euclidean),
79 2 => Ok(DistanceMetric::DotProduct),
80 _ => Err(AstraeaError::Deserialization(format!(
81 "unknown distance metric byte: {b}"
82 ))),
83 }
84}
85
86fn write_header<W: Write>(writer: &mut W, header: &HnswFileHeader) -> Result<()> {
88 writer.write_all(&header.magic.to_le_bytes())?;
89 writer.write_all(&header.version.to_le_bytes())?;
90 writer.write_all(&header.dimension.to_le_bytes())?;
91 writer.write_all(&[header.metric])?;
92 writer.write_all(&header.m.to_le_bytes())?;
93 writer.write_all(&header.m_max0.to_le_bytes())?;
94 writer.write_all(&header.ef_construction.to_le_bytes())?;
95 writer.write_all(&header.num_vectors.to_le_bytes())?;
96 writer.write_all(&header.num_layers.to_le_bytes())?;
97 Ok(())
98}
99
100fn read_header<R: Read>(reader: &mut R) -> Result<HnswFileHeader> {
102 let mut buf4 = [0u8; 4];
103 let mut buf8 = [0u8; 8];
104 let mut buf1 = [0u8; 1];
105
106 reader.read_exact(&mut buf4)?;
108 let magic = u32::from_le_bytes(buf4);
109 if magic != MAGIC {
110 return Err(AstraeaError::Deserialization(format!(
111 "invalid HNSW file magic: expected 0x{MAGIC:08X}, got 0x{magic:08X}"
112 )));
113 }
114
115 reader.read_exact(&mut buf4)?;
117 let version = u32::from_le_bytes(buf4);
118 if version != FORMAT_VERSION {
119 return Err(AstraeaError::Deserialization(format!(
120 "unsupported HNSW file version: expected {FORMAT_VERSION}, got {version}"
121 )));
122 }
123
124 reader.read_exact(&mut buf4)?;
126 let dimension = u32::from_le_bytes(buf4);
127
128 reader.read_exact(&mut buf1)?;
130 let metric = buf1[0];
131
132 reader.read_exact(&mut buf4)?;
134 let m = u32::from_le_bytes(buf4);
135
136 reader.read_exact(&mut buf4)?;
138 let m_max0 = u32::from_le_bytes(buf4);
139
140 reader.read_exact(&mut buf4)?;
142 let ef_construction = u32::from_le_bytes(buf4);
143
144 reader.read_exact(&mut buf8)?;
146 let num_vectors = u64::from_le_bytes(buf8);
147
148 reader.read_exact(&mut buf4)?;
150 let num_layers = u32::from_le_bytes(buf4);
151
152 Ok(HnswFileHeader {
153 magic,
154 version,
155 dimension,
156 metric,
157 m,
158 m_max0,
159 ef_construction,
160 num_vectors,
161 num_layers,
162 })
163}
164
165pub fn save_to_file(index: &HnswIndex, path: &Path) -> Result<()> {
174 let file = File::create(path)?;
175 let mut writer = BufWriter::new(file);
176
177 let dimension_u32 = u32::try_from(index.dimension()).map_err(|_| {
178 AstraeaError::Serialization(format!(
179 "index dimension {} exceeds u32::MAX and cannot be written to the HNSW file header",
180 index.dimension()
181 ))
182 })?;
183
184 let header = HnswFileHeader {
185 magic: MAGIC,
186 version: FORMAT_VERSION,
187 dimension: dimension_u32,
188 metric: metric_to_byte(index.metric()),
189 m: index.m() as u32,
190 m_max0: index.m_max0() as u32,
191 ef_construction: index.ef_construction() as u32,
192 num_vectors: index.len() as u64,
193 num_layers: index.num_layers() as u32,
194 };
195
196 write_header(&mut writer, &header)?;
197
198 bincode::serialize_into(&mut writer, index)
200 .map_err(|e| AstraeaError::Serialization(format!("bincode serialization failed: {e}")))?;
201
202 writer.flush()?;
203 Ok(())
204}
205
206pub fn load_from_file(path: &Path) -> Result<HnswIndex> {
226 let file = File::open(path)?;
227
228 let file_size = file.metadata()?.len();
232 if file_size > MAX_HNSW_BYTES {
233 return Err(AstraeaError::Deserialization(format!(
234 "HNSW file is too large ({file_size} bytes > {MAX_HNSW_BYTES} byte cap): \
235 refusing to load"
236 )));
237 }
238
239 let mut reader = BufReader::new(file);
240
241 let header = read_header(&mut reader)?;
245
246 let _metric = byte_to_metric(header.metric)?;
248
249 let body_limit = file_size.saturating_sub(HEADER_SIZE).max(1);
264 let index: HnswIndex = bincode::DefaultOptions::new()
265 .with_fixint_encoding()
266 .allow_trailing_bytes()
267 .with_limit(body_limit)
268 .deserialize_from(&mut reader)
269 .map_err(|e| {
270 AstraeaError::Deserialization(format!("bincode deserialization failed: {e}"))
271 })?;
272
273 if index.dimension() != header.dimension as usize {
275 return Err(AstraeaError::Deserialization(format!(
276 "header/body dimension mismatch: header says {}, body has {}",
277 header.dimension,
278 index.dimension()
279 )));
280 }
281
282 Ok(index)
283}
284
285pub fn load_from_file_with_dimension(path: &Path, expected_dimension: usize) -> Result<HnswIndex> {
296 let index = load_from_file(path)?;
297 let got = index.dimension();
298 if got != expected_dimension {
299 return Err(AstraeaError::DimensionMismatch {
300 expected: expected_dimension,
301 got,
302 });
303 }
304 Ok(index)
305}
306
307impl HnswIndex {
310 pub fn save(&self, path: &Path) -> Result<()> {
312 save_to_file(self, path)
313 }
314
315 pub fn load(path: &Path) -> Result<Self> {
317 load_from_file(path)
318 }
319
320 pub fn load_expecting_dimension(path: &Path, expected_dimension: usize) -> Result<Self> {
325 load_from_file_with_dimension(path, expected_dimension)
326 }
327}
328
329#[cfg(test)]
330mod tests {
331 use super::*;
332 use astraea_core::types::NodeId;
333 use rand::Rng;
334 use tempfile::NamedTempFile;
335
336 fn build_test_index(dim: usize, n: usize) -> HnswIndex {
338 let mut idx = HnswIndex::new(dim, DistanceMetric::Euclidean, 16, 200);
339 let mut rng = rand::thread_rng();
340 for i in 0..n {
341 let v: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>()).collect();
342 idx.insert(NodeId(i as u64), &v).unwrap();
343 }
344 idx
345 }
346
347 #[test]
348 fn test_round_trip_100_vectors() {
349 let dim = 32;
350 let n = 100;
351 let original = build_test_index(dim, n);
352
353 let tmp = NamedTempFile::new().unwrap();
355 original.save(tmp.path()).unwrap();
356
357 let loaded = HnswIndex::load(tmp.path()).unwrap();
359
360 assert_eq!(loaded.dimension(), original.dimension());
362 assert_eq!(loaded.metric(), original.metric());
363 assert_eq!(loaded.m(), original.m());
364 assert_eq!(loaded.m_max0(), original.m_max0());
365 assert_eq!(loaded.ef_construction(), original.ef_construction());
366 assert_eq!(loaded.len(), original.len());
367
368 let mut rng = rand::thread_rng();
370 let query: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>()).collect();
371 let k = 5;
372 let ef_search = 100;
373
374 let orig_results = original.search(&query, k, ef_search).unwrap();
375 let loaded_results = loaded.search(&query, k, ef_search).unwrap();
376
377 assert_eq!(orig_results.len(), loaded_results.len());
378 assert_eq!(orig_results[0].0, loaded_results[0].0);
380 assert!((orig_results[0].1 - loaded_results[0].1).abs() < 1e-6);
381 }
382
383 #[test]
384 fn test_round_trip_empty_index() {
385 let dim = 8;
386 let original = HnswIndex::new(dim, DistanceMetric::Cosine, 16, 200);
387 assert!(original.is_empty());
388
389 let tmp = NamedTempFile::new().unwrap();
390 original.save(tmp.path()).unwrap();
391
392 let loaded = HnswIndex::load(tmp.path()).unwrap();
393
394 assert_eq!(loaded.dimension(), dim);
395 assert_eq!(loaded.metric(), DistanceMetric::Cosine);
396 assert!(loaded.is_empty());
397 assert_eq!(loaded.len(), 0);
398
399 let results = loaded.search(&vec![0.0; dim], 5, 50).unwrap();
401 assert!(results.is_empty());
402 }
403
404 #[test]
405 fn test_invalid_magic_bytes() {
406 let dim = 4;
407 let original = build_test_index(dim, 5);
408
409 let tmp = NamedTempFile::new().unwrap();
410 original.save(tmp.path()).unwrap();
411
412 let mut data = std::fs::read(tmp.path()).unwrap();
414 data[0] = 0xFF;
415 data[1] = 0xFF;
416 data[2] = 0xFF;
417 data[3] = 0xFF;
418 std::fs::write(tmp.path(), &data).unwrap();
419
420 let result = HnswIndex::load(tmp.path());
421 assert!(result.is_err());
422 let err_msg = format!("{}", result.unwrap_err());
423 assert!(
424 err_msg.contains("invalid HNSW file magic"),
425 "expected magic error, got: {err_msg}"
426 );
427 }
428
429 #[test]
430 fn test_invalid_version() {
431 let dim = 4;
432 let original = build_test_index(dim, 5);
433
434 let tmp = NamedTempFile::new().unwrap();
435 original.save(tmp.path()).unwrap();
436
437 let mut data = std::fs::read(tmp.path()).unwrap();
439 let bad_version: u32 = 99;
440 data[4..8].copy_from_slice(&bad_version.to_le_bytes());
441 std::fs::write(tmp.path(), &data).unwrap();
442
443 let result = HnswIndex::load(tmp.path());
444 assert!(result.is_err());
445 let err_msg = format!("{}", result.unwrap_err());
446 assert!(
447 err_msg.contains("unsupported HNSW file version"),
448 "expected version error, got: {err_msg}"
449 );
450 }
451
452 #[test]
453 fn test_round_trip_cosine_metric() {
454 let dim = 16;
455 let n = 50;
456 let mut idx = HnswIndex::new(dim, DistanceMetric::Cosine, 8, 100);
457 let mut rng = rand::thread_rng();
458 for i in 0..n {
459 let v: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>() + 0.01).collect();
460 idx.insert(NodeId(i as u64), &v).unwrap();
461 }
462
463 let tmp = NamedTempFile::new().unwrap();
464 idx.save(tmp.path()).unwrap();
465
466 let loaded = HnswIndex::load(tmp.path()).unwrap();
467 assert_eq!(loaded.metric(), DistanceMetric::Cosine);
468 assert_eq!(loaded.len(), n);
469 assert_eq!(loaded.m(), 8);
470 assert_eq!(loaded.ef_construction(), 100);
471 }
472
473 #[test]
474 fn test_round_trip_dot_product_metric() {
475 let dim = 8;
476 let n = 20;
477 let mut idx = HnswIndex::new(dim, DistanceMetric::DotProduct, 12, 150);
478 let mut rng = rand::thread_rng();
479 for i in 0..n {
480 let v: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>()).collect();
481 idx.insert(NodeId(i as u64), &v).unwrap();
482 }
483
484 let tmp = NamedTempFile::new().unwrap();
485 idx.save(tmp.path()).unwrap();
486
487 let loaded = HnswIndex::load(tmp.path()).unwrap();
488 assert_eq!(loaded.metric(), DistanceMetric::DotProduct);
489 assert_eq!(loaded.len(), n);
490 }
491
492 #[test]
493 fn test_search_consistency_after_load() {
494 let dim = 16;
497 let n = 80;
498 let original = build_test_index(dim, n);
499
500 let tmp = NamedTempFile::new().unwrap();
501 original.save(tmp.path()).unwrap();
502 let loaded = HnswIndex::load(tmp.path()).unwrap();
503
504 let mut rng = rand::thread_rng();
505 for _ in 0..10 {
506 let query: Vec<f32> = (0..dim).map(|_| rng.r#gen::<f32>()).collect();
507 let orig_results = original.search(&query, 3, 100).unwrap();
508 let loaded_results = loaded.search(&query, 3, 100).unwrap();
509
510 assert_eq!(orig_results.len(), loaded_results.len());
511 for (o, l) in orig_results.iter().zip(loaded_results.iter()) {
512 assert_eq!(o.0, l.0, "node IDs should match");
513 assert!((o.1 - l.1).abs() < 1e-6, "distances should match");
514 }
515 }
516 }
517
518 #[test]
522 fn test_load_with_dimension_mismatch_returns_error() {
523 let dim = 128;
524 let original = build_test_index(dim, 10);
525 let tmp = NamedTempFile::new().unwrap();
526 original.save(tmp.path()).unwrap();
527
528 let result = HnswIndex::load_expecting_dimension(tmp.path(), 768);
529 assert!(
530 result.is_err(),
531 "expected DimensionMismatch error when loading 128-dim index expecting 768"
532 );
533 match result.unwrap_err() {
534 astraea_core::error::AstraeaError::DimensionMismatch { expected, got } => {
535 assert_eq!(expected, 768);
536 assert_eq!(got, 128);
537 }
538 other => panic!("expected DimensionMismatch, got: {other:?}"),
539 }
540 }
541
542 #[test]
544 fn test_load_with_dimension_matching_succeeds() {
545 let dim = 128;
546 let original = build_test_index(dim, 10);
547 let tmp = NamedTempFile::new().unwrap();
548 original.save(tmp.path()).unwrap();
549
550 let loaded = HnswIndex::load_expecting_dimension(tmp.path(), dim);
551 assert!(
552 loaded.is_ok(),
553 "loading at the matching dimension should succeed"
554 );
555 assert_eq!(loaded.unwrap().dimension(), dim);
556 }
557
558 #[test]
561 fn test_save_dimension_exceeding_u32_max_returns_error() {
562 let huge_dim: usize = (u32::MAX as usize) + 1;
565 let idx = HnswIndex::new(huge_dim, DistanceMetric::Euclidean, 16, 200);
566
567 let tmp = NamedTempFile::new().unwrap();
568 let result = idx.save(tmp.path());
569 assert!(
570 result.is_err(),
571 "saving an index with dimension > u32::MAX must fail"
572 );
573 match result.unwrap_err() {
574 astraea_core::error::AstraeaError::Serialization(msg) => {
575 assert!(
576 msg.contains("u32::MAX"),
577 "error message should mention u32::MAX, got: {msg}"
578 );
579 }
580 other => panic!("expected Serialization error, got: {other:?}"),
581 }
582 }
583
584 #[test]
589 fn test_corrupt_body_garbage_returns_err_not_abort() {
590 let dim: u32 = 4;
592 let mut data: Vec<u8> = Vec::new();
593
594 data.extend_from_slice(&MAGIC.to_le_bytes());
596 data.extend_from_slice(&FORMAT_VERSION.to_le_bytes());
597 data.extend_from_slice(&dim.to_le_bytes());
598 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));
607
608 let tmp = NamedTempFile::new().unwrap();
609 std::fs::write(tmp.path(), &data).unwrap();
610
611 let result = HnswIndex::load(tmp.path());
612 assert!(
613 result.is_err(),
614 "loading a file with a garbage body must return Err, not panic or abort"
615 );
616 }
617
618 #[test]
637 fn test_corrupt_body_huge_vector_count_returns_err_not_abort() {
638 let dim: u32 = 4;
639 let m: u32 = 16;
640 let m_max0: u32 = 32;
641 let ef_construction: u32 = 200;
642 let ml: f64 = 1.0_f64 / (m as f64).ln();
643
644 let mut data: Vec<u8> = Vec::new();
645
646 data.extend_from_slice(&MAGIC.to_le_bytes());
648 data.extend_from_slice(&FORMAT_VERSION.to_le_bytes());
649 data.extend_from_slice(&dim.to_le_bytes());
650 data.push(0u8); data.extend_from_slice(&m.to_le_bytes());
652 data.extend_from_slice(&m_max0.to_le_bytes());
653 data.extend_from_slice(&ef_construction.to_le_bytes());
654 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());
666 let tmp = NamedTempFile::new().unwrap();
669 std::fs::write(tmp.path(), &data).unwrap();
670
671 let result = HnswIndex::load(tmp.path());
673 assert!(
674 result.is_err(),
675 "loading a file claiming u64::MAX vectors must return Err, not abort"
676 );
677 match result.unwrap_err() {
679 AstraeaError::Deserialization(_) => {} other => panic!("expected Deserialization error, got: {other:?}"),
681 }
682 }
683
684 #[test]
690 fn test_round_trip_preserves_non_128_dimension_768() {
691 const DIM: usize = 768;
692 let mut idx = HnswIndex::new(DIM, DistanceMetric::Cosine, 16, 200);
693 let mut rng = rand::thread_rng();
694
695 for i in 0..5u64 {
697 let v: Vec<f32> = (0..DIM).map(|_| rng.r#gen::<f32>()).collect();
698 idx.insert(NodeId(i), &v).unwrap();
699 }
700 assert_eq!(idx.dimension(), DIM);
701
702 let tmp = NamedTempFile::new().unwrap();
703 idx.save(tmp.path()).unwrap();
704
705 let loaded = HnswIndex::load(tmp.path()).unwrap();
706
707 assert_eq!(
708 loaded.dimension(),
709 DIM,
710 "loaded index dimension must equal the saved 768, not be truncated or defaulted"
711 );
712 assert_eq!(loaded.metric(), DistanceMetric::Cosine);
713 assert_eq!(loaded.len(), 5);
714
715 let loaded2 = HnswIndex::load_expecting_dimension(tmp.path(), DIM).unwrap();
717 assert_eq!(loaded2.dimension(), DIM);
718
719 let wrong = HnswIndex::load_expecting_dimension(tmp.path(), 128);
721 match wrong {
722 Err(astraea_core::error::AstraeaError::DimensionMismatch { expected, got }) => {
723 assert_eq!(expected, 128);
724 assert_eq!(got, DIM);
725 }
726 other => panic!("expected DimensionMismatch(128, 768), got: {other:?}"),
727 }
728 }
729}