use std::{any::TypeId, future::Future, pin::Pin};
use crate::{Request, Response};
pub trait Middleware: Send + Sync + 'static {
async fn handle(&self, request: Request, next: Next<'_>) -> Response;
}
pub(crate) type MiddlewareFuture<'request> = Pin<Box<dyn Future<Output = Response> + 'request>>;
pub(crate) struct MiddlewareEntry {
type_id: TypeId,
service: Box<dyn MiddlewareService>,
}
impl MiddlewareEntry {
pub(crate) fn new<M: Middleware>(middleware: M) -> Self {
Self {
type_id: TypeId::of::<M>(),
service: Box::new(middleware),
}
}
pub(crate) fn type_id(&self) -> TypeId {
self.type_id
}
}
pub(crate) trait MiddlewareService: Send + Sync {
fn call<'request>(
&'request self,
request: Request,
next: Next<'request>,
) -> MiddlewareFuture<'request>;
}
impl<M: Middleware> MiddlewareService for M {
fn call<'request>(
&'request self,
request: Request,
next: Next<'request>,
) -> MiddlewareFuture<'request> {
Box::pin(self.handle(request, next))
}
}
pub struct Next<'next> {
middlewares: &'next [&'next MiddlewareEntry],
terminal: &'next dyn MiddlewareTerminal,
}
impl Next<'_> {
pub async fn run(self, request: Request) -> Response {
let Some((middleware, remaining)) = self.middlewares.split_first() else {
return self.terminal.call(request).await;
};
middleware
.service
.call(
request,
Next {
middlewares: remaining,
terminal: self.terminal,
},
)
.await
}
}
pub(crate) trait MiddlewareTerminal {
fn call(&self, request: Request) -> MiddlewareFuture<'_>;
}
pub(crate) async fn run(
middlewares: &[&MiddlewareEntry],
terminal: &dyn MiddlewareTerminal,
request: Request,
) -> Response {
Next {
middlewares,
terminal,
}
.run(request)
.await
}
#[cfg(test)]
mod tests {
use std::{
convert::Infallible,
future::Future,
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
},
task::{Context, Poll, Waker},
};
use crate::{
Config, FromRequest, Headers, Method, Middleware, Next, Path, Request, RequestStream,
Response, RouteMethods, Router, StreamError,
};
struct EmptyStream;
impl RequestStream for EmptyStream {
fn poll_next(
&mut self,
_context: &mut Context<'_>,
) -> Poll<Option<Result<(), StreamError>>> {
Poll::Ready(None)
}
fn chunk(&self) -> &[u8] {
&[]
}
}
struct ParentMiddleware(Arc<Mutex<Vec<&'static str>>>);
struct ChildMiddleware(Arc<Mutex<Vec<&'static str>>>);
struct RouteMiddleware(Arc<Mutex<Vec<&'static str>>>);
macro_rules! record_middleware {
($middleware:ident, $before:literal, $after:literal) => {
impl Middleware for $middleware {
async fn handle(&self, request: Request, next: Next<'_>) -> Response {
self.0.lock().unwrap().push($before);
let response = next.run(request).await;
self.0.lock().unwrap().push($after);
response
}
}
};
}
record_middleware!(ParentMiddleware, "parent:before", "parent:after");
record_middleware!(ChildMiddleware, "child:before", "child:after");
record_middleware!(RouteMiddleware, "route:before", "route:after");
struct Authentication(Arc<AtomicUsize>);
struct InjectHeader;
struct InjectedHeader(String);
struct RewriteRequest;
struct RequestTarget(String);
impl Middleware for Authentication {
async fn handle(&self, request: Request, next: Next<'_>) -> Response {
self.0.fetch_add(1, Ordering::Relaxed);
next.run(request).await
}
}
impl Middleware for InjectHeader {
async fn handle(&self, mut request: Request, next: Next<'_>) -> Response {
request.headers.set("X-Injected", "middleware").unwrap();
next.run(request).await
}
}
impl Middleware for RewriteRequest {
async fn handle(&self, mut request: Request, next: Next<'_>) -> Response {
request.method = Method::PATCH;
request.path = "/rewritten".to_owned();
request.query = Some("source=middleware".to_owned());
next.run(request).await
}
}
impl<'request> FromRequest<(&'request Request, &'request [u8])> for InjectedHeader {
type Error = Infallible;
async fn from_request(
input: (&'request Request, &'request [u8]),
) -> Result<Self, Self::Error> {
let value = input
.0
.headers
.get("X-Injected")
.and_then(|value| std::str::from_utf8(value).ok())
.unwrap_or_default()
.to_owned();
Ok(Self(value))
}
}
impl<'request> FromRequest<(&'request Request, &'request [u8])> for RequestTarget {
type Error = Infallible;
async fn from_request(
input: (&'request Request, &'request [u8]),
) -> Result<Self, Self::Error> {
Ok(Self(format!(
"{} {}?{}",
input.0.method.as_str(),
input.0.path,
input.0.query.as_deref().unwrap_or_default(),
)))
}
}
fn request(path: &str) -> Request {
Request::from_parts(
Method::GET,
path,
None,
Headers::new(),
Box::new(EmptyStream),
)
}
fn block_on<F: Future>(future: F) -> F::Output {
let mut future = std::pin::pin!(future);
let waker = Waker::noop();
let mut context = Context::from_waker(waker);
loop {
match future.as_mut().poll(&mut context) {
Poll::Ready(output) => return output,
Poll::Pending => std::thread::yield_now(),
}
}
}
#[test]
fn runs_parent_child_and_route_middleware_in_scope_order() {
let events = Arc::new(Mutex::new(Vec::new()));
let handler_events = Arc::clone(&events);
let child = Router::new(
Config::new().prefix("/v1"),
"/health"
.GET(move || {
let events = Arc::clone(&handler_events);
async move {
events.lock().unwrap().push("handler");
"ok"
}
})
.middleware(RouteMiddleware(Arc::clone(&events))),
)
.middleware(ChildMiddleware(Arc::clone(&events)))
.at("/mounted");
let router = Router::new(Config::new().prefix("/parent"), ())
.middleware(ParentMiddleware(Arc::clone(&events)))
.route(child);
let response = block_on(router.handle(request("/parent/mounted/v1/health")));
assert_eq!(response.body(), b"ok");
assert_eq!(
events.lock().unwrap().as_slice(),
[
"parent:before",
"child:before",
"route:before",
"handler",
"route:after",
"child:after",
"parent:after",
],
);
}
#[test]
fn child_middleware_only_runs_inside_the_child_prefix() {
let calls = Arc::new(AtomicUsize::new(0));
let child = Router::new(
Config::new().prefix("/api"),
"/inside".GET(|| async { "in" }),
)
.middleware(Authentication(Arc::clone(&calls)));
let router = Router::new(Config::new(), "/outside".GET(|| async { "out" })).route(child);
assert_eq!(block_on(router.handle(request("/outside"))).body(), b"out",);
assert_eq!(calls.load(Ordering::Relaxed), 0);
assert_eq!(
block_on(router.handle(request("/api/inside"))).body(),
b"in",
);
assert_eq!(calls.load(Ordering::Relaxed), 1);
assert_eq!(
block_on(router.handle(request("/api/missing"))).status(),
404,
);
assert_eq!(calls.load(Ordering::Relaxed), 2);
assert_eq!(block_on(router.handle(request("/missing"))).status(), 404,);
assert_eq!(calls.load(Ordering::Relaxed), 2);
}
#[test]
fn middleware_can_modify_request_headers_before_extraction() {
let router = Router::new(
Config::new(),
"/header".GET(|InjectedHeader(value): InjectedHeader| async move { value }),
)
.middleware(InjectHeader);
let response = block_on(router.handle(request("/header")));
assert_eq!(response.body(), b"middleware");
}
#[test]
fn middleware_can_modify_the_request_target_before_extraction() {
let router = Router::new(
Config::new(),
"/original".GET(|RequestTarget(value): RequestTarget| async move { value }),
)
.middleware(RewriteRequest);
let response = block_on(router.handle(request("/original")));
assert_eq!(response.body(), b"PATCH /rewritten?source=middleware");
}
#[test]
fn route_can_exclude_inherited_middleware_by_type() {
let calls = Arc::new(AtomicUsize::new(0));
let router = Router::new(
Config::new(),
(
"/private".GET(|| async { "private" }),
"/public"
.GET(|| async { "public" })
.without_middleware::<Authentication>(),
),
)
.middleware(Authentication(Arc::clone(&calls)));
assert_eq!(
block_on(router.handle(request("/private"))).body(),
b"private",
);
assert_eq!(
block_on(router.handle(request("/public"))).body(),
b"public",
);
assert_eq!(calls.load(Ordering::Relaxed), 1);
}
#[derive(crate::Schema)]
#[schema(rename_all = "camelCase")]
struct ProjectPath {
project_id: String,
}
struct ReadPath;
impl Middleware for ReadPath {
async fn handle(&self, request: Request, next: Next<'_>) -> Response {
let Path(path) = Path::<ProjectPath>::from_request((&request, &[][..]))
.await
.unwrap_or_else(|_| panic!("middleware must receive path parameters"));
let mut response = next.run(request).await;
response.set_header("X-Project", path.project_id);
response
}
}
#[test]
fn dynamic_prefix_params_are_available_to_middleware_and_handler() {
let child = Router::new(
Config::new().prefix("/project/:projectId"),
"/".GET(|Path(path): Path<ProjectPath>| async move { path.project_id }),
)
.middleware(ReadPath)
.at("/api");
let router = Router::new(Config::new(), child);
for method in [Method::GET, Method::HEAD] {
let mut input = request("/api/project/abc123");
input.method = method.clone();
let response = block_on(router.handle(input));
let (status, headers, body) = response.into_parts();
assert_eq!(status, 200);
assert_eq!(headers.get("x-project"), Some(b"abc123".as_slice()));
assert_eq!(
body.buffered(),
Some(if method == Method::HEAD {
b"".as_slice()
} else {
b"abc123".as_slice()
})
);
}
assert_eq!(
block_on(router.handle(request("/api/projects/abc123"))).status(),
404
);
}
struct ObserveParams;
#[test]
fn sibling_routers_do_not_inherit_each_others_middleware() {
let events = Arc::new(Mutex::new(Vec::new()));
let router = Router::new(
Config::new().prefix("/api"),
(
Router::new(
Config::new().prefix("/project/:projectId"),
"/".GET(|| async { "project" }),
)
.middleware(ParentMiddleware(Arc::clone(&events))),
Router::new(
Config::new().prefix("/project/:projectId"),
"/member".GET(|| async { "member" }),
)
.middleware(ChildMiddleware(Arc::clone(&events))),
),
);
assert_eq!(
block_on(router.handle(request("/api/project/abc123/member"))).body(),
b"member"
);
assert_eq!(*events.lock().unwrap(), ["child:before", "child:after"]);
events.lock().unwrap().clear();
assert_eq!(
block_on(router.handle(request("/api/project/abc123"))).body(),
b"project"
);
assert_eq!(*events.lock().unwrap(), ["parent:before", "parent:after"]);
}
impl Middleware for ObserveParams {
async fn handle(&self, request: Request, next: Next<'_>) -> Response {
let count = request.params().len();
let mut response = next.run(request).await;
response.set_header("X-Params", count.to_string());
response
}
}
#[test]
fn non_route_responses_clear_params_and_explicit_options_keeps_them() {
let router = Router::new(
Config::new(),
(
"/project/:projectId".GET(|| async { "ok" }),
"/explicit/:projectId".OPTIONS(|| async { "ok" }),
Router::new(Config::new().prefix("/fallback/:projectId"), ())
.fallback(|| async { "fallback" }),
),
)
.middleware(ObserveParams);
for (method, path, status, expected) in [
(Method::GET, "/unknown", 404, "0"),
(Method::POST, "/project/abc123", 405, "0"),
(Method::OPTIONS, "/project/abc123", 204, "0"),
(Method::OPTIONS, "/explicit/abc123", 200, "1"),
(Method::GET, "/fallback/abc123/missing", 200, "0"),
] {
let mut input = request(path);
input.method = method;
input.set_params(vec![("stale".into(), "value".into())]);
let (actual, headers, _) = block_on(router.handle(input)).into_parts();
assert_eq!(actual, status);
assert_eq!(headers.get("x-params"), Some(expected.as_bytes()));
}
}
#[test]
fn dynamic_scopes_keep_depth_order_and_route_override() {
let events = Arc::new(Mutex::new(Vec::new()));
let router = Router::new(
Config::new(),
Router::new(
Config::new().prefix("/project/:veryLongParameter"),
Router::new(
Config::new().prefix("/member"),
"/".POST(|| async { "ok" })
.without_middleware::<ChildMiddleware>()
.middleware(RouteMiddleware(Arc::clone(&events))),
)
.middleware(ChildMiddleware(Arc::clone(&events))),
)
.middleware(ParentMiddleware(Arc::clone(&events))),
);
let mut input = request("/project/abc123/member");
input.method = Method::POST;
assert_eq!(block_on(router.handle(input)).status(), 200);
assert_eq!(
*events.lock().unwrap(),
[
"parent:before",
"route:before",
"route:after",
"parent:after"
]
);
}
#[test]
fn rewriting_metadata_does_not_change_captured_params() {
let router = Router::new(
Config::new(),
"/project/:projectId"
.GET(|Path(path): Path<ProjectPath>| async move { path.project_id }),
)
.middleware(RewriteRequest);
assert_eq!(
block_on(router.handle(request("/project/abc123"))).body(),
b"abc123"
);
}
}