use crate::call::Call;
use crate::error::{Error, Result};
use crate::extract::FromCall;
use async_trait::async_trait;
use bytes::Bytes;
const MAX_PARTS: usize = 256;
#[derive(Debug, Clone)]
pub struct Part {
pub name: String,
pub filename: Option<String>,
pub content_type: Option<String>,
pub bytes: Bytes,
}
impl Part {
pub fn text(&self) -> Option<String> {
String::from_utf8(self.bytes.to_vec()).ok()
}
}
#[derive(Debug, Clone, Default)]
pub struct Multipart {
parts: Vec<Part>,
}
impl Multipart {
pub fn parts(&self) -> &[Part] {
&self.parts
}
pub fn part(&self, name: &str) -> Option<&Part> {
self.parts.iter().find(|p| p.name == name)
}
pub fn file(&self, name: &str) -> Option<&Part> {
self.parts
.iter()
.find(|p| p.name == name && p.filename.is_some())
}
pub fn field(&self, name: &str) -> Option<String> {
self.part(name).and_then(|p| p.text())
}
}
#[async_trait]
impl FromCall for Multipart {
async fn from_call(mut call: Call) -> Result<Self> {
let ct = call
.header(http::header::CONTENT_TYPE.as_str())
.unwrap_or("")
.to_string();
let boundary = boundary_of(&ct).ok_or_else(|| {
Error::new(
http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
"expected multipart/form-data with a boundary",
)
})?;
let body = call.try_receive_bytes().await?;
crate::extract::check_body_limit(&call, body.len())?;
parse(&body, &boundary).map(|parts| Multipart { parts })
}
}
fn boundary_of(content_type: &str) -> Option<String> {
let mut it = content_type.split(';');
let media = it.next()?.trim();
if !media.eq_ignore_ascii_case("multipart/form-data") {
return None;
}
for param in it {
let Some((k, v)) = param.split_once('=') else {
continue;
};
if k.trim().eq_ignore_ascii_case("boundary") {
let value = v.trim().trim_matches('"');
if value.is_empty() {
return None;
}
return Some(value.to_string());
}
}
None
}
const MAX_TRANSPORT_PADDING: usize = 64;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Tail {
Part(usize),
Close,
Content,
NeedMore,
}
fn delimiter_tail(after: &[u8], at_end: bool) -> Tail {
match (after.first(), after.get(1)) {
(Some(b'-'), Some(b'-')) => return Tail::Close,
(Some(b'-'), Some(_)) => return Tail::Content,
(Some(b'-'), None) => {
return if at_end {
Tail::Content
} else {
Tail::NeedMore
}
}
_ => {}
}
let mut i = 0;
loop {
match after.get(i) {
Some(b' ' | b'\t') if i < MAX_TRANSPORT_PADDING => i += 1,
Some(b'\r') => {
return match after.get(i + 1) {
Some(b'\n') => Tail::Part(i + 2),
Some(_) => Tail::Content,
None if at_end => Tail::Content,
None => Tail::NeedMore,
}
}
Some(_) => return Tail::Content,
None if at_end => return Tail::Content,
None => return Tail::NeedMore,
}
}
}
fn parse(body: &[u8], boundary: &str) -> Result<Vec<Part>> {
let delim = format!("\r\n--{boundary}").into_bytes();
let mut framed = Vec::with_capacity(body.len() + 2);
framed.extend_from_slice(b"\r\n");
framed.extend_from_slice(body);
let body: &[u8] = &framed;
let mut parts = Vec::new();
let mut part_at: Option<usize> = None;
let mut closed = false;
let mut at = 0;
while let Some(rel) = find(&body[at..], &delim) {
let hit = at + rel;
match delimiter_tail(&body[hit + delim.len()..], true) {
Tail::Part(skip) => {
if let Some(from) = part_at {
push_part(&mut parts, &body[from..hit])?;
}
at = hit + delim.len() + skip;
part_at = Some(at);
}
Tail::Close => {
if let Some(from) = part_at.take() {
push_part(&mut parts, &body[from..hit])?;
}
closed = true;
break;
}
Tail::Content | Tail::NeedMore => at = hit + 1,
}
}
if part_at.is_some() {
return Err(Error::bad_request("multipart body ended inside a part"));
}
if !closed {
return Err(Error::bad_request(
"multipart body has no boundary delimiter",
));
}
Ok(parts)
}
fn push_part(parts: &mut Vec<Part>, chunk: &[u8]) -> Result<()> {
let Some(split) = find(chunk, b"\r\n\r\n") else {
return Err(Error::bad_request(
"malformed multipart part: no header block",
));
};
let (head, content) = chunk.split_at(split);
let content = &content[4..];
let headers = std::str::from_utf8(head)
.map_err(|_| Error::bad_request("multipart headers are not valid UTF-8"))?;
let mut name = None;
let mut filename = None;
let mut content_type = None;
for line in headers.split("\r\n") {
let Some((k, v)) = line.split_once(':') else {
continue;
};
match k.trim().to_ascii_lowercase().as_str() {
"content-disposition" => {
name = param_of(v, "name");
filename = param_of(v, "filename");
}
"content-type" => content_type = Some(v.trim().to_string()),
_ => {}
}
}
let Some(name) = name else {
return Err(Error::bad_request(
"multipart part is missing a Content-Disposition name",
));
};
if parts.len() >= MAX_PARTS {
return Err(Error::new(
http::StatusCode::PAYLOAD_TOO_LARGE,
"too many multipart parts",
));
}
parts.push(Part {
name,
filename,
content_type,
bytes: Bytes::copy_from_slice(content),
});
Ok(())
}
const MAX_PART_HEADER_BYTES: usize = 8 * 1024;
const DEFAULT_MAX_FIELD_BYTES: usize = 8 * 1024 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Phase {
Preamble,
AfterDelimiter,
InPart,
Done,
}
pub struct MultipartStream {
stream: crate::call::BodyStream,
buffer: bytes::BytesMut,
delimiter: Vec<u8>,
phase: Phase,
drained: bool,
parts_seen: usize,
max_field_bytes: usize,
}
impl std::fmt::Debug for MultipartStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MultipartStream")
.field("phase", &self.phase)
.field("buffered", &self.buffer.len())
.field("parts_seen", &self.parts_seen)
.finish_non_exhaustive()
}
}
impl MultipartStream {
fn new(stream: crate::call::BodyStream, boundary: &str) -> Self {
let mut buffer = bytes::BytesMut::new();
buffer.extend_from_slice(b"\r\n");
Self {
stream,
buffer,
delimiter: format!("\r\n--{boundary}").into_bytes(),
phase: Phase::Preamble,
drained: false,
parts_seen: 0,
max_field_bytes: DEFAULT_MAX_FIELD_BYTES,
}
}
pub fn max_field_bytes(mut self, bytes: usize) -> Self {
self.max_field_bytes = bytes;
self
}
async fn fill(&mut self) -> Result<bool> {
if self.drained {
return Ok(false);
}
use futures_util::StreamExt;
match self.stream.next().await {
Some(Ok(chunk)) => {
self.buffer.extend_from_slice(&chunk);
Ok(true)
}
Some(Err(e)) => Err(e),
None => {
self.drained = true;
Ok(false)
}
}
}
pub async fn next_field(&mut self) -> Result<Option<Field<'_>>> {
if self.phase == Phase::InPart {
self.skip_rest_of_part().await?;
}
if self.phase == Phase::Preamble {
self.consume_through_delimiter().await?;
}
if self.phase == Phase::Done {
return Ok(None);
}
let (name, filename, content_type) = self.read_part_headers().await?;
self.parts_seen += 1;
if self.parts_seen > MAX_PARTS {
return Err(Error::new(
http::StatusCode::PAYLOAD_TOO_LARGE,
"too many multipart parts",
));
}
self.phase = Phase::InPart;
let max_field_bytes = self.max_field_bytes;
Ok(Some(Field {
name,
filename,
content_type,
max_field_bytes,
parser: self,
}))
}
async fn classify_at(&mut self, at: usize) -> Result<Tail> {
loop {
let after = at + self.delimiter.len();
let tail = delimiter_tail(&self.buffer[after..], self.drained);
if tail != Tail::NeedMore {
return Ok(tail);
}
self.fill().await?;
}
}
async fn consume_through_delimiter(&mut self) -> Result<()> {
loop {
if let Some(at) = find(&self.buffer, &self.delimiter) {
match self.classify_at(at).await? {
Tail::Part(skip) => {
let _ = self.buffer.split_to(at + self.delimiter.len() + skip);
self.phase = Phase::AfterDelimiter;
return Ok(());
}
Tail::Close => {
self.phase = Phase::Done;
return Ok(());
}
Tail::Content | Tail::NeedMore => {
let _ = self.buffer.split_to(at + 1);
continue;
}
}
}
let keep = self.delimiter.len().saturating_sub(1);
if self.buffer.len() > keep {
let drop_to = self.buffer.len() - keep;
let _ = self.buffer.split_to(drop_to);
}
if !self.fill().await? {
return Err(Error::bad_request(
"multipart body has no boundary delimiter",
));
}
}
}
async fn read_part_headers(&mut self) -> Result<(String, Option<String>, Option<String>)> {
let split = loop {
if let Some(at) = find(&self.buffer, b"\r\n\r\n") {
break at;
}
if self.buffer.len() > MAX_PART_HEADER_BYTES {
return Err(Error::new(
http::StatusCode::PAYLOAD_TOO_LARGE,
"multipart part headers are too large",
));
}
if !self.fill().await? {
return Err(Error::bad_request(
"malformed multipart part: no header block",
));
}
};
let head = self.buffer.split_to(split);
let _ = self.buffer.split_to(4);
let headers = std::str::from_utf8(&head)
.map_err(|_| Error::bad_request("multipart headers are not valid UTF-8"))?;
let mut name = None;
let mut filename = None;
let mut content_type = None;
for line in headers.split("\r\n") {
let Some((k, v)) = line.split_once(':') else {
continue;
};
match k.trim().to_ascii_lowercase().as_str() {
"content-disposition" => {
name = param_of(v, "name");
filename = param_of(v, "filename");
}
"content-type" => content_type = Some(v.trim().to_string()),
_ => {}
}
}
let name = name.ok_or_else(|| {
Error::bad_request("multipart part is missing a Content-Disposition name")
})?;
Ok((name, filename, content_type))
}
async fn skip_rest_of_part(&mut self) -> Result<()> {
while self.next_content_chunk().await?.is_some() {}
Ok(())
}
async fn next_content_chunk(&mut self) -> Result<Option<Bytes>> {
if self.phase != Phase::InPart {
return Ok(None);
}
let mut from = 0;
loop {
if let Some(rel) = find(&self.buffer[from..], &self.delimiter) {
let at = from + rel;
match self.classify_at(at).await? {
Tail::Part(skip) => {
let data = self.buffer.split_to(at).freeze();
let _ = self.buffer.split_to(self.delimiter.len() + skip);
self.phase = Phase::AfterDelimiter;
return Ok((!data.is_empty()).then_some(data));
}
Tail::Close => {
let data = self.buffer.split_to(at).freeze();
self.phase = Phase::Done;
return Ok((!data.is_empty()).then_some(data));
}
Tail::Content | Tail::NeedMore => {
from = at + 1;
continue;
}
}
}
let hold = self.delimiter.len().saturating_sub(1);
if self.buffer.len() > hold {
let take = self.buffer.len() - hold;
let data = self.buffer.split_to(take).freeze();
if !data.is_empty() {
return Ok(Some(data));
}
}
if !self.fill().await? {
return Err(Error::bad_request("multipart body ended inside a part"));
}
}
}
}
#[async_trait]
impl FromCall for MultipartStream {
async fn from_call(call: Call) -> Result<Self> {
let ct = call
.header(http::header::CONTENT_TYPE.as_str())
.unwrap_or("")
.to_string();
let boundary = boundary_of(&ct).ok_or_else(|| {
Error::new(
http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
"expected multipart/form-data with a boundary",
)
})?;
let crate::extract::Payload(stream) = crate::extract::Payload::from_call(call).await?;
Ok(MultipartStream::new(stream, &boundary))
}
}
pub struct Field<'a> {
name: String,
filename: Option<String>,
content_type: Option<String>,
max_field_bytes: usize,
parser: &'a mut MultipartStream,
}
impl std::fmt::Debug for Field<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Field")
.field("name", &self.name)
.field("filename", &self.filename)
.field("content_type", &self.content_type)
.finish_non_exhaustive()
}
}
impl Field<'_> {
pub fn name(&self) -> &str {
&self.name
}
pub fn filename(&self) -> Option<&str> {
self.filename.as_deref()
}
pub fn content_type(&self) -> Option<&str> {
self.content_type.as_deref()
}
pub async fn chunk(&mut self) -> Result<Option<Bytes>> {
self.parser.next_content_chunk().await
}
pub async fn bytes(&mut self) -> Result<Bytes> {
let mut out = bytes::BytesMut::new();
while let Some(chunk) = self.chunk().await? {
if out.len() + chunk.len() > self.max_field_bytes {
return Err(Error::new(
http::StatusCode::PAYLOAD_TOO_LARGE,
"multipart field too large",
));
}
out.extend_from_slice(&chunk);
}
Ok(out.freeze())
}
pub async fn text(&mut self) -> Result<String> {
let raw = self.bytes().await?;
String::from_utf8(raw.to_vec())
.map_err(|_| Error::bad_request("multipart field is not valid UTF-8"))
}
}
fn param_of(value: &str, key: &str) -> Option<String> {
for param in value.split(';') {
let Some((k, v)) = param.split_once('=') else {
continue;
};
if k.trim().eq_ignore_ascii_case(key) {
return Some(v.trim().trim_matches('"').to_string());
}
}
None
}
fn find(haystack: &[u8], needle: &[u8]) -> Option<usize> {
haystack.windows(needle.len()).position(|w| w == needle)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_delimiter_line_admits_only_spaces_and_tabs_before_its_crlf() {
assert_eq!(delimiter_tail(b"\r\nnext", true), Tail::Part(2));
assert_eq!(delimiter_tail(b" \t \r\nnext", true), Tail::Part(5));
assert_eq!(delimiter_tail(b"--\r\n", true), Tail::Close);
assert_eq!(delimiter_tail(b"z\r\nnext", true), Tail::Content);
assert_eq!(delimiter_tail(b"-z", true), Tail::Content);
assert_eq!(delimiter_tail(b"\rx", true), Tail::Content);
let over = vec![b' '; MAX_TRANSPORT_PADDING + 1];
assert_eq!(delimiter_tail(&over, true), Tail::Content);
assert_eq!(delimiter_tail(b"", false), Tail::NeedMore);
assert_eq!(delimiter_tail(b"\r", false), Tail::NeedMore);
assert_eq!(delimiter_tail(b"-", false), Tail::NeedMore);
assert_eq!(delimiter_tail(b"", true), Tail::Content);
assert_eq!(delimiter_tail(b"\r", true), Tail::Content);
}
#[test]
fn reads_the_boundary() {
assert_eq!(
boundary_of("multipart/form-data; boundary=abc").as_deref(),
Some("abc")
);
assert_eq!(
boundary_of("multipart/form-data; boundary=\"a b\"").as_deref(),
Some("a b")
);
assert_eq!(boundary_of("application/json"), None);
assert_eq!(boundary_of("multipart/form-data"), None);
}
}