use std::fmt;
#[derive(Debug)]
pub enum ReadError {
TooLarge { limit: usize, declared: Option<u64> },
Transport(reqwest::Error),
}
impl fmt::Display for ReadError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::TooLarge {
limit,
declared: Some(declared),
} => write!(
f,
"declared Content-Length {declared} exceeds the {limit} byte limit"
),
Self::TooLarge {
limit,
declared: None,
} => write!(f, "body exceeds the {limit} byte limit"),
Self::Transport(e) => write!(f, "read failed: {e}"),
}
}
}
impl std::error::Error for ReadError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Transport(e) => Some(e),
Self::TooLarge { .. } => None,
}
}
}
impl ReadError {
pub fn is_too_large(&self) -> bool {
matches!(self, Self::TooLarge { .. })
}
}
pub fn check_declared_length(response: &reqwest::Response, limit: usize) -> Result<(), ReadError> {
match response.content_length() {
Some(declared) if declared > limit as u64 => Err(ReadError::TooLarge {
limit,
declared: Some(declared),
}),
_ => Ok(()),
}
}
pub async fn read_bounded(
mut response: reqwest::Response,
limit: usize,
) -> Result<Vec<u8>, ReadError> {
check_declared_length(&response, limit)?;
let mut body = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(ReadError::Transport)? {
if body.len() + chunk.len() > limit {
return Err(ReadError::TooLarge {
limit,
declared: None,
});
}
body.extend_from_slice(&chunk);
}
Ok(body)
}
pub async fn read_preview(mut response: reqwest::Response, limit: usize) -> Preview {
let mut bytes = Vec::new();
let mut truncated = false;
while let Some(chunk) = response.chunk().await.ok().flatten() {
let room = limit.saturating_sub(bytes.len());
let take = chunk.len().min(room);
bytes.extend_from_slice(&chunk[..take]);
if take < chunk.len() {
truncated = true;
break;
}
}
Preview { bytes, truncated }
}
pub struct Preview {
pub bytes: Vec<u8>,
pub truncated: bool,
}
impl Preview {
pub fn to_message(&self) -> String {
let text = String::from_utf8_lossy(&self.bytes);
if self.truncated {
format!("{text}… (truncated)")
} else {
text.into_owned()
}
}
}
#[cfg(test)]
pub(crate) async fn flood_server(
chunk_bytes: usize,
max_chunks: usize,
) -> (String, tokio::task::JoinHandle<usize>) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener");
let url = format!("http://{}/", listener.local_addr().expect("test addr"));
let chunk = vec![b'x'; chunk_bytes];
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("test accept");
let mut discard = [0u8; 4096];
let _ = socket.read(&mut discard).await;
let head = "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n\
Transfer-Encoding: chunked\r\n\r\n";
if socket.write_all(head.as_bytes()).await.is_err() {
return 0;
}
let mut written = 0usize;
for _ in 0..max_chunks {
let framed = format!("{:x}\r\n", chunk.len());
if socket.write_all(framed.as_bytes()).await.is_err()
|| socket.write_all(&chunk).await.is_err()
|| socket.write_all(b"\r\n").await.is_err()
{
break;
}
written += chunk.len();
}
let _ = socket.write_all(b"0\r\n\r\n").await;
written
});
(url, handle)
}
#[cfg(test)]
pub(crate) fn assert_stopped_early(written: usize, attempted: usize, what: &str) {
assert!(
written < attempted,
"{what}: the peer got to write all {attempted} bytes — the cap must be \
enforced while streaming, not after the body is buffered"
);
let leaked = (written as f64 / attempted as f64) * 100.0;
if written * 2 >= attempted {
eprintln!(
"note: {what} stopped the peer at {written} of {attempted} bytes \
({leaked:.0}%) — bounded, but the reader is running well behind \
the writer"
);
}
}
#[cfg(test)]
mod tests {
use super::*;
const CHUNK: usize = 64 * 1024;
const CHUNKS: usize = 128;
#[tokio::test]
async fn a_chunked_body_is_refused_without_reading_it_all() {
let (url, server) = flood_server(CHUNK, CHUNKS).await;
let response = reqwest::Client::new().get(url).send().await.expect("head");
let err = read_bounded(response, 1024).await.expect_err("must refuse");
assert!(
matches!(err, ReadError::TooLarge { declared: None, .. }),
"a chunked response declares no length, so the refusal must come \
from the bytes themselves: {err}"
);
assert_stopped_early(
server.await.expect("test server"),
CHUNK * CHUNKS,
"read_bounded",
);
}
#[tokio::test]
async fn a_declared_oversize_body_is_refused_unread() {
use tokio::io::AsyncWriteExt;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener");
let url = format!("http://{}/", listener.local_addr().expect("test addr"));
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("test accept");
let mut discard = [0u8; 4096];
let _ = tokio::io::AsyncReadExt::read(&mut socket, &mut discard).await;
let _ = socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 5000\r\n\r\n")
.await;
tokio::time::sleep(std::time::Duration::from_secs(10)).await;
});
let response = reqwest::Client::new().get(url).send().await.expect("head");
let err = tokio::time::timeout(
std::time::Duration::from_secs(2),
read_bounded(response, 1024),
)
.await
.expect("must not wait on a body it has already refused")
.expect_err("must refuse");
assert!(matches!(
err,
ReadError::TooLarge {
limit: 1024,
declared: Some(5000)
}
));
server.abort();
}
#[tokio::test]
async fn a_body_exactly_at_the_limit_is_accepted() {
let (url, server) = flood_server(512, 2).await;
let response = reqwest::Client::new().get(url).send().await.expect("head");
let body = read_bounded(response, 1024).await.expect("accepted");
assert_eq!(body.len(), 1024);
assert!(body.iter().all(|b| *b == b'x'));
let _ = server.await;
}
#[tokio::test]
async fn a_preview_truncates_and_says_so() {
let (url, server) = flood_server(CHUNK, CHUNKS).await;
let response = reqwest::Client::new().get(url).send().await.expect("head");
let preview = read_preview(response, 16).await;
assert_eq!(preview.bytes.len(), 16);
assert!(preview.truncated);
assert_eq!(
preview.to_message(),
format!("{}… (truncated)", "x".repeat(16))
);
assert_stopped_early(
server.await.expect("test server"),
CHUNK * CHUNKS,
"read_preview",
);
}
#[tokio::test]
async fn a_preview_of_non_utf8_bytes_is_lossy_rather_than_lost() {
use tokio::io::AsyncWriteExt;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener");
let url = format!("http://{}/", listener.local_addr().expect("test addr"));
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("test accept");
let mut discard = [0u8; 4096];
let _ = tokio::io::AsyncReadExt::read(&mut socket, &mut discard).await;
let _ = socket
.write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 4\r\n\r\n")
.await;
let _ = socket.write_all(&[0xff, b'b', 0xfe, b'd']).await;
});
let response = reqwest::Client::new().get(url).send().await.expect("head");
let preview = read_preview(response, 512).await;
assert!(!preview.truncated);
assert_eq!(
preview.to_message(),
"\u{fffd}b\u{fffd}d",
"invalid bytes must become replacement characters, not an error"
);
let _ = server.await;
}
#[tokio::test]
async fn a_short_preview_is_not_marked_truncated() {
let (url, server) = flood_server(4, 1).await;
let response = reqwest::Client::new().get(url).send().await.expect("head");
let preview = read_preview(response, 512).await;
assert!(!preview.truncated);
assert_eq!(preview.to_message(), "xxxx");
let _ = server.await;
}
}