use std::borrow::Cow;
use std::collections::HashMap;
use crate::errors::OrionError;
use crate::storage::models::{Channel, ChannelProtocol};
struct RouteEntry {
channel_name: String,
methods: Vec<String>,
segments: Vec<RouteSegment>,
priority: i64,
}
impl RouteEntry {
fn canonical(&self) -> String {
canonical_segments(&self.segments)
}
}
#[derive(Debug, Clone)]
enum RouteSegment {
Static(String),
Param(String),
}
#[derive(Debug, Clone)]
pub struct RouteMatch {
pub channel_name: String,
pub params: HashMap<String, String>,
}
pub fn canonical_route(pattern: &str) -> String {
canonical_segments(&parse_route_pattern(pattern))
}
fn canonical_segments(segments: &[RouteSegment]) -> String {
let mut out = String::new();
for segment in segments {
out.push('/');
match segment {
RouteSegment::Static(s) => out.push_str(s),
RouteSegment::Param(_) => out.push_str("{}"),
}
}
if out.is_empty() {
out.push('/');
}
out
}
pub(crate) fn declared_route(ch: &Channel) -> Option<(String, Vec<String>)> {
let (segments, methods) = declared_segments(ch)?;
Some((canonical_segments(&segments), methods))
}
fn declared_segments(ch: &Channel) -> Option<(Vec<RouteSegment>, Vec<String>)> {
if ch.protocol != ChannelProtocol::Rest.as_str()
&& ch.protocol != ChannelProtocol::Http.as_str()
{
return None;
}
let pattern = ch.route_pattern.as_deref()?;
Some((
parse_route_pattern(pattern),
ch.methods().unwrap_or_default(),
))
}
pub fn methods_overlap(a: &[String], b: &[String]) -> bool {
if a.is_empty() || b.is_empty() {
return true;
}
a.iter()
.any(|x| b.iter().any(|y| x.eq_ignore_ascii_case(y)))
}
fn parse_route_pattern(pattern: &str) -> Vec<RouteSegment> {
pattern
.split('/')
.filter(|s| !s.is_empty())
.map(|seg| {
if seg.starts_with('{') && seg.ends_with('}') {
RouteSegment::Param(seg[1..seg.len() - 1].to_string())
} else {
RouteSegment::Static(seg.to_string())
}
})
.collect()
}
pub(crate) fn percent_decode_segment(segment: &str) -> Option<Cow<'_, str>> {
if !segment.contains('%') {
return Some(Cow::Borrowed(segment));
}
let bytes = segment.as_bytes();
let mut out = Vec::with_capacity(bytes.len());
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'%' {
let hi = bytes.get(i + 1).and_then(|b| (*b as char).to_digit(16))?;
let lo = bytes.get(i + 2).and_then(|b| (*b as char).to_digit(16))?;
out.push((hi * 16 + lo) as u8);
i += 3;
} else {
out.push(bytes[i]);
i += 1;
}
}
String::from_utf8(out).ok().map(Cow::Owned)
}
fn decode_path_parts(path: &str) -> Option<Vec<Cow<'_, str>>> {
path.split('/')
.filter(|s| !s.is_empty())
.map(percent_decode_segment)
.collect()
}
fn match_segments<S: AsRef<str>>(
segments: &[RouteSegment],
path_parts: &[S],
) -> Option<HashMap<String, String>> {
if segments.len() != path_parts.len() {
return None;
}
let param_count = segments
.iter()
.filter(|s| matches!(s, RouteSegment::Param(_)))
.count();
let mut params = HashMap::with_capacity(param_count);
for (seg, part) in segments.iter().zip(path_parts.iter()) {
match seg {
RouteSegment::Static(expected) => {
if expected.as_str() != part.as_ref() {
return None;
}
}
RouteSegment::Param(name) => {
params.insert(name.clone(), part.as_ref().to_string());
}
}
}
Some(params)
}
#[derive(Default)]
pub struct RouteTable {
entries: Vec<RouteEntry>,
by_first_static: HashMap<(usize, String), Vec<usize>>,
by_param_first: HashMap<usize, Vec<usize>>,
}
impl RouteTable {
pub(super) fn build<'a>(channels: impl IntoIterator<Item = &'a Channel>) -> Self {
let mut entries: Vec<RouteEntry> = channels
.into_iter()
.filter_map(|ch| {
let (segments, methods) = declared_segments(ch)?;
let methods: Vec<String> = methods.into_iter().map(|m| m.to_uppercase()).collect();
Some(RouteEntry {
channel_name: ch.name.clone(),
methods,
segments,
priority: ch.priority,
})
})
.collect();
entries.sort_by(|a, b| {
b.priority
.cmp(&a.priority)
.then_with(|| b.segments.len().cmp(&a.segments.len()))
.then_with(|| a.channel_name.cmp(&b.channel_name))
});
let table = Self::from_sorted_entries(entries);
table.warn_on_conflicts();
table
}
fn from_sorted_entries(entries: Vec<RouteEntry>) -> Self {
let mut by_first_static: HashMap<(usize, String), Vec<usize>> = HashMap::new();
let mut by_param_first: HashMap<usize, Vec<usize>> = HashMap::new();
for (i, entry) in entries.iter().enumerate() {
match entry.segments.first() {
Some(RouteSegment::Static(s)) => by_first_static
.entry((entry.segments.len(), s.clone()))
.or_default()
.push(i),
Some(RouteSegment::Param(_)) | None => by_param_first
.entry(entry.segments.len())
.or_default()
.push(i),
}
}
Self {
entries,
by_first_static,
by_param_first,
}
}
fn warn_on_conflicts(&self) {
let canonicals: Vec<String> = self.entries.iter().map(RouteEntry::canonical).collect();
for (i, entry) in self.entries.iter().enumerate() {
let canonical = &canonicals[i];
for (j, winner) in self.entries[..i].iter().enumerate() {
if &canonicals[j] == canonical && methods_overlap(&winner.methods, &entry.methods) {
tracing::warn!(
route = %canonical,
shadowed_channel = %entry.channel_name,
serving_channel = %winner.channel_name,
"Two active channels claim the same route; the shadowed one is \
unreachable by path (it is still reachable by name). Change its \
route_pattern, methods or priority."
);
break;
}
}
}
}
pub fn match_route(&self, method: &str, path: &str) -> Result<Option<RouteMatch>, OrionError> {
let Some(path_parts) = decode_path_parts(path) else {
return Err(OrionError::validation(
"Invalid percent-encoding in request path".to_string(),
));
};
let empty: Vec<usize> = Vec::new();
let static_bucket = path_parts
.first()
.and_then(|first| {
self.by_first_static
.get(&(path_parts.len(), first.as_ref().to_string()))
})
.unwrap_or(&empty);
let param_bucket = self.by_param_first.get(&path_parts.len()).unwrap_or(&empty);
let (mut a, mut b) = (0, 0);
while a < static_bucket.len() || b < param_bucket.len() {
let idx = match (static_bucket.get(a), param_bucket.get(b)) {
(Some(&x), Some(&y)) if x < y => {
a += 1;
x
}
(Some(_), Some(&y)) => {
b += 1;
y
}
(Some(&x), None) => {
a += 1;
x
}
(None, Some(&y)) => {
b += 1;
y
}
(None, None) => unreachable!("loop condition"),
};
let entry = &self.entries[idx];
if !entry.methods.is_empty()
&& !entry.methods.iter().any(|m| m.eq_ignore_ascii_case(method))
{
continue;
}
if let Some(params) = match_segments(&entry.segments, &path_parts) {
return Ok(Some(RouteMatch {
channel_name: entry.channel_name.clone(),
params,
}));
}
}
Ok(None)
}
#[cfg(test)]
fn match_route_linear(
&self,
method: &str,
path: &str,
) -> Result<Option<RouteMatch>, OrionError> {
let Some(path_parts) = decode_path_parts(path) else {
return Err(OrionError::validation(
"Invalid percent-encoding in request path".to_string(),
));
};
for entry in &self.entries {
if !entry.methods.is_empty()
&& !entry.methods.iter().any(|m| m.eq_ignore_ascii_case(method))
{
continue;
}
if let Some(params) = match_segments(&entry.segments, &path_parts) {
return Ok(Some(RouteMatch {
channel_name: entry.channel_name.clone(),
params,
}));
}
}
Ok(None)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_route_pattern_simple() {
let segments = parse_route_pattern("/orders");
assert_eq!(segments.len(), 1);
assert!(matches!(&segments[0], RouteSegment::Static(s) if s == "orders"));
}
#[test]
fn test_parse_route_pattern_with_params() {
let segments = parse_route_pattern("/orders/{id}/items/{item_id}");
assert_eq!(segments.len(), 4);
assert!(matches!(&segments[0], RouteSegment::Static(s) if s == "orders"));
assert!(matches!(&segments[1], RouteSegment::Param(s) if s == "id"));
assert!(matches!(&segments[2], RouteSegment::Static(s) if s == "items"));
assert!(matches!(&segments[3], RouteSegment::Param(s) if s == "item_id"));
}
fn parts(parts: &[&str]) -> Vec<String> {
parts.iter().map(|s| s.to_string()).collect()
}
#[test]
fn test_match_segments_exact() {
let segments = parse_route_pattern("/orders/{id}");
let params = match_segments(&segments, &parts(&["orders", "123"]));
assert!(params.is_some());
assert_eq!(params.expect("test").get("id").expect("test"), "123");
}
#[test]
fn test_match_segments_no_match() {
let segments = parse_route_pattern("/orders/{id}");
assert!(match_segments(&segments, &parts(&["users", "123"])).is_none());
assert!(match_segments(&segments, &parts(&["orders"])).is_none());
assert!(match_segments(&segments, &parts(&["orders", "123", "items"])).is_none());
}
#[test]
fn test_match_segments_is_case_sensitive() {
let segments = parse_route_pattern("/orders/{id}");
assert!(match_segments(&segments, &parts(&["ORDERS", "1"])).is_none());
assert!(match_segments(&segments, &parts(&["Orders", "1"])).is_none());
assert!(match_segments(&segments, &parts(&["orders", "1"])).is_some());
}
#[test]
fn test_route_table_match() {
let table = RouteTable::from_sorted_entries(vec![RouteEntry {
channel_name: "orders.get".to_string(),
methods: vec!["GET".to_string()],
segments: parse_route_pattern("/orders/{id}"),
priority: 0,
}]);
let result = table.match_route("GET", "orders/42").expect("valid path");
assert!(result.is_some());
let rm = result.expect("test");
assert_eq!(rm.channel_name, "orders.get");
assert_eq!(rm.params.get("id").expect("test"), "42");
}
#[test]
fn test_route_table_method_mismatch() {
let table = RouteTable::from_sorted_entries(vec![RouteEntry {
channel_name: "orders.get".to_string(),
methods: vec!["GET".to_string()],
segments: parse_route_pattern("/orders/{id}"),
priority: 0,
}]);
assert!(
table
.match_route("POST", "orders/42")
.expect("valid path")
.is_none()
);
}
fn orders_table() -> RouteTable {
RouteTable::from_sorted_entries(vec![RouteEntry {
channel_name: "orders.get".to_string(),
methods: vec!["GET".to_string()],
segments: parse_route_pattern("/orders/{id}"),
priority: 0,
}])
}
#[test]
fn test_params_are_percent_decoded_once() {
let m = orders_table()
.match_route("GET", "orders/a%2Fb")
.expect("valid path")
.expect("must match");
assert_eq!(m.params.get("id").expect("id"), "a/b");
let m = orders_table()
.match_route("GET", "orders/a%252Fb")
.expect("valid path")
.expect("must match");
assert_eq!(m.params.get("id").expect("id"), "a%2Fb");
}
#[test]
fn test_encoded_static_segment_matches() {
let m = orders_table()
.match_route("GET", "%6Frders/1")
.expect("valid path");
assert!(m.is_some());
let m = orders_table()
.match_route("GET", "%4Frders/1") .expect("valid path");
assert!(m.is_none());
}
#[test]
fn test_invalid_percent_sequence_is_rejected() {
for path in ["orders/a%ZZ", "orders/a%2", "orders/%", "orders/%G1"] {
assert!(
orders_table().match_route("GET", path).is_err(),
"{path} must be rejected"
);
}
assert!(orders_table().match_route("GET", "orders/%FF").is_err());
}
#[test]
fn test_route_table_priority_ordering() {
let table = RouteTable::from_sorted_entries(vec![
RouteEntry {
channel_name: "low".to_string(),
methods: vec![],
segments: parse_route_pattern("/items/{id}"),
priority: 0,
},
RouteEntry {
channel_name: "high".to_string(),
methods: vec![],
segments: parse_route_pattern("/items/{id}"),
priority: 10,
},
]);
assert_eq!(
table
.match_route("GET", "items/1")
.expect("valid path")
.expect("test")
.channel_name,
"low"
);
}
}
#[cfg(test)]
mod prop_tests {
use super::*;
use proptest::prelude::*;
fn table() -> RouteTable {
RouteTable::from_sorted_entries(vec![
RouteEntry {
channel_name: "users.get".to_string(),
methods: vec!["GET".to_string()],
segments: parse_route_pattern("/users/{id}/orders/{oid}"),
priority: 0,
},
RouteEntry {
channel_name: "static.post".to_string(),
methods: vec!["POST".to_string()],
segments: parse_route_pattern("/a/b/c"),
priority: 0,
},
])
}
fn wide_table() -> RouteTable {
RouteTable::from_sorted_entries(vec![
RouteEntry {
channel_name: "orders.high".to_string(),
methods: vec!["GET".to_string()],
segments: parse_route_pattern("/orders/{id}"),
priority: 10,
},
RouteEntry {
channel_name: "orders.low".to_string(),
methods: vec![],
segments: parse_route_pattern("/orders/{id}"),
priority: 0,
},
RouteEntry {
channel_name: "param.first".to_string(),
methods: vec!["GET".to_string()],
segments: parse_route_pattern("/{tenant}/orders"),
priority: 0,
},
RouteEntry {
channel_name: "deep".to_string(),
methods: vec!["POST".to_string()],
segments: parse_route_pattern("/a/b/c"),
priority: 0,
},
RouteEntry {
channel_name: "single".to_string(),
methods: vec!["GET".to_string()],
segments: parse_route_pattern("/orders"),
priority: 0,
},
])
}
proptest! {
#[test]
fn match_route_is_total(method in ".*", path in ".*") {
let _ = table().match_route(&method, &path);
}
#[test]
fn indexed_match_equals_the_linear_scan(
method in "(GET|POST|PUT|.*)",
path in "[a-c/{}%2F]{0,24}",
) {
let t = wide_table();
let indexed = t.match_route(&method, &path);
let linear = t.match_route_linear(&method, &path);
match (indexed, linear) {
(Ok(i), Ok(l)) => {
prop_assert_eq!(
i.as_ref().map(|m| (&m.channel_name, &m.params)),
l.as_ref().map(|m| (&m.channel_name, &m.params))
);
}
(Err(_), Err(_)) => {}
(i, l) => prop_assert!(false, "indexed={i:?} linear={l:?}"),
}
}
#[test]
fn indexed_match_equals_the_linear_scan_on_real_shapes(
first in "(orders|a|zzz)",
second in "[a-z]{1,4}",
method in "(GET|POST|DELETE)",
depth in 1usize..4,
) {
let t = wide_table();
let path = match depth {
1 => first.clone(),
2 => format!("{first}/{second}"),
_ => format!("{first}/{second}/c"),
};
let indexed = t.match_route(&method, &path).expect("valid path");
let linear = t.match_route_linear(&method, &path).expect("valid path");
prop_assert_eq!(
indexed.as_ref().map(|m| (&m.channel_name, &m.params)),
linear.as_ref().map(|m| (&m.channel_name, &m.params))
);
}
#[test]
fn extracted_params_round_trip(id in "[^/%]+", oid in "[^/%]+") {
let path = format!("users/{id}/orders/{oid}");
let m = table()
.match_route("GET", &path)
.expect("percent-free paths are always valid")
.expect("must match");
prop_assert_eq!(m.channel_name.as_str(), "users.get");
prop_assert_eq!(m.params.get("id").expect("id").as_str(), id.as_str());
prop_assert_eq!(m.params.get("oid").expect("oid").as_str(), oid.as_str());
}
#[test]
fn encoded_params_decode_to_the_original(id in "[^/%]+", oid in "[^/%]+") {
fn encode(s: &str) -> String {
s.bytes().map(|b| format!("%{b:02X}")).collect()
}
let path = format!("users/{}/orders/{}", encode(&id), encode(&oid));
let m = table()
.match_route("GET", &path)
.expect("fully-encoded segments are valid")
.expect("must match");
prop_assert_eq!(m.params.get("id").expect("id").as_str(), id.as_str());
prop_assert_eq!(m.params.get("oid").expect("oid").as_str(), oid.as_str());
}
#[test]
fn wrong_arity_never_matches(id in "[^/%]+") {
let short = format!("users/{id}");
let long = format!("users/{id}/orders/{id}/extra");
prop_assert!(table().match_route("GET", &short).expect("valid").is_none());
prop_assert!(table().match_route("GET", &long).expect("valid").is_none());
}
}
}