use std::io::{self, Write};
use super::super::io::retry::env_u64;
const DEFAULT_MAX_DECOMPRESSED_BYTES: u64 = 256 * 1024 * 1024;
#[must_use]
pub fn max_decompressed_bytes() -> u64 {
env_u64("VERSATILES_MAX_DECOMPRESSED_BYTES", DEFAULT_MAX_DECOMPRESSED_BYTES)
}
pub struct LimitedWriter {
buffer: Vec<u8>,
limit: u64,
}
impl LimitedWriter {
#[must_use]
pub fn new(limit: u64) -> Self {
Self {
buffer: Vec::new(),
limit,
}
}
#[must_use]
pub fn into_vec(self) -> Vec<u8> {
self.buffer
}
}
impl Write for LimitedWriter {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
if self.limit > 0 {
let written = self.buffer.len() as u64 + buf.len() as u64;
if written > self.limit {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"decompressed data exceeds the limit of {} bytes. \
If this input really is that large, raise VERSATILES_MAX_DECOMPRESSED_BYTES \
(or set it to 0 to remove the limit)",
self.limit
),
));
}
}
self.buffer.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn writes_up_to_the_limit() {
let mut writer = LimitedWriter::new(8);
writer.write_all(b"1234").unwrap();
writer.write_all(b"5678").unwrap();
assert_eq!(writer.into_vec(), b"12345678");
}
#[test]
fn refuses_to_grow_past_the_limit() {
let mut writer = LimitedWriter::new(8);
writer.write_all(b"12345").unwrap();
let error = writer.write_all(b"6789").unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
assert!(error.to_string().contains("exceeds the limit of 8 bytes"), "{error}");
assert_eq!(writer.into_vec(), b"12345");
}
#[test]
fn zero_means_unbounded() {
let mut writer = LimitedWriter::new(0);
writer.write_all(&vec![0u8; 10_000]).unwrap();
assert_eq!(writer.into_vec().len(), 10_000);
}
#[test]
fn default_limit_is_used_when_unset() {
assert_eq!(max_decompressed_bytes(), 256 * 1024 * 1024);
}
}