Skip to main content

ailake_index/
serialize.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2use ailake_core::{AilakeError, AilakeResult, VectorMetric};
3use serde::{Deserialize, Serialize};
4
5use crate::hnsw::{HnswConfig, HnswIndex};
6
7#[derive(Serialize, Deserialize)]
8struct HnswSnapshot {
9    m: usize,
10    ef_construction: usize,
11    max_elements: usize,
12    metric: u8,
13    dim: u32,
14    /// Row IDs parallel to flat_vecs (one entry per vector).
15    row_ids: Vec<u64>,
16    /// Contiguous vector storage: flat_vecs[i*dim..(i+1)*dim] = vector i.
17    flat_vecs: Vec<f32>,
18    // Graph structure (empty = old format, triggers brute-force fallback)
19    neighbors: Vec<Vec<Vec<usize>>>,
20    node_levels: Vec<usize>,
21    entry_point: Option<usize>,
22    max_layer: usize,
23}
24
25fn metric_to_u8(m: VectorMetric) -> u8 {
26    match m {
27        VectorMetric::Cosine => 0,
28        VectorMetric::Euclidean => 1,
29        VectorMetric::DotProduct => 2,
30        VectorMetric::NormalizedCosine => 3,
31    }
32}
33
34fn u8_to_metric(v: u8) -> AilakeResult<VectorMetric> {
35    match v {
36        0 => Ok(VectorMetric::Cosine),
37        1 => Ok(VectorMetric::Euclidean),
38        2 => Ok(VectorMetric::DotProduct),
39        3 => Ok(VectorMetric::NormalizedCosine),
40        _ => Err(AilakeError::Catalog(format!(
41            "HNSW index deserialization: unknown metric byte {v} (valid: 0=Cosine, 1=Euclidean, 2=DotProduct, 3=NormalizedCosine)"
42        ))),
43    }
44}
45
46pub struct HnswSerializer;
47
48impl HnswSerializer {
49    pub fn to_bytes(index: &HnswIndex) -> AilakeResult<Vec<u8>> {
50        let snap = HnswSnapshot {
51            m: index.config.m,
52            ef_construction: index.config.ef_construction,
53            max_elements: index.config.max_elements,
54            metric: metric_to_u8(index.metric),
55            dim: index.dim,
56            row_ids: index.row_ids.clone(),
57            flat_vecs: index.flat_vecs.clone(),
58            neighbors: index.neighbors.clone(),
59            node_levels: index.node_levels.clone(),
60            entry_point: index.entry_point,
61            max_layer: index.max_layer,
62        };
63        bincode::serialize(&snap).map_err(|e| AilakeError::Bincode(e.to_string()))
64    }
65
66    pub fn from_bytes(bytes: &[u8]) -> AilakeResult<HnswIndex> {
67        let snap: HnswSnapshot =
68            bincode::deserialize(bytes).map_err(|e| AilakeError::Bincode(e.to_string()))?;
69        let metric = u8_to_metric(snap.metric)?;
70        Ok(HnswIndex {
71            config: HnswConfig {
72                m: snap.m,
73                ef_construction: snap.ef_construction,
74                max_elements: snap.max_elements,
75            },
76            metric,
77            dim: snap.dim,
78            row_ids: snap.row_ids,
79            flat_vecs: snap.flat_vecs,
80            flat_vecs_f16: None, // populated at runtime by quantize_to_f16() if needed
81            neighbors: snap.neighbors,
82            node_levels: snap.node_levels,
83            entry_point: snap.entry_point,
84            max_layer: snap.max_layer,
85        })
86    }
87}
88
89#[cfg(test)]
90mod tests {
91    use super::*;
92    use crate::hnsw::HnswBuilder;
93    use ailake_core::RowId;
94
95    #[test]
96    fn serialize_roundtrip() {
97        let mut b = HnswBuilder::new(3, VectorMetric::Cosine, Default::default());
98        b.insert(RowId::new(0), vec![1.0, 0.0, 0.0]);
99        b.insert(RowId::new(1), vec![0.0, 1.0, 0.0]);
100        let idx = b.build();
101        let bytes = HnswSerializer::to_bytes(&idx).unwrap();
102        let idx2 = HnswSerializer::from_bytes(&bytes).unwrap();
103        assert_eq!(idx2.node_count(), 2);
104        assert_eq!(idx2.dim(), 3);
105        let r = idx2.search(&[1.0, 0.0, 0.0], 1, 50);
106        assert_eq!(r[0].0, RowId::new(0));
107    }
108
109    #[test]
110    fn serialize_preserves_graph() {
111        use rand::{rngs::StdRng, Rng, SeedableRng};
112        let mut rng = StdRng::seed_from_u64(7);
113        let mut b = HnswBuilder::new(8, VectorMetric::Euclidean, Default::default());
114        for i in 0..50u64 {
115            let v: Vec<f32> = (0..8).map(|_| rng.gen::<f32>()).collect();
116            b.insert(RowId::new(i), v);
117        }
118        let idx = b.build();
119        let query: Vec<f32> = (0..8).map(|_| rng.gen::<f32>()).collect();
120        let r1 = idx.search(&query, 5, 50);
121
122        let bytes = HnswSerializer::to_bytes(&idx).unwrap();
123        let idx2 = HnswSerializer::from_bytes(&bytes).unwrap();
124        let r2 = idx2.search(&query, 5, 50);
125
126        assert_eq!(r1.len(), r2.len());
127        for (a, b) in r1.iter().zip(r2.iter()) {
128            assert_eq!(a.0, b.0);
129        }
130    }
131}