use std::io;
use crate::beve::header;
use crate::beve::impls::{Block, NumericBytes};
use crate::error::{Error, ErrorCode, StreamError, StreamResult};
const CHUNK: usize = 1 << 20;
pub fn read_array_into<T, R>(out: &mut Vec<T>, reader: R) -> StreamResult<()>
where
T: NumericBytes,
R: io::Read,
{
let mut src = Source { reader, pos: 0 };
out.clear();
let result = src.array_head::<T>().and_then(|n| {
src.payload(out, n, (CHUNK / Block::<T>::WIDTH).max(1))?;
src.finish()
});
if result.is_err() {
out.clear();
}
result
}
pub fn from_reader_array<T, R>(reader: R) -> StreamResult<Vec<T>>
where
T: NumericBytes,
R: io::Read,
{
let mut out = Vec::new();
read_array_into(&mut out, reader)?;
Ok(out)
}
struct Source<R> {
reader: R,
pos: usize,
}
impl<R: io::Read> Source<R> {
fn fail(&self, at: usize, code: ErrorCode) -> StreamError {
StreamError::Parse(Error::new(code, at))
}
fn exact(&mut self, buf: &mut [u8]) -> StreamResult<()> {
match self.reader.read_exact(buf) {
Ok(()) => {
self.pos += buf.len();
Ok(())
}
Err(e) if e.kind() == io::ErrorKind::UnexpectedEof => {
Err(self.fail(self.pos, ErrorCode::UnexpectedEnd))
}
Err(e) => Err(StreamError::Io(e)),
}
}
fn byte(&mut self) -> StreamResult<u8> {
let mut b = [0u8; 1];
self.exact(&mut b)?;
Ok(b[0])
}
fn size(&mut self) -> StreamResult<u64> {
let b0 = self.byte()?;
let mut rest = [0u8; 7];
let rest = &mut rest[..header::size_extra(b0)];
self.exact(rest)?;
let mut v = u64::from(b0 >> 2);
for (k, &b) in rest.iter().enumerate() {
v |= u64::from(b) << (6 + 8 * k);
}
Ok(v)
}
fn count(&mut self) -> StreamResult<usize> {
let at = self.pos;
let n = self.size()?;
usize::try_from(n).map_err(|_| self.fail(at, ErrorCode::UnexpectedEnd))
}
fn array_head<T: NumericBytes>(&mut self) -> StreamResult<usize> {
let start = self.pos;
let h = self.byte()?;
if header::ty(h) != header::TY_TYPED_ARRAY {
return Err(self.fail(start, ErrorCode::ExpectedArray));
}
let (elem, n) = match (header::sub(h), header::count(h)) {
(header::CAT_OTHER, header::OTHER_BOOL | header::OTHER_STRING) => {
return Err(self.fail(start, ErrorCode::ElementTypeMismatch));
}
(header::CAT_OTHER, header::OTHER_ALIGNED) => {
let at = self.pos;
let inner = self.byte()?;
if header::ty(inner) != header::TY_TYPED_ARRAY
|| header::byte_width(header::sub(inner), header::count(inner)).is_none()
{
return Err(self.fail(at, ErrorCode::InvalidHeader));
}
let n = self.count()?;
let pad = usize::from(self.byte()?);
let mut skip = [0u8; 255];
self.exact(&mut skip[..pad])?;
(header::element_of(inner), n)
}
(header::CAT_OTHER, _) => return Err(self.fail(start, ErrorCode::InvalidHeader)),
(cat, count) if header::byte_width(cat, count).is_none() => {
return Err(self.fail(start, ErrorCode::InvalidHeader));
}
_ => (header::element_of(h), self.count()?),
};
if elem != T::ELEMENT {
return Err(self.fail(start, ErrorCode::ElementTypeMismatch));
}
Ok(n)
}
fn payload<T: NumericBytes>(
&mut self,
out: &mut Vec<T>,
n: usize,
per: usize,
) -> StreamResult<()> {
let width = Block::<T>::WIDTH;
while out.len() < n {
let base = out.len();
let end = n.min(base + per);
if out.capacity() < end {
let want = n.min((out.capacity() * 2).max(end));
out.reserve_exact(want - base);
}
unsafe {
core::ptr::write_bytes(out.as_mut_ptr().add(base), 0u8, end - base);
out.set_len(end);
}
let dst = unsafe {
core::slice::from_raw_parts_mut(
out.as_mut_ptr().add(base).cast::<u8>(),
(end - base) * width,
)
};
self.exact(dst)?;
if cfg!(target_endian = "big") && width > 1 {
for e in dst.chunks_exact_mut(width) {
e.reverse();
}
}
}
Ok(())
}
fn finish(&mut self) -> StreamResult<()> {
let mut b = [0u8; 1];
match self.reader.read_exact(&mut b) {
Err(e) if e.kind() == io::ErrorKind::UnexpectedEof => Ok(()),
Ok(()) => Err(self.fail(self.pos, ErrorCode::TrailingContent)),
Err(e) => Err(StreamError::Io(e)),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn block(values: &[f64]) -> Vec<u8> {
values.iter().flat_map(|v| v.to_le_bytes()).collect()
}
#[test]
fn the_chunk_seam_and_the_reserve_it_rests_on() {
for n in 0..40usize {
let values: Vec<f64> = (0..n).map(|i| i as f64).collect();
let bytes = block(&values);
let mut src = Source {
reader: &bytes[..],
pos: 0,
};
let mut out: Vec<f64> = Vec::new();
src.payload(&mut out, n, 3).unwrap();
assert_eq!(out, values, "n = {n}");
assert_eq!(out.capacity(), n, "n = {n}");
}
}
#[test]
fn a_short_payload_fails_in_whichever_chunk_it_stops_in() {
let values: Vec<f64> = (0..20).map(|i| i as f64).collect();
let bytes = block(&values);
for cut in 0..bytes.len() {
let mut src = Source {
reader: &bytes[..cut],
pos: 0,
};
let mut out: Vec<f64> = Vec::new();
let err = src.payload(&mut out, 20, 3).unwrap_err();
assert_eq!(
err.as_parse().unwrap().code,
ErrorCode::UnexpectedEnd,
"cut at {cut}"
);
}
}
#[test]
fn a_lying_count_reserves_on_what_arrived() {
let values: Vec<f64> = (0..10).map(|i| i as f64).collect();
let bytes = block(&values);
let mut src = Source {
reader: &bytes[..],
pos: 0,
};
let mut out: Vec<f64> = Vec::new();
assert!(src.payload(&mut out, usize::MAX, 3).is_err());
assert!(out.capacity() <= 23, "capacity {}", out.capacity());
}
}