mod response;
pub use response::IntoResponse;
use crate::handler::Handler;
use crate::proto::{Method, Request, Response, StatusCode};
enum Seg {
Lit(String),
Param(String),
Wildcard(String),
}
pub trait RouteHandler: Send + Sync {
fn call(&self, req: &Request) -> Response;
}
impl<F, R> RouteHandler for F
where
F: Fn(&Request) -> R + Send + Sync,
R: IntoResponse,
{
fn call(&self, req: &Request) -> Response {
(self)(req).into_response()
}
}
struct Route {
method: Method,
segs: Vec<Seg>,
handler: Box<dyn RouteHandler>,
}
impl Route {
fn matches(&self, path_segs: &[&str]) -> Option<Vec<(String, String)>> {
let mut params = Vec::new();
let mut i = 0;
for seg in &self.segs {
match seg {
Seg::Lit(lit) => {
if path_segs.get(i)? != lit {
return None;
}
i += 1;
}
Seg::Param(name) => {
let value = path_segs.get(i)?;
params.push((name.clone(), (*value).to_owned()));
i += 1;
}
Seg::Wildcard(name) => {
params.push((name.clone(), path_segs[i..].join("/")));
return Some(params);
}
}
}
if i == path_segs.len() {
Some(params)
} else {
None
}
}
}
#[derive(Default)]
pub struct Router {
routes: Vec<Route>,
fallback: Option<Box<dyn RouteHandler>>,
}
impl Router {
pub fn new() -> Router {
Router {
routes: Vec::new(),
fallback: None,
}
}
pub fn route<H>(mut self, method: Method, pattern: &str, handler: H) -> Router
where
H: RouteHandler + 'static,
{
self.routes.push(Route {
method,
segs: parse_pattern(pattern),
handler: Box::new(handler),
});
self
}
pub fn get<H: RouteHandler + 'static>(self, pattern: &str, handler: H) -> Router {
self.route(Method::Get, pattern, handler)
}
pub fn post<H: RouteHandler + 'static>(self, pattern: &str, handler: H) -> Router {
self.route(Method::Post, pattern, handler)
}
pub fn put<H: RouteHandler + 'static>(self, pattern: &str, handler: H) -> Router {
self.route(Method::Put, pattern, handler)
}
pub fn delete<H: RouteHandler + 'static>(self, pattern: &str, handler: H) -> Router {
self.route(Method::Delete, pattern, handler)
}
pub fn patch<H: RouteHandler + 'static>(self, pattern: &str, handler: H) -> Router {
self.route(Method::Patch, pattern, handler)
}
pub fn head<H: RouteHandler + 'static>(self, pattern: &str, handler: H) -> Router {
self.route(Method::Head, pattern, handler)
}
pub fn options<H: RouteHandler + 'static>(self, pattern: &str, handler: H) -> Router {
self.route(Method::Options, pattern, handler)
}
pub fn fallback<H: RouteHandler + 'static>(mut self, handler: H) -> Router {
self.fallback = Some(Box::new(handler));
self
}
fn dispatch(&self, req: &Request) -> Response {
let path_segs = segments(req.path());
let mut allowed: Vec<&str> = Vec::new();
for route in &self.routes {
let Some(params) = route.matches(&path_segs) else {
continue;
};
if &route.method != req.method() {
let token = route.method.as_str();
if !allowed.contains(&token) {
allowed.push(token);
}
continue;
}
if params.is_empty() {
return route.handler.call(req);
}
let mut routed = req.clone();
routed.set_params(params);
return route.handler.call(&routed);
}
if !allowed.is_empty() {
return Response::status(StatusCode::METHOD_NOT_ALLOWED)
.header("Allow", allowed.join(", "));
}
match &self.fallback {
Some(handler) => handler.call(req),
None => Response::status(StatusCode::NOT_FOUND),
}
}
}
impl Handler for Router {
fn handle(&self, req: &Request) -> Response {
self.dispatch(req)
}
}
fn segments(path: &str) -> Vec<&str> {
let trimmed = path.trim_matches('/');
if trimmed.is_empty() {
Vec::new()
} else {
trimmed.split('/').collect()
}
}
fn parse_pattern(pattern: &str) -> Vec<Seg> {
segments(pattern)
.into_iter()
.map(|seg| {
if let Some(name) = seg.strip_prefix(':') {
Seg::Param(name.to_owned())
} else if let Some(name) = seg.strip_prefix('*') {
Seg::Wildcard(if name.is_empty() {
"*".to_owned()
} else {
name.to_owned()
})
} else {
Seg::Lit(seg.to_owned())
}
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn req(method: Method, target: &str) -> Request {
Request::new(
method,
target.to_owned(),
crate::proto::Version::Http11,
crate::proto::Headers::new(),
Vec::new(),
)
}
fn body(resp: &Response) -> String {
String::from_utf8_lossy(resp.body_ref().as_bytes()).into_owned()
}
#[test]
fn static_route_and_404() {
let app = Router::new().get("/", |_: &Request| "root");
assert_eq!(body(&app.handle(&req(Method::Get, "/"))), "root");
assert_eq!(
app.handle(&req(Method::Get, "/missing")).status_code(),
StatusCode::NOT_FOUND
);
}
#[test]
fn path_params_captured() {
let app = Router::new().get("/users/:id", |r: &Request| {
format!("id={}", r.param("id").unwrap_or("?"))
});
assert_eq!(body(&app.handle(&req(Method::Get, "/users/42"))), "id=42");
assert_eq!(body(&app.handle(&req(Method::Get, "/users/42/"))), "id=42");
assert_eq!(
app.handle(&req(Method::Get, "/users/42/x")).status_code(),
StatusCode::NOT_FOUND
);
}
#[test]
fn wildcard_captures_remainder() {
let app = Router::new().get("/static/*path", |r: &Request| {
r.param("path").unwrap_or("").to_owned()
});
assert_eq!(
body(&app.handle(&req(Method::Get, "/static/css/app.css"))),
"css/app.css"
);
assert_eq!(body(&app.handle(&req(Method::Get, "/static/"))), "");
}
#[test]
fn method_mismatch_is_405_with_allow() {
let app = Router::new()
.get("/x", |_: &Request| "g")
.post("/x", |_: &Request| "p");
let resp = app.handle(&req(Method::Delete, "/x"));
assert_eq!(resp.status_code(), StatusCode::METHOD_NOT_ALLOWED);
let allow = resp.headers().get("allow").unwrap();
assert!(allow.contains("GET") && allow.contains("POST"));
}
#[test]
fn query_string_ignored_in_match() {
let app = Router::new().get("/search", |r: &Request| r.query().unwrap_or("").to_owned());
assert_eq!(body(&app.handle(&req(Method::Get, "/search?q=hi"))), "q=hi");
}
#[test]
fn fallback_used() {
let app = Router::new()
.get("/", |_: &Request| "root")
.fallback(|_: &Request| (StatusCode::FORBIDDEN, "fb"));
let resp = app.handle(&req(Method::Get, "/nope"));
assert_eq!(body(&resp), "fb");
assert_eq!(resp.status_code(), StatusCode::FORBIDDEN);
}
}