use std::cell::RefCell;
use std::sync::atomic::{AtomicUsize, Ordering};
use axum::extract::multipart::Field;
use axum::extract::{FromRequest, Request};
mod error;
pub use error::TypedMultipartError;
use error::truncate_reflected_value;
pub trait TryFromMultipartWithState<S: Send + Sync>: Sized {
fn try_from_multipart_with_state(
multipart: &mut axum::extract::Multipart,
state: &S,
) -> impl std::future::Future<Output = Result<Self, TypedMultipartError>> + Send;
}
pub struct MeteredField<'a> {
inner: Field<'a>,
}
impl<'a> MeteredField<'a> {
#[doc(hidden)]
#[must_use]
pub fn __from_field(inner: Field<'a>) -> Self {
Self { inner }
}
#[must_use]
pub fn name(&self) -> Option<&str> {
self.inner.name()
}
#[must_use]
pub fn file_name(&self) -> Option<&str> {
self.inner.file_name()
}
#[must_use]
pub fn content_type(&self) -> Option<&str> {
self.inner.content_type()
}
pub async fn chunk(&mut self) -> Result<Option<axum::body::Bytes>, TypedMultipartError> {
let next = self.inner.chunk().await?;
if let Some(chunk) = &next {
register_multipart_bytes(self.inner.name().unwrap_or_default(), chunk.len())?;
}
Ok(next)
}
pub async fn bytes(mut self) -> Result<axum::body::Bytes, TypedMultipartError> {
self.bytes_with_limit_inner(None, 0)
.await
.map(axum::body::Bytes::from)
}
pub async fn bytes_with_limit(
mut self,
limit_bytes: usize,
initial_capacity: usize,
) -> Result<axum::body::Bytes, TypedMultipartError> {
self.bytes_with_limit_inner(Some(limit_bytes), initial_capacity)
.await
.map(axum::body::Bytes::from)
}
async fn bytes_with_limit_inner(
&mut self,
limit: Option<usize>,
initial_capacity: usize,
) -> Result<Vec<u8>, TypedMultipartError> {
let capacity = limit.map_or(initial_capacity, |limit| initial_capacity.min(limit));
let mut acc: Vec<u8> = Vec::with_capacity(capacity);
while let Some(chunk) = self.chunk().await? {
if let Some(limit) = limit
&& acc.len().saturating_add(chunk.len()) > limit
{
return Err(TypedMultipartError::FieldTooLarge {
field_name: self.name().unwrap_or_default().to_string(),
limit_bytes: limit,
});
}
acc.extend_from_slice(&chunk);
}
Ok(acc)
}
}
impl From<&MeteredField<'_>> for FieldMetadata {
fn from(field: &MeteredField<'_>) -> Self {
Self::from(&field.inner)
}
}
pub trait TryFromFieldWithState<S: Send + Sync>: Sized {
fn try_from_field_with_state(
field: MeteredField<'_>,
limit_bytes: Option<usize>,
state: &S,
) -> impl std::future::Future<Output = Result<Self, TypedMultipartError>> + Send;
}
#[derive(Debug, Clone)]
pub struct FieldMetadata {
pub name: Option<String>,
pub file_name: Option<String>,
pub content_type: Option<String>,
pub headers: Option<axum::http::HeaderMap>,
}
impl FieldMetadata {
#[must_use]
pub const fn headers(&self) -> Option<&axum::http::HeaderMap> {
self.headers.as_ref()
}
#[must_use]
pub fn with_headers(mut self, headers: axum::http::HeaderMap) -> Self {
self.headers = Some(headers);
self
}
}
impl From<&Field<'_>> for FieldMetadata {
fn from(field: &Field<'_>) -> Self {
Self {
name: field.name().map(String::from),
file_name: field.file_name().map(String::from),
content_type: field.content_type().map(String::from),
headers: None,
}
}
}
#[derive(Debug)]
pub struct FieldData<T> {
pub metadata: FieldMetadata,
pub contents: T,
}
impl<T, S> TryFromFieldWithState<S> for FieldData<T>
where
T: TryFromFieldWithState<S> + Send,
S: Send + Sync,
{
async fn try_from_field_with_state(
field: MeteredField<'_>,
limit_bytes: Option<usize>,
state: &S,
) -> Result<Self, TypedMultipartError> {
let metadata = FieldMetadata::from(&field);
let contents = T::try_from_field_with_state(field, limit_bytes, state).await?;
Ok(Self { metadata, contents })
}
}
pub struct TypedMultipart<T>(pub T);
impl<T> std::ops::Deref for TypedMultipart<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T> std::ops::DerefMut for TypedMultipart<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.0
}
}
pub const DEFAULT_MULTIPART_MAX_TOTAL_BYTES: usize = 64 * 1024 * 1024;
pub const DEFAULT_MULTIPART_MAX_FIELDS: usize = 1024;
static DEFAULT_MULTIPART_TOTAL_LIMIT: AtomicUsize =
AtomicUsize::new(DEFAULT_MULTIPART_MAX_TOTAL_BYTES);
static DEFAULT_MULTIPART_FIELD_LIMIT: AtomicUsize = AtomicUsize::new(DEFAULT_MULTIPART_MAX_FIELDS);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MultipartLimits {
pub max_total_bytes: usize,
pub max_fields: usize,
}
impl MultipartLimits {
#[must_use]
pub const fn new(max_total_bytes: usize, max_fields: usize) -> Self {
Self {
max_total_bytes,
max_fields,
}
}
}
#[must_use]
pub fn default_multipart_limits() -> MultipartLimits {
MultipartLimits::new(
DEFAULT_MULTIPART_TOTAL_LIMIT.load(Ordering::Relaxed),
DEFAULT_MULTIPART_FIELD_LIMIT.load(Ordering::Relaxed),
)
}
pub fn set_default_multipart_limits(limits: MultipartLimits) -> MultipartLimits {
MultipartLimits::new(
DEFAULT_MULTIPART_TOTAL_LIMIT.swap(limits.max_total_bytes, Ordering::Relaxed),
DEFAULT_MULTIPART_FIELD_LIMIT.swap(limits.max_fields, Ordering::Relaxed),
)
}
#[derive(Debug)]
struct MultipartAggregateState {
limits: MultipartLimits,
total_bytes: usize,
fields: usize,
}
impl MultipartAggregateState {
const fn new(limits: MultipartLimits) -> Self {
Self {
limits,
total_bytes: 0,
fields: 0,
}
}
}
tokio::task_local! {
static MULTIPART_AGGREGATE: RefCell<MultipartAggregateState>;
}
pub fn register_multipart_part() -> Result<(), TypedMultipartError> {
MULTIPART_AGGREGATE
.try_with(|state| {
let mut state = state.borrow_mut();
state.fields = state.fields.saturating_add(1);
if state.fields > state.limits.max_fields {
return Err(TypedMultipartError::TooManyFields {
limit_fields: state.limits.max_fields,
});
}
Ok(())
})
.unwrap_or(Ok(()))
}
pub fn register_multipart_bytes(
field_name: &str,
chunk_len: usize,
) -> Result<(), TypedMultipartError> {
MULTIPART_AGGREGATE
.try_with(|state| {
let mut state = state.borrow_mut();
state.total_bytes = state.total_bytes.saturating_add(chunk_len);
if state.total_bytes > state.limits.max_total_bytes {
return Err(TypedMultipartError::RequestTooLarge {
field_name: field_name.to_owned(),
limit_bytes: state.limits.max_total_bytes,
});
}
Ok(())
})
.unwrap_or(Ok(()))
}
pub struct TypedMultipartWithLimits<
T,
const MAX_TOTAL_BYTES: usize,
const MAX_FIELDS: usize = DEFAULT_MULTIPART_MAX_FIELDS,
>(pub T);
async fn parse_typed_multipart_with_limits<T, S>(
req: Request,
state: &S,
limits: MultipartLimits,
) -> Result<T, TypedMultipartError>
where
T: TryFromMultipartWithState<S>,
S: Send + Sync + 'static,
{
let mut multipart = axum::extract::Multipart::from_request(req, state)
.await
.map_err(TypedMultipartError::from)?;
MULTIPART_AGGREGATE
.scope(
RefCell::new(MultipartAggregateState::new(limits)),
async move { T::try_from_multipart_with_state(&mut multipart, state).await },
)
.await
}
impl<T, S> FromRequest<S> for TypedMultipart<T>
where
T: TryFromMultipartWithState<S>,
S: Send + Sync + 'static,
{
type Rejection = TypedMultipartError;
async fn from_request(req: Request, state: &S) -> Result<Self, Self::Rejection> {
let value =
parse_typed_multipart_with_limits(req, state, default_multipart_limits()).await?;
Ok(Self(value))
}
}
impl<T, S, const MAX_TOTAL_BYTES: usize, const MAX_FIELDS: usize> FromRequest<S>
for TypedMultipartWithLimits<T, MAX_TOTAL_BYTES, MAX_FIELDS>
where
T: TryFromMultipartWithState<S>,
S: Send + Sync + 'static,
{
type Rejection = TypedMultipartError;
async fn from_request(req: Request, state: &S) -> Result<Self, Self::Rejection> {
let value = parse_typed_multipart_with_limits(
req,
state,
MultipartLimits::new(MAX_TOTAL_BYTES, MAX_FIELDS),
)
.await?;
Ok(Self(value))
}
}
struct FieldBytes<'a> {
field: MeteredField<'a>,
data: Vec<u8>,
}
async fn read_field_data(
mut field: MeteredField<'_>,
limit: Option<usize>,
initial_capacity: usize,
) -> Result<FieldBytes<'_>, TypedMultipartError> {
let buf = field
.bytes_with_limit_inner(limit, initial_capacity)
.await?;
Ok(FieldBytes { field, data: buf })
}
const DEFAULT_TINY_SCALAR_LIMIT_BYTES: usize = 256;
const TINY_SCALAR_INITIAL_CAPACITY_BYTES: usize = 16;
const STRING_INITIAL_CAPACITY_BYTES: usize = 64;
fn tiny_scalar_limit(limit_bytes: Option<usize>) -> usize {
limit_bytes.unwrap_or(DEFAULT_TINY_SCALAR_LIMIT_BYTES)
}
fn str_to_bool(s: &str) -> Option<bool> {
const TRUTHY: [&str; 5] = ["true", "yes", "y", "1", "on"];
const FALSY: [&str; 5] = ["false", "no", "n", "0", "off"];
let s = s.trim();
if TRUTHY.iter().any(|t| s.eq_ignore_ascii_case(t)) {
Some(true)
} else if FALSY.iter().any(|f| s.eq_ignore_ascii_case(f)) {
Some(false)
} else {
None
}
}
const DEFAULT_STRING_FIELD_LIMIT_BYTES: usize = 1024 * 1024;
pub const DEFAULT_TEMP_FILE_FIELD_LIMIT_BYTES: usize = 16 * 1024 * 1024;
static DEFAULT_TEMP_FILE_FIELD_LIMIT: AtomicUsize =
AtomicUsize::new(DEFAULT_TEMP_FILE_FIELD_LIMIT_BYTES);
#[must_use]
pub fn default_temp_file_field_limit_bytes() -> usize {
DEFAULT_TEMP_FILE_FIELD_LIMIT.load(Ordering::Relaxed)
}
pub fn set_default_temp_file_field_limit_bytes(limit_bytes: usize) -> usize {
DEFAULT_TEMP_FILE_FIELD_LIMIT.swap(limit_bytes, Ordering::Relaxed)
}
mod scalar_parsers;
impl<S: Send + Sync> TryFromFieldWithState<S> for tempfile::NamedTempFile {
async fn try_from_field_with_state(
mut field: MeteredField<'_>,
limit_bytes: Option<usize>,
_state: &S,
) -> Result<Self, TypedMultipartError> {
let (temp, std_file) = tokio::task::spawn_blocking(|| {
let temp = Self::new()?;
let std_file = temp.reopen()?;
Ok::<_, std::io::Error>((temp, std_file))
})
.await
.map_err(|e| TypedMultipartError::Other {
source: e.to_string(),
})?
.map_err(|e| TypedMultipartError::Other {
source: e.to_string(),
})?;
let mut file = tokio::fs::File::from_std(std_file);
let limit_bytes = limit_bytes.unwrap_or_else(default_temp_file_field_limit_bytes);
let mut total = 0usize;
while let Some(chunk) = field.chunk().await? {
total = total.saturating_add(chunk.len());
if total > limit_bytes {
return Err(TypedMultipartError::FieldTooLarge {
field_name: field.name().unwrap_or_default().to_string(),
limit_bytes,
});
}
tokio::io::AsyncWriteExt::write_all(&mut file, &chunk)
.await
.map_err(|e| TypedMultipartError::Other {
source: e.to_string(),
})?;
}
tokio::io::AsyncWriteExt::flush(&mut file)
.await
.map_err(|e| TypedMultipartError::Other {
source: e.to_string(),
})?;
Ok(temp)
}
}
#[cfg(test)]
mod tests;