Skip to main content

akar_storage/
npy_reader.rs

1//! NumPy NPY file format reader.
2//!
3//! Reads `.npy` files (NumPy array binary format, v1.0).
4//! Format: magic "\x93NUMPY", version byte, header_len (u16 LE),
5//!         Python dict header, then raw array data.
6//!
7//! Supports: f8 (float64), f4 (float32), i8 (int64), i4 (int32),
8//!           i2 (int16), i1 (int8), u8 (uint64), u4 (uint32),
9//!           u2 (uint16), u1 (uint8), b (bool).
10
11use akar_common::types::Value;
12use std::io::Read;
13
14/// Error type for NPY reader operations.
15#[derive(Debug)]
16pub enum NpyReaderError {
17    IoError(std::io::Error),
18    InvalidMagic,
19    InvalidVersion,
20    InvalidHeader,
21    TypeNotSupported(String),
22    ShapeMismatch { expected_rows: usize, actual: usize },
23}
24
25impl std::fmt::Display for NpyReaderError {
26    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
27        match self {
28            NpyReaderError::IoError(e) => write!(f, "NPY I/O error: {e}"),
29            NpyReaderError::InvalidMagic => write!(f, "Invalid NPY magic bytes"),
30            NpyReaderError::InvalidVersion => write!(f, "Unsupported NPY version"),
31            NpyReaderError::InvalidHeader => write!(f, "Invalid or unparseable NPY header"),
32            NpyReaderError::TypeNotSupported(t) => write!(f, "NPY dtype '{}' not supported", t),
33            NpyReaderError::ShapeMismatch { expected_rows, actual } => write!(
34                f,
35                "Shape mismatch: expected {} rows, file has {} elements",
36                expected_rows, actual
37            ),
38        }
39    }
40}
41
42/// Parsed NPY file header.
43#[derive(Debug)]
44struct NpyHeader {
45    _descr: String,
46    _fortran_order: bool,
47    shape: Vec<usize>,
48    data_offset: usize,
49    dtype: NpyDtype,
50}
51
52#[derive(Debug, Clone, PartialEq)]
53enum NpyDtype {
54    Float64,
55    Float32,
56    Int64,
57    Int32,
58    Int16,
59    Int8,
60    UInt64,
61    UInt32,
62    UInt16,
63    UInt8,
64    Bool,
65    String, // Fixed-length string
66}
67
68impl NpyDtype {
69    fn from_str(s: &str) -> Result<Self, NpyReaderError> {
70        // Normalize: strip byte order prefix and whitespace
71        let s = s.trim();
72        let s = if s.starts_with('<') || s.starts_with('>') || s.starts_with('=') || s.starts_with('|') {
73            &s[1..]
74        } else {
75            s
76        };
77        match s {
78            "f8" | "float64" => Ok(NpyDtype::Float64),
79            "f4" | "float32" => Ok(NpyDtype::Float32),
80            "i8" | "int64" => Ok(NpyDtype::Int64),
81            "i4" | "int32" => Ok(NpyDtype::Int32),
82            "i2" | "int16" => Ok(NpyDtype::Int16),
83            "i1" | "int8" => Ok(NpyDtype::Int8),
84            "u8" | "uint64" => Ok(NpyDtype::UInt64),
85            "u4" | "uint32" => Ok(NpyDtype::UInt32),
86            "u2" | "uint16" => Ok(NpyDtype::UInt16),
87            "u1" | "uint8" => Ok(NpyDtype::UInt8),
88            "b1" | "bool" => Ok(NpyDtype::Bool),
89            _ if s.starts_with('S') || s.starts_with('U') => Ok(NpyDtype::String),
90            _ => Err(NpyReaderError::TypeNotSupported(s.to_string())),
91        }
92    }
93
94    fn size(&self) -> usize {
95        match self {
96            NpyDtype::Float64 => 8,
97            NpyDtype::Float32 => 4,
98            NpyDtype::Int64 | NpyDtype::UInt64 => 8,
99            NpyDtype::Int32 | NpyDtype::UInt32 => 4,
100            NpyDtype::Int16 | NpyDtype::UInt16 => 2,
101            NpyDtype::Int8 | NpyDtype::UInt8 | NpyDtype::Bool => 1,
102            NpyDtype::String => 1, // variable, handle separately
103        }
104    }
105}
106
107/// Parse the NPY header from a byte buffer.
108fn parse_header(data: &[u8]) -> Result<NpyHeader, NpyReaderError> {
109    if data.len() < 10 {
110        return Err(NpyReaderError::InvalidMagic);
111    }
112
113    // Check magic: \x93NUMPY
114    if &data[0..6] != b"\x93NUMPY" {
115        return Err(NpyReaderError::InvalidMagic);
116    }
117
118    // Version byte
119    let version = data[6];
120    if version != 1 && version != 2 && version != 3 {
121        return Err(NpyReaderError::InvalidVersion);
122    }
123
124    // Header length (u16 LE)
125    let header_len = u16::from_le_bytes([data[8], data[9]]) as usize;
126    let data_offset = 10 + header_len;
127
128    if data.len() < data_offset {
129        return Err(NpyReaderError::InvalidHeader);
130    }
131
132    // Parse Python dict header
133    let header_str = std::str::from_utf8(&data[10..data_offset])
134        .map_err(|_| NpyReaderError::InvalidHeader)?
135        .trim()
136        .trim_end_matches('\n');
137
138    // Parse simple Python dict: {'descr': '...', 'fortran_order': bool, 'shape': (...)}
139    let descr = extract_py_str(header_str, "descr").unwrap_or("<f8".to_string());
140    let fortran_order = extract_py_bool(header_str, "fortran_order").unwrap_or(false);
141    let shape = extract_py_tuple(header_str, "shape").unwrap_or(vec![0]);
142
143    // Validate Fortran order
144    if fortran_order {
145        return Err(NpyReaderError::TypeNotSupported(
146            "Fortran-ordered arrays are not supported".into(),
147        ));
148    }
149
150    let dtype = NpyDtype::from_str(&descr)?;
151
152    Ok(NpyHeader {
153        _descr: descr,
154        _fortran_order: fortran_order,
155        shape,
156        data_offset,
157        dtype,
158    })
159}
160
161fn extract_py_str(data: &str, key: &str) -> Option<String> {
162    let key_pat = format!("'{}':", key);
163    let start = data.find(&key_pat)?;
164    let rest = &data[start + key_pat.len()..];
165    let rest = rest.trim();
166    let delim = rest.chars().next()?;
167    let start_inner = 1; // skip opening quote
168    let end_inner = rest[1..].find(delim)?;
169    Some(rest[start_inner..start_inner + end_inner].to_string())
170}
171
172fn extract_py_bool(data: &str, key: &str) -> Option<bool> {
173    let key_pat = format!("'{}':", key);
174    let start = data.find(&key_pat)?;
175    let rest = &data[start + key_pat.len()..].trim();
176    if rest.starts_with("True") {
177        Some(true)
178    } else {
179        Some(false)
180    }
181}
182
183fn extract_py_tuple(data: &str, key: &str) -> Option<Vec<usize>> {
184    let key_pat = format!("'{}':", key);
185    let start = data.find(&key_pat)?;
186    let rest = &data[start + key_pat.len()..].trim();
187
188    if !rest.starts_with('(') {
189        // Single integer
190        let end = rest.find(|c: char| !c.is_ascii_digit()).unwrap_or(rest.len());
191        let num: usize = rest[..end].parse().ok()?;
192        return Some(vec![num]);
193    }
194
195    let end_paren = rest.find(')')?;
196    let inner = &rest[1..end_paren];
197    // Remove trailing comma before closing paren
198    let inner = inner.trim_end_matches(',');
199    if inner.is_empty() {
200        return Some(vec![]);
201    }
202    let nums: Vec<usize> = inner.split(',').filter_map(|s| s.trim().parse().ok()).collect();
203    if nums.is_empty() { None } else { Some(nums) }
204}
205
206/// Read an NPY file and return its contents as `Vec<Value>`.
207///
208/// For 1D arrays, returns one Value per element.
209/// For multi-dimensional arrays, returns the total element count.
210pub fn read_npy(path: &str) -> Result<Vec<Value>, NpyReaderError> {
211    let mut file = std::fs::File::open(path).map_err(NpyReaderError::IoError)?;
212    let mut data = Vec::new();
213    file.read_to_end(&mut data).map_err(NpyReaderError::IoError)?;
214
215    let header = parse_header(&data)?;
216
217    let total_elements: usize = header.shape.iter().product();
218    let raw = &data[header.data_offset..];
219
220    read_values(raw, &header.dtype, total_elements)
221}
222
223fn read_values(raw: &[u8], dtype: &NpyDtype, count: usize) -> Result<Vec<Value>, NpyReaderError> {
224    let elem_size = dtype.size();
225    if count * elem_size > raw.len() {
226        return Err(NpyReaderError::ShapeMismatch {
227            expected_rows: count,
228            actual: raw.len() / elem_size.max(1),
229        });
230    }
231
232    Ok(compute_values(count, dtype, elem_size, raw))
233}
234
235/// Decode `count` little-endian elements of `dtype` from `raw`.
236///
237/// Caller must guarantee `raw` holds at least `count * elem_size` bytes
238/// (`read_values` enforces this before calling).
239fn compute_values(count: usize, dtype: &NpyDtype, elem_size: usize, raw: &[u8]) -> Vec<Value> {
240    let mut values = Vec::with_capacity(count);
241    for i in 0..count {
242        let offset = i * elem_size;
243        let val = match dtype {
244            NpyDtype::Float64 => {
245                let bytes: [u8; 8] = raw[offset..offset + 8].try_into().unwrap();
246                Value::Double(f64::from_le_bytes(bytes))
247            }
248            NpyDtype::Float32 => {
249                let bytes: [u8; 4] = raw[offset..offset + 4].try_into().unwrap();
250                Value::Float(f32::from_le_bytes(bytes))
251            }
252            NpyDtype::Int64 => {
253                let bytes: [u8; 8] = raw[offset..offset + 8].try_into().unwrap();
254                Value::Int64(i64::from_le_bytes(bytes))
255            }
256            NpyDtype::Int32 => {
257                let bytes: [u8; 4] = raw[offset..offset + 4].try_into().unwrap();
258                Value::Int32(i32::from_le_bytes(bytes))
259            }
260            NpyDtype::Int16 => {
261                let bytes: [u8; 2] = raw[offset..offset + 2].try_into().unwrap();
262                Value::Int16(i16::from_le_bytes(bytes))
263            }
264            NpyDtype::Int8 => Value::Int8(raw[offset] as i8),
265            NpyDtype::UInt64 => {
266                let bytes: [u8; 8] = raw[offset..offset + 8].try_into().unwrap();
267                Value::UInt64(u64::from_le_bytes(bytes))
268            }
269            NpyDtype::UInt32 => {
270                let bytes: [u8; 4] = raw[offset..offset + 4].try_into().unwrap();
271                Value::UInt32(u32::from_le_bytes(bytes))
272            }
273            NpyDtype::UInt16 => {
274                let bytes: [u8; 2] = raw[offset..offset + 2].try_into().unwrap();
275                Value::UInt16(u16::from_le_bytes(bytes))
276            }
277            NpyDtype::UInt8 => Value::UInt8(raw[offset]),
278            NpyDtype::Bool => Value::Bool(raw[offset] != 0),
279            NpyDtype::String => Value::Null, // not implemented
280        };
281        values.push(val);
282    }
283    values
284}
285
286#[cfg(test)]
287mod tests {
288    use super::*;
289
290    #[test]
291    fn test_parse_simple_header() {
292        let header_str = "{'descr': '<f8', 'fortran_order': False, 'shape': (3,), }";
293        let mut header_bytes = header_str.as_bytes().to_vec();
294        #[allow(clippy::manual_is_multiple_of)]
295        while (10 + header_bytes.len()) % 16 != 0 {
296            header_bytes.push(b' ');
297        }
298        header_bytes.push(b'\n');
299        let header_len = header_bytes.len() as u16;
300        let data_offset: usize = 10 + header_len as usize;
301
302        let mut buf = vec![];
303        buf.extend_from_slice(b"\x93NUMPY\x01\x00");
304        buf.extend_from_slice(&header_len.to_le_bytes());
305        buf.extend_from_slice(&header_bytes);
306        // Pad to data_offset + data
307        buf.resize(data_offset + 3 * 8, 0);
308        buf[data_offset..data_offset + 8].copy_from_slice(&1.0f64.to_le_bytes());
309        buf[data_offset + 8..data_offset + 16].copy_from_slice(&2.0f64.to_le_bytes());
310        buf[data_offset + 16..data_offset + 24].copy_from_slice(&3.0f64.to_le_bytes());
311
312        let header = parse_header(&buf).unwrap();
313        assert_eq!(header.shape, vec![3]);
314        assert_eq!(header.dtype, NpyDtype::Float64);
315
316        let vals = read_values(&buf[header.data_offset..], &header.dtype, 3).unwrap();
317        assert_eq!(vals.len(), 3);
318        assert_eq!(vals[0], Value::Double(1.0));
319        assert_eq!(vals[1], Value::Double(2.0));
320        assert_eq!(vals[2], Value::Double(3.0));
321    }
322
323    #[test]
324    fn test_npy_int32() {
325        let header_str = "{'descr': '<i4', 'fortran_order': False, 'shape': (2,), }";
326        let mut header_bytes = header_str.as_bytes().to_vec();
327        #[allow(clippy::manual_is_multiple_of)]
328        while (10 + header_bytes.len()) % 16 != 0 {
329            header_bytes.push(b' ');
330        }
331        header_bytes.push(b'\n');
332        let header_len = header_bytes.len() as u16;
333        let data_offset: usize = 10 + header_len as usize;
334
335        let mut buf = vec![];
336        buf.extend_from_slice(b"\x93NUMPY\x01\x00");
337        buf.extend_from_slice(&header_len.to_le_bytes());
338        buf.extend_from_slice(&header_bytes);
339        buf.resize(data_offset + 2 * 4, 0);
340        buf[data_offset..data_offset + 4].copy_from_slice(&42i32.to_le_bytes());
341        buf[data_offset + 4..data_offset + 8].copy_from_slice(&(-7i32).to_le_bytes());
342
343        let header = parse_header(&buf).unwrap();
344        let vals = read_values(&buf[header.data_offset..], &header.dtype, 2).unwrap();
345        assert_eq!(vals[0], Value::Int32(42));
346        assert_eq!(vals[1], Value::Int32(-7));
347    }
348}