use crate::error::{ArchToolkitError, Result};
pub async fn read_bounded_response_text<F>(
mut response: reqwest::Response,
maximum_bytes: usize,
resource_label: &str,
map_read_error: F,
) -> Result<String>
where
F: FnOnce(reqwest::Error) -> ArchToolkitError,
{
if let Some(length) = response.content_length()
&& length > u64::try_from(maximum_bytes).unwrap_or(u64::MAX)
{
return Err(response_too_large(
resource_label,
maximum_bytes,
usize::try_from(length).unwrap_or(usize::MAX),
));
}
let initial_capacity = response
.content_length()
.and_then(|length| usize::try_from(length).ok())
.unwrap_or(0);
let mut body = Vec::with_capacity(initial_capacity);
loop {
let chunk = match response.chunk().await {
Ok(Some(chunk)) => chunk,
Ok(None) => break,
Err(error) => return Err(map_read_error(error)),
};
let observed_length = body.len().saturating_add(chunk.len());
if observed_length > maximum_bytes {
return Err(response_too_large(
resource_label,
maximum_bytes,
observed_length,
));
}
body.extend_from_slice(&chunk);
}
String::from_utf8(body).map_err(|error| {
ArchToolkitError::Parse(format!(
"{resource_label} response body was not valid UTF-8: {error}"
))
})
}
fn response_too_large(
resource_label: &str,
maximum_bytes: usize,
actual_bytes: usize,
) -> ArchToolkitError {
ArchToolkitError::InputTooLong {
field: format!("{resource_label} response body"),
max_length: maximum_bytes,
actual_length: actual_bytes,
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
use std::time::Duration;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn fixture_read_error(error: reqwest::Error) -> ArchToolkitError {
ArchToolkitError::Parse(format!("fixture body read failed: {}", error.without_url()))
}
fn spawn_raw_response(response: Vec<u8>, linger: Duration) -> String {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind raw HTTP fixture");
let address = listener.local_addr().expect("raw HTTP fixture address");
thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("accept fixture request");
stream
.set_read_timeout(Some(Duration::from_secs(2)))
.expect("set fixture read timeout");
let mut request = [0_u8; 2048];
let _ = stream.read(&mut request);
stream.write_all(&response).expect("write fixture response");
stream.flush().expect("flush fixture response");
thread::sleep(linger);
});
format!("http://{address}/fixture")
}
fn response_without_length(body: &[u8]) -> Vec<u8> {
let mut response = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_vec();
response.extend_from_slice(body);
response
}
#[tokio::test]
async fn declared_oversize_is_rejected() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/body"))
.respond_with(ResponseTemplate::new(200).set_body_bytes(b"123456789"))
.mount(&server)
.await;
let response = reqwest::get(format!("{}/body", server.uri()))
.await
.expect("declared-length fixture response");
let error = read_bounded_response_text(response, 8, "declared fixture", fixture_read_error)
.await
.expect_err("declared oversize must fail");
assert_eq!(
error.to_string(),
"declared fixture response body exceeds maximum length of 8 bytes (got 9)"
);
assert!(matches!(
error,
ArchToolkitError::InputTooLong {
max_length: 8,
actual_length: 9,
..
}
));
}
#[tokio::test]
async fn dishonest_declared_oversize_is_rejected_early() {
let response =
b"HTTP/1.1 200 OK\r\nContent-Length: 9\r\nConnection: close\r\n\r\nx".to_vec();
let url = spawn_raw_response(response, Duration::from_secs(1));
let response = reqwest::get(url).await.expect("dishonest-length response");
let result = tokio::time::timeout(
Duration::from_millis(250),
read_bounded_response_text(response, 8, "dishonest fixture", fixture_read_error),
)
.await
.expect("declared oversize should reject before body completion");
assert!(matches!(
result,
Err(ArchToolkitError::InputTooLong {
max_length: 8,
actual_length: 9,
..
})
));
}
#[tokio::test]
async fn missing_length_oversize_is_rejected() {
let url = spawn_raw_response(
response_without_length(b"123456789"),
Duration::from_millis(0),
);
let response = reqwest::get(url).await.expect("missing-length response");
let error = read_bounded_response_text(response, 8, "missing fixture", fixture_read_error)
.await
.expect_err("missing-length oversize must fail");
assert!(matches!(
error,
ArchToolkitError::InputTooLong {
max_length: 8,
actual_length: 9,
..
}
));
}
#[tokio::test]
async fn dishonest_short_length_cannot_bypass_stream_limit() {
let response = b"HTTP/1.1 200 OK\r\nContent-Length: 1\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n9\r\n123456789\r\n0\r\n\r\n".to_vec();
let url = spawn_raw_response(response, Duration::from_millis(0));
let response = reqwest::get(url)
.await
.expect("dishonest short-length response");
let error =
read_bounded_response_text(response, 8, "dishonest short fixture", fixture_read_error)
.await
.expect_err("streamed overflow must override a dishonest short length");
assert!(matches!(
error,
ArchToolkitError::InputTooLong {
max_length: 8,
actual_length: 9,
..
}
));
}
#[tokio::test]
async fn chunked_overflow_stops_immediately() {
let response = b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n9\r\n123456789\r\n".to_vec();
let url = spawn_raw_response(response, Duration::from_secs(1));
let response = reqwest::get(url).await.expect("chunked fixture response");
let result = tokio::time::timeout(
Duration::from_millis(250),
read_bounded_response_text(response, 8, "chunked fixture", fixture_read_error),
)
.await
.expect("overflow should stop before the next chunk");
assert!(matches!(
result,
Err(ArchToolkitError::InputTooLong {
max_length: 8,
actual_length: 9,
..
})
));
}
#[tokio::test]
async fn exact_limit_is_accepted() {
let url = spawn_raw_response(
response_without_length(b"12345678"),
Duration::from_millis(0),
);
let response = reqwest::get(url).await.expect("exact-limit response");
let body = read_bounded_response_text(response, 8, "exact fixture", fixture_read_error)
.await
.expect("exact limit should succeed");
assert_eq!(body, "12345678");
}
#[tokio::test]
async fn invalid_utf8_is_rejected() {
let url = spawn_raw_response(
response_without_length(&[0xf0, 0x28, 0x8c]),
Duration::from_millis(0),
);
let response = reqwest::get(url).await.expect("invalid UTF-8 response");
let error = read_bounded_response_text(response, 8, "UTF-8 fixture", fixture_read_error)
.await
.expect_err("invalid UTF-8 must fail");
let message = error.to_string();
assert!(matches!(error, ArchToolkitError::Parse(_)));
assert!(message.contains("UTF-8 fixture"));
assert!(message.contains("not valid UTF-8"));
assert!(!message.contains('�'));
}
}