use axum::extract::Request;
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use std::future::Future;
use std::sync::Arc;
pub fn handler_as_middleware<H, F, R>(
handler: H,
) -> impl Fn(Request, Next) -> std::pin::Pin<Box<dyn std::future::Future<Output = Response> + Send>>
+ Clone
+ Send
+ Sync
+ 'static
where
H: Fn(Request) -> F + Send + Sync + 'static,
F: Future<Output = R> + Send + 'static,
R: IntoResponse + 'static,
{
let handler = Arc::new(handler);
move |req: Request, _next: Next| {
let handler = handler.clone();
Box::pin(async move {
let _ = _next;
handler(req).await.into_response()
})
}
}
pub fn middleware_as_handler<H, F>(
middleware: H,
) -> impl Fn(Request) -> std::pin::Pin<Box<dyn std::future::Future<Output = Response> + Send>>
+ Clone
+ Send
+ Sync
+ 'static
where
H: Fn(Request) -> F + Send + Sync + 'static,
F: Future<Output = Response> + Send + 'static,
{
let middleware = Arc::new(middleware);
move |req: Request| {
let middleware = middleware.clone();
Box::pin(async move { middleware(req).await })
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::http::StatusCode;
use axum::Router;
use http_body_util::BodyExt;
use tower::ServiceExt;
async fn read_body(resp: Response) -> String {
let bytes = resp.into_body().collect().await.unwrap().to_bytes();
String::from_utf8(bytes.to_vec()).unwrap()
}
fn make_request(method: &str, uri: &str) -> Request {
Request::builder()
.method(method)
.uri(uri)
.body(Body::empty())
.unwrap()
}
#[tokio::test]
async fn test_handler_as_middleware_calls_handler_returns_response() {
async fn my_handler(_req: Request) -> &'static str {
"handler called"
}
let app: Router = Router::new()
.route("/", axum::routing::get(|| async { "route" }))
.layer(axum::middleware::from_fn(handler_as_middleware(my_handler)));
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "handler called");
}
#[tokio::test]
async fn test_handler_as_middleware_does_not_call_route_handler() {
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
let next_called = Arc::new(AtomicBool::new(false));
let next_called_for_layer = next_called.clone();
async fn my_handler(_req: Request) -> &'static str {
"handler response"
}
let app: Router = Router::new()
.route(
"/",
axum::routing::get(move || {
let next_called = next_called_for_layer.clone();
async move {
next_called.store(true, Ordering::SeqCst);
"next called"
}
}),
)
.layer(axum::middleware::from_fn(handler_as_middleware(my_handler)));
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "handler response");
assert!(
!next_called.load(Ordering::SeqCst),
"route handler should NOT be called in handler_as_middleware"
);
}
#[tokio::test]
async fn test_handler_as_middleware_with_json_response() {
async fn json_handler(_req: Request) -> axum::Json<serde_json::Value> {
axum::Json(serde_json::json!({"code": 1, "msg": "ok"}))
}
let app: Router = Router::new()
.route("/", axum::routing::get(|| async { "route" }))
.layer(axum::middleware::from_fn(handler_as_middleware(
json_handler,
)));
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = read_body(resp).await;
assert!(body.contains("\"code\":1"));
assert!(body.contains("\"msg\":\"ok\""));
}
#[tokio::test]
async fn test_handler_as_middleware_with_status_code() {
async fn error_handler(_req: Request) -> (StatusCode, &'static str) {
(StatusCode::NOT_FOUND, "not found")
}
let app: Router = Router::new()
.route("/", axum::routing::get(|| async { "route" }))
.layer(axum::middleware::from_fn(handler_as_middleware(
error_handler,
)));
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
let body = read_body(resp).await;
assert_eq!(body, "not found");
}
#[tokio::test]
async fn test_handler_as_middleware_preserves_request_method() {
async fn method_handler(req: Request) -> String {
req.method().to_string()
}
let app: Router = Router::new()
.route("/", axum::routing::get(|| async { "route" }))
.layer(axum::middleware::from_fn(handler_as_middleware(
method_handler,
)));
let resp = app.oneshot(make_request("POST", "/")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "POST");
}
#[tokio::test]
async fn test_handler_as_middleware_preserves_request_uri() {
async fn uri_handler(req: Request) -> String {
req.uri().to_string()
}
let app: Router = Router::new()
.route("/api/test", axum::routing::get(|| async { "route" }))
.layer(axum::middleware::from_fn(handler_as_middleware(
uri_handler,
)));
let resp = app.oneshot(make_request("GET", "/api/test")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "/api/test");
}
#[tokio::test]
async fn test_handler_as_middleware_with_request_body() {
async fn body_handler(req: Request) -> String {
let bytes = req.into_body().collect().await.unwrap().to_bytes();
String::from_utf8(bytes.to_vec()).unwrap()
}
let app: Router = Router::new()
.route("/", axum::routing::get(|| async { "route" }))
.layer(axum::middleware::from_fn(handler_as_middleware(
body_handler,
)));
let req = Request::builder()
.method("POST")
.uri("/")
.body(Body::from("hello body"))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "hello body");
}
#[tokio::test]
async fn test_middleware_as_handler_calls_middleware() {
async fn my_middleware(_req: Request) -> Response {
Response::new(Body::from("middleware as handler"))
}
let app: Router = Router::new().route(
"/",
axum::routing::get(middleware_as_handler(my_middleware)),
);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "middleware as handler");
}
#[tokio::test]
async fn test_middleware_as_handler_with_status_code() {
async fn error_middleware(_req: Request) -> Response {
let mut resp = Response::new(Body::from("forbidden"));
*resp.status_mut() = StatusCode::FORBIDDEN;
resp
}
let app: Router = Router::new().route(
"/",
axum::routing::get(middleware_as_handler(error_middleware)),
);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
let body = read_body(resp).await;
assert_eq!(body, "forbidden");
}
#[tokio::test]
async fn test_middleware_as_handler_preserves_request_method() {
async fn method_middleware(req: Request) -> Response {
Response::new(Body::from(req.method().to_string()))
}
let app: Router = Router::new().route(
"/",
axum::routing::get(middleware_as_handler(method_middleware)),
);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "GET");
}
#[tokio::test]
async fn test_middleware_as_handler_with_json_response() {
async fn json_middleware(_req: Request) -> Response {
let body = serde_json::json!({"code": 1, "msg": "ok"});
Response::new(Body::from(body.to_string()))
}
let app: Router = Router::new().route(
"/",
axum::routing::get(middleware_as_handler(json_middleware)),
);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
let body = read_body(resp).await;
assert!(body.contains("\"code\":1"));
assert!(body.contains("\"msg\":\"ok\""));
}
#[tokio::test]
async fn test_handler_and_middleware_interchangeable() {
async fn my_handler(_req: Request) -> &'static str {
"handler"
}
let app: Router = Router::new()
.route("/", axum::routing::get(|| async { "route" }))
.layer(axum::middleware::from_fn(handler_as_middleware(my_handler)));
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "handler"); }
#[tokio::test]
async fn test_middleware_as_handler_used_as_route_handler() {
async fn terminal_middleware(_req: Request) -> Response {
Response::new(Body::from("terminal middleware as handler"))
}
let app: Router = Router::new().route(
"/api/test",
axum::routing::get(middleware_as_handler(terminal_middleware)),
);
let resp = app.oneshot(make_request("GET", "/api/test")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "terminal middleware as handler");
}
#[tokio::test]
async fn test_handler_as_middleware_layered_before_other_middleware() {
async fn first_middleware(req: Request, next: Next) -> Response {
next.run(req).await
}
async fn my_handler(_req: Request) -> String {
"handler called at order 1".to_string()
}
let app: Router = Router::new()
.route("/", axum::routing::get(|| async { "route" }))
.layer(axum::middleware::from_fn(handler_as_middleware(my_handler)))
.layer(axum::middleware::from_fn(first_middleware));
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "handler called at order 1");
}
}