checkpoint

This commit is contained in:
David Pedersen
2022-07-03 16:21:43 +02:00
parent a5b6b94530
commit 90e9b34736
5 changed files with 101 additions and 18 deletions
@@ -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>>()
+21 -3
View File
@@ -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
View File
@@ -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
}))
}
}
+5
View File
@@ -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!()
}
+18 -11
View File
@@ -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 {