Skip to main content

diskann_utils/
io.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6//! Read and write vectors in the DiskANN binary format.
7//!
8//! The binary format is:
9//! - 8-byte header
10//!   - `npoints` (u32 LE)
11//!   - `ndims` (u32 LE)
12//! - Payload: `npoints × ndims` elements of `T`, tightly packed in row-major order
13
14use std::io::{Read, Seek, Write};
15
16use diskann_wide::{LoHi, SplitJoin};
17use thiserror::Error;
18
19use crate::views::rowmajor::{self, Layout, Matrix, MatrixMut};
20
21/// Read a matrix of `T` from the DiskANN binary format (see [module docs](self)).
22///
23/// Validates that the reader contains enough data before allocating.
24pub 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
66/// Write a matrix of `T` in the DiskANN binary format (see [module docs](self)).
67///
68/// Returns the total number of bytes written.
69pub 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/// 8-byte header at the start of a DiskANN binary file: `npoints` and `ndims` as little-endian u32.
87#[derive(Debug, Clone, Copy, PartialEq, Eq)]
88pub struct Metadata {
89    npoints: u32,
90    ndims: u32,
91}
92
93impl Metadata {
94    /// Construct from any integer types that fit in `u32`.
95    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    /// Number of points as `usize`.
107    pub fn npoints(&self) -> usize {
108        self.npoints as usize
109    }
110
111    /// Number of points as `u32`.
112    pub fn npoints_u32(&self) -> u32 {
113        self.npoints
114    }
115
116    /// Number of dimensions as `usize`.
117    pub fn ndims(&self) -> usize {
118        self.ndims as usize
119    }
120
121    /// Number of dimensions as `u32`.
122    pub fn ndims_u32(&self) -> u32 {
123        self.ndims
124    }
125
126    /// Destructure into (`npoints`, `ndims`) as `usize`.
127    pub fn into_dims(&self) -> (usize, usize) {
128        (self.npoints(), self.ndims())
129    }
130
131    /// Deserialize the 8-byte header from a reader.
132    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    /// Serialize the 8-byte header to a writer. Returns the number of bytes written (always 8).
150    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/// Error type for [`read_bin`].
169#[derive(Debug, Error)]
170pub enum ReadBinError {
171    /// The reader has fewer bytes remaining than the header declares.
172    #[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    /// The dimensions do not describe a valid allocation (corrupt or malicious header).
185    #[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    /// Underlying IO failure.
195    #[error(transparent)]
196    Io(#[from] std::io::Error),
197}
198
199/// Error type for [`write_bin`].
200#[derive(Debug, Error)]
201pub enum SaveBinError {
202    /// Matrix dimensions exceed `u32::MAX` and cannot be represented in the binary header.
203    #[error("dimensions overflow u32: {nrows} rows × {ncols} cols")]
204    DimensionOverflow { nrows: usize, ncols: usize },
205
206    /// Underlying IO failure.
207    #[error(transparent)]
208    Io(#[from] std::io::Error),
209}
210
211///////////
212// Tests //
213///////////
214
215#[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        // Header says 10 points × 4 dims of f32, but only provide 8 bytes of payload
251        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        // Header with huge values that overflow usize multiplication
280        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        // no payload
307
308        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}