mirror of
https://github.com/tokio-rs/axum.git
synced 2026-08-29 00:00:18 +02:00
Add method filtering to route_with_tsr (#3586)
This commit is contained in:
@@ -9,8 +9,20 @@ and this project adheres to [Semantic Versioning].
|
|||||||
|
|
||||||
- **breaking:** Remove the deprecated `Host`, `Scheme` and `OptionalPath`
|
- **breaking:** Remove the deprecated `Host`, `Scheme` and `OptionalPath`
|
||||||
extractors ([#3599])
|
extractors ([#3599])
|
||||||
|
- **breaking:** Change `routing::RouterExt::route_with_tsr` to only redirect
|
||||||
|
the HTTP methods that the supplied `MethodRouter` handles. This allows the
|
||||||
|
following pattern which lead to a panic before because the two
|
||||||
|
`route_with_tsr` calls would both attempt to register a method-independent
|
||||||
|
redirect ([#3586]):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
Router::new()
|
||||||
|
.route_with_tsr("/path", get(/* handler */))
|
||||||
|
.route_with_tsr("/path", post(/* handler */))
|
||||||
|
```
|
||||||
|
|
||||||
[#3599]: https://github.com/tokio-rs/axum/pull/3599
|
[#3599]: https://github.com/tokio-rs/axum/pull/3599
|
||||||
|
[#3586]: https://github.com/tokio-rs/axum/pull/3586
|
||||||
|
|
||||||
# 0.12.5
|
# 0.12.5
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
use axum::{
|
use axum::{
|
||||||
extract::{OriginalUri, Request},
|
extract::{OriginalUri, Request},
|
||||||
response::{IntoResponse, Redirect, Response},
|
response::{IntoResponse, Redirect, Response},
|
||||||
routing::{any, MethodRouter},
|
routing::{any, on, MethodFilter, MethodRouter},
|
||||||
Router,
|
Router,
|
||||||
};
|
};
|
||||||
use http::{uri::PathAndQuery, StatusCode, Uri};
|
use http::{uri::PathAndQuery, StatusCode, Uri};
|
||||||
@@ -336,8 +336,9 @@ where
|
|||||||
Self: Sized,
|
Self: Sized,
|
||||||
{
|
{
|
||||||
validate_tsr_path(path);
|
validate_tsr_path(path);
|
||||||
|
let method_filter = method_router.method_filter();
|
||||||
self = self.route(path, method_router);
|
self = self.route(path, method_router);
|
||||||
add_tsr_redirect_route(self, path)
|
add_tsr_redirect_route(self, path, method_filter)
|
||||||
}
|
}
|
||||||
|
|
||||||
#[track_caller]
|
#[track_caller]
|
||||||
@@ -350,7 +351,7 @@ where
|
|||||||
{
|
{
|
||||||
validate_tsr_path(path);
|
validate_tsr_path(path);
|
||||||
self = self.route_service(path, service);
|
self = self.route_service(path, service);
|
||||||
add_tsr_redirect_route(self, path)
|
add_tsr_redirect_route(self, path, None)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -361,7 +362,11 @@ fn validate_tsr_path(path: &str) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn add_tsr_redirect_route<S>(router: Router<S>, path: &str) -> Router<S>
|
fn add_tsr_redirect_route<S>(
|
||||||
|
router: Router<S>,
|
||||||
|
path: &str,
|
||||||
|
method_filter: Option<MethodFilter>,
|
||||||
|
) -> Router<S>
|
||||||
where
|
where
|
||||||
S: Clone + Send + Sync + 'static,
|
S: Clone + Send + Sync + 'static,
|
||||||
{
|
{
|
||||||
@@ -379,11 +384,21 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(path_without_trailing_slash) = path.strip_suffix('/') {
|
let _slot;
|
||||||
router.route(path_without_trailing_slash, any(redirect_handler))
|
let redirect_path = if let Some(without_slash) = path.strip_suffix('/') {
|
||||||
|
without_slash
|
||||||
} else {
|
} else {
|
||||||
router.route(&format!("{path}/"), any(redirect_handler))
|
// FIXME: Can return `&format!(...)` directly when MSRV is updated
|
||||||
}
|
_slot = format!("{path}/");
|
||||||
|
&_slot
|
||||||
|
};
|
||||||
|
|
||||||
|
let method_router = match method_filter {
|
||||||
|
Some(f) => on(f, redirect_handler),
|
||||||
|
None => any(redirect_handler),
|
||||||
|
};
|
||||||
|
|
||||||
|
router.route(redirect_path, method_router)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Map the path of a `Uri`.
|
/// Map the path of a `Uri`.
|
||||||
@@ -417,7 +432,10 @@ mod sealed {
|
|||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use crate::test_helpers::*;
|
use crate::test_helpers::*;
|
||||||
use axum::{extract::Path, routing::get};
|
use axum::{
|
||||||
|
extract::Path,
|
||||||
|
routing::{get, post},
|
||||||
|
};
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_tsr() {
|
async fn test_tsr() {
|
||||||
@@ -500,6 +518,13 @@ mod tests {
|
|||||||
assert_eq!(res.headers()["location"], "/neko/nyan/");
|
assert_eq!(res.headers()["location"], "/neko/nyan/");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn tsr_independent_route_registration() {
|
||||||
|
let _: Router = Router::new()
|
||||||
|
.route_with_tsr("/x", get(|| async {}))
|
||||||
|
.route_with_tsr("/x", post(|| async {}));
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
#[should_panic = "Cannot add a trailing slash redirect route for `/`"]
|
#[should_panic = "Cannot add a trailing slash redirect route for `/`"]
|
||||||
fn tsr_at_root() {
|
fn tsr_at_root() {
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
(never returned `Err`) and be an uninhabited type if `with_graceful_shutdown` is not used
|
(never returned `Err`) and be an uninhabited type if `with_graceful_shutdown` is not used
|
||||||
(because it was already never terminating if that method wasn't used) ([#3601])
|
(because it was already never terminating if that method wasn't used) ([#3601])
|
||||||
- **added:** New `ListenerExt::limit_connections` allows limiting concurrent `axum::serve` connections ([#3489])
|
- **added:** New `ListenerExt::limit_connections` allows limiting concurrent `axum::serve` connections ([#3489])
|
||||||
|
- **added:** `MethodRouter::method_filter` ([#3586])
|
||||||
- **changed:** `serve` has an additional generic argument and can now work with any response body
|
- **changed:** `serve` has an additional generic argument and can now work with any response body
|
||||||
type, not just `axum::body::Body` ([#3205])
|
type, not just `axum::body::Body` ([#3205])
|
||||||
|
|
||||||
@@ -24,6 +25,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
[#3478]: https://github.com/tokio-rs/axum/pull/3478
|
[#3478]: https://github.com/tokio-rs/axum/pull/3478
|
||||||
[#3601]: https://github.com/tokio-rs/axum/pull/3601
|
[#3601]: https://github.com/tokio-rs/axum/pull/3601
|
||||||
[#3489]: https://github.com/tokio-rs/axum/pull/3489
|
[#3489]: https://github.com/tokio-rs/axum/pull/3489
|
||||||
|
[#3586]: https://github.com/tokio-rs/axum/pull/3586
|
||||||
|
|
||||||
# 0.8.8
|
# 0.8.8
|
||||||
|
|
||||||
|
|||||||
@@ -660,6 +660,54 @@ where
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Get a [`MethodFilter`] for the methods that this `MethodRouter` has
|
||||||
|
/// custom code for.
|
||||||
|
///
|
||||||
|
/// Note that `MethodRouter`'s [`Service`] implementation never fails (it
|
||||||
|
/// always creates an HTTP response) based on which HTTP method was used.
|
||||||
|
/// However, the information which methods have the default behavior of
|
||||||
|
/// returning HTTP 405 is stored, and can be queried with this method.
|
||||||
|
///
|
||||||
|
/// Returns `None` if the `MethodRouter` was constructed with [`any`] or
|
||||||
|
/// has had a [`fallback`][Self::fallback] set.
|
||||||
|
pub fn method_filter(&self) -> Option<MethodFilter> {
|
||||||
|
let Self {
|
||||||
|
get,
|
||||||
|
head,
|
||||||
|
delete,
|
||||||
|
options,
|
||||||
|
patch,
|
||||||
|
post,
|
||||||
|
put,
|
||||||
|
trace,
|
||||||
|
connect,
|
||||||
|
fallback,
|
||||||
|
allow_header: _,
|
||||||
|
} = self;
|
||||||
|
|
||||||
|
if !fallback.is_default() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
|
||||||
|
let filter = [
|
||||||
|
(get, MethodFilter::GET),
|
||||||
|
(head, MethodFilter::HEAD),
|
||||||
|
(delete, MethodFilter::DELETE),
|
||||||
|
(options, MethodFilter::OPTIONS),
|
||||||
|
(patch, MethodFilter::PATCH),
|
||||||
|
(post, MethodFilter::POST),
|
||||||
|
(put, MethodFilter::PUT),
|
||||||
|
(trace, MethodFilter::TRACE),
|
||||||
|
(connect, MethodFilter::CONNECT),
|
||||||
|
]
|
||||||
|
.into_iter()
|
||||||
|
.filter_map(|(ep, f)| ep.is_some().then_some(f))
|
||||||
|
.reduce(MethodFilter::or)
|
||||||
|
.expect("can't create a MethodRouter with all-default handlers");
|
||||||
|
|
||||||
|
Some(filter)
|
||||||
|
}
|
||||||
|
|
||||||
/// Add a fallback [`Handler`] if no custom one has been provided.
|
/// Add a fallback [`Handler`] if no custom one has been provided.
|
||||||
pub(crate) fn default_fallback<H, T>(self, handler: H) -> Self
|
pub(crate) fn default_fallback<H, T>(self, handler: H) -> Self
|
||||||
where
|
where
|
||||||
@@ -839,7 +887,7 @@ where
|
|||||||
panic!(
|
panic!(
|
||||||
"Overlapping method route. Cannot add two method routes that both handle \
|
"Overlapping method route. Cannot add two method routes that both handle \
|
||||||
`{method_name}`",
|
`{method_name}`",
|
||||||
)
|
);
|
||||||
}
|
}
|
||||||
*out = endpoint.clone();
|
*out = endpoint.clone();
|
||||||
for method in methods {
|
for method in methods {
|
||||||
@@ -1614,6 +1662,25 @@ mod tests {
|
|||||||
assert_eq!(text, "state");
|
assert_eq!(text, "state");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn method_filter() {
|
||||||
|
let router: MethodRouter = get(|| async {});
|
||||||
|
assert_eq!(router.method_filter(), Some(MethodFilter::GET));
|
||||||
|
|
||||||
|
let router: MethodRouter = get(|| async {}).head(|| async {}).post(|| async {});
|
||||||
|
assert_eq!(
|
||||||
|
router.method_filter(),
|
||||||
|
Some(
|
||||||
|
MethodFilter::GET
|
||||||
|
.or(MethodFilter::HEAD)
|
||||||
|
.or(MethodFilter::POST)
|
||||||
|
)
|
||||||
|
);
|
||||||
|
|
||||||
|
let router: MethodRouter = any(|| async {});
|
||||||
|
assert_eq!(router.method_filter(), None);
|
||||||
|
}
|
||||||
|
|
||||||
async fn call<S>(method: Method, svc: &mut S) -> (StatusCode, HeaderMap, String)
|
async fn call<S>(method: Method, svc: &mut S) -> (StatusCode, HeaderMap, String)
|
||||||
where
|
where
|
||||||
S: Service<Request, Error = Infallible>,
|
S: Service<Request, Error = Infallible>,
|
||||||
|
|||||||
@@ -740,6 +740,10 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn is_default(&self) -> bool {
|
||||||
|
matches!(self, Self::Default(..))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<S, E> Clone for Fallback<S, E> {
|
impl<S, E> Clone for Fallback<S, E> {
|
||||||
|
|||||||
Reference in New Issue
Block a user