diff --git a/axum/CHANGELOG.md b/axum/CHANGELOG.md index 37a98efb..4326f953 100644 --- a/axum/CHANGELOG.md +++ b/axum/CHANGELOG.md @@ -9,8 +9,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - **added:** Add `response::ErrorResponse` and `response::Result` for `IntoResponse`-based error handling ([#921]) +- **added:** Add `middleware::from_extractor` and deprecate `extract::extractor_middleware` ([#957]) -[#921]: https://github.com/tokio-rs/axum/pull/921 +[#921]: https://github.com/tokio-rs/axum/pull/921 +[#957]: https://github.com/tokio-rs/axum/pull/957 # 0.5.3 (19. April, 2022) diff --git a/axum/src/docs/middleware.md b/axum/src/docs/middleware.md index 8037fe47..afe438ef 100644 --- a/axum/src/docs/middleware.md +++ b/axum/src/docs/middleware.md @@ -168,9 +168,9 @@ Use [`axum::middleware::from_fn`] to write your middleware when: - You don't intend to publish your middleware as a crate for others to use. Middleware written like this are only compatible with axum. -## `axum::extract::extractor_middleware` +## `axum::middleware::from_extractor` -Use [`axum::extract::extractor_middleware`] to write your middleware when: +Use [`axum::middleware::from_extractor`] to write your middleware when: - You have a type that you sometimes want to use as an extractor and sometimes as a middleware. If you only need your type as a middleware prefer @@ -442,7 +442,7 @@ extensions you need. [`ServiceBuilder::map_response`]: tower::ServiceBuilder::map_response [`ServiceBuilder::then`]: tower::ServiceBuilder::then [`ServiceBuilder::and_then`]: tower::ServiceBuilder::and_then -[`axum::extract::extractor_middleware`]: crate::extract::extractor_middleware() +[`axum::middleware::from_extractor`]: crate::extract::extractor_middleware() [`Handler::layer`]: crate::handler::Handler::layer [`Router::layer`]: crate::routing::Router::layer [`MethodRouter::layer`]: crate::routing::MethodRouter::layer diff --git a/axum/src/extract/extractor_middleware.rs b/axum/src/extract/extractor_middleware.rs index aca7b6f3..9e3c1942 100644 --- a/axum/src/extract/extractor_middleware.rs +++ b/axum/src/extract/extractor_middleware.rs @@ -2,324 +2,15 @@ //! //! See [`extractor_middleware`] for more details. -use super::{FromRequest, RequestParts}; -use crate::{ - body::{Bytes, HttpBody}, - response::{IntoResponse, Response}, - BoxError, +use crate::middleware::from_extractor; + +pub use crate::middleware::{ + future::FromExtractorResponseFuture as ResponseFuture, FromExtractor as ExtractorMiddleware, + FromExtractorLayer as ExtractorMiddlewareLayer, }; -use futures_util::{future::BoxFuture, ready}; -use http::Request; -use pin_project_lite::pin_project; -use std::{ - fmt, - future::Future, - marker::PhantomData, - pin::Pin, - task::{Context, Poll}, -}; -use tower_layer::Layer; -use tower_service::Service; /// Convert an extractor into a middleware. -/// -/// If the extractor succeeds the value will be discarded and the inner service -/// will be called. If the extractor fails the rejection will be returned and -/// the inner service will _not_ be called. -/// -/// This can be used to perform validation of requests if the validation doesn't -/// produce any useful output, and run the extractor for several handlers -/// without repeating it in the function signature. -/// -/// Note that if the extractor consumes the request body, as `String` or -/// [`Bytes`] does, an empty body will be left in its place. Thus wont be -/// accessible to subsequent extractors or handlers. -/// -/// # Example -/// -/// ```rust -/// use axum::{ -/// extract::{extractor_middleware, FromRequest, RequestParts}, -/// routing::{get, post}, -/// Router, -/// }; -/// use http::StatusCode; -/// use async_trait::async_trait; -/// -/// // An extractor that performs authorization. -/// struct RequireAuth; -/// -/// #[async_trait] -/// impl FromRequest for RequireAuth -/// where -/// B: Send, -/// { -/// type Rejection = StatusCode; -/// -/// async fn from_request(req: &mut RequestParts) -> Result { -/// let auth_header = req -/// .headers() -/// .get(http::header::AUTHORIZATION) -/// .and_then(|value| value.to_str().ok()); -/// -/// match auth_header { -/// Some(auth_header) if token_is_valid(auth_header) => { -/// Ok(Self) -/// } -/// _ => Err(StatusCode::UNAUTHORIZED), -/// } -/// } -/// } -/// -/// fn token_is_valid(token: &str) -> bool { -/// // ... -/// # false -/// } -/// -/// async fn handler() { -/// // If we get here the request has been authorized -/// } -/// -/// async fn other_handler() { -/// // If we get here the request has been authorized -/// } -/// -/// let app = Router::new() -/// .route("/", get(handler)) -/// .route("/foo", post(other_handler)) -/// // The extractor will run before all routes -/// .route_layer(extractor_middleware::()); -/// # async { -/// # axum::Server::bind(&"".parse().unwrap()).serve(app.into_make_service()).await.unwrap(); -/// # }; -/// ``` +#[deprecated(note = "Please use `axum::middleware::from_extractor` instead")] pub fn extractor_middleware() -> ExtractorMiddlewareLayer { - ExtractorMiddlewareLayer(PhantomData) -} - -/// [`Layer`] that applies [`ExtractorMiddleware`] that runs an extractor and -/// discards the value. -/// -/// See [`extractor_middleware`] for more details. -/// -/// [`Layer`]: tower::Layer -pub struct ExtractorMiddlewareLayer(PhantomData E>); - -impl Clone for ExtractorMiddlewareLayer { - fn clone(&self) -> Self { - Self(PhantomData) - } -} - -impl fmt::Debug for ExtractorMiddlewareLayer { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.debug_struct("ExtractorMiddleware") - .field("extractor", &format_args!("{}", std::any::type_name::())) - .finish() - } -} - -impl Layer for ExtractorMiddlewareLayer { - type Service = ExtractorMiddleware; - - fn layer(&self, inner: S) -> Self::Service { - ExtractorMiddleware { - inner, - _extractor: PhantomData, - } - } -} - -/// Middleware that runs an extractor and discards the value. -/// -/// See [`extractor_middleware`] for more details. -pub struct ExtractorMiddleware { - inner: S, - _extractor: PhantomData E>, -} - -#[test] -fn traits() { - use crate::test_helpers::*; - assert_send::>(); - assert_sync::>(); -} - -impl Clone for ExtractorMiddleware -where - S: Clone, -{ - fn clone(&self) -> Self { - Self { - inner: self.inner.clone(), - _extractor: PhantomData, - } - } -} - -impl fmt::Debug for ExtractorMiddleware -where - S: fmt::Debug, -{ - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.debug_struct("ExtractorMiddleware") - .field("inner", &self.inner) - .field("extractor", &format_args!("{}", std::any::type_name::())) - .finish() - } -} - -impl Service> for ExtractorMiddleware -where - E: FromRequest + 'static, - ReqBody: Default + Send + 'static, - S: Service, Response = Response> + Clone, - ResBody: HttpBody + Send + 'static, - ResBody::Error: Into, -{ - type Response = Response; - type Error = S::Error; - type Future = ResponseFuture; - - #[inline] - fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { - self.inner.poll_ready(cx) - } - - fn call(&mut self, req: Request) -> Self::Future { - let extract_future = Box::pin(async move { - let mut req = super::RequestParts::new(req); - let extracted = E::from_request(&mut req).await; - (req, extracted) - }); - - ResponseFuture { - state: State::Extracting { - future: extract_future, - }, - svc: Some(self.inner.clone()), - } - } -} - -pin_project! { - /// Response future for [`ExtractorMiddleware`]. - #[allow(missing_debug_implementations)] - pub struct ResponseFuture - where - E: FromRequest, - S: Service>, - { - #[pin] - state: State, - svc: Option, - } -} - -pin_project! { - #[project = StateProj] - enum State - where - E: FromRequest, - S: Service>, - { - Extracting { future: BoxFuture<'static, (RequestParts, Result)> }, - Call { #[pin] future: S::Future }, - } -} - -impl Future for ResponseFuture -where - E: FromRequest, - S: Service, Response = Response>, - ReqBody: Default, - ResBody: HttpBody + Send + 'static, - ResBody::Error: Into, -{ - type Output = Result; - - fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { - loop { - let mut this = self.as_mut().project(); - - let new_state = match this.state.as_mut().project() { - StateProj::Extracting { future } => { - let (req, extracted) = ready!(future.as_mut().poll(cx)); - - match extracted { - Ok(_) => { - let mut svc = this.svc.take().expect("future polled after completion"); - let req = req.try_into_request().unwrap_or_default(); - let future = svc.call(req); - State::Call { future } - } - Err(err) => { - let res = err.into_response(); - return Poll::Ready(Ok(res)); - } - } - } - StateProj::Call { future } => { - return future - .poll(cx) - .map(|result| result.map(|response| response.map(crate::body::boxed))); - } - }; - - this.state.set(new_state); - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::{handler::Handler, routing::get, test_helpers::*, Router}; - use http::StatusCode; - - #[tokio::test] - async fn test_extractor_middleware() { - struct RequireAuth; - - #[async_trait::async_trait] - impl FromRequest for RequireAuth - where - B: Send, - { - type Rejection = StatusCode; - - async fn from_request(req: &mut RequestParts) -> Result { - if let Some(auth) = req - .headers() - .get("authorization") - .and_then(|v| v.to_str().ok()) - { - if auth == "secret" { - return Ok(Self); - } - } - - Err(StatusCode::UNAUTHORIZED) - } - } - - async fn handler() {} - - let app = Router::new().route( - "/", - get(handler.layer(extractor_middleware::())), - ); - - let client = TestClient::new(app); - - let res = client.get("/").send().await; - assert_eq!(res.status(), StatusCode::UNAUTHORIZED); - - let res = client - .get("/") - .header(http::header::AUTHORIZATION, "secret") - .send() - .await; - assert_eq!(res.status(), StatusCode::OK); - } + from_extractor() } diff --git a/axum/src/extract/mod.rs b/axum/src/extract/mod.rs index db7f974a..cde177a7 100644 --- a/axum/src/extract/mod.rs +++ b/axum/src/extract/mod.rs @@ -20,6 +20,7 @@ mod request_parts; pub use axum_core::extract::{FromRequest, RequestParts}; #[doc(inline)] +#[allow(deprecated)] pub use self::{ connect_info::ConnectInfo, content_length_limit::ContentLengthLimit, diff --git a/axum/src/middleware/from_extractor.rs b/axum/src/middleware/from_extractor.rs new file mode 100644 index 00000000..45bb9513 --- /dev/null +++ b/axum/src/middleware/from_extractor.rs @@ -0,0 +1,319 @@ +use crate::{ + body::{Bytes, HttpBody}, + extract::{FromRequest, RequestParts}, + response::{IntoResponse, Response}, + BoxError, +}; +use futures_util::{future::BoxFuture, ready}; +use http::Request; +use pin_project_lite::pin_project; +use std::{ + fmt, + future::Future, + marker::PhantomData, + pin::Pin, + task::{Context, Poll}, +}; +use tower_layer::Layer; +use tower_service::Service; + +/// Create a middleware from an extractor. +/// +/// If the extractor succeeds the value will be discarded and the inner service +/// will be called. If the extractor fails the rejection will be returned and +/// the inner service will _not_ be called. +/// +/// This can be used to perform validation of requests if the validation doesn't +/// produce any useful output, and run the extractor for several handlers +/// without repeating it in the function signature. +/// +/// Note that if the extractor consumes the request body, as `String` or +/// [`Bytes`] does, an empty body will be left in its place. Thus wont be +/// accessible to subsequent extractors or handlers. +/// +/// # Example +/// +/// ```rust +/// use axum::{ +/// extract::{FromRequest, RequestParts}, +/// middleware::from_extractor, +/// routing::{get, post}, +/// Router, +/// }; +/// use http::{header, StatusCode}; +/// use async_trait::async_trait; +/// +/// // An extractor that performs authorization. +/// struct RequireAuth; +/// +/// #[async_trait] +/// impl FromRequest for RequireAuth +/// where +/// B: Send, +/// { +/// type Rejection = StatusCode; +/// +/// async fn from_request(req: &mut RequestParts) -> Result { +/// let auth_header = req +/// .headers() +/// .get(header::AUTHORIZATION) +/// .and_then(|value| value.to_str().ok()); +/// +/// match auth_header { +/// Some(auth_header) if token_is_valid(auth_header) => { +/// Ok(Self) +/// } +/// _ => Err(StatusCode::UNAUTHORIZED), +/// } +/// } +/// } +/// +/// fn token_is_valid(token: &str) -> bool { +/// // ... +/// # false +/// } +/// +/// async fn handler() { +/// // If we get here the request has been authorized +/// } +/// +/// async fn other_handler() { +/// // If we get here the request has been authorized +/// } +/// +/// let app = Router::new() +/// .route("/", get(handler)) +/// .route("/foo", post(other_handler)) +/// // The extractor will run before all routes +/// .route_layer(from_extractor::()); +/// # async { +/// # axum::Server::bind(&"".parse().unwrap()).serve(app.into_make_service()).await.unwrap(); +/// # }; +/// ``` +pub fn from_extractor() -> FromExtractorLayer { + FromExtractorLayer(PhantomData) +} + +/// [`Layer`] that applies [`FromExtractor`] that runs an extractor and +/// discards the value. +/// +/// See [`from_extractor`] for more details. +/// +/// [`Layer`]: tower::Layer +pub struct FromExtractorLayer(PhantomData E>); + +impl Clone for FromExtractorLayer { + fn clone(&self) -> Self { + Self(PhantomData) + } +} + +impl fmt::Debug for FromExtractorLayer { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("FromExtractorLayer") + .field("extractor", &format_args!("{}", std::any::type_name::())) + .finish() + } +} + +impl Layer for FromExtractorLayer { + type Service = FromExtractor; + + fn layer(&self, inner: S) -> Self::Service { + FromExtractor { + inner, + _extractor: PhantomData, + } + } +} + +/// Middleware that runs an extractor and discards the value. +/// +/// See [`from_extractor`] for more details. +pub struct FromExtractor { + inner: S, + _extractor: PhantomData E>, +} + +#[test] +fn traits() { + use crate::test_helpers::*; + assert_send::>(); + assert_sync::>(); +} + +impl Clone for FromExtractor +where + S: Clone, +{ + fn clone(&self) -> Self { + Self { + inner: self.inner.clone(), + _extractor: PhantomData, + } + } +} + +impl fmt::Debug for FromExtractor +where + S: fmt::Debug, +{ + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("FromExtractor") + .field("inner", &self.inner) + .field("extractor", &format_args!("{}", std::any::type_name::())) + .finish() + } +} + +impl Service> for FromExtractor +where + E: FromRequest + 'static, + ReqBody: Default + Send + 'static, + S: Service, Response = Response> + Clone, + ResBody: HttpBody + Send + 'static, + ResBody::Error: Into, +{ + type Response = Response; + type Error = S::Error; + type Future = ResponseFuture; + + #[inline] + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, req: Request) -> Self::Future { + let extract_future = Box::pin(async move { + let mut req = RequestParts::new(req); + let extracted = E::from_request(&mut req).await; + (req, extracted) + }); + + ResponseFuture { + state: State::Extracting { + future: extract_future, + }, + svc: Some(self.inner.clone()), + } + } +} + +pin_project! { + /// Response future for [`FromExtractor`]. + #[allow(missing_debug_implementations)] + pub struct ResponseFuture + where + E: FromRequest, + S: Service>, + { + #[pin] + state: State, + svc: Option, + } +} + +pin_project! { + #[project = StateProj] + enum State + where + E: FromRequest, + S: Service>, + { + Extracting { future: BoxFuture<'static, (RequestParts, Result)> }, + Call { #[pin] future: S::Future }, + } +} + +impl Future for ResponseFuture +where + E: FromRequest, + S: Service, Response = Response>, + ReqBody: Default, + ResBody: HttpBody + Send + 'static, + ResBody::Error: Into, +{ + type Output = Result; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + loop { + let mut this = self.as_mut().project(); + + let new_state = match this.state.as_mut().project() { + StateProj::Extracting { future } => { + let (req, extracted) = ready!(future.as_mut().poll(cx)); + + match extracted { + Ok(_) => { + let mut svc = this.svc.take().expect("future polled after completion"); + let req = req.try_into_request().unwrap_or_default(); + let future = svc.call(req); + State::Call { future } + } + Err(err) => { + let res = err.into_response(); + return Poll::Ready(Ok(res)); + } + } + } + StateProj::Call { future } => { + return future + .poll(cx) + .map(|result| result.map(|response| response.map(crate::body::boxed))); + } + }; + + this.state.set(new_state); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{handler::Handler, routing::get, test_helpers::*, Router}; + use http::{header, StatusCode}; + + #[tokio::test] + async fn test_from_extractor() { + struct RequireAuth; + + #[async_trait::async_trait] + impl FromRequest for RequireAuth + where + B: Send, + { + type Rejection = StatusCode; + + async fn from_request(req: &mut RequestParts) -> Result { + if let Some(auth) = req + .headers() + .get(header::AUTHORIZATION) + .and_then(|v| v.to_str().ok()) + { + if auth == "secret" { + return Ok(Self); + } + } + + Err(StatusCode::UNAUTHORIZED) + } + } + + async fn handler() {} + + let app = Router::new().route("/", get(handler.layer(from_extractor::()))); + + let client = TestClient::new(app); + + let res = client.get("/").send().await; + assert_eq!(res.status(), StatusCode::UNAUTHORIZED); + + let res = client + .get("/") + .header(http::header::AUTHORIZATION, "secret") + .send() + .await; + assert_eq!(res.status(), StatusCode::OK); + } +} diff --git a/axum/src/middleware/mod.rs b/axum/src/middleware/mod.rs index 3273b7d4..f8be812b 100644 --- a/axum/src/middleware/mod.rs +++ b/axum/src/middleware/mod.rs @@ -2,13 +2,16 @@ //! #![doc = include_str!("../docs/middleware.md")] +mod from_extractor; mod from_fn; +pub use self::from_extractor::{from_extractor, FromExtractor, FromExtractorLayer}; pub use self::from_fn::{from_fn, FromFn, FromFnLayer, Next}; pub use crate::extension::AddExtension; pub mod future { //! Future types. + pub use super::from_extractor::ResponseFuture as FromExtractorResponseFuture; pub use super::from_fn::ResponseFuture as FromFnResponseFuture; }