use std::{
borrow::Cow,
io,
path::Path,
pin::Pin,
task::{Context, Poll},
};
use bytes::{Bytes, BytesMut};
use futures_util::{StreamExt, stream::Stream};
use http::header::HeaderValue;
use http_body::{Frame, SizeHint};
use http_body_util::StreamBody;
use mime::Mime;
use pin_project::pin_project;
use tokio::io::AsyncRead;
use tokio_util::io::ReaderStream;
use crate::{body::Body, error::BoxError};
type FrameStream = Pin<Box<dyn Stream<Item = Result<Frame<Bytes>, BoxError>> + Send + Sync>>;
#[pin_project]
struct SizedStreamBody<S> {
#[pin]
inner: StreamBody<S>,
exact_size: u64,
}
impl<S> http_body::Body for SizedStreamBody<S>
where
S: Stream<Item = Result<Frame<Bytes>, BoxError>>,
{
type Data = Bytes;
type Error = BoxError;
fn poll_frame(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
http_body::Body::poll_frame(self.project().inner, cx)
}
fn is_end_stream(&self) -> bool {
http_body::Body::is_end_stream(&self.inner)
}
fn size_hint(&self) -> SizeHint {
SizeHint::with_exact(self.exact_size)
}
}
fn exact_size_stream_body<S>(stream: S, exact_size: u64) -> Body
where
S: Stream<Item = Result<Frame<Bytes>, BoxError>> + Send + Sync + 'static,
{
Body::from_body(SizedStreamBody {
inner: StreamBody::new(stream),
exact_size,
})
}
#[must_use]
pub struct Form {
boundary: String,
parts: Vec<(Cow<'static, str>, Part)>,
}
impl Default for Form {
fn default() -> Self {
Self::new()
}
}
impl Form {
pub fn new() -> Self {
Self {
boundary: gen_boundary(),
parts: Vec::new(),
}
}
pub fn boundary(&self) -> &str {
&self.boundary
}
pub fn text<N, V>(self, name: N, value: V) -> Self
where
N: Into<Cow<'static, str>>,
V: Into<Cow<'static, str>>,
{
self.part(name, Part::text(value))
}
pub fn part<N>(mut self, name: N, part: Part) -> Self
where
N: Into<Cow<'static, str>>,
{
self.parts.push((name.into(), part));
self
}
pub async fn file<N, P>(self, name: N, path: P) -> io::Result<Self>
where
N: Into<Cow<'static, str>>,
P: AsRef<Path>,
{
Ok(self.part(name, Part::file(path).await?))
}
pub(crate) fn content_type(&self) -> HeaderValue {
HeaderValue::from_str(&format!("multipart/form-data; boundary={}", self.boundary))
.expect("multipart boundary should always be a valid header value")
}
pub(crate) fn into_body(self) -> Body {
let boundary = self.boundary;
if self.parts.iter().all(|(_, part)| part.data.is_in_memory()) {
let cap = self
.parts
.iter()
.map(|(name, part)| {
part.data.as_bytes().map_or(0, Bytes::len)
+ name.len()
+ part.file_name.as_ref().map_or(0, |f| f.len())
+ part.mime.as_ref().map_or(0, |m| m.as_ref().len())
+ boundary.len()
+ 96
})
.sum::<usize>()
+ boundary.len()
+ 8;
let mut buf = BytesMut::with_capacity(cap);
for (name, part) in &self.parts {
buf.extend_from_slice(&part.encode_header(&boundary, name));
if let Some(bytes) = part.data.as_bytes() {
buf.extend_from_slice(bytes);
}
buf.extend_from_slice(b"\r\n");
}
buf.extend_from_slice(format!("--{boundary}--\r\n").as_bytes());
return Body::from(buf.freeze());
}
let closing_boundary = Bytes::from(format!("--{boundary}--\r\n"));
let mut exact_len = Some(closing_boundary.len() as u64);
let mut streams: Vec<FrameStream> = Vec::with_capacity(self.parts.len() * 3 + 1);
for (name, part) in self.parts {
let header = part.encode_header(&boundary, &name);
exact_len = exact_len.and_then(|total| {
let data_len = part.len()?;
total
.checked_add(header.len() as u64)?
.checked_add(data_len)?
.checked_add(2)
});
streams.push(once_frame(header));
streams.push(part.data.into_stream());
streams.push(once_frame(Bytes::from_static(b"\r\n")));
}
streams.push(once_frame(closing_boundary));
let stream = futures_util::stream::iter(streams).flatten();
match exact_len {
Some(exact_len) => exact_size_stream_body(stream, exact_len),
None => Body::from_stream(stream),
}
}
}
#[must_use]
pub struct Part {
data: PartData,
length: Option<u64>,
file_name: Option<Cow<'static, str>>,
mime: Option<Mime>,
}
enum PartData {
Bytes(Bytes),
Stream(FrameStream),
}
impl PartData {
fn into_stream(self) -> FrameStream {
match self {
PartData::Bytes(bytes) => once_frame(bytes),
PartData::Stream(stream) => stream,
}
}
fn len(&self) -> Option<u64> {
match self {
PartData::Bytes(bytes) => Some(bytes.len() as u64),
PartData::Stream(_) => None,
}
}
fn is_in_memory(&self) -> bool {
matches!(self, PartData::Bytes(_))
}
fn as_bytes(&self) -> Option<&Bytes> {
match self {
PartData::Bytes(bytes) => Some(bytes),
PartData::Stream(_) => None,
}
}
}
impl Part {
pub fn text<T>(value: T) -> Self
where
T: Into<Cow<'static, str>>,
{
let bytes = match value.into() {
Cow::Borrowed(s) => Bytes::from_static(s.as_bytes()),
Cow::Owned(s) => Bytes::from(s),
};
Self::new(PartData::Bytes(bytes))
}
pub fn bytes<T>(value: T) -> Self
where
T: Into<Bytes>,
{
Self::new(PartData::Bytes(value.into()))
}
pub fn reader<R>(reader: R) -> Self
where
R: AsyncRead + Send + Sync + 'static,
{
let stream =
ReaderStream::new(reader).map(|res| res.map(Frame::data).map_err(BoxError::from));
Self::new(PartData::Stream(Box::pin(stream)))
}
pub async fn file<P>(path: P) -> io::Result<Self>
where
P: AsRef<Path>,
{
let path = path.as_ref();
let file_name = path
.file_name()
.map(|name| name.to_string_lossy().into_owned());
let mime = mime_guess::from_path(path).first();
let file = tokio::fs::File::open(path).await?;
let file_len = file.metadata().await?.len();
let mut part = Self::reader(file);
part.length = Some(file_len);
if let Some(file_name) = file_name {
part = part.file_name(file_name);
}
if let Some(mime) = mime {
part = part.mime(mime);
}
Ok(part)
}
fn new(data: PartData) -> Self {
let length = data.len();
Self {
data,
length,
file_name: None,
mime: None,
}
}
fn len(&self) -> Option<u64> {
self.length
}
pub fn file_name<T>(mut self, file_name: T) -> Self
where
T: Into<Cow<'static, str>>,
{
self.file_name = Some(file_name.into());
self
}
pub fn mime_str(self, mime: &str) -> Result<Self, mime::FromStrError> {
Ok(self.mime(mime.parse()?))
}
pub fn mime(mut self, mime: Mime) -> Self {
self.mime = Some(mime);
self
}
fn encode_header(&self, boundary: &str, name: &str) -> Bytes {
let cap = 96
+ boundary.len()
+ name.len()
+ self.file_name.as_ref().map_or(0, |f| f.len() + 16)
+ self.mime.as_ref().map_or(0, |m| m.as_ref().len() + 16);
let mut buf = BytesMut::with_capacity(cap);
buf.extend_from_slice(b"--");
buf.extend_from_slice(boundary.as_bytes());
buf.extend_from_slice(b"\r\nContent-Disposition: form-data; name=\"");
extend_escaped(&mut buf, name);
buf.extend_from_slice(b"\"");
if let Some(file_name) = &self.file_name {
buf.extend_from_slice(b"; filename=\"");
extend_escaped(&mut buf, file_name);
buf.extend_from_slice(b"\"");
}
buf.extend_from_slice(b"\r\n");
if let Some(mime) = &self.mime {
buf.extend_from_slice(b"Content-Type: ");
buf.extend_from_slice(mime.as_ref().as_bytes());
buf.extend_from_slice(b"\r\n");
}
buf.extend_from_slice(b"\r\n");
buf.freeze()
}
}
fn once_frame(bytes: Bytes) -> FrameStream {
Box::pin(futures_util::stream::once(
async move { Ok(Frame::data(bytes)) },
))
}
fn extend_escaped(buf: &mut BytesMut, value: &str) {
let bytes = value.as_bytes();
let mut start = 0;
for (i, &byte) in bytes.iter().enumerate() {
let replacement: &[u8] = match byte {
b'\\' => b"\\\\",
b'"' => b"\\\"",
b'\r' | b'\n' => b" ",
_ => continue,
};
buf.extend_from_slice(&bytes[start..i]);
buf.extend_from_slice(replacement);
start = i + 1;
}
buf.extend_from_slice(&bytes[start..]);
}
fn gen_boundary() -> String {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let rand = rand::random::<u64>();
let seq = COUNTER.fetch_add(1, Ordering::Relaxed);
format!("volo-http-boundary-{rand:016x}{seq:016x}")
}
#[cfg(test)]
mod tests {
use http_body_util::BodyExt;
use tempfile::NamedTempFile;
use super::*;
async fn body_to_string(body: Body) -> String {
let bytes = body.collect().await.unwrap().to_bytes();
String::from_utf8(bytes.to_vec()).unwrap()
}
#[tokio::test]
async fn encode_text_fields() {
let form = Form::new().text("key1", "val1").text("key2", "val2");
let boundary = form.boundary().to_owned();
let content_type = form.content_type();
let body = body_to_string(form.into_body()).await;
assert_eq!(
content_type.to_str().unwrap(),
format!("multipart/form-data; boundary={boundary}")
);
let expected = format!(
"--{boundary}\r\nContent-Disposition: form-data; \
name=\"key1\"\r\n\r\nval1\r\n--{boundary}\r\nContent-Disposition: form-data; \
name=\"key2\"\r\n\r\nval2\r\n--{boundary}--\r\n"
);
assert_eq!(body, expected);
}
#[tokio::test]
async fn encode_reader_part_with_metadata() {
let form = Form::new().part(
"file",
Part::reader(std::io::Cursor::new(b"file-content".to_vec()))
.file_name("a.txt")
.mime_str("text/plain")
.unwrap(),
);
let boundary = form.boundary().to_owned();
let body = body_to_string(form.into_body()).await;
let expected = format!(
"--{boundary}\r\nContent-Disposition: form-data; name=\"file\"; \
filename=\"a.txt\"\r\nContent-Type: \
text/plain\r\n\r\nfile-content\r\n--{boundary}--\r\n"
);
assert_eq!(body, expected);
}
#[tokio::test]
async fn in_memory_form_has_known_length() {
use http_body::Body as _;
let form = Form::new()
.text("key", "value")
.part("bytes", Part::bytes(&b"raw-bytes"[..]).file_name("a.bin"));
let body = form.into_body();
let hint = body.size_hint();
let exact = hint
.exact()
.expect("in-memory form should have a known length");
let encoded = body.collect().await.unwrap().to_bytes();
assert_eq!(exact, encoded.len() as u64);
}
#[tokio::test]
async fn streaming_form_has_unknown_length() {
use http_body::Body as _;
let form = Form::new().text("key", "value").part(
"file",
Part::reader(std::io::Cursor::new(b"streamed".to_vec())),
);
let body = form.into_body();
assert!(body.size_hint().exact().is_none());
}
#[tokio::test]
async fn file_backed_form_has_known_length() {
use http_body::Body as _;
let temp = NamedTempFile::new().unwrap();
tokio::fs::write(temp.path(), b"file-content-from-disk")
.await
.unwrap();
let form = Form::new()
.text("key", "value")
.part("file", Part::file(temp.path()).await.unwrap());
let body = form.into_body();
let exact = body
.size_hint()
.exact()
.expect("file-backed form should keep a known length");
let encoded = body.collect().await.unwrap().to_bytes();
assert_eq!(exact, encoded.len() as u64);
}
#[test]
fn boundaries_are_unique_and_valid() {
let a = gen_boundary();
let b = gen_boundary();
assert_ne!(a, b);
assert!(a.len() <= 70);
assert!(
a.bytes().all(|c| c.is_ascii_alphanumeric() || c == b'-'),
"boundary contains an invalid character: {a}"
);
}
#[test]
fn escape_special_chars() {
let mut buf = BytesMut::new();
extend_escaped(&mut buf, "a\"b\\c\r\nd");
assert_eq!(&buf[..], b"a\\\"b\\\\c d");
let mut buf = BytesMut::new();
extend_escaped(&mut buf, "plain_name.txt");
assert_eq!(&buf[..], b"plain_name.txt");
let mut buf = BytesMut::new();
extend_escaped(&mut buf, "");
assert_eq!(&buf[..], b"");
let mut buf = BytesMut::new();
extend_escaped(&mut buf, "\"ab\"");
assert_eq!(&buf[..], b"\\\"ab\\\"");
}
#[tokio::test]
async fn quoted_name_roundtrips_through_multer() {
let form = Form::new().part("my\"field", Part::text("value").file_name("a\"b.txt"));
let boundary = form.boundary().to_owned();
let bytes = form.into_body().collect().await.unwrap().to_bytes();
let stream =
futures_util::stream::once(async move { Ok::<_, std::convert::Infallible>(bytes) });
let mut multipart = multer::Multipart::new(stream, boundary);
let field = multipart.next_field().await.unwrap().unwrap();
assert_eq!(field.name().unwrap(), "my\"field");
assert_eq!(field.file_name().unwrap(), "a\"b.txt");
assert_eq!(field.bytes().await.unwrap(), &b"value"[..]);
}
}