Skip to main content

moirai_http/
request.rs

1//! Bounded HTTP/1.1 request-head parsing for server protocols.
2
3use std::io;
4
5use moirai_async::io::AsyncReadExt;
6
7pub(crate) const MAX_HEADER_SLOTS: usize = 128;
8
9/// A validated HTTP request line and header block.
10#[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    /// Returns the request method.
20    #[must_use]
21    pub fn method(&self) -> &str {
22        &self.method
23    }
24
25    /// Returns the origin-form request target.
26    #[must_use]
27    pub fn target(&self) -> &str {
28        &self.target
29    }
30
31    /// Returns the HTTP version token.
32    #[must_use]
33    pub fn version(&self) -> &str {
34        &self.version
35    }
36
37    /// Returns the first header value matching `name` case-insensitively.
38    #[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    /// Returns all validated headers in receive order.
48    #[must_use]
49    pub fn headers(&self) -> &[(String, String)] {
50        &self.headers
51    }
52
53    /// Returns the browser page origin, when the request supplied one.
54    #[must_use]
55    pub fn origin(&self) -> Option<&str> {
56        self.header("origin")
57    }
58}
59
60/// Read one bounded HTTP request head and preserve bytes after its delimiter.
61pub(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}