use std::io;
use moirai_async::io::{AsyncReadExt, AsyncWrite, AsyncWriteExt};
pub const DEFAULT_MAX_RESPONSE_BYTES: usize = 64 * 1024 * 1024;
#[derive(Debug, Clone)]
pub struct Response {
pub status: u16,
pub headers: Vec<(String, String)>,
pub body: Vec<u8>,
pub keep_alive: bool,
}
impl Response {
#[must_use]
pub fn header(&self, name: &str) -> Option<&str> {
let name = name.to_ascii_lowercase();
self.headers
.iter()
.find(|(k, _)| *k == name)
.map(|(_, v)| v.as_str())
}
}
pub async fn write_request<S: AsyncWrite + Unpin>(
stream: &mut S,
method: &str,
host: &str,
path: &str,
headers: &[(&str, &str)],
body: Option<&[u8]>,
) -> io::Result<()> {
let mut req = Vec::with_capacity(256);
req.extend_from_slice(method.as_bytes());
req.push(b' ');
req.extend_from_slice(path.as_bytes());
req.extend_from_slice(b" HTTP/1.1\r\n");
let has = |n: &str| headers.iter().any(|(k, _)| k.eq_ignore_ascii_case(n));
if !has("host") {
req.extend_from_slice(format!("Host: {host}\r\n").as_bytes());
}
for (k, v) in headers {
req.extend_from_slice(format!("{k}: {v}\r\n").as_bytes());
}
if let Some(b) = body
&& !has("content-length")
{
req.extend_from_slice(format!("Content-Length: {}\r\n", b.len()).as_bytes());
}
req.extend_from_slice(b"\r\n");
if let Some(b) = body {
req.extend_from_slice(b);
}
stream.write_all(&req).await?;
stream.flush().await
}
struct Buffered<'a, S> {
stream: &'a mut S,
buf: Vec<u8>,
pos: usize,
limit: usize,
}
impl<'a, S: AsyncReadExt + Unpin> Buffered<'a, S> {
fn new(stream: &'a mut S, limit: usize) -> Self {
Self {
stream,
buf: Vec::with_capacity(8192.min(limit.max(1))),
pos: 0,
limit,
}
}
fn available(&self) -> usize {
self.buf.len().saturating_sub(self.pos)
}
async fn fill(&mut self) -> io::Result<usize> {
if self.buf.len() >= self.limit {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"response exceeds maximum size",
));
}
let mut tmp = [0u8; 8192];
let n = self.stream.read(&mut tmp).await?;
#[expect(
clippy::indexing_slicing,
reason = "n <= tmp.len() per the Read trait contract"
)]
let src = &tmp[..n];
self.buf.extend_from_slice(src);
if self.buf.len() > self.limit {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"response exceeds maximum size",
));
}
Ok(n)
}
async fn read_crlf_line(&mut self) -> io::Result<String> {
loop {
let (_, tail) = self.buf.split_at(self.pos);
if let Some(rel) = find_crlf(tail) {
#[expect(
clippy::indexing_slicing,
reason = "rel < tail.len() per the CRLF search above"
)]
let line = tail[..rel].to_vec();
self.pos = self
.pos
.checked_add(rel)
.and_then(|after_line| after_line.checked_add(2))
.expect("invariant: CRLF line fits inside the buffered prefix");
return String::from_utf8(line).map_err(|_| {
io::Error::new(io::ErrorKind::InvalidData, "non-UTF8 header line")
});
}
if self.fill().await? == 0 {
return Err(eof("CRLF line"));
}
}
}
async fn read_n(&mut self, n: usize) -> io::Result<Vec<u8>> {
while self.available() < n {
if self.fill().await? == 0 {
return Err(eof("body"));
}
}
let (_, rest) = self.buf.split_at(self.pos);
#[expect(
clippy::indexing_slicing,
reason = "rest.len() >= n follows from available() >= n"
)]
let out = rest[..n].to_vec();
self.pos = self
.pos
.checked_add(n)
.expect("invariant: n <= available() was established by the fill loop");
Ok(out)
}
async fn read_to_eof(&mut self) -> io::Result<Vec<u8>> {
while self.fill().await? != 0 {}
let (_, tail) = self.buf.split_at(self.pos);
Ok(tail.to_vec())
}
async fn read_chunked(&mut self) -> io::Result<Vec<u8>> {
let mut body = Vec::new();
loop {
let line = self.read_crlf_line().await?;
let size_field = line.split(';').next().unwrap_or("").trim();
let size = usize::from_str_radix(size_field, 16)
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "bad chunk size"))?;
if size == 0 {
while !self.read_crlf_line().await?.is_empty() {}
break;
}
body.extend_from_slice(&self.read_n(size).await?);
let crlf = self.read_n(2).await?;
if crlf != b"\r\n" {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"missing chunk CRLF",
));
}
}
Ok(body)
}
}
pub async fn read_response<S: AsyncReadExt + Unpin>(
stream: &mut S,
is_head: bool,
max_response_bytes: usize,
) -> io::Result<Response> {
let mut r = Buffered::new(stream, max_response_bytes);
let (status, headers) = loop {
let mut header_storage = [httparse::EMPTY_HEADER; 96];
let mut resp = httparse::Response::new(&mut header_storage);
let (_, tail) = r.buf.split_at(r.pos);
match resp.parse(tail) {
Ok(httparse::Status::Complete(consumed)) => {
let status = resp
.code
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "no status code"))?;
let headers: Vec<(String, String)> = resp
.headers
.iter()
.map(|h| {
(
h.name.to_ascii_lowercase(),
String::from_utf8_lossy(h.value).into_owned(),
)
})
.collect();
r.pos = r
.pos
.checked_add(consumed)
.expect("invariant: header parse cannot consume beyond the buffer");
break (status, headers);
}
Ok(httparse::Status::Partial) => {
if r.fill().await? == 0 {
return Err(eof("response headers"));
}
}
Err(e) => {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("malformed response: {e}"),
));
}
}
};
let find = |n: &str| {
headers
.iter()
.find(|(k, _)| k == n)
.map(|(_, v)| v.as_str())
};
let chunked = find("transfer-encoding")
.map(|v| v.to_ascii_lowercase().contains("chunked"))
.unwrap_or(false);
let content_length: Option<usize> = match find("content-length") {
Some(v) => Some(v.trim().parse().map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("malformed Content-Length: {v:?}"),
)
})?),
None => None,
};
let conn_close = find("connection")
.map(|v| v.eq_ignore_ascii_case("close"))
.unwrap_or(false);
let bodyless = is_head || status == 204 || status == 304 || (100..200).contains(&status);
let (body, framed) = if bodyless {
(Vec::new(), true)
} else if chunked {
(r.read_chunked().await?, true)
} else if let Some(len) = content_length {
if len > max_response_bytes {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Content-Length exceeds maximum response size",
));
}
(r.read_n(len).await?, true)
} else {
(r.read_to_eof().await?, false)
};
Ok(Response {
status,
headers,
body,
keep_alive: framed && !conn_close,
})
}
fn find_crlf(buf: &[u8]) -> Option<usize> {
buf.windows(2).position(|w| w == b"\r\n")
}
fn eof(what: &str) -> io::Error {
io::Error::new(
io::ErrorKind::UnexpectedEof,
format!("connection closed while reading {what}"),
)
}
#[cfg(test)]
mod tests {
use super::*;
use std::pin::Pin;
use std::task::{Context, Poll};
use moirai_async::io::AsyncRead;
struct MockReader {
data: Vec<u8>,
pos: usize,
}
impl MockReader {
fn new(data: Vec<u8>) -> Self {
Self { data, pos: 0 }
}
}
impl AsyncRead for MockReader {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<io::Result<usize>> {
let (_, rest) = self.data.split_at(self.pos);
let n = rest.len().min(buf.len());
#[expect(clippy::indexing_slicing, reason = "n <= buf.len() by the min above")]
let dst = &mut buf[..n];
#[expect(clippy::indexing_slicing, reason = "n <= rest.len() by the min above")]
let src = &rest[..n];
dst.copy_from_slice(src);
self.pos = self
.pos
.checked_add(n)
.expect("invariant: n <= data.len() - pos");
Poll::Ready(Ok(n))
}
}
fn read(data: Vec<u8>, max: usize) -> io::Result<Response> {
moirai::block_on(read_response(&mut MockReader::new(data), false, max))
}
#[test]
fn oversized_content_length_is_rejected_up_front() {
let resp = b"HTTP/1.1 200 OK\r\nContent-Length: 999999999\r\n\r\n".to_vec();
let err = read(resp, 4096).expect_err("oversized Content-Length must be rejected");
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn eof_delimited_body_over_limit_is_rejected() {
let mut resp = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_vec();
resp.extend(std::iter::repeat_n(b'x', 64 * 1024));
let err = read(resp, 8 * 1024).expect_err("EOF body over the cap must be rejected");
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn chunked_body_over_limit_is_rejected() {
let mut resp = b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n".to_vec();
for _ in 0..64 {
resp.extend_from_slice(b"1000\r\n"); resp.extend(std::iter::repeat_n(b'y', 0x1000));
resp.extend_from_slice(b"\r\n");
}
resp.extend_from_slice(b"0\r\n\r\n");
let err = read(resp, 16 * 1024).expect_err("chunked body over the cap must be rejected");
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn malformed_content_length_is_invalid_data_not_eof_framing() {
for bad in ["abc", "-5", "18446744073709551616", "12abc", ""] {
let resp =
format!("HTTP/1.1 200 OK\r\nContent-Length: {bad}\r\n\r\nhello").into_bytes();
let err = read(resp, 64 * 1024)
.expect_err("garbage Content-Length must be rejected, not EOF-framed");
assert_eq!(err.kind(), io::ErrorKind::InvalidData, "value: {bad:?}");
}
}
#[test]
fn absent_content_length_still_uses_eof_framing() {
let resp = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\nstream-until-close".to_vec();
let parsed = read(resp, 64 * 1024).expect("EOF-framed response must parse");
assert_eq!(parsed.status, 200);
assert_eq!(parsed.body, b"stream-until-close");
assert!(!parsed.keep_alive, "EOF-framed body forbids reuse");
}
#[test]
fn well_framed_response_under_limit_parses() {
let resp = b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello".to_vec();
let parsed = read(resp, 64 * 1024).expect("valid response must parse");
assert_eq!(parsed.status, 200);
assert_eq!(parsed.body, b"hello");
assert_eq!(parsed.header("content-length"), Some("5"));
}
}