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 && !has("content-length")
69 {
70 req.extend_from_slice(format!("Content-Length: {}\r\n", b.len()).as_bytes());
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().saturating_sub(self.pos)
106 }
107
108 async fn fill(&mut self) -> io::Result<usize> {
114 if self.buf.len() >= self.limit {
115 return Err(io::Error::new(
116 io::ErrorKind::InvalidData,
117 "response exceeds maximum size",
118 ));
119 }
120 let mut tmp = [0u8; 8192];
121 let n = self.stream.read(&mut tmp).await?;
122 #[expect(
123 clippy::indexing_slicing,
124 reason = "n <= tmp.len() per the Read trait contract"
125 )]
126 let src = &tmp[..n];
127 self.buf.extend_from_slice(src);
128 if self.buf.len() > self.limit {
129 return Err(io::Error::new(
130 io::ErrorKind::InvalidData,
131 "response exceeds maximum size",
132 ));
133 }
134 Ok(n)
135 }
136
137 async fn read_crlf_line(&mut self) -> io::Result<String> {
139 loop {
140 let (_, tail) = self.buf.split_at(self.pos);
143 if let Some(rel) = find_crlf(tail) {
144 #[expect(
145 clippy::indexing_slicing,
146 reason = "rel < tail.len() per the CRLF search above"
147 )]
148 let line = tail[..rel].to_vec();
149 self.pos = self
153 .pos
154 .checked_add(rel)
155 .and_then(|after_line| after_line.checked_add(2))
156 .expect("invariant: CRLF line fits inside the buffered prefix");
157 return String::from_utf8(line).map_err(|_| {
158 io::Error::new(io::ErrorKind::InvalidData, "non-UTF8 header line")
159 });
160 }
161 if self.fill().await? == 0 {
162 return Err(eof("CRLF line"));
163 }
164 }
165 }
166
167 async fn read_n(&mut self, n: usize) -> io::Result<Vec<u8>> {
169 while self.available() < n {
170 if self.fill().await? == 0 {
171 return Err(eof("body"));
172 }
173 }
174 let (_, rest) = self.buf.split_at(self.pos);
176 #[expect(
177 clippy::indexing_slicing,
178 reason = "rest.len() >= n follows from available() >= n"
179 )]
180 let out = rest[..n].to_vec();
181 self.pos = self
182 .pos
183 .checked_add(n)
184 .expect("invariant: n <= available() was established by the fill loop");
185 Ok(out)
186 }
187
188 async fn read_to_eof(&mut self) -> io::Result<Vec<u8>> {
190 while self.fill().await? != 0 {}
191 let (_, tail) = self.buf.split_at(self.pos);
192 Ok(tail.to_vec())
193 }
194
195 async fn read_chunked(&mut self) -> io::Result<Vec<u8>> {
197 let mut body = Vec::new();
198 loop {
199 let line = self.read_crlf_line().await?;
200 let size_field = line.split(';').next().unwrap_or("").trim();
201 let size = usize::from_str_radix(size_field, 16)
202 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "bad chunk size"))?;
203 if size == 0 {
204 while !self.read_crlf_line().await?.is_empty() {}
206 break;
207 }
208 body.extend_from_slice(&self.read_n(size).await?);
209 let crlf = self.read_n(2).await?;
211 if crlf != b"\r\n" {
212 return Err(io::Error::new(
213 io::ErrorKind::InvalidData,
214 "missing chunk CRLF",
215 ));
216 }
217 }
218 Ok(body)
219 }
220}
221
222pub async fn read_response<S: AsyncReadExt + Unpin>(
227 stream: &mut S,
228 is_head: bool,
229 max_response_bytes: usize,
230) -> io::Result<Response> {
231 let mut r = Buffered::new(stream, max_response_bytes);
232
233 let (status, headers) = loop {
235 let mut header_storage = [httparse::EMPTY_HEADER; 96];
236 let mut resp = httparse::Response::new(&mut header_storage);
237 let (_, tail) = r.buf.split_at(r.pos);
240 match resp.parse(tail) {
241 Ok(httparse::Status::Complete(consumed)) => {
242 let status = resp
243 .code
244 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "no status code"))?;
245 let headers: Vec<(String, String)> = resp
246 .headers
247 .iter()
248 .map(|h| {
249 (
250 h.name.to_ascii_lowercase(),
251 String::from_utf8_lossy(h.value).into_owned(),
252 )
253 })
254 .collect();
255 r.pos = r
258 .pos
259 .checked_add(consumed)
260 .expect("invariant: header parse cannot consume beyond the buffer");
261 break (status, headers);
262 }
263 Ok(httparse::Status::Partial) => {
264 if r.fill().await? == 0 {
265 return Err(eof("response headers"));
266 }
267 }
268 Err(e) => {
269 return Err(io::Error::new(
270 io::ErrorKind::InvalidData,
271 format!("malformed response: {e}"),
272 ));
273 }
274 }
275 };
276
277 let find = |n: &str| {
278 headers
279 .iter()
280 .find(|(k, _)| k == n)
281 .map(|(_, v)| v.as_str())
282 };
283 let chunked = find("transfer-encoding")
284 .map(|v| v.to_ascii_lowercase().contains("chunked"))
285 .unwrap_or(false);
286 let content_length: Option<usize> = match find("content-length") {
290 Some(v) => Some(v.trim().parse().map_err(|_| {
291 io::Error::new(
292 io::ErrorKind::InvalidData,
293 format!("malformed Content-Length: {v:?}"),
294 )
295 })?),
296 None => None,
297 };
298 let conn_close = find("connection")
299 .map(|v| v.eq_ignore_ascii_case("close"))
300 .unwrap_or(false);
301 let bodyless = is_head || status == 204 || status == 304 || (100..200).contains(&status);
303
304 let (body, framed) = if bodyless {
305 (Vec::new(), true)
306 } else if chunked {
307 (r.read_chunked().await?, true)
308 } else if let Some(len) = content_length {
309 if len > max_response_bytes {
312 return Err(io::Error::new(
313 io::ErrorKind::InvalidData,
314 "Content-Length exceeds maximum response size",
315 ));
316 }
317 (r.read_n(len).await?, true)
318 } else {
319 (r.read_to_eof().await?, false)
321 };
322
323 Ok(Response {
324 status,
325 headers,
326 body,
327 keep_alive: framed && !conn_close,
328 })
329}
330
331fn find_crlf(buf: &[u8]) -> Option<usize> {
332 buf.windows(2).position(|w| w == b"\r\n")
333}
334
335fn eof(what: &str) -> io::Error {
336 io::Error::new(
337 io::ErrorKind::UnexpectedEof,
338 format!("connection closed while reading {what}"),
339 )
340}
341
342#[cfg(test)]
343mod tests {
344 use super::*;
345 use std::pin::Pin;
346 use std::task::{Context, Poll};
347
348 use moirai_async::io::AsyncRead;
349
350 struct MockReader {
353 data: Vec<u8>,
354 pos: usize,
355 }
356
357 impl MockReader {
358 fn new(data: Vec<u8>) -> Self {
359 Self { data, pos: 0 }
360 }
361 }
362
363 impl AsyncRead for MockReader {
364 fn poll_read(
365 mut self: Pin<&mut Self>,
366 _cx: &mut Context<'_>,
367 buf: &mut [u8],
368 ) -> Poll<io::Result<usize>> {
369 let (_, rest) = self.data.split_at(self.pos);
370 let n = rest.len().min(buf.len());
371 #[expect(clippy::indexing_slicing, reason = "n <= buf.len() by the min above")]
372 let dst = &mut buf[..n];
373 #[expect(clippy::indexing_slicing, reason = "n <= rest.len() by the min above")]
374 let src = &rest[..n];
375 dst.copy_from_slice(src);
376 self.pos = self
377 .pos
378 .checked_add(n)
379 .expect("invariant: n <= data.len() - pos");
380 Poll::Ready(Ok(n))
381 }
382 }
383
384 fn read(data: Vec<u8>, max: usize) -> io::Result<Response> {
385 moirai::block_on(read_response(&mut MockReader::new(data), false, max))
386 }
387
388 #[test]
389 fn oversized_content_length_is_rejected_up_front() {
390 let resp = b"HTTP/1.1 200 OK\r\nContent-Length: 999999999\r\n\r\n".to_vec();
393 let err = read(resp, 4096).expect_err("oversized Content-Length must be rejected");
394 assert_eq!(err.kind(), io::ErrorKind::InvalidData);
395 }
396
397 #[test]
398 fn eof_delimited_body_over_limit_is_rejected() {
399 let mut resp = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_vec();
402 resp.extend(std::iter::repeat_n(b'x', 64 * 1024));
403 let err = read(resp, 8 * 1024).expect_err("EOF body over the cap must be rejected");
404 assert_eq!(err.kind(), io::ErrorKind::InvalidData);
405 }
406
407 #[test]
408 fn chunked_body_over_limit_is_rejected() {
409 let mut resp = b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n".to_vec();
412 for _ in 0..64 {
413 resp.extend_from_slice(b"1000\r\n"); resp.extend(std::iter::repeat_n(b'y', 0x1000));
415 resp.extend_from_slice(b"\r\n");
416 }
417 resp.extend_from_slice(b"0\r\n\r\n");
418 let err = read(resp, 16 * 1024).expect_err("chunked body over the cap must be rejected");
419 assert_eq!(err.kind(), io::ErrorKind::InvalidData);
420 }
421
422 #[test]
423 fn malformed_content_length_is_invalid_data_not_eof_framing() {
424 for bad in ["abc", "-5", "18446744073709551616", "12abc", ""] {
428 let resp =
429 format!("HTTP/1.1 200 OK\r\nContent-Length: {bad}\r\n\r\nhello").into_bytes();
430 let err = read(resp, 64 * 1024)
431 .expect_err("garbage Content-Length must be rejected, not EOF-framed");
432 assert_eq!(err.kind(), io::ErrorKind::InvalidData, "value: {bad:?}");
433 }
434 }
435
436 #[test]
437 fn absent_content_length_still_uses_eof_framing() {
438 let resp = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\nstream-until-close".to_vec();
441 let parsed = read(resp, 64 * 1024).expect("EOF-framed response must parse");
442 assert_eq!(parsed.status, 200);
443 assert_eq!(parsed.body, b"stream-until-close");
444 assert!(!parsed.keep_alive, "EOF-framed body forbids reuse");
445 }
446
447 #[test]
448 fn well_framed_response_under_limit_parses() {
449 let resp = b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello".to_vec();
452 let parsed = read(resp, 64 * 1024).expect("valid response must parse");
453 assert_eq!(parsed.status, 200);
454 assert_eq!(parsed.body, b"hello");
455 assert_eq!(parsed.header("content-length"), Some("5"));
456 }
457}