From 11843addf6735eeb14e42fdfc7bc42630ac84663 Mon Sep 17 00:00:00 2001 From: David Pedersen Date: Mon, 31 May 2021 12:55:39 +0200 Subject: [PATCH] Just use async_trait for FromRequest --- src/extract.rs | 194 +++++++++++++++++++------------------------------ 1 file changed, 74 insertions(+), 120 deletions(-) diff --git a/src/extract.rs b/src/extract.rs index cf92d9fc..e516ed4e 100644 --- a/src/extract.rs +++ b/src/extract.rs @@ -1,46 +1,22 @@ use crate::{body::Body, Error}; +use async_trait::async_trait; use bytes::Bytes; -use futures_util::{future, ready}; use http::{header, Request, StatusCode}; -use pin_project::pin_project; use serde::de::DeserializeOwned; -use std::{ - collections::HashMap, - future::Future, - pin::Pin, - str::FromStr, - task::{Context, Poll}, -}; +use std::{collections::HashMap, str::FromStr}; +#[async_trait] pub trait FromRequest: Sized { - type Future: Future> + Send; - - fn from_request(req: &mut Request) -> Self::Future; + async fn from_request(req: &mut Request) -> Result; } +#[async_trait] impl FromRequest for Option where T: FromRequest, { - type Future = OptionFromRequestFuture; - - fn from_request(req: &mut Request) -> Self::Future { - OptionFromRequestFuture(T::from_request(req)) - } -} - -#[pin_project] -pub struct OptionFromRequestFuture(#[pin] F); - -impl Future for OptionFromRequestFuture -where - F: Future>, -{ - type Output = Result, Error>; - - fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { - let value = ready!(self.project().0.poll(cx)); - Poll::Ready(Ok(value.ok())) + async fn from_request(req: &mut Request) -> Result, Error> { + Ok(T::from_request(req).await.ok()) } } @@ -53,20 +29,15 @@ impl Query { } } +#[async_trait] impl FromRequest for Query where - T: DeserializeOwned + Send, + T: DeserializeOwned, { - type Future = future::Ready>; - - fn from_request(req: &mut Request) -> Self::Future { - let result = (|| { - let query = req.uri().query().ok_or(Error::QueryStringMissing)?; - let value = serde_urlencoded::from_str(query).map_err(Error::DeserializeQueryString)?; - Ok(Query(value)) - })(); - - future::ready(result) + async fn from_request(req: &mut Request) -> Result { + let query = req.uri().query().ok_or(Error::QueryStringMissing)?; + let value = serde_urlencoded::from_str(query).map_err(Error::DeserializeQueryString)?; + Ok(Query(value)) } } @@ -79,26 +50,22 @@ impl Json { } } +#[async_trait] impl FromRequest for Json where T: DeserializeOwned, { - type Future = future::BoxFuture<'static, Result>; - - fn from_request(req: &mut Request) -> Self::Future { + async fn from_request(req: &mut Request) -> Result { if has_content_type(&req, "application/json") { let body = std::mem::take(req.body_mut()); - Box::pin(async move { - let bytes = hyper::body::to_bytes(body) - .await - .map_err(Error::ConsumeRequestBody)?; - let value = - serde_json::from_slice(&bytes).map_err(Error::DeserializeRequestBody)?; - Ok(Json(value)) - }) + let bytes = hyper::body::to_bytes(body) + .await + .map_err(Error::ConsumeRequestBody)?; + let value = serde_json::from_slice(&bytes).map_err(Error::DeserializeRequestBody)?; + Ok(Json(value)) } else { - Box::pin(async { Err(Error::Status(StatusCode::BAD_REQUEST)) }) + Err(Error::Status(StatusCode::BAD_REQUEST)) } } } @@ -128,66 +95,58 @@ impl Extension { } } +#[async_trait] impl FromRequest for Extension where T: Clone + Send + Sync + 'static, { - type Future = future::Ready>; + async fn from_request(req: &mut Request) -> Result { + let value = req + .extensions() + .get::() + .ok_or_else(|| Error::MissingExtension { + type_name: std::any::type_name::(), + }) + .map(|x| x.clone())?; - fn from_request(req: &mut Request) -> Self::Future { - let result = (|| { - let value = req - .extensions() - .get::() - .ok_or_else(|| Error::MissingExtension { - type_name: std::any::type_name::(), - }) - .map(|x| x.clone())?; - Ok(Extension(value)) - })(); - - future::ready(result) + Ok(Extension(value)) } } +#[async_trait] impl FromRequest for Bytes { - type Future = future::BoxFuture<'static, Result>; - - fn from_request(req: &mut Request) -> Self::Future { + async fn from_request(req: &mut Request) -> Result { let body = std::mem::take(req.body_mut()); - Box::pin(async move { - let bytes = hyper::body::to_bytes(body) - .await - .map_err(Error::ConsumeRequestBody)?; - Ok(bytes) - }) + let bytes = hyper::body::to_bytes(body) + .await + .map_err(Error::ConsumeRequestBody)?; + + Ok(bytes) } } +#[async_trait] impl FromRequest for String { - type Future = future::BoxFuture<'static, Result>; - - fn from_request(req: &mut Request) -> Self::Future { + async fn from_request(req: &mut Request) -> Result { let body = std::mem::take(req.body_mut()); - Box::pin(async move { - let bytes = hyper::body::to_bytes(body) - .await - .map_err(Error::ConsumeRequestBody)? - .to_vec(); - let string = String::from_utf8(bytes).map_err(|_| Error::InvalidUtf8)?; - Ok(string) - }) + let bytes = hyper::body::to_bytes(body) + .await + .map_err(Error::ConsumeRequestBody)? + .to_vec(); + + let string = String::from_utf8(bytes).map_err(|_| Error::InvalidUtf8)?; + + Ok(string) } } +#[async_trait] impl FromRequest for Body { - type Future = future::Ready>; - - fn from_request(req: &mut Request) -> Self::Future { + async fn from_request(req: &mut Request) -> Result { let body = std::mem::take(req.body_mut()); - future::ok(body) + Ok(body) } } @@ -200,31 +159,28 @@ impl BytesMaxLength { } } +#[async_trait] impl FromRequest for BytesMaxLength { - type Future = future::BoxFuture<'static, Result>; - - fn from_request(req: &mut Request) -> Self::Future { + async fn from_request(req: &mut Request) -> Result { let content_length = req.headers().get(http::header::CONTENT_LENGTH).cloned(); let body = std::mem::take(req.body_mut()); - Box::pin(async move { - let content_length = - content_length.and_then(|value| value.to_str().ok()?.parse::().ok()); + 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(Error::PayloadTooLarge); - } - } else { - return Err(Error::LengthRequired); - }; + if let Some(length) = content_length { + if length > N { + return Err(Error::PayloadTooLarge); + } + } else { + return Err(Error::LengthRequired); + }; - let bytes = hyper::body::to_bytes(body) - .await - .map_err(Error::ConsumeRequestBody)?; + let bytes = hyper::body::to_bytes(body) + .await + .map_err(Error::ConsumeRequestBody)?; - Ok(BytesMaxLength(bytes)) - }) + Ok(BytesMaxLength(bytes)) } } @@ -249,16 +205,15 @@ impl UrlParamsMap { } } +#[async_trait] impl FromRequest for UrlParamsMap { - type Future = future::Ready>; - - fn from_request(req: &mut Request) -> Self::Future { + 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; - future::ok(Self(params.into_iter().collect())) + Ok(Self(params.into_iter().collect())) } else { panic!("no url params found for matched route. This is a bug in tower-web") } @@ -277,15 +232,14 @@ 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 Future = future::Ready>; - #[allow(non_snake_case)] - fn from_request(req: &mut Request) -> Self::Future { + async fn from_request(req: &mut Request) -> Result { let params = if let Some(params) = req .extensions_mut() .get_mut::>() @@ -299,7 +253,7 @@ macro_rules! impl_parse_url { let $head = if let Ok(x) = $head.parse::<$head>() { x } else { - return future::err(Error::InvalidUrlParam { + return Err(Error::InvalidUrlParam { type_name: std::any::type_name::<$head>(), }); }; @@ -308,13 +262,13 @@ macro_rules! impl_parse_url { let $tail = if let Ok(x) = $tail.parse::<$tail>() { x } else { - return future::err(Error::InvalidUrlParam { + return Err(Error::InvalidUrlParam { type_name: std::any::type_name::<$tail>(), }); }; )* - future::ok(UrlParams(($head, $($tail,)*))) + Ok(UrlParams(($head, $($tail,)*))) } else { panic!("wrong number of url params found for matched route. This is a bug in tower-web") }