use std::collections::HashMap;
use hyper::Method;
use crate::handler::Handler;
#[derive(Clone, Debug, Default)]
pub struct PathParams(pub HashMap<String, String>);
#[derive(Clone, Debug, Default)]
pub struct QueryParams(pub HashMap<String, String>);
#[derive(Default)]
pub struct Router<S> {
root: Node<S>,
}
struct Node<S> {
segment: String,
param_name: String,
is_wildcard: bool,
handlers: HashMap<Method, Handler<S>>,
children: Vec<Node<S>>,
}
impl<S> Default for Node<S> {
fn default() -> Self {
Node {
segment: String::new(),
param_name: String::new(),
is_wildcard: false,
handlers: HashMap::new(),
children: Vec::new(),
}
}
}
impl<S: Send + Sync + 'static> Router<S> {
pub fn new() -> Self {
Router { root: Node::default() }
}
pub fn insert(&mut self, method: Method, path: &str, handler: Handler<S>) {
let segments = split_path(path);
let mut node = &mut self.root;
for seg in segments {
if seg == "*" {
if let Some(idx) = node.children.iter().position(|c| c.is_wildcard) {
node = &mut node.children[idx];
} else {
node.children.push(Node {
segment: "*".to_string(),
param_name: String::new(),
is_wildcard: true,
handlers: HashMap::new(),
children: Vec::new(),
});
node = node.children.last_mut().unwrap();
}
} else if let Some(param_name) = seg.strip_prefix(':') {
if let Some(idx) = node.children.iter().position(|c| c.param_name == param_name) {
node = &mut node.children[idx];
} else {
node.children.push(Node {
segment: seg.to_string(),
param_name: param_name.to_string(),
is_wildcard: false,
handlers: HashMap::new(),
children: Vec::new(),
});
node = node.children.last_mut().unwrap();
}
} else {
if let Some(idx) = node.children.iter().position(|c| c.segment == seg) {
node = &mut node.children[idx];
} else {
node.children.push(Node {
segment: seg.to_string(),
param_name: String::new(),
is_wildcard: false,
handlers: HashMap::new(),
children: Vec::new(),
});
node = node.children.last_mut().unwrap();
}
}
}
node.handlers.insert(method, handler);
}
pub fn match_route<'a>(
&'a self,
method: &Method,
path: &str,
) -> Option<(&'a Handler<S>, PathParams)> {
let segments = split_path(path);
let mut params = PathParams::default();
let node = Self::find_node(&self.root, &segments, 0, method, &mut params)?;
node.handlers.get(method).map(|h| (h, params))
}
pub fn allowed_methods(&self, path: &str) -> Vec<Method> {
let segments = split_path(path);
let mut methods = std::collections::HashSet::new();
let mut params = PathParams::default();
Self::collect_allowed_methods(&self.root, &segments, 0, &mut params, &mut methods);
methods.into_iter().collect()
}
pub fn path_exists(&self, path: &str) -> bool {
!self.allowed_methods(path).is_empty()
}
fn find_node<'a>(
node: &'a Node<S>,
segments: &[String],
idx: usize,
method: &Method,
params: &mut PathParams,
) -> Option<&'a Node<S>> {
if idx == segments.len() {
return if node.handlers.contains_key(method) { Some(node) } else { None };
}
let seg = &segments[idx];
for child in &node.children {
if !child.is_wildcard && child.param_name.is_empty() && child.segment == *seg {
if let Some(found) = Self::find_node(child, segments, idx + 1, method, params) {
return Some(found);
}
}
}
for child in &node.children {
if !child.is_wildcard && !child.param_name.is_empty() {
let mut p = params.clone();
p.0.insert(child.param_name.clone(), seg.clone());
if let Some(found) = Self::find_node(child, segments, idx + 1, method, &mut p) {
*params = p;
return Some(found);
}
params.0.remove(&child.param_name);
}
}
for child in &node.children {
if child.is_wildcard && child.handlers.contains_key(method) {
params.0.insert("*".to_string(), segments[idx..].join("/"));
return Some(child);
}
}
None
}
fn collect_allowed_methods(
node: &Node<S>,
segments: &[String],
idx: usize,
params: &mut PathParams,
methods: &mut std::collections::HashSet<Method>,
) {
if idx == segments.len() {
methods.extend(node.handlers.keys().cloned());
return;
}
let seg = &segments[idx];
for child in &node.children {
if !child.is_wildcard && child.param_name.is_empty() && child.segment == *seg {
let mut p = params.clone();
Self::collect_allowed_methods(child, segments, idx + 1, &mut p, methods);
}
}
for child in &node.children {
if !child.is_wildcard && !child.param_name.is_empty() {
let mut p = params.clone();
p.0.insert(child.param_name.clone(), seg.clone());
Self::collect_allowed_methods(child, segments, idx + 1, &mut p, methods);
}
}
for child in &node.children {
if child.is_wildcard {
methods.extend(child.handlers.keys().cloned());
}
}
}
}
fn split_path(path: &str) -> Vec<String> {
path.trim_start_matches('/')
.split('/')
.filter(|s| !s.is_empty())
.map(|s| percent_encoding::percent_decode_str(s).decode_utf8_lossy().into_owned())
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::ServeError;
use hyper::Response;
use hyper::body::Bytes;
fn dummy_handler() -> Handler<()> {
crate::handler::handler(|_, _| async {
Ok::<_, ServeError>(Response::new(crate::handler::body(Bytes::from("ok"))))
})
}
#[test]
fn insert_and_match_static() {
let mut router = Router::new();
router.insert(Method::GET, "/hello", dummy_handler());
let (_, _) = router.match_route(&Method::GET, "/hello").unwrap();
}
#[test]
fn match_with_path_param() {
let mut router = Router::new();
router.insert(Method::GET, "/users/:id", dummy_handler());
let (_, params) = router.match_route(&Method::GET, "/users/42").unwrap();
assert_eq!(params.0.get("id").unwrap(), "42");
}
#[test]
fn match_with_wildcard() {
let mut router = Router::new();
router.insert(Method::GET, "/files/*", dummy_handler());
let (_, params) = router.match_route(&Method::GET, "/files/a/b/c").unwrap();
assert_eq!(params.0.get("*").unwrap(), "a/b/c");
}
#[test]
fn no_match_for_unregistered_route() {
let mut router = Router::new();
router.insert(Method::GET, "/hello", dummy_handler());
assert!(router.match_route(&Method::GET, "/world").is_none());
}
#[test]
fn method_mismatch_returns_none() {
let mut router = Router::new();
router.insert(Method::GET, "/hello", dummy_handler());
assert!(router.match_route(&Method::POST, "/hello").is_none());
}
#[test]
fn root_path_matches() {
let mut router = Router::new();
router.insert(Method::GET, "/", dummy_handler());
let (_, _) = router.match_route(&Method::GET, "/").unwrap();
}
#[test]
fn method_mismatch_on_static_branch_backtracks_to_param_sibling() {
let mut router = Router::new();
router.insert(Method::GET, "/users/:id", dummy_handler());
router.insert(Method::POST, "/users/new", dummy_handler());
let (_, params) = router.match_route(&Method::GET, "/users/new").unwrap();
assert_eq!(params.0.get("id").unwrap(), "new");
}
#[test]
fn static_branch_with_matching_method_still_wins_over_param_sibling() {
let mut router = Router::new();
router.insert(Method::GET, "/users/:id", dummy_handler());
router.insert(Method::GET, "/users/new", dummy_handler());
let (_, params) = router.match_route(&Method::GET, "/users/new").unwrap();
assert!(params.0.is_empty(), "static branch should win, not fall back to :id");
}
#[test]
fn allowed_methods_unions_across_ambiguous_branches() {
let mut router = Router::new();
router.insert(Method::GET, "/users/:id", dummy_handler());
router.insert(Method::POST, "/users/new", dummy_handler());
let mut allowed = router.allowed_methods("/users/new");
allowed.sort_by_key(|m| m.to_string());
assert_eq!(allowed, vec![Method::GET, Method::POST]);
}
#[test]
fn static_path_traversal_borrows_params_without_cloning() {
let mut router = Router::new();
router.insert(Method::GET, "/api/v1/users/profile", dummy_handler());
let (_, params) = router.match_route(&Method::GET, "/api/v1/users/profile").unwrap();
assert!(params.0.is_empty());
}
#[test]
fn param_backtracking_truncates_params_on_failure() {
let mut router = Router::new();
router.insert(Method::GET, "/users/:id", dummy_handler());
router.insert(Method::POST, "/users/new", dummy_handler());
let (_, params) = router.match_route(&Method::GET, "/users/new").unwrap();
assert_eq!(params.0.get("id").unwrap(), "new");
assert_eq!(params.0.len(), 1);
}
}