mirror of
https://github.com/tokio-rs/axum.git
synced 2026-08-30 00:00:32 +02:00
checkpoint
This commit is contained in:
@@ -58,6 +58,9 @@ where
|
|||||||
use futures_util::future::FutureExt;
|
use futures_util::future::FutureExt;
|
||||||
|
|
||||||
let handler = self.handler.clone();
|
let handler = self.handler.clone();
|
||||||
|
|
||||||
|
// TODO(david): this is duplicated in `axum/src/routing/mod.rs`
|
||||||
|
// extract into helper function
|
||||||
let State(state) = req
|
let State(state) = req
|
||||||
.extensions()
|
.extensions()
|
||||||
.get::<State<S>>()
|
.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>
|
impl<S, B> MethodRouter<S, B, Infallible, WithState>
|
||||||
where
|
where
|
||||||
B: Send + 'static,
|
B: Send + 'static,
|
||||||
@@ -1181,10 +1201,8 @@ where
|
|||||||
|
|
||||||
if req.extensions().get::<State<S>>().is_none() {
|
if req.extensions().get::<State<S>>().is_none() {
|
||||||
// the `unwrap` is safe because `self.state` is always some if `R = WithState`, which it is
|
// the `unwrap` is safe because `self.state` is always some if `R = WithState`, which it is
|
||||||
let prev = req
|
req.extensions_mut()
|
||||||
.extensions_mut()
|
|
||||||
.insert(State(state.as_ref().unwrap().clone()));
|
.insert(State(state.as_ref().unwrap().clone()));
|
||||||
debug_assert!(prev.is_none());
|
|
||||||
}
|
}
|
||||||
|
|
||||||
call!(req, method, HEAD, head);
|
call!(req, method, HEAD, head);
|
||||||
|
|||||||
+54
-4
@@ -21,7 +21,7 @@ use std::{
|
|||||||
sync::Arc,
|
sync::Arc,
|
||||||
task::{Context, Poll},
|
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_http::map_response_body::MapResponseBodyLayer;
|
||||||
use tower_layer::Layer;
|
use tower_layer::Layer;
|
||||||
use tower_service::Service;
|
use tower_service::Service;
|
||||||
@@ -168,13 +168,63 @@ where
|
|||||||
_marker: PhantomData,
|
_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>
|
pub fn map_state<F, OuterState>(self, f: F) -> Router<OuterState, B, MissingState>
|
||||||
where
|
where
|
||||||
// TODO(david): which Fn?
|
F: Fn(OuterState) -> InnerState + Clone + Send + Sync + 'static,
|
||||||
F: FnOnce(OuterState) -> S,
|
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;
|
let res = client.get("/").send().await;
|
||||||
assert_eq!(res.text().await, "foo");
|
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]
|
#[tokio::test]
|
||||||
async fn nesting_with_different_state() {
|
async fn nesting_with_different_state() {
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
struct State {
|
struct AppState {
|
||||||
inner: InnerState,
|
inner: InnerState,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
struct InnerState {}
|
struct InnerState {
|
||||||
|
value: &'static str,
|
||||||
|
}
|
||||||
|
|
||||||
impl From<State> for InnerState {
|
impl From<AppState> for InnerState {
|
||||||
fn from(state: State) -> Self {
|
fn from(state: AppState) -> Self {
|
||||||
state.inner
|
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()
|
let app = Router::with_state(AppState {
|
||||||
.state(State {
|
inner: InnerState { value: "inner" },
|
||||||
inner: InnerState {},
|
})
|
||||||
})
|
.nest("/a", inner_router.map_state(Into::into));
|
||||||
.nest("/", 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 {
|
macro_rules! nested_route_test {
|
||||||
|
|||||||
Reference in New Issue
Block a user