use super::Request;
use super::Response;
use super::method::Method;
use super::request::Params;
use super::sse::SseWriter;
use super::stream::StreamResponse;
#[cfg(feature = "ws")]
use super::websocket::WsConn;
use crate::RuntimeError;
use arrayvec::ArrayVec;
use std::collections::BTreeMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
pub type HandlerOutcome = Result<Response, RuntimeError>;
pub(super) type Handler =
Box<dyn Fn(&Request) -> Pin<Box<dyn Future<Output = HandlerOutcome> + Send>> + Send + Sync>;
pub(super) type SseHandler =
Arc<dyn Fn(&Request, &mut SseWriter) -> Result<(), RuntimeError> + Send + Sync>;
pub(super) type StreamHandler =
Box<dyn Fn(&Request) -> Pin<Box<dyn Future<Output = StreamResponse> + Send>> + Send + Sync>;
#[cfg(feature = "ws")]
pub(super) type WsHandler = Arc<dyn Fn(&Request, WsConn) -> Result<(), RuntimeError> + Send + Sync>;
pub(super) enum RouteHandler {
Async(Handler),
Stream(StreamHandler),
Sse(SseHandler),
#[cfg(feature = "ws")]
WebSocket(WsHandler),
Proxy {
backend: Arc<str>,
prefix: Arc<str>,
healthy: Option<Arc<AtomicBool>>,
},
ProxyStream {
backend: Arc<str>,
prefix: Arc<str>,
healthy: Option<Arc<AtomicBool>>,
},
}
type CaptureVec<'a> = Vec<&'a str>;
struct RegisteredHandler {
method: Method,
handler: RouteHandler,
capture_names: Box<[Arc<str>]>,
route: Arc<str>,
}
type MethodHandlers = Vec<RegisteredHandler>;
pub(super) const CANONICAL_METHODS: [Method; Method::COUNT] = [
Method::Get,
Method::Head,
Method::Post,
Method::Put,
Method::Patch,
Method::Delete,
Method::Options,
];
type MethodMask = u8;
const _: () = assert!(
Method::COUNT <= MethodMask::BITS as usize,
"MethodMask has fewer bits than there are methods"
);
const _: () = {
let mut listed = [false; Method::COUNT];
let mut position = 0;
while position < Method::COUNT {
let ordinal = CANONICAL_METHODS[position].ordinal();
assert!(!listed[ordinal], "CANONICAL_METHODS lists a method twice");
listed[ordinal] = true;
position += 1;
}
};
pub(super) enum RouteLookup<'n, 'p> {
Matched(Selected<'n, 'p>),
MethodMismatch {
route: Arc<str>,
allow: Arc<str>,
},
Unmatched,
}
pub(super) struct Selected<'n, 'p> {
pub(super) route: &'n Arc<str>,
pub(super) handler: &'n RouteHandler,
capture_names: &'n [Arc<str>],
captures: CaptureVec<'p>,
}
impl Selected<'_, '_> {
pub(super) fn bind_params(self) -> Params {
debug_assert_eq!(
self.capture_names.len(),
self.captures.len(),
"a pattern names one capture per segment its match collected"
);
self.capture_names
.iter()
.cloned()
.zip(self.captures)
.map(|(name, value)| (name, Box::from(value)))
.collect()
}
}
enum Segment {
Static(Box<str>),
Param(Box<str>),
Wildcard(Box<str>),
}
impl Segment {
fn render(&self, into: &mut String) {
match self {
Self::Static(name) => into.push_str(name),
Self::Param(name) => {
into.push(':');
into.push_str(name);
}
Self::Wildcard(name) => {
into.push('*');
into.push_str(name);
}
}
}
}
fn parse_segments(path: &str) -> Box<[Segment]> {
path.split('/')
.filter(|s| !s.is_empty())
.map(|s| match (s.strip_prefix('*'), s.strip_prefix(':')) {
(Some(name), _) => Segment::Wildcard(name.into()),
(_, Some(name)) => Segment::Param(name.into()),
_ => Segment::Static(s.into()),
})
.collect()
}
fn normalize_route(segments: &[Segment]) -> Arc<str> {
match segments.is_empty() {
true => Arc::from("/"),
false => {
let mut pattern = String::new();
for segment in segments {
pattern.push('/');
segment.render(&mut pattern);
}
Arc::from(pattern.as_str())
}
}
}
pub(crate) struct TrieNode {
static_children: BTreeMap<Box<str>, TrieNode>,
param_child: Option<Box<TrieNode>>,
wildcard: Option<MethodHandlers>,
handlers: MethodHandlers,
}
impl TrieNode {
pub(crate) fn new() -> Self {
Self {
static_children: BTreeMap::new(),
param_child: None,
wildcard: None,
handlers: Vec::new(),
}
}
pub(crate) fn insert_route(&mut self, method: Method, path: &str, handler: RouteHandler) {
let segments = parse_segments(path);
let capture_names = segments
.iter()
.filter_map(|segment| match segment {
Segment::Param(name) | Segment::Wildcard(name) => {
Some(Arc::<str>::from(name.as_ref()))
}
Segment::Static(_) => None,
})
.collect();
let registered = RegisteredHandler {
method,
handler,
capture_names,
route: normalize_route(&segments),
};
self.insert_segments(&segments, registered);
}
fn insert_segments(&mut self, segments: &[Segment], registered: RegisteredHandler) {
match segments.first() {
None => self.handlers.push(registered),
Some(Segment::Static(name)) => {
let child = self
.static_children
.entry(name.clone())
.or_insert_with(TrieNode::new);
child.insert_segments(&segments[1..], registered);
}
Some(Segment::Param(_)) => {
let child = self
.param_child
.get_or_insert_with(|| Box::new(TrieNode::new()));
child.insert_segments(&segments[1..], registered);
}
Some(Segment::Wildcard(_)) => {
self.wildcard.get_or_insert_with(Vec::new).push(registered);
}
}
}
pub(crate) fn freeze(self) -> FrozenNode {
let static_children: Box<[(Box<str>, FrozenNode)]> = self
.static_children
.into_iter()
.map(|(k, v)| (k, v.freeze()))
.collect();
FrozenNode {
static_children,
param_child: self.param_child.map(|node| Box::new(node.freeze())),
wildcard: self.wildcard.map(freeze_handlers),
handlers: freeze_handlers(self.handlers),
}
}
}
fn freeze_handlers(registered: MethodHandlers) -> FrozenMethodHandlers {
let mut by_method: MethodSlots = Default::default();
for entry in registered {
by_method[entry.method.ordinal()] = Some(FrozenHandler {
handler: entry.handler,
capture_names: entry.capture_names,
route: entry.route,
});
}
FrozenMethodHandlers {
claimed: freeze_claim(&by_method),
by_method,
}
}
fn freeze_claim(by_method: &MethodSlots) -> Option<Claimed> {
let route = CANONICAL_METHODS
.iter()
.find_map(|method| by_method[method.ordinal()].as_ref())
.map(|handler| Arc::clone(&handler.route))?;
let served = served_mask(by_method);
Some(Claimed {
route,
allow: render_allow(served),
served,
})
}
fn served_mask(by_method: &MethodSlots) -> MethodMask {
CANONICAL_METHODS
.into_iter()
.filter(|method| slot_serves(by_method, *method))
.fold(0, |mask, method| mask | method_bit(method))
}
fn method_bit(method: Method) -> MethodMask {
1 << method.ordinal()
}
fn render_allow(served: MethodMask) -> Arc<str> {
let rendered = CANONICAL_METHODS
.into_iter()
.filter(|method| served & method_bit(*method) != 0)
.fold(String::new(), |mut rendered, method| {
let separator = match rendered.is_empty() {
true => "",
false => ", ",
};
rendered.push_str(separator);
rendered.push_str(method.as_str());
rendered
});
Arc::from(rendered.as_str())
}
struct FrozenHandler {
handler: RouteHandler,
capture_names: Box<[Arc<str>]>,
route: Arc<str>,
}
type MethodSlots = [Option<FrozenHandler>; Method::COUNT];
struct Claimed {
route: Arc<str>,
allow: Arc<str>,
served: MethodMask,
}
struct FrozenMethodHandlers {
by_method: MethodSlots,
claimed: Option<Claimed>,
}
pub(crate) struct FrozenNode {
static_children: Box<[(Box<str>, FrozenNode)]>,
param_child: Option<Box<FrozenNode>>,
wildcard: Option<FrozenMethodHandlers>,
handlers: FrozenMethodHandlers,
}
type Accepts<'v, 'n> = &'v mut dyn FnMut(&'n FrozenMethodHandlers) -> bool;
type PathMatch<'n, 'p> = (&'n FrozenMethodHandlers, CaptureVec<'p>);
impl FrozenNode {
pub(super) fn lookup<'n, 'p>(
&'n self,
method: Option<Method>,
path: &'p str,
segments: &[&'p str],
) -> RouteLookup<'n, 'p> {
match method.and_then(|method| self.select(method, path, segments)) {
Some(selected) => RouteLookup::Matched(selected),
None => self.refuse_method(path, segments),
}
}
pub(super) fn select<'n, 'p>(
&'n self,
method: Method,
path: &'p str,
segments: &[&'p str],
) -> Option<Selected<'n, 'p>> {
let mut accepts =
|handlers: &'n FrozenMethodHandlers| slot_serves(&handlers.by_method, method);
let (handlers, mut captures) = self.resolve(&mut accepts, path, segments)?;
let selected = select_slot(&handlers.by_method, method)?;
captures.reverse();
Some(Selected {
route: &selected.route,
handler: &selected.handler,
capture_names: &selected.capture_names,
captures,
})
}
fn refuse_method<'n, 'p>(&'n self, path: &'p str, segments: &[&'p str]) -> RouteLookup<'n, 'p> {
let mut claims: Vec<&'n Claimed> = Vec::new();
let mut collect = |handlers: &'n FrozenMethodHandlers| {
claims.extend(handlers.claimed.as_ref());
false
};
let matched = self.resolve(&mut collect, path, segments);
debug_assert!(matched.is_none(), "a collecting pass accepts no node");
match claims.split_first() {
Some((first, rest)) => RouteLookup::MethodMismatch {
route: Arc::clone(&first.route),
allow: merge_allow(first, rest),
},
None => RouteLookup::Unmatched,
}
}
fn resolve<'n, 'p>(
&'n self,
accepts: Accepts<'_, 'n>,
path: &'p str,
segments: &[&'p str],
) -> Option<PathMatch<'n, 'p>> {
let Some(&segment) = segments.first() else {
return accepts(&self.handlers).then(|| (&self.handlers, CaptureVec::new()));
};
let rest = &segments[1..];
if let Some(matched) = self.resolve_static(accepts, path, segment, rest) {
return Some(matched);
}
if let Some(matched) = self.resolve_param(accepts, path, segment, rest) {
return Some(matched);
}
self.resolve_wildcard(accepts, path, segments)
}
fn resolve_static<'n, 'p>(
&'n self,
accepts: Accepts<'_, 'n>,
path: &'p str,
segment: &str,
rest: &[&'p str],
) -> Option<PathMatch<'n, 'p>> {
let idx = self
.static_children
.binary_search_by_key(&segment, |(k, _)| k)
.ok()?;
self.static_children[idx].1.resolve(accepts, path, rest)
}
fn resolve_param<'n, 'p>(
&'n self,
accepts: Accepts<'_, 'n>,
path: &'p str,
segment: &'p str,
rest: &[&'p str],
) -> Option<PathMatch<'n, 'p>> {
let child = self.param_child.as_ref()?;
let mut result = child.resolve(accepts, path, rest)?;
result.1.push(segment);
Some(result)
}
fn resolve_wildcard<'n, 'p>(
&'n self,
accepts: Accepts<'_, 'n>,
path: &'p str,
segments: &[&'p str],
) -> Option<PathMatch<'n, 'p>> {
let handlers = self.wildcard.as_ref()?;
accepts(handlers).then(|| (handlers, vec![wildcard_span(path, segments)]))
}
}
fn merge_allow(first: &Claimed, rest: &[&Claimed]) -> Arc<str> {
let served = rest
.iter()
.fold(first.served, |mask, claim| mask | claim.served);
match served == first.served {
true => Arc::clone(&first.allow),
false => render_allow(served),
}
}
fn slot_serves(by_method: &MethodSlots, method: Method) -> bool {
select_slot(by_method, method).is_some()
}
fn select_slot(by_method: &MethodSlots, method: Method) -> Option<&FrozenHandler> {
by_method[method.ordinal()]
.as_ref()
.or_else(|| match method {
Method::Head => by_method[Method::Get.ordinal()].as_ref(),
_ => None,
})
}
fn wildcard_span<'a>(path: &'a str, segments: &[&'a str]) -> &'a str {
match (segments.first(), segments.last()) {
(Some(first), Some(last)) => {
let start = first.as_ptr() as usize - path.as_ptr() as usize;
let end = last.as_ptr() as usize - path.as_ptr() as usize + last.len();
&path[start..end]
}
_ => "",
}
}
pub(super) const PATH_SEGMENT_LIMIT: usize = 32;
pub(super) fn split_path_segments(path: &str) -> Option<ArrayVec<&str, PATH_SEGMENT_LIMIT>> {
let mut segments = ArrayVec::new();
for seg in path.split('/').filter(|s| !s.is_empty()) {
match segments.try_push(seg) {
Ok(()) => {}
Err(_) => return None,
}
}
Some(segments)
}