use async_trait::async_trait;
use axol_http::body::BodyComponent;
use axol_http::header::HeaderMap;
use axol_http::mime::{BOUNDARY, Mime};
use axol_http::request::RequestPartsRef;
use axol_http::typed_headers::ContentType;
use axol_http::{Body, StatusCode};
use bytes::Bytes;
use futures::StreamExt;
use futures_util::stream::Stream;
use std::{
fmt,
pin::Pin,
task::{Context, Poll},
};
use crate::{Error, FromRequest, IntoResponse, Result};
#[cfg_attr(docsrs, doc(cfg(feature = "multipart")))]
#[derive(Debug)]
pub struct Multipart {
inner: multer::Multipart<'static>,
}
#[async_trait]
impl<'a> FromRequest<'a> for Multipart {
async fn from_request(request: RequestPartsRef<'a>, body: Body) -> Result<Self> {
let boundary = parse_boundary(request.headers).ok_or_else(|| {
Error::bad_request("Invalid `boundary` for `multipart/form-data` request")
})?;
let stream = body.into_stream().filter_map(|x| async move {
Some(match x {
Err(e) => Err(e),
Ok(BodyComponent::Trailers(_)) => return None,
Ok(BodyComponent::Data(data)) => Ok(data),
})
});
let multipart = multer::Multipart::new(stream, boundary);
Ok(Self { inner: multipart })
}
}
impl Multipart {
pub async fn next_field(&mut self) -> Result<Option<Field<'_>>> {
let field = self
.inner
.next_field()
.await
.map_err(MultipartError::from_multer)
.map_err(MultipartError::into_error)?;
if let Some(field) = field {
Ok(Some(Field {
inner: field,
_multipart: self,
}))
} else {
Ok(None)
}
}
}
#[derive(Debug)]
pub struct Field<'a> {
inner: multer::Field<'static>,
_multipart: &'a mut Multipart,
}
impl<'a> Stream for Field<'a> {
type Item = Result<Bytes>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Pin::new(&mut self.inner)
.poll_next(cx)
.map_err(MultipartError::from_multer)
.map_err(MultipartError::into_error)
}
}
impl<'a> Field<'a> {
pub fn name(&self) -> Option<&str> {
self.inner.name()
}
pub fn file_name(&self) -> Option<&str> {
self.inner.file_name()
}
pub fn content_type(&self) -> Option<&str> {
self.inner.content_type().map(|m| m.as_ref())
}
pub fn headers(&self) -> Result<HeaderMap> {
self.inner
.headers()
.clone()
.try_into()
.map_err(|_| Error::BadUtf8)
}
pub async fn bytes(self) -> Result<Bytes> {
self.inner
.bytes()
.await
.map_err(MultipartError::from_multer)
.map_err(MultipartError::into_error)
}
pub async fn text(self) -> Result<String> {
self.inner
.text()
.await
.map_err(MultipartError::from_multer)
.map_err(MultipartError::into_error)
}
pub async fn chunk(&mut self) -> Result<Option<Bytes>> {
self.inner
.chunk()
.await
.map_err(MultipartError::from_multer)
.map_err(MultipartError::into_error)
}
}
#[derive(Debug)]
struct MultipartError {
source: multer::Error,
}
impl MultipartError {
fn from_multer(multer: multer::Error) -> Self {
Self { source: multer }
}
pub fn body_text(&self) -> String {
self.source.to_string()
}
pub fn status(&self) -> StatusCode {
status_code_from_multer_error(&self.source)
}
}
fn status_code_from_multer_error(err: &multer::Error) -> StatusCode {
match err {
multer::Error::UnknownField { .. }
| multer::Error::IncompleteFieldData { .. }
| multer::Error::IncompleteHeaders
| multer::Error::ReadHeaderFailed(..)
| multer::Error::DecodeHeaderName { .. }
| multer::Error::DecodeContentType(..)
| multer::Error::NoBoundary
| multer::Error::DecodeHeaderValue { .. }
| multer::Error::NoMultipart
| multer::Error::IncompleteStream => StatusCode::BadRequest,
multer::Error::FieldSizeExceeded { .. } | multer::Error::StreamSizeExceeded { .. } => {
StatusCode::PayloadTooLarge
}
multer::Error::StreamReadFailed(err) => {
if let Some(err) = err.downcast_ref::<multer::Error>() {
return status_code_from_multer_error(err);
}
StatusCode::InternalServerError
}
_ => StatusCode::InternalServerError,
}
}
impl fmt::Display for MultipartError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "Error parsing `multipart/form-data` request")
}
}
impl std::error::Error for MultipartError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(&self.source)
}
}
impl MultipartError {
fn into_error(self) -> Error {
Error::Response((self.status(), self.body_text()).into_response().unwrap())
}
}
fn parse_boundary(headers: &HeaderMap) -> Option<String> {
let mime: Mime = headers.get_typed::<ContentType>()?.into();
Some(mime.get_param(BOUNDARY)?.as_str().to_string())
}