use crate::{body::Body, response::IntoResponse}; use async_trait::async_trait; use bytes::Bytes; use http::{header, Request, Response}; use rejection::{ BodyAlreadyTaken, FailedToBufferBody, InvalidJsonBody, InvalidUrlParam, InvalidUtf8, LengthRequired, MissingExtension, MissingJsonContentType, MissingRouteParams, PayloadTooLarge, QueryStringMissing, }; use serde::de::DeserializeOwned; use std::{collections::HashMap, convert::Infallible, str::FromStr}; pub mod rejection; #[async_trait] pub trait FromRequest: Sized { type Rejection: IntoResponse; async fn from_request(req: &mut Request) -> Result; } #[async_trait] impl FromRequest for Option where T: FromRequest, { type Rejection = Infallible; async fn from_request(req: &mut Request) -> Result, Self::Rejection> { Ok(T::from_request(req).await.ok()) } } #[derive(Debug, Clone, Copy, Default)] pub struct Query(pub T); #[async_trait] impl FromRequest for Query where T: DeserializeOwned, { type Rejection = QueryStringMissing; async fn from_request(req: &mut Request) -> Result { let query = req.uri().query().ok_or(QueryStringMissing(()))?; let value = serde_urlencoded::from_str(query).map_err(|_| QueryStringMissing(()))?; Ok(Query(value)) } } #[derive(Debug, Clone, Copy, Default)] pub struct Json(pub T); #[async_trait] impl FromRequest for Json where T: DeserializeOwned, { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { if has_content_type(req, "application/json") { let body = take_body(req).map_err(IntoResponse::into_response)?; let bytes = hyper::body::to_bytes(body) .await .map_err(InvalidJsonBody::from_err) .map_err(IntoResponse::into_response)?; let value = serde_json::from_slice(&bytes) .map_err(InvalidJsonBody::from_err) .map_err(IntoResponse::into_response)?; Ok(Json(value)) } else { Err(MissingJsonContentType(()).into_response()) } } } fn has_content_type(req: &Request, expected_content_type: &str) -> bool { let content_type = if let Some(content_type) = req.headers().get(header::CONTENT_TYPE) { content_type } else { return false; }; let content_type = if let Ok(content_type) = content_type.to_str() { content_type } else { return false; }; content_type.starts_with(expected_content_type) } #[derive(Debug, Clone, Copy)] pub struct Extension(pub T); #[async_trait] impl FromRequest for Extension where T: Clone + Send + Sync + 'static, { type Rejection = MissingExtension; async fn from_request(req: &mut Request) -> Result { let value = req .extensions() .get::() .ok_or(MissingExtension(())) .map(|x| x.clone())?; Ok(Extension(value)) } } #[async_trait] impl FromRequest for Bytes { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { let body = take_body(req).map_err(IntoResponse::into_response)?; let bytes = hyper::body::to_bytes(body) .await .map_err(FailedToBufferBody::from_err) .map_err(IntoResponse::into_response)?; Ok(bytes) } } #[async_trait] impl FromRequest for String { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { let body = take_body(req).map_err(IntoResponse::into_response)?; let bytes = hyper::body::to_bytes(body) .await .map_err(FailedToBufferBody::from_err) .map_err(IntoResponse::into_response)? .to_vec(); let string = String::from_utf8(bytes) .map_err(InvalidUtf8::from_err) .map_err(IntoResponse::into_response)?; Ok(string) } } #[async_trait] impl FromRequest for Body { type Rejection = BodyAlreadyTaken; async fn from_request(req: &mut Request) -> Result { take_body(req) } } #[derive(Debug, Clone)] pub struct BytesMaxLength(pub Bytes); #[async_trait] impl FromRequest for BytesMaxLength { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { let content_length = req.headers().get(http::header::CONTENT_LENGTH).cloned(); let body = take_body(req).map_err(|reject| reject.into_response())?; let content_length = content_length.and_then(|value| value.to_str().ok()?.parse::().ok()); if let Some(length) = content_length { if length > N { return Err(PayloadTooLarge(()).into_response()); } } else { return Err(LengthRequired(()).into_response()); }; let bytes = hyper::body::to_bytes(body) .await .map_err(|e| FailedToBufferBody::from_err(e).into_response())?; Ok(BytesMaxLength(bytes)) } } #[derive(Debug)] pub struct UrlParamsMap(HashMap); impl UrlParamsMap { pub fn get(&self, key: &str) -> Option<&str> { self.0.get(key).map(|s| &**s) } pub fn get_typed(&self, key: &str) -> Option where T: FromStr, { self.get(key)?.parse().ok() } } #[async_trait] impl FromRequest for UrlParamsMap { type Rejection = MissingRouteParams; async fn from_request(req: &mut Request) -> Result { if let Some(params) = req .extensions_mut() .get_mut::>() { let params = params.take().expect("params already taken").0; Ok(Self(params.into_iter().collect())) } else { Err(MissingRouteParams(())) } } } pub struct UrlParams(pub T); macro_rules! impl_parse_url { () => {}; ( $head:ident, $($tail:ident),* $(,)? ) => { #[async_trait] impl<$head, $($tail,)*> FromRequest for UrlParams<($head, $($tail,)*)> where $head: FromStr + Send, $( $tail: FromStr + Send, )* { type Rejection = Response; #[allow(non_snake_case)] async fn from_request(req: &mut Request) -> Result { let params = if let Some(params) = req .extensions_mut() .get_mut::>() { params.take().expect("params already taken").0 } else { return Err(MissingRouteParams(()).into_response()) }; if let [(_, $head), $((_, $tail),)*] = &*params { let $head = if let Ok(x) = $head.parse::<$head>() { x } else { return Err(InvalidUrlParam::new::<$head>().into_response()); }; $( let $tail = if let Ok(x) = $tail.parse::<$tail>() { x } else { return Err(InvalidUrlParam::new::<$tail>().into_response()); }; )* Ok(UrlParams(($head, $($tail,)*))) } else { return Err(MissingRouteParams(()).into_response()) } } } impl_parse_url!($($tail,)*); }; } impl_parse_url!(T1, T2, T3, T4, T5, T6); fn take_body(req: &mut Request) -> Result { struct BodyAlreadyTakenExt; if req.extensions_mut().insert(BodyAlreadyTakenExt).is_some() { Err(BodyAlreadyTaken(())) } else { let body = std::mem::take(req.body_mut()); Ok(body) } }