mirror of
https://github.com/tokio-rs/axum.git
synced 2026-08-29 00:00:18 +02:00
checkpoint
This commit is contained in:
@@ -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::<State<S>>()
|
||||
|
||||
@@ -672,6 +672,26 @@ impl<S, B, R> MethodRouter<S, B, Infallible, R> {
|
||||
}
|
||||
}
|
||||
|
||||
impl<S, B> MethodRouter<S, B, Infallible, MissingState> {
|
||||
pub(crate) fn change_state<S2>(self) -> MethodRouter<S2, B, Infallible, MissingState> {
|
||||
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<S, B> MethodRouter<S, B, Infallible, WithState>
|
||||
where
|
||||
B: Send + 'static,
|
||||
@@ -1181,10 +1201,8 @@ where
|
||||
|
||||
if req.extensions().get::<State<S>>().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);
|
||||
|
||||
+54
-4
@@ -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<InnerState, B> Router<InnerState, B, MissingState>
|
||||
where
|
||||
B: HttpBody + Send + 'static,
|
||||
{
|
||||
pub fn map_state<F, OuterState>(self, f: F) -> Router<OuterState, B, MissingState>
|
||||
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 `<Router as Service>::call`, so its
|
||||
// safe to ignore that it hasn't been provided yet
|
||||
Endpoint::MethodRouter(method_router.change_state::<OuterState>())
|
||||
}
|
||||
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::<State<OuterState>>()
|
||||
.unwrap_or_else(|| {
|
||||
panic!(
|
||||
"no state of type `{}` was found. Please file an issue",
|
||||
std::any::type_name::<State<OuterState>>()
|
||||
)
|
||||
})
|
||||
.clone();
|
||||
|
||||
let inner_state = f(outer_state);
|
||||
|
||||
req.extensions_mut().insert(State(inner_state));
|
||||
|
||||
req
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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!()
|
||||
}
|
||||
|
||||
@@ -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<State> for InnerState {
|
||||
fn from(state: State) -> Self {
|
||||
impl From<AppState> for InnerState {
|
||||
fn from(state: AppState) -> Self {
|
||||
state.inner
|
||||
}
|
||||
}
|
||||
|
||||
let inner_router = Router::<InnerState, Body, _>::new();
|
||||
let inner_router = Router::<InnerState, Body, _>::new().route(
|
||||
"/b",
|
||||
get(|State(state): State<InnerState>| async move { state.value }),
|
||||
);
|
||||
|
||||
let router_router = Router::<State, Body, _>::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 {
|
||||
|
||||
Reference in New Issue
Block a user