use std::borrow::Cow;
use std::fmt;
use axum::extract::multipart::{MultipartError, MultipartRejection};
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
#[derive(Debug)]
pub enum TypedMultipartError {
InvalidRequest {
source: MultipartRejection,
},
InvalidRequestBody {
source: MultipartError,
},
MissingField {
field_name: String,
},
WrongFieldType {
field_name: String,
wanted: Cow<'static, str>,
source: String,
},
DuplicateField {
field_name: String,
},
UnknownField {
field_name: String,
},
InvalidEnumValue {
field_name: String,
value: String,
},
NamelessField,
FieldTooLarge {
field_name: String,
limit_bytes: usize,
},
RequestTooLarge {
field_name: String,
limit_bytes: usize,
},
TooManyFields {
limit_fields: usize,
},
Other {
source: String,
},
}
pub(super) const MAX_REFLECTED_VALUE_CHARS: usize = 128;
pub(super) fn truncate_reflected_value(value: &str) -> std::borrow::Cow<'_, str> {
match value.char_indices().nth(MAX_REFLECTED_VALUE_CHARS) {
None => std::borrow::Cow::Borrowed(value),
Some((byte_idx, _)) => {
std::borrow::Cow::Owned(format!("{}... (truncated)", &value[..byte_idx]))
}
}
}
impl fmt::Display for TypedMultipartError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InvalidRequest { source } => {
write!(f, "Invalid multipart request: {source}")
}
Self::InvalidRequestBody { source } => {
write!(f, "Invalid multipart body: {source}")
}
Self::MissingField { field_name } => {
write!(f, "Missing field: `{field_name}`")
}
Self::WrongFieldType {
field_name,
wanted,
source,
} => {
write!(
f,
"Wrong type for field `{field_name}` (expected {wanted}): {source}"
)
}
Self::DuplicateField { field_name } => {
write!(f, "Duplicate field: `{field_name}`")
}
Self::UnknownField { field_name } => {
write!(f, "Unknown field: `{field_name}`")
}
Self::InvalidEnumValue { field_name, value } => {
write!(
f,
"Invalid enum value `{}` for field `{field_name}`",
truncate_reflected_value(value)
)
}
Self::NamelessField => write!(f, "Encountered a field without a name"),
Self::FieldTooLarge {
field_name,
limit_bytes,
} => {
write!(
f,
"Field `{field_name}` exceeds size limit of {limit_bytes} bytes"
)
}
Self::RequestTooLarge {
field_name,
limit_bytes,
} => write!(
f,
"Multipart request exceeds aggregate size limit of {limit_bytes} bytes while reading field `{field_name}`"
),
Self::TooManyFields { limit_fields } => {
write!(
f,
"Multipart request exceeds field count limit of {limit_fields}"
)
}
Self::Other { source } => write!(f, "{source}"),
}
}
}
impl std::error::Error for TypedMultipartError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::InvalidRequest { source } => Some(source),
Self::InvalidRequestBody { source } => Some(source),
Self::MissingField { .. }
| Self::WrongFieldType { .. }
| Self::DuplicateField { .. }
| Self::UnknownField { .. }
| Self::InvalidEnumValue { .. }
| Self::NamelessField
| Self::FieldTooLarge { .. }
| Self::RequestTooLarge { .. }
| Self::TooManyFields { .. }
| Self::Other { .. } => None,
}
}
}
impl TypedMultipartError {
#[must_use]
pub fn invalid_enum_value(field_name: String, value: &str) -> Self {
Self::InvalidEnumValue {
field_name,
value: truncate_reflected_value(value).into_owned(),
}
}
fn field_name(&self) -> Option<&str> {
match self {
Self::MissingField { field_name }
| Self::WrongFieldType { field_name, .. }
| Self::DuplicateField { field_name }
| Self::UnknownField { field_name }
| Self::InvalidEnumValue { field_name, .. }
| Self::FieldTooLarge { field_name, .. }
| Self::RequestTooLarge { field_name, .. } => Some(field_name),
Self::InvalidRequest { .. }
| Self::InvalidRequestBody { .. }
| Self::NamelessField
| Self::TooManyFields { .. }
| Self::Other { .. } => None,
}
}
pub(super) fn error_body(&self) -> Vec<u8> {
serde_json::to_vec(&MultipartErrorEnvelope {
errors: [MultipartOneError {
message: MultipartMessage(self),
path: self.field_name().unwrap_or(""),
}],
})
.unwrap_or_else(|_| br#"{"errors":[{"message":"serialization error","path":""}]}"#.to_vec())
}
}
const MULTIPART_INTERNAL_ERROR_MSG: &str = "internal error while processing multipart request";
struct MultipartMessage<'a>(&'a TypedMultipartError);
impl serde::Serialize for MultipartMessage<'_> {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
if matches!(self.0, TypedMultipartError::Other { .. }) {
serializer.serialize_str(MULTIPART_INTERNAL_ERROR_MSG)
} else {
serializer.collect_str(self.0)
}
}
}
#[derive(serde::Serialize)]
struct MultipartOneError<'a> {
message: MultipartMessage<'a>,
path: &'a str,
}
#[derive(serde::Serialize)]
struct MultipartErrorEnvelope<'a> {
errors: [MultipartOneError<'a>; 1],
}
impl IntoResponse for TypedMultipartError {
fn into_response(self) -> Response {
let status = match &self {
Self::InvalidRequest { source } => source.status(),
Self::InvalidRequestBody { source } => source.status(),
Self::MissingField { .. }
| Self::DuplicateField { .. }
| Self::UnknownField { .. }
| Self::InvalidEnumValue { .. }
| Self::NamelessField => StatusCode::BAD_REQUEST,
Self::WrongFieldType { .. } => StatusCode::UNPROCESSABLE_ENTITY,
Self::FieldTooLarge { .. }
| Self::RequestTooLarge { .. }
| Self::TooManyFields { .. } => StatusCode::PAYLOAD_TOO_LARGE,
Self::Other { .. } => StatusCode::INTERNAL_SERVER_ERROR,
};
let body = self.error_body();
(
status,
[(
axum::http::header::CONTENT_TYPE,
axum::http::HeaderValue::from_static("application/json"),
)],
body,
)
.into_response()
}
}
impl From<MultipartError> for TypedMultipartError {
fn from(source: MultipartError) -> Self {
Self::InvalidRequestBody { source }
}
}
impl From<MultipartRejection> for TypedMultipartError {
fn from(source: MultipartRejection) -> Self {
Self::InvalidRequest { source }
}
}