Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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" }

Expand Down
23 changes: 23 additions & 0 deletions examples/object.rs
Original file line number Diff line number Diff line change
@@ -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<dyn std::error::Error>> {
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(())
}
21 changes: 15 additions & 6 deletions src/header.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@

use std::io;
use core::fmt;
use core::num::TryFromIntError;

use py_literal::ParseError;
pub use py_literal::Value;
Expand Down Expand Up @@ -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<usize> {
pub fn num_bytes(&self) -> Result<Option<usize>, 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),
})
}
}
}
}
Expand Down
22 changes: 15 additions & 7 deletions src/read.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<usize>,
}

impl NpyHeader {
Expand Down Expand Up @@ -336,9 +336,9 @@ impl NpyHeader {

fn from_parts(dtype: DType, shape: Vec<u64>, order: Order) -> io::Result<NpyHeader> {
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 })
}
Expand Down Expand Up @@ -390,7 +390,11 @@ impl<R: io::Read, T: Deserialize> NpyReader<T, R> 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(())
Expand All @@ -416,7 +420,10 @@ impl<'a, T: Deserialize> NpyData<'a, T> {
pub fn from_bytes(bytes: &'a [u8]) -> io::Result<NpyData<'a, T>> {
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 })
}

Expand Down Expand Up @@ -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()
}

Expand Down
32 changes: 25 additions & 7 deletions src/type_str.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
}
Expand All @@ -207,6 +212,7 @@ impl TypeChar {
TypeChar::DateTime => "M",
TypeChar::ByteStr => "S",
TypeChar::UnicodeStr => "U",
TypeChar::Object => "O",
TypeChar::RawData => "V",
}
}
Expand Down Expand Up @@ -235,6 +241,7 @@ impl TypeChar {
// changes them to `|S1` and `|U1`.)
TypeChar::ByteStr |
TypeChar::UnicodeStr |
TypeChar::Object |
TypeChar::RawData => None,
}
}
Expand All @@ -253,6 +260,7 @@ impl TypeChar {
TypeChar::UnicodeStr => true,

TypeChar::ByteStr |
TypeChar::Object |
TypeChar::RawData => false,
}
}
Expand Down Expand Up @@ -292,6 +300,8 @@ fn type_str_num_bytes_as_usize(type_str: &TypeStr) -> Option<usize> {
TypeChar::RawData => Some(size_field),

TypeChar::UnicodeStr => size_field.checked_mul(4),

TypeChar::Object =>None,
}
}

Expand Down Expand Up @@ -453,7 +463,7 @@ mod parse {
fn from_str(input: &str) -> Result<Self, ParseTypeStrError> {
use self::ErrorKind::*;

if input.len() < 3 {
if input.len() < 2 {
bail!(SyntaxError);
}

Expand All @@ -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() {
Expand Down Expand Up @@ -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);
Expand Down
Binary file added test-data/object.npy
Binary file not shown.
Loading