use crate::courierust_bytes::Bytes;
use crate::courierust_error::{Error, Result};
use crate::courierust_http::header::{HeaderMap, HeaderName, HeaderValue};
use crate::courierust_http::method::Method;
use crate::courierust_http::status::StatusCode;
use crate::courierust_http::uri::PathAndQuery;
use crate::courierust_http::version::Version;
use crate::courierust_io::{BufReader, Read};
const MAX_LINE: usize = 64 * 1024;
const MAX_HEADERS: usize = 1024;
const MAX_HEADER_BLOCK: usize = 1024 * 1024;
#[derive(Debug, Clone)]
pub struct RequestLine {
pub method: Method,
pub target: PathAndQuery,
pub version: Version,
}
pub fn parse_request_line(line: &[u8]) -> Result<RequestLine> {
let line = trim_crlf(line);
let mut parts = line.split(|&b| b == b' ');
let method = parts
.next()
.filter(|m| !m.is_empty())
.ok_or_else(|| Error::protocol("missing method"))
.and_then(Method::from_bytes)?;
let target = parts
.next()
.filter(|t| !t.is_empty())
.ok_or_else(|| Error::protocol("missing request target"))
.and_then(PathAndQuery::from_bytes)?;
let version = parts
.next()
.filter(|v| !v.is_empty())
.ok_or_else(|| Error::protocol("missing HTTP version"))
.and_then(parse_version)?;
if parts.next().is_some() {
return Err(Error::protocol("malformed request line"));
}
Ok(RequestLine {
method,
target,
version,
})
}
pub fn parse_status_line(line: &[u8]) -> Result<(StatusCode, Version)> {
let line = trim_crlf(line);
let mut parts = line.splitn(3, |&b| b == b' ');
let version = parts
.next()
.filter(|v| !v.is_empty())
.ok_or_else(|| Error::protocol("missing version"))
.and_then(parse_version)?;
let code = parts
.next()
.ok_or_else(|| Error::protocol("missing status code"))?;
let code: u16 = std::str::from_utf8(code)
.map_err(|_| Error::protocol("non-ASCII status code"))?
.parse()
.map_err(|_| Error::protocol("invalid status code"))?;
if !(100..=599).contains(&code) {
return Err(Error::protocol("status code out of range"));
}
Ok((StatusCode::from_u16(code), version))
}
pub fn parse_version(v: &[u8]) -> Result<Version> {
Ok(match v {
b"HTTP/1.0" => Version::HTTP_10,
b"HTTP/1.1" => Version::HTTP_11,
b"HTTP/2" => Version::HTTP_2,
_ => return Err(Error::protocol("unsupported HTTP version")),
})
}
pub fn read_headers<R: Read>(reader: &mut BufReader<R>) -> Result<HeaderMap> {
let mut headers = HeaderMap::new();
let mut total = 0usize;
loop {
let line = reader.read_until(b'\n', MAX_LINE)?;
total += line.len();
if total > MAX_HEADER_BLOCK {
return Err(Error::overflow("header block too large"));
}
if line.len() >= MAX_LINE {
return Err(Error::overflow("header line too long"));
}
let trimmed = trim_crlf(&line);
if trimmed.is_empty() {
break; }
if headers.len() >= MAX_HEADERS {
return Err(Error::overflow("too many header fields"));
}
let (name, value) = split_header(trimmed)?;
headers.append(name, value);
}
Ok(headers)
}
pub fn read_headers_scratch<R: Read>(
reader: &mut BufReader<R>,
scratch: &mut crate::courierust_io::Scratch,
) -> Result<HeaderMap> {
let mut headers = HeaderMap::new();
let mut total = 0usize;
loop {
let line = scratch.line();
reader.read_until_into(b'\n', MAX_LINE, line)?;
total += line.len();
if total > MAX_HEADER_BLOCK {
return Err(Error::overflow("header block too large"));
}
if line.len() >= MAX_LINE {
return Err(Error::overflow("header line too long"));
}
let trimmed = trim_crlf(line);
if trimmed.is_empty() {
break; }
if headers.len() >= MAX_HEADERS {
return Err(Error::overflow("too many header fields"));
}
let (name, value) = split_header(trimmed)?;
headers.append(name, value);
}
Ok(headers)
}
pub(crate) fn split_header(line: &[u8]) -> Result<(HeaderName, HeaderValue)> {
let colon = line
.iter()
.position(|&b| b == b':')
.ok_or_else(|| Error::protocol("header line missing colon"))?;
let name = HeaderName::from_bytes(&line[..colon])?;
let mut start = colon + 1;
let mut end = line.len();
while start < end && (line[start] == b' ' || line[start] == b'\t') {
start += 1;
}
while end > start && (line[end - 1] == b' ' || line[end - 1] == b'\t') {
end -= 1;
}
let value = HeaderValue::from_bytes(&line[start..end])?;
Ok((name, value))
}
pub(crate) fn trim_crlf(line: &[u8]) -> &[u8] {
let mut end = line.len();
while end > 0 && (line[end - 1] == b'\n' || line[end - 1] == b'\r' || line[end - 1] == b' ') {
end -= 1;
}
&line[..end]
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BodyLen {
None,
Length(usize),
Chunked,
}
pub fn body_length(
headers: &HeaderMap,
method: Option<&Method>,
status: Option<StatusCode>,
) -> Result<BodyLen> {
let mut te_count = 0usize;
let mut chunked_pos: Option<usize> = None;
let mut any_te = false;
for v in headers.get_all("transfer-encoding") {
any_te = true;
let s = v
.to_str()
.map_err(|_| Error::protocol("invalid transfer-encoding"))?;
for tok in s.split(',') {
let tok = tok.trim();
if tok.is_empty() {
return Err(Error::protocol("invalid transfer-encoding"));
}
if tok.eq_ignore_ascii_case("chunked") {
if chunked_pos.is_some() {
return Err(Error::protocol("chunked repeated in transfer-encoding"));
}
chunked_pos = Some(te_count);
}
te_count += 1;
}
}
if any_te {
match chunked_pos {
Some(i) if i == te_count - 1 => return Ok(BodyLen::Chunked),
Some(_) => {
return Err(Error::protocol(
"chunked must be the final transfer-encoding",
))
}
None => return Err(Error::protocol("unsupported transfer-encoding")),
}
}
let cls: Vec<&HeaderValue> = headers.get_all("content-length").collect();
if !cls.is_empty() {
let parse_len = |v: &HeaderValue| -> Option<usize> {
let s = v.to_str().ok()?.trim();
s.parse::<usize>().ok()
};
let first = parse_len(cls[0]).ok_or_else(|| Error::protocol("invalid content-length"))?;
for v in cls.iter().skip(1) {
if parse_len(v) != Some(first) {
return Err(Error::protocol("conflicting content-length"));
}
}
return Ok(BodyLen::Length(first));
}
if let Some(m) = method {
if *m == Method::HEAD {
return Ok(BodyLen::None);
}
return Ok(BodyLen::None);
}
if let Some(s) = status {
if s.is_informational() || s == StatusCode::NO_CONTENT || s == StatusCode::NOT_MODIFIED {
return Ok(BodyLen::None);
}
}
Ok(BodyLen::None)
}
fn read_into<R: Read>(reader: &mut BufReader<R>, out: &mut Vec<u8>, total: usize) -> Result<()> {
out.reserve(total.saturating_sub(out.len()));
while out.len() < total {
let b = reader.fill_buf()?;
if b.is_empty() {
return Err(Error::eof());
}
let take = core::cmp::min(total - out.len(), b.len());
out.extend_from_slice(&b[..take]);
reader.consume(take);
}
Ok(())
}
pub fn read_body_fixed<R: Read>(
reader: &mut BufReader<R>,
len: usize,
max: usize,
) -> Result<Bytes> {
let mut out = Vec::new();
read_body_fixed_into(reader, len, max, &mut out)?;
Ok(Bytes::from(out))
}
pub fn read_body_fixed_scratch<R: Read>(
reader: &mut BufReader<R>,
len: usize,
max: usize,
scratch: &mut crate::courierust_io::Scratch,
) -> Result<Bytes> {
let out = scratch.body();
read_body_fixed_into(reader, len, max, out)?;
Ok(Bytes::from(core::mem::take(out)))
}
fn read_body_fixed_into<R: Read>(
reader: &mut BufReader<R>,
len: usize,
max: usize,
out: &mut Vec<u8>,
) -> Result<()> {
if len > max {
return Err(Error::overflow("body exceeds limit"));
}
read_into(reader, out, len)
}
pub fn read_body_chunked<R: Read>(reader: &mut BufReader<R>, max: usize) -> Result<Bytes> {
let mut out = Vec::new();
read_body_chunked_into(reader, max, &mut out)?;
Ok(Bytes::from(out))
}
pub fn read_body_chunked_scratch<R: Read>(
reader: &mut BufReader<R>,
max: usize,
scratch: &mut crate::courierust_io::Scratch,
) -> Result<Bytes> {
let out = scratch.body();
read_body_chunked_into(reader, max, out)?;
Ok(Bytes::from(core::mem::take(out)))
}
fn read_body_chunked_into<R: Read>(
reader: &mut BufReader<R>,
max: usize,
out: &mut Vec<u8>,
) -> Result<()> {
let mut line = Vec::new();
let mut trailer_total = 0usize;
loop {
line.clear();
reader.read_until_into(b'\n', MAX_LINE, &mut line)?;
let size = parse_chunk_size(trim_crlf(&line))
.ok_or_else(|| Error::protocol("invalid chunk size"))?;
if size == 0 {
loop {
line.clear();
reader.read_until_into(b'\n', MAX_LINE, &mut line)?;
trailer_total += line.len();
if trailer_total > MAX_HEADER_BLOCK {
return Err(Error::overflow("trailer section too large"));
}
if trim_crlf(&line).is_empty() {
break;
}
}
break;
}
if size > max.saturating_sub(out.len()) {
return Err(Error::overflow("body exceeds limit"));
}
read_into(reader, out, out.len() + size)?;
let mut crlf = [0u8; 2];
reader.read_exact_into(&mut crlf)?;
if crlf != *b"\r\n" {
return Err(Error::protocol("chunk terminator missing"));
}
}
Ok(())
}
pub(crate) fn parse_chunk_size(line: &[u8]) -> Option<usize> {
let before_ext = match line.iter().position(|&b| b == b';') {
Some(i) => &line[..i],
None => line,
};
let s = core::str::from_utf8(before_ext).ok()?;
let s = s.trim();
if s.is_empty() {
return None;
}
usize::from_str_radix(s, 16).ok()
}
pub fn write_request_head(
out: &mut Vec<u8>,
method: &Method,
target: &PathAndQuery,
version: Version,
headers: &HeaderMap,
) -> Result<()> {
out.extend_from_slice(method.as_str().as_bytes());
out.push(b' ');
out.extend_from_slice(target.as_str().as_bytes());
out.push(b' ');
out.extend_from_slice(version.wire_str().as_bytes());
out.extend_from_slice(b"\r\n");
write_headers(out, headers)
}
pub fn write_response_head(
out: &mut Vec<u8>,
status: StatusCode,
version: Version,
headers: &HeaderMap,
) -> Result<()> {
out.extend_from_slice(version.wire_str().as_bytes());
out.push(b' ');
let code = IToA::new(status.as_u16() as usize);
out.extend_from_slice(code.as_slice());
if let Some(reason) = status.canonical_reason() {
out.push(b' ');
out.extend_from_slice(reason.as_bytes());
}
out.extend_from_slice(b"\r\n");
write_headers(out, headers)
}
pub fn write_headers(out: &mut Vec<u8>, headers: &HeaderMap) -> Result<()> {
for (n, v) in headers.iter() {
if n.is_pseudo() {
continue; }
if v.as_bytes()
.iter()
.any(|&c| c == b'\r' || c == b'\n' || c == 0)
{
return Err(Error::invalid_header_value());
}
out.extend_from_slice(n.as_str().as_bytes());
out.extend_from_slice(b": ");
out.extend_from_slice(v.as_bytes());
out.extend_from_slice(b"\r\n");
}
out.extend_from_slice(b"\r\n");
Ok(())
}
pub fn encode_chunk(data: &[u8], out: &mut Vec<u8>) {
put_hex_usize(out, data.len());
out.extend_from_slice(b"\r\n");
out.extend_from_slice(data);
out.extend_from_slice(b"\r\n");
}
pub const CHUNKED_END: &[u8] = b"0\r\n\r\n";
pub struct IToA {
buf: [u8; 20],
len: usize,
}
impl IToA {
pub fn new(v: usize) -> Self {
let mut buf = [0u8; 20];
let mut i = buf.len();
if v == 0 {
i -= 1;
buf[i] = b'0';
} else {
let mut n = v;
while n > 0 {
i -= 1;
buf[i] = b'0' + (n % 10) as u8;
n /= 10;
}
}
Self {
buf,
len: buf.len() - i,
}
}
#[inline]
pub fn as_slice(&self) -> &[u8] {
&self.buf[self.buf.len() - self.len..]
}
}
const HEX_DIGITS: &[u8; 16] = b"0123456789abcdef";
fn put_hex_usize(out: &mut Vec<u8>, mut v: usize) {
let mut buf = [0u8; 16];
let mut i = buf.len();
if v == 0 {
out.push(b'0');
return;
}
while v > 0 {
i -= 1;
buf[i] = HEX_DIGITS[v & 0xf];
v >>= 4;
}
out.extend_from_slice(&buf[i..]);
}
fn connection_has_token(value: &str, tok: &str) -> bool {
value.split(',').any(|t| t.trim().eq_ignore_ascii_case(tok))
}
pub fn wants_close(headers: &HeaderMap) -> bool {
headers
.get("connection")
.and_then(|v| v.to_str().ok())
.map(|v| connection_has_token(v, "close"))
.unwrap_or(false)
}
pub fn keep_alive_requested(version: Version, headers: &HeaderMap) -> bool {
match headers.get("connection").and_then(|v| v.to_str().ok()) {
Some(c) if connection_has_token(c, "close") => false,
Some(c) if connection_has_token(c, "keep-alive") => true,
_ => version == Version::HTTP_11,
}
}
pub fn is_hop_by_hop(name: &str) -> bool {
matches!(
name,
"connection"
| "keep-alive"
| "proxy-connection"
| "transfer-encoding"
| "upgrade"
| "proxy-authenticate"
| "proxy-authorization"
| "te"
| "trailer"
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::courierust_error::ErrorKind;
use crate::courierust_http::header::{HeaderMap, HeaderName, HeaderValue};
use crate::courierust_http::method::Method;
use crate::courierust_io::{BufReader, SliceReader};
fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
let mut h = HeaderMap::new();
for (n, v) in pairs {
h.append(
HeaderName::from_bytes(n.as_bytes()).unwrap(),
HeaderValue::from_bytes(v.as_bytes()).unwrap(),
);
}
h
}
#[test]
fn conflicting_content_length_rejected() {
let h = headers(&[("content-length", "5"), ("content-length", "6")]);
assert!(body_length(&h, Some(&Method::POST), None).is_err());
}
#[test]
fn identical_content_length_accepted() {
let h = headers(&[("content-length", "5"), ("content-length", "5")]);
assert_eq!(
body_length(&h, Some(&Method::POST), None).unwrap(),
BodyLen::Length(5)
);
}
#[test]
fn transfer_encoding_wins_over_content_length() {
let h = headers(&[("transfer-encoding", "chunked"), ("content-length", "5")]);
assert_eq!(
body_length(&h, Some(&Method::POST), None).unwrap(),
BodyLen::Chunked
);
}
#[test]
fn chunk_size_overflow_rejected() {
let wire = b"ffffffffffffffff\r\nrest";
let mut reader = BufReader::new(SliceReader::new(wire), 64);
let mut out = Vec::new();
let r = read_body_chunked_into(&mut reader, 1024, &mut out);
assert!(r.is_err());
assert!(matches!(
r,
Err(Error {
kind: ErrorKind::Overflow,
..
})
));
}
#[test]
fn chunk_trailer_section_capped() {
let mut wire = b"0\r\n".to_vec();
let line = vec![b'x'; 1024];
for _ in 0..2 * 1024 {
wire.extend_from_slice(&line);
wire.push(b'\n');
}
let mut reader = BufReader::new(SliceReader::new(&wire), 4096);
let mut out = Vec::new();
let r = read_body_chunked_into(&mut reader, 1024, &mut out);
assert!(r.is_err());
assert!(matches!(
r,
Err(Error {
kind: ErrorKind::Overflow,
..
})
));
}
#[test]
fn itoa_formats_zero_and_large() {
assert_eq!(IToA::new(0).as_slice(), b"0");
assert_eq!(IToA::new(200).as_slice(), b"200");
assert_eq!(IToA::new(65535).as_slice(), b"65535");
}
#[test]
fn transfer_encoding_notchunked_rejected() {
let h = headers(&[("transfer-encoding", "notchunked")]);
assert!(body_length(&h, Some(&Method::POST), None).is_err());
}
#[test]
fn transfer_encoding_chunked_not_final_rejected() {
let h = headers(&[("transfer-encoding", "chunked, gzip")]);
assert!(body_length(&h, Some(&Method::POST), None).is_err());
}
#[test]
fn transfer_encoding_repeated_chunked_rejected() {
let h = headers(&[("transfer-encoding", "chunked, chunked")]);
assert!(body_length(&h, Some(&Method::POST), None).is_err());
}
#[test]
fn transfer_encoding_split_across_fields_accepted() {
let h = headers(&[
("transfer-encoding", "gzip"),
("transfer-encoding", "chunked"),
]);
assert_eq!(
body_length(&h, Some(&Method::POST), None).unwrap(),
BodyLen::Chunked
);
}
#[test]
fn request_line_extra_tokens_rejected() {
assert!(parse_request_line(b"GET / HTTP/1.1 garbage\r\n").is_err());
assert!(parse_request_line(b"GET / HTTP/1.1 extra more\r\n").is_err());
assert!(parse_request_line(b"GET /\r\n").is_err()); }
#[test]
fn request_line_normal_accepted() {
assert!(parse_request_line(b"GET /a?b HTTP/1.1\r\n").is_ok());
}
#[test]
fn chunk_size_whitespace_and_garbage() {
assert_eq!(parse_chunk_size(b"1A\r\n"), Some(0x1a));
assert_eq!(parse_chunk_size(b"1A ;ext\r\n"), Some(0x1a));
assert_eq!(parse_chunk_size(b"1A\t;ext\r\n"), Some(0x1a));
assert_eq!(parse_chunk_size(b" 1A \r\n"), Some(0x1a));
assert_eq!(parse_chunk_size(b"1A zzz\r\n"), None); assert_eq!(parse_chunk_size(b"\r\n"), None); assert_eq!(
parse_chunk_size(b"ffffffffffffffffffff\r\n"),
None );
}
#[test]
fn connection_token_boundary() {
let close = headers(&[("connection", "closex")]);
assert!(!wants_close(&close), "closex must not count as close");
let mixed = headers(&[("connection", "keep-alive, close")]);
assert!(wants_close(&mixed), "exact close token must count");
let ka = headers(&[("connection", "keep-aliveX")]);
assert!(
!keep_alive_requested(Version::HTTP_10, &ka),
"keep-aliveX must not count as keep-alive for HTTP/1.0"
);
let ok = headers(&[("connection", "keep-alive")]);
assert!(keep_alive_requested(Version::HTTP_10, &ok));
}
}