1use std::{cmp::min, str::FromStr};
8
9use bytes::{Buf, Bytes, BytesMut};
10use thiserror::Error;
11use tokio::io::{AsyncRead, AsyncReadExt};
12
13use crate::{
14 body::Body,
15 common::{HttpVersion, Method},
16 headers::{HeaderName, Headers, header_keys},
17 request::{Request, RequestTarget},
18};
19
20const HEADER_LIMIT: usize = 64 * 1024;
21const LINE_LIMIT: usize = 8 * 1024;
22
23pub fn needs_more_head(buffer: &[u8]) -> bool {
25 find_headers_end(buffer).is_none()
26}
27
28pub fn parse_request_head(buffer: &[u8]) -> Result<(Request, BodyMode, usize), ParseError> {
33 let head_end = find_headers_end(buffer).ok_or(ParseError::Incomplete)?;
34 if head_end > HEADER_LIMIT {
35 return Err(ParseError::HeaderTooLarge);
36 }
37
38 let head_bytes = &buffer[..head_end];
39 let head_str =
40 std::str::from_utf8(head_bytes).map_err(|_| ParseError::InvalidHeaderEncoding)?;
41
42 let mut lines = head_str.split("\r\n");
43 let request_line = lines.next().ok_or(ParseError::InvalidRequestLine)?;
44 if request_line.len() > LINE_LIMIT {
45 return Err(ParseError::RequestLineTooLong);
46 }
47
48 let (method, target, version) = parse_request_line(request_line)?;
49
50 let mut headers = Headers::new();
51 let mut host_value: Option<String> = None;
52 let mut host_seen = false;
53 let mut content_length: Option<u64> = None;
54 let mut transfer_encodings: Vec<String> = Vec::new();
55
56 for line in lines {
57 if line.is_empty() {
58 break;
59 }
60 if line.len() > LINE_LIMIT {
61 return Err(ParseError::HeaderLineTooLong);
62 }
63 if line.starts_with(' ') || line.starts_with('\t') {
64 return Err(ParseError::ObsoleteLineFolding);
65 }
66
67 let (name_str, value_str) = split_header_line(line)?;
68 if !is_field_name(name_str) {
69 return Err(ParseError::InvalidHeaderName);
70 }
71
72 let value = value_str.trim_matches(|c| matches!(c, ' ' | '\t'));
73 if value.len() > LINE_LIMIT {
74 return Err(ParseError::HeaderLineTooLong);
75 }
76 if contains_invalid_header_value(value) {
77 return Err(ParseError::InvalidHeaderValue);
78 }
79
80 let name = HeaderName::new(name_str);
81 let name_key = name.as_str();
82
83 match name_key {
84 header_keys::HOST => {
85 if host_seen {
86 return Err(ParseError::MultipleHostValues);
87 }
88 host_seen = true;
89 if !is_valid_host(value) {
90 return Err(ParseError::InvalidHost);
91 }
92 host_value = Some(value.to_string());
93 }
94 header_keys::CONTENT_LENGTH => {
95 let length = parse_content_length(value)?;
96 if let Some(existing) = content_length {
97 if existing != length {
98 return Err(ParseError::ConflictingContentLength);
99 }
100 } else {
101 content_length = Some(length);
102 }
103 }
104 header_keys::TRANSFER_ENCODING => {
105 let codings = parse_transfer_encoding(value)?;
106 transfer_encodings.extend(codings);
107 }
108 _ => {}
109 }
110
111 headers.append(name, value.to_string());
112 }
113
114 if host_value.is_none() {
115 return Err(ParseError::MissingHost);
116 }
117
118 let body_mode = determine_body_mode(&method, content_length, &transfer_encodings)?;
119
120 let mut request = Request::new(method, target);
121 request.set_version(version);
122 request.set_body(match body_mode {
123 BodyMode::None => Body::Empty,
124 BodyMode::Fixed(len) => Body::Fixed(len),
125 BodyMode::Chunked => Body::Chunked,
126 });
127 *request.headers_mut() = headers;
128
129 Ok((request, body_mode, head_end))
130}
131
132#[derive(Debug, Clone, Copy, PartialEq, Eq)]
134pub enum BodyMode {
135 None,
136 Fixed(u64),
137 Chunked,
138}
139
140pub fn body_reader<'a, R>(
142 mode: BodyMode,
143 reader: &'a mut R,
144 buffer: &'a mut BytesMut,
145) -> BodyReader<'a, R>
146where
147 R: AsyncRead + Unpin,
148{
149 match mode {
150 BodyMode::None => BodyReader {
151 inner: BodyReaderInner::Empty,
152 },
153 BodyMode::Fixed(remaining) => BodyReader {
154 inner: BodyReaderInner::Fixed(FixedBodyReader {
155 reader,
156 buffer,
157 remaining,
158 }),
159 },
160 BodyMode::Chunked => BodyReader {
161 inner: BodyReaderInner::Chunked(ChunkedBodyReader::new(reader, buffer)),
162 },
163 }
164}
165
166pub struct BodyReader<'a, R>
168where
169 R: AsyncRead + Unpin,
170{
171 inner: BodyReaderInner<'a, R>,
172}
173
174impl<'a, R> BodyReader<'a, R>
175where
176 R: AsyncRead + Unpin,
177{
178 pub async fn read_next(&mut self) -> Result<Option<Bytes>, BodyError> {
181 match &mut self.inner {
182 BodyReaderInner::Empty => Ok(None),
183 BodyReaderInner::Fixed(inner) => inner.read_next().await,
184 BodyReaderInner::Chunked(inner) => inner.read_next().await,
185 }
186 }
187
188 pub async fn drain(&mut self) -> Result<(), BodyError> {
190 while self.read_next().await?.is_some() {}
191 Ok(())
192 }
193
194 pub fn trailers(&self) -> Option<&Headers> {
196 match &self.inner {
197 BodyReaderInner::Chunked(inner) if inner.trailers_complete => Some(&inner.trailers),
198 _ => None,
199 }
200 }
201
202 pub fn is_finished(&self) -> bool {
204 match &self.inner {
205 BodyReaderInner::Empty => true,
206 BodyReaderInner::Fixed(inner) => inner.remaining == 0,
207 BodyReaderInner::Chunked(inner) => inner.state == ChunkState::Done,
208 }
209 }
210}
211
212enum BodyReaderInner<'a, R>
213where
214 R: AsyncRead + Unpin,
215{
216 Empty,
217 Fixed(FixedBodyReader<'a, R>),
218 Chunked(ChunkedBodyReader<'a, R>),
219}
220
221struct FixedBodyReader<'a, R>
222where
223 R: AsyncRead + Unpin,
224{
225 reader: &'a mut R,
226 buffer: &'a mut BytesMut,
227 remaining: u64,
228}
229
230impl<'a, R> FixedBodyReader<'a, R>
231where
232 R: AsyncRead + Unpin,
233{
234 async fn read_next(&mut self) -> Result<Option<Bytes>, BodyError> {
235 if self.remaining == 0 {
236 return Ok(None);
237 }
238
239 if !self.buffer.is_empty() {
240 let available = min(self.buffer.len() as u64, self.remaining) as usize;
241 let chunk = self.buffer.split_to(available).freeze();
242 self.remaining -= available as u64;
243 return Ok(Some(chunk));
244 }
245
246 let to_read = min(self.remaining, 8 * 1024) as usize;
247 let mut temp = vec![0_u8; to_read];
248 let read = self.reader.read(&mut temp).await?;
249 if read == 0 {
250 return Err(BodyError::UnexpectedEof);
251 }
252 self.remaining -= read as u64;
253 temp.truncate(read);
254 Ok(Some(Bytes::from(temp)))
255 }
256}
257
258struct ChunkedBodyReader<'a, R>
259where
260 R: AsyncRead + Unpin,
261{
262 reader: &'a mut R,
263 buffer: &'a mut BytesMut,
264 state: ChunkState,
265 current_chunk_remaining: u64,
266 trailers: Headers,
267 trailers_complete: bool,
268}
269
270impl<'a, R> ChunkedBodyReader<'a, R>
271where
272 R: AsyncRead + Unpin,
273{
274 fn new(reader: &'a mut R, buffer: &'a mut BytesMut) -> Self {
275 Self {
276 reader,
277 buffer,
278 state: ChunkState::ReadingSize,
279 current_chunk_remaining: 0,
280 trailers: Headers::new(),
281 trailers_complete: false,
282 }
283 }
284
285 async fn read_next(&mut self) -> Result<Option<Bytes>, BodyError> {
286 loop {
287 match self.state {
288 ChunkState::ReadingSize => {
289 let line = self.read_line().await?;
290 let size = parse_chunk_size(&line)?;
291 if size == 0 {
292 self.state = ChunkState::ReadingTrailers;
293 } else {
294 self.current_chunk_remaining = size;
295 self.state = ChunkState::ReadingData;
296 }
297 }
298 ChunkState::ReadingData => {
299 if self.current_chunk_remaining == 0 {
300 self.state = ChunkState::ExpectingCrLf;
301 continue;
302 }
303
304 if !self.buffer.is_empty() {
305 let available =
306 min(self.buffer.len() as u64, self.current_chunk_remaining) as usize;
307 let chunk = self.buffer.split_to(available).freeze();
308 self.current_chunk_remaining -= available as u64;
309 if self.current_chunk_remaining == 0 {
310 self.state = ChunkState::ExpectingCrLf;
311 }
312 return Ok(Some(chunk));
313 }
314
315 let to_read = min(self.current_chunk_remaining, 8 * 1024) as usize;
316 let mut temp = vec![0_u8; to_read];
317 let read = self.reader.read(&mut temp).await?;
318 if read == 0 {
319 return Err(BodyError::UnexpectedEof);
320 }
321 self.current_chunk_remaining -= read as u64;
322 temp.truncate(read);
323 if self.current_chunk_remaining == 0 {
324 self.state = ChunkState::ExpectingCrLf;
325 }
326 return Ok(Some(Bytes::from(temp)));
327 }
328 ChunkState::ExpectingCrLf => {
329 if self.buffer.len() < 2 {
330 let read = self.reader.read_buf(self.buffer).await?;
331 if read == 0 {
332 return Err(BodyError::UnexpectedEof);
333 }
334 continue;
335 }
336 if &self.buffer[..2] != b"\r\n" {
337 return Err(BodyError::InvalidChunk);
338 }
339 self.buffer.advance(2);
340 self.state = ChunkState::ReadingSize;
341 }
342 ChunkState::ReadingTrailers => {
343 let line = self.read_line().await?;
344 if line.is_empty() {
345 self.state = ChunkState::Done;
346 self.trailers_complete = true;
347 return Ok(None);
348 }
349 let (name_str, value_str) =
350 split_header_line(&line).map_err(BodyError::InvalidTrailer)?;
351 if !is_field_name(name_str) {
352 return Err(BodyError::InvalidTrailer(ParseError::InvalidHeaderName));
353 }
354 let value = value_str.trim_matches(|c| matches!(c, ' ' | '\t'));
355 if contains_invalid_header_value(value) {
356 return Err(BodyError::InvalidTrailer(ParseError::InvalidHeaderValue));
357 }
358 self.trailers
359 .append(HeaderName::new(name_str), value.to_string());
360 }
361 ChunkState::Done => return Ok(None),
362 }
363 }
364 }
365
366 async fn read_line(&mut self) -> Result<String, BodyError> {
367 loop {
368 if let Some(pos) = find_crlf(self.buffer) {
369 let mut line = self.buffer.split_to(pos + 2);
370 line.truncate(pos);
371 return String::from_utf8(line.to_vec()).map_err(|_| BodyError::InvalidChunk);
372 }
373 let read = self.reader.read_buf(self.buffer).await?;
374 if read == 0 {
375 return Err(BodyError::UnexpectedEof);
376 }
377 }
378 }
379}
380
381#[derive(Debug, Clone, Copy, PartialEq, Eq)]
382enum ChunkState {
383 ReadingSize,
384 ReadingData,
385 ExpectingCrLf,
386 ReadingTrailers,
387 Done,
388}
389
390#[derive(Debug, Error, PartialEq, Eq)]
392pub enum ParseError {
393 #[error("request head incomplete")]
394 Incomplete,
395 #[error("request line too long")]
396 RequestLineTooLong,
397 #[error("invalid request line")]
398 InvalidRequestLine,
399 #[error("invalid method token")]
400 InvalidMethod,
401 #[error("invalid HTTP version")]
402 InvalidVersion,
403 #[error("invalid request target")]
404 InvalidRequestTarget,
405 #[error("header section exceeds limit")]
406 HeaderTooLarge,
407 #[error("header line exceeds limit")]
408 HeaderLineTooLong,
409 #[error("obsolete line folding detected")]
410 ObsoleteLineFolding,
411 #[error("invalid header name")]
412 InvalidHeaderName,
413 #[error("invalid header value")]
414 InvalidHeaderValue,
415 #[error("invalid Host header value")]
416 InvalidHost,
417 #[error("multiple Host header values are not allowed")]
418 MultipleHostValues,
419 #[error("required Host header missing")]
420 MissingHost,
421 #[error("invalid Content-Length value")]
422 InvalidContentLength,
423 #[error("conflicting Content-Length values")]
424 ConflictingContentLength,
425 #[error("invalid Transfer-Encoding value")]
426 InvalidTransferEncoding,
427 #[error("Transfer-Encoding and Content-Length conflict")]
428 ConflictingLengthAndTransferEncoding,
429 #[error("unsupported Transfer-Encoding")]
430 UnsupportedTransferEncoding,
431 #[error("invalid header encoding")]
432 InvalidHeaderEncoding,
433}
434
435#[derive(Debug, Error)]
437pub enum BodyError {
438 #[error("unexpected end of stream")]
439 UnexpectedEof,
440 #[error("invalid chunked body")]
441 InvalidChunk,
442 #[error("invalid trailer: {0}")]
443 InvalidTrailer(ParseError),
444 #[error(transparent)]
445 Io(#[from] std::io::Error),
446}
447
448fn parse_request_line(line: &str) -> Result<(Method, RequestTarget, HttpVersion), ParseError> {
449 let bytes = line.as_bytes();
450 let first_space = bytes
451 .iter()
452 .position(|&b| b == b' ')
453 .ok_or(ParseError::InvalidRequestLine)?;
454 let method_str = &line[..first_space];
455 if method_str.is_empty() {
456 return Err(ParseError::InvalidMethod);
457 }
458
459 let rest = &line[first_space + 1..];
460 let second_space = rest
461 .as_bytes()
462 .iter()
463 .position(|&b| b == b' ')
464 .ok_or(ParseError::InvalidRequestLine)?;
465 let target_str = &rest[..second_space];
466 if target_str.is_empty() {
467 return Err(ParseError::InvalidRequestTarget);
468 }
469
470 let version_str = &rest[second_space + 1..];
471 if version_str.is_empty() || version_str.contains(' ') {
472 return Err(ParseError::InvalidVersion);
473 }
474
475 let method = Method::from_str(method_str).map_err(|_| ParseError::InvalidMethod)?;
476 let target = parse_request_target(&method, target_str)?;
477 let version = HttpVersion::from_str(version_str).map_err(|_| ParseError::InvalidVersion)?;
478 if !matches!(version, HttpVersion::Http10 | HttpVersion::Http11) {
479 return Err(ParseError::InvalidVersion);
480 }
481
482 Ok((method, target, version))
483}
484
485fn parse_request_target(method: &Method, target: &str) -> Result<RequestTarget, ParseError> {
486 if target == "*" {
487 if matches!(method, Method::Options) {
488 return Ok(RequestTarget::Asterisk);
489 }
490 return Err(ParseError::InvalidRequestTarget);
491 }
492
493 if target.starts_with('/') {
494 return Ok(RequestTarget::origin(target.to_string()));
495 }
496
497 if matches!(method, Method::Connect) {
498 if is_authority_form(target) {
499 return Ok(RequestTarget::Authority(target.to_string()));
500 }
501 return Err(ParseError::InvalidRequestTarget);
502 }
503
504 if target.contains("://") {
505 return Ok(RequestTarget::Absolute(target.to_string()));
506 }
507
508 Ok(RequestTarget::Origin(target.to_string()))
509}
510
511fn split_header_line(line: &str) -> Result<(&str, &str), ParseError> {
512 let (name, value) = line.split_once(':').ok_or(ParseError::InvalidHeaderName)?;
513 Ok((name, value))
514}
515
516fn parse_content_length(value: &str) -> Result<u64, ParseError> {
517 if value.is_empty() {
518 return Err(ParseError::InvalidContentLength);
519 }
520 value
521 .parse::<u64>()
522 .map_err(|_| ParseError::InvalidContentLength)
523}
524
525fn parse_transfer_encoding(value: &str) -> Result<Vec<String>, ParseError> {
526 let mut codings = Vec::new();
527 for coding in value.split(',') {
528 let token = coding.trim();
529 if token.is_empty() || !is_token(token) {
530 return Err(ParseError::InvalidTransferEncoding);
531 }
532 codings.push(token.to_ascii_lowercase());
533 }
534 if codings.is_empty() {
535 return Err(ParseError::InvalidTransferEncoding);
536 }
537 Ok(codings)
538}
539
540fn determine_body_mode(
541 method: &Method,
542 content_length: Option<u64>,
543 transfer_encodings: &[String],
544) -> Result<BodyMode, ParseError> {
545 if !transfer_encodings.is_empty() {
546 let chunked_positions: Vec<usize> = transfer_encodings
547 .iter()
548 .enumerate()
549 .filter_map(|(idx, coding)| coding.eq_ignore_ascii_case("chunked").then_some(idx))
550 .collect();
551
552 if chunked_positions.is_empty() {
553 return Err(ParseError::UnsupportedTransferEncoding);
554 }
555 if chunked_positions.len() > 1
556 || *chunked_positions.last().unwrap() != transfer_encodings.len() - 1
557 {
558 return Err(ParseError::InvalidTransferEncoding);
559 }
560 if content_length.is_some() {
561 return Err(ParseError::ConflictingLengthAndTransferEncoding);
562 }
563 return Ok(BodyMode::Chunked);
564 }
565
566 if let Some(length) = content_length {
567 return Ok(BodyMode::Fixed(length));
568 }
569
570 if matches!(method, Method::Get | Method::Head | Method::Trace) {
571 Ok(BodyMode::None)
572 } else {
573 Ok(BodyMode::None)
574 }
575}
576
577fn find_headers_end(buffer: &[u8]) -> Option<usize> {
578 buffer
579 .windows(4)
580 .position(|window| window == b"\r\n\r\n")
581 .map(|idx| idx + 4)
582}
583
584fn find_crlf(buffer: &BytesMut) -> Option<usize> {
585 buffer.windows(2).position(|window| window == b"\r\n")
586}
587
588fn is_field_name(name: &str) -> bool {
589 !name.is_empty() && name.bytes().all(is_tchar)
590}
591
592fn is_token(value: &str) -> bool {
593 !value.is_empty() && value.bytes().all(is_tchar)
594}
595
596const fn is_tchar(byte: u8) -> bool {
597 matches!(
598 byte,
599 b'!' | b'#' | b'$' | b'%' | b'&' | b'\'' | b'*' | b'+' | b'-' | b'.' | b'^' | b'_' | b'`'
600 | b'|' | b'~'
601 | b'0'..=b'9'
602 | b'A'..=b'Z'
603 | b'a'..=b'z'
604 )
605}
606
607fn contains_invalid_header_value(value: &str) -> bool {
608 value.bytes().any(|b| matches!(b, 0..=8 | 10..=31 | 127))
609}
610
611fn is_valid_host(value: &str) -> bool {
612 if value.is_empty() || value.contains(' ') || value.contains('\t') {
613 return false;
614 }
615
616 if value.starts_with('[') {
617 let Some(end) = value.find(']') else {
618 return false;
619 };
620 let addr = &value[1..end];
621 if addr.is_empty() || !addr.chars().all(|c| c.is_ascii_hexdigit() || c == ':') {
622 return false;
623 }
624 let remainder = &value[end + 1..];
625 if remainder.is_empty() {
626 return true;
627 }
628 if let Some(port) = remainder.strip_prefix(':') {
629 return !port.is_empty() && port.chars().all(|c| c.is_ascii_digit());
630 }
631 return false;
632 }
633
634 let mut parts = value.splitn(2, ':');
635 let host = parts.next().unwrap_or("");
636 if host.is_empty()
637 || !host
638 .chars()
639 .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '.')
640 {
641 return false;
642 }
643
644 if let Some(port) = parts.next() {
645 if port.is_empty() || !port.chars().all(|c| c.is_ascii_digit()) {
646 return false;
647 }
648 }
649
650 true
651}
652
653fn parse_chunk_size(line: &str) -> Result<u64, BodyError> {
654 let (size_str, _) = line.split_once(';').unwrap_or((line, ""));
655 u64::from_str_radix(size_str.trim(), 16).map_err(|_| BodyError::InvalidChunk)
656}
657
658fn is_authority_form(value: &str) -> bool {
659 if value.is_empty() {
660 return false;
661 }
662 if value.starts_with('[') {
663 if let Some(end) = value.find(']') {
664 let port = value[end + 1..].strip_prefix(':');
665 return port
666 .map(|p| !p.is_empty() && p.chars().all(|c| c.is_ascii_digit()))
667 .unwrap_or(false);
668 }
669 return false;
670 }
671 let mut parts = value.splitn(2, ':');
672 let host = parts.next().unwrap_or("");
673 let port = parts.next();
674 !host.is_empty()
675 && port
676 .map(|p| !p.is_empty() && p.chars().all(|c| c.is_ascii_digit()))
677 .unwrap_or(false)
678}
679
680#[cfg(test)]
681mod tests {
682 use super::*;
683 use bytes::BytesMut;
684
685 #[test]
686 fn detect_headers_end() {
687 let data = b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n";
688 assert_eq!(find_headers_end(data), Some(data.len()));
689 }
690
691 #[test]
692 fn parse_simple_request() {
693 let data = b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n";
694 let (request, mode, consumed) = parse_request_head(data).unwrap();
695 assert_eq!(consumed, data.len());
696 assert_eq!(request.method().as_str(), "GET");
697 assert_eq!(request.version(), HttpVersion::HTTP_1_1);
698 assert_eq!(mode, BodyMode::None);
699 }
700
701 #[tokio::test]
702 async fn fixed_body_reader_consumes_buffered_bytes() {
703 let mut buf = BytesMut::from(&b"hello"[..]);
704 let mut reader = tokio::io::empty();
705 let mut body = body_reader(BodyMode::Fixed(5), &mut reader, &mut buf);
706 let chunk = body.read_next().await.unwrap().unwrap();
707 assert_eq!(&chunk[..], b"hello");
708 assert!(body.read_next().await.unwrap().is_none());
709 }
710
711 #[tokio::test]
712 async fn chunked_reader_parses_small_chunk() {
713 let mut buf = BytesMut::from(&b"4\r\nRust\r\n0\r\n\r\n"[..]);
714 let mut reader = tokio::io::empty();
715 let mut body = body_reader(BodyMode::Chunked, &mut reader, &mut buf);
716 let chunk = body.read_next().await.unwrap().unwrap();
717 assert_eq!(&chunk[..], b"Rust");
718 assert!(body.read_next().await.unwrap().is_none());
719 }
720}