1use 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#[derive(Debug, Clone, PartialEq, Eq)]
27enum Segment {
28 Static(String),
30 Param(String),
32 Wildcard(String),
34}
35
36pub struct Route {
37 pub method: Method,
38 pub pattern: String,
40 pub name: Option<String>,
41 pub summary: Option<String>,
43 pub tag: Option<String>,
45 pub responses: Vec<(u16, String)>,
47 pub parameters: Vec<(String, String)>,
49 pub deprecated: bool,
50 pub version: Option<String>,
52 pub deprecated_at: Option<i64>,
54 pub sunset: Option<i64>,
56 segments: Vec<Segment>,
57 handler: Arc<dyn Handler>,
58 middleware: Arc<Vec<Arc<dyn Middleware>>>,
59}
60
61impl Route {
62 pub fn parameter_names(&self) -> Vec<String> {
68 self.segments
69 .iter()
70 .filter_map(|segment| match segment {
71 Segment::Param(name) | Segment::Wildcard(name) => Some(name.clone()),
72 Segment::Static(_) => None,
73 })
74 .collect()
75 }
76}
77
78impl Route {
79 fn specificity(&self) -> (usize, usize) {
81 let statics = self.segments.iter().filter(|s| matches!(s, Segment::Static(_))).count();
82 let wildcards = self.segments.iter().filter(|s| matches!(s, Segment::Wildcard(_))).count();
83 (wildcards, usize::MAX - statics)
85 }
86
87 fn match_path(&self, path: &str) -> Option<BTreeMap<String, String>> {
88 let mut params = BTreeMap::new();
89 let mut parts = split_path(path);
90 let mut index = 0;
91
92 while index < self.segments.len() {
93 match &self.segments[index] {
94 Segment::Wildcard(name) => {
95 let rest = parts.collect::<Vec<_>>().join("/");
97 params.insert(name.clone(), url::decode(&rest));
98 return Some(params);
99 }
100 Segment::Static(expected) => {
101 if parts.next()? != expected {
102 return None;
103 }
104 }
105 Segment::Param(name) => {
106 let value = parts.next()?;
107 if value.is_empty() {
108 return None;
109 }
110 params.insert(name.clone(), url::decode(value));
111 }
112 }
113 index += 1;
114 }
115
116 parts.next().is_none().then_some(params)
118 }
119}
120
121fn split_path(path: &str) -> impl Iterator<Item = &str> {
122 path.split('/').filter(|part| !part.is_empty())
123}
124
125fn parse_pattern(pattern: &str) -> Vec<Segment> {
126 split_path(pattern)
127 .map(|part| match part.strip_prefix('{').and_then(|p| p.strip_suffix('}')) {
128 Some(name) => match name.strip_suffix(":*") {
129 Some(name) => Segment::Wildcard(name.to_string()),
130 None => Segment::Param(name.to_string()),
131 },
132 None => Segment::Static(part.to_string()),
133 })
134 .collect()
135}
136
137#[derive(Default, Clone)]
144pub struct NamedRoutes(Vec<(String, Vec<Segment>)>);
145
146impl NamedRoutes {
147 pub fn url_for(&self, name: &str, params: &[(&str, &str)]) -> Option<String> {
148 let (_, segments) = self.0.iter().find(|(known, _)| known == name)?;
149 fill(segments, params)
150 }
151
152 pub fn names(&self) -> impl Iterator<Item = &str> {
154 self.0.iter().map(|(name, _)| name.as_str())
155 }
156}
157
158fn fill(segments: &[Segment], params: &[(&str, &str)]) -> Option<String> {
163 let lookup = |key: &str| params.iter().find(|(k, _)| *k == key).map(|(_, v)| *v);
164
165 let mut out = String::new();
166 for segment in segments {
167 out.push('/');
168 match segment {
169 Segment::Static(value) => out.push_str(value),
170 Segment::Param(name) => out.push_str(&url::encode(lookup(name)?)),
171 Segment::Wildcard(name) => out.push_str(lookup(name)?),
172 }
173 }
174 Some(if out.is_empty() { "/".to_string() } else { out })
175}
176
177#[derive(Default)]
179pub struct Router {
180 routes: Vec<Route>,
181 scope_prefix: String,
183 scope_version: Option<String>,
184 scope_middleware: Vec<Arc<dyn Middleware>>,
185 global_middleware: Vec<Arc<dyn Middleware>>,
187 fallback: Option<Arc<dyn Handler>>,
188}
189
190impl Router {
191 pub fn new() -> Self {
192 Self::default()
193 }
194
195 pub fn middleware(&mut self, middleware: impl Middleware) -> &mut Self {
198 if self.scope_prefix.is_empty() && self.scope_middleware.is_empty() {
199 self.global_middleware.push(Arc::new(middleware));
200 } else {
201 self.scope_middleware.push(Arc::new(middleware));
202 }
203 self
204 }
205
206 pub fn group(&mut self, prefix: &str, define: impl FnOnce(&mut Router)) -> &mut Self {
208 let mut child = Router {
209 scope_prefix: join_paths(&self.scope_prefix, prefix),
210 scope_middleware: self.scope_middleware.clone(),
211 scope_version: self.scope_version.clone(),
212 ..Router::default()
213 };
214 define(&mut child);
215
216 debug_assert!(child.global_middleware.is_empty() || !child.routes.is_empty());
218 self.routes.extend(child.routes);
219 self
220 }
221
222 pub fn version(&mut self, version: &str, define: impl FnOnce(&mut Router)) -> &mut Self {
235 let mut child = Router {
236 scope_prefix: join_paths(&self.scope_prefix, &format!("/{}", version.trim_start_matches('/'))),
237 scope_middleware: self.scope_middleware.clone(),
238 scope_version: Some(version.trim_start_matches('/').to_string()),
239 ..Router::default()
240 };
241 define(&mut child);
242 self.routes.extend(child.routes);
243 self
244 }
245
246 pub fn fallback(&mut self, handler: impl Handler) -> &mut Self {
248 self.fallback = Some(Arc::new(handler));
249 self
250 }
251
252 pub fn route(&mut self, method: Method, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
253 let full = join_paths(&self.scope_prefix, pattern);
254 self.routes.push(Route {
255 method,
256 segments: parse_pattern(&full),
257 pattern: full,
258 name: None,
259 summary: None,
260 tag: None,
261 responses: Vec::new(),
262 parameters: Vec::new(),
263 deprecated: false,
264 version: self.scope_version.clone(),
265 deprecated_at: None,
266 sunset: None,
267 handler: Arc::new(handler),
268 middleware: Arc::new(self.scope_middleware.clone()),
269 });
270 RouteHandle { index: self.routes.len() - 1, router: self }
271 }
272
273 pub fn get(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
274 self.route(Method::Get, pattern, handler)
275 }
276
277 pub fn post(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
278 self.route(Method::Post, pattern, handler)
279 }
280
281 pub fn put(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
282 self.route(Method::Put, pattern, handler)
283 }
284
285 pub fn patch(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
286 self.route(Method::Patch, pattern, handler)
287 }
288
289 pub fn delete(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
290 self.route(Method::Delete, pattern, handler)
291 }
292
293 pub fn resource<'r>(&'r mut self, base: &str) -> Resource<'r> {
295 Resource { base: base.trim_end_matches('/').to_string(), router: self }
296 }
297
298 pub fn finalize(&mut self) {
300 self.routes.sort_by_key(Route::specificity);
301 }
302
303 pub fn routes(&self) -> &[Route] {
304 &self.routes
305 }
306
307 pub fn url_for(&self, name: &str, params: &[(&str, &str)]) -> Option<String> {
309 let route = self.routes.iter().find(|route| route.name.as_deref() == Some(name))?;
310 fill(&route.segments, params)
311 }
312
313 pub fn named_routes(&self) -> NamedRoutes {
319 NamedRoutes(
320 self.routes
321 .iter()
322 .filter_map(|route| {
323 route.name.clone().map(|name| (name, route.segments.clone()))
324 })
325 .collect(),
326 )
327 }
328
329 pub async fn dispatch(&self, mut request: Request) -> Response {
331 let path = request.path().to_string();
332 let mut path_matched = false;
333 let mut allowed = Vec::new();
334
335 for route in &self.routes {
336 let Some(params) = route.match_path(&path) else { continue };
337 path_matched = true;
338
339 let usable = route.method == request.method()
341 || (request.method() == Method::Head && route.method == Method::Get);
342
343 if !usable {
344 allowed.push(route.method);
345 continue;
346 }
347
348 request.set_params(params);
349 request.route = Some(route.pattern.clone());
350 if let Some(version) = &route.version {
351 request.extend(crate::versioning::ApiVersion(version.clone()));
352 }
353
354 let mut stack = self.global_middleware.clone();
355 stack.extend(route.middleware.iter().cloned());
356 let response =
357 run_guarded(Next::new(Arc::new(stack), Arc::clone(&route.handler)), request).await;
358 return crate::versioning::stamp_lifecycle(route, response);
359 }
360
361 let response = if path_matched {
362 allowed.sort();
364 allowed.dedup();
365 let allow = allowed.iter().map(|m| m.as_str()).collect::<Vec<_>>().join(", ");
366 Response::new(Status::METHOD_NOT_ALLOWED).with_header("allow", allow).with_text(format!(
367 "{} is not allowed on {path}",
368 request.method()
369 ))
370 } else {
371 match &self.fallback {
372 Some(handler) => {
373 let stack = Arc::new(self.global_middleware.clone());
374 return Next::new(stack, Arc::clone(handler)).run(request).await;
375 }
376 None => Response::not_found(),
377 }
378 };
379
380 let stack = Arc::new(self.global_middleware.clone());
383 let endpoint: Arc<dyn Handler> = Arc::new(crate::handler::Fixed(response));
384 Next::new(stack, endpoint).run(request).await
385 }
386}
387
388async fn run_guarded(next: Next, request: Request) -> Response {
394 crate::panic::install_hook();
395 let started = std::time::Instant::now();
396
397 let probe = Request::new(request.method(), request.target().to_string())
400 .with_header("accept", request.header("accept").unwrap_or("text/html"));
401 let route = request.route().map(str::to_string);
402
403 let response = match crate::panic::catch(next.run(request)).await {
404 Ok(response) => response,
405 Err(message) => {
406 let location = crate::panic::take_location().map(|l| (l.file, l.line));
407 rustlavel_core::error!(
408 "panic in {} {}: {message}",
409 probe.method(),
410 probe.path()
411 );
412 crate::error_page::render(
413 &crate::error_page::Diagnostic::from_panic(message, location),
414 Some(&probe),
415 )
416 }
417 };
418
419 if rustlavel_core::events::has_subscribers() {
422 let mut event = rustlavel_core::Event::new("http.request")
423 .with("method", probe.method().as_str())
424 .with("path", probe.path())
425 .with("status", response.status.code())
426 .took(started.elapsed());
427 if let Some(route) = route {
428 event = event.with("route", route);
430 }
431 if let Some(id) = response.headers.get(crate::request_id::HEADER) {
434 event = event.with("request_id", id);
435 }
436 event.dispatch();
437 }
438
439 response
440}
441
442pub struct RouteHandle<'r> {
444 index: usize,
445 router: &'r mut Router,
446}
447
448impl RouteHandle<'_> {
449 pub fn name(self, name: &str) -> Self {
451 self.router.routes[self.index].name = Some(name.to_string());
452 self
453 }
454
455 pub fn describe(self, summary: &str) -> Self {
460 self.router.routes[self.index].summary = Some(summary.to_string());
461 self
462 }
463
464 pub fn tag(self, tag: &str) -> Self {
466 self.router.routes[self.index].tag = Some(tag.to_string());
467 self
468 }
469
470 pub fn responds(self, status: u16, description: &str) -> Self {
472 self.router.routes[self.index].responses.push((status, description.to_string()));
473 self
474 }
475
476 pub fn param(self, name: &str, description: &str) -> Self {
479 self.router.routes[self.index]
480 .parameters
481 .push((name.to_string(), description.to_string()));
482 self
483 }
484
485 pub fn middleware(self, middleware: impl Middleware) -> Self {
496 let stack = &mut self.router.routes[self.index].middleware;
497 Arc::make_mut(stack).push(Arc::new(middleware));
501 self
502 }
503
504 pub fn deprecated(self) -> Self {
506 self.router.routes[self.index].deprecated = true;
507 self
508 }
509
510 pub fn deprecated_at(self, date: &str) -> Self {
521 let when = crate::date::parse_ymd(date)
522 .unwrap_or_else(|| panic!("`{date}` is not a date; deprecated_at wants YYYY-MM-DD"));
523 let route = &mut self.router.routes[self.index];
524 route.deprecated = true;
525 route.deprecated_at = Some(when);
526 self
527 }
528
529 pub fn sunset(self, date: &str) -> Self {
540 let when = crate::date::parse_ymd(date)
541 .unwrap_or_else(|| panic!("`{date}` is not a date; sunset wants YYYY-MM-DD"));
542 let route = &mut self.router.routes[self.index];
543 route.deprecated = true;
544 route.sunset = Some(when);
545 self
546 }
547}
548
549pub struct Resource<'r> {
551 base: String,
552 router: &'r mut Router,
553}
554
555impl Resource<'_> {
556 pub fn index(self, handler: impl Handler) -> Self {
558 let (base, router) = (self.base.clone(), self.router);
559 router.get(&base, handler).name(&format!("{}.index", resource_name(&base)));
560 Resource { base, router }
561 }
562
563 pub fn store(self, handler: impl Handler) -> Self {
565 let (base, router) = (self.base.clone(), self.router);
566 router.post(&base, handler).name(&format!("{}.store", resource_name(&base)));
567 Resource { base, router }
568 }
569
570 pub fn show(self, handler: impl Handler) -> Self {
572 let (base, router) = (self.base.clone(), self.router);
573 let pattern = format!("{base}/{{id}}");
574 router.get(&pattern, handler).name(&format!("{}.show", resource_name(&base)));
575 Resource { base, router }
576 }
577
578 pub fn update(self, handler: impl Handler) -> Self {
580 let (base, router) = (self.base.clone(), self.router);
581 let pattern = format!("{base}/{{id}}");
582 router.put(&pattern, handler).name(&format!("{}.update", resource_name(&base)));
583 Resource { base, router }
584 }
585
586 pub fn destroy(self, handler: impl Handler) -> Self {
588 let (base, router) = (self.base.clone(), self.router);
589 let pattern = format!("{base}/{{id}}");
590 router.delete(&pattern, handler).name(&format!("{}.destroy", resource_name(&base)));
591 Resource { base, router }
592 }
593}
594
595fn resource_name(base: &str) -> String {
596 base.trim_matches('/').replace('/', ".")
597}
598
599fn join_paths(prefix: &str, path: &str) -> String {
600 let joined = format!("/{}/{}", prefix.trim_matches('/'), path.trim_matches('/'));
601 let cleaned = joined.replace("//", "/");
602 if cleaned.len() > 1 { cleaned.trim_end_matches('/').to_string() } else { "/".to_string() }
603}
604
605#[cfg(test)]
606mod tests {
607 #[tokio::test]
608 async fn per_route_middleware_does_not_leak_to_its_neighbours() {
609 use crate::testing::TestClient;
610 fn tag(name: &'static str) -> impl Middleware {
611 move |request: Request, next: Next| async move {
612 next.run(request).await.with_header("x-tag", name)
613 }
614 }
615
616 let mut router = Router::new();
617 router.group("/admin", |admin| {
618 admin.middleware(tag("group"));
619 admin.get("/users", |_req: Request| async { Response::text("users") }).middleware(tag("view"));
620 admin.get("/roles", |_req: Request| async { Response::text("roles") });
621 });
622 let client = TestClient::new(router);
623
624 assert_eq!(client.get("/admin/users").await.status(), 200);
627 assert_eq!(client.get("/admin/roles").await.status(), 200);
628
629 let mut counted = Router::new();
630 let hits = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
631 let counter = hits.clone();
632 counted.group("/admin", |admin| {
633 admin
634 .get("/users", |_req: Request| async { Response::text("users") })
635 .middleware(move |request: Request, next: Next| {
636 let counter = counter.clone();
637 async move {
638 counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
639 next.run(request).await
640 }
641 });
642 admin.get("/roles", |_req: Request| async { Response::text("roles") });
643 });
644 let client = TestClient::new(counted);
645 client.get("/admin/roles").await;
646 assert_eq!(hits.load(std::sync::atomic::Ordering::SeqCst), 0, "the guard belongs to /users only");
647 client.get("/admin/users").await;
648 assert_eq!(hits.load(std::sync::atomic::Ordering::SeqCst), 1);
649 }
650
651 use super::*;
652
653 async fn ok(_req: Request) -> &'static str {
654 "ok"
655 }
656
657 fn router_with(define: impl FnOnce(&mut Router)) -> Router {
658 let mut router = Router::new();
659 define(&mut router);
660 router.finalize();
661 router
662 }
663
664 #[tokio::test]
665 async fn matches_static_and_parameter_routes() {
666 let router = router_with(|r| {
667 r.get("/", ok);
668 r.get("/users/{id}", |req: Request| async move {
669 format!("user {}", req.param("id").unwrap())
670 });
671 });
672
673 assert_eq!(router.dispatch(Request::new(Method::Get, "/")).await.body_string(), "ok");
674 assert_eq!(
675 router.dispatch(Request::new(Method::Get, "/users/7")).await.body_string(),
676 "user 7"
677 );
678 assert_eq!(
679 router.dispatch(Request::new(Method::Get, "/nope")).await.status,
680 Status::NOT_FOUND
681 );
682 }
683
684 #[tokio::test]
685 async fn static_segments_win_over_parameters() {
686 let router = router_with(|r| {
687 r.get("/users/{id}", |_req: Request| async { "param" });
688 r.get("/users/new", |_req: Request| async { "static" });
689 });
690
691 assert_eq!(
692 router.dispatch(Request::new(Method::Get, "/users/new")).await.body_string(),
693 "static"
694 );
695 assert_eq!(
696 router.dispatch(Request::new(Method::Get, "/users/12")).await.body_string(),
697 "param"
698 );
699 }
700
701 #[tokio::test]
702 async fn wildcards_capture_the_remaining_path() {
703 let router = router_with(|r| {
704 r.get("/files/{path:*}", |req: Request| async move {
705 req.param("path").unwrap_or_default().to_string()
706 });
707 });
708
709 let response = router.dispatch(Request::new(Method::Get, "/files/css/app.css")).await;
710 assert_eq!(response.body_string(), "css/app.css");
711 }
712
713 #[tokio::test]
714 async fn parameters_are_percent_decoded() {
715 let router = router_with(|r| {
716 r.get("/tags/{tag}", |req: Request| async move { req.param("tag").unwrap().to_string() });
717 });
718
719 let response = router.dispatch(Request::new(Method::Get, "/tags/rust%20lang")).await;
720 assert_eq!(response.body_string(), "rust lang");
721 }
722
723 #[tokio::test]
724 async fn wrong_method_reports_405_with_allow() {
725 let router = router_with(|r| {
726 r.post("/users", ok);
727 });
728
729 let response = router.dispatch(Request::new(Method::Get, "/users")).await;
730 assert_eq!(response.status, Status::METHOD_NOT_ALLOWED);
731 assert_eq!(response.headers.get("allow"), Some("POST"));
732 }
733
734 #[tokio::test]
735 async fn head_is_served_by_the_get_route() {
736 let router = router_with(|r| {
737 r.get("/", ok);
738 });
739
740 assert_eq!(router.dispatch(Request::new(Method::Head, "/")).await.status, Status::OK);
741 }
742
743 #[tokio::test]
744 async fn groups_apply_prefix_and_middleware() {
745 let router = router_with(|r| {
746 r.group("/admin", |r| {
747 r.middleware(|req: Request, next: Next| async move {
748 next.run(req).await.with_header("x-guard", "on")
749 });
750 r.get("/dashboard", ok);
751 });
752 r.get("/public", ok);
753 });
754
755 let guarded = router.dispatch(Request::new(Method::Get, "/admin/dashboard")).await;
756 assert_eq!(guarded.headers.get("x-guard"), Some("on"));
757
758 let open = router.dispatch(Request::new(Method::Get, "/public")).await;
759 assert_eq!(open.headers.get("x-guard"), None);
760 }
761
762 #[tokio::test]
763 async fn global_middleware_also_sees_unmatched_requests() {
764 let router = router_with(|r| {
765 r.middleware(|req: Request, next: Next| async move {
766 next.run(req).await.with_header("x-seen", "1")
767 });
768 r.get("/", ok);
769 });
770
771 let missing = router.dispatch(Request::new(Method::Get, "/nope")).await;
772 assert_eq!(missing.status, Status::NOT_FOUND);
773 assert_eq!(missing.headers.get("x-seen"), Some("1"));
774 }
775
776 #[test]
777 fn builds_urls_from_named_routes() {
778 let router = router_with(|r| {
779 r.get("/users/{id}/posts/{slug}", ok).name("users.posts");
780 });
781
782 assert_eq!(
783 router.url_for("users.posts", &[("id", "7"), ("slug", "hello world")]).as_deref(),
784 Some("/users/7/posts/hello%20world")
785 );
786 assert_eq!(router.url_for("users.posts", &[("id", "7")]), None);
787 assert_eq!(router.url_for("missing", &[]), None);
788 }
789
790 #[tokio::test]
791 async fn resource_registers_the_rest_routes() {
792 let router = router_with(|r| {
793 r.resource("/posts").index(ok).store(ok).show(ok).update(ok).destroy(ok);
794 });
795
796 assert_eq!(router.routes().len(), 5);
797 assert_eq!(router.url_for("posts.show", &[("id", "3")]).as_deref(), Some("/posts/3"));
798 assert_eq!(router.dispatch(Request::new(Method::Delete, "/posts/3")).await.status, Status::OK);
799 }
800
801 #[tokio::test]
802 async fn fallback_replaces_the_default_404() {
803 let router = router_with(|r| {
804 r.fallback(|_req: Request| async { (404, "custom miss") });
805 });
806
807 let response = router.dispatch(Request::new(Method::Get, "/anything")).await;
808 assert_eq!(response.body_string(), "custom miss");
809 }
810
811 #[test]
812 fn documentation_rides_along_with_the_route() {
813 let router = router_with(|r| {
814 r.get("/users/{id}", ok)
815 .name("users.show")
816 .describe("Fetch one user")
817 .tag("Users")
818 .param("id", "The user's id")
819 .responds(200, "The user")
820 .responds(404, "No such user");
821 });
822
823 let route = &router.routes()[0];
824 assert_eq!(route.summary.as_deref(), Some("Fetch one user"));
825 assert_eq!(route.tag.as_deref(), Some("Users"));
826 assert_eq!(route.responses.len(), 2);
827 assert_eq!(route.parameter_names(), vec!["id"]);
828 assert!(!route.deprecated);
829 }
830
831 #[test]
832 fn path_parameters_are_known_even_when_undocumented() {
833 let router = router_with(|r| {
834 r.get("/teams/{team}/members/{member}", ok);
835 r.get("/files/{path:*}", ok);
836 });
837
838 let names: Vec<Vec<String>> =
839 router.routes().iter().map(Route::parameter_names).collect();
840 assert!(names.contains(&vec!["team".to_string(), "member".to_string()]));
841 assert!(names.contains(&vec!["path".to_string()]));
842 }
843
844 #[test]
845 fn joins_prefixes_without_doubling_slashes() {
846 assert_eq!(join_paths("/admin/", "/users"), "/admin/users");
847 assert_eq!(join_paths("", "/"), "/");
848 assert_eq!(join_paths("/admin", ""), "/admin");
849 }
850}