From 2fb00f5cc47b49f925b2a9c93a61b9e09451024d Mon Sep 17 00:00:00 2001 From: Patrick Gansterer Date: Fri, 10 Apr 2026 06:16:02 +0200 Subject: [PATCH] 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