use std::sync::Arc;
use axum::body::Body;
use axum::extract::Multipart;
use bytes::Bytes;
use futures::StreamExt;
use tokio::sync::mpsc;
pub const DEFAULT_MULTIPART_OUTER_CAPACITY: usize = 2;
pub const DEFAULT_FIELD_BYTES_CAPACITY: usize = 16;
pub const DEFAULT_BODY_CHANNEL_CAPACITY: usize = 16;
pub const DEFAULT_MAX_FIELD_BYTES: u64 = 256 * 1024 * 1024;
#[derive(Clone, Debug)]
pub struct MultipartStreamConfig {
pub outer_capacity: usize,
pub field_bytes_capacity: usize,
pub max_field_bytes: u64,
}
impl Default for MultipartStreamConfig {
fn default() -> Self {
Self {
outer_capacity: DEFAULT_MULTIPART_OUTER_CAPACITY,
field_bytes_capacity: DEFAULT_FIELD_BYTES_CAPACITY,
max_field_bytes: DEFAULT_MAX_FIELD_BYTES,
}
}
}
#[derive(Clone, Debug)]
pub struct BodyChannelConfig {
pub capacity: usize,
}
impl Default for BodyChannelConfig {
fn default() -> Self {
Self {
capacity: DEFAULT_BODY_CHANNEL_CAPACITY,
}
}
}
#[derive(Debug)]
pub struct MultipartField {
pub name: String,
pub filename: Option<String>,
pub content_type: Option<String>,
pub bytes: mpsc::Receiver<Result<Bytes, StreamError>>,
}
pub struct MultipartStream {
receiver: mpsc::Receiver<Result<MultipartField, StreamError>>,
}
impl MultipartStream {
pub fn start(multipart: Multipart, config: MultipartStreamConfig) -> Self {
let (outer_tx, outer_rx) = mpsc::channel(config.outer_capacity.max(1));
tokio::spawn(drive_multipart(multipart, outer_tx, config));
Self { receiver: outer_rx }
}
pub async fn next_field(&mut self) -> Result<Option<MultipartField>, StreamError> {
match self.receiver.recv().await {
Some(Ok(field)) => Ok(Some(field)),
Some(Err(error)) => Err(error),
None => Ok(None),
}
}
}
pub struct RequestBodyChannel {
receiver: mpsc::Receiver<Result<Bytes, StreamError>>,
}
impl RequestBodyChannel {
pub fn start(body: Body, config: BodyChannelConfig) -> Self {
let (tx, rx) = mpsc::channel(config.capacity.max(1));
tokio::spawn(drive_body(body, tx));
Self { receiver: rx }
}
pub async fn recv(&mut self) -> Result<Option<Bytes>, StreamError> {
match self.receiver.recv().await {
Some(Ok(bytes)) => Ok(Some(bytes)),
Some(Err(error)) => Err(error),
None => Ok(None),
}
}
}
#[derive(Debug, Clone)]
pub enum StreamError {
Multipart(String),
Body(String),
MissingFieldName,
FieldTooLarge {
field: Option<String>,
limit: u64,
},
}
impl std::fmt::Display for StreamError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Multipart(message) => write!(f, "multipart parse error: {message}"),
Self::Body(message) => write!(f, "body stream error: {message}"),
Self::MissingFieldName => write!(
f,
"multipart part is missing the required name= Content-Disposition param"
),
Self::FieldTooLarge { field, limit } => match field {
Some(name) => write!(
f,
"multipart field `{name}` exceeded max_field_bytes ({limit})"
),
None => write!(f, "multipart field exceeded max_field_bytes ({limit})"),
},
}
}
}
impl std::error::Error for StreamError {}
async fn drive_multipart(
mut multipart: Multipart,
outer_tx: mpsc::Sender<Result<MultipartField, StreamError>>,
config: MultipartStreamConfig,
) {
let inner_capacity = config.field_bytes_capacity.max(1);
let max_field_bytes = config.max_field_bytes;
loop {
let next = match multipart.next_field().await {
Ok(Some(field)) => field,
Ok(None) => return,
Err(error) => {
let _ = outer_tx
.send(Err(StreamError::Multipart(error.to_string())))
.await;
return;
}
};
let name = match next.name() {
Some(value) => value.to_string(),
None => {
if outer_tx
.send(Err(StreamError::MissingFieldName))
.await
.is_err()
{
return;
}
continue;
}
};
let filename = next.file_name().map(str::to_string);
let content_type = next.content_type().map(str::to_string);
let (inner_tx, inner_rx) = mpsc::channel::<Result<Bytes, StreamError>>(inner_capacity);
let field = MultipartField {
name: name.clone(),
filename,
content_type,
bytes: inner_rx,
};
if outer_tx.send(Ok(field)).await.is_err() {
return;
}
let inner_tx = Arc::new(inner_tx);
let outcome = pump_field_bytes(next, inner_tx.clone(), max_field_bytes, &name).await;
drop(inner_tx);
match outcome {
FieldOutcome::Complete => continue,
FieldOutcome::ConsumerDropped => return,
FieldOutcome::ParserFailed(message) => {
let _ = outer_tx.send(Err(StreamError::Multipart(message))).await;
return;
}
}
}
}
enum FieldOutcome {
Complete,
ConsumerDropped,
ParserFailed(String),
}
async fn pump_field_bytes(
mut field: axum::extract::multipart::Field<'_>,
inner_tx: Arc<mpsc::Sender<Result<Bytes, StreamError>>>,
max_field_bytes: u64,
field_name: &str,
) -> FieldOutcome {
let mut bytes_so_far: u64 = 0;
loop {
match field.chunk().await {
Ok(Some(chunk)) => {
let len = chunk.len() as u64;
if bytes_so_far.saturating_add(len) > max_field_bytes {
let _ = inner_tx
.send(Err(StreamError::FieldTooLarge {
field: Some(field_name.to_string()),
limit: max_field_bytes,
}))
.await;
while let Ok(Some(_)) = field.chunk().await {}
return FieldOutcome::Complete;
}
bytes_so_far += len;
if inner_tx.send(Ok(chunk)).await.is_err() {
return FieldOutcome::ConsumerDropped;
}
}
Ok(None) => return FieldOutcome::Complete,
Err(error) => return FieldOutcome::ParserFailed(error.to_string()),
}
}
}
async fn drive_body(body: Body, tx: mpsc::Sender<Result<Bytes, StreamError>>) {
let mut stream = body.into_data_stream();
while let Some(next) = stream.next().await {
let send_result = match next {
Ok(chunk) => tx.send(Ok(chunk)).await,
Err(error) => tx.send(Err(StreamError::Body(error.to_string()))).await,
};
if send_result.is_err() {
return;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::extract::Multipart;
use axum::http::{header, Method, Request, StatusCode};
use axum::response::{IntoResponse, Response};
use axum::routing::post;
use axum::Router;
use tower::ServiceExt;
async fn echo_multipart_handler(multipart: Multipart) -> Response {
let mut stream = MultipartStream::start(multipart, MultipartStreamConfig::default());
let mut summary = String::new();
while let Some(mut field) = match stream.next_field().await {
Ok(field) => field,
Err(error) => return (StatusCode::BAD_REQUEST, error.to_string()).into_response(),
} {
let mut total = 0u64;
while let Some(chunk_result) = field.bytes.recv().await {
match chunk_result {
Ok(chunk) => total += chunk.len() as u64,
Err(error) => {
return (StatusCode::PAYLOAD_TOO_LARGE, error.to_string()).into_response();
}
}
}
summary.push_str(&format!("{}:{}\n", field.name, total));
}
summary.into_response()
}
async fn echo_body_handler(req: Request<Body>) -> Response {
let (_parts, body) = req.into_parts();
let mut channel = RequestBodyChannel::start(body, BodyChannelConfig::default());
let mut total: u64 = 0;
while let Some(chunk_result) = match channel.recv().await {
Ok(value) => value.map(Ok),
Err(error) => Some(Err(error)),
} {
match chunk_result {
Ok(chunk) => total += chunk.len() as u64,
Err(error) => {
return (StatusCode::BAD_REQUEST, error.to_string()).into_response();
}
}
}
total.to_string().into_response()
}
fn build_app() -> Router {
Router::new()
.route("/multipart", post(echo_multipart_handler))
.route("/body", post(echo_body_handler))
}
fn boundary() -> &'static str {
"----streaming-unit-test"
}
fn multipart_body(fields: &[(&str, &[u8])]) -> Vec<u8> {
let mut out = Vec::new();
for (name, value) in fields {
out.extend_from_slice(format!("--{}\r\n", boundary()).as_bytes());
out.extend_from_slice(
format!("Content-Disposition: form-data; name=\"{name}\"\r\n\r\n").as_bytes(),
);
out.extend_from_slice(value);
out.extend_from_slice(b"\r\n");
}
out.extend_from_slice(format!("--{}--\r\n", boundary()).as_bytes());
out
}
async fn read_text(response: Response) -> String {
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
String::from_utf8(bytes.to_vec()).unwrap()
}
#[tokio::test]
async fn multipart_stream_yields_one_field_at_a_time() {
let body = multipart_body(&[("a", b"hello"), ("b", b"world!!")]);
let request = Request::builder()
.method(Method::POST)
.uri("/multipart")
.header(
header::CONTENT_TYPE,
format!("multipart/form-data; boundary={}", boundary()),
)
.body(Body::from(body))
.unwrap();
let response = build_app().oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let summary = read_text(response).await;
assert_eq!(summary, "a:5\nb:7\n");
}
#[tokio::test]
async fn multipart_stream_enforces_field_size_cap() {
async fn capped_handler(multipart: Multipart) -> Response {
let mut stream = MultipartStream::start(
multipart,
MultipartStreamConfig {
max_field_bytes: 4,
..Default::default()
},
);
let mut errors = Vec::new();
while let Some(mut field) = stream.next_field().await.unwrap() {
while let Some(chunk_result) = field.bytes.recv().await {
if let Err(error) = chunk_result {
errors.push(error.to_string());
}
}
}
errors.join("|").into_response()
}
let body = multipart_body(&[("big", b"this body is more than four bytes")]);
let app = Router::new().route("/c", post(capped_handler));
let request = Request::builder()
.method(Method::POST)
.uri("/c")
.header(
header::CONTENT_TYPE,
format!("multipart/form-data; boundary={}", boundary()),
)
.body(Body::from(body))
.unwrap();
let response = app.oneshot(request).await.unwrap();
let summary = read_text(response).await;
assert!(
summary.contains("max_field_bytes (4)"),
"expected size-cap error, got `{summary}`"
);
assert!(
summary.contains("`big`"),
"expected field name in error, got `{summary}`"
);
}
#[tokio::test]
async fn body_channel_streams_chunked_upload() {
let payload = vec![0xABu8; 1024];
let request = Request::builder()
.method(Method::POST)
.uri("/body")
.body(Body::from(payload.clone()))
.unwrap();
let response = build_app().oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let summary = read_text(response).await;
assert_eq!(summary, "1024");
}
}