Implement tower::Layer for Extension (#801)

* Implement `tower::Layer` for `Extension`

* changelog
This commit is contained in:
David Pedersen
2022-03-01 00:39:22 +01:00
committed by GitHub
parent 0d05b5e31f
commit a2b568c7c1
17 changed files with 51 additions and 43 deletions
+2 -2
View File
@@ -45,8 +45,8 @@ where
/// [request extensions]: https://docs.rs/http/latest/http/struct.Extensions.html
#[derive(Clone, Copy, Debug)]
pub struct AddExtension<S, T> {
inner: S,
value: T,
pub(crate) inner: S,
pub(crate) value: T,
}
impl<S, T> AddExtension<S, T> {
+2 -2
View File
@@ -5,7 +5,7 @@
//! [`Router::into_make_service_with_connect_info`]: crate::routing::Router::into_make_service_with_connect_info
use super::{Extension, FromRequest, RequestParts};
use crate::{AddExtension, AddExtensionLayer};
use crate::AddExtension;
use async_trait::async_trait;
use hyper::server::conn::AddrStream;
use std::{
@@ -104,7 +104,7 @@ where
fn call(&mut self, target: T) -> Self::Future {
let connect_info = ConnectInfo(C::connect_info(target));
let svc = AddExtensionLayer::new(connect_info).layer(self.svc.clone());
let svc = Extension(connect_info).layer(self.svc.clone());
ResponseFuture::new(ready(Ok(svc)))
}
}
+16 -3
View File
@@ -12,7 +12,6 @@ use std::ops::Deref;
///
/// ```rust,no_run
/// use axum::{
/// AddExtensionLayer,
/// extract::Extension,
/// routing::get,
/// Router,
@@ -33,7 +32,7 @@ use std::ops::Deref;
/// let app = Router::new().route("/", get(handler))
/// // Add middleware that inserts the state into all incoming request's
/// // extensions.
/// .layer(AddExtensionLayer::new(state));
/// .layer(Extension(state));
/// # async {
/// # axum::Server::bind(&"".parse().unwrap()).serve(app.into_make_service()).await.unwrap();
/// # };
@@ -58,7 +57,7 @@ where
.get::<T>()
.ok_or_else(|| {
MissingExtension::from_err(format!(
"Extension of type `{}` was not found. Perhaps you forgot to add it? See `axum::AddExtensionLayer`.",
"Extension of type `{}` was not found. Perhaps you forgot to add it? See `axum::extract::Extension`.",
std::any::type_name::<T>()
))
})
@@ -95,3 +94,17 @@ where
res
}
}
impl<S, T> tower_layer::Layer<S> for Extension<T>
where
T: Clone + Send + Sync + 'static,
{
type Service = crate::AddExtension<S, T>;
fn layer(&self, inner: S) -> Self::Service {
crate::AddExtension {
inner,
value: self.0.clone(),
}
}
}
+3 -6
View File
@@ -212,9 +212,10 @@ where
mod tests {
use crate::{
body::Body,
extract::Extension,
routing::{get, post},
test_helpers::*,
AddExtensionLayer, Router,
Router,
};
use http::{Method, Request, StatusCode};
@@ -247,11 +248,7 @@ mod tests {
parts.extensions.get::<Ext>().unwrap();
}
let client = TestClient::new(
Router::new()
.route("/", get(handler))
.layer(AddExtensionLayer::new(Ext)),
);
let client = TestClient::new(Router::new().route("/", get(handler)).layer(Extension(Ext)));
let res = client.get("/").header("x-foo", "123").send().await;
assert_eq!(res.status(), StatusCode::OK);
+3 -6
View File
@@ -173,13 +173,11 @@
//!
//! ## Using request extensions
//!
//! The easiest way to extract state in handlers is using [`AddExtension`]
//! middleware (applied with [`AddExtensionLayer`]) and the
//! [`Extension`](crate::extract::Extension) extractor:
//! The easiest way to extract state in handlers is using [`Extension`](crate::extract::Extension)
//! as layer and extractor:
//!
//! ```rust,no_run
//! use axum::{
//! AddExtensionLayer,
//! extract::Extension,
//! routing::get,
//! Router,
@@ -194,7 +192,7 @@
//!
//! let app = Router::new()
//! .route("/", get(handler))
//! .layer(AddExtensionLayer::new(shared_state));
//! .layer(Extension(shared_state));
//!
//! async fn handler(
//! Extension(state): Extension<Arc<State>>,
@@ -217,7 +215,6 @@
//!
//! ```rust,no_run
//! use axum::{
//! AddExtensionLayer,
//! Json,
//! extract::{Extension, Path},
//! routing::{get, post},
+2 -2
View File
@@ -102,11 +102,11 @@ use tower_service::Service;
/// ```rust
/// use axum::{
/// Router,
/// extract::Extension,
/// http::{Request, StatusCode},
/// routing::get,
/// response::IntoResponse,
/// middleware::{self, Next},
/// AddExtensionLayer,
/// };
/// use tower::ServiceBuilder;
///
@@ -129,7 +129,7 @@ use tower_service::Service;
/// .route("/", get(|| async { /* ... */ }))
/// .layer(
/// ServiceBuilder::new()
/// .layer(AddExtensionLayer::new(state))
/// .layer(Extension(state))
/// .layer(middleware::from_fn(my_middleware)),
/// );
/// # let app: Router = app;