1use 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: Vec<u64>,
16 flat_vecs: Vec<f32>,
18 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, 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 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 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 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 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 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 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 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], 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}