1use akar_common::types::Value;
12use std::io::Read;
13
14#[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#[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, }
67
68impl NpyDtype {
69 fn from_str(s: &str) -> Result<Self, NpyReaderError> {
70 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, }
104 }
105}
106
107fn parse_header(data: &[u8]) -> Result<NpyHeader, NpyReaderError> {
109 if data.len() < 10 {
110 return Err(NpyReaderError::InvalidMagic);
111 }
112
113 if &data[0..6] != b"\x93NUMPY" {
115 return Err(NpyReaderError::InvalidMagic);
116 }
117
118 let version = data[6];
120 if version != 1 && version != 2 && version != 3 {
121 return Err(NpyReaderError::InvalidVersion);
122 }
123
124 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 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 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 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; 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 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 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
206pub 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
235fn 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, };
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 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}