use core::cmp;
#[cfg(not(feature = "std"))]
use core::fmt;
#[cfg(feature = "std")]
use std::io;
pub trait Read {
type Error;
fn unexpected_eof() -> Self::Error;
fn read_all(&mut self, buf: &mut[u8]) -> Result<usize, Self::Error>;
fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), Self::Error> {
if buf.len() != self.read_all(buf)? {
Err(Self::unexpected_eof())
}
else {
Ok(())
}
}
fn take(self, limit: u64) -> Take<Self>
where Self: Sized
{
Take { inner: self, limit }
}
#[inline]
fn by_ref(&mut self) -> &mut Self {
self
}
}
pub(crate) fn discard_to_end<R: Read, const BUF: usize>(rd: &mut R) -> Result<(), R::Error> {
use core::mem::MaybeUninit;
assert!(BUF != 0);
let mut data = MaybeUninit::<[u8; BUF]>::uninit();
let buf: &mut [u8; BUF] = unsafe {
data.assume_init_mut()
};
while 0 != rd.read_all(buf)? {}
Ok(())
}
#[derive(Debug)]
pub struct Take<R> {
limit: u64,
inner: R,
}
impl<R> Take<R> {
#[inline]
pub fn limit(&self) -> u64 {
self.limit
}
pub fn into_inner(self) -> R {
self.inner
}
pub fn get_ref(&self) -> &R {
&self.inner
}
pub fn get_mut(&mut self) -> &mut R {
&mut self.inner
}
}
impl<R: Read> Read for Take<R> {
type Error = R::Error;
#[inline]
fn unexpected_eof() -> Self::Error {
R::unexpected_eof()
}
#[inline]
fn read_all(&mut self, buf: &mut [u8]) -> Result<usize, Self::Error> {
if self.limit == 0 {
return Ok(0);
}
let max = cmp::min(buf.len() as u64, self.limit) as usize;
let n = self.inner.read_all(&mut buf[..max])?;
self.limit = self.limit.checked_sub(n as u64).expect("number of read bytes exceeds limit");
Ok(n)
}
#[inline]
fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), Self::Error> {
let len = buf.len();
if len as u64 > self.limit {
return Err(Self::unexpected_eof());
}
self.inner.read_exact(buf)?;
self.limit -= len as u64;
Ok(())
}
}
#[cfg(feature = "std")]
impl<R: io::Read> Read for R {
type Error = io::Error;
fn unexpected_eof() -> Self::Error {
io::Error::new(io::ErrorKind::UnexpectedEof, "failed to fill whole buffer")
}
fn read_all(&mut self, mut buf: &mut[u8]) -> Result<usize, Self::Error> {
let orig_len = buf.len();
loop {
match self.read(buf) {
Ok(0) => break,
Ok(n) if n < buf.len() => {
buf = &mut buf[n..];
},
Ok(..) => return Ok(orig_len),
Err(ref e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e)
}
}
Ok(orig_len - buf.len())
}
#[inline]
fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), Self::Error> {
io::Read::read_exact(self, buf)
}
}
#[cfg(not(feature = "std"))]
#[cfg_attr(docsrs, doc(cfg(not(feature = "std"))))]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UnexpectedEofError;
#[cfg(not(feature = "std"))]
impl core::error::Error for UnexpectedEofError {}
#[cfg(not(feature = "std"))]
impl<R: Read + ?Sized> Read for &mut R {
type Error = R::Error;
fn unexpected_eof() -> Self::Error {
R::unexpected_eof()
}
fn read_all(&mut self, buf: &mut[u8]) -> Result<usize, Self::Error> {
R::read_all(*self, buf)
}
fn read_exact(&mut self, buf: &mut[u8]) -> Result<(), Self::Error> {
R::read_exact(*self, buf)
}
}
#[cfg(not(feature = "std"))]
impl<R: Read + ?Sized> Read for alloc::boxed::Box<R> {
type Error = R::Error;
fn unexpected_eof() -> Self::Error {
R::unexpected_eof()
}
fn read_all(&mut self, buf: &mut[u8]) -> Result<usize, Self::Error> {
R::read_all(self, buf)
}
fn read_exact(&mut self, buf: &mut[u8]) -> Result<(), Self::Error> {
R::read_exact(self, buf)
}
}
#[cfg(not(feature = "std"))]
impl fmt::Display for UnexpectedEofError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("failed to fill whole buffer")
}
}
#[cfg(not(feature = "std"))]
impl Read for &'_[u8] {
type Error = UnexpectedEofError;
#[inline]
fn unexpected_eof() -> Self::Error {
UnexpectedEofError
}
#[inline]
fn read_all(&mut self, buf: &mut[u8]) -> Result<usize, Self::Error> {
let amt = cmp::min(buf.len(), self.len());
let (a, b) = self.split_at(amt);
if amt == 1 {
buf[0] = a[0];
} else {
buf[..amt].copy_from_slice(a);
}
*self = b;
Ok(amt)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(not(feature = "std"))]
#[test]
fn stub_io_no_std_works() {
use alloc::{boxed::Box, string::ToString};
assert_eq!((&[][..]).read_all(&mut[]).unwrap(), 0);
assert_eq!((&[][..]).read_all(&mut[0]).unwrap(), 0);
assert_eq!((&[][..]).read_exact(&mut[0]).unwrap_err(), UnexpectedEofError);
assert_eq!((&[][..]).take(0).read_exact(&mut[0]).unwrap_err(), UnexpectedEofError);
let mut data: Box<&[u8]> = Box::new(&[]);
assert_eq!(data.read_all(&mut[]).unwrap(), 0);
assert_eq!(data.read_all(&mut[0]).unwrap(), 0);
assert_eq!(data.read_exact(&mut[0]).unwrap_err(), UnexpectedEofError);
assert_eq!(<Box<&[u8]> as Read>::unexpected_eof(), UnexpectedEofError);
let data: &mut &mut _ = &mut &mut data;
assert_eq!(data.read_all(&mut[]).unwrap(), 0);
assert_eq!(data.read_all(&mut[0]).unwrap(), 0);
assert_eq!(data.read_exact(&mut[0]).unwrap_err(), UnexpectedEofError);
assert_eq!(<&mut &[u8] as Read>::unexpected_eof(), UnexpectedEofError);
assert_eq!(UnexpectedEofError.to_string(), "failed to fill whole buffer");
}
#[cfg(feature = "std")]
#[test]
fn stub_io_std_works() {
use std::io;
#[derive(Debug, Eq, PartialEq)]
struct Interrupted(bool);
impl io::Read for Interrupted {
fn read(&mut self, _buf: &mut[u8]) -> io::Result<usize> {
if !self.0 {
self.0 = true;
Err(io::ErrorKind::Interrupted.into())
}
else {
Err(io::ErrorKind::UnexpectedEof.into())
}
}
}
assert_eq!((&[][..]).read_all(&mut[]).unwrap(), 0);
assert_eq!((&[][..]).read_all(&mut[0]).unwrap(), 0);
assert_eq!(Interrupted(false).read_all(&mut[0]).unwrap_err().kind(), io::ErrorKind::UnexpectedEof);
assert_eq!(Interrupted(true).take(0).get_ref(), &Interrupted(true));
assert_eq!(Interrupted(true).take(0).get_mut(), &mut Interrupted(true));
}
}