1#[cfg(test)]
2mod test;
3
4use crate::io::Write;
5use std::{
6 fs::File,
7 io::{self, BufReader, BufWriter, Error, ErrorKind, Read, Result},
8 mem::size_of,
9 path::Path,
10 str::from_utf8,
11};
12
13pub trait NpyType: Copy {
14 const DESCR: &'static str;
15 const SIZE: usize;
16 fn write_le(self, buffer: &mut Vec<u8>);
17 fn read_le(bytes: &[u8]) -> Self;
18 #[cfg(target_endian = "little")]
19 fn read_from<R: Read>(file: &mut R, count: usize) -> Result<Vec<Self>> {
20 let mut data = Vec::<Self>::with_capacity(count);
21 unsafe {
22 let bytes =
23 std::slice::from_raw_parts_mut(data.as_mut_ptr() as *mut u8, count * Self::SIZE);
24 file.read_exact(bytes)
25 .map_err(|_| invalid("truncated .npy data".into()))?;
26 data.set_len(count);
27 }
28 Ok(data)
29 }
30 #[cfg(target_endian = "big")]
31 fn read_from<R: Read>(file: &mut R, count: usize) -> Result<Vec<Self>> {
32 let mut bytes = vec![0u8; count * Self::SIZE];
33 file.read_exact(&mut bytes)
34 .map_err(|_| invalid("truncated .npy data".into()))?;
35 Ok(Self::read_le_all(&bytes))
36 }
37 #[allow(clippy::chunks_exact_to_as_chunks)]
38 fn read_le_all(bytes: &[u8]) -> Vec<Self> {
39 bytes.chunks_exact(Self::SIZE).map(Self::read_le).collect()
40 }
41 #[cfg(target_endian = "little")]
42 fn write_le_all<W: io::Write>(data: &[Self], file: &mut W) -> Result<()> {
43 let bytes = unsafe {
44 std::slice::from_raw_parts(data.as_ptr() as *const u8, std::mem::size_of_val(data))
45 };
46 file.write_all(bytes)
47 }
48 #[cfg(target_endian = "big")]
49 fn write_le_all<W: io::Write>(data: &[Self], file: &mut W) -> Result<()> {
50 file.write_all(&Self::write_le_bytes(data))
51 }
52 fn write_le_bytes(data: &[Self]) -> Vec<u8> {
53 let mut buffer = Vec::with_capacity(std::mem::size_of_val(data));
54 for &value in data {
55 value.write_le(&mut buffer);
56 }
57 buffer
58 }
59}
60
61macro_rules! npy_type {
62 ($type:ty, $descr:literal) => {
63 impl NpyType for $type {
64 const DESCR: &'static str = $descr;
65 const SIZE: usize = size_of::<$type>();
66 fn write_le(self, buffer: &mut Vec<u8>) {
67 buffer.extend_from_slice(&self.to_le_bytes());
68 }
69 fn read_le(bytes: &[u8]) -> Self {
70 Self::from_le_bytes(bytes.try_into().unwrap())
71 }
72 }
73 };
74}
75npy_type!(u8, "|u1");
76npy_type!(i8, "|i1");
77npy_type!(u16, "<u2");
78npy_type!(i16, "<i2");
79npy_type!(u32, "<u4");
80npy_type!(i32, "<i4");
81npy_type!(u64, "<u8");
82npy_type!(i64, "<i8");
83npy_type!(f32, "<f4");
84npy_type!(f64, "<f8");
85
86pub struct Npy<T> {
94 pub data: Vec<T>,
95 pub shape: Vec<usize>,
96 pub fortran_order: bool,
97}
98
99impl<T, P> Write<P> for Npy<T>
100where
101 T: NpyType,
102 P: AsRef<Path>,
103{
104 type Error = Error;
105 fn write(&self, path: P) -> Result<()> {
106 self.write_to(&mut BufWriter::new(File::create(path)?))
107 }
108}
109
110impl<T: NpyType> Npy<T> {
111 pub fn write_to<W: io::Write>(&self, file: &mut W) -> Result<()> {
112 write_npy(&self.data, &self.shape, self.fortran_order, file)
113 }
114}
115
116fn write_npy<T: NpyType, W: io::Write>(
117 data: &[T],
118 shape: &[usize],
119 fortran_order: bool,
120 file: &mut W,
121) -> Result<()> {
122 let order = if fortran_order { "True" } else { "False" };
123 let dims: String = shape.iter().map(|d| format!("{d}, ")).collect();
124 let mut header = format!(
125 "{{'descr': '{}', 'fortran_order': {order}, 'shape': ({dims}), }}",
126 T::DESCR
127 );
128 let pad = (64 - (10 + header.len() + 1) % 64) % 64;
129 header.push_str(&" ".repeat(pad));
130 header.push('\n');
131 file.write_all(b"\x93NUMPY")?;
132 file.write_all(&[1, 0])?;
133 file.write_all(&(header.len() as u16).to_le_bytes())?;
134 file.write_all(header.as_bytes())?;
135 T::write_le_all(data, file)?;
136 file.flush()
137}
138
139impl<T: NpyType> Npy<T> {
140 pub fn read<P: AsRef<Path>>(path: P) -> Result<Self> {
141 let mut file = BufReader::new(File::open(path)?);
142 Self::read_from(&mut file)
143 }
144 pub fn read_from<R: Read>(file: &mut R) -> Result<Self> {
145 let mut prefix = [0u8; 10];
146 file.read_exact(&mut prefix)
147 .map_err(|_| invalid("not a .npy file".into()))?;
148 if &prefix[..6] != b"\x93NUMPY" {
149 return Err(invalid("not a .npy file".into()));
150 }
151 let header_length = match prefix[6] {
152 1 => u16::from_le_bytes([prefix[8], prefix[9]]) as usize,
153 2 => {
154 let mut rest = [0u8; 2];
155 file.read_exact(&mut rest)?;
156 u32::from_le_bytes([prefix[8], prefix[9], rest[0], rest[1]]) as usize
157 }
158 other => return Err(invalid(format!("unsupported .npy version {other}"))),
159 };
160 let mut header_bytes = vec![0u8; header_length];
161 file.read_exact(&mut header_bytes)?;
162 let header =
163 from_utf8(&header_bytes).map_err(|_| invalid("non-UTF-8 .npy header".into()))?;
164 let descr = quoted(header, "'descr':").ok_or_else(|| invalid("no descr".into()))?;
165 if descr.starts_with('>') || descr.get(1..) != T::DESCR.get(1..) {
166 return Err(invalid(format!(
167 "dtype {descr} does not match {}",
168 T::DESCR
169 )));
170 }
171 let fortran_order = header.contains("'fortran_order': True");
172 let shape = shape(header)?;
173 let count: usize = shape.iter().product();
174 let data = T::read_from(file, count)?;
175 Ok(Npy {
176 data,
177 shape,
178 fortran_order,
179 })
180 }
181}
182
183fn quoted<'a>(header: &'a str, key: &str) -> Option<&'a str> {
184 let at = header.find(key)? + key.len();
185 let open = header[at..].find('\'')? + at + 1;
186 let close = header[open..].find('\'')? + open;
187 Some(&header[open..close])
188}
189
190fn shape(header: &str) -> Result<Vec<usize>> {
191 let at = header
192 .find("'shape':")
193 .ok_or_else(|| invalid("no shape".into()))?;
194 let open = header[at..]
195 .find('(')
196 .ok_or_else(|| invalid("malformed shape".into()))?
197 + at
198 + 1;
199 let close = header[open..]
200 .find(')')
201 .ok_or_else(|| invalid("malformed shape".into()))?
202 + open;
203 header[open..close]
204 .split(',')
205 .map(str::trim)
206 .filter(|dim| !dim.is_empty())
207 .map(|dim| {
208 dim.parse()
209 .map_err(|_| invalid(format!("bad shape entry {dim}")))
210 })
211 .collect()
212}
213
214fn invalid(message: String) -> Error {
215 Error::new(ErrorKind::InvalidData, message)
216}