use crate::types::Http1Error;
const MAX_TRAILER_BUF_SIZE: usize = 8192;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ChunkedState {
SizeLine,
ChunkData,
ChunkDataEnd,
Trailer,
Done,
}
#[derive(Debug, Clone)]
pub struct ChunkedDecoder {
pub state: ChunkedState,
pub remaining: u64,
pub current_chunk_size: u64,
pub total_decoded: u64,
pub max_total: u64,
pub saw_last: bool,
size_buf: [u8; 64],
size_len: usize,
trailer_buf: Vec<u8>,
}
impl ChunkedDecoder {
#[inline]
pub fn new(max_total: u64) -> Self {
Self {
state: ChunkedState::SizeLine,
remaining: 0,
current_chunk_size: 0,
total_decoded: 0,
max_total,
saw_last: false,
size_buf: [0u8; 64],
size_len: 0,
trailer_buf: Vec::new(),
}
}
#[inline]
pub fn reset(&mut self) {
self.state = ChunkedState::SizeLine;
self.remaining = 0;
self.current_chunk_size = 0;
self.total_decoded = 0;
self.saw_last = false;
self.size_len = 0;
self.trailer_buf.clear();
}
#[inline]
pub fn is_done(&self) -> bool {
matches!(self.state, ChunkedState::Done)
}
pub fn feed(&mut self, input: &[u8]) -> Result<(Vec<u8>, usize), Http1Error> {
let mut output = Vec::with_capacity(input.len());
if self.size_len > 0 {
let buffered = self.size_len;
let combined_len = buffered.checked_add(input.len()).ok_or_else(|| {
Http1Error::ChunkedError("buffered input length overflow".into())
})?;
let mut combined = Vec::with_capacity(combined_len);
combined.extend_from_slice(&self.size_buf[..buffered]);
combined.extend_from_slice(input);
self.size_len = 0;
let consumed = self.feed_combined(&combined, &mut output)?;
let own_consumed = consumed.saturating_sub(buffered);
return Ok((output, own_consumed));
}
if !self.trailer_buf.is_empty() {
let buffered = self.trailer_buf.len();
let combined_len = buffered.checked_add(input.len()).ok_or_else(|| {
Http1Error::ChunkedError("buffered trailer input length overflow".into())
})?;
let mut combined = Vec::with_capacity(combined_len);
combined.extend_from_slice(&self.trailer_buf);
combined.extend_from_slice(input);
self.trailer_buf.clear();
let consumed = self.feed_combined(&combined, &mut output)?;
let own_consumed = consumed.saturating_sub(buffered);
return Ok((output, own_consumed));
}
let consumed = self.feed_combined(input, &mut output)?;
Ok((output, consumed))
}
fn feed_combined(
&mut self,
input: &[u8],
output: &mut Vec<u8>,
) -> Result<usize, Http1Error> {
let mut consumed = 0;
while consumed < input.len() && !matches!(self.state, ChunkedState::Done) {
match self.state {
ChunkedState::SizeLine => {
let rest = &input[consumed..];
if let Some(crlf) = find_crlf(rest) {
let line = &rest[..crlf];
self.parse_size_line(line)?;
consumed += crlf + 2;
} else {
let available = rest.len().min(self.size_buf.len() - self.size_len);
if available > 0 {
self.size_buf[self.size_len..self.size_len + available]
.copy_from_slice(&rest[..available]);
self.size_len += available;
consumed += available;
}
if self.size_len >= self.size_buf.len() {
return Err(Http1Error::ChunkedError(
"chunk-size line too long".into(),
));
}
break;
}
}
ChunkedState::ChunkData => {
let take = (input.len() - consumed) as u64;
let take = take.min(self.remaining);
if take == 0 {
break;
}
output.extend_from_slice(
&input[consumed..consumed + take as usize],
);
consumed += take as usize;
self.remaining -= take;
self.total_decoded = self.total_decoded.checked_add(take).ok_or(Http1Error::BodyTooLarge)?;
if self.total_decoded > self.max_total {
return Err(Http1Error::BodyTooLarge);
}
if self.remaining == 0 {
self.state = ChunkedState::ChunkDataEnd;
}
}
ChunkedState::ChunkDataEnd => {
let rest = &input[consumed..];
if rest.len() < 2 {
break;
}
if rest[0] == b'\r' && rest[1] == b'\n' {
consumed += 2;
self.state = ChunkedState::SizeLine;
} else {
return Err(Http1Error::ChunkedError(
"missing CRLF after chunk-data".into(),
));
}
}
ChunkedState::Trailer => {
let rest = &input[consumed..];
if rest.is_empty() {
break;
}
if let Some(crlf) = find_crlf(rest) {
consumed += crlf + 2;
if crlf == 0 {
self.state = ChunkedState::Done;
}
} else {
self.trailer_buf.extend_from_slice(rest);
if self.trailer_buf.len() > MAX_TRAILER_BUF_SIZE {
return Err(Http1Error::ChunkedError(
"trailer buffer exceeds maximum size".into(),
));
}
consumed += rest.len();
break;
}
}
ChunkedState::Done => {
return Err(Http1Error::ChunkedError(
"chunked decoder in terminal state".into(),
));
}
}
}
Ok(consumed)
}
fn parse_size_line(&mut self, line: &[u8]) -> Result<(), Http1Error> {
let parsed = crate::parser::parse_chunk_size(line)?;
let val = u64::try_from(parsed).map_err(|_| Http1Error::BodyTooLarge)?;
self.current_chunk_size = val;
if val == 0 {
self.saw_last = true;
self.state = ChunkedState::Trailer;
} else {
if val > self.max_total {
return Err(Http1Error::BodyTooLarge);
}
if self
.total_decoded
.checked_add(val)
.map(|sum| sum > self.max_total)
.unwrap_or(true)
{
return Err(Http1Error::BodyTooLarge);
}
self.remaining = val;
self.state = ChunkedState::ChunkData;
}
Ok(())
}
}
#[inline]
fn find_crlf(buf: &[u8]) -> Option<usize> {
let mut i = 0;
while i + 1 < buf.len() {
if buf[i] == b'\r' && buf[i + 1] == b'\n' {
return Some(i);
}
i += 1;
}
None
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_simple_chunked() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\r\n0\r\n\r\n";
let (out, consumed) = dec.feed(input).unwrap();
assert_eq!(out, b"hello");
assert_eq!(consumed, input.len());
assert!(dec.is_done());
}
#[test]
fn test_multiple_chunks() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\r\n6\r\n world\r\n0\r\n\r\n";
let (out, consumed) = dec.feed(input).unwrap();
assert_eq!(out, b"hello world");
assert_eq!(consumed, input.len());
assert!(dec.is_done());
}
#[test]
fn test_incremental() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\r\n0\r\n\r\n";
let (out1, c1) = dec.feed(&input[..6]).unwrap();
assert_eq!(c1, 6, "first feed should consume 6 bytes");
let (out2, _c2) = dec.feed(&input[c1..]).unwrap();
let mut all = out1;
all.extend_from_slice(&out2);
assert_eq!(all, b"hello");
assert!(dec.is_done());
}
#[test]
fn test_chunk_extension() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5;foo=bar\r\nhello\r\n0\r\n\r\n";
let (out, _) = dec.feed(input).unwrap();
assert_eq!(out, b"hello");
assert!(dec.is_done());
}
#[test]
fn test_invalid_hex() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"ZZ\r\nx\r\n0\r\n\r\n";
let r = dec.feed(input);
assert!(r.is_err());
}
#[test]
fn test_too_large() {
let mut dec = ChunkedDecoder::new(2);
let input = b"5\r\nhello\r\n0\r\n\r\n";
let r = dec.feed(input);
assert!(r.is_err());
}
#[test]
fn test_missing_crlf() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\rBAD";
let r = dec.feed(input);
assert!(r.is_err());
}
#[test]
fn test_empty_chunk_zero_size() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"0\r\n\r\n";
let (out, consumed) = dec.feed(input).unwrap();
assert!(out.is_empty());
assert_eq!(consumed, input.len());
assert!(dec.is_done());
assert!(dec.saw_last);
}
#[test]
fn test_single_byte_chunk() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"1\r\nA\r\n0\r\n\r\n";
let (out, _) = dec.feed(input).unwrap();
assert_eq!(out, b"A");
assert!(dec.is_done());
}
#[test]
fn test_large_hex_chunk_size() {
let mut dec = ChunkedDecoder::new(1024 * 1024);
let size = 0xFF;
let mut input = Vec::new();
input.extend_from_slice(format!("{:X}\r\n", size).as_bytes());
input.extend_from_slice(&vec![b'x'; size]);
input.extend_from_slice(b"\r\n0\r\n\r\n");
let (out, _) = dec.feed(&input).unwrap();
assert_eq!(out.len(), size);
assert!(dec.is_done());
}
#[test]
fn test_chunk_size_with_trailing_spaces() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5 \r\nhello\r\n0\r\n\r\n";
let (out, _) = dec.feed(input).unwrap();
assert_eq!(out, b"hello");
assert!(dec.is_done());
}
#[test]
fn test_chunk_size_with_tabs() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\t\r\nhello\r\n0\r\n\r\n";
let (out, _) = dec.feed(input).unwrap();
assert_eq!(out, b"hello");
assert!(dec.is_done());
}
#[test]
fn test_empty_chunk_size_line() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"\r\nhello\r\n0\r\n\r\n";
let r = dec.feed(input);
assert!(r.is_err());
}
#[test]
fn test_chunk_size_line_too_long() {
let mut dec = ChunkedDecoder::new(1024);
let mut long_line = vec![b'A'; 100];
long_line.extend_from_slice(b"\r\n");
let r = dec.feed(&long_line);
assert!(r.is_err());
}
#[test]
fn test_trailer_state_received() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\r\n0\r\n\r\n";
let (out, _) = dec.feed(input).unwrap();
assert_eq!(out, b"hello");
assert!(dec.is_done());
}
#[test]
fn test_trailer_empty_immediate_done() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\r\n0\r\n";
dec.feed(input).unwrap();
let _ = dec.feed(b"\r\n");
assert!(dec.is_done());
}
#[test]
fn test_reset_decoder() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\r\n0\r\n\r\n";
dec.feed(input).unwrap();
assert!(dec.is_done());
dec.reset();
assert_eq!(dec.state, ChunkedState::SizeLine);
assert_eq!(dec.total_decoded, 0);
assert!(!dec.saw_last);
assert_eq!(dec.current_chunk_size, 0);
}
#[test]
fn test_chunked_state_variants() {
let states = [
ChunkedState::SizeLine,
ChunkedState::ChunkData,
ChunkedState::ChunkDataEnd,
ChunkedState::Trailer,
ChunkedState::Done,
];
for (i, s) in states.iter().enumerate() {
assert_eq!(*s, states[i]);
}
assert_ne!(ChunkedState::SizeLine, ChunkedState::Done);
}
#[test]
fn test_decoder_clone() {
let dec = ChunkedDecoder::new(1024);
let dec2 = dec.clone();
assert_eq!(dec.state, dec2.state);
assert_eq!(dec.max_total, dec2.max_total);
}
#[test]
fn test_total_decoded_tracking() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\r\n6\r\n world\r\n0\r\n\r\n";
let (_, _) = dec.feed(input).unwrap();
assert_eq!(dec.total_decoded, 11);
}
#[test]
fn test_uppercase_hex() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"A\r\n0123456789\r\n0\r\n\r\n";
let (out, _) = dec.feed(input).unwrap();
assert_eq!(out.len(), 10);
assert!(dec.is_done());
}
#[test]
fn test_lowercase_hex() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"a\r\n0123456789\r\n0\r\n\r\n";
let (out, _) = dec.feed(input).unwrap();
assert_eq!(out.len(), 10);
assert!(dec.is_done());
}
#[test]
fn test_negative_chunk_size_rejected() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"-5\r\nhello\r\n0\r\n\r\n";
let r = dec.feed(input);
assert!(r.is_err());
}
#[test]
fn test_body_too_large_accumulated() {
let mut dec = ChunkedDecoder::new(10);
let input = b"6\r\nhello \r\n6\r\nworld!\r\n0\r\n\r\n";
let r = dec.feed(input);
assert!(r.is_err());
}
#[test]
fn test_total_decoded_overflow_fail_closed() {
let mut dec = ChunkedDecoder::new(u64::MAX);
let input = b"5\r\nhello\r\nFFFFFFFFFFFFFFFC\r\nx\r\n0\r\n\r\n";
let r = dec.feed(input);
assert!(
matches!(r, Err(Http1Error::BodyTooLarge)),
"total_decoded 溢出必须返回 BodyTooLarge,实际 {r:?}"
);
}
#[test]
fn test_incremental_small_chunks() {
let mut dec = ChunkedDecoder::new(1024);
let input = b"5\r\nhello\r\n0\r\n\r\n";
let mut output = Vec::new();
let mut pos = 0;
while pos < input.len() {
let end = (pos + 2).min(input.len());
let (out, consumed) = dec.feed(&input[pos..end]).unwrap();
output.extend_from_slice(&out);
pos += consumed;
if consumed == 0 {
pos += 1;
}
}
assert_eq!(output, b"hello");
assert!(dec.is_done());
}
}