use futures_util::StreamExt as _;
pub const ANSWER: usize = 16 * 1024 * 1024;
pub const METADATA: usize = 1024 * 1024;
#[derive(Debug)]
pub enum IntakeError {
Declared { limit: usize, declared: u64 },
Exceeded { limit: usize },
Transport(reqwest::Error),
}
impl std::fmt::Display for IntakeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Declared { limit, declared } => write!(
f,
"the answer declared {declared} bytes and this plane reads at most {limit}"
),
Self::Exceeded { limit } => write!(
f,
"the answer grew past {limit} bytes, which is the most this plane reads"
),
Self::Transport(e) => write!(f, "the answer could not be read: {e}"),
}
}
}
impl std::error::Error for IntakeError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Transport(e) => Some(e),
_ => None,
}
}
}
impl IntakeError {
#[must_use]
pub const fn is_refusal(&self) -> bool {
matches!(self, Self::Declared { .. } | Self::Exceeded { .. })
}
#[must_use]
pub const fn transport(&self) -> Option<&reqwest::Error> {
match self {
Self::Transport(e) => Some(e),
_ => None,
}
}
}
#[derive(Debug)]
pub struct Meter {
limit: usize,
seen: usize,
}
impl Meter {
#[must_use]
pub const fn new(limit: usize) -> Self {
Self { limit, seen: 0 }
}
pub fn charge(&mut self, bytes: usize) -> Result<(), IntakeError> {
self.seen = self.seen.saturating_add(bytes);
if self.seen > self.limit {
return Err(IntakeError::Exceeded { limit: self.limit });
}
Ok(())
}
}
pub async fn read(response: reqwest::Response, limit: usize) -> Result<Vec<u8>, IntakeError> {
let declared = response
.headers()
.get(reqwest::header::CONTENT_LENGTH)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<u64>().ok());
if let Some(declared) = declared
&& declared > limit as u64
{
return Err(IntakeError::Declared { limit, declared });
}
let mut meter = Meter::new(limit);
let mut body = Vec::with_capacity(
declared
.and_then(|n| usize::try_from(n).ok())
.unwrap_or(0)
.min(limit),
);
let mut stream = response.bytes_stream();
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(IntakeError::Transport)?;
meter.charge(chunk.len())?;
body.extend_from_slice(&chunk);
}
Ok(body)
}
pub async fn read_text(response: reqwest::Response, limit: usize) -> Result<String, IntakeError> {
let bytes = read(response, limit).await?;
Ok(String::from_utf8_lossy(&bytes).into_owned())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_budget_refuses_the_chunk_that_crosses_it() {
let mut meter = Meter::new(10);
assert!(meter.charge(6).is_ok());
let err = meter.charge(6).expect_err("11 bytes is past 10");
assert!(err.is_refusal());
assert!(matches!(err, IntakeError::Exceeded { limit: 10 }));
}
#[test]
fn a_budget_admits_an_answer_of_exactly_the_ceiling() {
let mut meter = Meter::new(10);
assert!(meter.charge(10).is_ok());
assert!(meter.charge(1).is_err());
}
fn response(body: &'static str, declared: Option<u64>) -> reqwest::Response {
let mut builder = http::Response::builder();
if let Some(n) = declared {
builder = builder.header(http::header::CONTENT_LENGTH, n);
}
reqwest::Response::from(builder.body(body).expect("a response"))
}
#[tokio::test]
async fn a_declared_oversize_is_refused_before_a_byte_is_read() {
let err = read(response("hello", Some(9_000)), 10)
.await
.expect_err("9000 declared against a ceiling of 10");
assert!(matches!(
err,
IntakeError::Declared {
limit: 10,
declared: 9_000
}
));
assert!(err.is_refusal());
}
#[tokio::test]
async fn a_body_that_understates_itself_is_still_refused() {
let err = read(response("hello world", Some(2)), 4)
.await
.expect_err("eleven bytes past a ceiling of four");
assert!(
matches!(err, IntakeError::Exceeded { limit: 4 }),
"the header said it would fit; the bytes are what decides"
);
}
#[tokio::test]
async fn an_answer_within_the_ceiling_is_returned_whole() {
let body = read(response("hello", Some(5)), 16).await.expect("it fits");
assert_eq!(body, b"hello");
}
#[test]
fn an_answer_big_enough_to_overflow_is_still_refused() {
let mut meter = Meter::new(10);
assert!(meter.charge(usize::MAX).is_err());
}
}