diff --git a/src/write.rs b/src/write.rs index ce66c60..65c8c64 100644 --- a/src/write.rs +++ b/src/write.rs @@ -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), + TwoDimensional(u64), + OneDimensional, +} struct DataFromBuilder { order: Order, dtype: DType, - shape: Option>, + shape: ShapeHint, _marker: PhantomData, // contravariant } @@ -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())) } @@ -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`, 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::Writer>> + where + T: Serialize, + Self: HasDType + 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())) } @@ -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()) } @@ -369,7 +396,10 @@ pub struct NpyWriter { 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 }, @@ -412,7 +442,7 @@ impl NpyWriter { 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, @@ -449,7 +479,7 @@ impl NpyWriter { })); } }, - 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))?; @@ -461,6 +491,32 @@ impl NpyWriter { self.fw.write_all(&::std::iter::repeat(b' ').take(FILLER_FOR_UNKNOWN_SIZE.len() - length.len()).collect::>())?; 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::>())?; + self.fw.seek(SeekFrom::Start(end_pos))?; + } } self.fw.flush()?; Ok(()) @@ -482,7 +538,7 @@ fn write_header( 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); @@ -511,7 +567,7 @@ fn write_header( Ok((shape_info, version_props)) } -fn create_dict(dtype: &DType, order: Order, shape: Option<&[u64]>) -> (Vec, ShapeInfo) { +fn create_dict(dtype: &DType, order: Order, shape: &ShapeHint) -> (Vec, ShapeInfo) { let mut header: Vec = vec![]; header.extend(&b"{'descr': "[..]); header.extend(dtype.descr().as_bytes()); @@ -522,18 +578,35 @@ fn create_dict(dtype: &DType, order: Order, shape: Option<&[u64]>) -> (Vec, } 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) @@ -713,6 +786,21 @@ pub(crate) fn to_writer_1d(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(order: Order, record_len: u64, data: &[T]) -> io::Result> { + 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(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(writer: W, data: &[T], shape: &[u64]) -> io::Result<()> { @@ -732,6 +820,17 @@ pub(crate) fn to_writer_1d_with_seeking(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::*; @@ -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::()?, 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::()?, 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::()?, 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::()?, 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::()?, 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![]); @@ -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![]);