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/Cargo.toml b/Cargo.toml index 286cd5f..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 @@ -50,6 +51,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" } @@ -57,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/pickle.rs b/examples/pickle.rs new file mode 100644 index 0000000..16ded83 --- /dev/null +++ b/examples/pickle.rs @@ -0,0 +1,72 @@ +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<()> { + 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)?; + + 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(()) +} + +// 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(()) +} + +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 662b493..5a88b1a 100644 --- a/src/header.rs +++ b/src/header.rs @@ -112,7 +112,8 @@ 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`. + /// 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) => ty.num_bytes(), @@ -123,6 +124,20 @@ impl DType { }, } } + + /// `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.uses_pickled_array(), + DType::Array(_, inner) => inner.uses_pickled_array(), + DType::Record(fields) => fields.iter().any(|field| field.dtype.uses_pickled_array()), + } + } } fn convert_list_to_record_fields(values: &[Value]) -> io::Result> { @@ -374,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(), @@ -384,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(), @@ -393,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![ @@ -472,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', ', + 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,11 +367,16 @@ 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 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), + 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 }) + Ok(NpyHeader { dtype, shape, strides, order, n_records, item_size, uses_pickled_array }) } } @@ -384,13 +420,20 @@ 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(); assert!(index <= len, "index out of bounds for seeking (the index is {} but the len is {})", index, len); 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::new(io::ErrorKind::Unsupported, NO_SEEK_MSG))?; + reader.seek(io::SeekFrom::Current(delta * item_size as i64))?; *current_index = index; } Ok(()) @@ -413,10 +456,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 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 }) } @@ -437,11 +493,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)) @@ -454,10 +512,13 @@ 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_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/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..c78adb6 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`. -* ` 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 an array using this DType is stored using `pickle`. + /// + /// 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. @@ -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, } @@ -170,6 +177,12 @@ pub enum TypeChar { /// See [`type_matchup_docs`][`crate::type_matchup_docs`] for information on which types can /// use this for serialization. UnicodeStr, + /// Code `O`. Represents an arbitrary python object. + /// + /// 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. /// /// Can use [`crate::FixedSizeBytes`] for serialization, or some other types; see @@ -190,6 +203,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 +221,7 @@ impl TypeChar { TypeChar::DateTime => "M", TypeChar::ByteStr => "S", TypeChar::UnicodeStr => "U", + TypeChar::Object => "O", TypeChar::RawData => "V", } } @@ -235,6 +250,7 @@ impl TypeChar { // changes them to `|S1` and `|U1`.) TypeChar::ByteStr | TypeChar::UnicodeStr | + TypeChar::Object | TypeChar::RawData => None, } } @@ -253,6 +269,7 @@ impl TypeChar { TypeChar::UnicodeStr => true, TypeChar::ByteStr | + TypeChar::Object | TypeChar::RawData => false, } } @@ -292,6 +309,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, } } @@ -390,7 +409,10 @@ 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, self.size)?; + write!(f, "{}{}", self.endianness, self.type_char)?; + if !self.uses_pickled_array() { + write!(f, "{}", self.size)?; + } if let Some(time_units) = self.time_units { write!(f, "[{}]", time_units)?; } @@ -453,7 +475,7 @@ mod parse { fn from_str(input: &str) -> Result { use self::ErrorKind::*; - if input.len() < 3 { + if input.len() < 2 { bail!(SyntaxError); } @@ -476,13 +498,14 @@ 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 = match (type_char, size) { + (TypeChar::Object, "") => 0, + (TypeChar::Object, _) | (_, "") => bail!(SyntaxError), + _ => match size.parse() { + Ok(size) => size, + Err(e) => bail!(ParseIntError(e)), + }, }; let time_units = if remainder.is_empty() { @@ -572,9 +595,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); @@ -646,6 +671,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, @@ -679,6 +714,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..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}; @@ -38,6 +40,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, @@ -52,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() } } @@ -66,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.** @@ -102,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, { @@ -122,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, { @@ -132,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. @@ -196,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. @@ -215,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() } } @@ -246,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 {} @@ -271,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]) => { @@ -370,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, @@ -458,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': "[..]); @@ -684,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() { @@ -773,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(()) + } } diff --git a/test-data/pickle.npy b/test-data/pickle.npy new file mode 100644 index 0000000..30db4cd Binary files /dev/null and b/test-data/pickle.npy differ