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        Self::validate_snapshot(&snap)?;
71        Ok(HnswIndex {
72            config: HnswConfig {
73                m: snap.m,
74                ef_construction: snap.ef_construction,
75                max_elements: snap.max_elements,
76            },
77            metric,
78            dim: snap.dim,
79            row_ids: snap.row_ids,
80            flat_vecs: snap.flat_vecs,
81            flat_vecs_f16: None, // populated at runtime by quantize_to_f16() if needed
82            neighbors: snap.neighbors,
83            node_levels: snap.node_levels,
84            entry_point: snap.entry_point,
85            max_layer: snap.max_layer,
86        })
87    }
88
89    /// Checks the invariants `HnswIndex`'s unsafe search path (`VisitedTracker::visit`'s
90    /// `get_unchecked_mut`) relies on, since `snap` comes from untrusted bytes (disk/S3/IPC)
91    /// and bincode deserialization alone doesn't guarantee index-range consistency.
92    fn validate_snapshot(snap: &HnswSnapshot) -> AilakeResult<()> {
93        let n = snap.row_ids.len();
94        if let Some(ep) = snap.entry_point {
95            if ep >= n {
96                return Err(AilakeError::Bincode(format!(
97                    "corrupt HNSW graph: entry_point {ep} out of bounds (n={n})"
98                )));
99            }
100        }
101        // Empty neighbors = old format, triggers brute-force fallback (see HnswSnapshot doc).
102        if !snap.neighbors.is_empty() {
103            if snap.neighbors.len() != n {
104                return Err(AilakeError::Bincode(format!(
105                    "corrupt HNSW graph: neighbors.len()={} != row_ids.len()={n}",
106                    snap.neighbors.len()
107                )));
108            }
109            if snap.node_levels.len() != n {
110                return Err(AilakeError::Bincode(format!(
111                    "corrupt HNSW graph: node_levels.len()={} != row_ids.len()={n}",
112                    snap.node_levels.len()
113                )));
114            }
115            // HnswBuilder always pushes `vec![Vec::new(); l + 1]` for a node at level `l`
116            // (hnsw.rs), so this must hold exactly — a mismatch means a layer index derived
117            // from node_levels would index out of bounds into neighbors[i] during search.
118            for (i, per_node) in snap.neighbors.iter().enumerate() {
119                if per_node.len() != snap.node_levels[i] + 1 {
120                    return Err(AilakeError::Bincode(format!(
121                        "corrupt HNSW graph: node {i} has node_levels={} but neighbors[{i}].len()={}",
122                        snap.node_levels[i],
123                        per_node.len()
124                    )));
125                }
126                for per_layer in per_node {
127                    for &nb in per_layer {
128                        if nb >= n {
129                            return Err(AilakeError::Bincode(format!(
130                                "corrupt HNSW graph: neighbor index {nb} out of bounds (n={n})"
131                            )));
132                        }
133                    }
134                }
135            }
136            // max_layer must equal the highest level any node was actually built at
137            // (HnswBuilder only ever raises it to `l` when inserting a node at level `l`,
138            // hnsw.rs). An inflated max_layer drives `for lc in (1..=self.max_layer).rev()`
139            // in the search hot path into an effectively unbounded loop.
140            let max_node_level = snap.node_levels.iter().copied().max().unwrap_or(0);
141            if snap.max_layer != max_node_level {
142                return Err(AilakeError::Bincode(format!(
143                    "corrupt HNSW graph: max_layer={} != max(node_levels)={max_node_level}",
144                    snap.max_layer
145                )));
146            }
147        } else if snap.max_layer != 0 {
148            return Err(AilakeError::Bincode(format!(
149                "corrupt HNSW graph: max_layer={} but neighbors is empty (old format)",
150                snap.max_layer
151            )));
152        }
153        let expected_flat_len = n * snap.dim as usize;
154        if snap.flat_vecs.len() != expected_flat_len {
155            return Err(AilakeError::Bincode(format!(
156                "corrupt HNSW graph: flat_vecs.len()={} != row_ids.len()*dim={expected_flat_len}",
157                snap.flat_vecs.len()
158            )));
159        }
160        Ok(())
161    }
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167    use crate::hnsw::HnswBuilder;
168    use ailake_core::RowId;
169
170    #[test]
171    fn serialize_roundtrip() {
172        let mut b = HnswBuilder::new(3, VectorMetric::Cosine, Default::default());
173        b.insert(RowId::new(0), vec![1.0, 0.0, 0.0]);
174        b.insert(RowId::new(1), vec![0.0, 1.0, 0.0]);
175        let idx = b.build();
176        let bytes = HnswSerializer::to_bytes(&idx).unwrap();
177        let idx2 = HnswSerializer::from_bytes(&bytes).unwrap();
178        assert_eq!(idx2.node_count(), 2);
179        assert_eq!(idx2.dim(), 3);
180        let r = idx2.search(&[1.0, 0.0, 0.0], 1, 50);
181        assert_eq!(r[0].0, RowId::new(0));
182    }
183
184    #[test]
185    fn serialize_preserves_graph() {
186        use rand::{rngs::StdRng, Rng, SeedableRng};
187        let mut rng = StdRng::seed_from_u64(7);
188        let mut b = HnswBuilder::new(8, VectorMetric::Euclidean, Default::default());
189        for i in 0..50u64 {
190            let v: Vec<f32> = (0..8).map(|_| rng.gen::<f32>()).collect();
191            b.insert(RowId::new(i), v);
192        }
193        let idx = b.build();
194        let query: Vec<f32> = (0..8).map(|_| rng.gen::<f32>()).collect();
195        let r1 = idx.search(&query, 5, 50);
196
197        let bytes = HnswSerializer::to_bytes(&idx).unwrap();
198        let idx2 = HnswSerializer::from_bytes(&bytes).unwrap();
199        let r2 = idx2.search(&query, 5, 50);
200
201        assert_eq!(r1.len(), r2.len());
202        for (a, b) in r1.iter().zip(r2.iter()) {
203            assert_eq!(a.0, b.0);
204        }
205    }
206
207    #[test]
208    fn from_bytes_rejects_out_of_bounds_neighbor_index() {
209        let snap = HnswSnapshot {
210            m: 16,
211            ef_construction: 150,
212            max_elements: 100,
213            metric: metric_to_u8(VectorMetric::Cosine),
214            dim: 3,
215            row_ids: vec![0, 1],
216            flat_vecs: vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0],
217            // node 0's neighbor list at layer 0 points at node index 99, which
218            // doesn't exist — corrupt/malicious bytes should be rejected, not
219            // silently accepted and later fed to an unchecked array access.
220            neighbors: vec![vec![vec![99]], vec![vec![]]],
221            node_levels: vec![0, 0],
222            entry_point: Some(0),
223            max_layer: 0,
224        };
225        let bytes = bincode::serialize(&snap).unwrap();
226        let err = HnswSerializer::from_bytes(&bytes).err().unwrap();
227        assert!(err.to_string().contains("out of bounds"), "{err}");
228    }
229
230    #[test]
231    fn from_bytes_rejects_inflated_max_layer() {
232        // Structurally consistent otherwise (bounds/lengths all check out), but max_layer
233        // claims a level far above what node_levels actually reaches. Uncaught, this drives
234        // `for lc in (1..=self.max_layer).rev()` in the search hot path into an effectively
235        // unbounded loop — a DoS via a single crafted/corrupted graph.
236        let snap = HnswSnapshot {
237            m: 16,
238            ef_construction: 150,
239            max_elements: 100,
240            metric: metric_to_u8(VectorMetric::Cosine),
241            dim: 3,
242            row_ids: vec![0, 1],
243            flat_vecs: vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0],
244            neighbors: vec![vec![vec![]], vec![vec![]]],
245            node_levels: vec![0, 0],
246            entry_point: Some(0),
247            max_layer: 1_000_000,
248        };
249        let bytes = bincode::serialize(&snap).unwrap();
250        let err = HnswSerializer::from_bytes(&bytes).err().unwrap();
251        assert!(err.to_string().contains("max_layer"), "{err}");
252    }
253
254    #[test]
255    fn from_bytes_rejects_node_levels_neighbors_mismatch() {
256        // node 0 claims level 5 (so a query could traverse layers 1..=5 for it) but only has
257        // a single per-layer neighbor list (layer 0) — HnswBuilder never produces this
258        // shape (neighbors[i].len() == node_levels[i] + 1 always), so this is corrupt/
259        // malicious input. Uncaught, `neighbors[c.idx][layer]` in search_layer indexes out
260        // of bounds and panics the first time a query traverses through this node above
261        // layer 0.
262        let snap = HnswSnapshot {
263            m: 16,
264            ef_construction: 150,
265            max_elements: 100,
266            metric: metric_to_u8(VectorMetric::Cosine),
267            dim: 3,
268            row_ids: vec![0, 1],
269            flat_vecs: vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0],
270            neighbors: vec![vec![vec![]], vec![vec![]]],
271            node_levels: vec![5, 0],
272            entry_point: Some(0),
273            max_layer: 5,
274        };
275        let bytes = bincode::serialize(&snap).unwrap();
276        let err = HnswSerializer::from_bytes(&bytes).err().unwrap();
277        assert!(err.to_string().contains("node_levels"), "{err}");
278    }
279
280    #[test]
281    fn from_bytes_rejects_out_of_bounds_entry_point() {
282        let snap = HnswSnapshot {
283            m: 16,
284            ef_construction: 150,
285            max_elements: 100,
286            metric: metric_to_u8(VectorMetric::Cosine),
287            dim: 3,
288            row_ids: vec![0],
289            flat_vecs: vec![1.0, 0.0, 0.0],
290            neighbors: vec![],
291            node_levels: vec![],
292            entry_point: Some(7),
293            max_layer: 0,
294        };
295        let bytes = bincode::serialize(&snap).unwrap();
296        let err = HnswSerializer::from_bytes(&bytes).err().unwrap();
297        assert!(err.to_string().contains("entry_point"), "{err}");
298    }
299
300    #[test]
301    fn from_bytes_rejects_flat_vecs_length_mismatch() {
302        let snap = HnswSnapshot {
303            m: 16,
304            ef_construction: 150,
305            max_elements: 100,
306            metric: metric_to_u8(VectorMetric::Cosine),
307            dim: 3,
308            row_ids: vec![0, 1],
309            flat_vecs: vec![1.0, 0.0, 0.0], // only 1 vector's worth for 2 row_ids
310            neighbors: vec![],
311            node_levels: vec![],
312            entry_point: None,
313            max_layer: 0,
314        };
315        let bytes = bincode::serialize(&snap).unwrap();
316        let err = HnswSerializer::from_bytes(&bytes).err().unwrap();
317        assert!(err.to_string().contains("flat_vecs"), "{err}");
318    }
319}