use std::borrow::Cow;
use crate::error::SeaboredDeError;
pub type ReadResult<'data, T> = Result<T, SeaboredDeError<'data>>;
pub trait Read<'data> {
fn peek_byte(&mut self) -> ReadResult<'data, u8>;
fn advance(&mut self, n: usize) -> ReadResult<'data, ()>;
fn read_byte(&mut self) -> ReadResult<'data, u8>;
fn read_slice<'a>(&'a mut self, len: usize) -> ReadResult<'data, Cow<'data, [u8]>>;
fn read_array<const N: usize>(&mut self) -> ReadResult<'data, Cow<'data, [u8; N]>>;
#[inline]
fn read_be_u16(&mut self) -> ReadResult<'data, u16> {
self.read_array().map(|arr| u16::from_be_bytes(*arr))
}
#[inline]
fn read_be_u32(&mut self) -> ReadResult<'data, u32> {
self.read_array().map(|arr| u32::from_be_bytes(*arr))
}
#[inline]
fn read_be_u64(&mut self) -> ReadResult<'data, u64> {
self.read_array().map(|arr| u64::from_be_bytes(*arr))
}
}
impl<'data, R: Read<'data> + ?Sized> Read<'data> for &mut R {
#[inline(always)]
fn peek_byte(&mut self) -> ReadResult<'data, u8> {
(**self).peek_byte()
}
#[inline(always)]
fn advance(&mut self, n: usize) -> ReadResult<'data, ()> {
(**self).advance(n)
}
#[inline(always)]
fn read_byte(&mut self) -> ReadResult<'data, u8> {
(**self).read_byte()
}
#[inline(always)]
fn read_slice<'a>(&'a mut self, len: usize) -> ReadResult<'data, Cow<'data, [u8]>> {
(**self).read_slice(len)
}
#[inline(always)]
fn read_array<const N: usize>(&mut self) -> ReadResult<'data, Cow<'data, [u8; N]>> {
(**self).read_array()
}
}
impl<'data, R: Read<'data> + ?Sized> Read<'data> for Box<R> {
#[inline(always)]
fn peek_byte(&mut self) -> ReadResult<'data, u8> {
(**self).peek_byte()
}
#[inline(always)]
fn advance(&mut self, n: usize) -> ReadResult<'data, ()> {
(**self).advance(n)
}
#[inline(always)]
fn read_byte(&mut self) -> ReadResult<'data, u8> {
(**self).read_byte()
}
#[inline(always)]
fn read_slice<'a>(&'a mut self, len: usize) -> ReadResult<'data, Cow<'data, [u8]>> {
(**self).read_slice(len)
}
#[inline(always)]
fn read_array<const N: usize>(&mut self) -> ReadResult<'data, Cow<'data, [u8; N]>> {
(**self).read_array()
}
}
#[derive(Debug)]
pub struct DepthAwareReader<'data, R: Read<'data>> {
reader: R,
limit: usize,
depth: usize,
_marker: std::marker::PhantomData<&'data ()>,
}
impl<'data, R: Read<'data>> DepthAwareReader<'data, R> {
pub const DEFAULT_LIMIT: usize = 256;
#[inline(always)]
pub fn from_reader(reader: R) -> Self {
Self {
reader,
limit: Self::DEFAULT_LIMIT,
depth: 0,
_marker: Default::default(),
}
}
#[inline(always)]
pub fn from_reader_with_limit(reader: R, limit: usize) -> Self {
Self {
reader,
limit,
depth: 0,
_marker: Default::default(),
}
}
#[inline(always)]
pub fn enter(&mut self) -> Result<DepthAwareReaderGuard<'_, 'data, R>, SeaboredDeError<'data>> {
self.depth += 1;
if self.depth <= self.limit {
Ok(DepthAwareReaderGuard::new(self))
} else {
Err(SeaboredDeError::AllowedDepthOverflow {
depth: self.depth,
limit: self.limit,
})
}
}
}
pub struct DepthAwareReaderGuard<'a, 'data, R: Read<'data>> {
rdr: &'a mut DepthAwareReader<'data, R>,
accumulated_depth: usize,
}
impl<'a, 'data, R: Read<'data>> DepthAwareReaderGuard<'a, 'data, R> {
fn new(rdr: &'a mut DepthAwareReader<'data, R>) -> Self {
Self {
rdr,
accumulated_depth: 1,
}
}
pub fn enter(mut self) -> Result<Self, SeaboredDeError<'data>> {
self.rdr.depth += 1;
self.accumulated_depth += 1;
if self.rdr.depth <= self.rdr.limit {
Ok(self)
} else {
Err(SeaboredDeError::AllowedDepthOverflow {
depth: self.rdr.depth,
limit: self.rdr.limit,
})
}
}
}
impl<'a, 'data, R: Read<'data>> Drop for DepthAwareReaderGuard<'a, 'data, R> {
#[inline(always)]
fn drop(&mut self) {
self.rdr.depth -= self.accumulated_depth;
}
}
impl<'data, R: Read<'data>> Read<'data> for DepthAwareReader<'data, R> {
#[inline(always)]
fn peek_byte(&mut self) -> ReadResult<'data, u8> {
self.reader.peek_byte()
}
#[inline(always)]
fn advance(&mut self, n: usize) -> ReadResult<'data, ()> {
self.reader.advance(n)
}
#[inline(always)]
fn read_byte(&mut self) -> ReadResult<'data, u8> {
self.reader.read_byte()
}
#[inline(always)]
fn read_slice<'a>(&'a mut self, len: usize) -> ReadResult<'data, Cow<'data, [u8]>> {
self.reader.read_slice(len)
}
#[inline(always)]
fn read_array<const N: usize>(&mut self) -> ReadResult<'data, Cow<'data, [u8; N]>> {
self.reader.read_array()
}
}
impl<'data> Read<'data> for &'data [u8] {
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn peek_byte(&mut self) -> ReadResult<'data, u8> {
if self.is_empty() {
return Err(SeaboredDeError::IoKind(std::io::ErrorKind::UnexpectedEof));
}
Ok(unsafe { *self.as_ptr() })
}
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn advance(&mut self, n: usize) -> ReadResult<'data, ()> {
if n > self.len() {
return Err(SeaboredDeError::IoKind(std::io::ErrorKind::UnexpectedEof));
}
*self = unsafe { std::slice::from_raw_parts(self.as_ptr().add(n), self.len() - n) };
Ok(())
}
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn read_byte(&mut self) -> ReadResult<'data, u8> {
if self.is_empty() {
return Err(SeaboredDeError::IoKind(std::io::ErrorKind::UnexpectedEof));
}
let ptr = self.as_ptr();
let b = unsafe { *ptr };
*self = unsafe { std::slice::from_raw_parts(ptr.add(1), self.len() - 1) };
Ok(b)
}
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn read_slice<'a>(&'a mut self, len: usize) -> ReadResult<'data, Cow<'data, [u8]>> {
let Some((start, end)) = self.split_at_checked(len) else {
return Err(SeaboredDeError::IoKind(std::io::ErrorKind::UnexpectedEof));
};
*self = end;
Ok(Cow::Borrowed(start))
}
#[inline]
fn read_array<const N: usize>(&mut self) -> ReadResult<'data, Cow<'data, [u8; N]>> {
let Some((start, end)) = self.split_at_checked(N) else {
return Err(SeaboredDeError::IoKind(std::io::ErrorKind::UnexpectedEof));
};
*self = end;
Ok(Cow::Borrowed(unsafe { &*start.as_ptr().cast() }))
}
}
#[derive(Debug)]
#[repr(transparent)]
pub struct StdReader<T: std::io::Read + std::io::Seek>(T);
impl<T: std::io::Read + std::io::Seek> StdReader<T> {
#[inline(always)]
pub fn new(reader: T) -> Self {
Self::from(reader)
}
#[inline(always)]
pub fn into_inner(self) -> T {
self.0
}
}
impl<T: std::io::Read + std::io::Seek> From<T> for StdReader<T> {
#[inline(always)]
fn from(value: T) -> Self {
Self(value)
}
}
impl<'data, T: std::io::Read + std::io::Seek> Read<'data> for StdReader<T> {
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn peek_byte(&mut self) -> ReadResult<'data, u8> {
let b = self.read_byte()?;
self.0.seek_relative(-1)?;
Ok(b)
}
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn advance(&mut self, n: usize) -> ReadResult<'data, ()> {
self.0.seek_relative(n as i64)?;
Ok(())
}
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn read_byte(&mut self) -> ReadResult<'data, u8> {
let mut b = 0;
self.0.read_exact(std::slice::from_mut(&mut b))?;
Ok(b)
}
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn read_slice<'a>(&'a mut self, len: usize) -> ReadResult<'data, Cow<'data, [u8]>> {
let mut buf = vec![0; len];
self.0.read_exact(&mut buf)?;
Ok(Cow::Owned(buf))
}
#[cfg_attr(feature = "inline-nontrivial", inline)]
fn read_array<const N: usize>(&mut self) -> ReadResult<'data, Cow<'data, [u8; N]>> {
let mut arr = [0; N];
self.0.read_exact(&mut arr)?;
Ok(Cow::Owned(arr))
}
}
impl<T: std::io::Read + std::io::Seek> std::io::Read for StdReader<T> {
#[inline(always)]
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
self.0.read(buf)
}
}
impl<T: std::io::Read + std::io::Seek> std::io::Seek for StdReader<T> {
#[inline(always)]
fn seek(&mut self, pos: std::io::SeekFrom) -> std::io::Result<u64> {
self.0.seek(pos)
}
}
impl<T: std::io::BufRead + std::io::Seek> std::io::BufRead for StdReader<T> {
#[inline(always)]
fn fill_buf(&mut self) -> std::io::Result<&[u8]> {
self.0.fill_buf()
}
#[inline(always)]
fn consume(&mut self, amount: usize) {
self.0.consume(amount)
}
}