use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use std::fmt;
use thiserror::Error;
use tokio::sync::RwLock;
use tokio::time::timeout;
use serde::{Serialize, Deserialize};
use crate::config::Config;
#[derive(Error, Debug)]
pub enum ProxyError {
#[error("HTTP client error: {0}")]
ClientError(#[from] reqwest::Error),
#[error("IO error: {0}")]
IoError(#[from] std::io::Error),
#[error("request timed out after {0:?}")]
Timeout(Duration),
#[error("routing error: {0}")]
RoutingError(String),
#[error("filter error: {0}")]
FilterError(String),
#[error("configuration error: {0}")]
ConfigError(String),
#[error("{0}")]
Other(String),
}
impl From<crate::config::error::ConfigError> for ProxyError {
fn from(err: crate::config::error::ConfigError) -> Self {
ProxyError::ConfigError(err.to_string())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "UPPERCASE")]
pub enum HttpMethod {
Get,
Post,
Put,
Delete,
Head,
Options,
Patch,
Trace,
Connect,
}
impl fmt::Display for HttpMethod {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
HttpMethod::Get => write!(f, "GET"),
HttpMethod::Post => write!(f, "POST"),
HttpMethod::Put => write!(f, "PUT"),
HttpMethod::Delete => write!(f, "DELETE"),
HttpMethod::Head => write!(f, "HEAD"),
HttpMethod::Options => write!(f, "OPTIONS"),
HttpMethod::Patch => write!(f, "PATCH"),
HttpMethod::Trace => write!(f, "TRACE"),
HttpMethod::Connect => write!(f, "CONNECT"),
}
}
}
impl From<&reqwest::Method> for HttpMethod {
fn from(method: &reqwest::Method) -> Self {
match *method {
reqwest::Method::GET => HttpMethod::Get,
reqwest::Method::POST => HttpMethod::Post,
reqwest::Method::PUT => HttpMethod::Put,
reqwest::Method::DELETE => HttpMethod::Delete,
reqwest::Method::HEAD => HttpMethod::Head,
reqwest::Method::OPTIONS => HttpMethod::Options,
reqwest::Method::PATCH => HttpMethod::Patch,
reqwest::Method::TRACE => HttpMethod::Trace,
reqwest::Method::CONNECT => HttpMethod::Connect,
_ => HttpMethod::Get, }
}
}
impl From<HttpMethod> for reqwest::Method {
fn from(method: HttpMethod) -> Self {
match method {
HttpMethod::Get => reqwest::Method::GET,
HttpMethod::Post => reqwest::Method::POST,
HttpMethod::Put => reqwest::Method::PUT,
HttpMethod::Delete => reqwest::Method::DELETE,
HttpMethod::Head => reqwest::Method::HEAD,
HttpMethod::Options => reqwest::Method::OPTIONS,
HttpMethod::Patch => reqwest::Method::PATCH,
HttpMethod::Trace => reqwest::Method::TRACE,
HttpMethod::Connect => reqwest::Method::CONNECT,
}
}
}
#[derive(Debug, Clone)]
pub struct ProxyRequest {
pub method: HttpMethod,
pub path: String,
pub query: Option<String>,
pub headers: reqwest::header::HeaderMap,
pub body: Vec<u8>,
pub context: Arc<RwLock<RequestContext>>,
}
#[derive(Debug, Clone)]
pub struct ProxyResponse {
pub status: u16,
pub headers: reqwest::header::HeaderMap,
pub body: Vec<u8>,
pub context: Arc<RwLock<ResponseContext>>,
}
#[derive(Debug, Default, Clone)]
pub struct RequestContext {
pub client_ip: Option<String>,
pub start_time: Option<std::time::Instant>,
pub attributes: std::collections::HashMap<String, serde_json::Value>,
}
#[derive(Debug, Default, Clone)]
pub struct ResponseContext {
pub receive_time: Option<std::time::Instant>,
pub attributes: std::collections::HashMap<String, serde_json::Value>,
}
#[derive(Debug)]
pub struct ProxyCore {
pub config: Arc<Config>,
pub client: reqwest::Client,
pub router: Arc<dyn Router>,
pub global_filters: Arc<RwLock<Vec<Arc<dyn Filter>>>>,
pub route_filters: Arc<RwLock<HashMap<String, Vec<Arc<dyn Filter>>>>>,
}
impl ProxyCore {
pub fn new(config: Arc<Config>, router: Arc<dyn Router>) -> Result<Self, ProxyError> {
let timeout_secs: u64 = config.get_or_default("proxy.timeout", 30)?;
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.build()
.map_err(ProxyError::ClientError)?;
Ok(Self {
config,
client,
router,
global_filters: Arc::new(RwLock::new(Vec::new())),
route_filters: Arc::new(RwLock::new(HashMap::new())),
})
}
pub async fn add_global_filter(&self, filter: Arc<dyn Filter>) {
let mut filters = self.global_filters.write().await;
filters.push(filter);
}
pub async fn add_route_filter(&self, route_id: &str, filter: Arc<dyn Filter>) {
let mut filters = self.route_filters.write().await;
let route_filters = filters.entry(route_id.to_string()).or_insert_with(Vec::new);
route_filters.push(filter);
}
pub async fn process_request(&self, mut request: ProxyRequest) -> Result<ProxyResponse, ProxyError> {
for filter in self.global_filters.read().await.iter() {
if filter.filter_type().is_pre() || filter.filter_type().is_both() {
request = filter.pre_filter(request).await?;
}
}
let route = self.router.route(&request).await?;
if let Some(route_filters) = self.route_filters.read().await.get(&route.id) {
for filter in route_filters.iter() {
if filter.filter_type().is_pre() || filter.filter_type().is_both() {
request = filter.pre_filter(request).await?;
}
}
}
let target_url = url::Url::parse(&route.target_base_url)
.map_err(|e| ProxyError::RoutingError(format!("Invalid target URL: {}", e)))?;
let forwarded_path = if target_url.path() == "/" {
request.path.clone()
} else {
request.path.clone()
};
let mut new_url = target_url.clone();
new_url.set_path(&forwarded_path);
let mut builder = self.client.request(
request.method.into(),
new_url.as_str()
);
if let Some(query) = &request.query {
builder = builder.query(&[(query, "")]);
}
let request_clone = request.clone();
builder = builder.headers(request_clone.headers);
if !request.body.is_empty() {
builder = builder.body(request_clone.body);
}
let timeout_secs: u64 = self.config.get_or_default("proxy.timeout", 30)?;
let timeout_duration = Duration::from_secs(timeout_secs);
let resp = match timeout(timeout_duration, builder.send()).await {
Ok(Ok(resp)) => resp,
Ok(Err(e)) => return Err(ProxyError::ClientError(e)),
Err(_) => return Err(ProxyError::Timeout(timeout_duration)),
};
let status = resp.status().as_u16();
let headers = resp.headers().clone();
let body = resp.bytes().await.map_err(ProxyError::ClientError)?.to_vec();
let mut response = ProxyResponse {
status,
headers,
body,
context: Arc::new(RwLock::new(ResponseContext::default())),
};
let mut context = response.context.write().await;
context.receive_time = Some(std::time::Instant::now());
drop(context);
if let Some(route_filters) = self.route_filters.read().await.get(&route.id) {
for filter in route_filters.iter() {
if filter.filter_type().is_post() || filter.filter_type().is_both() {
response = filter.post_filter(request.clone(), response).await?;
}
}
}
for filter in self.global_filters.read().await.iter() {
if filter.filter_type().is_post() || filter.filter_type().is_both() {
response = filter.post_filter(request.clone(), response).await?;
}
}
Ok(response)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FilterType {
Pre,
Post,
Both,
}
impl FilterType {
pub fn is_pre(&self) -> bool {
matches!(self, FilterType::Pre | FilterType::Both)
}
pub fn is_post(&self) -> bool {
matches!(self, FilterType::Post | FilterType::Both)
}
pub fn is_both(&self) -> bool {
matches!(self, FilterType::Both)
}
}
#[async_trait::async_trait]
pub trait Filter: fmt::Debug + Send + Sync {
fn filter_type(&self) -> FilterType;
fn name(&self) -> &str;
async fn pre_filter(&self, request: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
Ok(request)
}
async fn post_filter(&self, _request: ProxyRequest, response: ProxyResponse) -> Result<ProxyResponse, ProxyError> {
Ok(response)
}
}
#[derive(Debug, Clone)]
pub struct Route {
pub id: String,
pub target_base_url: String,
pub path_pattern: String,
pub filter_ids: Vec<String>,
}
#[async_trait::async_trait]
pub trait Router: fmt::Debug + Send + Sync {
async fn route(&self, request: &ProxyRequest) -> Result<Route, ProxyError>;
async fn get_routes(&self) -> Vec<Route>;
async fn add_route(&self, route: Route) -> Result<(), ProxyError>;
async fn remove_route(&self, route_id: &str) -> Result<(), ProxyError>;
}