1use std::io;
4
5use moirai_async::io::AsyncReadExt;
6
7pub(crate) const MAX_HEADER_SLOTS: usize = 128;
8
9#[derive(Debug, Clone, PartialEq, Eq)]
11pub struct HttpRequestHead {
12 method: String,
13 target: String,
14 version: String,
15 headers: Vec<(String, String)>,
16}
17
18impl HttpRequestHead {
19 #[must_use]
21 pub fn method(&self) -> &str {
22 &self.method
23 }
24
25 #[must_use]
27 pub fn target(&self) -> &str {
28 &self.target
29 }
30
31 #[must_use]
33 pub fn version(&self) -> &str {
34 &self.version
35 }
36
37 #[must_use]
39 pub fn header(&self, name: &str) -> Option<&str> {
40 let name = name.to_ascii_lowercase();
41 self.headers
42 .iter()
43 .find(|(key, _)| key == &name)
44 .map(|(_, value)| value.as_str())
45 }
46
47 #[must_use]
49 pub fn headers(&self) -> &[(String, String)] {
50 &self.headers
51 }
52
53 #[must_use]
55 pub fn origin(&self) -> Option<&str> {
56 self.header("origin")
57 }
58}
59
60pub(crate) async fn read_request_head<S: AsyncReadExt + Unpin>(
62 stream: &mut S,
63 max_header_bytes: usize,
64 max_header_count: usize,
65) -> io::Result<(HttpRequestHead, Vec<u8>)> {
66 if max_header_bytes == 0 || max_header_count == 0 || max_header_count > MAX_HEADER_SLOTS {
67 return Err(io::Error::new(
68 io::ErrorKind::InvalidInput,
69 "HTTP request limits are outside the supported bounds",
70 ));
71 }
72
73 let mut bytes = Vec::with_capacity(max_header_bytes.min(4096));
74 loop {
75 if let Some(consumed) = header_end(&bytes) {
76 let head = bytes
77 .get(..consumed)
78 .ok_or_else(|| io::Error::other("request delimiter exceeded buffered bytes"))?;
79 let remainder = bytes
80 .get(consumed..)
81 .ok_or_else(|| io::Error::other("request remainder exceeded buffered bytes"))?
82 .to_vec();
83 return parse_request_head(head, remainder, max_header_count);
84 }
85 if bytes.len() >= max_header_bytes {
86 return Err(io::Error::new(
87 io::ErrorKind::InvalidData,
88 "HTTP request headers exceed configured byte bound",
89 ));
90 }
91
92 let available = max_header_bytes
93 .checked_sub(bytes.len())
94 .ok_or_else(|| io::Error::other("HTTP request buffer exceeded its configured bound"))?;
95 let mut chunk = [0u8; 1024];
96 let read_len = chunk.len().min(available);
97 let target = chunk
98 .get_mut(..read_len)
99 .ok_or_else(|| io::Error::other("HTTP read slice exceeded its buffer"))?;
100 let count = stream.read(target).await?;
101 if count == 0 {
102 return Err(io::Error::new(
103 io::ErrorKind::UnexpectedEof,
104 "connection closed before HTTP request headers completed",
105 ));
106 }
107 if count > available {
108 return Err(io::Error::new(
109 io::ErrorKind::InvalidData,
110 "HTTP request headers exceed configured byte bound",
111 ));
112 }
113 let chunk = chunk
114 .get(..count)
115 .ok_or_else(|| io::Error::other("HTTP read count exceeded buffer"))?;
116 bytes.extend_from_slice(chunk);
117 }
118}
119
120fn parse_request_head(
121 bytes: &[u8],
122 remainder: Vec<u8>,
123 max_header_count: usize,
124) -> io::Result<(HttpRequestHead, Vec<u8>)> {
125 let mut storage = [httparse::EMPTY_HEADER; MAX_HEADER_SLOTS];
126 let mut request = httparse::Request::new(&mut storage);
127 let parsed = request.parse(bytes).map_err(|error| {
128 io::Error::new(
129 io::ErrorKind::InvalidData,
130 format!("malformed HTTP request head: {error}"),
131 )
132 })?;
133 let consumed = match parsed {
134 httparse::Status::Complete(consumed) => consumed,
135 httparse::Status::Partial => {
136 return Err(io::Error::new(
137 io::ErrorKind::InvalidData,
138 "HTTP request delimiter was not parsed",
139 ));
140 }
141 };
142 if consumed != bytes.len() {
143 return Err(io::Error::new(
144 io::ErrorKind::InvalidData,
145 "HTTP request parser did not consume the complete head",
146 ));
147 }
148 if request.headers.len() > max_header_count {
149 return Err(io::Error::new(
150 io::ErrorKind::InvalidData,
151 "HTTP request header count exceeds configured bound",
152 ));
153 }
154
155 let method = request
156 .method
157 .ok_or_else(|| invalid_request("HTTP request has no method"))?;
158 let target = request
159 .path
160 .ok_or_else(|| invalid_request("HTTP request has no target"))?;
161 let version = request
162 .version
163 .ok_or_else(|| invalid_request("HTTP request has no version"))?;
164 let version = format!("HTTP/1.{version}");
165 if !target.starts_with('/') {
166 return Err(invalid_request("HTTP request target must be origin-form"));
167 }
168 validate_token(method, "HTTP method")?;
169 validate_target(target)?;
170
171 let mut headers = Vec::with_capacity(request.headers.len());
172 for header in request.headers {
173 validate_header_name(header.name)?;
174 let value = std::str::from_utf8(header.value)
175 .map_err(|_| invalid_request("HTTP header value is not ASCII UTF-8"))?;
176 if value
177 .bytes()
178 .any(|byte| byte < 0x20 && byte != b'\t' || byte == 0x7f)
179 {
180 return Err(invalid_request("HTTP header value contains a control byte"));
181 }
182 let name = header.name.to_ascii_lowercase();
183 if headers.iter().any(|(existing, _)| existing == &name) {
184 return Err(invalid_request("duplicate HTTP header is not accepted"));
185 }
186 headers.push((name, value.trim().to_owned()));
187 }
188
189 Ok((
190 HttpRequestHead {
191 method: method.to_owned(),
192 target: target.to_owned(),
193 version,
194 headers,
195 },
196 remainder,
197 ))
198}
199
200fn validate_token(value: &str, what: &str) -> io::Result<()> {
201 if value.is_empty()
202 || !value
203 .bytes()
204 .all(|byte| byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(&byte))
205 {
206 return Err(invalid_request(what));
207 }
208 Ok(())
209}
210
211fn validate_header_name(name: &str) -> io::Result<()> {
212 validate_token(name, "HTTP header name is not a token")
213}
214
215fn validate_target(target: &str) -> io::Result<()> {
216 if target
217 .bytes()
218 .any(|byte| byte < 0x20 || byte == 0x7f || byte == b' ')
219 {
220 return Err(invalid_request(
221 "HTTP request target contains a control byte",
222 ));
223 }
224 Ok(())
225}
226
227fn invalid_request(message: &str) -> io::Error {
228 io::Error::new(io::ErrorKind::InvalidData, message)
229}
230
231fn header_end(bytes: &[u8]) -> Option<usize> {
232 bytes
233 .windows(4)
234 .position(|window| window == b"\r\n\r\n")
235 .and_then(|position| position.checked_add(4))
236}
237
238#[cfg(test)]
239mod tests {
240 use super::*;
241 use std::collections::VecDeque;
242 use std::pin::Pin;
243 use std::task::{Context, Poll};
244
245 struct Input {
246 bytes: VecDeque<u8>,
247 }
248
249 impl Input {
250 fn new(bytes: &[u8]) -> Self {
251 Self {
252 bytes: bytes.iter().copied().collect(),
253 }
254 }
255 }
256
257 impl moirai_async::io::AsyncRead for Input {
258 fn poll_read(
259 mut self: Pin<&mut Self>,
260 _cx: &mut Context<'_>,
261 output: &mut [u8],
262 ) -> Poll<io::Result<usize>> {
263 let count = output.len().min(self.bytes.len()).min(3);
264 for slot in output.iter_mut().take(count) {
265 let Some(byte) = self.bytes.pop_front() else {
266 return Poll::Ready(Ok(0));
267 };
268 *slot = byte;
269 }
270 Poll::Ready(Ok(count))
271 }
272 }
273
274 #[test]
275 fn parser_preserves_buffered_bytes_after_request_head() {
276 let mut input = Input::new(b"GET /socket HTTP/1.1\r\nUpgrade: websocket\r\n\r\nframe");
277 let (head, remainder) = moirai::block_on(read_request_head(&mut input, 512, 8))
278 .expect("request head must parse");
279 assert_eq!(head.method(), "GET");
280 assert_eq!(head.target(), "/socket");
281 assert_eq!(head.version(), "HTTP/1.1");
282 assert_eq!(head.header("upgrade"), Some("websocket"));
283 assert_eq!(remainder, b"f");
284 }
285
286 #[test]
287 fn parser_rejects_duplicate_headers_and_oversized_heads() {
288 for input in [
289 b"GET / HTTP/1.1\r\nX-Test: one\r\nx-test: two\r\n\r\n".as_slice(),
290 b"GET / HTTP/1.1\r\nX-Test: one".as_slice(),
291 ] {
292 let mut reader = Input::new(input);
293 let error = moirai::block_on(read_request_head(&mut reader, 32, 8))
294 .expect_err("invalid head must fail");
295 assert!(matches!(
296 error.kind(),
297 io::ErrorKind::InvalidData | io::ErrorKind::UnexpectedEof
298 ));
299 }
300 }
301
302 #[test]
303 fn parser_accepts_near_limit_head_with_pipelined_bytes() {
304 let prefix = b"GET / HTTP/1.1\r\nX-Pad: ";
305 let suffix = b"\r\n\r\n";
306 let target_head_length: usize = 1024;
307 let padding = target_head_length
308 .checked_sub(prefix.len() + suffix.len())
309 .expect("test head target exceeds fixed prefix");
310 let mut input = Vec::with_capacity(target_head_length + 4);
311 input.extend_from_slice(prefix);
312 input.extend(std::iter::repeat_n(b'x', padding));
313 input.extend_from_slice(suffix);
314 input.extend_from_slice(b"next");
315
316 let mut reader = Input::new(&input);
317 let (head, remainder) = moirai::block_on(read_request_head(&mut reader, 1024, 8))
318 .expect("near-limit head must parse");
319 assert_eq!(head.target(), "/");
320 assert!(remainder.is_empty());
321 }
322}