use crate::head::Head;
use crate::header::{HeaderId, HeaderVec};
use crate::limits::MAX_HEADERS_CEILING;
use crate::{ByteStr, Limits, Method, Version};
use bytes::Bytes;
use memchr::memmem;
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
pub enum ParseError {
#[error("bare CR in message head")]
BareCr,
#[error("bare LF in message head")]
BareLf,
#[error("obs-fold in message head")]
ObsFold,
#[error("whitespace before colon in field name")]
WhitespaceBeforeColon,
#[error("invalid field name")]
InvalidHeaderName,
#[error("invalid request line")]
InvalidRequestLine,
#[error("invalid method token")]
InvalidMethod,
#[error("request target is not valid UTF-8")]
NonUtf8Target,
#[error("invalid character in request target")]
InvalidTarget,
#[error("unsupported HTTP version")]
UnsupportedVersion,
#[error("message head too large")]
HeadTooLarge,
#[error("too many header fields")]
TooManyHeaders,
}
impl ParseError {
#[inline]
pub fn status(&self) -> u16 {
match self {
ParseError::HeadTooLarge | ParseError::TooManyHeaders => 431,
ParseError::UnsupportedVersion => 505,
_ => 400,
}
}
}
#[inline]
pub fn find_head_end(buf: &[u8]) -> Option<usize> {
memmem::find(buf, b"\r\n\r\n").map(|i| i + 4)
}
pub fn prescan(head: &[u8]) -> Result<(), ParseError> {
for i in memchr::memchr_iter(b'\r', head) {
if head.get(i + 1) != Some(&b'\n') {
return Err(ParseError::BareCr);
}
if matches!(head.get(i + 2), Some(b' ') | Some(b'\t')) {
return Err(ParseError::ObsFold);
}
}
for i in memchr::memchr_iter(b'\n', head) {
if i == 0 || head[i - 1] != b'\r' {
return Err(ParseError::BareLf);
}
}
Ok(())
}
const _: () = assert!(MAX_HEADERS_CEILING == 128);
const HEADER_SCRATCH_TIER1: usize = 32;
pub fn parse_head(buf: &Bytes, limits: &Limits) -> Result<Option<(Head, usize)>, ParseError> {
debug_assert!(limits.max_headers <= MAX_HEADERS_CEILING);
let Some(head_len) = find_head_end(buf) else {
if buf.len() > limits.max_head_bytes {
return Err(ParseError::HeadTooLarge);
}
return Ok(None);
};
if head_len > limits.max_head_bytes {
return Err(ParseError::HeadTooLarge);
}
let region = &buf[..head_len];
prescan(region)?;
let mut scratch1 = [httparse::EMPTY_HEADER; HEADER_SCRATCH_TIER1];
let mut req1 = httparse::Request::new(&mut scratch1);
match req1.parse(region) {
Ok(parsed) => return finish_parse_head(buf, req1, parsed, head_len, limits),
Err(httparse::Error::TooManyHeaders) if limits.max_headers > HEADER_SCRATCH_TIER1 => {
}
Err(e) => return Err(map_httparse(e)),
}
let mut scratch2 = [httparse::EMPTY_HEADER; MAX_HEADERS_CEILING];
let mut req2 = httparse::Request::new(&mut scratch2);
let parsed = req2.parse(region).map_err(map_httparse)?;
finish_parse_head(buf, req2, parsed, head_len, limits)
}
fn finish_parse_head(
buf: &Bytes,
req: httparse::Request<'_, '_>,
parsed: httparse::Status<usize>,
head_len: usize,
limits: &Limits,
) -> Result<Option<(Head, usize)>, ParseError> {
let httparse::Status::Complete(consumed) = parsed else {
return Err(ParseError::InvalidRequestLine);
};
debug_assert_eq!(consumed, head_len);
let version = req
.version
.and_then(Version::from_httparse)
.ok_or(ParseError::UnsupportedVersion)?;
let method_token = req.method.ok_or(ParseError::InvalidRequestLine)?;
let method = match Method::from_bytes(method_token.as_bytes()) {
Some(m) => m,
None => {
if method_token.is_empty() {
return Err(ParseError::InvalidMethod);
}
Method::Other(
ByteStr::from_utf8(buf.slice_ref(method_token.as_bytes()))
.map_err(|_| ParseError::InvalidMethod)?,
)
}
};
let target_token = req.path.ok_or(ParseError::InvalidRequestLine)?;
validate_target(target_token.as_bytes())?;
validate_target_form(target_token.as_bytes(), &method)?;
let target = ByteStr::from_utf8(buf.slice_ref(target_token.as_bytes()))
.map_err(|_| ParseError::NonUtf8Target)?;
let field_count = req.headers.len();
if field_count > limits.max_headers {
return Err(ParseError::TooManyHeaders);
}
let mut headers = HeaderVec::with_capacity(field_count);
for h in req.headers.iter() {
let name = h.name.as_bytes();
let id = match HeaderId::from_bytes(name) {
Some(id) => id,
None => {
let lowered = if name.iter().any(|b| b.is_ascii_uppercase()) {
Bytes::from(name.to_ascii_lowercase())
} else {
buf.slice_ref(name)
};
HeaderId::Other(
ByteStr::from_utf8(lowered).map_err(|_| ParseError::InvalidHeaderName)?,
)
}
};
headers.push((id, buf.slice_ref(h.value)));
}
Ok(Some((
Head::new(method, target, version, headers),
head_len,
)))
}
fn validate_target(target: &[u8]) -> Result<(), ParseError> {
if target.is_empty() {
return Err(ParseError::InvalidTarget);
}
for &b in target {
let ok = b.is_ascii_alphanumeric()
|| matches!(
b,
b'-' | b'.'
| b'_'
| b'~'
| b':'
| b'/'
| b'?'
| b'['
| b']'
| b'@'
| b'!'
| b'$'
| b'&'
| b'\''
| b'('
| b')'
| b'*'
| b'+'
| b','
| b';'
| b'='
| b'%'
);
if !ok {
return Err(ParseError::InvalidTarget);
}
}
Ok(())
}
fn validate_target_form(target: &[u8], method: &Method) -> Result<(), ParseError> {
if target == b"*" {
return if matches!(method, Method::Options) {
Ok(())
} else {
Err(ParseError::InvalidTarget)
};
}
if matches!(method, Method::Connect) {
let has_scheme = target.windows(3).any(|w| w == b"://");
let has_colon = target.contains(&b':');
return if !has_scheme && has_colon && !target.starts_with(b"/") {
Ok(())
} else {
Err(ParseError::InvalidTarget)
};
}
if target.starts_with(b"/") {
return Ok(());
}
if let Some(i) = target.windows(3).position(|w| w == b"://") {
let scheme = &target[..i];
let valid_scheme = scheme.first().is_some_and(|b| b.is_ascii_alphabetic())
&& scheme
.iter()
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'+' | b'-' | b'.'));
return if valid_scheme {
Ok(())
} else {
Err(ParseError::InvalidTarget)
};
}
Err(ParseError::InvalidTarget)
}
fn map_httparse(e: httparse::Error) -> ParseError {
match e {
httparse::Error::HeaderName => ParseError::WhitespaceBeforeColon,
httparse::Error::Version => ParseError::UnsupportedVersion,
httparse::Error::Token => ParseError::InvalidMethod,
httparse::Error::TooManyHeaders => ParseError::TooManyHeaders,
_ => ParseError::InvalidRequestLine,
}
}
#[cfg(test)]
mod tests {
use super::*;
const OK: &[u8] = b"GET / HTTP/1.1\r\nHost: a\r\n\r\n";
#[test]
fn find_head_end_locates_the_terminator() {
assert_eq!(find_head_end(OK), Some(OK.len()));
let with_body = b"GET / HTTP/1.1\r\nHost: a\r\n\r\nBODY";
assert_eq!(find_head_end(with_body), Some(with_body.len() - 4));
}
#[test]
fn find_head_end_reports_incomplete() {
assert_eq!(find_head_end(b""), None);
assert_eq!(find_head_end(b"GET / HTTP/1.1\r\n"), None);
assert_eq!(find_head_end(b"GET / HTTP/1.1\r\nHost: a\r\n"), None);
assert_eq!(find_head_end(b"GET / HTTP/1.1\r\nHost: a\r\n\r"), None);
}
#[test]
fn prescan_accepts_strict_crlf() {
assert!(prescan(OK).is_ok());
assert!(prescan(b"GET / HTTP/1.1\r\nHost: a\r\nAccept: */*\r\n\r\n").is_ok());
assert!(prescan(b"GET / HTTP/1.1\r\n\r\n").is_ok());
}
#[test]
fn prescan_rejects_bare_lf() {
assert_eq!(
prescan(b"GET / HTTP/1.1\nHost: a\r\n\r\n"),
Err(ParseError::BareLf)
);
assert_eq!(
prescan(b"GET / HTTP/1.1\r\nHost: a\n\r\n"),
Err(ParseError::BareLf)
);
}
#[test]
fn prescan_rejects_bare_cr() {
assert_eq!(
prescan(b"GET / HTTP/1.1\rHost: a\r\n\r\n"),
Err(ParseError::BareCr)
);
assert_eq!(
prescan(b"GET / HTTP/1.1\r\nHost: a\rb\r\n\r\n"),
Err(ParseError::BareCr)
);
}
#[test]
fn prescan_rejects_obs_fold() {
assert_eq!(
prescan(b"GET / HTTP/1.1\r\nHost: a\r\n b\r\n\r\n"),
Err(ParseError::ObsFold)
);
assert_eq!(
prescan(b"GET / HTTP/1.1\r\nHost: a\r\n\tb\r\n\r\n"),
Err(ParseError::ObsFold)
);
}
#[test]
fn prescan_does_not_mistake_the_terminator_for_a_fold() {
assert!(prescan(b"GET / HTTP/1.1\r\nHost: a\r\n\r\n").is_ok());
}
#[test]
fn status_codes_map_correctly() {
assert_eq!(ParseError::BareCr.status(), 400);
assert_eq!(ParseError::BareLf.status(), 400);
assert_eq!(ParseError::ObsFold.status(), 400);
assert_eq!(ParseError::WhitespaceBeforeColon.status(), 400);
assert_eq!(ParseError::InvalidRequestLine.status(), 400);
assert_eq!(ParseError::NonUtf8Target.status(), 400);
assert_eq!(ParseError::InvalidTarget.status(), 400);
assert_eq!(ParseError::HeadTooLarge.status(), 431);
assert_eq!(ParseError::TooManyHeaders.status(), 431);
assert_eq!(ParseError::UnsupportedVersion.status(), 505);
}
fn parse(raw: &'static [u8]) -> Result<Option<(Head, usize)>, ParseError> {
parse_head(&Bytes::from_static(raw), &Limits::default())
}
#[test]
fn parses_a_minimal_request() {
let (head, n) = parse(b"GET / HTTP/1.1\r\nHost: a.example\r\n\r\n")
.unwrap()
.unwrap();
assert_eq!(n, 35);
assert_eq!(head.method, Method::Get);
assert_eq!(head.target().as_str(), "/");
assert_eq!(head.version, Version::Http11);
assert_eq!(head.headers.len(), 1);
assert_eq!(head.get_str(&HeaderId::Host), Some("a.example"));
}
#[test]
fn reports_incomplete_heads() {
assert!(parse(b"GET / HTTP/1.1\r\n").unwrap().is_none());
assert!(parse(b"GE").unwrap().is_none());
assert!(parse(b"").unwrap().is_none());
}
#[test]
fn header_values_share_the_read_buffer() {
let buf = Bytes::from_static(b"GET /a?b=1 HTTP/1.1\r\nHost: a.example\r\n\r\n");
let (head, _) = parse_head(&buf, &Limits::default()).unwrap().unwrap();
let host = head.get(&HeaderId::Host).unwrap();
let base = buf.as_ptr() as usize;
let addr = host.as_ptr() as usize;
assert!(
addr >= base && addr < base + buf.len(),
"header value must point into the read buffer, not a copy"
);
let target = head.target().as_bytes().as_ptr() as usize;
assert!(target >= base && target < base + buf.len());
}
#[test]
fn splits_path_and_query() {
let (head, _) = parse(b"GET /a/b?x=1&y=2 HTTP/1.1\r\nHost: a\r\n\r\n")
.unwrap()
.unwrap();
assert_eq!(head.path(), "/a/b");
assert_eq!(head.query(), Some("x=1&y=2"));
let (head, _) = parse(b"GET /a/b HTTP/1.1\r\nHost: a\r\n\r\n")
.unwrap()
.unwrap();
assert_eq!(head.path(), "/a/b");
assert_eq!(head.query(), None);
let (head, _) = parse(b"GET /a? HTTP/1.1\r\nHost: a\r\n\r\n")
.unwrap()
.unwrap();
assert_eq!(head.path(), "/a");
assert_eq!(head.query(), Some(""));
}
#[test]
fn unknown_method_and_header_become_other_variants() {
let (head, _) = parse(b"PROPFIND / HTTP/1.1\r\nHost: a\r\nX-Req-Id: z\r\n\r\n")
.unwrap()
.unwrap();
assert_eq!(head.method.as_str(), "PROPFIND");
assert!(matches!(head.method, Method::Other(_)));
let id = HeaderId::Other(ByteStr::from_static("x-req-id"));
assert_eq!(head.get_str(&id), Some("z"));
}
#[test]
fn custom_header_names_are_lowercased() {
let (head, _) = parse(b"GET / HTTP/1.1\r\nHost: a\r\nX-Req-Id: z\r\n\r\n")
.unwrap()
.unwrap();
match &head.headers[1].0 {
HeaderId::Other(name) => assert_eq!(name.as_str(), "x-req-id"),
other => panic!("expected Other, got {other:?}"),
}
}
#[test]
fn http_10_is_accepted_and_http_12_is_not() {
let (head, _) = parse(b"GET / HTTP/1.0\r\n\r\n").unwrap().unwrap();
assert_eq!(head.version, Version::Http10);
assert_eq!(
parse(b"GET / HTTP/1.2\r\nHost: a\r\n\r\n"),
Err(ParseError::UnsupportedVersion)
);
}
#[test]
fn rejects_invalid_target_characters() {
for raw in [
&b"GET /a\\b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a<b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a>b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a{b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a|b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a^b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a\"b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a#frag HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a?q=1#frag HTTP/1.1\r\nHost: a\r\n\r\n"[..],
] {
assert_eq!(
parse_head(&Bytes::copy_from_slice(raw), &Limits::default()),
Err(ParseError::InvalidTarget),
"should have rejected: {}",
String::from_utf8_lossy(raw)
);
}
for raw in [
&b"GET /a\x01b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a\x7fb HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a\x00b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a\xc3\xa9 HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a\xffb HTTP/1.1\r\nHost: a\r\n\r\n"[..],
] {
let err = parse_head(&Bytes::copy_from_slice(raw), &Limits::default())
.expect_err("a non-URI byte in a target must be rejected");
assert_eq!(err.status(), 400, "{err:?}");
}
}
#[test]
fn rejects_targets_matching_no_permitted_form() {
for raw in [
&b"GET foo/bar HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET a.example:80 HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET ://a.example/x HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET 1http://a.example/x HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET * HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"CONNECT /x HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"CONNECT http://a.example/x HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"CONNECT a.example HTTP/1.1\r\nHost: a\r\n\r\n"[..],
] {
assert_eq!(
parse_head(&Bytes::copy_from_slice(raw), &Limits::default()),
Err(ParseError::InvalidTarget),
"should have rejected: {}",
String::from_utf8_lossy(raw)
);
}
}
#[test]
fn accepts_every_legitimate_target_form() {
for raw in [
&b"GET / HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a/b?x=1&y=2 HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a%20b HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /~user/a.b_c-d HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET /a?q=1;2,3!$&'()*+= HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"OPTIONS * HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"GET http://a.example:8080/x HTTP/1.1\r\nHost: a.example\r\n\r\n"[..],
&b"GET http://[::1]:8080/x HTTP/1.1\r\nHost: a\r\n\r\n"[..],
&b"CONNECT a.example:443 HTTP/1.1\r\nHost: a.example\r\n\r\n"[..],
] {
assert!(
parse_head(&Bytes::copy_from_slice(raw), &Limits::default()).is_ok(),
"should have accepted: {}",
String::from_utf8_lossy(raw)
);
}
}
#[test]
fn rejects_whitespace_before_colon() {
assert_eq!(
parse(b"GET / HTTP/1.1\r\nHost : a\r\n\r\n"),
Err(ParseError::WhitespaceBeforeColon)
);
}
#[test]
fn rejects_malformed_request_lines() {
assert!(parse(b"GET\r\nHost: a\r\n\r\n").is_err());
assert!(parse(b"GET /\r\nHost: a\r\n\r\n").is_err());
assert!(parse(b"/ HTTP/1.1\r\nHost: a\r\n\r\n").is_err());
}
#[test]
fn prescan_runs_before_tokenization() {
assert_eq!(
parse(b"GET / HTTP/1.1\nHost: a\r\n\r\n"),
Err(ParseError::BareLf)
);
}
#[test]
fn enforces_head_size_limit() {
let limits = Limits {
max_head_bytes: 40,
..Default::default()
};
let long = Bytes::from_static(
b"GET /aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa HTTP/1.1\r\nHost: a\r\n\r\n",
);
assert_eq!(parse_head(&long, &limits), Err(ParseError::HeadTooLarge));
}
#[test]
fn enforces_header_count_limit() {
let limits = Limits {
max_headers: 2,
..Default::default()
};
let raw = Bytes::from_static(b"GET / HTTP/1.1\r\nHost: a\r\nA: 1\r\nB: 2\r\n\r\n");
assert_eq!(parse_head(&raw, &limits), Err(ParseError::TooManyHeaders));
}
#[test]
fn head_crossing_tier1_but_within_max_headers_still_parses() {
let mut raw = String::from("GET / HTTP/1.1\r\n");
for i in 0..33 {
raw.push_str(&format!("x-{i}: v\r\n"));
}
raw.push_str("\r\n");
let buf = Bytes::from(raw.into_bytes());
let limits = Limits::default();
assert!(limits.max_headers > HEADER_SCRATCH_TIER1);
let (head, _) = parse_head(&buf, &limits).unwrap().unwrap();
assert_eq!(head.headers.len(), 33);
}
#[test]
fn head_exceeding_max_headers_after_crossing_tier1_still_rejects() {
let limits = Limits {
max_headers: 40,
..Default::default()
};
assert!(limits.max_headers > HEADER_SCRATCH_TIER1);
let mut raw = String::from("GET / HTTP/1.1\r\n");
for i in 0..41 {
raw.push_str(&format!("x-{i}: v\r\n"));
}
raw.push_str("\r\n");
let buf = Bytes::from(raw.into_bytes());
assert_eq!(parse_head(&buf, &limits), Err(ParseError::TooManyHeaders));
}
#[test]
fn keep_alive_defaults_follow_the_version() {
let (head, _) = parse(b"GET / HTTP/1.1\r\nHost: a\r\n\r\n")
.unwrap()
.unwrap();
assert!(head.is_keep_alive(), "HTTP/1.1 persists by default");
let (head, _) = parse(b"GET / HTTP/1.0\r\n\r\n").unwrap().unwrap();
assert!(!head.is_keep_alive(), "HTTP/1.0 closes by default");
}
#[test]
fn connection_tokens_override_the_default() {
let (head, _) = parse(b"GET / HTTP/1.1\r\nHost: a\r\nConnection: close\r\n\r\n")
.unwrap()
.unwrap();
assert!(!head.is_keep_alive());
let (head, _) = parse(b"GET / HTTP/1.0\r\nConnection: keep-alive\r\n\r\n")
.unwrap()
.unwrap();
assert!(head.is_keep_alive());
}
#[test]
fn connection_token_matching_is_case_insensitive_and_list_aware() {
let (head, _) =
parse(b"GET / HTTP/1.1\r\nHost: a\r\nConnection: Keep-Alive, Upgrade\r\n\r\n")
.unwrap()
.unwrap();
assert!(head.connection_has_token("upgrade"));
assert!(head.connection_has_token("keep-alive"));
assert!(!head.connection_has_token("close"));
assert!(head.is_keep_alive());
}
}