mirror of
https://github.com/tokio-rs/axum.git
synced 2026-09-03 00:00:07 +02:00
Implement tower::Layer for Extension (#801)
* Implement `tower::Layer` for `Extension` * changelog
This commit is contained in:
@@ -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> {
|
||||
|
||||
@@ -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)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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},
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user