diff --git a/axum/src/handler/into_extension_service.rs b/axum/src/handler/into_extension_service.rs index 19e68281..1d9e54ef 100644 --- a/axum/src/handler/into_extension_service.rs +++ b/axum/src/handler/into_extension_service.rs @@ -58,6 +58,9 @@ where use futures_util::future::FutureExt; let handler = self.handler.clone(); + + // TODO(david): this is duplicated in `axum/src/routing/mod.rs` + // extract into helper function let State(state) = req .extensions() .get::>() diff --git a/axum/src/routing/method_routing.rs b/axum/src/routing/method_routing.rs index 0a84e39b..e9cf34cf 100644 --- a/axum/src/routing/method_routing.rs +++ b/axum/src/routing/method_routing.rs @@ -672,6 +672,26 @@ impl MethodRouter { } } +impl MethodRouter { + pub(crate) fn change_state(self) -> MethodRouter { + debug_assert!(self.state.is_none()); + MethodRouter { + state: None, + get: self.get, + head: self.head, + delete: self.delete, + options: self.options, + patch: self.patch, + post: self.post, + put: self.put, + trace: self.trace, + fallback: self.fallback, + allow_header: self.allow_header, + _marker: PhantomData, + } + } +} + impl MethodRouter where B: Send + 'static, @@ -1181,10 +1201,8 @@ where if req.extensions().get::>().is_none() { // the `unwrap` is safe because `self.state` is always some if `R = WithState`, which it is - let prev = req - .extensions_mut() + req.extensions_mut() .insert(State(state.as_ref().unwrap().clone())); - debug_assert!(prev.is_none()); } call!(req, method, HEAD, head); diff --git a/axum/src/routing/mod.rs b/axum/src/routing/mod.rs index e3851c2b..8bb15948 100644 --- a/axum/src/routing/mod.rs +++ b/axum/src/routing/mod.rs @@ -21,7 +21,7 @@ use std::{ sync::Arc, task::{Context, Poll}, }; -use tower::{layer::layer_fn, ServiceBuilder}; +use tower::{layer::layer_fn, util::MapRequestLayer, ServiceBuilder}; use tower_http::map_response_body::MapResponseBodyLayer; use tower_layer::Layer; use tower_service::Service; @@ -168,13 +168,63 @@ where _marker: PhantomData, } } +} +impl Router +where + B: HttpBody + Send + 'static, +{ pub fn map_state(self, f: F) -> Router where - // TODO(david): which Fn? - F: FnOnce(OuterState) -> S, + F: Fn(OuterState) -> InnerState + Clone + Send + Sync + 'static, + OuterState: Clone + Send + Sync + 'static, + InnerState: Send + Sync + 'static, { - todo!() + debug_assert!(self.state.is_none()); + + let routes = self + .routes + .into_iter() + .map(|(route_id, endpoint)| { + let endpoint = match endpoint { + Endpoint::MethodRouter(method_router) => { + // the state will be provided later in `::call`, so its + // safe to ignore that it hasn't been provided yet + Endpoint::MethodRouter(method_router.change_state::()) + } + Endpoint::Route(route) => Endpoint::Route(route), + }; + (route_id, endpoint) + }) + .collect(); + + Router { + state: None, + routes, + node: self.node, + fallback: self.fallback, + _marker: PhantomData, + } + .layer(MapRequestLayer::new(move |mut req: Request<_>| { + // TODO(david): this is duplicated in `axum/src/handler/into_extension_service.rs` + // extract into helper function + let State(outer_state) = req + .extensions() + .get::>() + .unwrap_or_else(|| { + panic!( + "no state of type `{}` was found. Please file an issue", + std::any::type_name::>() + ) + }) + .clone(); + + let inner_state = f(outer_state); + + req.extensions_mut().insert(State(inner_state)); + + req + })) } } diff --git a/axum/src/routing/tests/mod.rs b/axum/src/routing/tests/mod.rs index 37c6dcd7..a00a5b7f 100644 --- a/axum/src/routing/tests/mod.rs +++ b/axum/src/routing/tests/mod.rs @@ -695,3 +695,8 @@ async fn extracting_state() { let res = client.get("/").send().await; assert_eq!(res.text().await, "foo"); } + +#[tokio::test] +async fn extracting_wrong_state_type_doesnt_compile() { + todo!() +} diff --git a/axum/src/routing/tests/nest.rs b/axum/src/routing/tests/nest.rs index efb751ad..99fd5972 100644 --- a/axum/src/routing/tests/nest.rs +++ b/axum/src/routing/tests/nest.rs @@ -389,28 +389,35 @@ async fn nest_with_and_without_trailing() { #[tokio::test] async fn nesting_with_different_state() { #[derive(Clone)] - struct State { + struct AppState { inner: InnerState, } #[derive(Clone)] - struct InnerState {} + struct InnerState { + value: &'static str, + } - impl From for InnerState { - fn from(state: State) -> Self { + impl From for InnerState { + fn from(state: AppState) -> Self { state.inner } } - let inner_router = Router::::new(); + let inner_router = Router::::new().route( + "/b", + get(|State(state): State| async move { state.value }), + ); - let router_router = Router::::new() - .state(State { - inner: InnerState {}, - }) - .nest("/", inner_router.map_state(Into::into)); + let app = Router::with_state(AppState { + inner: InnerState { value: "inner" }, + }) + .nest("/a", inner_router.map_state(Into::into)); - todo!(); + let client = TestClient::new(app); + + let res = client.get("/a/b").send().await; + assert_eq!(res.text().await, "inner"); } macro_rules! nested_route_test {