Skip to main content

serverkit/
middleware.rs

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}