use std::fs::File;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use super::response::BodyKind;
use super::response::STREAM_CHUNK;
use super::{Headers, Method, Request, Response, StatusCode, Version};
use crate::error::{Error, Result};
use crate::proto::read_at_exact;
#[derive(Debug, Clone, Copy)]
pub struct Limits {
pub max_header_bytes: usize,
pub max_body_bytes: usize,
}
impl Default for Limits {
fn default() -> Limits {
Limits {
max_header_bytes: 64 * 1024,
max_body_bytes: 16 * 1024 * 1024,
}
}
}
#[derive(Debug, Clone, Copy)]
struct Pending {
version: Version,
keep_alive: bool,
is_head: bool,
}
#[derive(Debug)]
struct FileBody {
file: Arc<File>,
offset: u64,
remaining: u64,
}
#[derive(Debug)]
pub struct H1Conn {
inbuf: Vec<u8>,
outbuf: Vec<u8>,
body_stream: Option<FileBody>,
limits: Limits,
pending: Option<Pending>,
head_scanned: usize,
chunk: Option<ChunkDecoder>,
closed: bool,
interim_sent: bool,
server_name: Option<String>,
}
impl Default for H1Conn {
fn default() -> H1Conn {
H1Conn::new(Limits::default())
}
}
impl H1Conn {
pub fn new(limits: Limits) -> H1Conn {
H1Conn {
inbuf: Vec::new(),
outbuf: Vec::new(),
body_stream: None,
limits,
pending: None,
head_scanned: 0,
chunk: None,
closed: false,
interim_sent: false,
server_name: Some(concat!("httpsd/", env!("CARGO_PKG_VERSION")).to_owned()),
}
}
pub fn set_server_name(&mut self, name: Option<String>) {
self.server_name = name;
}
pub fn feed(&mut self, data: &[u8]) {
self.inbuf.extend_from_slice(data);
}
pub fn take_out(&mut self) -> Vec<u8> {
if !self.outbuf.is_empty() {
return std::mem::take(&mut self.outbuf);
}
let Some(fb) = self.body_stream.as_mut() else {
return Vec::new();
};
let want = (fb.remaining as usize).min(STREAM_CHUNK);
let mut buf = vec![0u8; want];
match read_at_exact(&fb.file, fb.offset, &mut buf) {
Ok(n) => {
buf.truncate(n);
fb.offset += n as u64;
fb.remaining -= n as u64;
if n < want {
self.body_stream = None;
self.closed = true;
} else if fb.remaining == 0 {
self.body_stream = None;
}
buf
}
Err(_) => {
self.body_stream = None;
self.closed = true;
Vec::new()
}
}
}
pub fn has_output(&self) -> bool {
!self.outbuf.is_empty() || self.body_stream.is_some()
}
pub fn wants_close(&self) -> bool {
self.closed
}
pub fn awaiting_response(&self) -> bool {
self.pending.is_some()
}
pub fn poll_request(&mut self) -> Result<Option<Request>> {
if self.closed || self.pending.is_some() {
return Ok(None);
}
let search_from = self.head_scanned.saturating_sub(3);
let head_end = match find_subslice(&self.inbuf[search_from..], b"\r\n\r\n") {
Some(rel) => {
let end = search_from + rel;
self.head_scanned = end + 3;
end
}
None => {
self.head_scanned = self.inbuf.len();
if self.inbuf.len() > self.limits.max_header_bytes {
return Err(self.fail(StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE, "headers"));
}
return Ok(None);
}
};
let header_block_len = head_end; if header_block_len > self.limits.max_header_bytes {
return Err(self.fail(StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE, "headers"));
}
let body_start = head_end + 4;
let head = &self.inbuf[..header_block_len];
let (method, target, version, headers) = match parse_head(head) {
Ok(parts) => parts,
Err(e) => {
let status = match &e {
Error::BadRequest(_) => StatusCode::BAD_REQUEST,
_ => StatusCode::BAD_REQUEST,
};
return Err(self.fail(status, "request line/headers"));
}
};
if version.is_none() {
return Err(self.fail(StatusCode::HTTP_VERSION_NOT_SUPPORTED, "version"));
}
let version = version.unwrap();
let framing = match body_framing(&headers) {
Ok(f) => f,
Err(()) => return Err(self.fail(StatusCode::BAD_REQUEST, "body framing")),
};
let body: Vec<u8>;
let consumed_total: usize;
match framing {
BodyFraming::None => {
body = Vec::new();
consumed_total = body_start;
}
BodyFraming::Length(len) => {
if len > self.limits.max_body_bytes {
return Err(self.fail(StatusCode::PAYLOAD_TOO_LARGE, "body"));
}
if self.inbuf.len() < body_start + len {
self.maybe_send_continue(&headers);
return Ok(None);
}
body = self.inbuf[body_start..body_start + len].to_vec();
consumed_total = body_start + len;
}
BodyFraming::Chunked => {
let mut dec = self.chunk.take().unwrap_or_default();
match dec.advance(&self.inbuf[body_start..], self.limits.max_body_bytes) {
Ok(Some(used)) => {
body = std::mem::take(&mut dec.out);
consumed_total = body_start + used;
}
Ok(None) => {
self.chunk = Some(dec);
self.maybe_send_continue(&headers);
return Ok(None);
}
Err(status) => return Err(self.fail(status, "chunked body")),
}
}
}
self.inbuf.drain(..consumed_total);
self.head_scanned = 0;
self.chunk = None;
self.interim_sent = false;
let keep_alive = negotiate_keep_alive(version, &headers);
let is_head = method.is_head();
self.pending = Some(Pending {
version,
keep_alive,
is_head,
});
Ok(Some(Request::new(method, target, version, headers, body)))
}
pub fn respond(&mut self, resp: Response) {
let meta = self
.pending
.take()
.expect("respond() called with no request in flight");
self.serialize(meta, resp);
}
fn fail(&mut self, status: StatusCode, what: &'static str) -> Error {
let meta = Pending {
version: Version::Http11,
keep_alive: false,
is_head: false,
};
let resp = Response::status(status);
self.pending = None;
self.head_scanned = 0;
self.chunk = None;
self.serialize(meta, resp);
self.closed = true;
match status.code() {
413 | 431 => Error::TooLarge(what),
_ => Error::BadRequest(what),
}
}
fn maybe_send_continue(&mut self, headers: &Headers) {
if !self.interim_sent && headers.contains_token("expect", "100-continue") {
self.outbuf
.extend_from_slice(b"HTTP/1.1 100 Continue\r\n\r\n");
self.interim_sent = true;
}
}
fn serialize(&mut self, meta: Pending, resp: Response) {
let (status, mut headers, body) = resp.into_parts();
let bodyless = status.is_bodyless();
let omit_body = bodyless || meta.is_head;
for h in HOP_BY_HOP_HEADERS {
headers.remove(h);
}
if !bodyless {
headers.set("Content-Length", body.len().to_string());
} else {
headers.remove("Content-Length");
}
let keep_alive = meta.keep_alive && !self.closed;
headers.set(
"Connection",
if keep_alive { "keep-alive" } else { "close" },
);
if let Some(server) = &self.server_name {
headers.set_if_absent("Server", server.clone());
}
headers.set_if_absent("Date", http_date(now_secs()));
let line = format!(
"{} {} {}\r\n",
meta.version.as_str(),
status.code(),
status.reason()
);
self.outbuf.extend_from_slice(line.as_bytes());
for (name, value) in headers.iter() {
if !is_token(name)
|| value
.bytes()
.any(|b| b == b'\r' || b == b'\n' || b == b'\0')
{
continue;
}
self.outbuf.extend_from_slice(name.as_bytes());
self.outbuf.extend_from_slice(b": ");
self.outbuf.extend_from_slice(value.as_bytes());
self.outbuf.extend_from_slice(b"\r\n");
}
self.outbuf.extend_from_slice(b"\r\n");
if !omit_body {
match body.into_kind() {
BodyKind::Bytes(bytes) => self.outbuf.extend_from_slice(&bytes),
BodyKind::File { file, offset, len } if len > 0 => {
self.body_stream = Some(FileBody {
file,
offset,
remaining: len,
});
}
BodyKind::File { .. } => {}
}
}
if !keep_alive {
self.closed = true;
}
}
}
enum BodyFraming {
None,
Length(usize),
Chunked,
}
fn body_framing(headers: &Headers) -> std::result::Result<BodyFraming, ()> {
let chunked = headers.contains_token("transfer-encoding", "chunked");
let has_te = headers.contains("transfer-encoding");
let has_cl = headers.contains("content-length");
if has_te && has_cl {
return Err(());
}
if chunked {
return Ok(BodyFraming::Chunked);
}
if has_te {
return Err(());
}
let mut len: Option<usize> = None;
for v in headers.get_all("content-length") {
let v = v.trim();
if v.is_empty() || !v.bytes().all(|b| b.is_ascii_digit()) {
return Err(());
}
let parsed: usize = v.parse().map_err(|_| ())?;
match len {
Some(prev) if prev != parsed => return Err(()),
_ => len = Some(parsed),
}
}
match len {
Some(0) | None => Ok(BodyFraming::None),
Some(n) => Ok(BodyFraming::Length(n)),
}
}
fn negotiate_keep_alive(version: Version, headers: &Headers) -> bool {
if headers.contains_token("connection", "close") {
return false;
}
if headers.contains_token("connection", "keep-alive") {
return true;
}
version.default_keep_alive()
}
fn parse_head(head: &[u8]) -> Result<(Method, String, Option<Version>, Headers)> {
for (i, &b) in head.iter().enumerate() {
if b == b'\r' {
if head.get(i + 1) != Some(&b'\n') {
return Err(Error::BadRequest("bare CR in header block"));
}
} else if b == b'\n' && (i == 0 || head[i - 1] != b'\r') {
return Err(Error::BadRequest("bare LF in header block"));
}
}
let text = std::str::from_utf8(head).map_err(|_| Error::BadRequest("non-UTF-8 header"))?;
let mut lines = text.split("\r\n");
let request_line = lines.next().ok_or(Error::BadRequest("empty request"))?;
let mut parts = request_line.split(' ');
let method = parts.next().ok_or(Error::BadRequest("no method"))?;
let target = parts.next().ok_or(Error::BadRequest("no target"))?;
let version_tok = parts.next().ok_or(Error::BadRequest("no version"))?;
if parts.next().is_some() {
return Err(Error::BadRequest("trailing request-line tokens"));
}
if method.is_empty() || target.is_empty() {
return Err(Error::BadRequest("empty request-line token"));
}
let method = Method::parse(method);
let version = Version::parse(version_tok);
let mut headers = Headers::new();
for line in lines {
if line.is_empty() {
continue;
}
if line.starts_with(' ') || line.starts_with('\t') {
return Err(Error::BadRequest("obsolete header folding"));
}
if headers.len() >= MAX_HEADER_FIELDS {
return Err(Error::BadRequest("too many header fields"));
}
let (name, value) = line
.split_once(':')
.ok_or(Error::BadRequest("header without colon"))?;
if !is_token(name) {
return Err(Error::BadRequest("invalid header name"));
}
let value = value.trim();
if !is_valid_field_value(value) {
return Err(Error::BadRequest("invalid header value"));
}
headers.append(name, value);
}
Ok((method, target.to_owned(), version, headers))
}
#[derive(Debug, Default)]
struct ChunkDecoder {
pos: usize,
out: Vec<u8>,
state: ChunkState,
}
#[derive(Debug, Default)]
enum ChunkState {
#[default]
Size,
Data { remaining: usize },
DataCrlf,
Trailer { start: usize },
}
impl ChunkDecoder {
fn advance(
&mut self,
data: &[u8],
max_body: usize,
) -> std::result::Result<Option<usize>, StatusCode> {
loop {
match self.state {
ChunkState::Size => {
let eol = match find_subslice(&data[self.pos..], b"\r\n") {
Some(eol) => eol,
None => {
if data.len() - self.pos > MAX_CHUNK_LINE_BYTES {
return Err(StatusCode::BAD_REQUEST);
}
return Ok(None);
}
};
if eol > MAX_CHUNK_LINE_BYTES {
return Err(StatusCode::BAD_REQUEST);
}
let size_line = &data[self.pos..self.pos + eol];
let hex = match size_line.iter().position(|&b| b == b';') {
Some(i) => &size_line[..i],
None => size_line,
};
let hex = std::str::from_utf8(hex).map_err(|_| StatusCode::BAD_REQUEST)?;
let hex = hex.trim();
if hex.is_empty() || !hex.bytes().all(|b| b.is_ascii_hexdigit()) {
return Err(StatusCode::BAD_REQUEST);
}
let size =
usize::from_str_radix(hex, 16).map_err(|_| StatusCode::BAD_REQUEST)?;
if size > max_body {
return Err(StatusCode::PAYLOAD_TOO_LARGE);
}
let after_size = self.pos + eol + 2;
if size == 0 {
self.pos = after_size;
self.state = ChunkState::Trailer { start: after_size };
continue;
}
match self.out.len().checked_add(size) {
Some(total) if total <= max_body => {}
_ => return Err(StatusCode::PAYLOAD_TOO_LARGE),
}
self.pos = after_size;
self.state = ChunkState::Data { remaining: size };
}
ChunkState::Data { remaining } => {
let avail = data.len() - self.pos;
let take = remaining.min(avail);
self.out.extend_from_slice(&data[self.pos..self.pos + take]);
self.pos += take;
let left = remaining - take;
if left > 0 {
self.state = ChunkState::Data { remaining: left };
return Ok(None);
}
self.state = ChunkState::DataCrlf;
}
ChunkState::DataCrlf => {
if data.len() < self.pos + 2 {
return Ok(None);
}
if &data[self.pos..self.pos + 2] != b"\r\n" {
return Err(StatusCode::BAD_REQUEST);
}
self.pos += 2;
self.state = ChunkState::Size;
}
ChunkState::Trailer { start } => {
loop {
if self.pos.saturating_sub(start) > MAX_TRAILER_BYTES {
return Err(StatusCode::BAD_REQUEST);
}
match find_subslice(&data[self.pos..], b"\r\n") {
None => {
if data.len() - start > MAX_TRAILER_BYTES {
return Err(StatusCode::BAD_REQUEST);
}
return Ok(None);
}
Some(0) => return Ok(Some(self.pos + 2)),
Some(eol) => self.pos += eol + 2,
}
}
}
}
}
}
}
const MAX_HEADER_FIELDS: usize = 100;
const MAX_CHUNK_LINE_BYTES: usize = 16 * 1024;
const MAX_TRAILER_BYTES: usize = 8 * 1024;
const HOP_BY_HOP_HEADERS: [&str; 7] = [
"transfer-encoding",
"connection",
"keep-alive",
"upgrade",
"te",
"trailer",
"proxy-connection",
];
fn is_tchar(b: u8) -> bool {
b.is_ascii_alphanumeric()
|| matches!(
b,
b'!' | b'#'
| b'$'
| b'%'
| b'&'
| b'\''
| b'*'
| b'+'
| b'-'
| b'.'
| b'^'
| b'_'
| b'`'
| b'|'
| b'~'
)
}
fn is_token(s: &str) -> bool {
!s.is_empty() && s.bytes().all(is_tchar)
}
fn is_valid_field_value(s: &str) -> bool {
s.bytes().all(|b| (b >= 0x20 && b != 0x7f) || b == b'\t')
}
fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option<usize> {
let nlen = needle.len();
if nlen == 0 || haystack.len() < nlen {
return None;
}
let first = needle[0];
let last_start = haystack.len() - nlen;
let mut i = 0;
while i <= last_start {
match haystack[i..=last_start].iter().position(|&b| b == first) {
Some(off) => {
let cand = i + off;
if haystack[cand..cand + nlen] == *needle {
return Some(cand);
}
i = cand + 1;
}
None => return None,
}
}
None
}
fn now_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
pub(crate) fn http_date(secs: u64) -> String {
const WDAY: [&str; 7] = ["Sun", "Mon", "Tue", "Wed", "Thu", "Fri", "Sat"];
const MON: [&str; 12] = [
"Jan", "Feb", "Mar", "Apr", "May", "Jun", "Jul", "Aug", "Sep", "Oct", "Nov", "Dec",
];
let days = (secs / 86_400) as i64;
let tod = secs % 86_400;
let (h, mi, s) = (tod / 3600, (tod % 3600) / 60, tod % 60);
let z = days + 719_468;
let era = if z >= 0 { z } else { z - 146_096 } / 146_097;
let doe = z - era * 146_097; let yoe = (doe - doe / 1460 + doe / 36_524 - doe / 146_096) / 365; let mut year = yoe + era * 400;
let doy = doe - (365 * yoe + yoe / 4 - yoe / 100); let mp = (5 * doy + 2) / 153; let day = doy - (153 * mp + 2) / 5 + 1; let month = if mp < 10 { mp + 3 } else { mp - 9 }; if month <= 2 {
year += 1;
}
let wday = ((days % 7 + 7) % 7 + 4) % 7;
format!(
"{}, {:02} {} {:04} {:02}:{:02}:{:02} GMT",
WDAY[wday as usize],
day,
MON[(month - 1) as usize],
year,
h,
mi,
s,
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::proto::Body;
fn temp_file(data: &[u8]) -> Arc<File> {
use std::io::Write;
let path = std::env::temp_dir().join(format!(
"httpsd-h1-stream-{}-{}",
std::process::id(),
COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
));
let mut f = File::create(&path).unwrap();
f.write_all(data).unwrap();
f.sync_all().unwrap();
let opened = File::open(&path).unwrap();
let _ = std::fs::remove_file(&path); Arc::new(opened)
}
static COUNTER: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
fn drain_all(conn: &mut H1Conn) -> Vec<u8> {
let mut out = Vec::new();
while conn.has_output() {
let chunk = conn.take_out();
if chunk.is_empty() {
break;
}
out.extend_from_slice(&chunk);
}
out
}
#[test]
fn streams_multichunk_file_with_correct_length() {
let n = 5 * STREAM_CHUNK + 123;
let data: Vec<u8> = (0..n).map(|i| (i % 251) as u8).collect();
let file = temp_file(&data);
let mut c = H1Conn::default();
let _ = drive(&mut c, b"GET / HTTP/1.1\r\nConnection: close\r\n\r\n").unwrap();
c.respond(Response::new(StatusCode::OK).body(Body::file(file, 0, n as u64)));
let out = drain_all(&mut c);
let split = find_subslice(&out, b"\r\n\r\n").unwrap() + 4;
let head = String::from_utf8(out[..split].to_vec()).unwrap();
assert!(
head.contains(&format!("Content-Length: {n}\r\n")),
"head: {head}"
);
assert_eq!(&out[split..], &data[..], "streamed body must be byte-exact");
}
#[test]
fn streams_file_range_span_only() {
let data: Vec<u8> = (0..(2 * STREAM_CHUNK)).map(|i| (i % 256) as u8).collect();
let file = temp_file(&data);
let (start, len) = (1000u64, (STREAM_CHUNK + 7) as u64);
let mut c = H1Conn::default();
let _ = drive(&mut c, b"GET / HTTP/1.1\r\nConnection: close\r\n\r\n").unwrap();
c.respond(Response::new(StatusCode::PARTIAL_CONTENT).body(Body::file(file, start, len)));
let out = drain_all(&mut c);
let split = find_subslice(&out, b"\r\n\r\n").unwrap() + 4;
assert_eq!(
&out[split..],
&data[start as usize..(start + len) as usize],
"range body must be exactly the requested span"
);
}
#[test]
fn head_file_sends_length_but_no_body() {
let data = vec![7u8; 3 * STREAM_CHUNK];
let file = temp_file(&data);
let mut c = H1Conn::default();
let _ = drive(&mut c, b"HEAD / HTTP/1.1\r\nConnection: close\r\n\r\n").unwrap();
c.respond(Response::new(StatusCode::OK).body(Body::file(file, 0, data.len() as u64)));
let out = drain_all(&mut c);
let text = String::from_utf8(out).unwrap();
assert!(
text.contains(&format!("Content-Length: {}\r\n", data.len())),
"head: {text}"
);
assert!(text.ends_with("\r\n\r\n"), "HEAD must send no body: {text}");
}
fn drive(conn: &mut H1Conn, input: &[u8]) -> Option<Request> {
conn.feed(input);
conn.poll_request().unwrap()
}
#[test]
fn parses_simple_get() {
let mut c = H1Conn::default();
let req = drive(&mut c, b"GET /hello?x=1 HTTP/1.1\r\nHost: a\r\n\r\n").unwrap();
assert_eq!(req.method(), &Method::Get);
assert_eq!(req.path(), "/hello");
assert_eq!(req.query(), Some("x=1"));
assert_eq!(req.host(), Some("a"));
assert!(req.body().is_empty());
}
#[test]
fn waits_for_full_body() {
let mut c = H1Conn::default();
c.feed(b"POST / HTTP/1.1\r\nContent-Length: 5\r\n\r\nab");
assert!(c.poll_request().unwrap().is_none());
c.feed(b"cde");
let req = c.poll_request().unwrap().unwrap();
assert_eq!(req.body(), b"abcde");
}
#[test]
fn decodes_chunked() {
let mut c = H1Conn::default();
let req = drive(
&mut c,
b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n0\r\n\r\n",
)
.unwrap();
assert_eq!(req.body(), b"hello");
}
#[test]
fn keep_alive_default_by_version() {
let mut c = H1Conn::default();
let req = drive(&mut c, b"GET / HTTP/1.1\r\n\r\n").unwrap();
assert!(negotiate_keep_alive(req.version(), req.headers()));
let mut c = H1Conn::default();
let req = drive(&mut c, b"GET / HTTP/1.0\r\n\r\n").unwrap();
assert!(!negotiate_keep_alive(req.version(), req.headers()));
}
#[test]
fn serializes_response_with_framing() {
let mut c = H1Conn::default();
let _ = drive(&mut c, b"GET / HTTP/1.1\r\nConnection: close\r\n\r\n").unwrap();
c.respond(Response::text("hi"));
let out = String::from_utf8(c.take_out()).unwrap();
assert!(out.starts_with("HTTP/1.1 200 OK\r\n"));
assert!(out.contains("Content-Length: 2\r\n"));
assert!(out.contains("Connection: close\r\n"));
assert!(out.ends_with("\r\n\r\nhi"));
assert!(c.wants_close());
}
#[test]
fn head_omits_body_keeps_length() {
let mut c = H1Conn::default();
let _ = drive(&mut c, b"HEAD / HTTP/1.1\r\n\r\n").unwrap();
c.respond(Response::text("hello"));
let out = String::from_utf8(c.take_out()).unwrap();
assert!(out.contains("Content-Length: 5\r\n"));
assert!(out.ends_with("\r\n\r\n")); }
#[test]
fn rejects_te_and_cl_together() {
let mut c = H1Conn::default();
c.feed(b"POST / HTTP/1.1\r\nContent-Length: 1\r\nTransfer-Encoding: chunked\r\n\r\n");
assert!(c.poll_request().is_err());
assert!(c.wants_close());
let out = String::from_utf8(c.take_out()).unwrap();
assert!(out.starts_with("HTTP/1.1 400"));
}
#[test]
fn chunked_with_real_trailer_parses_and_drains() {
let mut c = H1Conn::default();
let req = drive(
&mut c,
b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n0\r\nA: 1\r\n\r\n",
)
.unwrap();
assert_eq!(req.body(), b"hello");
assert!(c.inbuf.is_empty());
}
#[test]
fn chunked_partial_trailer_waits() {
let mut c = H1Conn::default();
c.feed(b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n0\r\nA: 1\r\n");
assert!(c.poll_request().unwrap().is_none());
c.feed(b"\r\n");
let req = c.poll_request().unwrap().unwrap();
assert_eq!(req.body(), b"hello");
assert!(c.inbuf.is_empty());
}
#[test]
fn chunked_trailer_bytes_not_smuggled_as_next_request() {
let mut c = H1Conn::default();
c.feed(
b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n\
0\r\nX: y\r\n\r\nGET /evil HTTP/1.1\r\nHost: a\r\n\r\n",
);
let _ = c.poll_request().unwrap().unwrap();
c.respond(Response::text("ok"));
let next = c.poll_request().unwrap().unwrap();
assert_eq!(next.path(), "/evil");
}
#[test]
fn chunked_huge_size_is_rejected_not_panicking() {
let mut c = H1Conn::default();
c.feed(b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\nfffffffffffffff0\r\n");
let err = c.poll_request();
assert!(err.is_err());
let out = String::from_utf8(c.take_out()).unwrap();
assert!(out.starts_with("HTTP/1.1 413"));
}
#[test]
fn chunked_non_hex_size_is_rejected() {
let mut c = H1Conn::default();
c.feed(b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n+5\r\nhello\r\n0\r\n\r\n");
assert!(c.poll_request().is_err());
}
#[test]
fn chunked_oversized_size_line_is_rejected() {
let mut c = H1Conn::default();
c.feed(b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n");
let mut huge = b"1;".to_vec();
huge.extend(std::iter::repeat_n(b'a', MAX_CHUNK_LINE_BYTES + 16));
c.feed(&huge);
assert!(c.poll_request().is_err());
}
#[test]
fn rejects_bare_lf_in_headers() {
let mut c = H1Conn::default();
c.feed(b"GET / HTTP/1.1\r\nHost: a\nX: y\r\n\r\n");
assert!(c.poll_request().is_err());
}
#[test]
fn rejects_control_char_in_header_value() {
let mut c = H1Conn::default();
c.feed(b"GET / HTTP/1.1\r\nX: a\x01b\r\n\r\n");
assert!(c.poll_request().is_err());
}
#[test]
fn rejects_non_token_header_name() {
let mut c = H1Conn::default();
c.feed(b"GET / HTTP/1.1\r\nBad Name: x\r\n\r\n");
assert!(c.poll_request().is_err());
}
#[test]
fn rejects_too_many_header_fields() {
let mut c = H1Conn::default();
let mut req = b"GET / HTTP/1.1\r\n".to_vec();
for i in 0..(MAX_HEADER_FIELDS + 5) {
req.extend_from_slice(format!("X-{i}: v\r\n").as_bytes());
}
req.extend_from_slice(b"\r\n");
c.feed(&req);
assert!(c.poll_request().is_err());
}
#[test]
fn rejects_non_digit_content_length() {
let mut c = H1Conn::default();
c.feed(b"POST / HTTP/1.1\r\nContent-Length: +5\r\n\r\nhello");
assert!(c.poll_request().is_err());
}
#[test]
fn serialize_strips_handler_transfer_encoding_and_injection() {
let mut c = H1Conn::default();
let _ = drive(&mut c, b"GET / HTTP/1.1\r\nConnection: close\r\n\r\n").unwrap();
let mut resp = Response::text("hi");
resp.headers_mut().set("Transfer-Encoding", "chunked");
resp.headers_mut().set("X-Evil", "a\r\nInjected: 1");
c.respond(resp);
let out = String::from_utf8(c.take_out()).unwrap();
assert!(!out.to_ascii_lowercase().contains("transfer-encoding"));
assert!(!out.contains("Injected: 1"));
assert!(out.contains("Content-Length: 2\r\n"));
}
#[test]
fn chunked_byte_by_byte_matches_all_at_once() {
let req_bytes: &[u8] = b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n\
5;ext=1\r\nhello\r\n6\r\n world\r\n0\r\nX: y\r\n\r\n";
let mut c_all = H1Conn::default();
c_all.feed(req_bytes);
let req_all = c_all.poll_request().unwrap().unwrap();
assert_eq!(req_all.body(), b"hello world");
assert!(c_all.inbuf.is_empty());
let mut c_inc = H1Conn::default();
let mut got = None;
for &b in req_bytes {
c_inc.feed(&[b]);
if let Some(r) = c_inc.poll_request().unwrap() {
got = Some(r);
}
}
let req_inc = got.expect("incremental feed never completed the request");
assert_eq!(req_inc.body(), req_all.body());
assert_eq!(req_inc.method(), req_all.method());
assert_eq!(req_inc.path(), req_all.path());
assert!(
c_inc.inbuf.is_empty(),
"inbuf must be drained clean after incremental parse"
);
assert!(c_inc.chunk.is_none());
assert_eq!(c_inc.head_scanned, 0);
}
#[test]
fn large_header_block_byte_by_byte_parses() {
let mut req = b"GET / HTTP/1.1\r\n".to_vec();
for i in 0..50 {
req.extend_from_slice(format!("X-Pad-{i}: {}\r\n", "v".repeat(800)).as_bytes());
}
req.extend_from_slice(b"Host: a\r\n\r\n");
let mut c = H1Conn::default();
let mut got = None;
for &b in &req {
c.feed(&[b]);
if let Some(r) = c.poll_request().unwrap() {
got = Some(r);
}
}
let req = got.expect("terminator never found under byte-by-byte feed");
assert_eq!(req.method(), &Method::Get);
assert_eq!(req.host(), Some("a"));
assert!(c.inbuf.is_empty());
}
#[test]
fn chunked_huge_size_rejected_under_incremental_feed() {
let bytes: &[u8] =
b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\nfffffffffffffff0\r\n";
let mut c = H1Conn::default();
let mut err = false;
for &b in bytes {
c.feed(&[b]);
match c.poll_request() {
Ok(_) => {}
Err(_) => {
err = true;
break;
}
}
}
assert!(err, "malicious chunk size must be rejected (no panic)");
let out = String::from_utf8(c.take_out()).unwrap();
assert!(out.starts_with("HTTP/1.1 413"));
}
#[test]
fn chunked_large_body_byte_by_byte_completes() {
const N: usize = 200 * 1024;
let mut body = Vec::with_capacity(N);
for i in 0..N {
body.push(b'a' + (i % 26) as u8);
}
let mut req = b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n".to_vec();
req.extend_from_slice(format!("{:x}\r\n", N).as_bytes());
req.extend_from_slice(&body);
req.extend_from_slice(b"\r\n0\r\n\r\n");
let mut c = H1Conn::default();
let mut got = None;
for &b in &req {
c.feed(&[b]);
if let Some(r) = c.poll_request().unwrap() {
got = Some(r);
}
}
let req = got.expect("large chunked body never completed");
assert_eq!(req.body().len(), N);
assert_eq!(req.body(), &body[..]);
assert!(c.inbuf.is_empty());
}
#[test]
fn find_subslice_linear_correctness() {
assert_eq!(find_subslice(b"", b"\r\n"), None);
assert_eq!(find_subslice(b"\r\n", b"\r\n"), Some(0));
assert_eq!(find_subslice(b"ab\r\ncd", b"\r\n"), Some(2));
assert_eq!(find_subslice(b"a\rb\r\nc", b"\r\n"), Some(3));
assert_eq!(find_subslice(b"x\r\n\r\ny", b"\r\n\r\n"), Some(1));
assert_eq!(find_subslice(b"\r\n\r", b"\r\n\r\n"), None);
assert_eq!(find_subslice(b"abc", b"abc"), Some(0));
assert_eq!(find_subslice(b"aabc", b"abc"), Some(1));
assert_eq!(find_subslice(b"ab", b"abc"), None);
assert_eq!(find_subslice(b"hello", b"l"), Some(2));
assert_eq!(find_subslice(b"hello", b""), None);
}
#[test]
fn http_date_known_value() {
assert_eq!(http_date(784_111_777), "Sun, 06 Nov 1994 08:49:37 GMT");
}
}