1#[cfg(test)]
12mod tests;
13
14use std::collections::HashMap;
15use std::sync::Arc;
16use std::time::{Duration, Instant};
17use std::{fmt, mem};
18use thiserror::Error;
19use crate::security::{ProviderConfig, SecurityChain, SecurityProvider, SecurityStage};
20use tokio::sync::RwLock;
21use tokio::time::timeout;
22use serde::{Serialize, Deserialize};
23
24use crate::config::Config;
25
26#[derive(Error, Debug)]
28pub enum ProxyError {
29 #[error("HTTP client error: {0}")]
31 ClientError(#[from] reqwest::Error),
32
33 #[error("IO error: {0}")]
35 IoError(#[from] std::io::Error),
36
37 #[error("request timed out after {0:?}")]
39 Timeout(Duration),
40
41 #[error("routing error: {0}")]
43 RoutingError(String),
44
45 #[error("filter error: {0}")]
47 FilterError(String),
48
49 #[error("configuration error: {0}")]
51 ConfigError(String),
52
53 #[error("security error: {0}")]
55 SecurityError(String),
56
57 #[error("{0}")]
59 Other(String),
60}
61
62impl From<crate::config::error::ConfigError> for ProxyError {
63 fn from(err: crate::config::error::ConfigError) -> Self {
64 ProxyError::ConfigError(err.to_string())
65 }
66}
67
68impl From<globset::Error> for ProxyError {
69 fn from(e: globset::Error) -> Self {
70 ProxyError::SecurityError(e.to_string())
71 }
72}
73
74impl From<jsonwebtoken::errors::Error> for ProxyError {
75 fn from(e: jsonwebtoken::errors::Error) -> Self {
76 ProxyError::SecurityError(e.to_string())
77 }
78}
79
80#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
82#[serde(rename_all = "UPPERCASE")]
83pub enum HttpMethod {
84 Get,
85 Post,
86 Put,
87 Delete,
88 Head,
89 Options,
90 Patch,
91 Trace,
92 Connect,
93}
94
95impl fmt::Display for HttpMethod {
96 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
97 match self {
98 HttpMethod::Get => write!(f, "GET"),
99 HttpMethod::Post => write!(f, "POST"),
100 HttpMethod::Put => write!(f, "PUT"),
101 HttpMethod::Delete => write!(f, "DELETE"),
102 HttpMethod::Head => write!(f, "HEAD"),
103 HttpMethod::Options => write!(f, "OPTIONS"),
104 HttpMethod::Patch => write!(f, "PATCH"),
105 HttpMethod::Trace => write!(f, "TRACE"),
106 HttpMethod::Connect => write!(f, "CONNECT"),
107 }
108 }
109}
110
111impl From<&reqwest::Method> for HttpMethod {
112 fn from(method: &reqwest::Method) -> Self {
113 match *method {
114 reqwest::Method::GET => HttpMethod::Get,
115 reqwest::Method::POST => HttpMethod::Post,
116 reqwest::Method::PUT => HttpMethod::Put,
117 reqwest::Method::DELETE => HttpMethod::Delete,
118 reqwest::Method::HEAD => HttpMethod::Head,
119 reqwest::Method::OPTIONS => HttpMethod::Options,
120 reqwest::Method::PATCH => HttpMethod::Patch,
121 reqwest::Method::TRACE => HttpMethod::Trace,
122 reqwest::Method::CONNECT => HttpMethod::Connect,
123 _ => HttpMethod::Get, }
125 }
126}
127
128impl From<HttpMethod> for reqwest::Method {
129 fn from(method: HttpMethod) -> Self {
130 match method {
131 HttpMethod::Get => reqwest::Method::GET,
132 HttpMethod::Post => reqwest::Method::POST,
133 HttpMethod::Put => reqwest::Method::PUT,
134 HttpMethod::Delete => reqwest::Method::DELETE,
135 HttpMethod::Head => reqwest::Method::HEAD,
136 HttpMethod::Options => reqwest::Method::OPTIONS,
137 HttpMethod::Patch => reqwest::Method::PATCH,
138 HttpMethod::Trace => reqwest::Method::TRACE,
139 HttpMethod::Connect => reqwest::Method::CONNECT,
140 }
141 }
142}
143
144#[derive(Debug)]
146pub struct ProxyRequest {
147 pub method: HttpMethod,
148 pub path: String,
149 pub query: Option<String>,
150 pub headers: reqwest::header::HeaderMap,
151 pub body: reqwest::Body,
152 pub context: Arc<RwLock<RequestContext>>,
153}
154
155impl Clone for ProxyRequest {
156 fn clone(&self) -> Self {
157 Self {
159 method: self.method,
160 path: self.path.clone(),
161 query: self.query.clone(),
162 headers: self.headers.clone(),
163 body: reqwest::Body::from(""),
164 context: self.context.clone(),
165 }
166 }
167}
168
169#[derive(Debug)]
171pub struct ProxyResponse {
172 pub status: u16,
173 pub headers: reqwest::header::HeaderMap,
174 pub body: reqwest::Body,
175 pub context: Arc<RwLock<ResponseContext>>,
176}
177
178#[derive(Debug, Default, Clone)]
180pub struct RequestContext {
181 pub client_ip: Option<String>,
183 pub start_time: Option<std::time::Instant>,
185 pub attributes: std::collections::HashMap<String, serde_json::Value>,
187}
188
189#[derive(Debug, Default, Clone)]
191pub struct ResponseContext {
192 pub receive_time: Option<std::time::Instant>,
194 pub attributes: std::collections::HashMap<String, serde_json::Value>,
196}
197
198#[derive(Debug)]
200pub struct ProxyCore {
201 pub config: Arc<Config>,
203 pub client: reqwest::Client,
205 pub router: Arc<dyn Router>,
207 pub global_filters: Arc<RwLock<Vec<Arc<dyn Filter>>>>,
209 pub security_chain: Arc<RwLock<SecurityChain>>,
211}
212
213impl ProxyCore {
214 pub async fn new(config: Arc<Config>, router: Arc<dyn Router>) -> Result<Self, ProxyError> {
216 let timeout_secs: u64 = config.get_or_default("proxy.timeout", 30)?;
218
219 let client = reqwest::Client::builder()
220 .timeout(Duration::from_secs(timeout_secs))
221 .build()
222 .map_err(ProxyError::ClientError)?;
223
224 let security_config = config
225 .get::<Vec<ProviderConfig>>("proxy.security_chain")
226 .unwrap_or_default();
227
228 let security_chain = SecurityChain::from_configs(
229 security_config.unwrap_or_default()
230 ).await?;
231
232 Ok(Self {
233 config,
234 client,
235 router,
236 global_filters: Arc::new(RwLock::new(Vec::new())),
237 security_chain: Arc::new(RwLock::new(security_chain)),
238 })
239 }
240
241 pub async fn add_global_filter(&self, filter: Arc<dyn Filter>) {
243 let mut filters = self.global_filters.write().await;
244 filters.push(filter);
245 }
246
247 pub async fn add_security_provider(&self, p: Arc<dyn SecurityProvider>) {
249 self.security_chain.write().await.add(p);
250 }
251
252 pub async fn process_request(
254 &self,
255 mut request: ProxyRequest,
256 ) -> Result<ProxyResponse, ProxyError> {
257 let overall_start = Instant::now();
258
259 let mut request = self.security_chain.read().await.apply_pre(request).await?;
261
262 for f in self.global_filters.read().await.iter() {
264 if f.filter_type().is_pre() || f.filter_type().is_both() {
265 request = f.pre_filter(request).await?;
266 }
267 }
268 let route = self.router.route(&request).await?;
269 let route_filters = route.filters.clone().unwrap_or_default();
270 for f in &route_filters {
271 if f.filter_type().is_pre() || f.filter_type().is_both() {
272 request = f.pre_filter(request).await?;
273 }
274 }
275
276 let url = format!("{}{}", route.target_base_url, request.path);
278 let outbound_body = mem::replace(&mut request.body, reqwest::Body::from(""));
279
280 let mut builder = self
281 .client
282 .request(request.method.into(), &url)
283 .headers(request.headers.clone())
284 .body(outbound_body);
285
286 if let Some(q) = &request.query {
287 builder = builder.query(&[(q, "")]);
288 }
289
290 let timeout_dur =
292 Duration::from_secs(self.config.get_or_default("proxy.timeout", 30)?);
293
294 let upstream_start = Instant::now();
295 let resp = timeout(timeout_dur, builder.send())
296 .await
297 .map_err(|_| ProxyError::Timeout(timeout_dur))?
298 .map_err(ProxyError::ClientError)?;
299 let upstream_elapsed = upstream_start.elapsed();
300
301 let status = resp.status().as_u16();
303 let headers = resp.headers().clone();
304 let body = reqwest::Body::wrap_stream(resp.bytes_stream());
305
306 let mut proxy_resp = ProxyResponse {
307 status,
308 headers,
309 body,
310 context: Arc::new(RwLock::new(ResponseContext::default())),
311 };
312 proxy_resp.context.write().await.receive_time = Some(Instant::now());
313
314 for f in &route_filters {
316 if f.filter_type().is_post() || f.filter_type().is_both() {
317 proxy_resp = f.post_filter(request.clone(), proxy_resp).await?;
318 }
319 }
320 for f in self.global_filters.read().await.iter() {
321 if f.filter_type().is_post() || f.filter_type().is_both() {
322 proxy_resp = f.post_filter(request.clone(), proxy_resp).await?;
323 }
324 }
325
326 proxy_resp = self.security_chain.read().await.apply_post(request.clone(), proxy_resp).await?;
328
329 let overall_elapsed = overall_start.elapsed();
331 let internal_elapsed = overall_elapsed.saturating_sub(upstream_elapsed);
332
333 log::debug!(
334 "[timing] {} {} -> {} | total={:?} upstream={:?} internal={:?}",
335 request.method,
336 request.path,
337 proxy_resp.status,
338 overall_elapsed,
339 upstream_elapsed,
340 internal_elapsed
341 );
342
343 Ok(proxy_resp)
344 }
345}
346
347#[derive(Debug, Clone, Copy, PartialEq, Eq)]
349pub enum FilterType {
350 Pre,
352 Post,
354 Both,
356}
357
358impl FilterType {
359 pub fn is_pre(&self) -> bool {
361 matches!(self, FilterType::Pre | FilterType::Both)
362 }
363
364 pub fn is_post(&self) -> bool {
366 matches!(self, FilterType::Post | FilterType::Both)
367 }
368
369 pub fn is_both(&self) -> bool {
371 matches!(self, FilterType::Both)
372 }
373}
374
375#[async_trait::async_trait]
377pub trait Filter: fmt::Debug + Send + Sync {
378 fn filter_type(&self) -> FilterType;
380
381 fn name(&self) -> &str;
383
384 async fn pre_filter(&self, request: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
386 Ok(request)
388 }
389
390 async fn post_filter(&self, _request: ProxyRequest, response: ProxyResponse) -> Result<ProxyResponse, ProxyError> {
392 Ok(response)
394 }
395}
396
397#[derive(Debug, Clone)]
399pub struct Route {
400 pub id: String,
402 pub target_base_url: String,
404 pub path_pattern: String,
406 pub filters: Option<Vec<Arc<dyn Filter>>>,
408}
409
410#[async_trait::async_trait]
412pub trait Router: fmt::Debug + Send + Sync {
413 async fn route(&self, request: &ProxyRequest) -> Result<Route, ProxyError>;
415
416 async fn get_routes(&self) -> Vec<Route>;
418
419 async fn add_route(&self, route: Route) -> Result<(), ProxyError>;
421
422 async fn remove_route(&self, route_id: &str) -> Result<(), ProxyError>;
424}