From 5deef529ae3aa33d42c7ec3d95aabf2d9402fd7f Mon Sep 17 00:00:00 2001 From: David Pedersen Date: Sun, 3 Jul 2022 16:57:44 +0200 Subject: [PATCH] checkpoint --- axum/src/extract/state.rs | 1 + axum/src/middleware/from_extractor.rs | 5 --- axum/src/middleware/from_fn.rs | 52 +++++++++++++++++---------- 3 files changed, 35 insertions(+), 23 deletions(-) diff --git a/axum/src/extract/state.rs b/axum/src/extract/state.rs index 0801a113..9c2c5516 100644 --- a/axum/src/extract/state.rs +++ b/axum/src/extract/state.rs @@ -3,6 +3,7 @@ use async_trait::async_trait; use std::convert::Infallible; /// TODO(david): docs +// TODO(david): document how to extract this from middleware #[derive(Clone, Copy, Debug, Default)] pub struct State(pub S); diff --git a/axum/src/middleware/from_extractor.rs b/axum/src/middleware/from_extractor.rs index ab9b0205..52db8733 100644 --- a/axum/src/middleware/from_extractor.rs +++ b/axum/src/middleware/from_extractor.rs @@ -326,9 +326,4 @@ mod tests { .await; assert_eq!(res.status(), StatusCode::OK); } - - #[test] - fn extracting_state() { - todo!() - } } diff --git a/axum/src/middleware/from_fn.rs b/axum/src/middleware/from_fn.rs index 3b6f6c94..dfefcebe 100644 --- a/axum/src/middleware/from_fn.rs +++ b/axum/src/middleware/from_fn.rs @@ -259,7 +259,7 @@ macro_rules! impl_service { impl Service> for FromFn where F: FnMut($($ty),*, Next) -> Fut + Clone + Send + 'static, - $( $ty: FromRequest + Send, )* + $( $ty: FromRequest<(), ReqBody> + Send, )* Fut: Future + Send + 'static, Out: IntoResponse + 'static, S: Service, Response = Response, Error = Infallible> @@ -286,7 +286,7 @@ macro_rules! impl_service { let mut f = self.f.clone(); let future = Box::pin(async move { - let mut parts = RequestParts::new(req); + let mut parts = RequestParts::new((), req); $( let $ty = match $ty::from_request(&mut parts).await { Ok(value) => value, @@ -370,9 +370,8 @@ impl fmt::Debug for ResponseFuture { #[cfg(test)] mod tests { use super::*; - use crate::{body::Empty, routing::get, Router}; + use crate::{extract::State, routing::get, test_helpers::TestClient, Router}; use http::{HeaderMap, StatusCode}; - use tower::ServiceExt; #[tokio::test] async fn basic() { @@ -391,22 +390,39 @@ mod tests { .route("/", get(handle)) .layer(from_fn(insert_header)); - let res = app - .oneshot( - Request::builder() - .uri("/") - .body(body::boxed(Empty::new())) - .unwrap(), - ) - .await - .unwrap(); + let client = TestClient::new(app); + + let res = client.get("/").send().await; + assert_eq!(res.status(), StatusCode::OK); - let body = hyper::body::to_bytes(res).await.unwrap(); - assert_eq!(&body[..], b"ok"); + assert_eq!(res.text().await, "ok"); } - #[test] - fn extracting_state() { - todo!() + #[tokio::test] + async fn extracting_state() { + async fn access_state(req: Request, next: Next) -> impl IntoResponse { + let State(state) = req.extensions().get::>().unwrap().clone(); + state.value + } + + async fn handle() { + panic!() + } + + #[derive(Clone)] + struct AppState { + value: &'static str, + } + + let app = Router::with_state(AppState { value: "foo" }) + .route("/", get(handle)) + .layer(from_fn(access_state)); + + let client = TestClient::new(app); + + let res = client.get("/").send().await; + + assert_eq!(res.status(), StatusCode::OK); + assert_eq!(res.text().await, "foo"); } }