use std::ops::Range;
use std::task::{Context, Poll, ready};
use arrayvec::ArrayVec;
use bytes::Bytes;
use futures_core::Stream;
use memchr::memmem;
use crate::Error;
use crate::multipart::{Multipart, State};
use crate::utils::trim_ows;
pub const MAX_HEADERS: usize = 3;
#[derive(Debug)]
pub struct HeaderBlock {
spans: ArrayVec<HeaderSpan, MAX_HEADERS>,
next: usize,
data_start: usize,
}
#[derive(Debug, Clone)]
struct HeaderSpan {
name: Range<usize>,
value: Range<usize>,
}
fn parse_header_block(block: &[u8]) -> Result<HeaderBlock, Error> {
let mut headers = [httparse::EMPTY_HEADER; MAX_HEADERS];
match httparse::parse_headers(block, &mut headers) {
Ok(httparse::Status::Complete((data_start, parsed))) => {
let mut spans = ArrayVec::<HeaderSpan, MAX_HEADERS>::new();
for header in parsed {
let name_start = header.name.as_ptr() as usize - block.as_ptr() as usize;
let value_start = header.value.as_ptr() as usize - block.as_ptr() as usize;
let _ = spans.try_push(HeaderSpan {
name: name_start..name_start.saturating_add(header.name.len()),
value: value_start..value_start.saturating_add(header.value.len()),
});
}
Ok(HeaderBlock {
spans,
next: 0,
data_start,
})
}
Ok(httparse::Status::Partial) => Err(Error::IncompleteStream),
Err(httparse::Error::TooManyHeaders) => parse_header_block_fallback(block),
Err(_) => Err(Error::InvalidFormat),
}
}
fn parse_header_block_fallback(block: &[u8]) -> Result<HeaderBlock, Error> {
if !block.ends_with(b"\r\n\r\n") {
return Err(Error::InvalidFormat);
}
let blank = block.len().saturating_sub(4);
let header_text = &block[..blank];
let mut spans = ArrayVec::<HeaderSpan, MAX_HEADERS>::new();
let mut rest = header_text;
while !rest.is_empty() && spans.len() < MAX_HEADERS {
let (line, next_rest) = match memmem::find(rest, b"\r\n") {
Some(line_end) => (&rest[..line_end], &rest[line_end.saturating_add(2)..]),
None => (rest, &b""[..]),
};
let colon = memchr::memchr(b':', line).ok_or(Error::InvalidFormat)?;
let name = &line[..colon];
if name.is_empty() {
return Err(Error::InvalidFormat);
}
let value = trim_ows(&line[colon.saturating_add(1)..]);
let name_start = name.as_ptr() as usize - block.as_ptr() as usize;
let name_end = name_start.saturating_add(name.len());
let value_start = value.as_ptr() as usize - block.as_ptr() as usize;
let value_end = value_start.saturating_add(value.len());
let _ = spans.try_push(HeaderSpan {
name: name_start..name_end,
value: value_start..value_end,
});
rest = next_rest;
}
Ok(HeaderBlock {
spans,
next: 0,
data_start: blank.saturating_add(4),
})
}
impl<S> Multipart<S>
where
S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
{
pub(super) fn next_header_inner(&mut self) -> Result<Option<httparse::Header<'_>>, Error> {
let Some(block) = self.headers.as_mut() else {
return Ok(None);
};
let Some(span) = block.spans.get(block.next).cloned() else {
self.finish_headers();
return Ok(None);
};
block.next = block.next.saturating_add(1);
let buffer = self.buffer_mut()?;
let name_bytes = buffer.buf.get(span.name.clone()).ok_or(Error::InvalidFormat)?;
let value = buffer.buf.get(span.value.clone()).ok_or(Error::InvalidFormat)?;
let name = std::str::from_utf8(name_bytes).map_err(|_| Error::InvalidFormat)?;
Ok(Some(httparse::Header { name, value }))
}
pub(super) fn finish_headers(&mut self) {
let Some(block) = self.headers.take() else {
return;
};
if self.state == State::ReadingPartHeaders {
let _ = self.buffer_mut().map(|buf| buf.buf.split_to(block.data_start));
self.state = State::ReadingPartData;
}
self.part_handed_out = false;
}
pub(super) fn poll_ensure_headers(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
if self.headers.is_some() || self.state == State::ReadingPartData {
return Poll::Ready(Ok(()));
}
if self.state != State::ReadingPartHeaders {
return Poll::Ready(Ok(()));
}
let max = self.max_buffer_size;
let end = {
let buffer = match self.buffer_mut() {
Ok(buffer) => buffer,
Err(err) => return Poll::Ready(Err(err)),
};
ready!(buffer.poll_read_header_block(max, cx))?
};
let block = match self.buffer_mut() {
Ok(buffer) => parse_header_block(&buffer.buf[..end]),
Err(err) => return Poll::Ready(Err(err)),
};
match block {
Ok(block) => {
self.headers = Some(block);
Poll::Ready(Ok(()))
}
Err(err) => Poll::Ready(Err(err)),
}
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::panic, clippy::unreachable, clippy::unwrap_used)]
mod tests {
use super::*;
use std::fmt::Write as _;
use futures_util::stream;
use crate::Boundary;
fn parser_from_chunks(
chunks: Vec<Result<&'static [u8], std::io::Error>>,
) -> Multipart<impl Stream<Item = Result<Bytes, Error>> + Send + Sync> {
Multipart::new(
stream::iter(
chunks
.into_iter()
.map(|item| item.map(Bytes::from_static).map_err(Error::stream_read_failed)),
),
&Boundary::new(b"boundary").unwrap(),
4096,
)
}
#[test]
fn header_parser_error_and_fallback_paths() {
assert!(matches!(parse_header_block(b"X: y\r"), Err(Error::IncompleteStream)));
assert!(matches!(parse_header_block(b"\0bad"), Err(Error::InvalidFormat)));
let mut block = String::new();
for idx in 0..40 {
let _ = write!(block, "X-{idx}: value-{idx}\r\n");
}
block.push_str("\r\n");
let parsed = parse_header_block(block.as_bytes()).unwrap();
assert_eq!(parsed.spans.len(), MAX_HEADERS);
assert!(matches!(parse_header_block_fallback(b": y\r\n\r\n"), Err(Error::InvalidFormat)));
assert!(matches!(parse_header_block_fallback(b"X: y"), Err(Error::InvalidFormat)));
assert!(matches!(parse_header_block_fallback(b"X: y\r\n"), Err(Error::InvalidFormat)));
assert!(parse_header_block_fallback(b"X: y\r\n\r\n").is_ok());
}
#[test]
fn fallback_stops_at_the_cap_without_reading_surplus_fields() {
let mut block = String::from("X-0: 0\r\nX-1: 1\r\nX-2: 2\r\nnot a header line\r\n");
block.push_str("\r\n");
let parsed = parse_header_block_fallback(block.as_bytes()).unwrap();
assert_eq!(parsed.spans.len(), MAX_HEADERS);
assert_eq!(parsed.data_start, block.len());
assert!(matches!(parse_header_block(block.as_bytes()), Err(Error::InvalidFormat)));
}
#[test]
fn header_offsets_come_from_the_parsed_block() {
assert!(matches!(parse_header_block(b"X: y"), Err(Error::IncompleteStream)));
assert!(matches!(parse_header_block(b"X: y\r\n"), Err(Error::IncompleteStream)));
let parsed = parse_header_block(b"X:\r\n\r\n").unwrap();
assert_eq!(parsed.spans.len(), 1);
assert_eq!(parsed.spans[0].name, 0..1);
assert_eq!(parsed.spans[0].value, 2..2);
assert_eq!(parsed.data_start, 6);
let parsed = parse_header_block(b"\r\n").unwrap();
assert!(parsed.spans.is_empty());
assert_eq!(parsed.data_start, 2);
let parsed = parse_header_block(b"\r\n\r\n").unwrap();
assert!(parsed.spans.is_empty());
assert_eq!(parsed.data_start, 2);
let mut many = String::from("X-0: 0\r\nX-1: 1\r\nX-2: 2\r\nX-3: 3\r\n");
let len = many.len();
many.push_str("\r\n");
let parsed = parse_header_block_fallback(many.as_bytes()).unwrap();
assert_eq!(parsed.spans.len(), MAX_HEADERS);
assert_eq!(parsed.data_start, len + 2);
}
#[test]
fn finish_headers_keeps_non_header_state() {
let mut mp = parser_from_chunks(vec![Ok(b"--boundary\r\nX: y\r\n\r\ndata\r\n--boundary--\r\n")]);
mp.state = State::ReadingBoundary;
mp.headers = Some(HeaderBlock {
spans: ArrayVec::new(),
next: 0,
data_start: 0,
});
mp.finish_headers();
assert_eq!(mp.state, State::ReadingBoundary);
}
}