diff --git a/axum-extra/CHANGELOG.md b/axum-extra/CHANGELOG.md index a93e88e0..b1dd34c3 100644 --- a/axum-extra/CHANGELOG.md +++ b/axum-extra/CHANGELOG.md @@ -12,9 +12,11 @@ and this project adheres to [Semantic Versioning]. - **changed:** For methods that accept some `S: Service`, the bounds have been relaxed so the response type must implement `IntoResponse` rather than being a literal `Response` +- **added:** Support chaining handlers with `HandlerCallWithExtractors::or` ([#1170]) - **change:** axum-extra's MSRV is now 1.60 ([#1239]) [#1119]: https://github.com/tokio-rs/axum/pull/1119 +[#1170]: https://github.com/tokio-rs/axum/pull/1170 [#1239]: https://github.com/tokio-rs/axum/pull/1239 # 0.3.5 (27. June, 2022) diff --git a/axum-extra/src/handler/mod.rs b/axum-extra/src/handler/mod.rs new file mode 100644 index 00000000..5a5ccf1f --- /dev/null +++ b/axum-extra/src/handler/mod.rs @@ -0,0 +1,192 @@ +//! Additional handler utilities. + +use axum::{ + extract::{FromRequest, RequestParts}, + handler::Handler, + response::{IntoResponse, Response}, +}; +use futures_util::future::{BoxFuture, FutureExt, Map}; +use std::{future::Future, marker::PhantomData}; + +mod or; + +pub use self::or::Or; + +/// Trait for async functions that can be used to handle requests. +/// +/// This trait is similar to [`Handler`] but rather than taking the request it takes the extracted +/// inputs. +/// +/// The drawbacks of this trait is that you cannot apply middleware to individual handlers like you +/// can with [`Handler::layer`]. +pub trait HandlerCallWithExtractors: Sized { + /// The type of future calling this handler returns. + type Future: Future + Send + 'static; + + /// Call the handler with the extracted inputs. + fn call(self, extractors: T) -> >::Future; + + /// Conver this `HandlerCallWithExtractors` into [`Handler`]. + fn into_handler(self) -> IntoHandler { + IntoHandler { + handler: self, + _marker: PhantomData, + } + } + + /// Chain two handlers together, running the second one if the first one rejects. + /// + /// Note that this only moves to the next handler if an extractor fails. The response from + /// handlers are not considered. + /// + /// # Example + /// + /// ``` + /// use axum_extra::handler::HandlerCallWithExtractors; + /// use axum::{ + /// Router, + /// async_trait, + /// routing::get, + /// extract::FromRequest, + /// }; + /// + /// // handlers for varying levels of access + /// async fn admin(admin: AdminPermissions) { + /// // request came from an admin + /// } + /// + /// async fn user(user: User) { + /// // we have a `User` + /// } + /// + /// async fn guest() { + /// // `AdminPermissions` and `User` failed, so we're just a guest + /// } + /// + /// // extractors for checking permissions + /// struct AdminPermissions {} + /// + /// #[async_trait] + /// impl FromRequest for AdminPermissions { + /// // check for admin permissions... + /// # type Rejection = (); + /// # async fn from_request(req: &mut axum::extract::RequestParts) -> Result { + /// # todo!() + /// # } + /// } + /// + /// struct User {} + /// + /// #[async_trait] + /// impl FromRequest for User { + /// // check for a logged in user... + /// # type Rejection = (); + /// # async fn from_request(req: &mut axum::extract::RequestParts) -> Result { + /// # todo!() + /// # } + /// } + /// + /// let app = Router::new().route( + /// "/users/:id", + /// get( + /// // first try `admin`, if that rejects run `user`, finally falling back + /// // to `guest` + /// admin.or(user).or(guest) + /// ) + /// ); + /// # let _: Router = app; + /// ``` + fn or(self, rhs: R) -> Or + where + R: HandlerCallWithExtractors, + { + Or { + lhs: self, + rhs, + _marker: PhantomData, + } + } +} + +macro_rules! impl_handler_call_with { + ( $($ty:ident),* $(,)? ) => { + #[allow(non_snake_case)] + impl HandlerCallWithExtractors<($($ty,)*), B> for F + where + F: FnOnce($($ty,)*) -> Fut, + Fut: Future + Send + 'static, + Fut::Output: IntoResponse, + { + // this puts `futures_util` in our public API but thats fine in axum-extra + type Future = Map Response>; + + fn call( + self, + ($($ty,)*): ($($ty,)*), + ) -> >::Future { + self($($ty,)*).map(IntoResponse::into_response) + } + } + }; +} + +impl_handler_call_with!(); +impl_handler_call_with!(T1); +impl_handler_call_with!(T1, T2); +impl_handler_call_with!(T1, T2, T3); +impl_handler_call_with!(T1, T2, T3, T4); +impl_handler_call_with!(T1, T2, T3, T4, T5); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8, T9); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13, T14); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13, T14, T15); +impl_handler_call_with!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13, T14, T15, T16); + +/// A [`Handler`] created from a [`HandlerCallWithExtractors`]. +/// +/// Created with [`HandlerCallWithExtractors::into_handler`]. +#[allow(missing_debug_implementations)] +pub struct IntoHandler { + handler: H, + _marker: PhantomData (T, B)>, +} + +impl Handler for IntoHandler +where + H: HandlerCallWithExtractors + Clone + Send + 'static, + T: FromRequest + Send + 'static, + T::Rejection: Send, + B: Send + 'static, +{ + type Future = BoxFuture<'static, Response>; + + fn call(self, req: http::Request) -> Self::Future { + Box::pin(async move { + let mut req = RequestParts::new(req); + match req.extract::().await { + Ok(t) => self.handler.call(t).await, + Err(rejection) => rejection.into_response(), + } + }) + } +} + +impl Copy for IntoHandler where H: Copy {} + +impl Clone for IntoHandler +where + H: Clone, +{ + fn clone(&self) -> Self { + Self { + handler: self.handler.clone(), + _marker: self._marker, + } + } +} diff --git a/axum-extra/src/handler/or.rs b/axum-extra/src/handler/or.rs new file mode 100644 index 00000000..195dcfd8 --- /dev/null +++ b/axum-extra/src/handler/or.rs @@ -0,0 +1,151 @@ +use super::HandlerCallWithExtractors; +use crate::Either; +use axum::{ + extract::{FromRequest, RequestParts}, + handler::Handler, + http::Request, + response::{IntoResponse, Response}, +}; +use futures_util::future::{BoxFuture, Either as EitherFuture, FutureExt, Map}; +use http::StatusCode; +use std::{future::Future, marker::PhantomData}; + +/// [`Handler`] that runs one [`Handler`] and if that rejects it'll fallback to another +/// [`Handler`]. +/// +/// Created with [`HandlerCallWithExtractors::or`](super::HandlerCallWithExtractors::or). +#[allow(missing_debug_implementations)] +pub struct Or { + pub(super) lhs: L, + pub(super) rhs: R, + pub(super) _marker: PhantomData (Lt, Rt, B)>, +} + +impl HandlerCallWithExtractors, B> for Or +where + L: HandlerCallWithExtractors + Send + 'static, + R: HandlerCallWithExtractors + Send + 'static, + Rt: Send + 'static, + Lt: Send + 'static, + B: Send + 'static, +{ + // this puts `futures_util` in our public API but thats fine in axum-extra + type Future = EitherFuture< + Map::Output) -> Response>, + Map::Output) -> Response>, + >; + + fn call( + self, + extractors: Either, + ) -> , B>>::Future { + match extractors { + Either::Left(lt) => self + .lhs + .call(lt) + .map(IntoResponse::into_response as _) + .left_future(), + Either::Right(rt) => self + .rhs + .call(rt) + .map(IntoResponse::into_response as _) + .right_future(), + } + } +} + +impl Handler<(Lt, Rt), B> for Or +where + L: HandlerCallWithExtractors + Clone + Send + 'static, + R: HandlerCallWithExtractors + Clone + Send + 'static, + Lt: FromRequest + Send + 'static, + Rt: FromRequest + Send + 'static, + Lt::Rejection: Send, + Rt::Rejection: Send, + B: Send + 'static, +{ + // this puts `futures_util` in our public API but thats fine in axum-extra + type Future = BoxFuture<'static, Response>; + + fn call(self, req: Request) -> Self::Future { + Box::pin(async move { + let mut req = RequestParts::new(req); + + if let Ok(lt) = req.extract::().await { + return self.lhs.call(lt).await; + } + + if let Ok(rt) = req.extract::().await { + return self.rhs.call(rt).await; + } + + StatusCode::NOT_FOUND.into_response() + }) + } +} + +impl Copy for Or +where + L: Copy, + R: Copy, +{ +} + +impl Clone for Or +where + L: Clone, + R: Clone, +{ + fn clone(&self) -> Self { + Self { + lhs: self.lhs.clone(), + rhs: self.rhs.clone(), + _marker: self._marker, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_helpers::*; + use axum::{ + extract::{Path, Query}, + routing::get, + Router, + }; + use serde::Deserialize; + + #[tokio::test] + async fn works() { + #[derive(Deserialize)] + struct Params { + a: String, + } + + async fn one(Path(id): Path) -> String { + id.to_string() + } + + async fn two(Query(params): Query) -> String { + params.a + } + + async fn three() -> &'static str { + "fallback" + } + + let app = Router::new().route("/:id", get(one.or(two).or(three))); + + let client = TestClient::new(app); + + let res = client.get("/123").send().await; + assert_eq!(res.text().await, "123"); + + let res = client.get("/foo?a=bar").send().await; + assert_eq!(res.text().await, "bar"); + + let res = client.get("/foo").send().await; + assert_eq!(res.text().await, "fallback"); + } +} diff --git a/axum-extra/src/lib.rs b/axum-extra/src/lib.rs index a63080eb..6f7ac4e2 100644 --- a/axum-extra/src/lib.rs +++ b/axum-extra/src/lib.rs @@ -63,14 +63,61 @@ #![cfg_attr(docsrs, feature(doc_cfg, doc_auto_cfg))] #![cfg_attr(test, allow(clippy::float_cmp))] +use axum::{ + async_trait, + extract::{FromRequest, RequestParts}, + response::IntoResponse, +}; + pub mod body; pub mod extract; +pub mod handler; pub mod response; pub mod routing; #[cfg(feature = "json-lines")] pub mod json_lines; +/// Combines two extractors or responses into a single type. +#[derive(Debug, Copy, Clone)] +pub enum Either { + /// A value of type L. + Left(L), + /// A value of type R. + Right(R), +} + +#[async_trait] +impl FromRequest for Either +where + L: FromRequest, + R: FromRequest, + B: Send, +{ + type Rejection = R::Rejection; + + async fn from_request(req: &mut RequestParts) -> Result { + if let Ok(l) = req.extract().await { + return Ok(Either::Left(l)); + } + + Ok(Either::Right(req.extract().await?)) + } +} + +impl IntoResponse for Either +where + L: IntoResponse, + R: IntoResponse, +{ + fn into_response(self) -> axum::response::Response { + match self { + Self::Left(inner) => inner.into_response(), + Self::Right(inner) => inner.into_response(), + } + } +} + #[cfg(feature = "typed-routing")] #[doc(hidden)] pub mod __private {