1use std::{any::TypeId, future::Future, pin::Pin};
2
3use crate::{Request, Response};
4
5pub trait Middleware: Send + Sync + 'static {
6 async fn handle(&self, request: Request, next: Next<'_>) -> Response;
7}
8
9pub(crate) type MiddlewareFuture<'request> = Pin<Box<dyn Future<Output = Response> + 'request>>;
10
11pub(crate) struct MiddlewareEntry {
12 type_id: TypeId,
13 service: Box<dyn MiddlewareService>,
14}
15
16impl MiddlewareEntry {
17 pub(crate) fn new<M: Middleware>(middleware: M) -> Self {
18 Self {
19 type_id: TypeId::of::<M>(),
20 service: Box::new(middleware),
21 }
22 }
23
24 pub(crate) fn type_id(&self) -> TypeId {
25 self.type_id
26 }
27}
28
29pub(crate) trait MiddlewareService: Send + Sync {
30 fn call<'request>(
31 &'request self,
32 request: Request,
33 next: Next<'request>,
34 ) -> MiddlewareFuture<'request>;
35}
36
37impl<M: Middleware> MiddlewareService for M {
38 fn call<'request>(
39 &'request self,
40 request: Request,
41 next: Next<'request>,
42 ) -> MiddlewareFuture<'request> {
43 Box::pin(self.handle(request, next))
44 }
45}
46
47pub struct Next<'next> {
48 middlewares: &'next [&'next MiddlewareEntry],
49 terminal: &'next dyn MiddlewareTerminal,
50}
51
52impl Next<'_> {
53 pub async fn run(self, request: Request) -> Response {
54 let Some((middleware, remaining)) = self.middlewares.split_first() else {
55 return self.terminal.call(request).await;
56 };
57
58 middleware
59 .service
60 .call(
61 request,
62 Next {
63 middlewares: remaining,
64 terminal: self.terminal,
65 },
66 )
67 .await
68 }
69}
70
71pub(crate) trait MiddlewareTerminal {
72 fn call(&self, request: Request) -> MiddlewareFuture<'_>;
73}
74
75pub(crate) async fn run(
76 middlewares: &[&MiddlewareEntry],
77 terminal: &dyn MiddlewareTerminal,
78 request: Request,
79) -> Response {
80 Next {
81 middlewares,
82 terminal,
83 }
84 .run(request)
85 .await
86}
87
88#[cfg(test)]
89mod tests {
90 use std::{
91 convert::Infallible,
92 future::Future,
93 sync::{
94 Arc, Mutex,
95 atomic::{AtomicUsize, Ordering},
96 },
97 task::{Context, Poll, Waker},
98 };
99
100 use crate::{
101 Config, FromRequest, Headers, Method, Middleware, Next, Request, RequestStream, Response,
102 RouteMethods, Router, StreamError,
103 };
104
105 struct EmptyStream;
106
107 impl RequestStream for EmptyStream {
108 fn poll_next(
109 &mut self,
110 _context: &mut Context<'_>,
111 ) -> Poll<Option<Result<(), StreamError>>> {
112 Poll::Ready(None)
113 }
114
115 fn chunk(&self) -> &[u8] {
116 &[]
117 }
118 }
119
120 struct ParentMiddleware(Arc<Mutex<Vec<&'static str>>>);
121 struct ChildMiddleware(Arc<Mutex<Vec<&'static str>>>);
122 struct RouteMiddleware(Arc<Mutex<Vec<&'static str>>>);
123
124 macro_rules! record_middleware {
125 ($middleware:ident, $before:literal, $after:literal) => {
126 impl Middleware for $middleware {
127 async fn handle(&self, request: Request, next: Next<'_>) -> Response {
128 self.0.lock().unwrap().push($before);
129 let response = next.run(request).await;
130 self.0.lock().unwrap().push($after);
131 response
132 }
133 }
134 };
135 }
136
137 record_middleware!(ParentMiddleware, "parent:before", "parent:after");
138 record_middleware!(ChildMiddleware, "child:before", "child:after");
139 record_middleware!(RouteMiddleware, "route:before", "route:after");
140
141 struct Authentication(Arc<AtomicUsize>);
142
143 struct InjectHeader;
144
145 struct InjectedHeader(String);
146
147 struct RewriteRequest;
148
149 struct RequestTarget(String);
150
151 impl Middleware for Authentication {
152 async fn handle(&self, request: Request, next: Next<'_>) -> Response {
153 self.0.fetch_add(1, Ordering::Relaxed);
154 next.run(request).await
155 }
156 }
157
158 impl Middleware for InjectHeader {
159 async fn handle(&self, mut request: Request, next: Next<'_>) -> Response {
160 request.headers.set("X-Injected", "middleware").unwrap();
161 next.run(request).await
162 }
163 }
164
165 impl Middleware for RewriteRequest {
166 async fn handle(&self, mut request: Request, next: Next<'_>) -> Response {
167 request.method = Method::PATCH;
168 request.path = "/rewritten".to_owned();
169 request.query = Some("source=middleware".to_owned());
170 next.run(request).await
171 }
172 }
173
174 impl<'request> FromRequest<(&'request Request, &'request [u8])> for InjectedHeader {
175 type Error = Infallible;
176
177 async fn from_request(
178 input: (&'request Request, &'request [u8]),
179 ) -> Result<Self, Self::Error> {
180 let value = input
181 .0
182 .headers
183 .get("X-Injected")
184 .and_then(|value| std::str::from_utf8(value).ok())
185 .unwrap_or_default()
186 .to_owned();
187
188 Ok(Self(value))
189 }
190 }
191
192 impl<'request> FromRequest<(&'request Request, &'request [u8])> for RequestTarget {
193 type Error = Infallible;
194
195 async fn from_request(
196 input: (&'request Request, &'request [u8]),
197 ) -> Result<Self, Self::Error> {
198 Ok(Self(format!(
199 "{} {}?{}",
200 input.0.method.as_str(),
201 input.0.path,
202 input.0.query.as_deref().unwrap_or_default(),
203 )))
204 }
205 }
206
207 fn request(path: &str) -> Request {
208 Request::from_parts(
209 Method::GET,
210 path,
211 None,
212 Headers::new(),
213 Box::new(EmptyStream),
214 )
215 }
216
217 fn block_on<F: Future>(future: F) -> F::Output {
218 let mut future = std::pin::pin!(future);
219 let waker = Waker::noop();
220 let mut context = Context::from_waker(waker);
221
222 loop {
223 match future.as_mut().poll(&mut context) {
224 Poll::Ready(output) => return output,
225 Poll::Pending => std::thread::yield_now(),
226 }
227 }
228 }
229
230 #[test]
231 fn runs_parent_child_and_route_middleware_in_scope_order() {
232 let events = Arc::new(Mutex::new(Vec::new()));
233 let handler_events = Arc::clone(&events);
234 let child = Router::new(
235 Config::new().prefix("/v1"),
236 "/health"
237 .GET(move || {
238 let events = Arc::clone(&handler_events);
239
240 async move {
241 events.lock().unwrap().push("handler");
242 "ok"
243 }
244 })
245 .middleware(RouteMiddleware(Arc::clone(&events))),
246 )
247 .middleware(ChildMiddleware(Arc::clone(&events)))
248 .at("/mounted");
249 let router = Router::new(Config::new().prefix("/parent"), ())
250 .middleware(ParentMiddleware(Arc::clone(&events)))
251 .route(child);
252
253 let response = block_on(router.handle(request("/parent/mounted/v1/health")));
254
255 assert_eq!(response.body(), b"ok");
256 assert_eq!(
257 events.lock().unwrap().as_slice(),
258 [
259 "parent:before",
260 "child:before",
261 "route:before",
262 "handler",
263 "route:after",
264 "child:after",
265 "parent:after",
266 ],
267 );
268 }
269
270 #[test]
271 fn child_middleware_only_runs_inside_the_child_prefix() {
272 let calls = Arc::new(AtomicUsize::new(0));
273 let child = Router::new(
274 Config::new().prefix("/api"),
275 "/inside".GET(|| async { "in" }),
276 )
277 .middleware(Authentication(Arc::clone(&calls)));
278 let router = Router::new(Config::new(), "/outside".GET(|| async { "out" })).route(child);
279
280 assert_eq!(block_on(router.handle(request("/outside"))).body(), b"out",);
281 assert_eq!(calls.load(Ordering::Relaxed), 0);
282
283 assert_eq!(
284 block_on(router.handle(request("/api/inside"))).body(),
285 b"in",
286 );
287 assert_eq!(calls.load(Ordering::Relaxed), 1);
288
289 assert_eq!(
290 block_on(router.handle(request("/api/missing"))).status(),
291 404,
292 );
293 assert_eq!(calls.load(Ordering::Relaxed), 2);
294
295 assert_eq!(block_on(router.handle(request("/missing"))).status(), 404,);
296 assert_eq!(calls.load(Ordering::Relaxed), 2);
297 }
298
299 #[test]
300 fn middleware_can_modify_request_headers_before_extraction() {
301 let router = Router::new(
302 Config::new(),
303 "/header".GET(|InjectedHeader(value): InjectedHeader| async move { value }),
304 )
305 .middleware(InjectHeader);
306
307 let response = block_on(router.handle(request("/header")));
308
309 assert_eq!(response.body(), b"middleware");
310 }
311
312 #[test]
313 fn middleware_can_modify_the_request_target_before_extraction() {
314 let router = Router::new(
315 Config::new(),
316 "/original".GET(|RequestTarget(value): RequestTarget| async move { value }),
317 )
318 .middleware(RewriteRequest);
319
320 let response = block_on(router.handle(request("/original")));
321
322 assert_eq!(response.body(), b"PATCH /rewritten?source=middleware");
323 }
324
325 #[test]
326 fn route_can_exclude_inherited_middleware_by_type() {
327 let calls = Arc::new(AtomicUsize::new(0));
328 let router = Router::new(
329 Config::new(),
330 (
331 "/private".GET(|| async { "private" }),
332 "/public"
333 .GET(|| async { "public" })
334 .without_middleware::<Authentication>(),
335 ),
336 )
337 .middleware(Authentication(Arc::clone(&calls)));
338
339 assert_eq!(
340 block_on(router.handle(request("/private"))).body(),
341 b"private",
342 );
343 assert_eq!(
344 block_on(router.handle(request("/public"))).body(),
345 b"public",
346 );
347 assert_eq!(calls.load(Ordering::Relaxed), 1);
348 }
349}