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)]
139pub struct Router {
140 routes: Vec<Route>,
141 scope_prefix: String,
143 scope_version: Option<String>,
144 scope_middleware: Vec<Arc<dyn Middleware>>,
145 global_middleware: Vec<Arc<dyn Middleware>>,
147 fallback: Option<Arc<dyn Handler>>,
148}
149
150impl Router {
151 pub fn new() -> Self {
152 Self::default()
153 }
154
155 pub fn middleware(&mut self, middleware: impl Middleware) -> &mut Self {
158 if self.scope_prefix.is_empty() && self.scope_middleware.is_empty() {
159 self.global_middleware.push(Arc::new(middleware));
160 } else {
161 self.scope_middleware.push(Arc::new(middleware));
162 }
163 self
164 }
165
166 pub fn group(&mut self, prefix: &str, define: impl FnOnce(&mut Router)) -> &mut Self {
168 let mut child = Router {
169 scope_prefix: join_paths(&self.scope_prefix, prefix),
170 scope_middleware: self.scope_middleware.clone(),
171 scope_version: self.scope_version.clone(),
172 ..Router::default()
173 };
174 define(&mut child);
175
176 debug_assert!(child.global_middleware.is_empty() || !child.routes.is_empty());
178 self.routes.extend(child.routes);
179 self
180 }
181
182 pub fn version(&mut self, version: &str, define: impl FnOnce(&mut Router)) -> &mut Self {
195 let mut child = Router {
196 scope_prefix: join_paths(&self.scope_prefix, &format!("/{}", version.trim_start_matches('/'))),
197 scope_middleware: self.scope_middleware.clone(),
198 scope_version: Some(version.trim_start_matches('/').to_string()),
199 ..Router::default()
200 };
201 define(&mut child);
202 self.routes.extend(child.routes);
203 self
204 }
205
206 pub fn fallback(&mut self, handler: impl Handler) -> &mut Self {
208 self.fallback = Some(Arc::new(handler));
209 self
210 }
211
212 pub fn route(&mut self, method: Method, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
213 let full = join_paths(&self.scope_prefix, pattern);
214 self.routes.push(Route {
215 method,
216 segments: parse_pattern(&full),
217 pattern: full,
218 name: None,
219 summary: None,
220 tag: None,
221 responses: Vec::new(),
222 parameters: Vec::new(),
223 deprecated: false,
224 version: self.scope_version.clone(),
225 deprecated_at: None,
226 sunset: None,
227 handler: Arc::new(handler),
228 middleware: Arc::new(self.scope_middleware.clone()),
229 });
230 RouteHandle { index: self.routes.len() - 1, router: self }
231 }
232
233 pub fn get(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
234 self.route(Method::Get, pattern, handler)
235 }
236
237 pub fn post(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
238 self.route(Method::Post, pattern, handler)
239 }
240
241 pub fn put(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
242 self.route(Method::Put, pattern, handler)
243 }
244
245 pub fn patch(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
246 self.route(Method::Patch, pattern, handler)
247 }
248
249 pub fn delete(&mut self, pattern: &str, handler: impl Handler) -> RouteHandle<'_> {
250 self.route(Method::Delete, pattern, handler)
251 }
252
253 pub fn resource<'r>(&'r mut self, base: &str) -> Resource<'r> {
255 Resource { base: base.trim_end_matches('/').to_string(), router: self }
256 }
257
258 pub fn finalize(&mut self) {
260 self.routes.sort_by_key(Route::specificity);
261 }
262
263 pub fn routes(&self) -> &[Route] {
264 &self.routes
265 }
266
267 pub fn url_for(&self, name: &str, params: &[(&str, &str)]) -> Option<String> {
269 let route = self.routes.iter().find(|route| route.name.as_deref() == Some(name))?;
270 let lookup = |key: &str| params.iter().find(|(k, _)| *k == key).map(|(_, v)| *v);
271
272 let mut out = String::new();
273 for segment in &route.segments {
274 out.push('/');
275 match segment {
276 Segment::Static(value) => out.push_str(value),
277 Segment::Param(name) => out.push_str(&url::encode(lookup(name)?)),
278 Segment::Wildcard(name) => out.push_str(lookup(name)?),
279 }
280 }
281 Some(if out.is_empty() { "/".to_string() } else { out })
282 }
283
284 pub async fn dispatch(&self, mut request: Request) -> Response {
286 let path = request.path().to_string();
287 let mut path_matched = false;
288 let mut allowed = Vec::new();
289
290 for route in &self.routes {
291 let Some(params) = route.match_path(&path) else { continue };
292 path_matched = true;
293
294 let usable = route.method == request.method()
296 || (request.method() == Method::Head && route.method == Method::Get);
297
298 if !usable {
299 allowed.push(route.method);
300 continue;
301 }
302
303 request.set_params(params);
304 request.route = Some(route.pattern.clone());
305 if let Some(version) = &route.version {
306 request.extend(crate::versioning::ApiVersion(version.clone()));
307 }
308
309 let mut stack = self.global_middleware.clone();
310 stack.extend(route.middleware.iter().cloned());
311 let response =
312 run_guarded(Next::new(Arc::new(stack), Arc::clone(&route.handler)), request).await;
313 return crate::versioning::stamp_lifecycle(route, response);
314 }
315
316 let response = if path_matched {
317 allowed.sort();
319 allowed.dedup();
320 let allow = allowed.iter().map(|m| m.as_str()).collect::<Vec<_>>().join(", ");
321 Response::new(Status::METHOD_NOT_ALLOWED).with_header("allow", allow).with_text(format!(
322 "{} is not allowed on {path}",
323 request.method()
324 ))
325 } else {
326 match &self.fallback {
327 Some(handler) => {
328 let stack = Arc::new(self.global_middleware.clone());
329 return Next::new(stack, Arc::clone(handler)).run(request).await;
330 }
331 None => Response::not_found(),
332 }
333 };
334
335 let stack = Arc::new(self.global_middleware.clone());
338 let endpoint: Arc<dyn Handler> = Arc::new(crate::handler::Fixed(response));
339 Next::new(stack, endpoint).run(request).await
340 }
341}
342
343async fn run_guarded(next: Next, request: Request) -> Response {
349 crate::panic::install_hook();
350 let started = std::time::Instant::now();
351
352 let probe = Request::new(request.method(), request.target().to_string())
355 .with_header("accept", request.header("accept").unwrap_or("text/html"));
356 let route = request.route().map(str::to_string);
357
358 let response = match crate::panic::catch(next.run(request)).await {
359 Ok(response) => response,
360 Err(message) => {
361 let location = crate::panic::take_location().map(|l| (l.file, l.line));
362 rustlavel_core::error!(
363 "panic in {} {}: {message}",
364 probe.method(),
365 probe.path()
366 );
367 crate::error_page::render(
368 &crate::error_page::Diagnostic::from_panic(message, location),
369 Some(&probe),
370 )
371 }
372 };
373
374 if rustlavel_core::events::has_subscribers() {
377 let mut event = rustlavel_core::Event::new("http.request")
378 .with("method", probe.method().as_str())
379 .with("path", probe.path())
380 .with("status", response.status.code())
381 .took(started.elapsed());
382 if let Some(route) = route {
383 event = event.with("route", route);
385 }
386 if let Some(id) = response.headers.get(crate::request_id::HEADER) {
389 event = event.with("request_id", id);
390 }
391 event.dispatch();
392 }
393
394 response
395}
396
397pub struct RouteHandle<'r> {
399 index: usize,
400 router: &'r mut Router,
401}
402
403impl RouteHandle<'_> {
404 pub fn name(self, name: &str) -> Self {
406 self.router.routes[self.index].name = Some(name.to_string());
407 self
408 }
409
410 pub fn describe(self, summary: &str) -> Self {
415 self.router.routes[self.index].summary = Some(summary.to_string());
416 self
417 }
418
419 pub fn tag(self, tag: &str) -> Self {
421 self.router.routes[self.index].tag = Some(tag.to_string());
422 self
423 }
424
425 pub fn responds(self, status: u16, description: &str) -> Self {
427 self.router.routes[self.index].responses.push((status, description.to_string()));
428 self
429 }
430
431 pub fn param(self, name: &str, description: &str) -> Self {
434 self.router.routes[self.index]
435 .parameters
436 .push((name.to_string(), description.to_string()));
437 self
438 }
439
440 pub fn middleware(self, middleware: impl Middleware) -> Self {
451 let stack = &mut self.router.routes[self.index].middleware;
452 Arc::make_mut(stack).push(Arc::new(middleware));
456 self
457 }
458
459 pub fn deprecated(self) -> Self {
461 self.router.routes[self.index].deprecated = true;
462 self
463 }
464
465 pub fn deprecated_at(self, date: &str) -> Self {
476 let when = crate::date::parse_ymd(date)
477 .unwrap_or_else(|| panic!("`{date}` is not a date; deprecated_at wants YYYY-MM-DD"));
478 let route = &mut self.router.routes[self.index];
479 route.deprecated = true;
480 route.deprecated_at = Some(when);
481 self
482 }
483
484 pub fn sunset(self, date: &str) -> Self {
495 let when = crate::date::parse_ymd(date)
496 .unwrap_or_else(|| panic!("`{date}` is not a date; sunset wants YYYY-MM-DD"));
497 let route = &mut self.router.routes[self.index];
498 route.deprecated = true;
499 route.sunset = Some(when);
500 self
501 }
502}
503
504pub struct Resource<'r> {
506 base: String,
507 router: &'r mut Router,
508}
509
510impl Resource<'_> {
511 pub fn index(self, handler: impl Handler) -> Self {
513 let (base, router) = (self.base.clone(), self.router);
514 router.get(&base, handler).name(&format!("{}.index", resource_name(&base)));
515 Resource { base, router }
516 }
517
518 pub fn store(self, handler: impl Handler) -> Self {
520 let (base, router) = (self.base.clone(), self.router);
521 router.post(&base, handler).name(&format!("{}.store", resource_name(&base)));
522 Resource { base, router }
523 }
524
525 pub fn show(self, handler: impl Handler) -> Self {
527 let (base, router) = (self.base.clone(), self.router);
528 let pattern = format!("{base}/{{id}}");
529 router.get(&pattern, handler).name(&format!("{}.show", resource_name(&base)));
530 Resource { base, router }
531 }
532
533 pub fn update(self, handler: impl Handler) -> Self {
535 let (base, router) = (self.base.clone(), self.router);
536 let pattern = format!("{base}/{{id}}");
537 router.put(&pattern, handler).name(&format!("{}.update", resource_name(&base)));
538 Resource { base, router }
539 }
540
541 pub fn destroy(self, handler: impl Handler) -> Self {
543 let (base, router) = (self.base.clone(), self.router);
544 let pattern = format!("{base}/{{id}}");
545 router.delete(&pattern, handler).name(&format!("{}.destroy", resource_name(&base)));
546 Resource { base, router }
547 }
548}
549
550fn resource_name(base: &str) -> String {
551 base.trim_matches('/').replace('/', ".")
552}
553
554fn join_paths(prefix: &str, path: &str) -> String {
555 let joined = format!("/{}/{}", prefix.trim_matches('/'), path.trim_matches('/'));
556 let cleaned = joined.replace("//", "/");
557 if cleaned.len() > 1 { cleaned.trim_end_matches('/').to_string() } else { "/".to_string() }
558}
559
560#[cfg(test)]
561mod tests {
562 #[tokio::test]
563 async fn per_route_middleware_does_not_leak_to_its_neighbours() {
564 use crate::testing::TestClient;
565 fn tag(name: &'static str) -> impl Middleware {
566 move |request: Request, next: Next| async move {
567 next.run(request).await.with_header("x-tag", name)
568 }
569 }
570
571 let mut router = Router::new();
572 router.group("/admin", |admin| {
573 admin.middleware(tag("group"));
574 admin.get("/users", |_req: Request| async { Response::text("users") }).middleware(tag("view"));
575 admin.get("/roles", |_req: Request| async { Response::text("roles") });
576 });
577 let client = TestClient::new(router);
578
579 assert_eq!(client.get("/admin/users").await.status(), 200);
582 assert_eq!(client.get("/admin/roles").await.status(), 200);
583
584 let mut counted = Router::new();
585 let hits = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
586 let counter = hits.clone();
587 counted.group("/admin", |admin| {
588 admin
589 .get("/users", |_req: Request| async { Response::text("users") })
590 .middleware(move |request: Request, next: Next| {
591 let counter = counter.clone();
592 async move {
593 counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
594 next.run(request).await
595 }
596 });
597 admin.get("/roles", |_req: Request| async { Response::text("roles") });
598 });
599 let client = TestClient::new(counted);
600 client.get("/admin/roles").await;
601 assert_eq!(hits.load(std::sync::atomic::Ordering::SeqCst), 0, "the guard belongs to /users only");
602 client.get("/admin/users").await;
603 assert_eq!(hits.load(std::sync::atomic::Ordering::SeqCst), 1);
604 }
605
606 use super::*;
607
608 async fn ok(_req: Request) -> &'static str {
609 "ok"
610 }
611
612 fn router_with(define: impl FnOnce(&mut Router)) -> Router {
613 let mut router = Router::new();
614 define(&mut router);
615 router.finalize();
616 router
617 }
618
619 #[tokio::test]
620 async fn matches_static_and_parameter_routes() {
621 let router = router_with(|r| {
622 r.get("/", ok);
623 r.get("/users/{id}", |req: Request| async move {
624 format!("user {}", req.param("id").unwrap())
625 });
626 });
627
628 assert_eq!(router.dispatch(Request::new(Method::Get, "/")).await.body_string(), "ok");
629 assert_eq!(
630 router.dispatch(Request::new(Method::Get, "/users/7")).await.body_string(),
631 "user 7"
632 );
633 assert_eq!(
634 router.dispatch(Request::new(Method::Get, "/nope")).await.status,
635 Status::NOT_FOUND
636 );
637 }
638
639 #[tokio::test]
640 async fn static_segments_win_over_parameters() {
641 let router = router_with(|r| {
642 r.get("/users/{id}", |_req: Request| async { "param" });
643 r.get("/users/new", |_req: Request| async { "static" });
644 });
645
646 assert_eq!(
647 router.dispatch(Request::new(Method::Get, "/users/new")).await.body_string(),
648 "static"
649 );
650 assert_eq!(
651 router.dispatch(Request::new(Method::Get, "/users/12")).await.body_string(),
652 "param"
653 );
654 }
655
656 #[tokio::test]
657 async fn wildcards_capture_the_remaining_path() {
658 let router = router_with(|r| {
659 r.get("/files/{path:*}", |req: Request| async move {
660 req.param("path").unwrap_or_default().to_string()
661 });
662 });
663
664 let response = router.dispatch(Request::new(Method::Get, "/files/css/app.css")).await;
665 assert_eq!(response.body_string(), "css/app.css");
666 }
667
668 #[tokio::test]
669 async fn parameters_are_percent_decoded() {
670 let router = router_with(|r| {
671 r.get("/tags/{tag}", |req: Request| async move { req.param("tag").unwrap().to_string() });
672 });
673
674 let response = router.dispatch(Request::new(Method::Get, "/tags/rust%20lang")).await;
675 assert_eq!(response.body_string(), "rust lang");
676 }
677
678 #[tokio::test]
679 async fn wrong_method_reports_405_with_allow() {
680 let router = router_with(|r| {
681 r.post("/users", ok);
682 });
683
684 let response = router.dispatch(Request::new(Method::Get, "/users")).await;
685 assert_eq!(response.status, Status::METHOD_NOT_ALLOWED);
686 assert_eq!(response.headers.get("allow"), Some("POST"));
687 }
688
689 #[tokio::test]
690 async fn head_is_served_by_the_get_route() {
691 let router = router_with(|r| {
692 r.get("/", ok);
693 });
694
695 assert_eq!(router.dispatch(Request::new(Method::Head, "/")).await.status, Status::OK);
696 }
697
698 #[tokio::test]
699 async fn groups_apply_prefix_and_middleware() {
700 let router = router_with(|r| {
701 r.group("/admin", |r| {
702 r.middleware(|req: Request, next: Next| async move {
703 next.run(req).await.with_header("x-guard", "on")
704 });
705 r.get("/dashboard", ok);
706 });
707 r.get("/public", ok);
708 });
709
710 let guarded = router.dispatch(Request::new(Method::Get, "/admin/dashboard")).await;
711 assert_eq!(guarded.headers.get("x-guard"), Some("on"));
712
713 let open = router.dispatch(Request::new(Method::Get, "/public")).await;
714 assert_eq!(open.headers.get("x-guard"), None);
715 }
716
717 #[tokio::test]
718 async fn global_middleware_also_sees_unmatched_requests() {
719 let router = router_with(|r| {
720 r.middleware(|req: Request, next: Next| async move {
721 next.run(req).await.with_header("x-seen", "1")
722 });
723 r.get("/", ok);
724 });
725
726 let missing = router.dispatch(Request::new(Method::Get, "/nope")).await;
727 assert_eq!(missing.status, Status::NOT_FOUND);
728 assert_eq!(missing.headers.get("x-seen"), Some("1"));
729 }
730
731 #[test]
732 fn builds_urls_from_named_routes() {
733 let router = router_with(|r| {
734 r.get("/users/{id}/posts/{slug}", ok).name("users.posts");
735 });
736
737 assert_eq!(
738 router.url_for("users.posts", &[("id", "7"), ("slug", "hello world")]).as_deref(),
739 Some("/users/7/posts/hello%20world")
740 );
741 assert_eq!(router.url_for("users.posts", &[("id", "7")]), None);
742 assert_eq!(router.url_for("missing", &[]), None);
743 }
744
745 #[tokio::test]
746 async fn resource_registers_the_rest_routes() {
747 let router = router_with(|r| {
748 r.resource("/posts").index(ok).store(ok).show(ok).update(ok).destroy(ok);
749 });
750
751 assert_eq!(router.routes().len(), 5);
752 assert_eq!(router.url_for("posts.show", &[("id", "3")]).as_deref(), Some("/posts/3"));
753 assert_eq!(router.dispatch(Request::new(Method::Delete, "/posts/3")).await.status, Status::OK);
754 }
755
756 #[tokio::test]
757 async fn fallback_replaces_the_default_404() {
758 let router = router_with(|r| {
759 r.fallback(|_req: Request| async { (404, "custom miss") });
760 });
761
762 let response = router.dispatch(Request::new(Method::Get, "/anything")).await;
763 assert_eq!(response.body_string(), "custom miss");
764 }
765
766 #[test]
767 fn documentation_rides_along_with_the_route() {
768 let router = router_with(|r| {
769 r.get("/users/{id}", ok)
770 .name("users.show")
771 .describe("Fetch one user")
772 .tag("Users")
773 .param("id", "The user's id")
774 .responds(200, "The user")
775 .responds(404, "No such user");
776 });
777
778 let route = &router.routes()[0];
779 assert_eq!(route.summary.as_deref(), Some("Fetch one user"));
780 assert_eq!(route.tag.as_deref(), Some("Users"));
781 assert_eq!(route.responses.len(), 2);
782 assert_eq!(route.parameter_names(), vec!["id"]);
783 assert!(!route.deprecated);
784 }
785
786 #[test]
787 fn path_parameters_are_known_even_when_undocumented() {
788 let router = router_with(|r| {
789 r.get("/teams/{team}/members/{member}", ok);
790 r.get("/files/{path:*}", ok);
791 });
792
793 let names: Vec<Vec<String>> =
794 router.routes().iter().map(Route::parameter_names).collect();
795 assert!(names.contains(&vec!["team".to_string(), "member".to_string()]));
796 assert!(names.contains(&vec!["path".to_string()]));
797 }
798
799 #[test]
800 fn joins_prefixes_without_doubling_slashes() {
801 assert_eq!(join_paths("/admin/", "/users"), "/admin/users");
802 assert_eq!(join_paths("", "/"), "/");
803 assert_eq!(join_paths("/admin", ""), "/admin");
804 }
805}