1use std::fmt;
29
30const MAGIC: &[u8; 4] = b"VRB1";
32
33const HEADER_LEN: usize = 16;
35
36const ID_WIDTH: u8 = 8;
38
39#[derive(Debug, Clone, PartialEq)]
42pub struct RawBulk {
43 pub ids: Vec<u64>,
45 pub vectors: Vec<f32>,
47 pub dimension: usize,
49}
50
51#[derive(Debug, Clone, PartialEq, Eq)]
53pub enum VrbError {
54 TooShort {
56 got: usize,
58 },
59 BadMagic,
61 BadIdWidth(u8),
63 ReservedNotZero,
65 Overflow,
67 LengthMismatch {
69 got: usize,
71 expected: usize,
73 },
74}
75
76impl fmt::Display for VrbError {
77 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
78 match self {
79 Self::TooShort { got } => {
80 write!(f, "body too short: {got} bytes (header needs {HEADER_LEN})")
81 }
82 Self::BadMagic => write!(f, "bad magic: expected b\"VRB1\""),
83 Self::BadIdWidth(w) => {
84 write!(
85 f,
86 "unsupported id_width {w}: only {ID_WIDTH} (u64) is supported"
87 )
88 }
89 Self::ReservedNotZero => write!(f, "reserved header bytes must be zero"),
90 Self::Overflow => write!(f, "overflow computing body length"),
91 Self::LengthMismatch { got, expected } => {
92 write!(f, "body length {got} != expected {expected}")
93 }
94 }
95 }
96}
97
98impl std::error::Error for VrbError {}
99
100fn parse_header(body: &[u8]) -> Result<(usize, usize), VrbError> {
105 if body.len() < HEADER_LEN {
106 return Err(VrbError::TooShort { got: body.len() });
107 }
108 if &body[0..4] != MAGIC {
109 return Err(VrbError::BadMagic);
110 }
111 let count = u32::from_le_bytes([body[4], body[5], body[6], body[7]]) as usize;
112 let dim = u32::from_le_bytes([body[8], body[9], body[10], body[11]]) as usize;
113 if body[12] != ID_WIDTH {
114 return Err(VrbError::BadIdWidth(body[12]));
115 }
116 if body[13] != 0 || body[14] != 0 || body[15] != 0 {
117 return Err(VrbError::ReservedNotZero);
118 }
119 Ok((count, dim))
120}
121
122fn expected_body_len(count: usize, dim: usize) -> Result<usize, VrbError> {
126 let ids_bytes = count.checked_mul(8).ok_or(VrbError::Overflow)?;
127 let vec_elems = count.checked_mul(dim).ok_or(VrbError::Overflow)?;
128 let vec_bytes = vec_elems.checked_mul(4).ok_or(VrbError::Overflow)?;
129 HEADER_LEN
130 .checked_add(ids_bytes)
131 .and_then(|h| h.checked_add(vec_bytes))
132 .ok_or(VrbError::Overflow)
133}
134
135fn decode_ids(body: &[u8], count: usize) -> Vec<u64> {
140 let start = HEADER_LEN;
141 let end = start + count * 8;
142 body[start..end]
143 .chunks_exact(8)
144 .map(|c| u64::from_le_bytes([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]))
145 .collect()
146}
147
148fn decode_vectors(body: &[u8], count: usize, dim: usize) -> Vec<f32> {
150 let start = HEADER_LEN + count * 8;
151 let end = start + count * dim * 4;
152 body[start..end]
153 .chunks_exact(4)
154 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
155 .collect()
156}
157
158pub fn decode(body: &[u8]) -> Result<RawBulk, VrbError> {
168 let (count, dim) = parse_header(body)?;
169 let expected = expected_body_len(count, dim)?;
170 if body.len() != expected {
171 return Err(VrbError::LengthMismatch {
172 got: body.len(),
173 expected,
174 });
175 }
176 Ok(RawBulk {
177 ids: decode_ids(body, count),
178 vectors: decode_vectors(body, count, dim),
179 dimension: dim,
180 })
181}
182
183#[must_use]
190pub fn encode(ids: &[u64], vectors: &[f32], dimension: usize) -> Vec<u8> {
191 let count = ids.len();
192 let mut buf = Vec::with_capacity(HEADER_LEN + count * 8 + vectors.len() * 4);
193 buf.extend_from_slice(MAGIC);
194 buf.extend_from_slice(&u32::try_from(count).unwrap_or(u32::MAX).to_le_bytes());
198 buf.extend_from_slice(&u32::try_from(dimension).unwrap_or(u32::MAX).to_le_bytes());
199 buf.push(ID_WIDTH);
200 buf.extend_from_slice(&[0u8; 3]);
201 for id in ids {
202 buf.extend_from_slice(&id.to_le_bytes());
203 }
204 for v in vectors {
205 buf.extend_from_slice(&v.to_le_bytes());
206 }
207 buf
208}
209
210#[cfg(test)]
211#[allow(clippy::float_cmp)]
212mod tests {
213 use super::*;
214
215 #[test]
216 fn roundtrip_decode_encode() {
217 let ids = [1u64, 2, 3];
218 let vectors = [0.1f32, 0.2, 0.3, 0.4, 0.5, 0.6];
219 let body = encode(&ids, &vectors, 2);
220 let raw = decode(&body).expect("valid body decodes");
221 assert_eq!(raw.ids, vec![1, 2, 3]);
222 assert_eq!(raw.vectors, vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6]);
223 assert_eq!(raw.dimension, 2);
224 }
225
226 #[test]
227 fn encode_is_deterministic_and_pinned() {
228 let ids = [7u64, 42];
229 let vectors = [1.0f32, 2.0, 3.0, 4.0];
230 let a = encode(&ids, &vectors, 2);
231 let b = encode(&ids, &vectors, 2);
232 assert_eq!(a, b, "encoding must be deterministic");
233 assert_eq!(&a[0..4], b"VRB1");
234 assert_eq!(&a[4..8], &2u32.to_le_bytes());
235 assert_eq!(&a[8..12], &2u32.to_le_bytes());
236 assert_eq!(a[12], 8);
237 assert_eq!(&a[13..16], &[0, 0, 0]);
238 }
239
240 #[test]
241 fn empty_batch_roundtrips() {
242 let body = encode(&[], &[], 4);
243 let raw = decode(&body).expect("empty batch decodes");
244 assert!(raw.ids.is_empty());
245 assert!(raw.vectors.is_empty());
246 assert_eq!(raw.dimension, 4);
247 }
248
249 #[test]
250 fn bad_magic_rejected() {
251 let mut body = encode(&[1], &[0.0, 0.0], 2);
252 body[0] = b'X';
253 assert_eq!(decode(&body), Err(VrbError::BadMagic));
254 }
255
256 #[test]
257 fn short_body_rejected() {
258 let body = vec![0u8; 4];
259 assert_eq!(decode(&body), Err(VrbError::TooShort { got: 4 }));
260 }
261
262 #[test]
263 fn bad_id_width_rejected() {
264 let mut body = encode(&[1], &[0.0, 0.0], 2);
265 body[12] = 4; assert_eq!(decode(&body), Err(VrbError::BadIdWidth(4)));
267 }
268
269 #[test]
270 fn reserved_not_zero_rejected() {
271 let mut body = encode(&[1], &[0.0, 0.0], 2);
272 body[13] = 1;
273 assert_eq!(decode(&body), Err(VrbError::ReservedNotZero));
274 }
275
276 #[test]
277 fn length_mismatch_rejected() {
278 let mut body = encode(&[1, 2], &[0.0, 0.0, 0.0, 0.0], 2);
279 body.pop(); match decode(&body) {
281 Err(VrbError::LengthMismatch { .. }) => {}
282 other => panic!("expected LengthMismatch, got {other:?}"),
283 }
284 }
285
286 #[test]
291 fn overflow_count_dim_rejected() {
292 let mut body = Vec::with_capacity(HEADER_LEN);
293 body.extend_from_slice(MAGIC);
294 body.extend_from_slice(&u32::MAX.to_le_bytes()); body.extend_from_slice(&u32::MAX.to_le_bytes()); body.push(ID_WIDTH);
297 body.extend_from_slice(&[0u8; 3]);
298 assert_eq!(decode(&body), Err(VrbError::Overflow));
299 }
300
301 #[test]
304 fn error_display_and_trait_cover_all_variants() {
305 let cases: [VrbError; 6] = [
306 VrbError::TooShort { got: 3 },
307 VrbError::BadMagic,
308 VrbError::BadIdWidth(4),
309 VrbError::ReservedNotZero,
310 VrbError::Overflow,
311 VrbError::LengthMismatch {
312 got: 10,
313 expected: 16,
314 },
315 ];
316 let rendered: Vec<String> = cases.iter().map(ToString::to_string).collect();
317 assert!(rendered.iter().all(|s| !s.is_empty()));
318 let unique: std::collections::HashSet<&String> = rendered.iter().collect();
320 assert_eq!(unique.len(), cases.len());
321 assert!(rendered[0].contains("too short"));
323 assert!(rendered[2].contains("id_width 4"));
324 let err: &dyn std::error::Error = &cases[1];
326 assert_eq!(err.to_string(), "bad magic: expected b\"VRB1\"");
327 }
328}