use std::convert::Infallible;
use crate::storage::StorageError;
pub const DEFAULT_MAX_MULTIPART_BYTES: usize = 32 * 1024 * 1024;
#[derive(Clone, Debug)]
pub struct FilePart {
pub field_name: String,
pub filename: Option<String>,
pub content_type: Option<String>,
pub bytes: Vec<u8>,
}
#[derive(Debug, Default)]
pub struct MultipartForm {
pub fields: Vec<(String, String)>,
pub files: Vec<FilePart>,
}
impl MultipartForm {
pub fn field(&self, name: &str) -> Option<&str> {
self.fields
.iter()
.rev()
.find(|(k, _)| k == name)
.map(|(_, v)| v.as_str())
}
pub fn iter_fields(&self) -> impl Iterator<Item = (&str, &str)> {
self.fields.iter().map(|(k, v)| (k.as_str(), v.as_str()))
}
}
#[derive(Debug)]
pub enum MultipartError {
MissingBoundary,
Parse(String),
TooLarge {
limit: usize,
actual: usize,
},
}
impl std::fmt::Display for MultipartError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
MultipartError::MissingBoundary => {
write!(f, "multipart: Content-Type has no boundary parameter")
}
MultipartError::Parse(s) => write!(f, "multipart: parse error: {s}"),
MultipartError::TooLarge { limit, actual } => write!(
f,
"multipart: body {actual}B exceeds configured cap of {limit}B"
),
}
}
}
impl std::error::Error for MultipartError {}
#[derive(Debug)]
pub enum MultipartUploadError {
Multipart(MultipartError),
Storage(StorageError),
NoStorageBackend,
}
impl std::fmt::Display for MultipartUploadError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
MultipartUploadError::Multipart(e) => write!(f, "{e}"),
MultipartUploadError::Storage(e) => write!(f, "{e}"),
MultipartUploadError::NoStorageBackend => write!(
f,
"multipart upload: no Storage backend registered; add StoragePlugin \
or call umbral::storage::set_storage"
),
}
}
}
impl std::error::Error for MultipartUploadError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
MultipartUploadError::Multipart(e) => Some(e),
MultipartUploadError::Storage(e) => Some(e),
MultipartUploadError::NoStorageBackend => None,
}
}
}
impl From<MultipartError> for MultipartUploadError {
fn from(e: MultipartError) -> Self {
MultipartUploadError::Multipart(e)
}
}
impl From<StorageError> for MultipartUploadError {
fn from(e: StorageError) -> Self {
MultipartUploadError::Storage(e)
}
}
pub fn is_multipart(content_type: &str) -> bool {
content_type
.trim_start()
.to_ascii_lowercase()
.starts_with("multipart/form-data")
}
pub async fn parse_multipart(
content_type_header: &str,
body: impl Into<bytes::Bytes>,
) -> Result<MultipartForm, MultipartError> {
parse_multipart_capped(content_type_header, body, DEFAULT_MAX_MULTIPART_BYTES).await
}
pub async fn parse_multipart_capped(
content_type_header: &str,
body: impl Into<bytes::Bytes>,
max_bytes: usize,
) -> Result<MultipartForm, MultipartError> {
let boundary =
multer::parse_boundary(content_type_header).map_err(|_| MultipartError::MissingBoundary)?;
let body: bytes::Bytes = body.into();
if body.len() > max_bytes {
return Err(MultipartError::TooLarge {
limit: max_bytes,
actual: body.len(),
});
}
let stream = futures_util::stream::once(async move { Ok::<_, Infallible>(body) });
let mut multipart = multer::Multipart::new(stream, boundary);
let mut form = MultipartForm::default();
let mut accumulated: usize = 0;
while let Some(field) = multipart
.next_field()
.await
.map_err(|e| MultipartError::Parse(e.to_string()))?
{
let field_name = field.name().map(str::to_owned).unwrap_or_default();
let filename = field.file_name().map(str::to_owned);
let content_type = field.content_type().map(|m| m.to_string());
if filename.is_some() {
let bytes = field
.bytes()
.await
.map_err(|e| MultipartError::Parse(e.to_string()))?;
accumulated = accumulated.saturating_add(bytes.len());
if accumulated > max_bytes {
return Err(MultipartError::TooLarge {
limit: max_bytes,
actual: accumulated,
});
}
form.files.push(FilePart {
field_name,
filename,
content_type,
bytes: bytes.to_vec(),
});
} else {
let value = field
.text()
.await
.map_err(|e| MultipartError::Parse(e.to_string()))?;
accumulated = accumulated.saturating_add(value.len());
if accumulated > max_bytes {
return Err(MultipartError::TooLarge {
limit: max_bytes,
actual: accumulated,
});
}
form.fields.push((field_name, value));
}
}
Ok(form)
}
pub async fn parse_and_store_multipart(
content_type_header: &str,
body: impl Into<bytes::Bytes>,
) -> Result<Vec<(String, String)>, MultipartUploadError> {
let form = parse_multipart(content_type_header, body).await?;
let mut pairs: Vec<(String, String)> = Vec::new();
for file in &form.files {
if file.bytes.is_empty() {
continue;
}
let backend =
crate::storage::storage_opt().ok_or(MultipartUploadError::NoStorageBackend)?;
let filename = file
.filename
.as_deref()
.filter(|s| !s.is_empty())
.unwrap_or(&file.field_name);
let content_type = file
.content_type
.as_deref()
.unwrap_or("application/octet-stream");
let stored = backend.store(filename, content_type, &file.bytes).await?;
pairs.push((file.field_name.clone(), stored.key));
}
pairs.extend(form.fields);
Ok(pairs)
}
#[cfg(test)]
mod tests {
use super::*;
const BOUNDARY: &str = "X-UMBRAL-BOUNDARY";
type PartSpec<'a> = (&'a str, Option<&'a str>, Option<&'a str>, &'a [u8]);
fn build_body(parts: &[PartSpec<'_>]) -> Vec<u8> {
let mut out = Vec::new();
for (name, filename, content_type, value) in parts {
out.extend_from_slice(format!("--{BOUNDARY}\r\n").as_bytes());
match filename {
Some(fname) => {
out.extend_from_slice(
format!(
"Content-Disposition: form-data; name=\"{name}\"; filename=\"{fname}\"\r\n"
)
.as_bytes(),
);
if let Some(ct) = content_type {
out.extend_from_slice(format!("Content-Type: {ct}\r\n").as_bytes());
}
}
None => {
out.extend_from_slice(
format!("Content-Disposition: form-data; name=\"{name}\"\r\n").as_bytes(),
);
}
}
out.extend_from_slice(b"\r\n");
out.extend_from_slice(value);
out.extend_from_slice(b"\r\n");
}
out.extend_from_slice(format!("--{BOUNDARY}--\r\n").as_bytes());
out
}
fn ct_header() -> String {
format!("multipart/form-data; boundary={BOUNDARY}")
}
#[test]
fn is_multipart_matches_form_data_content_types() {
assert!(is_multipart("multipart/form-data; boundary=abc"));
assert!(is_multipart("multipart/form-data"));
assert!(is_multipart(" Multipart/Form-Data; boundary=Z")); assert!(!is_multipart("application/x-www-form-urlencoded"));
assert!(!is_multipart("application/json"));
assert!(!is_multipart("multipart/mixed; boundary=abc"));
}
#[tokio::test]
async fn parse_separates_text_and_file_parts() {
let png = b"\x89PNG\r\n\x1a\nfake-image-bytes";
let body = build_body(&[
("title", None, None, b"Hello"),
("cover", Some("p.png"), Some("image/png"), png),
]);
let form = parse_multipart(&ct_header(), body).await.unwrap();
assert_eq!(
form.fields,
vec![("title".to_string(), "Hello".to_string())]
);
assert_eq!(form.files.len(), 1);
let file = &form.files[0];
assert_eq!(file.field_name, "cover");
assert_eq!(file.filename.as_deref(), Some("p.png"));
assert_eq!(file.content_type.as_deref(), Some("image/png"));
assert_eq!(file.bytes, png);
}
#[tokio::test]
async fn parse_preserves_repeated_text_field_names() {
let body = build_body(&[
("tags", None, None, b"red"),
("tags", None, None, b"blue"),
("name", None, None, b"shirt"),
]);
let form = parse_multipart(&ct_header(), body).await.unwrap();
assert_eq!(
form.fields,
vec![
("tags".to_string(), "red".to_string()),
("tags".to_string(), "blue".to_string()),
("name".to_string(), "shirt".to_string()),
]
);
assert_eq!(form.field("tags"), Some("blue"));
assert_eq!(form.field("name"), Some("shirt"));
assert_eq!(form.field("missing"), None);
assert_eq!(form.iter_fields().filter(|(k, _)| *k == "tags").count(), 2);
}
#[tokio::test]
async fn parse_keeps_binary_bytes_intact() {
let raw: Vec<u8> = vec![0x00, 0xFF, 0xFE, 0x80, 0x01, 0x7F];
assert!(std::str::from_utf8(&raw).is_err());
let body = build_body(&[(
"blob",
Some("data.bin"),
Some("application/octet-stream"),
&raw,
)]);
let form = parse_multipart(&ct_header(), body).await.unwrap();
assert_eq!(form.files.len(), 1);
assert_eq!(form.files[0].bytes, raw, "raw bytes must round-trip");
}
#[tokio::test]
async fn capped_rejects_body_over_the_limit() {
let big = vec![b'x'; 4096];
let body = build_body(&[(
"blob",
Some("big.bin"),
Some("application/octet-stream"),
&big,
)]);
let err = parse_multipart_capped(&ct_header(), body, 1024)
.await
.unwrap_err();
match err {
MultipartError::TooLarge { limit, actual } => {
assert_eq!(limit, 1024);
assert!(actual > 1024, "reports the offending size, got {actual}");
}
other => panic!("expected TooLarge, got {other:?}"),
}
}
#[tokio::test]
async fn capped_allows_body_under_the_limit() {
let small = b"hello";
let body = build_body(&[(
"blob",
Some("s.bin"),
Some("application/octet-stream"),
small,
)]);
let form = parse_multipart_capped(&ct_header(), body, 1024)
.await
.unwrap();
assert_eq!(form.files.len(), 1);
assert_eq!(form.files[0].bytes, small);
}
#[tokio::test]
async fn capped_counts_text_field_bytes_too() {
let big = vec![b'a'; 4096];
let body = build_body(&[("notes", None, None, &big)]);
let err = parse_multipart_capped(&ct_header(), body, 512)
.await
.unwrap_err();
assert!(matches!(err, MultipartError::TooLarge { .. }));
}
#[tokio::test]
async fn parse_errors_on_missing_boundary() {
let body = build_body(&[("title", None, None, b"Hi")]);
let err = parse_multipart("multipart/form-data", body)
.await
.unwrap_err();
assert!(matches!(err, MultipartError::MissingBoundary));
}
}