use std::{
io::{ErrorKind, Read},
mem,
};
use crate::{
Error, Result,
varint::{max_of_last_byte, varint_max},
};
const READ_CHUNK: usize = 64 * 1024;
pub struct SkipRead<R> {
stack: SkipStack<R>,
depth: usize,
failure: Option<ErrorKind>,
}
impl<R: Read> SkipRead<R> {
pub fn new_value(inner: R, len: usize) -> Self {
let stack = SkipStack::SkipBlock(SkipBlock::exact(SkipStack::Base(inner), len));
Self { stack, depth: 1, failure: None }
}
pub fn new(inner: R) -> Self {
Self { stack: SkipStack::Base(inner), depth: 0, failure: None }
}
fn note_failure<T>(&mut self, res: Result<T>) -> Result<T> {
if let Err(Error::Io(err)) = &res {
self.failure.get_or_insert(err.kind());
}
res
}
pub fn read_u8(&mut self) -> Result<u8> {
let res = self.stack.read_byte();
self.note_failure(res)
}
pub fn read(&mut self, cnt: usize) -> Result<Vec<u8>> {
let res = self.stack.read(cnt);
self.note_failure(res)
}
pub fn start_skippable(&mut self) {
let this = mem::replace(&mut self.stack, SkipStack::Dummy);
self.stack = SkipStack::SkipBlock(SkipBlock::new(this));
self.depth += 1;
}
pub fn start_empty_block(&mut self) {
let this = mem::replace(&mut self.stack, SkipStack::Dummy);
self.stack = SkipStack::SkipBlock(SkipBlock::exact(this, 0));
self.depth += 1;
}
pub fn end_skippable(&mut self) -> Result<()> {
match mem::replace(&mut self.stack, SkipStack::Dummy) {
SkipStack::Base(_) => panic!("no skip block is open"),
SkipStack::SkipBlock(sb) => self.stack = self.note_failure(sb.finish())?,
SkipStack::Dummy => unreachable!(),
}
self.depth -= 1;
Ok(())
}
pub fn depth(&self) -> usize {
self.depth
}
pub fn pop_to(&mut self, depth: usize) -> Result<()> {
assert!(depth <= self.depth, "cannot return to a depth that is not open");
if let Some(kind) = self.failure {
return Err(Error::Io(kind.into()));
}
while self.depth > depth {
self.end_skippable()?;
}
Ok(())
}
pub fn into_inner(self) -> R {
self.stack.into_inner()
}
pub fn read_skippable_block(&mut self) -> Result<Vec<u8>> {
self.start_skippable();
let SkipStack::SkipBlock(sb) = &mut self.stack else { unreachable!() };
let res = sb.read_all();
let data = self.note_failure(res)?;
self.end_skippable()?;
Ok(data)
}
pub fn block_exhausted(&mut self) -> Result<bool> {
let res = match &mut self.stack {
SkipStack::SkipBlock(sb) => sb.exhausted(),
SkipStack::Base(_) | SkipStack::Dummy => unreachable!("no block to be at the end of"),
};
self.note_failure(res)
}
pub fn read_rest(&mut self) -> Result<Vec<u8>> {
let res = match &mut self.stack {
SkipStack::SkipBlock(sb) => sb.read_all(),
SkipStack::Base(_) | SkipStack::Dummy => unreachable!("no block to read to the end of"),
};
self.note_failure(res)
}
}
enum SkipStack<R> {
Base(R),
SkipBlock(SkipBlock<R>),
Dummy,
}
impl<R: Read> SkipStack<R> {
pub fn read(&mut self, ct: usize) -> Result<Vec<u8>> {
match self {
Self::Base(base) => {
let mut buf = Vec::new();
while buf.len() < ct {
let chunk = (ct - buf.len()).min(READ_CHUNK);
let start = buf.len();
buf.resize(start + chunk, 0);
base.read_exact(&mut buf[start..])?;
}
Ok(buf)
}
Self::SkipBlock(sb) => sb.read(ct),
Self::Dummy => unreachable!(),
}
}
fn read_byte(&mut self) -> Result<u8> {
match self {
Self::Base(base) => {
let mut buf = [0u8; 1];
base.read_exact(&mut buf)?;
Ok(buf[0])
}
Self::SkipBlock(sb) => sb.read_byte(),
Self::Dummy => unreachable!(),
}
}
fn try_take_varint_u16(&mut self) -> Result<u16> {
let mut out = 0;
for i in 0..varint_max::<u16>() {
let val = self.read_byte()?;
let carry = (val & 0x7F) as u16;
out |= carry << (7 * i);
if (val & 0x80) == 0 {
if i == varint_max::<u16>() - 1 && val > max_of_last_byte::<u16>() {
return Err(Error::BadVarint);
} else {
return Ok(out);
}
}
}
Err(Error::BadVarint)
}
fn into_inner(self) -> R {
match self {
SkipStack::Base(base) => base,
SkipStack::SkipBlock(sb) => sb.inner.into_inner(),
SkipStack::Dummy => unreachable!(),
}
}
}
struct SkipBlock<R> {
inner: Box<SkipStack<R>>,
remaining: usize,
has_next_block: bool,
}
impl<R: Read> SkipBlock<R> {
const MAX_LEN: usize = u16::MAX as usize;
fn new(inner: SkipStack<R>) -> Self {
Self { inner: Box::new(inner), remaining: 0, has_next_block: true }
}
fn exact(inner: SkipStack<R>, len: usize) -> Self {
Self { inner: Box::new(inner), remaining: len, has_next_block: false }
}
fn update_remaining(&mut self) -> Result<()> {
if self.remaining > 0 || !self.has_next_block {
return Ok(());
}
self.remaining = self.inner.try_take_varint_u16()?.into();
self.has_next_block = self.remaining == Self::MAX_LEN;
Ok(())
}
fn read_byte(&mut self) -> Result<u8> {
self.update_remaining()?;
if self.remaining == 0 {
return Err(Error::EndOfBlock);
}
let byte = self.inner.read_byte()?;
self.remaining -= 1;
Ok(byte)
}
fn read(&mut self, mut ct: usize) -> Result<Vec<u8>> {
self.update_remaining()?;
if self.remaining >= ct {
let buf = self.inner.read(ct)?;
self.remaining -= ct;
return Ok(buf);
}
let mut buf = Vec::with_capacity(ct.min(READ_CHUNK));
while ct > 0 {
self.update_remaining()?;
if self.remaining == 0 {
return Err(Error::EndOfBlock);
}
let n = ct.min(self.remaining);
buf.extend(&self.inner.read(n)?);
self.remaining -= n;
ct -= n;
}
Ok(buf)
}
fn finish(mut self) -> Result<SkipStack<R>> {
loop {
self.update_remaining()?;
if self.remaining > 0 {
self.inner.read(self.remaining)?;
self.remaining = 0;
} else {
break;
}
}
Ok(*self.inner)
}
fn exhausted(&mut self) -> Result<bool> {
self.update_remaining()?;
Ok(self.remaining == 0)
}
fn read_all(&mut self) -> Result<Vec<u8>> {
let mut buf = Vec::new();
loop {
self.update_remaining()?;
if self.remaining == 0 {
break;
}
let chunk = self.inner.read(self.remaining)?;
buf.extend(chunk);
self.remaining = 0;
}
Ok(buf)
}
}