diff --git a/axum-extra/Cargo.toml b/axum-extra/Cargo.toml index 13517e66..6501a50d 100644 --- a/axum-extra/Cargo.toml +++ b/axum-extra/Cargo.toml @@ -15,6 +15,7 @@ erased-json = ["serde", "serde_json"] [dependencies] axum = { path = "../axum", version = "0.3" } +http = "0.2" mime = "0.3" tower-service = "0.3" diff --git a/axum-extra/src/extract/cached.rs b/axum-extra/src/extract/cached.rs new file mode 100644 index 00000000..613a79d3 --- /dev/null +++ b/axum-extra/src/extract/cached.rs @@ -0,0 +1,235 @@ +use axum::{ + async_trait, + body::{boxed, BoxBody}, + extract::{ + rejection::{ExtensionRejection, ExtensionsAlreadyExtracted}, + Extension, FromRequest, RequestParts, + }, + http::Response, + response::IntoResponse, +}; +use std::{ + fmt, + ops::{Deref, DerefMut}, +}; + +/// Cache results of other extractors. +/// +/// `Cached` wraps another extractor and caches its result in [request extensions]. +/// +/// This is useful if you have a tree of extractors that share common sub-extractors that +/// you only want to run once, perhaps because they're expensive. +/// +/// The cache purely type based so you can only cache one value of each type. The cache is also +/// local to the current request and not reused across requests. +/// +/// # Example +/// +/// ```rust +/// use axum_extra::extract::Cached; +/// use axum::{ +/// async_trait, +/// extract::{FromRequest, RequestParts}, +/// body::{self, BoxBody}, +/// response::IntoResponse, +/// http::{StatusCode, Response}, +/// }; +/// +/// #[derive(Clone)] +/// struct Session { /* ... */ } +/// +/// #[async_trait] +/// impl FromRequest for Session +/// where +/// B: Send, +/// { +/// type Rejection = (StatusCode, String); +/// +/// async fn from_request(req: &mut RequestParts) -> Result { +/// // load session... +/// # unimplemented!() +/// } +/// } +/// +/// struct CurrentUser { /* ... */ } +/// +/// #[async_trait] +/// impl FromRequest for CurrentUser +/// where +/// B: Send, +/// { +/// type Rejection = Response; +/// +/// async fn from_request(req: &mut RequestParts) -> Result { +/// // loading a `CurrentUser` requires first loading the `Session` +/// // +/// // by using `Cached` we avoid extracting the session more than +/// // once, in case other extractors for the same request also loads the session +/// let session: Session = Cached::::from_request(req) +/// .await +/// .map_err(|err| err.into_response().map(body::boxed))? +/// .0; +/// +/// // load user from session... +/// # unimplemented!() +/// } +/// } +/// +/// // handler that extracts the current user and the session +/// // +/// // the session will only be loaded once, even though `CurrentUser` +/// // also loads it +/// async fn handler( +/// current_user: CurrentUser, +/// // we have to use `Cached` here otherwise the +/// // cached session would not be used +/// Cached(session): Cached, +/// ) { +/// // ... +/// } +/// ``` +/// +/// [request extensions]: http::Extensions +#[derive(Debug, Clone, Default)] +pub struct Cached(pub T); + +#[derive(Clone)] +struct CachedEntry(T); + +#[async_trait] +impl FromRequest for Cached +where + B: Send, + T: FromRequest + Clone + Send + Sync + 'static, +{ + type Rejection = CachedRejection; + + async fn from_request(req: &mut RequestParts) -> Result { + match Extension::>::from_request(req).await { + Ok(Extension(CachedEntry(value))) => Ok(Self(value)), + Err(ExtensionRejection::ExtensionsAlreadyExtracted(err)) => { + Err(CachedRejection::ExtensionsAlreadyExtracted(err)) + } + Err(_) => { + let value = T::from_request(req).await.map_err(CachedRejection::Inner)?; + + req.extensions_mut() + .ok_or_else(|| { + CachedRejection::ExtensionsAlreadyExtracted( + ExtensionsAlreadyExtracted::default(), + ) + })? + .insert(CachedEntry(value.clone())); + + Ok(Self(value)) + } + } + } +} + +impl Deref for Cached { + type Target = T; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl DerefMut for Cached { + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.0 + } +} + +/// Rejection used for [`Cached`]. +/// +/// Contains one variant for each way the [`Cached`] extractor can fail. +#[derive(Debug)] +#[non_exhaustive] +pub enum CachedRejection { + #[allow(missing_docs)] + ExtensionsAlreadyExtracted(ExtensionsAlreadyExtracted), + #[allow(missing_docs)] + Inner(R), +} + +impl IntoResponse for CachedRejection +where + R: IntoResponse, +{ + type Body = BoxBody; + type BodyError = ::Error; + + fn into_response(self) -> Response { + match self { + Self::ExtensionsAlreadyExtracted(inner) => inner.into_response().map(boxed), + Self::Inner(inner) => inner.into_response().map(boxed), + } + } +} + +impl fmt::Display for CachedRejection +where + R: fmt::Display, +{ + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::ExtensionsAlreadyExtracted(inner) => write!(f, "{}", inner), + Self::Inner(inner) => write!(f, "{}", inner), + } + } +} + +impl std::error::Error for CachedRejection +where + R: std::error::Error + 'static, +{ + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + Self::ExtensionsAlreadyExtracted(inner) => Some(inner), + Self::Inner(inner) => Some(inner), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::http::Request; + use std::{ + convert::Infallible, + sync::atomic::{AtomicU32, Ordering}, + time::Instant, + }; + + #[tokio::test] + async fn works() { + static COUNTER: AtomicU32 = AtomicU32::new(0); + + #[derive(Clone, Debug, PartialEq, Eq)] + struct Extractor(Instant); + + #[async_trait] + impl FromRequest for Extractor + where + B: Send, + { + type Rejection = Infallible; + + async fn from_request(_req: &mut RequestParts) -> Result { + COUNTER.fetch_add(1, Ordering::SeqCst); + Ok(Self(Instant::now())) + } + } + + let mut req = RequestParts::new(Request::new(())); + + let first = Cached::::from_request(&mut req).await.unwrap().0; + assert_eq!(COUNTER.load(Ordering::SeqCst), 1); + + let second = Cached::::from_request(&mut req).await.unwrap().0; + assert_eq!(COUNTER.load(Ordering::SeqCst), 1); + + assert_eq!(first, second); + } +} diff --git a/axum-extra/src/extract/mod.rs b/axum-extra/src/extract/mod.rs new file mode 100644 index 00000000..434fafa5 --- /dev/null +++ b/axum-extra/src/extract/mod.rs @@ -0,0 +1,11 @@ +//! Additional extractors. + +mod cached; + +pub use self::cached::Cached; + +pub mod rejection { + //! Rejection response types. + + pub use super::cached::CachedRejection; +} diff --git a/axum-extra/src/lib.rs b/axum-extra/src/lib.rs index e7f8384e..bb16461f 100644 --- a/axum-extra/src/lib.rs +++ b/axum-extra/src/lib.rs @@ -43,5 +43,6 @@ #![cfg_attr(docsrs, feature(doc_cfg))] #![cfg_attr(test, allow(clippy::float_cmp))] +pub mod extract; pub mod response; pub mod routing; diff --git a/axum/src/macros.rs b/axum/src/macros.rs index 52946a89..2a9e9554 100644 --- a/axum/src/macros.rs +++ b/axum/src/macros.rs @@ -76,6 +76,12 @@ macro_rules! define_rejection { } impl std::error::Error for $name {} + + impl Default for $name { + fn default() -> Self { + Self + } + } }; (