#[cfg(test)]
mod tests;
use std::{
fmt,
io::{self, Write},
};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum WriterError {
LimitExceeded {
limit: u64,
},
InvalidWriteCount {
offered: usize,
written: usize,
},
}
impl fmt::Display for WriterError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::LimitExceeded { limit } => write!(f, "output exceeds {limit} bytes"),
Self::InvalidWriteCount { offered, written } => {
write!(
f,
"writer accepted {written} bytes from a {offered}-byte buffer"
)
}
}
}
}
impl std::error::Error for WriterError {}
pub struct BoundedWriter<W> {
inner: W,
limit: u64,
bytes: u64,
exceeded: bool,
}
impl<W> BoundedWriter<W> {
#[must_use]
pub const fn new(inner: W, limit: u64) -> Self {
Self {
inner,
limit,
bytes: 0,
exceeded: false,
}
}
#[must_use]
pub const fn bytes_written(&self) -> u64 {
self.bytes
}
#[must_use]
pub const fn limit_exceeded(&self) -> bool {
self.exceeded
}
#[must_use]
pub fn into_inner(self) -> W {
self.inner
}
}
impl<W: Write> Write for BoundedWriter<W> {
fn write(&mut self, buffer: &[u8]) -> io::Result<usize> {
if !u64::try_from(buffer.len()).is_ok_and(|length| length <= self.limit - self.bytes) {
self.exceeded = true;
return Err(io::Error::other(WriterError::LimitExceeded {
limit: self.limit,
}));
}
let written = self.inner.write(buffer)?;
if written > buffer.len() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
WriterError::InvalidWriteCount {
offered: buffer.len(),
written,
},
));
}
self.bytes += written as u64;
Ok(written)
}
fn flush(&mut self) -> io::Result<()> {
self.inner.flush()
}
}