Skip to content
Merged
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
209 changes: 196 additions & 13 deletions src/write.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,12 +13,18 @@ use crate::read::Order;

// Long enough to accomodate a large integer followed by ",), }".
// Used when no shape is provided.
const FILLER_FOR_UNKNOWN_SIZE: &'static [u8] = &[b'*'; 19];
const FILLER_FOR_UNKNOWN_SIZE: &'static [u8] = &[b'*'; 20];

enum ShapeHint {
NDimensional(Vec<u64>),
TwoDimensional(u64),
OneDimensional,
}

struct DataFromBuilder<T: ?Sized> {
order: Order,
dtype: DType,
shape: Option<Vec<u64>>,
shape: ShapeHint,
_marker: PhantomData<fn(&T)>, // contravariant
}

Expand Down Expand Up @@ -124,7 +130,7 @@ pub mod write_options {
NpyWriter::_begin(DataFromBuilder {
dtype: self.__get_dtype(),
order: self.__get_order(),
shape: Some(self.__get_shape()),
shape: ShapeHint::NDimensional(self.__get_shape()),
_marker: PhantomData,
}, MaybeSeek::Isnt(self.__into_writer()))
}
Expand All @@ -145,7 +151,28 @@ pub mod write_options {
NpyWriter::_begin(DataFromBuilder {
dtype: self.__get_dtype(),
order: self.__get_order(),
shape: None,
shape: ShapeHint::OneDimensional,
_marker: PhantomData,
}, MaybeSeek::new_seek(self.__into_writer()))
}

