From 2fb00f5cc47b49f925b2a9c93a61b9e09451024d Mon Sep 17 00:00:00 2001 From: Patrick Gansterer Date: Fri, 10 Apr 2026 06:16:02 +0200 Subject: [PATCH 1/7] add support for object values Pickeld values are stored with a "|O" TypeStr. --- Cargo.toml | 1 + examples/object.rs | 23 +++++++++++++++++++++++ src/header.rs | 21 +++++++++++++++------ src/read.rs | 22 +++++++++++++++------- src/type_str.rs | 32 +++++++++++++++++++++++++------- test-data/object.npy | Bin 0 -> 275 bytes 6 files changed, 79 insertions(+), 20 deletions(-) create mode 100644 examples/object.rs create mode 100644 test-data/object.npy diff --git a/Cargo.toml b/Cargo.toml index 286cd5f..2f91e7a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -50,6 +50,7 @@ zip = { version = "0.6", default-features = false,features = ["deflate"],optiona # # Also, sprs has an ndarray dependency that might not be the most recent. ndarray = { version = "0.15" } +serde-pickle = { version = "1.2.0" } sprs = { version = "0.11", default-features = false } bencher = { version = "0.1" } diff --git a/examples/object.rs b/examples/object.rs new file mode 100644 index 0000000..edbab57 --- /dev/null +++ b/examples/object.rs @@ -0,0 +1,23 @@ +use std::fs::File; +use std::io; + +// test-data/object.npy is generated by this Python code: +// +// import numpy as np +// np.save('test-data/object.npy', {}) + +fn main() -> Result<(), Box> { + let mut file = io::BufReader::new(File::open("test-data/object.npy")?); + + let header = npyz::NpyHeader::from_reader(&mut file).unwrap(); + + let npyz::DType::Plain(type_str) = header.dtype() else { + panic!() + }; + assert_eq!(type_str.type_char(), npyz::TypeChar::Object); + + let value = serde_pickle::value_from_reader(&mut file, Default::default()); + eprintln!("{:?}", value?); + + Ok(()) +} diff --git a/src/header.rs b/src/header.rs index 662b493..2efc09a 100644 --- a/src/header.rs +++ b/src/header.rs @@ -1,6 +1,7 @@ use std::io; use core::fmt; +use core::num::TryFromIntError; use py_literal::ParseError; pub use py_literal::Value; @@ -113,14 +114,22 @@ impl DType { /// Get the number of bytes that each item of this type occupies. /// /// If this value overflows the plaform's `usize` datatype, returns `None`. - pub fn num_bytes(&self) -> Option { + pub fn num_bytes(&self) -> Result, TryFromIntError> { match self { - DType::Plain(ty) => ty.num_bytes(), - DType::Array(n, inner) => inner.num_bytes()?.checked_mul(usize::try_from(*n).ok()?), - DType::Record(fields) => { - fields.iter().map(|field| field.dtype.num_bytes()) - .fold(Some(0), |a, b| a?.checked_add(b?)) + DType::Plain(ty) => Ok(ty.num_bytes()), + DType::Array(n, inner) => match inner.num_bytes()? { + Some(num_bytes) => Ok(num_bytes.checked_mul(usize::try_from(*n)?)), + None => Ok(None), }, + DType::Record(fields) => { + fields + .iter() + .map(|field| field.dtype.num_bytes()) + .fold(Ok(Some(0)), |a, b| match (a?, b?) { + (Some(a), Some(b)) => Ok(a.checked_add(b)), + _ => Ok(None), + }) + } } } } diff --git a/src/read.rs b/src/read.rs index 9fe5698..94aaa99 100644 --- a/src/read.rs +++ b/src/read.rs @@ -159,7 +159,7 @@ pub struct NpyHeader { /// Total number of elements, pre-computed from the shape. n_records: u64, /// Item size in bytes. - item_size: usize, + item_size: Option, } impl NpyHeader { @@ -336,9 +336,9 @@ impl NpyHeader { fn from_parts(dtype: DType, shape: Vec, order: Order) -> io::Result { let n_records = shape.iter().product(); - let item_size = dtype.num_bytes().ok_or_else(|| { - invalid_data(format_args!("dtype is larger than usize!")) - })?; + let item_size = dtype + .num_bytes() + .or_else(|_| Err(invalid_data(format_args!("dtype is larger than usize!"))))?; let strides = strides(order, &shape); Ok(NpyHeader { dtype, shape, strides, order, n_records, item_size }) } @@ -390,7 +390,11 @@ impl NpyReader where R: io::Seek { let (reader, current_index) = &mut self.reader_and_current_index; let delta = index as i64 - *current_index as i64; if delta != 0 { - reader.seek(io::SeekFrom::Current(delta * self.header.item_size as i64))?; + let item_size = self + .header + .item_size + .ok_or(io::Error::from(io::ErrorKind::Unsupported))?; + reader.seek(io::SeekFrom::Current(delta * item_size as i64))?; *current_index = index; } Ok(()) @@ -416,7 +420,10 @@ impl<'a, T: Deserialize> NpyData<'a, T> { pub fn from_bytes(bytes: &'a [u8]) -> io::Result> { let inner = NpyFile::new(bytes)?.data().map_err(invalid_data)?; - assert_eq!(inner.header.item_size as u64 * inner.header.n_records, inner.reader().len() as u64); + assert_eq!( + inner.header.item_size.unwrap_or_default() as u64 * inner.header.n_records, + inner.reader().len() as u64 + ); Ok(NpyData { inner }) } @@ -457,7 +464,8 @@ impl<'a, T: Deserialize> NpyData<'a, T> { /// Panics if the bytes stored for the element are invalid for the dtype, /// or if the index is out of bounds. pub fn get_unchecked(&self, i: usize) -> T { - let item_bytes = &self.get_data_slice()[i * self.inner.header.item_size..]; + let item_size = self.inner.header.item_size.unwrap(); + let item_bytes = &self.get_data_slice()[i * item_size..]; self.inner.type_reader.read_one(item_bytes).unwrap() } diff --git a/src/type_str.rs b/src/type_str.rs index 4cab1da..9da8c13 100644 --- a/src/type_str.rs +++ b/src/type_str.rs @@ -170,6 +170,10 @@ pub enum TypeChar { /// See [`type_matchup_docs`][`crate::type_matchup_docs`] for information on which types can /// use this for serialization. UnicodeStr, + /// Code `O`. + /// + /// Notice that numpy uses this to store pickled values. + Object, /// Code `V`. Represents a binary blob of `size` bytes. /// /// Can use [`crate::FixedSizeBytes`] for serialization, or some other types; see @@ -190,6 +194,7 @@ impl TypeChar { 'M' => Some(TypeChar::DateTime), 'S' | 'a' => Some(TypeChar::ByteStr), 'U' => Some(TypeChar::UnicodeStr), + 'O' => Some(TypeChar::Object), 'V' => Some(TypeChar::RawData), _ => None, } @@ -207,6 +212,7 @@ impl TypeChar { TypeChar::DateTime => "M", TypeChar::ByteStr => "S", TypeChar::UnicodeStr => "U", + TypeChar::Object => "O", TypeChar::RawData => "V", } } @@ -235,6 +241,7 @@ impl TypeChar { // changes them to `|S1` and `|U1`.) TypeChar::ByteStr | TypeChar::UnicodeStr | + TypeChar::Object | TypeChar::RawData => None, } } @@ -253,6 +260,7 @@ impl TypeChar { TypeChar::UnicodeStr => true, TypeChar::ByteStr | + TypeChar::Object | TypeChar::RawData => false, } } @@ -292,6 +300,8 @@ fn type_str_num_bytes_as_usize(type_str: &TypeStr) -> Option { TypeChar::RawData => Some(size_field), TypeChar::UnicodeStr => size_field.checked_mul(4), + + TypeChar::Object =>None, } } @@ -453,7 +463,7 @@ mod parse { fn from_str(input: &str) -> Result { use self::ErrorKind::*; - if input.len() < 3 { + if input.len() < 2 { bail!(SyntaxError); } @@ -476,13 +486,19 @@ mod parse { remainder.bytes().position(|b| !b.is_ascii_digit()) .unwrap_or(remainder.len()) }; - if size_end == 0 { - bail!(SyntaxError); - } let (size, remainder) = remainder.split_at(size_end); - let size = match size.parse() { - Err(e) => bail!(ParseIntError(e)), // probably overflow - Ok(v) => v, + let size = if type_char == TypeChar::Object { + match size { + "" => 0, + _ => bail!(SyntaxError), + } + } else if size_end == 0 { + bail!(SyntaxError); + } else { + match size.parse() { + Err(e) => bail!(ParseIntError(e)), // probably overflow + Ok(v) => v, + } }; let time_units = if remainder.is_empty() { @@ -572,9 +588,11 @@ mod parse { check_err!("", SyntaxError); check_err!(">", SyntaxError); check_err!(">i", SyntaxError); + check_err!("|O4", SyntaxError); check_ok!(">i8"); check_ok!(">c16"); check_err!(">i8garbage", SyntaxError); + check_ok!("|O"); // length-zero integer check_err!(">m[us]", SyntaxError); diff --git a/test-data/object.npy b/test-data/object.npy new file mode 100644 index 0000000000000000000000000000000000000000..2831eafd48019f9994811ebb6bd24358f48a749e GIT binary patch literal 275 zcmbR27wQ`j$;eQ~P_3SlTAW;@Zl$1Js?sT~Xu&?A;tnp;q*7oVJ8l&Y6onp2XQSX7i)Ii-guz9=<0 zKd-o?s5H4`%H%0MtYDQ>df4+)AW9}r@n&e9;>?&drF}}!6b)}i%?#!q)|8UUf>e-t oCVz`5ogEM<<{rkHDSm!_UjKmrOnBRuOiAjDol-j`5y;R30QMqR?f?J) literal 0 HcmV?d00001 From 238a87ea478144d176bc6c55ce039fb61910a7b9 Mon Sep 17 00:00:00 2001 From: Michael Lamparski Date: Sat, 11 Apr 2026 17:24:29 -0400 Subject: [PATCH 2/7] tweaks and fixes on object typestr --- src/header.rs | 33 +++++++++++++++++---------------- src/read.rs | 46 +++++++++++++++++++++++++++++++++------------- src/type_str.rs | 42 ++++++++++++++++++++++-------------------- 3 files changed, 72 insertions(+), 49 deletions(-) diff --git a/src/header.rs b/src/header.rs index 2efc09a..c804741 100644 --- a/src/header.rs +++ b/src/header.rs @@ -1,7 +1,6 @@ use std::io; use core::fmt; -use core::num::TryFromIntError; use py_literal::ParseError; pub use py_literal::Value; @@ -113,23 +112,25 @@ impl DType { /// Get the number of bytes that each item of this type occupies. /// - /// If this value overflows the plaform's `usize` datatype, returns `None`. - pub fn num_bytes(&self) -> Result, TryFromIntError> { + /// If this size is not fixed (e.g. `|O`) or would overflow the platform's `usize` type, returns `None`. + /// You can differentiate between these two error cases by calling [`Self::has_variable_size`]. + pub fn num_bytes(&self) -> Option { match self { - DType::Plain(ty) => Ok(ty.num_bytes()), - DType::Array(n, inner) => match inner.num_bytes()? { - Some(num_bytes) => Ok(num_bytes.checked_mul(usize::try_from(*n)?)), - None => Ok(None), - }, + DType::Plain(ty) => ty.num_bytes(), + DType::Array(n, inner) => inner.num_bytes()?.checked_mul(usize::try_from(*n).ok()?), DType::Record(fields) => { - fields - .iter() - .map(|field| field.dtype.num_bytes()) - .fold(Ok(Some(0)), |a, b| match (a?, b?) { - (Some(a), Some(b)) => Ok(a.checked_add(b)), - _ => Ok(None), - }) - } + fields.iter().map(|field| field.dtype.num_bytes()) + .fold(Some(0), |a, b| a?.checked_add(b?)) + }, + } + } + + /// `true` if the DType may have a variable number of bytes for each element. + pub fn has_variable_size(&self) -> bool { + match self { + DType::Plain(ty) => ty.has_variable_size(), + DType::Array(_, inner) => inner.has_variable_size(), + DType::Record(fields) => fields.iter().any(|field| field.dtype.has_variable_size()), } } } diff --git a/src/read.rs b/src/read.rs index 94aaa99..f7d7937 100644 --- a/src/read.rs +++ b/src/read.rs @@ -158,7 +158,7 @@ pub struct NpyHeader { order: Order, /// Total number of elements, pre-computed from the shape. n_records: u64, - /// Item size in bytes. + /// Item size in bytes. `None` if elements are variable size. item_size: Option, } @@ -336,9 +336,13 @@ impl NpyHeader { fn from_parts(dtype: DType, shape: Vec, order: Order) -> io::Result { let n_records = shape.iter().product(); - let item_size = dtype - .num_bytes() - .or_else(|_| Err(invalid_data(format_args!("dtype is larger than usize!"))))?; + let item_size = match dtype.has_variable_size() { + true => None, + false => match dtype.num_bytes() { + Some(num) => Some(num), + None => Err(invalid_data(format_args!("dtype is larger than usize!")))?, + }, + }; let strides = strides(order, &shape); Ok(NpyHeader { dtype, shape, strides, order, n_records, item_size }) } @@ -384,6 +388,8 @@ impl NpyReader where R: io::Seek { /// /// Panics if the index is greater than [`Self::total_len`]. pub fn seek_to(&mut self, index: u64) -> io::Result<()> { + const NO_SEEK_MSG: &str = "array with variable size elements does not support seeking"; + let len = self.total_len(); assert!(index <= len, "index out of bounds for seeking (the index is {} but the len is {})", index, len); @@ -393,7 +399,7 @@ impl NpyReader where R: io::Seek { let item_size = self .header .item_size - .ok_or(io::Error::from(io::ErrorKind::Unsupported))?; + .ok_or(io::Error::new(io::ErrorKind::Unsupported, NO_SEEK_MSG))?; reader.seek(io::SeekFrom::Current(delta * item_size as i64))?; *current_index = index; } @@ -417,13 +423,23 @@ impl NpyReader where R: io::Seek { #[allow(deprecated)] impl<'a, T: Deserialize> NpyData<'a, T> { /// Deserialize a NPY file represented as bytes + /// + /// Returns `Err` if the header cannot be parsed. + /// + /// Panics if the buffer is not the correct size for the payload indicated by the header. pub fn from_bytes(bytes: &'a [u8]) -> io::Result> { let inner = NpyFile::new(bytes)?.data().map_err(invalid_data)?; - assert_eq!( - inner.header.item_size.unwrap_or_default() as u64 * inner.header.n_records, - inner.reader().len() as u64 - ); + if let Some(item_size) = inner.header.item_size { + // FIXME: This is unfair. The caller cannot verify that their buffer has the correct size in + // advance as they do not know the array shape or where the header ends. This should go the + // same way as any "corrupt" file: io::Error. + assert_eq!( + item_size as u64 * inner.header.n_records, + inner.reader().len() as u64, + ); + } + Ok(NpyData { inner }) } @@ -444,11 +460,13 @@ impl<'a, T: Deserialize> NpyData<'a, T> { /// Gets a single data-record with the specified flat index. /// - /// Returns None if the index is out of bounds. + /// Returns `None` if the index is out of bounds. /// /// # Panics /// - /// Panics if the bytes stored for the element are invalid for the dtype. + /// Panics in the following cases: + /// * The bytes stored for the element are invalid for the dtype. + /// * The type forbids random access due to having variable-size items. pub fn get(&self, i: usize) -> Option { if i < self.len() { Some(self.get_unchecked(i)) @@ -461,8 +479,10 @@ impl<'a, T: Deserialize> NpyData<'a, T> { /// /// # Panics /// - /// Panics if the bytes stored for the element are invalid for the dtype, - /// or if the index is out of bounds. + /// Panics in the following cases: + /// * The index is out of bounds. + /// * The bytes stored for the element are invalid for the dtype. + /// * The type forbids random access due to having variable-size items. pub fn get_unchecked(&self, i: usize) -> T { let item_size = self.inner.header.item_size.unwrap(); let item_bytes = &self.get_data_slice()[i * item_size..]; diff --git a/src/type_str.rs b/src/type_str.rs index 9da8c13..8cbf0a2 100644 --- a/src/type_str.rs +++ b/src/type_str.rs @@ -32,25 +32,29 @@ impl TypeStr { pub fn endianness(&self) -> Endianness { self.endianness } /// Extract the type character from the type string. - /// - /// For most **(but not all!)** types, this is the number of bytes that a single value occupies. - /// For the `U` type, it is the number of code units. pub fn type_char(&self) -> TypeChar { self.type_char } /// Extract the "size" field from the type string. This is the number that appears after the type character. /// - /// For most **(but not all!)** types, this is the number of bytes that a single value occupies. - /// For the `U` type, it is the number of code units. + /// * For **most** types, this is the number of bytes that a single value occupies. + /// * For the [`'U'`](TypeChar::UnicodeStr) type, it is the number of code units. + /// * For the [`'O'`](TypeChar::Object) type, this returns zero. (the actual string does not contain a number here) pub fn size_field(&self) -> u64 { self.size } - /// Extract the time units, if this type string has any. Only [`TypeChar::TimeDelta`] and - /// [`TypeChar::DateTime`] have time units. + /// Extract the time units, if this type string has any. Only [`'m'`](TypeChar::TimeDelta) and + /// [`'M'`](TypeChar::DateTime) have time units. pub fn time_units(&self) -> Option { self.time_units } /// Get the number of bytes for a single value. /// - /// If this value would overflow the platform's `usize` type, returns `None`. + /// If this size is not fixed (e.g. `|O`) or would overflow the platform's `usize` type, returns `None`. + /// You can differentiate between these two error cases by calling [`Self::has_variable_size`]. pub fn num_bytes(&self) -> Option { type_str_num_bytes_as_usize(self) } + + /// Returns true if each element in the array can vary in size. + /// + /// This is true if and only if the type character is [`'O'`](TypeChar::Object). + pub fn has_variable_size(&self) -> bool { matches!(self.type_char, TypeChar::Object) } } /// Represents the first character in a [`TypeStr`], which describes endianness. @@ -62,7 +66,10 @@ pub enum Endianness { Big, /// Code `|`. Used when endianness is irrelevant. /// - /// Only valid when the size is `1`, or the type character is [`TypeChar::ByteStr`]. + /// Only used in the following cases: + /// * The size is `1`. + /// * The type character is [`TypeChar::ByteStr`]. + /// * The type character is [`TypeChar::Object`]. Irrelevant, } @@ -487,17 +494,12 @@ mod parse { .unwrap_or(remainder.len()) }; let (size, remainder) = remainder.split_at(size_end); - let size = if type_char == TypeChar::Object { - match size { - "" => 0, - _ => bail!(SyntaxError), - } - } else if size_end == 0 { - bail!(SyntaxError); - } else { - match size.parse() { - Err(e) => bail!(ParseIntError(e)), // probably overflow - Ok(v) => v, + let size = match (type_char, size) { + (TypeChar::Object, "") => 0, + (TypeChar::Object, _) | (_, "") => bail!(SyntaxError), + _ => match size.parse() { + Ok(size) => size, + Err(e) => bail!(ParseIntError(e)), } }; From 1fa67ca451496a52e122d53e2d61050a75e4f6ad Mon Sep 17 00:00:00 2001 From: Michael Lamparski Date: Sat, 11 Apr 2026 19:40:57 -0400 Subject: [PATCH 3/7] more tests, fix object display --- src/header.rs | 54 +++++++++++++++++++++++++++++++++++++++++++++++-- src/type_str.rs | 16 ++++++++++++++- src/write.rs | 2 ++ 3 files changed, 69 insertions(+), 3 deletions(-) diff --git a/src/header.rs b/src/header.rs index c804741..fe9f504 100644 --- a/src/header.rs +++ b/src/header.rs @@ -102,6 +102,11 @@ impl DType { DType::Plain(ty) } + /// Construct a scalar `DType` of type `|O`, for pickled data. + pub fn new_pickled() -> Self { + DType::Plain("|O".parse().unwrap()) + } + /// Return a `TypeStr` only if the `DType` is a primitive scalar. (no arrays or record types) pub(crate) fn as_scalar(&self) -> Option<&TypeStr> { match self { @@ -384,7 +389,7 @@ mod tests { } #[test] - fn converts_simple_description_to_record_dtype() -> TestResult { + fn converts_simple_description_to_plain_dtype() -> TestResult { let dtype = ">f8"; assert_eq!( DType::from_descr(&Value::String(dtype.to_string())).unwrap(), @@ -394,7 +399,7 @@ mod tests { } #[test] - fn converts_non_endian_description_to_record_dtype() -> TestResult { + fn converts_non_endian_description_to_plain_dtype() -> TestResult { let dtype = "|u1"; assert_eq!( DType::from_descr(&Value::String(dtype.to_string())).unwrap(), @@ -403,6 +408,17 @@ mod tests { Ok(()) } + + #[test] + fn converts_object_description_to_plain_dtype() -> TestResult { + let dtype = "|O"; + assert_eq!( + DType::from_descr(&Value::String(dtype.to_string())).unwrap(), + DType::Plain(dtype.parse()?), + ); + Ok(()) + } + #[test] fn converts_record_description_to_record_dtype() -> TestResult { let descr = parse("[('a', ' TestResult { + let descr = parse("[('a', ' TestResult { let original_dtype = DType::Record(vec![ @@ -482,6 +515,23 @@ mod tests { Ok(()) } + #[test] + fn test_size_of_packed_data() -> TestResult { + let descr = parse("[('a', 'i4')]"); + let dtype = DType::from_descr(&descr).unwrap(); + assert_eq!(dtype.num_bytes(), Some(6)); + Ok(()) + } + + #[test] + fn test_variable_size() -> TestResult { + let descr = parse("[('a', ' fmt::Result { - write!(f, "{}{}{}", self.endianness, self.type_char, self.size)?; + write!(f, "{}{}", self.endianness, self.type_char)?; + if !self.has_variable_size() { + write!(f, "{}", self.size)?; + } if let Some(time_units) = self.time_units { write!(f, "[{}]", time_units)?; } @@ -666,6 +669,16 @@ mod tests { "|S13", ); + assert_eq!( + TypeStr { + endianness: Endianness::Irrelevant, + type_char: TypeChar::Object, + size: 0, + time_units: None, + }.to_string(), + "|O", + ); + assert_eq!( TypeStr { endianness: Endianness::Big, @@ -699,6 +712,7 @@ mod tests { check_roundtrip!("|S0"); check_roundtrip!("U3"); + check_roundtrip!("|O"); check_roundtrip!("m8[ms]"); } diff --git a/src/write.rs b/src/write.rs index eecfa14..9476c9a 100644 --- a/src/write.rs +++ b/src/write.rs @@ -38,6 +38,8 @@ pub mod write_options { /// Construction of an [`NpyWriter`] always begins here, with a call to [`WriteOptions::new`]. /// Then the methods of the [`WriterBuilder`] trait must be used to supply options. /// See that trait for more details. + /// + /// (Notice: This means you will have to `use npyz::WriterBuilder` to work with this...) #[derive(Debug)] pub struct WriteOptions { order: Order, From fb03c7662630e8c354c7f7577dfd7d11375eec28 Mon Sep 17 00:00:00 2001 From: Michael Lamparski Date: Sun, 12 Apr 2026 11:46:49 -0400 Subject: [PATCH 4/7] Flesh out object dtype I had plans for a `"pickle"` feature, PickleReader/ PickleWriter helpers, impls for serde_pickle::Value... and then I learned that numpy pickles the ENTIRE array object and not its items. So here's what we're left with. --- Cargo.toml | 7 +++-- examples/object.rs | 23 --------------- examples/pickle.rs | 43 +++++++++++++++++++++++++++ src/header.rs | 22 +++++++------- src/lib.rs | 2 +- src/read.rs | 39 ++++++++++++++++++++++-- src/serialize/mod.rs | 3 +- src/serialize/traits.rs | 6 +++- src/type_matchup_docs.rs | 62 +++++++++++++++++++++++++++++---------- src/type_str.rs | 8 +++-- test-data/object.npy | Bin 275 -> 0 bytes test-data/pickle.npy | Bin 0 -> 317 bytes 12 files changed, 153 insertions(+), 62 deletions(-) delete mode 100644 examples/object.rs create mode 100644 examples/pickle.rs delete mode 100644 test-data/object.npy create mode 100644 test-data/pickle.npy diff --git a/Cargo.toml b/Cargo.toml index 2f91e7a..2c44071 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -39,10 +39,11 @@ optional = true default-features = false [target.'cfg(not(target_arch = "wasm32"))'.dependencies] -zip = { version = "0.6",optional = true} +zip = { version = "0.6", optional = true } [target.'cfg(target_arch = "wasm32")'.dependencies] -zip = { version = "0.6", default-features = false,features = ["deflate"],optional = true} +zip = { version = "0.6", default-features = false, features = ["deflate"], optional = true } + [dev-dependencies] # For examples ONLY. We don't want to provide a public interface because ndarray undergoes @@ -58,7 +59,7 @@ bencher = { version = "0.1" } zip = { version = "0.6", default-features = true } # NOTICE: also in dependencies [target.'cfg(target_arch = "wasm32")'.dev-dependencies] -zip = { version = "0.6", default-features = false,features = ["deflate"]} +zip = { version = "0.6", default-features = false, features = ["deflate"] } wasm-bindgen = "0.2" wasm-bindgen-test = "0.3" diff --git a/examples/object.rs b/examples/object.rs deleted file mode 100644 index edbab57..0000000 --- a/examples/object.rs +++ /dev/null @@ -1,23 +0,0 @@ -use std::fs::File; -use std::io; - -// test-data/object.npy is generated by this Python code: -// -// import numpy as np -// np.save('test-data/object.npy', {}) - -fn main() -> Result<(), Box> { - let mut file = io::BufReader::new(File::open("test-data/object.npy")?); - - let header = npyz::NpyHeader::from_reader(&mut file).unwrap(); - - let npyz::DType::Plain(type_str) = header.dtype() else { - panic!() - }; - assert_eq!(type_str.type_char(), npyz::TypeChar::Object); - - let value = serde_pickle::value_from_reader(&mut file, Default::default()); - eprintln!("{:?}", value?); - - Ok(()) -} diff --git a/examples/pickle.rs b/examples/pickle.rs new file mode 100644 index 0000000..f3ce2f6 --- /dev/null +++ b/examples/pickle.rs @@ -0,0 +1,43 @@ +use std::io; +use std::fs::File; + +use serde_pickle; + +// test-data/pickle.npy is generated by this Python code: +// +// import numpy as np +// a = np.array([ +// [1, 3.5, "number"], +// [{}, {}, None], +// ], dtype=object) +// np.save('test-data/pickle.npy', a) + +fn main() -> io::Result<()> { + let file = io::BufReader::new(File::open("test-data/pickle.npy")?); + let npy = npyz::NpyFile::new(file)?; + + assert_eq!(npy.shape(), &[2, 3]); + println!("DType: {}", npy.dtype().descr()); + println!("Shape: {:?}", npy.shape()); + println!("Strides: {:?}", npy.strides()); + + // When an array involves objects, the **entire** array is written in a single pickle. + // npyz can't help us any further here. + let file = npy.into_inner(); + let de_options = Default::default(); + let value = serde_pickle::value_from_reader(file, de_options).map_err(pickle_err_to_io_err)?; + + // The pickled value begins with a bunch of array metadata that's tough to make sense of, + // but at the end you will see the data, a flat list of the objects written by the python + // snippet at the top of this file. + println!("Pickled value: {:#?}", value); + + Ok(()) +} + +fn pickle_err_to_io_err(error: serde_pickle::Error) -> io::Error { + match error { + serde_pickle::Error::Io(e) => e, + error => io::Error::new(io::ErrorKind::Other, error.to_string()), + } +} diff --git a/src/header.rs b/src/header.rs index fe9f504..5a88b1a 100644 --- a/src/header.rs +++ b/src/header.rs @@ -102,11 +102,6 @@ impl DType { DType::Plain(ty) } - /// Construct a scalar `DType` of type `|O`, for pickled data. - pub fn new_pickled() -> Self { - DType::Plain("|O".parse().unwrap()) - } - /// Return a `TypeStr` only if the `DType` is a primitive scalar. (no arrays or record types) pub(crate) fn as_scalar(&self) -> Option<&TypeStr> { match self { @@ -130,12 +125,17 @@ impl DType { } } - /// `true` if the DType may have a variable number of bytes for each element. - pub fn has_variable_size(&self) -> bool { + /// `true` if an array using this `DType` is stored using `pickle`. + /// + /// This is true if and only if the type contains an [`'O'`](TypeChar::Object) dtype. + /// + /// Pickled arrays present unique challenges, and most of the `npyz` crate does not support them + /// beyond parsing or writing the file header. + pub fn uses_pickled_array(&self) -> bool { match self { - DType::Plain(ty) => ty.has_variable_size(), - DType::Array(_, inner) => inner.has_variable_size(), - DType::Record(fields) => fields.iter().any(|field| field.dtype.has_variable_size()), + DType::Plain(ty) => ty.uses_pickled_array(), + DType::Array(_, inner) => inner.uses_pickled_array(), + DType::Record(fields) => fields.iter().any(|field| field.dtype.uses_pickled_array()), } } } @@ -528,7 +528,7 @@ mod tests { let descr = parse("[('a', ', + uses_pickled_array: bool, } impl NpyHeader { @@ -229,6 +230,11 @@ impl NpyFile { pub fn header(&self) -> &NpyHeader { &self.header } + + /// Recover the wrapped [`io::Read`], which now points at the beginning of the raw data bytes. + pub fn into_inner(self) -> R { + self.reader + } } // Provided for backwards compatibility. @@ -269,6 +275,25 @@ impl NpyHeader { pub fn len(&self) -> u64 { self.n_records } + + /// `true` if an array using this `DType` is stored using [`pickle`](https://docs.python.org/3/library/pickle.html). + /// + /// This is true if and only if the type contains an [`'O'`](crate::TypeChar::Object) dtype. + /// + /// Pickled arrays present unique challenges, and you will be unable to extract items + /// from this array using this crate. However, you will be able to extract array metadata + /// such as [`Self::shape`]. + pub fn uses_pickled_array(&self) -> bool { + self.uses_pickled_array + } + + fn forbid_pickle(&self) -> Result<(), DTypeError> { + if self.uses_pickled_array { + Err(DTypeError(DTypeErrorKind::RequiresPickle)) + } else { + Ok(()) + } + } } impl NpyFile { @@ -287,6 +312,8 @@ impl NpyFile { /// The returned type implements [`Iterator`]`>`, and provides additional methods /// for random access when `R: Seek`. See [`NpyReader`] for more details. pub fn data(self) -> Result, DTypeError> { + self.forbid_pickle()?; + let NpyFile { reader, header } = self; let type_reader = T::reader(&header.dtype)?; Ok(NpyReader { type_reader, header, reader_and_current_index: (reader, 0) }) @@ -296,6 +323,10 @@ impl NpyFile { /// /// This fallible form of the function returns `self` on error, so that you can try again with a different `T`. pub fn try_data(self) -> Result, Self> { + if self.uses_pickled_array { + return Err(self); + } + let type_reader = match T::reader(&self.header.dtype) { Ok(r) => r, Err(_) => return Err(self), @@ -336,7 +367,8 @@ impl NpyHeader { fn from_parts(dtype: DType, shape: Vec, order: Order) -> io::Result { let n_records = shape.iter().product(); - let item_size = match dtype.has_variable_size() { + let uses_pickled_array = dtype.uses_pickled_array(); + let item_size = match uses_pickled_array { true => None, false => match dtype.num_bytes() { Some(num) => Some(num), @@ -344,7 +376,7 @@ impl NpyHeader { }, }; let strides = strides(order, &shape); - Ok(NpyHeader { dtype, shape, strides, order, n_records, item_size }) + Ok(NpyHeader { dtype, shape, strides, order, n_records, item_size, uses_pickled_array }) } } @@ -388,6 +420,7 @@ impl NpyReader where R: io::Seek { /// /// Panics if the index is greater than [`Self::total_len`]. pub fn seek_to(&mut self, index: u64) -> io::Result<()> { + // NOTE: This error may be unreachable due to pickle errors taking priority. const NO_SEEK_MSG: &str = "array with variable size elements does not support seeking"; let len = self.total_len(); diff --git a/src/serialize/mod.rs b/src/serialize/mod.rs index 5cfe95c..47c8b58 100644 --- a/src/serialize/mod.rs +++ b/src/serialize/mod.rs @@ -9,7 +9,8 @@ mod test_helpers; pub use traits::{Serialize, Deserialize, AutoSerialize}; pub use traits::{TypeRead, TypeWrite, TypeWriteDyn, TypeReadDyn, DTypeError}; -use traits::{helper, ErrorKind}; +pub(crate) use traits::ErrorKind; +use traits::helper; #[macro_use] mod traits; diff --git a/src/serialize/traits.rs b/src/serialize/traits.rs index 04ea1e9..cbaab35 100644 --- a/src/serialize/traits.rs +++ b/src/serialize/traits.rs @@ -225,6 +225,7 @@ pub(crate) enum ErrorKind { verb: &'static str, }, UsizeOverflow(u64), + RequiresPickle, } impl std::error::Error for DTypeError {} @@ -292,6 +293,9 @@ impl fmt::Display for DTypeError { ErrorKind::UsizeOverflow(value) => { write!(f, "cannot cast {} as usize", value) }, + ErrorKind::RequiresPickle => { + write!(f, "this dtype uses a pickled array, which npyz's read/write APIs do not currently support") + } } } } @@ -384,4 +388,4 @@ mod tests { writer.write_one(&mut buf, &4000).unwrap(); assert_eq!(reader.read_one(&buf[..]).unwrap(), 4000); } -} \ No newline at end of file +} diff --git a/src/type_matchup_docs.rs b/src/type_matchup_docs.rs index 5c52644..e1fc6ff 100644 --- a/src/type_matchup_docs.rs +++ b/src/type_matchup_docs.rs @@ -17,27 +17,30 @@ Integers and floats correspond to simple dtypes: ### Integers -* The rust types `i8`, `i16`, `i32`, `i64` use type code `i`. -* The rust types `u8`, `u16`, `u32`, `u64` use type code `u`. +* The rust types `i8`, `i16`, `i32`, `i64` use type code [`i`](TypeChar::Int). +* The rust types `u8`, `u16`, `u32`, `u64` use type code [`u`](TypeChar::Uint). **Notice:** numpy does not support 128-bit integers ### Floats -* The rust types `f32`, `f64` use type code `f`. +* The rust types `f32`, `f64` use type code [`f`](TypeChar::Float). * When the **`"half"`** feature is enabled, [`f16`] is also supported. **Notice:** numpy *does* have 128-bit floats, but it is not currently supported by `npyz`. +(also, these might just be 80-bit floats on `X86_64`; there's a similar `np.float96` on 32 bit...) + ### Complex -When the **`"complex"`** feature is enabled, rust types [`Complex32`] and [`Complex64`] may use type code `c`. +When the **`"complex"`** feature is enabled, rust types [`Complex32`] and [`Complex64`] may use type code +[`c`](TypeChar::Complex). **Notice:** numpy does have have complex numbers backed by 128-bit floats, but this is not supported by `npyz`. ### Bool -The rust type `bool` may be serialized as `|b1`. +The rust type `bool` may be serialized as [`|b1`](TypeChar::Bool). ### Endianness @@ -49,8 +52,8 @@ must match the size of the rust type used. There are two type codes used by numpy for time and date. -* `m`: A `numpy.timedelta64`. -* `M`: A `numpy.datetime64`. +* [`m`](TypeChar::TimeDelta): A `numpy.timedelta64`. +* [`M`](TypeChar::DateTime): A `numpy.datetime64`. Both of these are represented as 8-byte signed integers, and therefore can use **`i64`** in rust. @@ -63,9 +66,9 @@ nanoseconds. There are three type codes for variable-sized strings of data found in npy files: -* `|VN`: A fixed-size array of `N` bytes. -* `|SN` (or `|aN`): A possibly-null-terminated sequence of bytes of length `<= N`. -* ` Option { type_str_num_bytes_as_usize(self) } - /// Returns true if each element in the array can vary in size. + /// Returns true if an array using this DType is stored using `pickle`. /// /// This is true if and only if the type character is [`'O'`](TypeChar::Object). - pub fn has_variable_size(&self) -> bool { matches!(self.type_char, TypeChar::Object) } + /// + /// Pickled arrays present unique challenges, and most of the `npyz` crate does not support them. + pub fn uses_pickled_array(&self) -> bool { matches!(self.type_char, TypeChar::Object) } } /// Represents the first character in a [`TypeStr`], which describes endianness. @@ -408,7 +410,7 @@ impl fmt::Display for TimeUnits { impl fmt::Display for TypeStr { fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { write!(f, "{}{}", self.endianness, self.type_char)?; - if !self.has_variable_size() { + if !self.uses_pickled_array() { write!(f, "{}", self.size)?; } if let Some(time_units) = self.time_units { diff --git a/test-data/object.npy b/test-data/object.npy deleted file mode 100644 index 2831eafd48019f9994811ebb6bd24358f48a749e..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 275 zcmbR27wQ`j$;eQ~P_3SlTAW;@Zl$1Js?sT~Xu&?A;tnp;q*7oVJ8l&Y6onp2XQSX7i)Ii-guz9=<0 zKd-o?s5H4`%H%0MtYDQ>df4+)AW9}r@n&e9;>?&drF}}!6b)}i%?#!q)|8UUf>e-t oCVz`5ogEM<<{rkHDSm!_UjKmrOnBRuOiAjDol-j`5y;R30QMqR?f?J) diff --git a/test-data/pickle.npy b/test-data/pickle.npy new file mode 100644 index 0000000000000000000000000000000000000000..30db4cd20e94966a84e968ec165c383f68108439 GIT binary patch literal 317 zcmbV`O-sW-5QcXXt-98)f3UZW5OUCqcu^24ba7h>LOe*>Y*rN7Bs19v!3KJ<#j`(H zSL-kEIlM6MJPh-x?(fDAl)RERk=xp!xa49n5-}G~B|6l_w8Y&0)B`=Mt?%n+U0FXz zXE8rNjd{oa4_k(&Xy#R$m=bL=Z)WaABkGn-(VDWT9X7@>ARCJn`DP1Ll6MGhXa?aL zwkEJ0Nh$*wuAuj)=B}+QgPk-Wgp4j=R}x9rur~n~$uOn9UBJFlYWU^(4SD6DpM5;S dfxmWp-S>(o?fLLeW)|VFgQE_Pm345ECEvV$R{H<| literal 0 HcmV?d00001 From 8e15397aeb42e963cd41110e422328a84525f9be Mon Sep 17 00:00:00 2001 From: Michael Lamparski Date: Sun, 12 Apr 2026 13:49:28 -0400 Subject: [PATCH 5/7] add ability to write header only --- CHANGELOG.md | 14 +++- examples/pickle.rs | 31 ++++++- src/type_matchup_docs.rs | 5 +- src/type_str.rs | 8 +- src/write.rs | 173 ++++++++++++++++++++++++++++++++------- 5 files changed, 194 insertions(+), 37 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f37759a..4961b88 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,9 +4,19 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). -## [Unreleased] +## [Unreleased] (major bump needed: 0.9.0) -No unreleased changes yet! + + +### Changed +- The `T: Serialize` bound has been removed from the `WriterBuilder` trait. This may introduce places in downstream code where a `T: npyz::Serialize` bound must now be explicitly added. + +### Added +- Limited support for pickled arrays. See `examples/pickle.rs`. Thanks, @paroga! + - `TypeChar::Object` (numpy's `"|O"`) + - `{TypeStr, DType NpyHeader}::uses_pickled_array()` boolean function. + - `NpyFile` can open files with pickled arrays, but it cannot extract items. + - `WriteOptions::new_header_only`, `WriterBuilder::write_header_only` makes it possible to write headers wwith object dtypes. ## [0.8.4] - 2025-05-14 diff --git a/examples/pickle.rs b/examples/pickle.rs index f3ce2f6..16ded83 100644 --- a/examples/pickle.rs +++ b/examples/pickle.rs @@ -13,6 +13,13 @@ use serde_pickle; // np.save('test-data/pickle.npy', a) fn main() -> io::Result<()> { + read_example()?; + write_example()?; + Ok(()) +} + +// Example of reading a pickled ndarray's header. +fn read_example() -> Result<(), io::Error> { let file = io::BufReader::new(File::open("test-data/pickle.npy")?); let npy = npyz::NpyFile::new(file)?; @@ -30,7 +37,29 @@ fn main() -> io::Result<()> { // The pickled value begins with a bunch of array metadata that's tough to make sense of, // but at the end you will see the data, a flat list of the objects written by the python // snippet at the top of this file. - println!("Pickled value: {:#?}", value); + println!("Pickled value: {:?}", value); + + Ok(()) +} + +// Example of writing a pickled ndarray's header. +fn write_example() -> Result<(), io::Error> { + use npyz::WriterBuilder; + + let type_str = "|O".parse().unwrap(); + let dtype = npyz::DType::new_scalar(type_str); + + let mut file = io::BufWriter::new(File::create("examples/output/pickle-corrupt.npy")?); + + // With npyz it is possible to write the *header*... + npyz::WriteOptions::new_header_only() + .dtype(dtype) + .shape(&[2, 3]) + .writer(&mut file) + .write_header_only()?; + + // And beyond that I have positively no idea what you should do to write the payload. + // Have fun! Ok(()) } diff --git a/src/type_matchup_docs.rs b/src/type_matchup_docs.rs index e1fc6ff..70592cf 100644 --- a/src/type_matchup_docs.rs +++ b/src/type_matchup_docs.rs @@ -211,7 +211,8 @@ Things that you *can* do with pickled arrays: * Parse and format [`DType`] strings containing `|O`. * Read and write [`NpyHeader`], leaving the file cursor at the beginning of the pickle. -* Open an [`NpyFile`]. + * Reading can be done by opening an [`NpyFile`]. + * Writing can be done via [`WriteOptions::new_header_only`] and [`WriterBuilder::write_header_only`]. * Get [`shape`](NpyHeader::shape), [`strides`](NpyHeader::strides), [`len`](NpyHeader::len), [`order`](NpyHeader::order), [`uses_pickled_array`](NpyHeader::uses_pickled_array) from an [`NpyHeader`] or [`NpyFile`]. @@ -223,7 +224,7 @@ demonstration of how you could read one of these arrays using [`serde_pickle`](h #[expect(unused)] // used by docstring use crate::{FixedSizeBytes, TypeStr, Deserialize, Serialize, AutoSerialize, DType, TypeChar}; #[expect(unused)] // used by docstring -use crate::{NpyFile, NpyHeader}; +use crate::{NpyFile, NpyHeader, WriteOptions, WriterBuilder}; #[cfg(feature = "arrayvec")] diff --git a/src/type_str.rs b/src/type_str.rs index 64b7c7c..32dcac5 100644 --- a/src/type_str.rs +++ b/src/type_str.rs @@ -53,8 +53,6 @@ impl TypeStr { /// Returns true if an array using this DType is stored using `pickle`. /// - /// This is true if and only if the type character is [`'O'`](TypeChar::Object). - /// /// Pickled arrays present unique challenges, and most of the `npyz` crate does not support them. pub fn uses_pickled_array(&self) -> bool { matches!(self.type_char, TypeChar::Object) } } @@ -179,9 +177,11 @@ pub enum TypeChar { /// See [`type_matchup_docs`][`crate::type_matchup_docs`] for information on which types can /// use this for serialization. UnicodeStr, - /// Code `O`. + /// Code `O`. Represents an arbitrary python object. /// - /// Notice that numpy uses this to store pickled values. + /// Arbitrary python objects are stored using the [`pickle` protocol](https://docs.python.org/3/library/pickle.html). + /// Most of the `npyz` crate does not support these. See [`type_matchup_docs`][`crate::type_matchup_docs`] + /// for more information. Object, /// Code `V`. Represents a binary blob of `size` bytes. /// diff --git a/src/write.rs b/src/write.rs index 9476c9a..ce66c60 100644 --- a/src/write.rs +++ b/src/write.rs @@ -3,6 +3,8 @@ use std::fs::File; use std::path::Path; use std::marker::PhantomData; +use std::convert::Infallible as Never; + use byteorder::{WriteBytesExt, LittleEndian}; use crate::serialize::{AutoSerialize, Serialize, TypeWrite}; @@ -54,6 +56,17 @@ pub mod write_options { }} } + impl WriteOptions { + /// Construct an almost empty Writer configuration, which will not be used to write any items. + /// + /// You will be able to call [`WriterBuilder::write_header_only`], but not [`WriterBuilder::begin_1d`] + /// or [`WriterBuilder::begin_nd`]. + pub fn new_header_only() -> Self { WriteOptions { + order: Order::C, + _marker: PhantomData, + }} + } + impl Default for WriteOptions { fn default() -> Self { Self::new() } } @@ -68,7 +81,7 @@ pub mod write_options { /// /// The majority of methods return a type that also implements [`WriterBuilder`]; they are meant /// to be chained together to construct a full config. - pub trait WriterBuilder: Sized { + pub trait WriterBuilder: Sized { /// Calls [`Self::dtype`] with the default dtype for the type to be serialized. /// /// **Calling this method or [`Self::dtype`] is required.** @@ -104,6 +117,7 @@ pub mod write_options { /// Begin writing an array of the previously supplied [`shape`][Self::shape]. fn begin_nd(self) -> io::Result::Writer>> where + T: Serialize, Self: HasDType + HasWriter + HasShape, ::Writer: Write, { @@ -124,6 +138,7 @@ pub mod write_options { /// validated against the number of elements written. This may change in the future. fn begin_1d(self) -> io::Result::Writer>> where + T: Serialize, Self: HasDType + HasWriter, ::Writer: Write + Seek, { @@ -134,6 +149,27 @@ pub mod write_options { _marker: PhantomData, }, MaybeSeek::new_seek(self.__into_writer())) } + + /// Write the header, and recover the writer, which now points to where the data should be written. + /// + /// (If the original writer object was passed in by reference, it too is guaranteed to point to this location + /// after calling this function, so you do not necessarily need to use the returned writer.) + /// + /// This is the only build method you can call if the WriteOptions were + /// created with [`WriteOptions::new_header_only`] + fn write_header_only(mut self) -> io::Result<::Writer> + where + Self: HasDType + HasWriter + HasShape, + ::Writer: Write, + { + let dtype = self.__get_dtype(); + let order = self.__get_order(); + let shape = self.__get_shape(); + + write_header(self.__writer_mut(), &dtype, order, Some(shape.as_slice()))?; + + Ok(self.__into_writer()) + } } /// Return type of [`WriterBuilder::writer`]. It represents a config with a known output stream. @@ -198,6 +234,8 @@ pub mod write_options { type Writer; #[doc(hidden)] fn __into_writer(self) -> Self::Writer; + #[doc(hidden)] + fn __writer_mut(&mut self) -> &mut Self::Writer; } // NOTE: This mainly exists to prevent the accidental usage of `.writer()` when working with NPZ files. @@ -217,22 +255,22 @@ pub mod write_options { /// so you do not need to add one. pub trait MissingWriter {} - impl WriterBuilder for WriteOptions { + impl WriterBuilder for WriteOptions { fn order(mut self, order: Order) -> Self { self.order = order; self } fn __get_order(&self) -> Order { self.order } } - impl> WriterBuilder for WithWriter { + impl> WriterBuilder for WithWriter { fn order(mut self, order: Order) -> Self { self.inner = self.inner.order(order); self } fn __get_order(&self) -> Order { self.inner.__get_order() } } - impl> WriterBuilder for WithDType { + impl> WriterBuilder for WithDType { fn order(mut self, order: Order) -> Self { self.inner = self.inner.order(order); self } fn __get_order(&self) -> Order { self.inner.__get_order() } } - impl> WriterBuilder for WithShape { + impl> WriterBuilder for WithShape { fn order(mut self, order: Order) -> Self { self.inner = self.inner.order(order); self } fn __get_order(&self) -> Order { self.inner.__get_order() } } @@ -248,6 +286,7 @@ pub mod write_options { impl HasWriter for WithWriter { type Writer = W; fn __into_writer(self) -> Self::Writer { self.writer } + fn __writer_mut(&mut self) -> &mut Self::Writer { &mut self.writer } } impl MissingWriter for WriteOptions {} @@ -273,6 +312,7 @@ pub mod write_options { impl<$($impl_generics)*> HasWriter for $Self where $inner: HasWriter { type Writer = $inner::Writer; fn __into_writer(self) -> Self::Writer { self.inner.__into_writer() } + fn __writer_mut(&mut self) -> &mut Self::Writer { self.inner.__writer_mut() } } }; (@single [$inner:ident] [$($impl_generics:tt)*] [$Self:ty] [MissingWriter]) => { @@ -372,29 +412,7 @@ impl NpyWriter { MaybeSeek::Isnt(_) => None, }; - if let DType::Array(..) = dtype { - panic!("the outermost dtype cannot be an array (got: {:?})", dtype); - } - - let (dict_text, shape_info) = create_dict(&dtype, order, shape.as_deref()); - let (header_text, version, version_props) = determine_required_version_and_pad_header(dict_text); - - fw.write_all(&[0x93u8])?; - fw.write_all(b"NUMPY")?; - fw.write_all(&[version.0, version.1])?; - - assert_eq!((header_text.len() + version_props.bytes_before_text()) % 16, 0); - match version_props.header_size_type { - HeaderSizeType::U16 => { - assert!(header_text.len() <= u16::MAX as usize); - fw.write_u16::(header_text.len() as u16)?; - }, - HeaderSizeType::U32 => { - assert!(header_text.len() <= u32::MAX as usize); - fw.write_u32::(header_text.len() as u32)?; - }, - } - fw.write_all(&header_text)?; + let (shape_info, version_props) = write_header(&mut fw, &dtype, order, shape.as_deref())?; let writer = match Row::writer(&dtype) { Ok(writer) => writer, @@ -460,6 +478,39 @@ impl NpyWriter { } } +fn write_header( + fw: &mut W, + dtype: &DType, + order: Order, + shape: Option<&[u64]>, +) -> io::Result<(ShapeInfo, VersionProps)> { + if let DType::Array(..) = dtype { + panic!("the outermost dtype cannot be an array (got: {:?})", dtype); + } + + let (dict_text, shape_info) = create_dict(dtype, order, shape); + let (header_text, version, version_props) = determine_required_version_and_pad_header(dict_text); + + fw.write_all(&[0x93u8])?; + fw.write_all(b"NUMPY")?; + fw.write_all(&[version.0, version.1])?; + + assert_eq!((header_text.len() + version_props.bytes_before_text()) % 16, 0); + match version_props.header_size_type { + HeaderSizeType::U16 => { + assert!(header_text.len() <= u16::MAX as usize); + fw.write_u16::(header_text.len() as u16)?; + }, + HeaderSizeType::U32 => { + assert!(header_text.len() <= u32::MAX as usize); + fw.write_u32::(header_text.len() as u32)?; + }, + } + fw.write_all(&header_text)?; + + Ok((shape_info, version_props)) +} + fn create_dict(dtype: &DType, order: Order, shape: Option<&[u64]>) -> (Vec, ShapeInfo) { let mut header: Vec = vec![]; header.extend(&b"{'descr': "[..]); @@ -686,6 +737,7 @@ mod tests { use super::*; use std::io::{self, Cursor}; use crate::NpyFile; + use crate::header::Field; fn bytestring_contains(haystack: &[u8], needle: &[u8]) -> bool { if needle.is_empty() { @@ -775,4 +827,69 @@ mod tests { Ok(()) } + + #[test] + fn write_header_only_positions() -> io::Result<()> { + let mut cursor = Cursor::new(vec![]); + + let dtype = DType::new_scalar("|O".parse().unwrap()); + + // bookend with garbage to make sure positions are correct. + cursor.write(b"ABCD")?; + let header_start_pos = cursor.position(); + + let returned_writer = WriteOptions::new_header_only() + .shape(&[2, 3]) + .dtype(dtype.clone()) + .writer(&mut cursor) + .write_header_only()?; + returned_writer.write(b"dcba")?; + + cursor.set_position(header_start_pos); + + // NpyFile should stop reading exactly where we stopped writing. + let npy = NpyFile::new(&mut cursor)?; + assert_eq!(npy.dtype(), dtype); + assert_eq!(npy.shape(), &[2, 3]); + assert_eq!(npy.len(), 6); + assert!(npy.uses_pickled_array()); + + let mut trailing_bytes = vec![]; + std::io::Read::read_to_end(&mut cursor, &mut trailing_bytes)?; + assert_eq!(&trailing_bytes[..], b"dcba"); + Ok(()) + } + + #[test] + fn write_header_only_funny_type() -> io::Result<()> { + let mut cursor = Cursor::new(vec![]); + + let dtype = DType::Record(vec![ + Field { + name: "parent".to_string(), + dtype: DType::Record(vec![ + Field { + name: "child".to_string(), + dtype: DType::Plain("|O".parse().unwrap()), + }, + ]), + } + ]); + + WriteOptions::new_header_only() + .shape(&[2, 3]) + .dtype(dtype.clone()) + .writer(&mut cursor) + .write_header_only()?; + + cursor.set_position(0); + + // We should get back the same complicated DType. + let npy = NpyFile::new(&mut cursor)?; + assert_eq!(npy.dtype(), dtype); + assert_eq!(npy.shape(), &[2, 3]); + assert_eq!(npy.len(), 6); + assert!(npy.uses_pickled_array()); + Ok(()) + } } From 78583b7f9f9090150247b29634fe116373eebf46 Mon Sep 17 00:00:00 2001 From: Michael Lamparski Date: Sun, 12 Apr 2026 14:10:18 -0400 Subject: [PATCH 6/7] clean up formatting in lines added to type_str --- src/type_str.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/type_str.rs b/src/type_str.rs index 32dcac5..57896a8 100644 --- a/src/type_str.rs +++ b/src/type_str.rs @@ -310,7 +310,7 @@ fn type_str_num_bytes_as_usize(type_str: &TypeStr) -> Option { TypeChar::UnicodeStr => size_field.checked_mul(4), - TypeChar::Object =>None, + TypeChar::Object => None, } } @@ -505,7 +505,7 @@ mod parse { _ => match size.parse() { Ok(size) => size, Err(e) => bail!(ParseIntError(e)), - } + }, }; let time_units = if remainder.is_empty() { From c6e464d98227d9b1fccc384616f6271d0ff94a22 Mon Sep 17 00:00:00 2001 From: Michael Lamparski Date: Sun, 12 Apr 2026 14:15:39 -0400 Subject: [PATCH 7/7] Fix markdown bug --- src/type_matchup_docs.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/type_matchup_docs.rs b/src/type_matchup_docs.rs index 70592cf..c78adb6 100644 --- a/src/type_matchup_docs.rs +++ b/src/type_matchup_docs.rs @@ -216,7 +216,7 @@ Things that you *can* do with pickled arrays: * Get [`shape`](NpyHeader::shape), [`strides`](NpyHeader::strides), [`len`](NpyHeader::len), [`order`](NpyHeader::order), [`uses_pickled_array`](NpyHeader::uses_pickled_array) from an [`NpyHeader`] or [`NpyFile`]. -**See [`examples/pickle.rs`](https://github.com/ExpHP/npyz/blob/master/examples/pickle.rs) for a +**See [`examples/pickle.rs`](https://github.com/ExpHP/npyz/blob/master/examples/pickle.rs)** for a demonstration of how you could read one of these arrays using [`serde_pickle`](https://docs.rs/serde-pickle/latest/serde_pickle/). **/