use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::RwLock;
use regex::Regex;
use serde::{Serialize, Deserialize};
use crate::config::Config;
use crate::core::{ProxyRequest, ProxyError, Router, Route};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteConfig {
pub id: String,
pub target: String,
pub path: String,
#[serde(default)]
pub filters: Vec<String>,
#[serde(default = "default_priority")]
pub priority: i32,
}
fn default_priority() -> i32 {
0
}
#[derive(Debug)]
pub struct PathRouter {
routes: RwLock<HashMap<String, Route>>,
patterns: RwLock<Vec<(Regex, Route)>>,
config: Arc<Config>,
}
impl PathRouter {
pub async fn new(config: Arc<Config>) -> Result<Self, ProxyError> {
let router = Self {
routes: RwLock::new(HashMap::new()),
patterns: RwLock::new(Vec::new()),
config,
};
router.load_routes_from_config().await?;
Ok(router)
}
async fn load_routes_from_config(&self) -> Result<(), ProxyError> {
let route_configs: Option<Vec<RouteConfig>> = self.config.get("routes")?;
if let Some(route_configs) = route_configs {
let mut sorted_routes = route_configs;
sorted_routes.sort_by(|a, b| b.priority.cmp(&a.priority));
for route_config in sorted_routes {
let route = Route {
id: route_config.id,
target_base_url: route_config.target,
path_pattern: route_config.path,
filter_ids: route_config.filters,
};
self.add_route(route).await?;
}
}
if self.routes.read().await.is_empty() {
if let Ok(Some(default_target)) = self.config.get::<String>("proxy.default_target") {
let default_route = Route {
id: "default".to_string(),
target_base_url: default_target,
path_pattern: ".*".to_string(), filter_ids: Vec::new(),
};
self.add_route(default_route).await?;
}
}
Ok(())
}
fn compile_pattern(&self, pattern: &str) -> Result<Regex, ProxyError> {
let mut regex_pattern = "^".to_string();
let mut chars = pattern.chars().peekable();
while let Some(c) = chars.next() {
match c {
':' => {
let mut param_name = String::new();
while let Some(&next_char) = chars.peek() {
if next_char.is_alphanumeric() || next_char == '_' {
param_name.push(chars.next().unwrap());
} else {
break;
}
}
regex_pattern.push_str(&format!("([^/]+)"));
},
'*' => {
regex_pattern.push_str("(.*)");
},
'.' | '^' | '$' | '|' | '+' | '?' | '(' | ')' | '[' | ']' | '{' | '}' | '\\' => {
regex_pattern.push('\\');
regex_pattern.push(c);
},
_ => {
regex_pattern.push(c);
}
}
}
regex_pattern.push('$');
Regex::new(®ex_pattern)
.map_err(|e| ProxyError::RoutingError(format!("Invalid route pattern '{}': {}", pattern, e)))
}
}
#[async_trait]
impl Router for PathRouter {
async fn route(&self, request: &ProxyRequest) -> Result<Route, ProxyError> {
let path = &request.path;
let patterns = self.patterns.read().await;
for (regex, route) in patterns.iter() {
if regex.is_match(path) {
return Ok(route.clone());
}
}
if let Ok(Some(default_target)) = self.config.get::<String>("proxy.default_target") {
return Ok(Route {
id: "default".to_string(),
target_base_url: default_target,
path_pattern: ".*".to_string(),
filter_ids: Vec::new(),
});
}
Err(ProxyError::RoutingError(format!("No route found for path: {}", path)))
}
async fn get_routes(&self) -> Vec<Route> {
self.routes.read().await.values().cloned().collect()
}
async fn add_route(&self, route: Route) -> Result<(), ProxyError> {
let regex = self.compile_pattern(&route.path_pattern)?;
{
let mut routes = self.routes.write().await;
routes.insert(route.id.clone(), route.clone());
}
{
let mut patterns = self.patterns.write().await;
patterns.push((regex, route));
patterns.sort_by(|(_, a), (_, b)| {
b.path_pattern.len().cmp(&a.path_pattern.len())
});
}
Ok(())
}
async fn remove_route(&self, route_id: &str) -> Result<(), ProxyError> {
{
let mut routes = self.routes.write().await;
if routes.remove(route_id).is_none() {
return Err(ProxyError::RoutingError(format!("Route not found: {}", route_id)));
}
}
{
let mut patterns = self.patterns.write().await;
patterns.retain(|(_, route)| route.id != route_id);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{Config, ConfigProvider};
use serde_json::{json, Value};
#[derive(Debug)]
struct MockConfigProvider {
routes: Vec<RouteConfig>,
}
impl ConfigProvider for MockConfigProvider {
fn get_raw(&self, key: &str) -> Result<Option<Value>, crate::config::ConfigError> {
if key == "routes" {
Ok(Some(json!(self.routes)))
} else if key == "proxy.default_target" {
Ok(Some(json!("http://default-target.com")))
} else {
Ok(None)
}
}
fn has(&self, key: &str) -> bool {
key == "routes" || key == "proxy.default_target"
}
fn provider_name(&self) -> &str {
"mock_provider"
}
}
#[tokio::test]
async fn test_route_matching() {
let routes = vec![
RouteConfig {
id: "api".to_string(),
target: "http://api-service.com".to_string(),
path: "/api/:version/*".to_string(),
filters: vec!["logging".to_string()],
priority: 10,
},
RouteConfig {
id: "web".to_string(),
target: "http://web-service.com".to_string(),
path: "/*".to_string(),
filters: vec![],
priority: 0,
},
];
let provider = MockConfigProvider { routes };
let config = Config::builder().with_provider(provider).build();
let router = PathRouter::new(Arc::new(config)).await.unwrap();
let api_request = ProxyRequest {
method: crate::core::HttpMethod::Get,
path: "/api/v1/users".to_string(),
query: None,
headers: reqwest::header::HeaderMap::new(),
body: Vec::new(),
context: Arc::new(RwLock::new(crate::core::RequestContext::default())),
};
let api_route = router.route(&api_request).await.unwrap();
assert_eq!(api_route.id, "api");
assert_eq!(api_route.target_base_url, "http://api-service.com");
let web_request = ProxyRequest {
method: crate::core::HttpMethod::Get,
path: "/home".to_string(),
query: None,
headers: reqwest::header::HeaderMap::new(),
body: Vec::new(),
context: Arc::new(RwLock::new(crate::core::RequestContext::default())),
};
let web_route = router.route(&web_request).await.unwrap();
assert_eq!(web_route.id, "web");
assert_eq!(web_route.target_base_url, "http://web-service.com");
}
#[tokio::test]
async fn test_default_route() {
let provider = MockConfigProvider { routes: vec![] };
let config = Config::builder().with_provider(provider).build();
let router = PathRouter::new(Arc::new(config)).await.unwrap();
let request = ProxyRequest {
method: crate::core::HttpMethod::Get,
path: "/some/path".to_string(),
query: None,
headers: reqwest::header::HeaderMap::new(),
body: Vec::new(),
context: Arc::new(RwLock::new(crate::core::RequestContext::default())),
};
let route = router.route(&request).await.unwrap();
assert_eq!(route.id, "default");
assert_eq!(route.target_base_url, "http://default-target.com");
}
}