1use std::io::{Read, Seek, Write};
15
16use diskann_wide::{LoHi, SplitJoin};
17use thiserror::Error;
18
19use crate::views::rowmajor::{self, Layout, Matrix, MatrixMut};
20
21pub fn read_bin<T>(reader: &mut (impl Read + Seek)) -> Result<rowmajor::Owned<T>, ReadBinError>
25where
26 T: bytemuck::Pod,
27{
28 let metadata = Metadata::read(reader)?;
29 let (npoints, ndims) = (metadata.npoints(), metadata.ndims());
30 let type_size = std::mem::size_of::<T>();
31
32 let layout = match Layout::<T>::new(npoints, ndims) {
33 Ok(layout) => layout,
34 Err(_) => {
35 return Err(ReadBinError::Overflow {
36 npoints: metadata.npoints_u32(),
37 ndims: metadata.ndims_u32(),
38 type_size,
39 })
40 }
41 };
42
43 let expected_bytes = layout.num_elements() * std::mem::size_of::<T>();
44
45 let data_start = reader.stream_position()?;
46 let end = reader.seek(std::io::SeekFrom::End(0))?;
47 let available = end - data_start;
48 reader.seek(std::io::SeekFrom::Start(data_start))?;
49
50 if available < expected_bytes as u64 {
51 return Err(ReadBinError::SizeMismatch {
52 expected: expected_bytes as u64,
53 available,
54 npoints: metadata.npoints_u32(),
55 ndims: metadata.ndims_u32(),
56 type_size,
57 });
58 }
59
60 let mut data =
61 rowmajor::Owned::from_element_with_layout(layout, <T as bytemuck::Zeroable>::zeroed());
62 reader.read_exact(bytemuck::must_cast_slice_mut::<T, u8>(data.as_mut_slice()))?;
63 Ok(data)
64}
65
66pub fn write_bin<T>(
70 data: rowmajor::Ref<'_, T>,
71 writer: &mut impl Write,
72) -> Result<usize, SaveBinError>
73where
74 T: bytemuck::Pod,
75{
76 let metadata =
77 Metadata::new(data.nrows(), data.ncols()).map_err(|_| SaveBinError::DimensionOverflow {
78 nrows: data.nrows(),
79 ncols: data.ncols(),
80 })?;
81 let bytes = metadata.write(writer)?;
82 writer.write_all(bytemuck::must_cast_slice::<T, u8>(data.as_slice()))?;
83 Ok(bytes + std::mem::size_of_val(data.as_slice()))
84}
85
86#[derive(Debug, Clone, Copy, PartialEq, Eq)]
88pub struct Metadata {
89 npoints: u32,
90 ndims: u32,
91}
92
93impl Metadata {
94 pub fn new<T, U>(npoints: T, ndims: U) -> Result<Self, MetadataError<T::Error, U::Error>>
96 where
97 T: TryInto<u32>,
98 U: TryInto<u32>,
99 {
100 Ok(Self {
101 npoints: npoints.try_into().map_err(MetadataError::NumPoints)?,
102 ndims: ndims.try_into().map_err(MetadataError::Dim)?,
103 })
104 }
105
106 pub fn npoints(&self) -> usize {
108 self.npoints as usize
109 }
110
111 pub fn npoints_u32(&self) -> u32 {
113 self.npoints
114 }
115
116 pub fn ndims(&self) -> usize {
118 self.ndims as usize
119 }
120
121 pub fn ndims_u32(&self) -> u32 {
123 self.ndims
124 }
125
126 pub fn into_dims(&self) -> (usize, usize) {
128 (self.npoints(), self.ndims())
129 }
130
131 pub fn read<R>(reader: &mut R) -> std::io::Result<Self>
133 where
134 R: Read,
135 {
136 let mut bytes = [0u8; 8];
137 reader.read_exact(&mut bytes)?;
138
139 let LoHi {
140 lo: npts_bytes,
141 hi: ndims_bytes,
142 } = bytes.split();
143
144 let npoints = u32::from_le_bytes(npts_bytes);
145 let ndims = u32::from_le_bytes(ndims_bytes);
146 Ok(Metadata { npoints, ndims })
147 }
148
149 pub fn write<W>(&self, writer: &mut W) -> std::io::Result<usize>
151 where
152 W: Write,
153 {
154 let bytes: [u8; 8] = LoHi::new(self.npoints.to_le_bytes(), self.ndims.to_le_bytes()).join();
155 writer.write_all(&bytes)?;
156 Ok(2 * std::mem::size_of::<u32>())
157 }
158}
159
160#[derive(Debug, Error)]
161pub enum MetadataError<T, U> {
162 #[error("num points conversion")]
163 NumPoints(#[source] T),
164 #[error("dim conversion")]
165 Dim(#[source] U),
166}
167
168#[derive(Debug, Error)]
170pub enum ReadBinError {
171 #[error(
173 "binary data too short: header declares {npoints} points × {ndims} dims × {type_size} bytes = \
174 {expected} bytes, but only {available} bytes available"
175 )]
176 SizeMismatch {
177 expected: u64,
178 available: u64,
179 npoints: u32,
180 ndims: u32,
181 type_size: usize,
182 },
183
184 #[error(
186 "header dimensions overflow: {npoints} points × {ndims} dims × {type_size} bytes overflows"
187 )]
188 Overflow {
189 npoints: u32,
190 ndims: u32,
191 type_size: usize,
192 },
193
194 #[error(transparent)]
196 Io(#[from] std::io::Error),
197}
198
199#[derive(Debug, Error)]
201pub enum SaveBinError {
202 #[error("dimensions overflow u32: {nrows} rows × {ncols} cols")]
204 DimensionOverflow { nrows: usize, ncols: usize },
205
206 #[error(transparent)]
208 Io(#[from] std::io::Error),
209}
210
211#[cfg(test)]
216mod tests {
217 use std::io::Cursor;
218
219 use crate::assert_contains;
220
221 use super::*;
222
223 #[test]
224 fn round_trip_f32() {
225 let mut counter = 1.0f32;
226 let matrix = rowmajor::Owned::<f32>::from_fn(3, 4, |_| {
227 let v = counter;
228 counter += 1.0;
229 v
230 });
231
232 assert_eq!(
233 matrix.as_slice(),
234 &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0]
235 );
236
237 let mut buf = Vec::new();
238 let written = write_bin(matrix.as_view(), &mut buf).unwrap();
239 assert_eq!(written, 8 + 3 * 4 * 4);
240
241 let mut cursor = Cursor::new(&buf);
242 let loaded = read_bin::<f32>(&mut cursor).unwrap();
243 assert_eq!(loaded.nrows(), 3);
244 assert_eq!(loaded.ncols(), 4);
245 assert_eq!(loaded.as_slice(), matrix.as_slice());
246 }
247
248 #[test]
249 fn read_bin_size_mismatch() {
250 let mut buf = Vec::new();
252 let metadata = Metadata::new(10u32, 4u32).unwrap();
253 metadata.write(&mut buf).unwrap();
254 buf.extend_from_slice(&[0u8; 8]);
255
256 let mut cursor = Cursor::new(&buf);
257 let err = read_bin::<f32>(&mut cursor).unwrap_err();
258
259 match err {
260 ReadBinError::SizeMismatch {
261 expected,
262 available,
263 npoints,
264 ndims,
265 type_size,
266 } => {
267 assert_eq!(expected, 10 * 4 * 4);
268 assert_eq!(available, 8);
269 assert_eq!(npoints, 10);
270 assert_eq!(ndims, 4);
271 assert_eq!(type_size, 4);
272 }
273 other => panic!("expected SizeMismatch, got: {other}"),
274 }
275 }
276
277 #[test]
278 fn read_bin_overflow() {
279 let mut buf = Vec::new();
281 buf.extend_from_slice(&u32::MAX.to_le_bytes());
282 buf.extend_from_slice(&u32::MAX.to_le_bytes());
283
284 let mut cursor = Cursor::new(&buf);
285 let err = read_bin::<f32>(&mut cursor).unwrap_err();
286
287 match err {
288 ReadBinError::Overflow {
289 npoints,
290 ndims,
291 type_size,
292 } => {
293 assert_eq!(npoints, u32::MAX);
294 assert_eq!(ndims, u32::MAX);
295 assert_eq!(type_size, 4);
296 }
297 other => panic!("expected Overflow, got: {other}"),
298 }
299 }
300
301 #[test]
302 fn read_bin_error_message_is_informative() {
303 let mut buf = Vec::new();
304 let metadata = Metadata::new(100u32, 32u32).unwrap();
305 metadata.write(&mut buf).unwrap();
306 let mut cursor = Cursor::new(&buf);
309 let err = read_bin::<f32>(&mut cursor).unwrap_err();
310 let msg = err.to_string();
311
312 assert_contains!(msg, "100 points", "missing npoints");
313 assert_contains!(msg, "32 dims", "missing ndims");
314 assert_contains!(msg, "12800 bytes", "missing expected");
315 assert_contains!(msg, "0 bytes available", "missing available");
316 }
317
318 #[test]
319 fn metadata_read_write_round_trip() {
320 let mut buf = Vec::new();
321 let metadata = Metadata::new(200u32, 128u32).unwrap();
322 metadata.write(&mut buf).unwrap();
323
324 let mut cursor = Cursor::new(&buf);
325 let loaded = Metadata::read(&mut cursor).unwrap();
326 assert_eq!(loaded, metadata);
327 }
328}