1use std::collections::HashMap;
2use std::fmt::{Display, Write};
3
4use std::convert::Infallible;
5
6use anyhow::{anyhow, bail};
7use axum::Router;
8use axum::extract::Request;
9use axum::handler::Handler;
10use axum::middleware::from_fn;
11use axum::response::IntoResponse;
12use axum::routing::{self, MethodRouter, Route};
13use tower::{Layer, Service};
14
15use crate::AppState;
16
17#[derive(Default)]
32pub struct Routes {
33 router: Router<AppState>,
34 names: Vec<(String, String)>,
35 last_path: Option<String>,
36 listing: Vec<RouteInfo>,
37 fallback: Option<MethodRouter<AppState>>,
38 domains: Vec<(String, Routes)>,
39}
40
41pub(crate) struct RouteParts {
43 pub router: Router<AppState>,
44 pub names: Vec<(String, String)>,
45 pub listing: Vec<RouteInfo>,
46 pub fallback: Option<MethodRouter<AppState>>,
47 pub domains: Vec<(String, Routes)>,
48}
49
50#[derive(Debug, Clone, PartialEq, Eq)]
52#[non_exhaustive]
53pub struct RouteInfo {
54 pub method: String,
56 pub path: String,
58 pub name: Option<String>,
60 pub module: String,
62 pub middleware: Vec<String>,
64 pub domain: Option<String>,
66}
67
68macro_rules! method {
69 ($($method:ident),*) => {$(
70 pub fn $method<H, T>(self, path: &str, handler: H) -> Self
72 where
73 H: Handler<T, AppState>,
74 T: 'static,
75 {
76 self.add(path, routing::$method(handler), &stringify!($method).to_uppercase())
77 }
78 )*};
79}
80
81impl Routes {
82 pub fn new() -> Self {
84 Self::default()
85 }
86
87 method!(get, post, put, patch, delete);
88
89 pub fn resource(mut self, path: &str, name: &str, resource: Resource) -> Self {
115 assert!(
116 path.starts_with('/') && !path.ends_with('/'),
117 "Routes::resource(\"{path}\", …): the path must start with `/` and not end with one"
118 );
119 let member = format!("{path}/{{id}}");
120 let actions = [
121 ("index", "GET", path.to_owned(), resource.index),
122 ("create", "GET", format!("{path}/new"), resource.create),
123 ("store", "POST", path.to_owned(), resource.store),
124 ("show", "GET", member.clone(), resource.show),
125 ("edit", "GET", format!("{member}/edit"), resource.edit),
126 ("update", "PUT", member.clone(), resource.update),
127 ("destroy", "DELETE", member, resource.destroy),
128 ];
129 for (action, method, at, router) in actions {
130 if let Some(router) = router {
131 let method = if action == "update" {
132 "PUT|PATCH"
133 } else {
134 method
135 };
136 self = self
137 .add(&at, router, method)
138 .name(&format!("{name}.{action}"));
139 }
140 }
141 self
142 }
143
144 pub fn route(self, path: &str, method_router: MethodRouter<AppState>) -> Self {
146 self.add(path, method_router, "*")
147 }
148
149 pub fn view(self, path: &str, template: &str) -> Self {
159 let template = template.to_owned();
160 self.add(
161 path,
162 axum::routing::get(move || {
163 let template = template.clone();
164 async move { crate::view(&template, minijinja::context! {}) }
165 }),
166 "GET",
167 )
168 }
169
170 pub fn redirect(self, path: &str, to: &str) -> Self {
173 let to = to.to_owned();
174 self.add(
175 path,
176 axum::routing::any(move || {
177 let to = to.clone();
178 async move { redirect_with(axum::http::StatusCode::FOUND, &to) }
179 }),
180 "*",
181 )
182 }
183
184 pub fn permanent_redirect(self, path: &str, to: &str) -> Self {
186 let to = to.to_owned();
187 self.add(
188 path,
189 axum::routing::any(move || {
190 let to = to.clone();
191 async move { redirect_with(axum::http::StatusCode::MOVED_PERMANENTLY, &to) }
192 }),
193 "*",
194 )
195 }
196
197 fn add(mut self, path: &str, method_router: MethodRouter<AppState>, method: &str) -> Self {
198 self.router = self.router.route(path, method_router);
199 self.last_path = Some(path.to_owned());
200 self.listing.push(RouteInfo {
201 method: method.to_owned(),
202 path: path.to_owned(),
203 name: None,
204 module: String::new(),
205 middleware: Vec::new(),
206 domain: None,
207 });
208 self
209 }
210
211 fn mark(mut self, middleware: &str) -> Self {
213 for route in &mut self.listing {
214 route.middleware.push(middleware.to_owned());
215 }
216 self
217 }
218
219 pub fn name(mut self, name: &str) -> Self {
225 let path = self
226 .last_path
227 .clone()
228 .expect("Routes::name() must follow a route");
229 self.names.push((name.to_owned(), path));
230 let path = self.last_path.clone();
234 for route in self.listing.iter_mut().rev() {
235 if Some(&route.path) != path.as_ref() {
236 break;
237 }
238 if route.name.is_none() {
239 route.name = Some(name.to_owned());
240 }
241 }
242 self
243 }
244
245 pub fn require_auth(self) -> Self {
248 self.route_layer(from_fn(crate::auth::require_auth))
249 .mark("auth")
250 }
251
252 pub fn require_verified(self) -> Self {
255 self.route_layer(from_fn(crate::auth::require_verified))
256 .mark("verified")
257 }
258
259 pub fn require_gate(self, name: &str) -> Self {
263 self.requirement(
264 crate::auth::Requirement::Gate(name.to_owned()),
265 "gate",
266 name,
267 )
268 }
269
270 pub fn require_role(self, role: &str) -> Self {
273 self.requirement(
274 crate::auth::Requirement::Role(role.to_owned()),
275 "role",
276 role,
277 )
278 }
279
280 pub fn require_permission(self, permission: &str) -> Self {
283 self.requirement(
284 crate::auth::Requirement::Permission(permission.to_owned()),
285 "permission",
286 permission,
287 )
288 }
289
290 pub fn require_ability(self, ability: &str) -> Self {
293 self.requirement(
294 crate::auth::Requirement::Ability(ability.to_owned()),
295 "ability",
296 ability,
297 )
298 }
299
300 fn requirement(self, requirement: crate::auth::Requirement, kind: &str, name: &str) -> Self {
301 let requirement = std::sync::Arc::new(requirement);
302 self.route_layer(from_fn(
303 move |req: Request, next: axum::middleware::Next| {
304 crate::auth::require(requirement.clone(), req, next)
305 },
306 ))
307 .mark(&format!("{kind}:{name}"))
308 }
309
310 pub fn require_password_confirmed(self) -> Self {
314 self.route_layer(from_fn(crate::auth::require_password_confirmed))
315 .route_layer(from_fn(crate::auth::require_auth))
316 .mark("password.confirm")
317 }
318
319 pub fn guest_only(self) -> Self {
322 self.route_layer(from_fn(crate::auth::guest_only))
323 .mark("guest")
324 }
325
326 pub fn throttle(self, max: u32, per: std::time::Duration) -> Self {
329 let covered: Vec<String> = self
330 .listing
331 .iter()
332 .map(|route| format!("{} {}", route.method, route.path))
333 .collect();
334 let id =
335 crate::webhook::sha256_hex(format!("{}|{max}|{}", covered.join(","), per.as_secs()));
336 let limiter = std::sync::Arc::new(crate::rate_limit::Limiter::new(
337 id[..16].to_owned(),
338 max,
339 per,
340 ));
341 self.route_layer(from_fn(
342 move |req: Request, next: axum::middleware::Next| {
343 let limiter = limiter.clone();
344 async move { crate::rate_limit::check(&limiter, req, next).await }
345 },
346 ))
347 .mark(&format!("throttle:{max}/{}s", per.as_secs()))
348 }
349
350 pub fn throttle_by(self, name: &str) -> Self {
353 let owned = name.to_owned();
354 self.route_layer(from_fn(
355 move |req: Request, next: axum::middleware::Next| {
356 let name = owned.clone();
357 async move { crate::rate_limit::check_named(&name, req, next).await }
358 },
359 ))
360 .mark(&format!("throttle:{name}"))
361 }
362
363 pub fn webhook<W: crate::webhook::Webhook>(self, path: &str) -> Self {
368 let mut routes = self
369 .post(path, crate::webhook::receive::<W>)
370 .name(&format!("webhooks.{}", W::PROVIDER));
371 if let Some(route) = routes.listing.last_mut() {
372 route.middleware.push("no-csrf".into());
373 route.middleware.push(format!("webhook:{}", W::PROVIDER));
374 }
375 routes
376 }
377
378 pub fn etag(self) -> Self {
385 self.route_layer(from_fn(
386 |req: Request, next: axum::middleware::Next| async move {
387 let mut res = next.run(req).await;
388 res.extensions_mut().insert(crate::security::WantsEtag);
389 res
390 },
391 ))
392 .mark("etag")
393 }
394
395 pub fn without_csrf(self) -> Self {
399 self.mark("no-csrf")
400 }
401
402 pub fn cors(self, origins: &[&str]) -> Self {
408 use axum::http::{HeaderName, HeaderValue, Method, header};
409 use tower_http::cors::{AllowOrigin, CorsLayer};
410
411 let origin = if origins.contains(&"*") {
412 AllowOrigin::any()
413 } else {
414 AllowOrigin::list(
415 origins
416 .iter()
417 .filter_map(|o| HeaderValue::from_str(o.trim_end_matches('/')).ok()),
418 )
419 };
420 let layer = CorsLayer::new()
421 .allow_origin(origin)
422 .allow_methods([
423 Method::GET,
424 Method::POST,
425 Method::PUT,
426 Method::PATCH,
427 Method::DELETE,
428 Method::OPTIONS,
429 ])
430 .allow_headers([
431 header::CONTENT_TYPE,
432 header::AUTHORIZATION,
433 header::ACCEPT,
434 HeaderName::from_static(crate::CSRF_HEADER),
435 ])
436 .max_age(std::time::Duration::from_secs(3600));
437 self.cors_layer(layer)
438 }
439
440 pub fn cors_layer(mut self, layer: tower_http::cors::CorsLayer) -> Self {
443 self.router = self.router.layer(layer);
446 self.mark("cors")
447 }
448
449 pub fn route_layer<L>(mut self, layer: L) -> Self
451 where
452 L: Layer<Route> + Clone + Send + Sync + 'static,
453 L::Service: Service<Request> + Clone + Send + Sync + 'static,
454 <L::Service as Service<Request>>::Response: IntoResponse + 'static,
455 <L::Service as Service<Request>>::Error: Into<Infallible> + 'static,
456 <L::Service as Service<Request>>::Future: Send + 'static,
457 {
458 self.router = self.router.route_layer(layer);
459 self
460 }
461
462 pub fn merge(mut self, other: impl Into<Routes>) -> Self {
464 let other = other.into();
465 self.router = self.router.merge(other.router);
466 self.names.extend(other.names);
467 self.listing.extend(other.listing);
468 self.domains.extend(other.domains);
469 assert!(
470 self.fallback.is_none() || other.fallback.is_none(),
471 "Routes::merge: both sides have a fallback"
472 );
473 self.fallback = self.fallback.or(other.fallback);
474 self.last_path = None;
475 self
476 }
477
478 pub fn fallback<H, T>(mut self, handler: H) -> Self
492 where
493 H: Handler<T, AppState>,
494 T: 'static,
495 {
496 assert!(
497 self.fallback.is_none(),
498 "Routes::fallback: these routes have a fallback already"
499 );
500 self.fallback = Some(routing::any(handler));
501 self
502 }
503
504 pub fn domain(mut self, pattern: &str, routes: impl Into<Routes>) -> Self {
530 self.domains.push((pattern.to_owned(), routes.into()));
531 self.last_path = None;
532 self
533 }
534
535 pub fn group(mut self, path: &str, name: &str, routes: impl Into<Routes>) -> Self {
560 assert!(
561 path.starts_with('/') && !path.ends_with('/'),
562 "Routes::group(\"{path}\", …): the prefix must start with `/` and not end with one"
563 );
564 let routes = routes.into();
565 assert!(
566 routes.domains.is_empty() && routes.fallback.is_none(),
567 "Routes::group(\"{path}\", …): put `domain` and `fallback` outside path groups"
568 );
569 let join = |inner: &str| {
570 if inner == "/" {
571 path.to_owned()
572 } else {
573 format!("{path}{inner}")
574 }
575 };
576 self.router = self.router.nest(path, routes.router);
577 self.names.extend(
578 routes
579 .names
580 .into_iter()
581 .map(|(route_name, route_path)| (format!("{name}{route_name}"), join(&route_path))),
582 );
583 self.listing
584 .extend(routes.listing.into_iter().map(|info| RouteInfo {
585 path: join(&info.path),
586 name: info.name.map(|route_name| format!("{name}{route_name}")),
587 ..info
588 }));
589 self.last_path = None;
590 self
591 }
592
593 pub(crate) fn into_parts(self) -> RouteParts {
594 RouteParts {
595 router: self.router,
596 names: self.names,
597 listing: self.listing,
598 fallback: self.fallback,
599 domains: self.domains,
600 }
601 }
602}
603
604impl From<Router<AppState>> for Routes {
605 fn from(router: Router<AppState>) -> Self {
606 Self {
607 router,
608 ..Self::default()
609 }
610 }
611}
612
613#[derive(Debug, Default)]
615pub struct RouteTable {
616 paths: HashMap<String, String>,
617 domains: HashMap<String, String>,
619 methods: HashMap<String, String>,
621}
622
623impl RouteTable {
624 pub(crate) fn insert(&mut self, name: String, path: String) -> anyhow::Result<()> {
625 if let Some(existing) = self.paths.get(&name) {
626 bail!("route name `{name}` is used for both `{existing}` and `{path}`");
627 }
628 self.paths.insert(name, path);
629 Ok(())
630 }
631
632 pub fn path(&self, name: &str) -> Option<&str> {
634 self.paths.get(name).map(String::as_str)
635 }
636
637 pub(crate) fn set_method(&mut self, name: &str, method: &str) {
639 self.methods.insert(name.to_owned(), method.to_owned());
640 }
641
642 pub(crate) fn set_domain(&mut self, name: &str, domain: &str) {
644 self.domains.insert(name.to_owned(), domain.to_owned());
645 }
646
647 pub(crate) fn name_of(
654 &self,
655 path: &str,
656 domain: Option<&str>,
657 method: Option<&str>,
658 ) -> Option<&str> {
659 let on_path = || {
660 self.paths
661 .iter()
662 .filter(move |(name, p)| {
663 p.as_str() == path && self.domains.get(*name).map(String::as_str) == domain
664 })
665 .map(|(name, _)| name.as_str())
666 };
667 let wanted = method.map(|m| if m == "HEAD" { "GET" } else { m });
668 let by_method = wanted.and_then(|wanted| {
669 on_path()
670 .filter(|name| {
671 self.methods.get(*name).is_some_and(|listed| {
672 listed == "*" || listed.split('|').any(|m| m == wanted)
673 })
674 })
675 .min()
676 });
677 by_method.or_else(|| on_path().min())
678 }
679
680 pub fn url(&self, name: &str, params: &[&dyn Display]) -> anyhow::Result<String> {
682 let pattern = self
683 .path(name)
684 .ok_or_else(|| anyhow!("route `{name}` is not defined"))?;
685 fill(pattern, params).map_err(|err| anyhow!("route `{name}`: {err}"))
686 }
687}
688
689fn fill(pattern: &str, params: &[&dyn Display]) -> anyhow::Result<String> {
690 let mut out = String::with_capacity(pattern.len());
691 let mut params = params.iter();
692 let mut rest = pattern;
693
694 while let Some(start) = rest.find('{') {
695 out.push_str(&rest[..start]);
696 let end = rest[start..]
697 .find('}')
698 .ok_or_else(|| anyhow!("unclosed `{{` in `{pattern}`"))?
699 + start;
700 let placeholder = &rest[start + 1..end];
701 let value = params
702 .next()
703 .ok_or_else(|| anyhow!("missing value for `{{{placeholder}}}`"))?
704 .to_string();
705 encode(&mut out, &value, placeholder.starts_with('*'));
706 rest = &rest[end + 1..];
707 }
708 out.push_str(rest);
709
710 if params.next().is_some() {
711 bail!("too many parameters for `{pattern}`");
712 }
713 Ok(out)
714}
715
716pub(crate) fn encode(out: &mut String, value: &str, keep_slashes: bool) {
717 for byte in value.bytes() {
718 match byte {
719 b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~' => {
720 out.push(byte as char)
721 }
722 b'/' if keep_slashes => out.push('/'),
723 _ => {
724 let _ = write!(out, "%{byte:02X}");
725 }
726 }
727 }
728}
729
730#[cfg(test)]
731mod tests {
732 use super::*;
733
734 fn table() -> RouteTable {
735 let mut t = RouteTable::default();
736 t.insert("home".into(), "/".into()).unwrap();
737 t.insert("products.show".into(), "/products/{id}".into())
738 .unwrap();
739 t.insert("docs".into(), "/docs/{*path}".into()).unwrap();
740 t
741 }
742
743 #[test]
744 fn builds_urls() {
745 let t = table();
746 assert_eq!(t.url("home", &[]).unwrap(), "/");
747 assert_eq!(t.url("products.show", &[&42]).unwrap(), "/products/42");
748 assert_eq!(
749 t.url("products.show", &[&"a b/c"]).unwrap(),
750 "/products/a%20b%2Fc"
751 );
752 assert_eq!(
753 t.url("docs", &[&"guide/intro"]).unwrap(),
754 "/docs/guide/intro"
755 );
756 }
757
758 #[test]
759 fn rejects_bad_calls() {
760 let t = table();
761 assert!(t.url("missing", &[]).is_err());
762 assert!(t.url("products.show", &[]).is_err());
763 assert!(t.url("home", &[&1]).is_err());
764 }
765
766 #[test]
767 fn rejects_duplicate_names() {
768 let mut t = table();
769 assert!(t.insert("home".into(), "/home".into()).is_err());
770 }
771}
772
773#[derive(Default)]
776#[must_use = "a resource does nothing until given to Routes::resource"]
777pub struct Resource {
778 index: Option<MethodRouter<AppState>>,
779 create: Option<MethodRouter<AppState>>,
780 store: Option<MethodRouter<AppState>>,
781 show: Option<MethodRouter<AppState>>,
782 edit: Option<MethodRouter<AppState>>,
783 update: Option<MethodRouter<AppState>>,
784 destroy: Option<MethodRouter<AppState>>,
785}
786
787macro_rules! resource_action {
788 ($($action:ident => $router:expr),* $(,)?) => {$(
789 pub fn $action<H, T>(mut self, handler: H) -> Self
791 where
792 H: Handler<T, AppState>,
793 T: 'static,
794 {
795 let make: fn(H) -> MethodRouter<AppState> = $router;
796 self.$action = Some(make(handler));
797 self
798 }
799 )*};
800}
801
802impl Resource {
803 pub fn new() -> Self {
805 Self::default()
806 }
807
808 resource_action!(
809 index => routing::get,
810 create => routing::get,
811 store => routing::post,
812 show => routing::get,
813 edit => routing::get,
814 update => |h| routing::put(h.clone()).patch(h),
815 destroy => routing::delete,
816 );
817}
818
819pub(crate) fn route_name_matches(name: &str, pattern: &str) -> bool {
822 let mut parts = pattern.split('*');
823 let first = parts.next().unwrap_or_default();
824 let Some(mut rest) = name.strip_prefix(first) else {
825 return false;
826 };
827 let parts: Vec<&str> = parts.collect();
828 let Some((last, middle)) = parts.split_last() else {
829 return rest.is_empty(); };
831 for part in middle {
832 match rest.find(part) {
833 Some(at) => rest = &rest[at + part.len()..],
834 None => return false,
835 }
836 }
837 rest.ends_with(last)
838}
839
840#[derive(Debug, Clone)]
854pub struct CurrentRoute {
855 name: Option<String>,
856 path: Option<String>,
857}
858
859impl CurrentRoute {
860 pub fn name(&self) -> Option<&str> {
862 self.name.as_deref()
863 }
864
865 pub fn path(&self) -> Option<&str> {
867 self.path.as_deref()
868 }
869
870 pub fn is(&self, pattern: &str) -> bool {
873 self.name
874 .as_deref()
875 .is_some_and(|name| route_name_matches(name, pattern))
876 }
877
878 pub(crate) fn of(
879 method: &axum::http::Method,
880 extensions: &axum::http::Extensions,
881 state: &AppState,
882 ) -> Self {
883 let path = extensions
884 .get::<axum::extract::MatchedPath>()
885 .map(|p| p.as_str().to_owned());
886 let domain = extensions
887 .get::<crate::domain::MatchedDomain>()
888 .map(|d| &*d.0);
889 let name = path
890 .as_deref()
891 .and_then(|p| state.routes.name_of(p, domain, Some(method.as_str())))
892 .map(str::to_owned);
893 Self { name, path }
894 }
895}
896
897impl<S: Send + Sync> axum::extract::FromRequestParts<S> for CurrentRoute {
898 type Rejection = Infallible;
899
900 async fn from_request_parts(
901 parts: &mut axum::http::request::Parts,
902 _: &S,
903 ) -> Result<Self, Infallible> {
904 Ok(match parts.extensions.get::<AppState>() {
905 Some(state) => Self::of(&parts.method, &parts.extensions, state),
906 None => Self {
907 name: None,
908 path: None,
909 },
910 })
911 }
912}
913
914#[cfg(test)]
915mod route_name_tests {
916 use super::route_name_matches;
917
918 #[test]
919 fn patterns_match_like_laravel() {
920 assert!(route_name_matches("admin.users.index", "admin.*"));
921 assert!(route_name_matches("admin.users.index", "*.index"));
922 assert!(route_name_matches("admin.users.index", "admin.*.index"));
923 assert!(route_name_matches("products.show", "products.show"));
924 assert!(!route_name_matches("products.show", "products"));
925 assert!(!route_name_matches("shop.products.show", "products.*"));
926 assert!(route_name_matches("anything", "*"));
927 assert!(!route_name_matches("admin", "admin.*"));
928 assert!(!route_name_matches("admin.users.index", "admin.*.posts.*"));
930 assert!(route_name_matches(
931 "admin.users.posts.edit",
932 "admin.*.posts.*"
933 ));
934 }
935
936 #[tokio::test]
939 async fn current_route_without_the_app_is_none() {
940 use axum::extract::FromRequestParts;
941 let (mut parts, ()) = axum::http::Request::builder()
942 .uri("/x")
943 .body(())
944 .unwrap()
945 .into_parts();
946 let current = super::CurrentRoute::from_request_parts(&mut parts, &())
947 .await
948 .unwrap();
949 assert_eq!(current.name(), None);
950 }
951}
952
953fn redirect_with(status: axum::http::StatusCode, to: &str) -> axum::response::Response {
955 match axum::http::HeaderValue::from_str(to) {
956 Ok(location) => (status, [(axum::http::header::LOCATION, location)]).into_response(),
957 Err(_) => axum::http::StatusCode::INTERNAL_SERVER_ERROR.into_response(),
958 }
959}