use std::io::{Cursor, Error, ErrorKind, IoSlice, Result, Seek, SeekFrom, Write};
pub struct LimitedCursor {
inner: Cursor<Vec<u8>>,
limit: u64,
}
impl LimitedCursor {
pub fn new(limit: u64) -> Self {
Self {
inner: Cursor::new(vec![]),
limit,
}
}
pub fn into_inner(self) -> Vec<u8> {
self.inner.into_inner()
}
}
impl Seek for LimitedCursor {
fn seek(&mut self, pos: SeekFrom) -> Result<u64> {
self.inner.seek(pos)
}
fn stream_position(&mut self) -> Result<u64> {
self.inner.stream_position()
}
}
impl Write for LimitedCursor {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
let length = self.inner.position() + buf.len() as u64;
if length > self.limit {
return Err(Error::new(
ErrorKind::Other,
format!(
"Write of {} bytes at position {} would exceed size limit {}",
buf.len(),
self.inner.position(),
self.limit
),
));
}
self.inner.write(buf)
}
fn flush(&mut self) -> Result<()> {
self.inner.flush()
}
}
pub struct LimitedWrite<T>
where
T: Write,
{
inner: T,
limit: u64,
}
impl<T> LimitedWrite<T>
where
T: Write,
{
pub fn new(inner: T, limit: u64) -> Self {
Self { inner, limit }
}
pub fn remaining(&self) -> u64 {
self.limit
}
fn check_limit(&mut self, len: u64) -> Result<()> {
if len as u64 > self.limit {
return Err(Error::new(
ErrorKind::Other,
format!(
"Write of {} bytes would exceed size limit {}",
len, self.limit
),
));
}
self.limit -= len;
Ok(())
}
}
impl<T> Write for LimitedWrite<T>
where
T: Write,
{
fn write(&mut self, buf: &[u8]) -> Result<usize> {
self.check_limit(buf.len() as u64)?;
self.inner.write(buf)
}
fn write_vectored(&mut self, bufs: &[IoSlice<'_>]) -> Result<usize> {
for buf in bufs {
self.check_limit(buf.len() as u64)?;
}
self.inner.write_vectored(bufs)
}
fn flush(&mut self) -> Result<()> {
self.inner.flush()
}
fn write_all(&mut self, buf: &[u8]) -> Result<()> {
self.check_limit(buf.len() as u64)?;
self.inner.write_all(buf)
}
}