1use std::collections::HashMap;
16
17use smallvec::SmallVec;
18
19use std::sync::Arc;
20use std::sync::atomic::{AtomicU64, Ordering};
21
22use super::extract::{Extensions, Request};
23use super::handler::Handler;
24use super::middleware::{
25 RequestLogSink, RequestTracePolicy, resolve_trace_id, trace_request, wall_clock_now,
26};
27use super::response::{IntoResponse, Response, StatusCode};
28use crate::Cx;
29use crate::service::Layer;
30use crate::types::{
31 Budget, Time,
32 id::{next_bootstrap_region_id, next_bootstrap_task_id},
33};
34
35const METHOD_GET: &str = "GET";
38const METHOD_POST: &str = "POST";
39const METHOD_PUT: &str = "PUT";
40const METHOD_DELETE: &str = "DELETE";
41const METHOD_PATCH: &str = "PATCH";
42const METHOD_HEAD: &str = "HEAD";
43const METHOD_OPTIONS: &str = "OPTIONS";
44
45#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
51pub struct RouteInfo {
52 pub method: String,
54 pub pattern: String,
56 pub handler_name: &'static str,
58 pub mount_prefix: Option<String>,
62}
63
64pub struct MethodRouter {
68 handlers: HashMap<String, Box<dyn Handler>>,
69 method_not_allowed: Box<dyn Handler>,
70}
71
72impl MethodRouter {
73 fn new() -> Self {
75 Self {
76 handlers: HashMap::with_capacity(4),
77 method_not_allowed: Box::new(MethodNotAllowedHandler::new(String::new())),
78 }
79 }
80
81 fn on(mut self, method: &str, handler: impl Handler) -> Self {
83 self.handlers
84 .insert(method.to_uppercase(), Box::new(handler));
85 self.method_not_allowed = Box::new(MethodNotAllowedHandler::new(self.allow_header()));
86 self
87 }
88
89 #[must_use]
91 pub fn get(self, handler: impl Handler) -> Self {
92 self.on(METHOD_GET, handler)
93 }
94
95 #[must_use]
97 pub fn post(self, handler: impl Handler) -> Self {
98 self.on(METHOD_POST, handler)
99 }
100
101 #[must_use]
103 pub fn put(self, handler: impl Handler) -> Self {
104 self.on(METHOD_PUT, handler)
105 }
106
107 #[must_use]
109 pub fn delete(self, handler: impl Handler) -> Self {
110 self.on(METHOD_DELETE, handler)
111 }
112
113 #[must_use]
115 pub fn patch(self, handler: impl Handler) -> Self {
116 self.on(METHOD_PATCH, handler)
117 }
118
119 #[must_use]
121 pub fn head(self, handler: impl Handler) -> Self {
122 self.on(METHOD_HEAD, handler)
123 }
124
125 #[must_use]
127 pub fn options(self, handler: impl Handler) -> Self {
128 self.on(METHOD_OPTIONS, handler)
129 }
130
131 fn map_handlers(&mut self, wrap: &dyn Fn(Box<dyn Handler>) -> Box<dyn Handler>) {
136 let handlers = std::mem::take(&mut self.handlers);
137 self.handlers = handlers
138 .into_iter()
139 .map(|(method, handler)| (method, wrap(handler)))
140 .collect();
141 let method_not_allowed = std::mem::replace(
142 &mut self.method_not_allowed,
143 Box::new(MethodNotAllowedHandler::new(String::new())),
144 );
145 self.method_not_allowed = wrap(method_not_allowed);
146 }
147
148 #[must_use]
150 pub fn methods(&self) -> Vec<String> {
151 sorted_methods(self.handlers.keys().map(String::as_str))
152 }
153
154 fn allow_header(&self) -> String {
155 self.methods().join(", ")
156 }
157
158 fn route_entries(&self, pattern: &str, mount_prefix: Option<&str>) -> Vec<RouteInfo> {
159 let mut entries = self
160 .handlers
161 .iter()
162 .map(|(method, handler)| RouteInfo {
163 method: method.clone(),
164 pattern: pattern.to_string(),
165 handler_name: handler.handler_name(),
166 mount_prefix: mount_prefix.map(ToOwned::to_owned),
167 })
168 .collect::<Vec<_>>();
169 entries.sort_by(|left, right| {
170 compare_methods(&left.method, &right.method)
171 .then_with(|| left.handler_name.cmp(right.handler_name))
172 });
173 entries
174 }
175
176 async fn dispatch(&self, cx: &Cx, req: Request) -> Response {
178 if let Some(handler) = self.handlers.get(&req.method) {
180 return handler.call(cx, req).await;
181 }
182 let upper = req.method.to_uppercase();
184 match self.handlers.get(&upper) {
185 Some(handler) => handler.call(cx, req).await,
186 None => self.method_not_allowed.call(cx, req).await,
187 }
188 }
189}
190
191struct MethodNotAllowedHandler {
192 allow: String,
193}
194
195impl MethodNotAllowedHandler {
196 fn new(allow: String) -> Self {
197 Self { allow }
198 }
199}
200
201impl Handler for MethodNotAllowedHandler {
202 fn call(
203 &self,
204 _cx: &Cx,
205 _req: Request,
206 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Response> + Send + '_>> {
207 let allow = self.allow.clone();
208 Box::pin(async move {
209 let mut resp = StatusCode::METHOD_NOT_ALLOWED.into_response();
210 if !allow.is_empty() {
211 resp.set_header("allow", allow);
212 }
213 resp
214 })
215 }
216}
217
218fn sorted_methods<'a>(methods: impl IntoIterator<Item = &'a str>) -> Vec<String> {
219 let mut methods = methods
220 .into_iter()
221 .map(str::to_string)
222 .collect::<Vec<String>>();
223 methods.sort_by(|left, right| compare_methods(left, right));
224 methods
225}
226
227fn compare_methods(left: &str, right: &str) -> std::cmp::Ordering {
228 method_sort_key(left).cmp(&method_sort_key(right))
229}
230
231fn method_sort_key(method: &str) -> (u8, &str) {
232 match method {
233 METHOD_GET => (0, method),
234 METHOD_POST => (1, method),
235 METHOD_PUT => (2, method),
236 METHOD_DELETE => (3, method),
237 METHOD_PATCH => (4, method),
238 METHOD_HEAD => (5, method),
239 METHOD_OPTIONS => (6, method),
240 _ => (7, method),
241 }
242}
243
244pub fn get(handler: impl Handler) -> MethodRouter {
248 MethodRouter::new().get(handler)
249}
250
251pub fn post(handler: impl Handler) -> MethodRouter {
253 MethodRouter::new().post(handler)
254}
255
256pub fn put(handler: impl Handler) -> MethodRouter {
258 MethodRouter::new().put(handler)
259}
260
261pub fn delete(handler: impl Handler) -> MethodRouter {
263 MethodRouter::new().delete(handler)
264}
265
266pub fn patch(handler: impl Handler) -> MethodRouter {
268 MethodRouter::new().patch(handler)
269}
270
271#[derive(Debug, Clone)]
275struct RoutePattern {
276 #[allow(dead_code)] raw: String,
279 segments: Vec<Segment>,
281}
282
283#[derive(Debug, Clone)]
284struct RouteMatch {
285 params: HashMap<String, String>,
286 specificity: RouteSpecificity,
287}
288
289#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
290struct RouteSpecificity {
291 exact_path: bool,
292 literal_segments: usize,
293 param_segments: usize,
294 total_segments: usize,
295}
296
297#[derive(Debug, Clone)]
298enum Segment {
299 Literal(String),
300 Param(String),
301 Wildcard,
302}
303
304impl RoutePattern {
305 fn parse(pattern: &str) -> Self {
307 let segments = pattern
308 .split('/')
309 .filter(|s| !s.is_empty())
310 .map(|s| {
311 s.strip_prefix(':').map_or_else(
312 || {
313 if s == "*" {
314 Segment::Wildcard
315 } else {
316 Segment::Literal(s.to_string())
317 }
318 },
319 |param| Segment::Param(param.to_string()),
320 )
321 })
322 .collect();
323
324 Self {
325 raw: pattern.to_string(),
326 segments,
327 }
328 }
329
330 fn matches(&self, path: &str) -> Option<RouteMatch> {
332 if path.contains("//") {
348 return None;
349 }
350 let path_segments: SmallVec<[&str; 8]> =
351 path.split('/').filter(|s| !s.is_empty()).collect();
352
353 let has_wildcard = self
355 .segments
356 .last()
357 .is_some_and(|s| matches!(s, Segment::Wildcard));
358
359 if has_wildcard {
360 if path_segments.len() < self.segments.len() - 1 {
361 return None;
362 }
363 } else if path_segments.len() != self.segments.len() {
364 return None;
365 }
366
367 let mut params = HashMap::with_capacity(2);
368
369 for (i, segment) in self.segments.iter().enumerate() {
370 match segment {
371 Segment::Literal(lit) => {
372 if path_segments.get(i) != Some(&lit.as_str()) {
373 return None;
374 }
375 }
376 Segment::Param(name) => {
377 if let Some(&value) = path_segments.get(i) {
378 params.insert(name.clone(), value.to_string());
379 } else {
380 return None;
381 }
382 }
383 Segment::Wildcard => {
384 let rest = path_segments[i..].join("/");
386 params.insert("*".to_string(), rest);
387 return Some(RouteMatch {
388 params,
389 specificity: self.specificity(),
390 });
391 }
392 }
393 }
394
395 Some(RouteMatch {
396 params,
397 specificity: self.specificity(),
398 })
399 }
400
401 fn specificity(&self) -> RouteSpecificity {
402 let mut literal_segments = 0;
403 let mut param_segments = 0;
404 let mut exact_path = true;
405
406 for segment in &self.segments {
407 match segment {
408 Segment::Literal(_) => literal_segments += 1,
409 Segment::Param(_) => param_segments += 1,
410 Segment::Wildcard => exact_path = false,
411 }
412 }
413
414 RouteSpecificity {
415 exact_path,
416 literal_segments,
417 param_segments,
418 total_segments: self.segments.len(),
419 }
420 }
421}
422
423pub struct Router {
453 routes: Vec<(RoutePattern, MethodRouter)>,
454 nested: Vec<(String, Self)>,
455 fallback: Option<Box<dyn Handler>>,
456 extensions: Extensions,
457 default_trace: Option<DefaultTrace>,
458}
459
460struct DefaultTrace {
463 policy: RequestTracePolicy,
464 time_getter: fn() -> Time,
465 counter: Arc<AtomicU64>,
466 sink: Option<RequestLogSink>,
467}
468
469impl Default for DefaultTrace {
470 fn default() -> Self {
471 Self {
472 policy: RequestTracePolicy {
476 duration_header: None,
477 trace_header: None,
478 },
479 time_getter: wall_clock_now,
480 counter: Arc::new(AtomicU64::new(1)),
481 sink: None,
482 }
483 }
484}
485
486impl Default for Router {
487 fn default() -> Self {
488 Self {
489 routes: Vec::new(),
490 nested: Vec::new(),
491 fallback: None,
492 extensions: Extensions::new(),
493 default_trace: Some(DefaultTrace::default()),
494 }
495 }
496}
497
498impl Router {
499 #[must_use]
501 pub fn new() -> Self {
502 Self::default()
503 }
504
505 #[must_use]
507 pub fn route(mut self, pattern: &str, method_router: MethodRouter) -> Self {
508 self.routes
509 .push((RoutePattern::parse(pattern), method_router));
510 self
511 }
512
513 #[must_use]
515 pub fn nest(mut self, prefix: &str, router: Self) -> Self {
516 self.nested.push((prefix.to_string(), router));
517 self
518 }
519
520 #[must_use]
522 pub fn fallback(mut self, handler: impl Handler) -> Self {
523 self.fallback = Some(Box::new(handler));
524 self
525 }
526
527 #[must_use]
570 pub fn layer<L>(mut self, layer: L) -> Self
571 where
572 L: Layer<Box<dyn Handler>>,
573 L::Service: Handler,
574 {
575 let wrap =
576 move |handler: Box<dyn Handler>| -> Box<dyn Handler> { Box::new(layer.layer(handler)) };
577 self.apply_wrap(&wrap);
578 self
579 }
580
581 fn apply_wrap(&mut self, wrap: &dyn Fn(Box<dyn Handler>) -> Box<dyn Handler>) {
583 for (_, method_router) in &mut self.routes {
584 method_router.map_handlers(wrap);
585 }
586 if let Some(fallback) = self.fallback.take() {
587 self.fallback = Some(wrap(fallback));
588 }
589 for (_, nested) in &mut self.nested {
590 nested.apply_wrap(wrap);
591 }
592 }
593
594 #[must_use]
598 pub fn with_state<T>(mut self, state: T) -> Self
599 where
600 T: Clone + Send + Sync + 'static,
601 {
602 self.extensions.insert_typed(state);
603 self
604 }
605
606 #[must_use]
614 pub fn without_default_trace(mut self) -> Self {
615 self.default_trace = None;
616 self
617 }
618
619 #[must_use]
622 pub fn with_default_trace_policy(mut self, policy: RequestTracePolicy) -> Self {
623 self.default_trace
624 .get_or_insert_with(DefaultTrace::default)
625 .policy = policy;
626 self
627 }
628
629 #[must_use]
632 pub fn with_default_trace_time_getter(mut self, time_getter: fn() -> Time) -> Self {
633 self.default_trace
634 .get_or_insert_with(DefaultTrace::default)
635 .time_getter = time_getter;
636 self
637 }
638
639 #[must_use]
644 pub fn with_default_trace_record_sink(mut self, sink: RequestLogSink) -> Self {
645 self.default_trace
646 .get_or_insert_with(DefaultTrace::default)
647 .sink = Some(sink);
648 self
649 }
650
651 #[must_use]
656 pub fn handle(&self, req: Request) -> Response {
657 let cx = Cx::new(
658 next_bootstrap_region_id(),
659 next_bootstrap_task_id(),
660 Budget::INFINITE,
661 );
662 futures_lite::future::block_on(self.handle_with_cx(&cx, req))
663 }
664
665 #[must_use]
677 pub async fn handle_with_cx(&self, cx: &Cx, mut req: Request) -> Response {
678 if let Some(trace) = &self.default_trace {
679 Self::ensure_request_id(&mut req, &trace.counter);
680 return trace_request(
681 &trace.policy,
682 trace.time_getter,
683 trace.sink.as_ref(),
684 cx,
685 req,
686 |req| self.handle_inner(cx, req),
687 )
688 .await;
689 }
690 self.handle_inner(cx, req).await
691 }
692
693 fn ensure_request_id(req: &mut Request, counter: &AtomicU64) {
698 if resolve_trace_id(req).is_none() {
699 let id = format!("req-{}", counter.fetch_add(1, Ordering::Relaxed));
700 req.extensions.insert("request_id", id.clone());
701 req.headers.insert("x-request-id".to_string(), id);
702 }
703 }
704
705 async fn handle_inner(&self, cx: &Cx, mut req: Request) -> Response {
707 req.extensions.extend_from(&self.extensions);
708
709 let mut best_route: Option<(RouteSpecificity, &MethodRouter, HashMap<String, String>)> =
713 None;
714 for (pattern, method_router) in &self.routes {
715 if let Some(route_match) = pattern.matches(&req.path) {
716 match &best_route {
717 Some((best_specificity, _, _))
718 if *best_specificity >= route_match.specificity => {}
719 _ => {
720 best_route =
721 Some((route_match.specificity, method_router, route_match.params));
722 }
723 }
724 }
725 }
726 if let Some((_, method_router, params)) = best_route {
727 req.path_params = params;
728 return method_router.dispatch(cx, req).await;
729 }
730
731 let mut best_nested_match: Option<(usize, &Self, String)> = None;
733 for (prefix, router) in &self.nested {
734 if let Some(sub_path) = strip_prefix(&req.path, prefix) {
735 let normalized_len = prefix.trim_end_matches('/').len();
736 match &best_nested_match {
737 Some((best_len, _, _)) if *best_len >= normalized_len => {}
738 _ => best_nested_match = Some((normalized_len, router, sub_path)),
739 }
740 }
741 }
742 if let Some((_, router, sub_path)) = best_nested_match {
743 req.path = sub_path;
744 return Box::pin(router.handle_inner(cx, req)).await;
747 }
748
749 if let Some(handler) = &self.fallback {
751 return handler.call(cx, req).await;
752 }
753
754 StatusCode::NOT_FOUND.into_response()
755 }
756
757 #[must_use]
759 pub fn route_count(&self) -> usize {
760 self.routes.len()
761 }
762
763 #[must_use]
765 pub fn nested_router_count(&self) -> usize {
766 self.nested.len()
767 }
768
769 #[must_use]
771 pub fn has_fallback(&self) -> bool {
772 self.fallback.is_some()
773 }
774
775 #[must_use]
781 pub fn routes(&self) -> Vec<RouteInfo> {
782 let mut entries = Vec::new();
783 self.collect_routes("", None, &mut entries);
784 entries.sort_by(|left, right| {
785 left.pattern
786 .cmp(&right.pattern)
787 .then_with(|| compare_methods(&left.method, &right.method))
788 .then_with(|| left.handler_name.cmp(right.handler_name))
789 });
790 entries
791 }
792
793 fn collect_routes(
794 &self,
795 prefix: &str,
796 mount_prefix: Option<&str>,
797 entries: &mut Vec<RouteInfo>,
798 ) {
799 for (pattern, method_router) in &self.routes {
800 let full_pattern = join_route_pattern(prefix, &pattern.raw);
801 entries.extend(method_router.route_entries(&full_pattern, mount_prefix));
802 }
803
804 for (nested_prefix, router) in &self.nested {
805 let full_prefix = join_route_pattern(prefix, nested_prefix);
806 router.collect_routes(&full_prefix, Some(&full_prefix), entries);
807 }
808 }
809}
810
811fn join_route_pattern(prefix: &str, pattern: &str) -> String {
812 let prefix = normalize_route_pattern(prefix);
813 let pattern = normalize_route_pattern(pattern);
814
815 if prefix == "/" {
816 return pattern;
817 }
818 if pattern == "/" {
819 return prefix;
820 }
821
822 format!(
823 "{}/{}",
824 prefix.trim_end_matches('/'),
825 pattern.trim_start_matches('/')
826 )
827}
828
829fn normalize_route_pattern(pattern: &str) -> String {
830 if pattern.is_empty() || pattern == "/" {
831 "/".to_string()
832 } else if pattern.starts_with('/') {
833 pattern.to_string()
834 } else {
835 format!("/{pattern}")
836 }
837}
838
839fn strip_prefix(path: &str, prefix: &str) -> Option<String> {
841 let normalized_path = if path.is_empty() { "/" } else { path };
842
843 if prefix.trim_matches('/').is_empty() {
844 return normalized_path
845 .starts_with('/')
846 .then(|| normalized_path.to_string());
847 }
848
849 let requires_slash_boundary = prefix.ends_with('/');
850 let normalized_prefix = prefix.trim_end_matches('/');
851
852 if normalized_path == normalized_prefix {
853 if requires_slash_boundary {
854 return None;
855 }
856 return Some("/".to_string());
857 }
858
859 let rest = normalized_path.strip_prefix(normalized_prefix)?;
860 let rest = rest.strip_prefix('/')?;
861 if rest.starts_with('/') {
862 return None;
863 }
864
865 Some(if rest.is_empty() {
866 "/".to_string()
867 } else {
868 format!("/{rest}")
869 })
870}
871
872#[cfg(test)]
875mod tests {
876 #![allow(
877 clippy::pedantic,
878 clippy::nursery,
879 clippy::expect_fun_call,
880 clippy::map_unwrap_or,
881 clippy::cast_possible_wrap,
882 clippy::future_not_send
883 )]
884 use super::*;
885 use crate::web::handler::FnHandler;
886
887 fn ok_handler() -> &'static str {
888 "ok"
889 }
890
891 fn not_found_handler() -> StatusCode {
892 StatusCode::NOT_FOUND
893 }
894
895 fn created_handler() -> StatusCode {
896 StatusCode::CREATED
897 }
898
899 #[test]
900 fn route_exact_match() {
901 let router = Router::new().route("/", get(FnHandler::new(ok_handler)));
902
903 let resp = router.handle(Request::new("GET", "/"));
904 assert_eq!(resp.status, StatusCode::OK);
905 }
906
907 #[test]
908 fn route_not_found() {
909 let router = Router::new().route("/", get(FnHandler::new(ok_handler)));
910
911 let resp = router.handle(Request::new("GET", "/missing"));
912 assert_eq!(resp.status, StatusCode::NOT_FOUND);
913 }
914
915 #[test]
916 fn route_method_not_allowed() {
917 let router = Router::new().route("/", get(FnHandler::new(ok_handler)));
918
919 let resp = router.handle(Request::new("POST", "/"));
920 assert_eq!(resp.status, StatusCode::METHOD_NOT_ALLOWED);
921 assert_eq!(resp.header_value("allow"), Some("GET"));
922 }
923
924 #[test]
925 fn route_with_params() {
926 use crate::web::extract::Path;
927 use crate::web::handler::FnHandler1;
928
929 fn get_user(Path(id): Path<String>) -> String {
930 format!("user:{id}")
931 }
932
933 let router = Router::new().route(
934 "/users/:id",
935 get(FnHandler1::<_, Path<String>>::new(get_user)),
936 );
937
938 let resp = router.handle(Request::new("GET", "/users/42"));
939 assert_eq!(resp.status, StatusCode::OK);
940 }
941
942 #[test]
943 fn route_with_typed_path_and_query_extractors() {
944 use crate::web::extract::{Path, Query};
945 use crate::web::handler::FnHandler2;
946
947 #[derive(serde::Deserialize)]
948 struct UserPath {
949 id: u64,
950 }
951
952 #[derive(serde::Deserialize)]
953 struct Pagination {
954 page: u32,
955 active: bool,
956 }
957
958 fn handler(Path(path): Path<UserPath>, Query(query): Query<Pagination>) -> String {
959 format!("id:{} page:{} active:{}", path.id, query.page, query.active)
960 }
961
962 let router = Router::new().route(
963 "/users/:id",
964 get(FnHandler2::<_, Path<UserPath>, Query<Pagination>>::new(
965 handler,
966 )),
967 );
968
969 let req = Request::new("GET", "/users/42").with_query("page=3&active=true");
970 let resp = router.handle(req);
971 assert_eq!(resp.status, StatusCode::OK);
972 assert_eq!(resp.body.as_ref(), b"id:42 page:3 active:true");
973 }
974
975 #[test]
976 fn route_with_typed_query_error_returns_400() {
977 use crate::web::extract::Query;
978 use crate::web::handler::FnHandler1;
979
980 #[derive(serde::Deserialize)]
981 #[allow(dead_code)] struct Pagination {
983 page: u32,
984 }
985
986 fn handler(Query(_query): Query<Pagination>) -> &'static str {
987 "ok"
988 }
989
990 let router = Router::new().route(
991 "/items",
992 get(FnHandler1::<_, Query<Pagination>>::new(handler)),
993 );
994
995 let req = Request::new("GET", "/items").with_query("page=not-a-number");
996 let resp = router.handle(req);
997 assert_eq!(resp.status, StatusCode::BAD_REQUEST);
998 }
999
1000 #[test]
1001 fn route_with_typed_state() {
1002 use crate::web::extract::State;
1003 use crate::web::handler::FnHandler1;
1004
1005 #[derive(Clone)]
1006 struct AppState {
1007 greeting: &'static str,
1008 }
1009
1010 fn greet(State(state): State<AppState>) -> String {
1011 state.greeting.to_string()
1012 }
1013
1014 let router = Router::new()
1015 .route("/", get(FnHandler1::<_, State<AppState>>::new(greet)))
1016 .with_state(AppState { greeting: "hello" });
1017
1018 let resp = router.handle(Request::new("GET", "/"));
1019 assert_eq!(resp.status, StatusCode::OK);
1020 assert_eq!(resp.body.as_ref(), b"hello");
1021 }
1022
1023 #[test]
1024 fn route_with_typed_state_missing_returns_500() {
1025 use crate::web::extract::State;
1026 use crate::web::handler::FnHandler1;
1027
1028 #[derive(Clone)]
1029 struct AppState;
1030
1031 fn handler(State(_state): State<AppState>) -> &'static str {
1032 "ok"
1033 }
1034
1035 let router = Router::new().route("/", get(FnHandler1::<_, State<AppState>>::new(handler)));
1036
1037 let resp = router.handle(Request::new("GET", "/"));
1038 assert_eq!(resp.status, StatusCode::INTERNAL_SERVER_ERROR);
1039 }
1040
1041 #[test]
1042 fn route_with_multiple_typed_states() {
1043 use crate::web::extract::State;
1044 use crate::web::handler::FnHandler2;
1045
1046 #[derive(Clone)]
1047 struct AppState {
1048 name: &'static str,
1049 }
1050
1051 #[derive(Clone)]
1052 struct FeatureFlags {
1053 beta: bool,
1054 }
1055
1056 fn handler(State(app): State<AppState>, State(flags): State<FeatureFlags>) -> String {
1057 format!("{}:{}", app.name, flags.beta)
1058 }
1059
1060 let router = Router::new()
1061 .route(
1062 "/",
1063 get(FnHandler2::<_, State<AppState>, State<FeatureFlags>>::new(
1064 handler,
1065 )),
1066 )
1067 .with_state(AppState { name: "router" })
1068 .with_state(FeatureFlags { beta: true });
1069
1070 let resp = router.handle(Request::new("GET", "/"));
1071 assert_eq!(resp.status, StatusCode::OK);
1072 assert_eq!(resp.body.as_ref(), b"router:true");
1073 }
1074
1075 #[test]
1076 fn route_with_state_same_type_last_insert_wins() {
1077 use crate::web::extract::State;
1078 use crate::web::handler::FnHandler1;
1079
1080 #[derive(Clone)]
1081 struct AppState {
1082 value: &'static str,
1083 }
1084
1085 fn handler(State(app): State<AppState>) -> String {
1086 app.value.to_string()
1087 }
1088
1089 let router = Router::new()
1090 .route("/", get(FnHandler1::<_, State<AppState>>::new(handler)))
1091 .with_state(AppState { value: "first" })
1092 .with_state(AppState { value: "second" });
1093
1094 let resp = router.handle(Request::new("GET", "/"));
1095 assert_eq!(resp.status, StatusCode::OK);
1096 assert_eq!(resp.body.as_ref(), b"second");
1097 }
1098
1099 #[test]
1100 fn route_multiple_methods() {
1101 fn post_handler() -> StatusCode {
1102 StatusCode::CREATED
1103 }
1104
1105 let router = Router::new().route(
1106 "/items",
1107 get(FnHandler::new(ok_handler)).post(FnHandler::new(post_handler)),
1108 );
1109
1110 let resp_get = router.handle(Request::new("GET", "/items"));
1111 assert_eq!(resp_get.status, StatusCode::OK);
1112
1113 let resp_post = router.handle(Request::new("POST", "/items"));
1114 assert_eq!(resp_post.status, StatusCode::CREATED);
1115
1116 let resp_patch = router.handle(Request::new("PATCH", "/items"));
1117 assert_eq!(resp_patch.status, StatusCode::METHOD_NOT_ALLOWED);
1118 assert_eq!(resp_patch.header_value("allow"), Some("GET, POST"));
1119 }
1120
1121 #[test]
1122 fn route_priority_literal_before_param() {
1123 use crate::web::extract::Path;
1124 use crate::web::handler::FnHandler1;
1125
1126 fn param_handler(Path(_id): Path<String>) -> StatusCode {
1127 StatusCode::CREATED
1128 }
1129
1130 let router = Router::new()
1131 .route("/users/me", get(FnHandler::new(ok_handler)))
1132 .route(
1133 "/users/:id",
1134 get(FnHandler1::<_, Path<String>>::new(param_handler)),
1135 );
1136
1137 let resp = router.handle(Request::new("GET", "/users/me"));
1138 assert_eq!(resp.status, StatusCode::OK);
1139 }
1140
1141 #[test]
1142 fn route_priority_param_before_literal() {
1143 use crate::web::extract::Path;
1144 use crate::web::handler::FnHandler1;
1145
1146 fn param_handler(Path(_id): Path<String>) -> StatusCode {
1147 StatusCode::CREATED
1148 }
1149
1150 let router = Router::new()
1151 .route(
1152 "/users/:id",
1153 get(FnHandler1::<_, Path<String>>::new(param_handler)),
1154 )
1155 .route("/users/me", get(FnHandler::new(ok_handler)));
1156
1157 let resp = router.handle(Request::new("GET", "/users/me"));
1158 assert_eq!(resp.status, StatusCode::OK);
1159 }
1160
1161 #[test]
1162 fn route_priority_literal_before_wildcard() {
1163 use crate::web::extract::Path;
1164 use crate::web::handler::FnHandler1;
1165
1166 fn wildcard_handler(
1167 Path(_params): Path<std::collections::HashMap<String, String>>,
1168 ) -> StatusCode {
1169 StatusCode::ACCEPTED
1170 }
1171
1172 let router = Router::new()
1173 .route("/files/static", get(FnHandler::new(ok_handler)))
1174 .route(
1175 "/files/*",
1176 get(FnHandler1::<
1177 _,
1178 Path<std::collections::HashMap<String, String>>,
1179 >::new(wildcard_handler)),
1180 );
1181
1182 let resp = router.handle(Request::new("GET", "/files/static"));
1183 assert_eq!(resp.status, StatusCode::OK);
1184 }
1185
1186 #[test]
1187 fn route_priority_wildcard_cannot_shadow_literal() {
1188 use crate::web::extract::Path;
1189 use crate::web::handler::FnHandler1;
1190
1191 fn wildcard_handler(
1192 Path(_params): Path<std::collections::HashMap<String, String>>,
1193 ) -> StatusCode {
1194 StatusCode::ACCEPTED
1195 }
1196
1197 let router = Router::new()
1198 .route(
1199 "/files/*",
1200 get(FnHandler1::<
1201 _,
1202 Path<std::collections::HashMap<String, String>>,
1203 >::new(wildcard_handler))
1204 .post(FnHandler1::<
1205 _,
1206 Path<std::collections::HashMap<String, String>>,
1207 >::new(wildcard_handler)),
1208 )
1209 .route("/files/static", get(FnHandler::new(ok_handler)));
1210
1211 let resp = router.handle(Request::new("GET", "/files/static"));
1212 assert_eq!(resp.status, StatusCode::OK);
1213
1214 let resp = router.handle(Request::new("POST", "/files/static"));
1215 assert_eq!(resp.status, StatusCode::METHOD_NOT_ALLOWED);
1216 }
1217
1218 #[test]
1219 fn route_priority_wildcard_cannot_shadow_parameter_auth_path() {
1220 use crate::web::extract::Path;
1221 use crate::web::handler::FnHandler1;
1222
1223 fn public_wildcard(
1224 Path(_params): Path<std::collections::HashMap<String, String>>,
1225 ) -> StatusCode {
1226 StatusCode::OK
1227 }
1228
1229 fn protected_param(Path(_tenant): Path<String>) -> StatusCode {
1230 StatusCode::UNAUTHORIZED
1231 }
1232
1233 let router = Router::new()
1234 .route(
1235 "/admin/*",
1236 get(FnHandler1::<
1237 _,
1238 Path<std::collections::HashMap<String, String>>,
1239 >::new(public_wildcard))
1240 .post(FnHandler1::<
1241 _,
1242 Path<std::collections::HashMap<String, String>>,
1243 >::new(public_wildcard)),
1244 )
1245 .route(
1246 "/admin/:tenant/secret",
1247 get(FnHandler1::<_, Path<String>>::new(protected_param)),
1248 );
1249
1250 let resp = router.handle(Request::new("GET", "/admin/acme/secret"));
1251 assert_eq!(resp.status, StatusCode::UNAUTHORIZED);
1252
1253 let resp = router.handle(Request::new("POST", "/admin/acme/secret"));
1254 assert_eq!(resp.status, StatusCode::METHOD_NOT_ALLOWED);
1255 }
1256
1257 #[test]
1258 fn nested_router() {
1259 let api = Router::new().route("/users", get(FnHandler::new(ok_handler)));
1260
1261 let app = Router::new().nest("/api/v1", api);
1262
1263 let resp = app.handle(Request::new("GET", "/api/v1/users"));
1264 assert_eq!(resp.status, StatusCode::OK);
1265
1266 let resp = app.handle(Request::new("GET", "/other"));
1267 assert_eq!(resp.status, StatusCode::NOT_FOUND);
1268 }
1269
1270 #[test]
1271 fn nested_router_top_level_priority() {
1272 let api = Router::new().route("/users", get(FnHandler::new(created_handler)));
1273
1274 let app = Router::new()
1275 .route("/api/v1/users", get(FnHandler::new(ok_handler)))
1276 .nest("/api/v1", api);
1277
1278 let resp = app.handle(Request::new("POST", "/api/v1/users"));
1279 assert_eq!(resp.status, StatusCode::METHOD_NOT_ALLOWED);
1280 }
1281
1282 #[test]
1283 fn nested_router_typed_state_override_prefers_nested_router() {
1284 use crate::web::extract::State;
1285 use crate::web::handler::FnHandler1;
1286
1287 #[derive(Clone)]
1288 struct AppState {
1289 greeting: &'static str,
1290 }
1291
1292 fn handler(State(state): State<AppState>) -> String {
1293 state.greeting.to_string()
1294 }
1295
1296 let api = Router::new()
1297 .route("/", get(FnHandler1::<_, State<AppState>>::new(handler)))
1298 .with_state(AppState { greeting: "nested" });
1299
1300 let app = Router::new()
1301 .with_state(AppState { greeting: "parent" })
1302 .nest("/api", api);
1303
1304 let resp = app.handle(Request::new("GET", "/api/"));
1305 assert_eq!(resp.status, StatusCode::OK);
1306 assert_eq!(resp.body.as_ref(), b"nested");
1307 }
1308
1309 #[test]
1310 fn nested_router_trailing_slash_prefix() {
1311 let api = Router::new().route("/users", get(FnHandler::new(ok_handler)));
1312
1313 let app = Router::new().nest("/api/v1/", api);
1314
1315 let resp = app.handle(Request::new("GET", "/api/v1/users/"));
1316 assert_eq!(resp.status, StatusCode::OK);
1317 }
1318
1319 #[test]
1320 fn nested_router_trailing_slash_prefix_rejects_slashless_boundary() {
1321 let api = Router::new().route("/", get(FnHandler::new(created_handler)));
1322
1323 let app = Router::new()
1324 .nest("/api/v1/", api)
1325 .fallback(FnHandler::new(ok_handler));
1326
1327 let resp = app.handle(Request::new("GET", "/api/v1"));
1328 assert_eq!(resp.status, StatusCode::OK);
1329
1330 let resp = app.handle(Request::new("GET", "/api/v1/"));
1331 assert_eq!(resp.status, StatusCode::CREATED);
1332 }
1333
1334 #[test]
1335 fn nested_router_prefers_most_specific_prefix() {
1336 let broad = Router::new().route("/health", get(FnHandler::new(ok_handler)));
1337 let specific = Router::new().route("/users", get(FnHandler::new(created_handler)));
1338
1339 let app = Router::new().nest("/api", broad).nest("/api/v1", specific);
1341
1342 let resp = app.handle(Request::new("GET", "/api/v1/users"));
1343 assert_eq!(resp.status, StatusCode::CREATED);
1344 }
1345
1346 #[test]
1347 fn fallback_handler() {
1348 let router = Router::new()
1349 .route("/", get(FnHandler::new(ok_handler)))
1350 .fallback(FnHandler::new(not_found_handler));
1351
1352 let resp = router.handle(Request::new("GET", "/missing"));
1353 assert_eq!(resp.status, StatusCode::NOT_FOUND);
1354 }
1355
1356 #[test]
1357 fn route_pattern_matching() {
1358 let pattern = RoutePattern::parse("/users/:id");
1359 let params = pattern.matches("/users/42").unwrap().params;
1360 assert_eq!(params.get("id").unwrap(), "42");
1361
1362 assert!(pattern.matches("/users").is_none());
1363 assert!(pattern.matches("/users/42/extra").is_none());
1364 }
1365
1366 #[test]
1367 fn route_pattern_multiple_params() {
1368 let pattern = RoutePattern::parse("/users/:uid/posts/:pid");
1369 let params = pattern.matches("/users/1/posts/99").unwrap().params;
1370 assert_eq!(params.get("uid").unwrap(), "1");
1371 assert_eq!(params.get("pid").unwrap(), "99");
1372 }
1373
1374 #[test]
1375 fn route_pattern_wildcard() {
1376 let pattern = RoutePattern::parse("/files/*");
1377 let params = pattern.matches("/files/a/b/c").unwrap().params;
1378 assert_eq!(params.get("*").unwrap(), "a/b/c");
1379 }
1380
1381 #[test]
1382 fn route_pattern_wildcard_empty_rest() {
1383 use crate::web::extract::Path;
1384 use crate::web::handler::FnHandler1;
1385
1386 fn wildcard_handler(
1387 Path(params): Path<std::collections::HashMap<String, String>>,
1388 ) -> String {
1389 params.get("*").cloned().unwrap_or_default()
1390 }
1391
1392 let router = Router::new().route(
1393 "/files/*",
1394 get(FnHandler1::<
1395 _,
1396 Path<std::collections::HashMap<String, String>>,
1397 >::new(wildcard_handler)),
1398 );
1399
1400 let resp = router.handle(Request::new("GET", "/files"));
1401 assert_eq!(resp.status, StatusCode::OK);
1402 assert_eq!(std::str::from_utf8(&resp.body).unwrap(), "");
1403 }
1404
1405 #[test]
1406 fn route_pattern_literal_only() {
1407 let pattern = RoutePattern::parse("/health");
1408 assert!(pattern.matches("/health").is_some());
1409 assert!(pattern.matches("/other").is_none());
1410 }
1411
1412 #[test]
1413 fn route_trailing_slash_matches() {
1414 let router = Router::new().route("/users", get(FnHandler::new(ok_handler)));
1415
1416 let resp = router.handle(Request::new("GET", "/users/"));
1417 assert_eq!(resp.status, StatusCode::OK);
1418 }
1419
1420 #[test]
1421 fn router_route_count() {
1422 let router = Router::new()
1423 .route("/a", get(FnHandler::new(ok_handler)))
1424 .route("/b", get(FnHandler::new(ok_handler)));
1425 assert_eq!(router.route_count(), 2);
1426 }
1427
1428 #[test]
1429 fn router_routes_lists_direct_and_nested_entries_deterministically() {
1430 let api = Router::new()
1431 .route(
1432 "/users",
1433 post(FnHandler::new(created_handler)).get(FnHandler::new(ok_handler)),
1434 )
1435 .route("/", delete(FnHandler::new(not_found_handler)));
1436
1437 let router = Router::new()
1438 .route(
1439 "/items",
1440 post(FnHandler::new(created_handler)).get(FnHandler::new(ok_handler)),
1441 )
1442 .route("/items/:id", delete(FnHandler::new(not_found_handler)))
1443 .nest("/api", api);
1444
1445 let routes = router.routes();
1446 let serialized = serde_json::to_value(&routes).expect("route info must serialize");
1447 assert_eq!(serialized[0]["method"], "DELETE");
1448 assert_eq!(serialized[0]["pattern"], "/api");
1449 assert_eq!(serialized[0]["handler_name"], "FnHandler");
1450 assert_eq!(serialized[0]["mount_prefix"], "/api");
1451
1452 let got = routes
1453 .into_iter()
1454 .map(|route| {
1455 (
1456 route.method,
1457 route.pattern,
1458 route.handler_name,
1459 route.mount_prefix,
1460 )
1461 })
1462 .collect::<Vec<_>>();
1463
1464 assert_eq!(
1465 got,
1466 vec![
1467 (
1468 "DELETE".to_string(),
1469 "/api".to_string(),
1470 "FnHandler",
1471 Some("/api".to_string())
1472 ),
1473 (
1474 "GET".to_string(),
1475 "/api/users".to_string(),
1476 "FnHandler",
1477 Some("/api".to_string())
1478 ),
1479 (
1480 "POST".to_string(),
1481 "/api/users".to_string(),
1482 "FnHandler",
1483 Some("/api".to_string())
1484 ),
1485 ("GET".to_string(), "/items".to_string(), "FnHandler", None),
1486 ("POST".to_string(), "/items".to_string(), "FnHandler", None),
1487 (
1488 "DELETE".to_string(),
1489 "/items/:id".to_string(),
1490 "FnHandler",
1491 None
1492 ),
1493 ]
1494 );
1495 }
1496
1497 #[test]
1498 fn strip_prefix_basic() {
1499 assert_eq!(
1500 strip_prefix("/api/v1/users", "/api/v1"),
1501 Some("/users".to_string())
1502 );
1503 assert_eq!(strip_prefix("/api/v1", "/api/v1"), Some("/".to_string()));
1504 assert_eq!(strip_prefix("/api/v1/", "/api/v1"), Some("/".to_string()));
1505 assert!(strip_prefix("/other", "/api/v1").is_none());
1506 }
1507
1508 #[test]
1509 fn strip_prefix_boundary_mismatch() {
1510 assert!(strip_prefix("/apix/users", "/api").is_none());
1511 assert!(strip_prefix("/apiary", "/api").is_none());
1512 }
1513
1514 #[test]
1515 fn strip_prefix_trailing_slash_prefix_requires_declared_boundary() {
1516 assert_eq!(
1517 strip_prefix("/api/v1/users", "/api/v1/"),
1518 Some("/users".to_string())
1519 );
1520 assert_eq!(strip_prefix("/api/v1/", "/api/v1/"), Some("/".to_string()));
1521 assert!(strip_prefix("/api/v1", "/api/v1/").is_none());
1522 }
1523
1524 #[test]
1525 fn strip_prefix_rejects_empty_segment_at_mount_boundary() {
1526 assert!(strip_prefix("/api//users", "/api").is_none());
1527 assert!(strip_prefix("/api//users", "/api/").is_none());
1528 }
1529
1530 mod route_precedence_audit {
1537 use super::*;
1538 use crate::web::handler::FnHandler;
1539
1540 fn literal_handler() -> StatusCode {
1541 StatusCode::OK
1542 }
1543
1544 fn param_handler() -> StatusCode {
1545 StatusCode::ACCEPTED
1546 }
1547
1548 fn wildcard_handler() -> StatusCode {
1549 StatusCode::CREATED
1550 }
1551
1552 #[test]
1557 fn audit_literal_beats_parameter_core_requirement() {
1558 let router1 = Router::new()
1560 .route("/users/me", get(FnHandler::new(literal_handler)))
1561 .route("/users/:id", get(FnHandler::new(param_handler)))
1562 .route("/users/*", get(FnHandler::new(wildcard_handler)));
1563
1564 let resp1 = router1.handle(Request::new("GET", "/users/me"));
1565 assert_eq!(
1566 resp1.status,
1567 StatusCode::OK,
1568 "Literal route '/users/me' must win over '/users/:id' when registered first"
1569 );
1570
1571 let router2 = Router::new()
1573 .route("/users/:id", get(FnHandler::new(param_handler)))
1574 .route("/users/*", get(FnHandler::new(wildcard_handler)))
1575 .route("/users/me", get(FnHandler::new(literal_handler)));
1576
1577 let resp2 = router2.handle(Request::new("GET", "/users/me"));
1578 assert_eq!(
1579 resp2.status,
1580 StatusCode::OK,
1581 "Literal route '/users/me' must win over '/users/:id' regardless of registration order"
1582 );
1583
1584 let resp3 = router2.handle(Request::new("GET", "/users/someone"));
1587 assert_eq!(
1588 resp3.status,
1589 StatusCode::ACCEPTED,
1590 "Parameter route should still handle non-literal single-segment users"
1591 );
1592
1593 let resp4 = router2.handle(Request::new("GET", "/users/some/path"));
1594 assert_eq!(
1595 resp4.status,
1596 StatusCode::CREATED,
1597 "Wildcard route should remain the least-specific fallback"
1598 );
1599 }
1600
1601 #[test]
1606 fn audit_multiple_literal_segments_precedence() {
1607 use crate::web::extract::Path;
1608 use crate::web::handler::FnHandler1;
1609
1610 fn param_handler(Path(_params): Path<HashMap<String, String>>) -> StatusCode {
1611 StatusCode::ACCEPTED
1612 }
1613
1614 let router = Router::new()
1615 .route(
1616 "/api/:version/users",
1617 get(FnHandler1::<_, Path<HashMap<String, String>>>::new(
1618 param_handler,
1619 )),
1620 )
1621 .route("/api/v1/users", get(FnHandler::new(literal_handler)))
1622 .route(
1623 "/api/:version/:resource",
1624 get(FnHandler1::<_, Path<HashMap<String, String>>>::new(
1625 param_handler,
1626 )),
1627 );
1628
1629 let resp = router.handle(Request::new("GET", "/api/v1/users"));
1631 assert_eq!(
1632 resp.status,
1633 StatusCode::OK,
1634 "Route with more literal segments '/api/v1/users' must win over '/api/:version/users'"
1635 );
1636 }
1637
1638 #[test]
1642 fn audit_route_specificity_calculation() {
1643 let literal_route = RoutePattern::parse("/users/me/profile");
1644 let mixed_route = RoutePattern::parse("/users/:id/profile");
1645 let param_route = RoutePattern::parse("/users/:id/:section");
1646 let wildcard_route = RoutePattern::parse("/users/*");
1647
1648 let literal_spec = literal_route.specificity();
1649 let mixed_spec = mixed_route.specificity();
1650 let param_spec = param_route.specificity();
1651 let wildcard_spec = wildcard_route.specificity();
1652
1653 assert_eq!(
1655 literal_spec.literal_segments, 3,
1656 "Literal route should have 3 literal segments"
1657 );
1658 assert_eq!(
1659 mixed_spec.literal_segments, 2,
1660 "Mixed route should have 2 literal segments"
1661 );
1662 assert_eq!(
1663 param_spec.literal_segments, 1,
1664 "Param route should have 1 literal segment"
1665 );
1666 assert_eq!(
1667 wildcard_spec.literal_segments, 1,
1668 "Wildcard route should have 1 literal segment"
1669 );
1670
1671 assert_eq!(
1673 literal_spec.param_segments, 0,
1674 "Literal route should have 0 parameter segments"
1675 );
1676 assert_eq!(
1677 mixed_spec.param_segments, 1,
1678 "Mixed route should have 1 parameter segment"
1679 );
1680 assert_eq!(
1681 param_spec.param_segments, 2,
1682 "Param route should have 2 parameter segments"
1683 );
1684 assert_eq!(
1685 wildcard_spec.param_segments, 0,
1686 "Wildcard route should have 0 parameter segments (wildcard is separate)"
1687 );
1688
1689 assert!(
1691 literal_spec > mixed_spec,
1692 "Literal route must be more specific than mixed route"
1693 );
1694 assert!(
1695 mixed_spec > param_spec,
1696 "Mixed route must be more specific than parameter route"
1697 );
1698 assert!(
1699 param_spec > wildcard_spec,
1700 "Parameter route must be more specific than wildcard route"
1701 );
1702 }
1703
1704 #[test]
1708 fn audit_complex_precedence_scenarios() {
1709 fn route_a() -> &'static str {
1710 "route_a"
1711 }
1712 fn route_b() -> &'static str {
1713 "route_b"
1714 }
1715 fn route_c() -> &'static str {
1716 "route_c"
1717 }
1718
1719 let router = Router::new()
1720 .route("/api/v1/users/me", get(FnHandler::new(route_a)))
1722 .route("/api/v1/users/:id", get(FnHandler::new(route_b)))
1724 .route("/api/:version/users/:id", get(FnHandler::new(route_c)))
1726 .route("/api/*", get(FnHandler::new(|| "wildcard")));
1728
1729 let resp = router.handle(Request::new("GET", "/api/v1/users/me"));
1730 assert_eq!(resp.status, StatusCode::OK);
1731 let body = String::from_utf8(resp.body.to_vec()).unwrap();
1732 assert_eq!(body, "route_a", "Most specific literal route should win");
1733
1734 let resp2 = router.handle(Request::new("GET", "/api/v1/users/123"));
1736 assert_eq!(resp2.status, StatusCode::OK);
1737 let body2 = String::from_utf8(resp2.body.to_vec()).unwrap();
1738 assert_eq!(
1739 body2, "route_b",
1740 "Parameter route should handle non-literal values"
1741 );
1742
1743 let resp3 = router.handle(Request::new("GET", "/api/v2/users/123"));
1744 assert_eq!(resp3.status, StatusCode::OK);
1745 let body3 = String::from_utf8(resp3.body.to_vec()).unwrap();
1746 assert_eq!(
1747 body3, "route_c",
1748 "Less-specific parameter route should handle non-v1 versions"
1749 );
1750 }
1751
1752 #[test]
1756 fn audit_similar_literal_paths_distinction() {
1757 let router = Router::new()
1758 .route("/users/me", get(FnHandler::new(|| "me")))
1759 .route("/users/menu", get(FnHandler::new(|| "menu")))
1760 .route("/users/metrics", get(FnHandler::new(|| "metrics")));
1761
1762 let resp_me = router.handle(Request::new("GET", "/users/me"));
1764 assert_eq!(String::from_utf8(resp_me.body.to_vec()).unwrap(), "me");
1765
1766 let resp_menu = router.handle(Request::new("GET", "/users/menu"));
1767 assert_eq!(String::from_utf8(resp_menu.body.to_vec()).unwrap(), "menu");
1768
1769 let resp_metrics = router.handle(Request::new("GET", "/users/metrics"));
1770 assert_eq!(
1771 String::from_utf8(resp_metrics.body.to_vec()).unwrap(),
1772 "metrics"
1773 );
1774 }
1775
1776 #[test]
1780 fn audit_precedence_across_http_methods() {
1781 use crate::web::extract::Path;
1782 use crate::web::handler::FnHandler1;
1783
1784 fn literal_get() -> &'static str {
1785 "literal_get"
1786 }
1787 fn literal_post() -> &'static str {
1788 "literal_post"
1789 }
1790 fn param_get(Path(_): Path<String>) -> &'static str {
1791 "param_get"
1792 }
1793 fn param_post(Path(_): Path<String>) -> &'static str {
1794 "param_post"
1795 }
1796
1797 let router = Router::new()
1798 .route(
1799 "/users/:id",
1800 get(FnHandler1::<_, Path<String>>::new(param_get)).post(FnHandler1::<
1801 _,
1802 Path<String>,
1803 >::new(
1804 param_post
1805 )),
1806 )
1807 .route(
1808 "/users/me",
1809 get(FnHandler::new(literal_get)).post(FnHandler::new(literal_post)),
1810 );
1811
1812 let resp_get = router.handle(Request::new("GET", "/users/me"));
1814 assert_eq!(
1815 String::from_utf8(resp_get.body.to_vec()).unwrap(),
1816 "literal_get"
1817 );
1818
1819 let resp_post = router.handle(Request::new("POST", "/users/me"));
1821 assert_eq!(
1822 String::from_utf8(resp_post.body.to_vec()).unwrap(),
1823 "literal_post"
1824 );
1825 }
1826
1827 #[test]
1831 fn audit_parameter_routes_capture_when_appropriate() {
1832 use crate::web::extract::Path;
1833 use crate::web::handler::FnHandler1;
1834
1835 fn param_handler(Path(id): Path<String>) -> String {
1836 format!("captured:{}", id)
1837 }
1838
1839 let router = Router::new()
1840 .route(
1841 "/users/me",
1842 get(FnHandler::new(|| "literal:me".to_string())),
1843 )
1844 .route(
1845 "/users/:id",
1846 get(FnHandler1::<_, Path<String>>::new(param_handler)),
1847 );
1848
1849 let resp_me = router.handle(Request::new("GET", "/users/me"));
1851 assert_eq!(
1852 String::from_utf8(resp_me.body.to_vec()).unwrap(),
1853 "literal:me"
1854 );
1855
1856 let resp_123 = router.handle(Request::new("GET", "/users/123"));
1858 assert_eq!(
1859 String::from_utf8(resp_123.body.to_vec()).unwrap(),
1860 "captured:123"
1861 );
1862
1863 let resp_admin = router.handle(Request::new("GET", "/users/admin"));
1864 assert_eq!(
1865 String::from_utf8(resp_admin.body.to_vec()).unwrap(),
1866 "captured:admin"
1867 );
1868 }
1869 }
1870
1871 mod layering {
1874 use super::*;
1875 use std::sync::{Arc, Mutex};
1876
1877 #[derive(Clone)]
1879 struct RecordingLayer {
1880 name: &'static str,
1881 log: Arc<Mutex<Vec<String>>>,
1882 }
1883
1884 impl RecordingLayer {
1885 fn new(name: &'static str, log: Arc<Mutex<Vec<String>>>) -> Self {
1886 Self { name, log }
1887 }
1888 }
1889
1890 struct RecordingMiddleware<H> {
1891 inner: H,
1892 name: &'static str,
1893 log: Arc<Mutex<Vec<String>>>,
1894 }
1895
1896 impl<H: Handler> Layer<H> for RecordingLayer {
1897 type Service = RecordingMiddleware<H>;
1898
1899 fn layer(&self, inner: H) -> Self::Service {
1900 RecordingMiddleware {
1901 inner,
1902 name: self.name,
1903 log: Arc::clone(&self.log),
1904 }
1905 }
1906 }
1907
1908 impl<H: Handler> Handler for RecordingMiddleware<H> {
1909 fn call(
1910 &self,
1911 cx: &Cx,
1912 req: Request,
1913 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Response> + Send + '_>>
1914 {
1915 let cx = cx.clone();
1916 Box::pin(async move {
1917 self.log
1918 .lock()
1919 .expect("log lock")
1920 .push(format!("enter:{}", self.name));
1921 let resp = self.inner.call(&cx, req).await;
1922 self.log
1923 .lock()
1924 .expect("log lock")
1925 .push(format!("exit:{}", self.name));
1926 resp
1927 })
1928 }
1929 }
1930
1931 fn recording_handler(log: Arc<Mutex<Vec<String>>>) -> impl Handler {
1932 struct H(Arc<Mutex<Vec<String>>>);
1933 impl Handler for H {
1934 fn call(
1935 &self,
1936 _cx: &Cx,
1937 _req: Request,
1938 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Response> + Send + '_>>
1939 {
1940 Box::pin(async move {
1941 self.0.lock().expect("log lock").push("handler".to_string());
1942 StatusCode::OK.into_response()
1943 })
1944 }
1945 }
1946 H(log)
1947 }
1948
1949 #[test]
1954 fn layer_execution_order_golden() {
1955 let log = Arc::new(Mutex::new(Vec::new()));
1956 let router = Router::new()
1957 .route("/traced", get(recording_handler(Arc::clone(&log))))
1958 .layer(RecordingLayer::new("auth", Arc::clone(&log)))
1959 .layer(RecordingLayer::new("trace", Arc::clone(&log)))
1960 .layer(RecordingLayer::new("request_id", Arc::clone(&log)));
1961
1962 let resp = router.handle(Request::new("GET", "/traced"));
1963 assert_eq!(resp.status, StatusCode::OK);
1964
1965 let golden = vec![
1966 "enter:request_id".to_string(),
1967 "enter:trace".to_string(),
1968 "enter:auth".to_string(),
1969 "handler".to_string(),
1970 "exit:auth".to_string(),
1971 "exit:trace".to_string(),
1972 "exit:request_id".to_string(),
1973 ];
1974 assert_eq!(
1975 *log.lock().expect("log lock"),
1976 golden,
1977 "onion ordering: last-added layer must be outermost"
1978 );
1979 }
1980
1981 #[test]
1982 fn routes_added_after_layer_are_not_wrapped() {
1983 let log = Arc::new(Mutex::new(Vec::new()));
1984 let router = Router::new()
1985 .route("/wrapped", get(FnHandler::new(ok_handler)))
1986 .layer(RecordingLayer::new("mw", Arc::clone(&log)))
1987 .route("/bare", get(FnHandler::new(ok_handler)));
1988
1989 let _ = router.handle(Request::new("GET", "/bare"));
1990 assert!(
1991 log.lock().expect("log lock").is_empty(),
1992 "route added after .layer() must not be wrapped"
1993 );
1994
1995 let _ = router.handle(Request::new("GET", "/wrapped"));
1996 assert_eq!(
1997 *log.lock().expect("log lock"),
1998 vec!["enter:mw".to_string(), "exit:mw".to_string()],
1999 "route added before .layer() must be wrapped"
2000 );
2001 }
2002
2003 #[test]
2004 fn layer_wraps_fallback() {
2005 let log = Arc::new(Mutex::new(Vec::new()));
2006 let router = Router::new()
2007 .fallback(FnHandler::new(not_found_handler))
2008 .layer(RecordingLayer::new("mw", Arc::clone(&log)));
2009
2010 let resp = router.handle(Request::new("GET", "/nope"));
2011 assert_eq!(resp.status, StatusCode::NOT_FOUND);
2012 assert_eq!(
2013 *log.lock().expect("log lock"),
2014 vec!["enter:mw".to_string(), "exit:mw".to_string()],
2015 "fallback handler must be wrapped by .layer()"
2016 );
2017 }
2018
2019 #[test]
2020 fn layer_wraps_method_not_allowed() {
2021 let log = Arc::new(Mutex::new(Vec::new()));
2022 let router = Router::new()
2023 .route("/traced", get(FnHandler::new(ok_handler)))
2024 .layer(RecordingLayer::new("mw", Arc::clone(&log)));
2025
2026 let resp = router.handle(Request::new("POST", "/traced"));
2027 assert_eq!(resp.status, StatusCode::METHOD_NOT_ALLOWED);
2028 assert_eq!(resp.header_value("allow"), Some("GET"));
2029 assert_eq!(
2030 *log.lock().expect("log lock"),
2031 vec!["enter:mw".to_string(), "exit:mw".to_string()],
2032 "method-not-allowed must be wrapped by .layer()"
2033 );
2034 }
2035
2036 #[test]
2037 fn cors_layer_handles_preflight_without_options_route() {
2038 use crate::web::middleware::{CorsLayer, CorsPolicy};
2039
2040 let router = Router::new()
2041 .route("/cors", get(FnHandler::new(ok_handler)))
2042 .layer(CorsLayer::new(CorsPolicy::default()));
2043
2044 let resp = router.handle(
2045 Request::new("OPTIONS", "/cors")
2046 .with_header("Origin", "https://example.com")
2047 .with_header("Access-Control-Request-Method", "POST")
2048 .with_header("Access-Control-Request-Headers", "content-type"),
2049 );
2050
2051 assert_eq!(resp.status, StatusCode::NO_CONTENT);
2052 assert_eq!(
2053 resp.headers.get("access-control-allow-origin"),
2054 Some(&"*".to_string())
2055 );
2056 assert!(resp.headers.contains_key("access-control-allow-methods"));
2057 assert!(resp.headers.contains_key("access-control-allow-headers"));
2058 }
2059
2060 #[test]
2061 fn layer_wraps_nested_routers() {
2062 let log = Arc::new(Mutex::new(Vec::new()));
2063 let api = Router::new().route("/users", get(FnHandler::new(ok_handler)));
2064 let router = Router::new()
2065 .nest("/api", api)
2066 .layer(RecordingLayer::new("mw", Arc::clone(&log)));
2067
2068 let resp = router.handle(Request::new("GET", "/api/users"));
2069 assert_eq!(resp.status, StatusCode::OK);
2070 assert_eq!(
2071 *log.lock().expect("log lock"),
2072 vec!["enter:mw".to_string(), "exit:mw".to_string()],
2073 "nested router handlers must be wrapped by .layer()"
2074 );
2075 }
2076
2077 #[test]
2078 fn builtin_middleware_layers_compose_on_router() {
2079 use crate::web::middleware::{
2080 AuthLayer, AuthPolicy, HeaderOverwrite, SetResponseHeaderLayer,
2081 };
2082
2083 let router = Router::new()
2084 .route("/secure", get(FnHandler::new(ok_handler)))
2085 .layer(AuthLayer::new(AuthPolicy::exact_bearer("tok")))
2086 .layer(SetResponseHeaderLayer::new(
2087 "x-frame-options",
2088 "DENY",
2089 HeaderOverwrite::Always,
2090 ));
2091
2092 let resp = router.handle(Request::new("GET", "/secure"));
2094 assert_eq!(resp.status, StatusCode::UNAUTHORIZED);
2095 assert_eq!(
2096 resp.headers.get("x-frame-options").map(String::as_str),
2097 Some("DENY")
2098 );
2099
2100 let resp = router
2102 .handle(Request::new("GET", "/secure").with_header("authorization", "Bearer tok"));
2103 assert_eq!(resp.status, StatusCode::OK);
2104 assert_eq!(
2105 resp.headers.get("x-frame-options").map(String::as_str),
2106 Some("DENY")
2107 );
2108 }
2109
2110 mod extension_lifecycle {
2113 use super::*;
2114 use crate::web::extract::Extension;
2115 use crate::web::handler::FnHandler1;
2116 use std::sync::atomic::{AtomicU64, Ordering};
2117 use std::sync::{Arc, Mutex, Weak};
2118
2119 #[derive(Clone)]
2120 struct RequestStamp {
2121 serial: u64,
2122 }
2123
2124 #[derive(Clone)]
2126 struct StampLayer {
2127 counter: Arc<AtomicU64>,
2128 }
2129
2130 struct StampMiddleware<H> {
2131 inner: H,
2132 counter: Arc<AtomicU64>,
2133 }
2134
2135 impl<H: Handler> Layer<H> for StampLayer {
2136 type Service = StampMiddleware<H>;
2137
2138 fn layer(&self, inner: H) -> Self::Service {
2139 StampMiddleware {
2140 inner,
2141 counter: Arc::clone(&self.counter),
2142 }
2143 }
2144 }
2145
2146 impl<H: Handler> Handler for StampMiddleware<H> {
2147 fn call(
2148 &self,
2149 cx: &Cx,
2150 mut req: Request,
2151 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Response> + Send + '_>>
2152 {
2153 let serial = self.counter.fetch_add(1, Ordering::SeqCst) + 1;
2154 req.extensions.insert_typed(RequestStamp { serial });
2155 self.inner.call(cx, req)
2156 }
2157 }
2158
2159 fn stamp_echo_handler() -> impl Handler {
2160 FnHandler1::<_, Extension<RequestStamp>>::new(
2161 |Extension(stamp): Extension<RequestStamp>| format!("serial:{}", stamp.serial),
2162 )
2163 }
2164
2165 #[test]
2168 fn middleware_insert_handler_extract_no_cross_request_bleed() {
2169 let router = Router::new()
2170 .route("/stamped", get(stamp_echo_handler()))
2171 .layer(StampLayer {
2172 counter: Arc::new(AtomicU64::new(0)),
2173 });
2174
2175 let r1 = router.handle(Request::new("GET", "/stamped"));
2176 let r2 = router.handle(Request::new("GET", "/stamped"));
2177 let r3 = router.handle(Request::new("GET", "/stamped"));
2178 assert_eq!(String::from_utf8(r1.body.to_vec()).unwrap(), "serial:1");
2179 assert_eq!(String::from_utf8(r2.body.to_vec()).unwrap(), "serial:2");
2180 assert_eq!(String::from_utf8(r3.body.to_vec()).unwrap(), "serial:3");
2181 }
2182
2183 #[test]
2186 fn missing_extension_is_internal_server_error() {
2187 let router = Router::new().route("/stamped", get(stamp_echo_handler()));
2188 let resp = router.handle(Request::new("GET", "/stamped"));
2189 assert_eq!(resp.status, StatusCode::INTERNAL_SERVER_ERROR);
2190 }
2191
2192 #[derive(Clone)]
2193 struct ProbeExt(#[allow(dead_code)] Arc<()>);
2194
2195 #[test]
2198 fn extension_dropped_when_request_region_completes() {
2199 use crate::web::request_region::{RegionOutcome, RequestRegion};
2200
2201 let probe = Arc::new(());
2202 let weak: Weak<()> = Arc::downgrade(&probe);
2203
2204 let mut req = Request::new("GET", "/probe");
2205 req.extensions.insert_typed(ProbeExt(probe));
2206
2207 let handler = FnHandler1::<_, Extension<ProbeExt>>::new(
2208 |Extension(_probe): Extension<ProbeExt>| "ok",
2209 );
2210
2211 let cx = Cx::for_testing();
2212 let region = RequestRegion::new(&cx, req);
2213 let outcome = futures_lite::future::block_on(region.run_handler(&handler));
2214 assert!(matches!(outcome, RegionOutcome::Ok(_)));
2215
2216 assert!(
2217 weak.upgrade().is_none(),
2218 "extension value must drop with the request when the region run completes"
2219 );
2220 }
2221
2222 #[test]
2225 fn per_request_extensions_do_not_accumulate() {
2226 let weaks: Arc<Mutex<Vec<Weak<()>>>> = Arc::new(Mutex::new(Vec::new()));
2227
2228 #[derive(Clone)]
2229 struct ProbeLayer {
2230 weaks: Arc<Mutex<Vec<Weak<()>>>>,
2231 }
2232
2233 struct ProbeMiddleware<H> {
2234 inner: H,
2235 weaks: Arc<Mutex<Vec<Weak<()>>>>,
2236 }
2237
2238 impl<H: Handler> Layer<H> for ProbeLayer {
2239 type Service = ProbeMiddleware<H>;
2240
2241 fn layer(&self, inner: H) -> Self::Service {
2242 ProbeMiddleware {
2243 inner,
2244 weaks: Arc::clone(&self.weaks),
2245 }
2246 }
2247 }
2248
2249 impl<H: Handler> Handler for ProbeMiddleware<H> {
2250 fn call(
2251 &self,
2252 cx: &Cx,
2253 mut req: Request,
2254 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Response> + Send + '_>>
2255 {
2256 let probe = Arc::new(());
2257 self.weaks
2258 .lock()
2259 .expect("weaks lock")
2260 .push(Arc::downgrade(&probe));
2261 req.extensions.insert_typed(ProbeExt(probe));
2262 self.inner.call(cx, req)
2263 }
2264 }
2265
2266 let router = Router::new()
2267 .route(
2268 "/probe",
2269 get(FnHandler1::<_, Extension<ProbeExt>>::new(
2270 |Extension(_p): Extension<ProbeExt>| "ok",
2271 )),
2272 )
2273 .layer(ProbeLayer {
2274 weaks: Arc::clone(&weaks),
2275 });
2276
2277 for _ in 0..3 {
2278 let resp = router.handle(Request::new("GET", "/probe"));
2279 assert_eq!(resp.status, StatusCode::OK);
2280 }
2281
2282 let weaks = weaks.lock().expect("weaks lock");
2283 assert_eq!(weaks.len(), 3);
2284 assert!(
2285 weaks.iter().all(|w| w.upgrade().is_none()),
2286 "every request's extension must be dropped once its dispatch completes"
2287 );
2288 }
2289 }
2290 }
2291
2292 mod default_trace {
2295 use super::*;
2296 use crate::CancelReason;
2297 use crate::web::middleware::{RequestLogRecord, RequestLogSink};
2298 use std::sync::{Arc, Mutex};
2299
2300 fn fixed_time() -> Time {
2301 Time::from_millis(7_000)
2302 }
2303
2304 fn collecting_sink() -> (RequestLogSink, Arc<Mutex<Vec<RequestLogRecord>>>) {
2305 let records = Arc::new(Mutex::new(Vec::new()));
2306 let sink_records = Arc::clone(&records);
2307 let sink: RequestLogSink = Arc::new(move |record: &RequestLogRecord| {
2308 sink_records
2309 .lock()
2310 .expect("records lock")
2311 .push(record.clone());
2312 });
2313 (sink, records)
2314 }
2315
2316 fn boom_handler() -> Response {
2317 StatusCode::INTERNAL_SERVER_ERROR.into_response()
2318 }
2319
2320 #[test]
2324 fn default_trace_golden_scrubbed() {
2325 let (sink, records) = collecting_sink();
2326 let router = Router::new()
2327 .route("/ok", get(FnHandler::new(ok_handler)))
2328 .route("/boom", get(FnHandler::new(boom_handler)))
2329 .with_default_trace_time_getter(fixed_time)
2330 .with_default_trace_record_sink(sink);
2331
2332 let _ = router.handle(Request::new("GET", "/ok"));
2333 let _ = router.handle(Request::new("GET", "/boom"));
2334 let _ = router.handle(Request::new("POST", "/ok"));
2335 let _ = router.handle(Request::new("GET", "/missing"));
2336
2337 let got =
2338 serde_json::to_value(&*records.lock().expect("records lock")).expect("serialize");
2339 let want = serde_json::json!([
2340 {
2341 "method": "GET",
2342 "path": "/ok",
2343 "status": 200,
2344 "severity": "ok",
2345 "duration_ms": 0,
2346 "request_id": "req-1",
2347 "cancelled": false,
2348 "cancel_reason": null
2349 },
2350 {
2351 "method": "GET",
2352 "path": "/boom",
2353 "status": 500,
2354 "severity": "server_error",
2355 "duration_ms": 0,
2356 "request_id": "req-2",
2357 "cancelled": false,
2358 "cancel_reason": null
2359 },
2360 {
2361 "method": "POST",
2362 "path": "/ok",
2363 "status": 405,
2364 "severity": "client_error",
2365 "duration_ms": 0,
2366 "request_id": "req-3",
2367 "cancelled": false,
2368 "cancel_reason": null
2369 },
2370 {
2371 "method": "GET",
2372 "path": "/missing",
2373 "status": 404,
2374 "severity": "client_error",
2375 "duration_ms": 0,
2376 "request_id": "req-4",
2377 "cancelled": false,
2378 "cancel_reason": null
2379 }
2380 ]);
2381 assert_eq!(got, want, "request log schema golden drifted");
2382 }
2383
2384 #[test]
2387 fn cancelled_request_logs_499_with_reason() {
2388 let (sink, records) = collecting_sink();
2389 let router = Router::new()
2390 .route("/slow", get(FnHandler::new(ok_handler)))
2391 .with_default_trace_time_getter(fixed_time)
2392 .with_default_trace_record_sink(sink);
2393
2394 let cx = Cx::for_testing();
2395 cx.set_cancel_requested(true);
2396 cx.set_cancel_reason(CancelReason::user("client disconnected"));
2397 let _resp = futures_lite::future::block_on(
2398 router.handle_with_cx(&cx, Request::new("GET", "/slow")),
2399 );
2400
2401 let records = records.lock().expect("records lock");
2402 assert_eq!(records.len(), 1);
2403 let record = &records[0];
2404 assert_eq!(record.status, 499, "cancelled requests must log as 499");
2405 assert_eq!(record.severity, "cancelled");
2406 assert!(record.cancelled);
2407 let reason = record.cancel_reason.as_deref().expect("cancel reason");
2408 assert!(
2409 reason.contains("client disconnected"),
2410 "cancel reason must carry the message, got: {reason}"
2411 );
2412 }
2413
2414 #[test]
2416 fn without_default_trace_emits_nothing() {
2417 let (sink, records) = collecting_sink();
2418 let router = Router::new()
2419 .with_default_trace_record_sink(sink)
2420 .without_default_trace()
2421 .route("/ok", get(FnHandler::new(ok_handler)));
2422
2423 let resp = router.handle(Request::new("GET", "/ok"));
2424 assert_eq!(resp.status, StatusCode::OK);
2425 assert!(
2426 records.lock().expect("records lock").is_empty(),
2427 "opted-out router must not emit trace records"
2428 );
2429 }
2430
2431 #[test]
2433 fn client_request_id_is_propagated() {
2434 let (sink, records) = collecting_sink();
2435 let router = Router::new()
2436 .route("/ok", get(FnHandler::new(ok_handler)))
2437 .with_default_trace_time_getter(fixed_time)
2438 .with_default_trace_record_sink(sink);
2439
2440 let _ = router
2441 .handle(Request::new("GET", "/ok").with_header("x-request-id", "client-abc-123"));
2442
2443 let records = records.lock().expect("records lock");
2444 assert_eq!(records[0].request_id.as_deref(), Some("client-abc-123"));
2445 }
2446
2447 #[test]
2449 fn default_policy_does_not_mutate_response_headers() {
2450 let router = Router::new().route("/ok", get(FnHandler::new(ok_handler)));
2451 let resp = router.handle(Request::new("GET", "/ok"));
2452 assert!(!resp.headers.contains_key("x-response-time-ms"));
2453 assert!(!resp.headers.contains_key("x-trace-id"));
2454 }
2455
2456 #[test]
2458 fn opt_in_policy_stamps_headers() {
2459 let router = Router::new()
2460 .route("/ok", get(FnHandler::new(ok_handler)))
2461 .with_default_trace_policy(RequestTracePolicy::default())
2462 .with_default_trace_time_getter(fixed_time);
2463
2464 let resp = router.handle(Request::new("GET", "/ok"));
2465 assert_eq!(
2466 resp.headers.get("x-response-time-ms").map(String::as_str),
2467 Some("0")
2468 );
2469 assert_eq!(
2470 resp.headers.get("x-trace-id").map(String::as_str),
2471 Some("req-1")
2472 );
2473 }
2474 }
2475}