Skip to main content

rustlavel_http/
router.rs

1//! The router: what an application's `routes/web.rs` fills in.
2//!
3//! ```ignore
4//! pub fn routes(r: &mut Router) {
5//!     r.get("/", home);
6//!     r.get("/users/{id}", show).name("users.show");
7//!
8//!     r.group("/admin", |r| {
9//!         r.middleware(auth);
10//!         r.get("/dashboard", dashboard);
11//!     });
12//! }
13//! ```
14
15use crate::handler::Handler;
16use crate::method::Method;
17use crate::middleware::{Middleware, Next};
18use crate::request::Request;
19use crate::response::Response;
20use crate::status::Status;
21use crate::url;
22use std::collections::BTreeMap;
23use std::sync::Arc;
24
25/// One piece of a route pattern.
26#[derive(Debug, Clone, PartialEq, Eq)]
27enum Segment {
28    /// A literal path segment.
29    Static(String),
30    /// `{id}` — matches exactly one segment and captures it.
31    Param(String),
32    /// `{path:*}` — matches the rest of the path, slashes included.
33    Wildcard(String),
34}
35
36pub struct Route {
37    pub method: Method,
38    /// The pattern as written, used for `route:list` and metrics labels.
39    pub pattern: String,
40    pub name: Option<String>,
41    /// What this route does, in one line. Feeds generated API documentation.
42    pub summary: Option<String>,
43    /// A grouping label, so generated docs are not one flat list.
44    pub tag: Option<String>,
45    /// Documented responses: status, and what it means.
46    pub responses: Vec<(u16, String)>,
47    /// Documented parameters: name, and what it is.
48    pub parameters: Vec<(String, String)>,
49    pub deprecated: bool,
50    segments: Vec<Segment>,
51    handler: Arc<dyn Handler>,
52    middleware: Arc<Vec<Arc<dyn Middleware>>>,
53}
54
55impl Route {
56    /// The parameter names this route captures, in order.
57    ///
58    /// Generated documentation needs these even when the author documented
59    /// none, because a path parameter is required whether or not it is
60    /// described.
61    pub fn parameter_names(&self) -> Vec<String> {
62        self.segments
63            .iter()
64            .filter_map(|segment| match segment {
65                Segment::Param(name) | Segment::Wildcard(name) => Some(name.clone()),
66                Segment::Static(_) => None,
67            })
68            .collect()
69    }
70}
71
72impl Route {
73    /// How specific this route is, so `/users/new` is tried before `/users/{id}`.
74    fn specificity(&self) -> (usize, usize) {
75        let statics = self.segments.iter().filter(|s| matches!(s, Segment::Static(_))).count();
76        let wildcards = self.segments.iter().filter(|s| matches!(s, Segment::Wildcard(_))).count();
77        // More static segments first; any wildcard sinks to the bottom.
78        (wildcards, usize::MAX - statics)
79    }
80
81    fn match_path(&self, path: &str) -> Option<BTreeMap<String, String>> {
82        let mut params = BTreeMap::new();
83        let mut parts = split_path(path);
84        let mut index = 0;
85
86        while index < self.segments.len() {
87            match &self.segments[index] {
88                Segment::Wildcard(name) => {
89                    // Consumes everything that is left, including nothing.
90                    let rest = parts.collect::<Vec<_>>().join("/");
91                    params.insert(name.clone(), url::decode(&rest));
92                    return Some(params);
93                }
94                Segment::Static(expected) => {
95                    if parts.next()? != expected {
96                        return None;
97                    }
98                }
99                Segment::Param(name) => {
100                    let value = parts.next()?;
101                    if value.is_empty() {
102                        return None;
103                    }
104                    params.insert(name.clone(), url::decode(value));
105                }
106            }
107            index += 1;
108        }
109
110        // Every pattern segment matched; the path must be exhausted too.
111        parts.next().is_none().then_some(params)
112    }
113}
114
115fn split_path(path: &str) -> impl Iterator<Item = &str> {
116    path.split('/').filter(|part| !part.is_empty())
117}
118
119fn parse_pattern(pattern: &str) -> Vec<Segment> {
120    split_path(pattern)
121        .map(|part| match part.strip_prefix('{').and_then(|p| p.strip_suffix('}')) {
122            Some(name) => match name.strip_suffix(":*") {
123                Some(name) => Segment::Wildcard(name.to_string()),
124                None => Segment::Param(name.to_string()),
125            },
126            None => Segment::Static(part.to_string()),
127        })
128        .collect()
129}
130
131/// Collects routes, then answers requests.
132#[derive(Default)]
133pub struct Router {
134    routes: Vec<Route>,
135    /// Prefix and middleware of the group currently being defined.
136    scope_prefix: String,
137    scope_middleware: Vec<Arc<dyn Middleware>>,
138    /// Runs for every request, whatever the route.
139    global_middleware: Vec<Arc<dyn Middleware>>,
140    fallback: Option<Arc<dyn Handler>>,
141}
142
143impl Router {
144    pub fn new() -> Self {
145        Self::default()
146    }
147
148    /// Add middleware. Inside a `group` it applies to that group's routes;
149    /// at the top level it applies to every request, including 404s.
150    pub fn middleware(&mut self, middleware: impl Middleware) -> &mut Self {
151        if self.scope_prefix.is_empty() && self.scope_middleware.is_empty() {
152            self.global_middleware.push(Arc::new(middleware));
153        } else {
154            self.scope_middleware.push(Arc::new(middleware));
155        }
156        self
157    }
158
159    /// Register routes under a shared prefix and middleware stack.
160    pub fn group(&mut self, prefix: &str, define: impl FnOnce(&mut Router)) -> &mut Self {
161        let mut child = Router {
162            scope_prefix: join_paths(&self.scope_prefix, prefix),
163            scope_middleware: self.scope_middleware.clone(),
164            ..Router::default()
165        };
166        define(&mut child);
167
168        // A group's global-looking middleware belongs to that group only.
169        debug_assert!(child.global_middleware.is_empty() || !child.routes.is_empty());
170        self.routes.extend(child.routes);
171        self
172    }
173
174    /// The response when nothing matched. Defaults to a plain 404.
175    pub fn fallback(&mut self, handler: impl Handler) -> &mut Self {
176        self.fallback = Some(Arc::new(handler));
177        self
178    }
179
180    pub fn route(&mut self, method: Method, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
181        let full = join_paths(&self.scope_prefix, pattern);
182        self.routes.push(Route {
183            method,
184            segments: parse_pattern(&full),
185            pattern: full,
186            name: None,
187            summary: None,
188            tag: None,
189            responses: Vec::new(),
190            parameters: Vec::new(),
191            deprecated: false,
192            handler: Arc::new(handler),
193            middleware: Arc::new(self.scope_middleware.clone()),
194        });
195        RouteHandle { index: self.routes.len() - 1, router: self }
196    }
197
198    pub fn get(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
199        self.route(Method::Get, pattern, handler)
200    }
201
202    pub fn post(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
203        self.route(Method::Post, pattern, handler)
204    }
205
206    pub fn put(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
207        self.route(Method::Put, pattern, handler)
208    }
209
210    pub fn patch(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
211        self.route(Method::Patch, pattern, handler)
212    }
213
214    pub fn delete(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
215        self.route(Method::Delete, pattern, handler)
216    }
217
218    /// Start a RESTful resource: `r.resource("/posts").index(..).show(..)`.
219    pub fn resource<'r>(&'r mut self, base: &str) -> Resource<'r> {
220        Resource { base: base.trim_end_matches('/').to_string(), router: self }
221    }
222
223    /// Sort routes so lookups are deterministic. Called once before serving.
224    pub fn finalize(&mut self) {
225        self.routes.sort_by_key(Route::specificity);
226    }
227
228    pub fn routes(&self) -> &[Route] {
229        &self.routes
230    }
231
232    /// Build a URL from a named route: `url_for("users.show", &[("id", "7")])`.
233    pub fn url_for(&self, name: &str, params: &[(&str, &str)]) -> Option<String> {
234        let route = self.routes.iter().find(|route| route.name.as_deref() == Some(name))?;
235        let lookup = |key: &str| params.iter().find(|(k, _)| *k == key).map(|(_, v)| *v);
236
237        let mut out = String::new();
238        for segment in &route.segments {
239            out.push('/');
240            match segment {
241                Segment::Static(value) => out.push_str(value),
242                Segment::Param(name) => out.push_str(&url::encode(lookup(name)?)),
243                Segment::Wildcard(name) => out.push_str(lookup(name)?),
244            }
245        }
246        Some(if out.is_empty() { "/".to_string() } else { out })
247    }
248
249    /// Match a request and run it through the pipeline.
250    pub async fn dispatch(&self, mut request: Request) -> Response {
251        let path = request.path().to_string();
252        let mut path_matched = false;
253        let mut allowed = Vec::new();
254
255        for route in &self.routes {
256            let Some(params) = route.match_path(&path) else { continue };
257            path_matched = true;
258
259            // HEAD is served by the GET route, minus the body.
260            let usable = route.method == request.method()
261                || (request.method() == Method::Head && route.method == Method::Get);
262
263            if !usable {
264                allowed.push(route.method);
265                continue;
266            }
267
268            request.set_params(params);
269            request.route = Some(route.pattern.clone());
270
271            let mut stack = self.global_middleware.clone();
272            stack.extend(route.middleware.iter().cloned());
273            return run_guarded(Next::new(Arc::new(stack), Arc::clone(&route.handler)), request).await;
274        }
275
276        let response = if path_matched {
277            // The path exists but not for this verb: 405, and say what is allowed.
278            allowed.sort();
279            allowed.dedup();
280            let allow = allowed.iter().map(|m| m.as_str()).collect::<Vec<_>>().join(", ");
281            Response::new(Status::METHOD_NOT_ALLOWED).with_header("allow", allow).with_text(format!(
282                "{} is not allowed on {path}",
283                request.method()
284            ))
285        } else {
286            match &self.fallback {
287                Some(handler) => {
288                    let stack = Arc::new(self.global_middleware.clone());
289                    return Next::new(stack, Arc::clone(handler)).run(request).await;
290                }
291                None => Response::not_found(),
292            }
293        };
294
295        // Global middleware still observes unmatched requests, so logging and
296        // Telescope see the 404s too.
297        let stack = Arc::new(self.global_middleware.clone());
298        let endpoint: Arc<dyn Handler> = Arc::new(crate::handler::Fixed(response));
299        Next::new(stack, endpoint).run(request).await
300    }
301}
302
303/// Run the pipeline, turning a panic into the error page.
304///
305/// This lives at the router rather than in the server so a panicking handler
306/// fails the same way in a test as it does in production — a test that would
307/// otherwise abort the whole run instead reports a 500.
308async fn run_guarded(next: Next, request: Request) -> Response {
309    crate::panic::install_hook();
310    let started = std::time::Instant::now();
311
312    // A panicking handler still needs to render a page describing the request,
313    // but the request itself has been moved into the pipeline by then.
314    let probe = Request::new(request.method(), request.target().to_string())
315        .with_header("accept", request.header("accept").unwrap_or("text/html"));
316    let route = request.route().map(str::to_string);
317
318    let response = match crate::panic::catch(next.run(request)).await {
319        Ok(response) => response,
320        Err(message) => {
321            let location = crate::panic::take_location().map(|l| (l.file, l.line));
322            rustlavel_core::error!(
323                "panic in {} {}: {message}",
324                probe.method(),
325                probe.path()
326            );
327            crate::error_page::render(
328                &crate::error_page::Diagnostic::from_panic(message, location),
329                Some(&probe),
330            )
331        }
332    };
333
334    // Dispatched here rather than in the server, so instrumentation sees the
335    // same events under the test client as it does over a socket.
336    if rustlavel_core::events::has_subscribers() {
337        let mut event = rustlavel_core::Event::new("http.request")
338            .with("method", probe.method().as_str())
339            .with("path", probe.path())
340            .with("status", response.status.code())
341            .took(started.elapsed());
342        if let Some(route) = route {
343            // The pattern, not the path: one metric series per route, not per id.
344            event = event.with("route", route);
345        }
346        event.dispatch();
347    }
348
349    response
350}
351
352/// Returned by `get`/`post`/… so a route can be named after registration.
353pub struct RouteHandle<'r> {
354    index: usize,
355    router: &'r mut Router,
356}
357
358impl RouteHandle<'_> {
359    /// Name the route for `url_for` and `route:list`.
360    pub fn name(self, name: &str) -> Self {
361        self.router.routes[self.index].name = Some(name.to_string());
362        self
363    }
364
365    /// Say what this route does, in one line.
366    ///
367    /// Documentation is attached here rather than kept in a separate file, so
368    /// it cannot drift away from the route it describes.
369    pub fn describe(self, summary: &str) -> Self {
370        self.router.routes[self.index].summary = Some(summary.to_string());
371        self
372    }
373
374    /// Group this route under a heading in generated documentation.
375    pub fn tag(self, tag: &str) -> Self {
376        self.router.routes[self.index].tag = Some(tag.to_string());
377        self
378    }
379
380    /// Document a response this route can return.
381    pub fn responds(self, status: u16, description: &str) -> Self {
382        self.router.routes[self.index].responses.push((status, description.to_string()));
383        self
384    }
385
386    /// Describe a parameter. Undescribed path parameters are still documented,
387    /// just without prose.
388    pub fn param(self, name: &str, description: &str) -> Self {
389        self.router.routes[self.index]
390            .parameters
391            .push((name.to_string(), description.to_string()));
392        self
393    }
394
395    /// Mark the route as deprecated in generated documentation.
396    pub fn deprecated(self) -> Self {
397        self.router.routes[self.index].deprecated = true;
398        self
399    }
400}
401
402/// The seven RESTful routes, registered one at a time.
403pub struct Resource<'r> {
404    base: String,
405    router: &'r mut Router,
406}
407
408impl Resource<'_> {
409    /// `GET /posts`
410    pub fn index(self, handler: impl Handler) -> Self {
411        let (base, router) = (self.base.clone(), self.router);
412        router.get(&base, handler).name(&format!("{}.index", resource_name(&base)));
413        Resource { base, router }
414    }
415
416    /// `POST /posts`
417    pub fn store(self, handler: impl Handler) -> Self {
418        let (base, router) = (self.base.clone(), self.router);
419        router.post(&base, handler).name(&format!("{}.store", resource_name(&base)));
420        Resource { base, router }
421    }
422
423    /// `GET /posts/{id}`
424    pub fn show(self, handler: impl Handler) -> Self {
425        let (base, router) = (self.base.clone(), self.router);
426        let pattern = format!("{base}/{{id}}");
427        router.get(&pattern, handler).name(&format!("{}.show", resource_name(&base)));
428        Resource { base, router }
429    }
430
431    /// `PUT /posts/{id}`
432    pub fn update(self, handler: impl Handler) -> Self {
433        let (base, router) = (self.base.clone(), self.router);
434        let pattern = format!("{base}/{{id}}");
435        router.put(&pattern, handler).name(&format!("{}.update", resource_name(&base)));
436        Resource { base, router }
437    }
438
439    /// `DELETE /posts/{id}`
440    pub fn destroy(self, handler: impl Handler) -> Self {
441        let (base, router) = (self.base.clone(), self.router);
442        let pattern = format!("{base}/{{id}}");
443        router.delete(&pattern, handler).name(&format!("{}.destroy", resource_name(&base)));
444        Resource { base, router }
445    }
446}
447
448fn resource_name(base: &str) -> String {
449    base.trim_matches('/').replace('/', ".")
450}
451
452fn join_paths(prefix: &str, path: &str) -> String {
453    let joined = format!("/{}/{}", prefix.trim_matches('/'), path.trim_matches('/'));
454    let cleaned = joined.replace("//", "/");
455    if cleaned.len() > 1 { cleaned.trim_end_matches('/').to_string() } else { "/".to_string() }
456}
457
458#[cfg(test)]
459mod tests {
460    use super::*;
461
462    async fn ok(_req: Request) -> &'static str {
463        "ok"
464    }
465
466    fn router_with(define: impl FnOnce(&mut Router)) -> Router {
467        let mut router = Router::new();
468        define(&mut router);
469        router.finalize();
470        router
471    }
472
473    #[tokio::test]
474    async fn matches_static_and_parameter_routes() {
475        let router = router_with(|r| {
476            r.get("/", ok);
477            r.get("/users/{id}", |req: Request| async move {
478                format!("user {}", req.param("id").unwrap())
479            });
480        });
481
482        assert_eq!(router.dispatch(Request::new(Method::Get, "/")).await.body_string(), "ok");
483        assert_eq!(
484            router.dispatch(Request::new(Method::Get, "/users/7")).await.body_string(),
485            "user 7"
486        );
487        assert_eq!(
488            router.dispatch(Request::new(Method::Get, "/nope")).await.status,
489            Status::NOT_FOUND
490        );
491    }
492
493    #[tokio::test]
494    async fn static_segments_win_over_parameters() {
495        let router = router_with(|r| {
496            r.get("/users/{id}", |_req: Request| async { "param" });
497            r.get("/users/new", |_req: Request| async { "static" });
498        });
499
500        assert_eq!(
501            router.dispatch(Request::new(Method::Get, "/users/new")).await.body_string(),
502            "static"
503        );
504        assert_eq!(
505            router.dispatch(Request::new(Method::Get, "/users/12")).await.body_string(),
506            "param"
507        );
508    }
509
510    #[tokio::test]
511    async fn wildcards_capture_the_remaining_path() {
512        let router = router_with(|r| {
513            r.get("/files/{path:*}", |req: Request| async move {
514                req.param("path").unwrap_or_default().to_string()
515            });
516        });
517
518        let response = router.dispatch(Request::new(Method::Get, "/files/css/app.css")).await;
519        assert_eq!(response.body_string(), "css/app.css");
520    }
521
522    #[tokio::test]
523    async fn parameters_are_percent_decoded() {
524        let router = router_with(|r| {
525            r.get("/tags/{tag}", |req: Request| async move { req.param("tag").unwrap().to_string() });
526        });
527
528        let response = router.dispatch(Request::new(Method::Get, "/tags/rust%20lang")).await;
529        assert_eq!(response.body_string(), "rust lang");
530    }
531
532    #[tokio::test]
533    async fn wrong_method_reports_405_with_allow() {
534        let router = router_with(|r| {
535            r.post("/users", ok);
536        });
537
538        let response = router.dispatch(Request::new(Method::Get, "/users")).await;
539        assert_eq!(response.status, Status::METHOD_NOT_ALLOWED);
540        assert_eq!(response.headers.get("allow"), Some("POST"));
541    }
542
543    #[tokio::test]
544    async fn head_is_served_by_the_get_route() {
545        let router = router_with(|r| {
546            r.get("/", ok);
547        });
548
549        assert_eq!(router.dispatch(Request::new(Method::Head, "/")).await.status, Status::OK);
550    }
551
552    #[tokio::test]
553    async fn groups_apply_prefix_and_middleware() {
554        let router = router_with(|r| {
555            r.group("/admin", |r| {
556                r.middleware(|req: Request, next: Next| async move {
557                    next.run(req).await.with_header("x-guard", "on")
558                });
559                r.get("/dashboard", ok);
560            });
561            r.get("/public", ok);
562        });
563
564        let guarded = router.dispatch(Request::new(Method::Get, "/admin/dashboard")).await;
565        assert_eq!(guarded.headers.get("x-guard"), Some("on"));
566
567        let open = router.dispatch(Request::new(Method::Get, "/public")).await;
568        assert_eq!(open.headers.get("x-guard"), None);
569    }
570
571    #[tokio::test]
572    async fn global_middleware_also_sees_unmatched_requests() {
573        let router = router_with(|r| {
574            r.middleware(|req: Request, next: Next| async move {
575                next.run(req).await.with_header("x-seen", "1")
576            });
577            r.get("/", ok);
578        });
579
580        let missing = router.dispatch(Request::new(Method::Get, "/nope")).await;
581        assert_eq!(missing.status, Status::NOT_FOUND);
582        assert_eq!(missing.headers.get("x-seen"), Some("1"));
583    }
584
585    #[test]
586    fn builds_urls_from_named_routes() {
587        let router = router_with(|r| {
588            r.get("/users/{id}/posts/{slug}", ok).name("users.posts");
589        });
590
591        assert_eq!(
592            router.url_for("users.posts", &[("id", "7"), ("slug", "hello world")]).as_deref(),
593            Some("/users/7/posts/hello%20world")
594        );
595        assert_eq!(router.url_for("users.posts", &[("id", "7")]), None);
596        assert_eq!(router.url_for("missing", &[]), None);
597    }
598
599    #[tokio::test]
600    async fn resource_registers_the_rest_routes() {
601        let router = router_with(|r| {
602            r.resource("/posts").index(ok).store(ok).show(ok).update(ok).destroy(ok);
603        });
604
605        assert_eq!(router.routes().len(), 5);
606        assert_eq!(router.url_for("posts.show", &[("id", "3")]).as_deref(), Some("/posts/3"));
607        assert_eq!(router.dispatch(Request::new(Method::Delete, "/posts/3")).await.status, Status::OK);
608    }
609
610    #[tokio::test]
611    async fn fallback_replaces_the_default_404() {
612        let router = router_with(|r| {
613            r.fallback(|_req: Request| async { (404, "custom miss") });
614        });
615
616        let response = router.dispatch(Request::new(Method::Get, "/anything")).await;
617        assert_eq!(response.body_string(), "custom miss");
618    }
619
620    #[test]
621    fn documentation_rides_along_with_the_route() {
622        let router = router_with(|r| {
623            r.get("/users/{id}", ok)
624                .name("users.show")
625                .describe("Fetch one user")
626                .tag("Users")
627                .param("id", "The user's id")
628                .responds(200, "The user")
629                .responds(404, "No such user");
630        });
631
632        let route = &router.routes()[0];
633        assert_eq!(route.summary.as_deref(), Some("Fetch one user"));
634        assert_eq!(route.tag.as_deref(), Some("Users"));
635        assert_eq!(route.responses.len(), 2);
636        assert_eq!(route.parameter_names(), vec!["id"]);
637        assert!(!route.deprecated);
638    }
639
640    #[test]
641    fn path_parameters_are_known_even_when_undocumented() {
642        let router = router_with(|r| {
643            r.get("/teams/{team}/members/{member}", ok);
644            r.get("/files/{path:*}", ok);
645        });
646
647        let names: Vec<Vec<String>> =
648            router.routes().iter().map(Route::parameter_names).collect();
649        assert!(names.contains(&vec!["team".to_string(), "member".to_string()]));
650        assert!(names.contains(&vec!["path".to_string()]));
651    }
652
653    #[test]
654    fn joins_prefixes_without_doubling_slashes() {
655        assert_eq!(join_paths("/admin/", "/users"), "/admin/users");
656        assert_eq!(join_paths("", "/"), "/");
657        assert_eq!(join_paths("/admin", ""), "/admin");
658    }
659}