1use std::io;
5
6use moirai_async::io::{AsyncReadExt, AsyncWrite, AsyncWriteExt};
7
8pub const DEFAULT_MAX_RESPONSE_BYTES: usize = 64 * 1024 * 1024;
14
15#[derive(Debug, Clone)]
17pub struct Response {
18 pub status: u16,
20 pub headers: Vec<(String, String)>,
22 pub body: Vec<u8>,
24 pub keep_alive: bool,
26}
27
28impl Response {
29 #[must_use]
31 pub fn header(&self, name: &str) -> Option<&str> {
32 let name = name.to_ascii_lowercase();
33 self.headers
34 .iter()
35 .find(|(k, _)| *k == name)
36 .map(|(_, v)| v.as_str())
37 }
38}
39
40pub async fn write_request<S: AsyncWrite + Unpin>(
47 stream: &mut S,
48 method: &str,
49 host: &str,
50 path: &str,
51 headers: &[(&str, &str)],
52 body: Option<&[u8]>,
53) -> io::Result<()> {
54 let mut req = Vec::with_capacity(256);
55 req.extend_from_slice(method.as_bytes());
56 req.push(b' ');
57 req.extend_from_slice(path.as_bytes());
58 req.extend_from_slice(b" HTTP/1.1\r\n");
59
60 let has = |n: &str| headers.iter().any(|(k, _)| k.eq_ignore_ascii_case(n));
61 if !has("host") {
62 req.extend_from_slice(format!("Host: {host}\r\n").as_bytes());
63 }
64 for (k, v) in headers {
65 req.extend_from_slice(format!("{k}: {v}\r\n").as_bytes());
66 }
67 if let Some(b) = body {
68 if !has("content-length") {
69 req.extend_from_slice(format!("Content-Length: {}\r\n", b.len()).as_bytes());
70 }
71 }
72 req.extend_from_slice(b"\r\n");
73 if let Some(b) = body {
74 req.extend_from_slice(b);
75 }
76
77 stream.write_all(&req).await?;
78 stream.flush().await
79}
80
81struct Buffered<'a, S> {
83 stream: &'a mut S,
84 buf: Vec<u8>,
85 pos: usize,
86 limit: usize,
90}
91
92impl<'a, S: AsyncReadExt + Unpin> Buffered<'a, S> {
93 fn new(stream: &'a mut S, limit: usize) -> Self {
94 Self {
95 stream,
96 buf: Vec::with_capacity(8192.min(limit.max(1))),
97 pos: 0,
98 limit,
99 }
100 }
101
102 fn available(&self) -> usize {
103 self.buf.len() - self.pos
104 }
105
106 async fn fill(&mut self) -> io::Result<usize> {
112 if self.buf.len() >= self.limit {
113 return Err(io::Error::new(
114 io::ErrorKind::InvalidData,
115 "response exceeds maximum size",
116 ));
117 }
118 let mut tmp = [0u8; 8192];
119 let n = self.stream.read(&mut tmp).await?;
120 self.buf.extend_from_slice(&tmp[..n]);
121 if self.buf.len() > self.limit {
122 return Err(io::Error::new(
123 io::ErrorKind::InvalidData,
124 "response exceeds maximum size",
125 ));
126 }
127 Ok(n)
128 }
129
130 async fn read_crlf_line(&mut self) -> io::Result<String> {
132 loop {
133 if let Some(rel) = find_crlf(&self.buf[self.pos..]) {
134 let line = self.buf[self.pos..self.pos + rel].to_vec();
135 self.pos += rel + 2;
136 return String::from_utf8(line).map_err(|_| {
137 io::Error::new(io::ErrorKind::InvalidData, "non-UTF8 header line")
138 });
139 }
140 if self.fill().await? == 0 {
141 return Err(eof("CRLF line"));
142 }
143 }
144 }
145
146 async fn read_n(&mut self, n: usize) -> io::Result<Vec<u8>> {
148 while self.available() < n {
149 if self.fill().await? == 0 {
150 return Err(eof("body"));
151 }
152 }
153 let out = self.buf[self.pos..self.pos + n].to_vec();
154 self.pos += n;
155 Ok(out)
156 }
157
158 async fn read_to_eof(&mut self) -> io::Result<Vec<u8>> {
160 while self.fill().await? != 0 {}
161 Ok(self.buf[self.pos..].to_vec())
162 }
163
164 async fn read_chunked(&mut self) -> io::Result<Vec<u8>> {
166 let mut body = Vec::new();
167 loop {
168 let line = self.read_crlf_line().await?;
169 let size_field = line.split(';').next().unwrap_or("").trim();
170 let size = usize::from_str_radix(size_field, 16)
171 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "bad chunk size"))?;
172 if size == 0 {
173 while !self.read_crlf_line().await?.is_empty() {}
175 break;
176 }
177 body.extend_from_slice(&self.read_n(size).await?);
178 let crlf = self.read_n(2).await?;
180 if crlf != b"\r\n" {
181 return Err(io::Error::new(
182 io::ErrorKind::InvalidData,
183 "missing chunk CRLF",
184 ));
185 }
186 }
187 Ok(body)
188 }
189}
190
191pub async fn read_response<S: AsyncReadExt + Unpin>(
196 stream: &mut S,
197 is_head: bool,
198 max_response_bytes: usize,
199) -> io::Result<Response> {
200 let mut r = Buffered::new(stream, max_response_bytes);
201
202 let (status, headers) = loop {
204 let mut header_storage = [httparse::EMPTY_HEADER; 96];
205 let mut resp = httparse::Response::new(&mut header_storage);
206 match resp.parse(&r.buf[r.pos..]) {
207 Ok(httparse::Status::Complete(consumed)) => {
208 let status = resp
209 .code
210 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "no status code"))?;
211 let headers: Vec<(String, String)> = resp
212 .headers
213 .iter()
214 .map(|h| {
215 (
216 h.name.to_ascii_lowercase(),
217 String::from_utf8_lossy(h.value).into_owned(),
218 )
219 })
220 .collect();
221 r.pos += consumed;
222 break (status, headers);
223 }
224 Ok(httparse::Status::Partial) => {
225 if r.fill().await? == 0 {
226 return Err(eof("response headers"));
227 }
228 }
229 Err(e) => {
230 return Err(io::Error::new(
231 io::ErrorKind::InvalidData,
232 format!("malformed response: {e}"),
233 ))
234 }
235 }
236 };
237
238 let find = |n: &str| {
239 headers
240 .iter()
241 .find(|(k, _)| k == n)
242 .map(|(_, v)| v.as_str())
243 };
244 let chunked = find("transfer-encoding")
245 .map(|v| v.to_ascii_lowercase().contains("chunked"))
246 .unwrap_or(false);
247 let content_length: Option<usize> = match find("content-length") {
251 Some(v) => Some(v.trim().parse().map_err(|_| {
252 io::Error::new(
253 io::ErrorKind::InvalidData,
254 format!("malformed Content-Length: {v:?}"),
255 )
256 })?),
257 None => None,
258 };
259 let conn_close = find("connection")
260 .map(|v| v.eq_ignore_ascii_case("close"))
261 .unwrap_or(false);
262 let bodyless = is_head || status == 204 || status == 304 || (100..200).contains(&status);
264
265 let (body, framed) = if bodyless {
266 (Vec::new(), true)
267 } else if chunked {
268 (r.read_chunked().await?, true)
269 } else if let Some(len) = content_length {
270 if len > max_response_bytes {
273 return Err(io::Error::new(
274 io::ErrorKind::InvalidData,
275 "Content-Length exceeds maximum response size",
276 ));
277 }
278 (r.read_n(len).await?, true)
279 } else {
280 (r.read_to_eof().await?, false)
282 };
283
284 Ok(Response {
285 status,
286 headers,
287 body,
288 keep_alive: framed && !conn_close,
289 })
290}
291
292fn find_crlf(buf: &[u8]) -> Option<usize> {
293 buf.windows(2).position(|w| w == b"\r\n")
294}
295
296fn eof(what: &str) -> io::Error {
297 io::Error::new(
298 io::ErrorKind::UnexpectedEof,
299 format!("connection closed while reading {what}"),
300 )
301}
302
303#[cfg(test)]
304mod tests {
305 use super::*;
306 use std::pin::Pin;
307 use std::task::{Context, Poll};
308
309 use moirai_async::io::AsyncRead;
310
311 struct MockReader {
314 data: Vec<u8>,
315 pos: usize,
316 }
317
318 impl MockReader {
319 fn new(data: Vec<u8>) -> Self {
320 Self { data, pos: 0 }
321 }
322 }
323
324 impl AsyncRead for MockReader {
325 fn poll_read(
326 mut self: Pin<&mut Self>,
327 _cx: &mut Context<'_>,
328 buf: &mut [u8],
329 ) -> Poll<io::Result<usize>> {
330 let remaining = self.data.len() - self.pos;
331 let n = remaining.min(buf.len());
332 buf[..n].copy_from_slice(&self.data[self.pos..self.pos + n]);
333 self.pos += n;
334 Poll::Ready(Ok(n))
335 }
336 }
337
338 fn read(data: Vec<u8>, max: usize) -> io::Result<Response> {
339 moirai::block_on(read_response(&mut MockReader::new(data), false, max))
340 }
341
342 #[test]
343 fn oversized_content_length_is_rejected_up_front() {
344 let resp = b"HTTP/1.1 200 OK\r\nContent-Length: 999999999\r\n\r\n".to_vec();
347 let err = read(resp, 4096).expect_err("oversized Content-Length must be rejected");
348 assert_eq!(err.kind(), io::ErrorKind::InvalidData);
349 }
350
351 #[test]
352 fn eof_delimited_body_over_limit_is_rejected() {
353 let mut resp = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_vec();
356 resp.extend(std::iter::repeat_n(b'x', 64 * 1024));
357 let err = read(resp, 8 * 1024).expect_err("EOF body over the cap must be rejected");
358 assert_eq!(err.kind(), io::ErrorKind::InvalidData);
359 }
360
361 #[test]
362 fn chunked_body_over_limit_is_rejected() {
363 let mut resp = b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n".to_vec();
366 for _ in 0..64 {
367 resp.extend_from_slice(b"1000\r\n"); resp.extend(std::iter::repeat_n(b'y', 0x1000));
369 resp.extend_from_slice(b"\r\n");
370 }
371 resp.extend_from_slice(b"0\r\n\r\n");
372 let err = read(resp, 16 * 1024).expect_err("chunked body over the cap must be rejected");
373 assert_eq!(err.kind(), io::ErrorKind::InvalidData);
374 }
375
376 #[test]
377 fn malformed_content_length_is_invalid_data_not_eof_framing() {
378 for bad in ["abc", "-5", "18446744073709551616", "12abc", ""] {
382 let resp =
383 format!("HTTP/1.1 200 OK\r\nContent-Length: {bad}\r\n\r\nhello").into_bytes();
384 let err = read(resp, 64 * 1024)
385 .expect_err("garbage Content-Length must be rejected, not EOF-framed");
386 assert_eq!(err.kind(), io::ErrorKind::InvalidData, "value: {bad:?}");
387 }
388 }
389
390 #[test]
391 fn absent_content_length_still_uses_eof_framing() {
392 let resp = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\nstream-until-close".to_vec();
395 let parsed = read(resp, 64 * 1024).expect("EOF-framed response must parse");
396 assert_eq!(parsed.status, 200);
397 assert_eq!(parsed.body, b"stream-until-close");
398 assert!(!parsed.keep_alive, "EOF-framed body forbids reuse");
399 }
400
401 #[test]
402 fn well_framed_response_under_limit_parses() {
403 let resp = b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello".to_vec();
406 let parsed = read(resp, 64 * 1024).expect("valid response must parse");
407 assert_eq!(parsed.status, 200);
408 assert_eq!(parsed.body, b"hello");
409 assert_eq!(parsed.header("content-length"), Some("5"));
410 }
411}