/// Begin writing a 2d array, of length to be inferred from the number of elements written.
///
/// Notice that, in contrast to [`Self::begin_nd`], this method requires [`Seek`]. If you have
/// a `Vec<u8>`, you can wrap it in an [`io::Cursor`] to satisfy this requirement.
///
/// **Note:** At present, any [`shape`][Self::shape] you *did* happen to provide will be ignored and not
/// validated against the number of elements written. This may change in the future.
fn begin_2d(self, record_len: u64) -> io::Result<NpyWriter<T, <Self as HasWriter>::Writer>>
where
T: Serialize,
Self: HasDType + HasWriter,
<Self as HasWriter>::Writer: Write + Seek,
{
NpyWriter::_begin(DataFromBuilder {
dtype: self.__get_dtype(),
order: self.__get_order(),
shape: ShapeHint::TwoDimensional(record_len),
_marker: PhantomData,
}, MaybeSeek::new_seek(self.__into_writer()))
}
Expand All @@ -166,7 +193,7 @@ pub mod write_options {
let order = self.__get_order();
let shape = self.__get_shape();

write_header(self.__writer_mut(), &dtype, order, Some(shape.as_slice()))?;
write_header(self.__writer_mut(), &dtype, order, &ShapeHint::NDimensional(shape))?;

Ok(self.__into_writer())
}
Expand Down Expand Up @@ -369,7 +396,10 @@ pub struct NpyWriter<Row: Serialize + ?Sized, W: Write> {

enum ShapeInfo {
// No shape was written; we'll return to write a 1D shape on `finish()`.
Automatic { offset_in_header_text: u64 },
Automatic1D { offset_in_header_text: u64 },
// Only partial shape was written; we'll return to write the missing part of shape on `finish()`.
// We need to keep track of the `Order` to be able to writer the missing part correctly.
Automatic2D { offset_in_header_text: u64, record_len: u64, order: Order },
// The complete shape has already been written.
// Raise an error on `finish()` if the wrong number of elements is given.
Known { expected_num_items: u64 },
Expand Down Expand Up @@ -412,7 +442,7 @@ impl<Row: Serialize + ?Sized , W: Write> NpyWriter<Row, W> {
MaybeSeek::Isnt(_) => None,
};

let (shape_info, version_props) = write_header(&mut fw, &dtype, order, shape.as_deref())?;
let (shape_info, version_props) = write_header(&mut fw, &dtype, order, &shape)?;

let writer = match Row::writer(&dtype) {
Ok(writer) => writer,
Expand Down Expand Up @@ -449,7 +479,7 @@ impl<Row: Serialize + ?Sized , W: Write> NpyWriter<Row, W> {
}));
}
},
ShapeInfo::Automatic { offset_in_header_text } => {
ShapeInfo::Automatic1D { offset_in_header_text } => {
// Write the size to the header
let shape_pos = self.start_pos.unwrap() + self.version_props.bytes_before_text() as u64 + offset_in_header_text;
let end_pos = self.fw.seek(SeekFrom::Current(0))?;
Expand All @@ -461,6 +491,32 @@ impl<Row: Serialize + ?Sized , W: Write> NpyWriter<Row, W> {
self.fw.write_all(&::std::iter::repeat(b' ').take(FILLER_FOR_UNKNOWN_SIZE.len() - length.len()).collect::<Vec<_>>())?;
self.fw.seek(SeekFrom::Start(end_pos))?;
},
ShapeInfo::Automatic2D { offset_in_header_text, record_len, order } => {
if self.num_items % record_len != 0 {
return Err(io::Error::new(io::ErrorKind::InvalidData, {
format!("{} item(s) is not divisible by {}!", self.num_items, record_len)
}));
}

// Write the size to the header
let shape_pos = self.start_pos.unwrap() + self.version_props.bytes_before_text() as u64 + offset_in_header_text;
let end_pos = self.fw.seek(SeekFrom::Current(0))?;

self.fw.seek(SeekFrom::Start(shape_pos))?;
let length = format!("{}", self.num_items / record_len);

match order {
Order::C => {
write!(self.fw, "{}, {}), }}", length, record_len)?;
},
Order::Fortran => {
write!(self.fw, "{}, {}), }}", record_len, length)?;
}
};

self.fw.write_all(&::std::iter::repeat(b' ').take(FILLER_FOR_UNKNOWN_SIZE.len() - length.len()).collect::<Vec<_>>())?;
self.fw.seek(SeekFrom::Start(end_pos))?;
}
}
self.fw.flush()?;
Ok(())
Expand All @@ -482,7 +538,7 @@ fn write_header<W: Write>(
fw: &mut W,
dtype: &DType,
order: Order,
shape: Option<&[u64]>,
shape: &ShapeHint,
) -> io::Result<(ShapeInfo, VersionProps)> {
if let DType::Array(..) = dtype {
panic!("the outermost dtype cannot be an array (got: {:?})", dtype);
Expand Down Expand Up @@ -511,7 +567,7 @@ fn write_header<W: Write>(
Ok((shape_info, version_props))
}

fn create_dict(dtype: &DType, order: Order, shape: Option<&[u64]>) -> (Vec<u8>, ShapeInfo) {
fn create_dict(dtype: &DType, order: Order, shape: &ShapeHint) -> (Vec<u8>, ShapeInfo) {
let mut header: Vec<u8> = vec![];
header.extend(&b"{'descr': "[..]);
header.extend(dtype.descr().as_bytes());
Expand All @@ -522,18 +578,35 @@ fn create_dict(dtype: &DType, order: Order, shape: Option<&[u64]>) -> (Vec<u8>,
}
header.extend(&b", 'shape': ("[..]);
let shape_info = match shape {
Some(shape) => {
ShapeHint::NDimensional(shape) => {
for x in shape {
write!(header, "{}, ", x).unwrap();
}
header.extend(&b"), }"[..]);
ShapeInfo::Known { expected_num_items: shape.iter().product() }
},
None => {
ShapeHint::TwoDimensional(record_len) => {
let shape_offset = header.len() as u64;

match order {
Order::C => {
header.extend(FILLER_FOR_UNKNOWN_SIZE);
write!(header, ", {}), }}", record_len).unwrap();
},
Order::Fortran => {
write!(header, "{}, ", record_len).unwrap();
header.extend(FILLER_FOR_UNKNOWN_SIZE);
header.extend(b"), }");
}
};

ShapeInfo::Automatic2D { offset_in_header_text: shape_offset, record_len: *record_len, order }
},
ShapeHint::OneDimensional => {
let shape_offset = header.len() as u64;
header.extend(FILLER_FOR_UNKNOWN_SIZE);
header.extend(&b",), }"[..]);
ShapeInfo::Automatic { offset_in_header_text: shape_offset }
ShapeInfo::Automatic1D { offset_in_header_text: shape_offset }
},
};
(header, shape_info)
Expand Down Expand Up @@ -713,6 +786,21 @@ pub(crate) fn to_writer_1d<W: io::Write + io::Seek, T: AutoSerialize>(writer: W,
to_writer_1d_with_seeking(writer, data)
}

/// Quick API for writing a 2D array to a vector of bytes.
#[cfg(test)]
pub(crate) fn to_bytes_2d<T: AutoSerialize>(order: Order, record_len: u64, data: &[T]) -> io::Result<Vec<u8>> {
let mut cursor = io::Cursor::new(vec![]);
to_writer_2d(&mut cursor, order, record_len, data)?;
Ok(cursor.into_inner())
}

/// Quick API for writing a 2D array to an io::Write.
#[cfg(test)]
pub(crate) fn to_writer_2d<W: io::Write + io::Seek, T: AutoSerialize>(writer: W, order: Order, record_len: u64, data: &[T]) -> io::Result<()> {
// we might change this later and/or remove the Seek bound from the current function, but for now this will do
to_writer_2d_with_seeking(writer, order, record_len, data)
}

/// Quick API for writing an n-d array to an io::Write.
#[cfg(test)]
pub(crate) fn to_writer_nd<W: io::Write, T: AutoSerialize>(writer: W, data: &[T], shape: &[u64]) -> io::Result<()> {
Expand All @@ -732,6 +820,17 @@ pub(crate) fn to_writer_1d_with_seeking<W: io::Write + io::Seek, T: AutoSerializ
writer.finish()
}

/// Quick API for writing a 2D array to an io::Write in a manner which makes use of io::Seek.
///
/// (tests will use this instead of 'to_writer_2d' if their purpose is to test the correctness of seek behavior,
/// so that changing 'to_writer_2d' to be Seek-less won't affect these tests)
#[cfg(test)]
pub(crate) fn to_writer_2d_with_seeking<W: io::Write + io::Seek, T: AutoSerialize>(writer: W, order: Order, record_len: u64, data: &[T]) -> io::Result<()> {
let mut writer = WriteOptions::new().default_dtype().order(order).writer(writer).begin_2d(record_len)?;
writer.extend(data)?;
writer.finish()
}

#[cfg(test)]
mod tests {
use super::*;
Expand Down Expand Up @@ -781,6 +880,72 @@ mod tests {
Ok(())
}

#[test]
fn write_2d_simple() -> io::Result<()> {
// 3 columns
let raw_buffer = to_bytes_2d(Order::C, 3, &[1.0, 3.0, 5.0, -4.0, 3.0, 2.0])?;

let reader = NpyFile::new(&raw_buffer[..])?;
assert_eq!(reader.shape(), &[2, 3][..]);
assert_eq!(reader.into_vec::<f64>()?, vec![1.0, 3.0, 5.0, -4.0, 3.0, 2.0]);

// same but with 2 columns
let raw_buffer = to_bytes_2d(Order::C, 2, &[1.0, 3.0, 5.0, -4.0, 3.0, 2.0])?;

let reader = NpyFile::new(&raw_buffer[..])?;
assert_eq!(reader.shape(), &[3, 2][..]);
assert_eq!(reader.into_vec::<f64>()?, vec![1.0, 3.0, 5.0, -4.0, 3.0, 2.0]);

Ok(())
}

#[test]
fn write_2d_simple_fortran_order() -> io::Result<()> {
// 3 columns
let raw_buffer = to_bytes_2d(Order::Fortran, 3, &[1.0, 3.0, 5.0, -4.0, 3.0, 2.0])?;

let reader = NpyFile::new(&raw_buffer[..])?;
assert_eq!(reader.order(), Order::Fortran);
assert_eq!(reader.shape(), &[3, 2][..]);
assert_eq!(reader.into_vec::<f64>()?, vec![1.0, 3.0, 5.0, -4.0, 3.0, 2.0]);

// same but with 2 columns
let raw_buffer = to_bytes_2d(Order::Fortran, 2, &[1.0, 3.0, 5.0, -4.0, 3.0, 2.0])?;

let reader = NpyFile::new(&raw_buffer[..])?;
assert_eq!(reader.order(), Order::Fortran);
assert_eq!(reader.shape(), &[2, 3][..]);
assert_eq!(reader.into_vec::<f64>()?, vec![1.0, 3.0, 5.0, -4.0, 3.0, 2.0]);

Ok(())
}

#[test]
fn write_2d_in_the_middle() -> io::Result<()> {
let mut cursor = Cursor::new(vec![]);

let prefix = b"lorem ipsum dolor sit amet.";
let suffix = b"and they lived happily ever after.";

// write to the cursor both before and after writing the file
cursor.write_all(prefix)?;
to_writer_2d_with_seeking(&mut cursor, Order::C, 3, &[1.0, 3.0, 5.0, 6.0, -1.0, 4.0])?;
cursor.write_all(suffix)?;

// check that the seeking did not interfere with our extra writes
let raw_buffer = cursor.into_inner();
assert!(raw_buffer.starts_with(prefix));
assert!(raw_buffer.ends_with(suffix));

// check the bytes written by `OutFile`
let written_bytes = &raw_buffer[prefix.len()..raw_buffer.len() - suffix.len()];
let reader = NpyFile::new(&written_bytes[..])?;
assert_eq!(reader.shape(), &[2, 3][..]);
assert_eq!(reader.into_vec::<f64>()?, vec![1.0, 3.0, 5.0, 6.0, -1.0, 4.0]);

Ok(())
}

#[test]
fn implicit_finish() -> io::Result<()> {
let mut cursor = Cursor::new(vec![]);
Expand Down Expand Up @@ -828,6 +993,24 @@ mod tests {
Ok(())
}

#[test]
fn write_2d_wrong_len() -> io::Result<()> {
let try_writing = |elems: &[i32]| -> io::Result<()> {
let mut buf = io::Cursor::new(vec![]);
let mut writer = WriteOptions::new().default_dtype().writer(&mut buf).begin_2d(3)?;
for &x in elems {
writer.push(&x)?;
}
writer.finish()?;
Ok(())
};
assert!(try_writing(&[00, 01, 02, 10, 11]).is_err());
assert!(try_writing(&[00, 01, 02, 10, 11, 12]).is_ok());
assert!(try_writing(&[00, 01, 02, 10, 11, 12, 20]).is_err());

Ok(())
}

#[test]
fn write_header_only_positions() -> io::Result<()> {
let mut cursor = Cursor::new(vec![]);
Expand Down
Loading