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 #[cfg(feature = "opentelemetry")]
249 pub async fn add_opentelemetry_filter(&self, config: &crate::opentelemetry::OpenTelemetryConfig) -> Result<(), ProxyError> {
250 let filter = Arc::new(crate::opentelemetry::OpenTelemetryFilter::new(config.clone()));
251 self.add_global_filter(filter).await;
252 Ok(())
253 }
254
255 pub async fn add_security_provider(&self, p: Arc<dyn SecurityProvider>) {
257 self.security_chain.write().await.add(p);
258 }
259
260 pub async fn process_request(
262 &self,
263 request: ProxyRequest,
264 ) -> Result<ProxyResponse, ProxyError> {
265 let overall_start = Instant::now();
266 let method = request.method.to_string();
267 let path = request.path.clone();
268
269 log::trace!("Processing request: {} {}", method, path);
270
271 let mut request = match self.security_chain.read().await.apply_pre(request).await {
273 Ok(req) => {
274 log::trace!("Security pre-auth passed for {} {}", method, path);
275 req
276 },
277 Err(e) => {
278 log::warn!("Security pre-auth failed for {} {}: {}", method, path, e);
279 return Err(e);
280 }
281 };
282
283 for f in self.global_filters.read().await.iter() {
285 if f.filter_type().is_pre() || f.filter_type().is_both() {
286 log::trace!("Applying global pre-filter: {}", f.name());
287 match f.pre_filter(request).await {
288 Ok(req) => request = req,
289 Err(e) => {
290 log::error!("Global pre-filter '{}' failed: {}", f.name(), e);
291 return Err(e);
292 }
293 }
294 }
295 }
296
297 let route = match self.router.route(&request).await {
298 Ok(r) => {
299 log::debug!("Request {} {} matched route: {}", method, path, r.id);
300 r
301 },
302 Err(e) => {
303 log::warn!("No route found for {} {}: {}", method, path, e);
304 return Err(e);
305 }
306 };
307
308 let route_filters = route.filters.clone().unwrap_or_default();
309 for f in &route_filters {
310 if f.filter_type().is_pre() || f.filter_type().is_both() {
311 log::trace!("Applying route pre-filter: {}", f.name());
312 match f.pre_filter(request).await {
313 Ok(req) => request = req,
314 Err(e) => {
315 log::error!("Route pre-filter '{}' failed: {}", f.name(), e);
316 return Err(e);
317 }
318 }
319 }
320 }
321
322 let url = format!("{}{}", route.target_base_url, request.path);
324 log::debug!("Forwarding to target: {}", url);
325 let outbound_body = mem::replace(&mut request.body, reqwest::Body::from(""));
326
327 let mut builder = self
328 .client
329 .request(request.method.into(), &url)
330 .headers(request.headers.clone())
331 .body(outbound_body);
332
333 if let Some(q) = &request.query {
334 builder = builder.query(&[(q, "")]);
335 }
336
337 let timeout_dur =
339 Duration::from_secs(self.config.get_or_default("proxy.timeout", 30).unwrap_or_else(|e| {
340 log::error!("Failed to get timeout config: {}", e);
341 30 }));
343
344 let upstream_start = Instant::now();
345 log::trace!("Sending request to upstream with timeout: {:?}", timeout_dur);
346
347 let resp = match timeout(timeout_dur, builder.send()).await {
348 Ok(result) => match result {
349 Ok(response) => response,
350 Err(e) => {
351 log::error!("Upstream request failed: {}", e);
352 return Err(ProxyError::ClientError(e));
353 }
354 },
355 Err(_) => {
356 log::warn!("Request to {} timed out after {:?}", url, timeout_dur);
357 return Err(ProxyError::Timeout(timeout_dur));
358 }
359 };
360
361 let upstream_elapsed = upstream_start.elapsed();
362 log::trace!("Received response from upstream in {:?}", upstream_elapsed);
363
364 let status = resp.status().as_u16();
366 let headers = resp.headers().clone();
367 let body = reqwest::Body::wrap_stream(resp.bytes_stream());
368
369 let mut proxy_resp = ProxyResponse {
370 status,
371 headers,
372 body,
373 context: Arc::new(RwLock::new(ResponseContext::default())),
374 };
375 proxy_resp.context.write().await.receive_time = Some(Instant::now());
376
377 log::debug!("Upstream responded with status: {}", status);
378
379 for f in &route_filters {
381 if f.filter_type().is_post() || f.filter_type().is_both() {
382 log::trace!("Applying route post-filter: {}", f.name());
383 match f.post_filter(request.clone(), proxy_resp).await {
384 Ok(resp) => proxy_resp = resp,
385 Err(e) => {
386 log::error!("Route post-filter '{}' failed: {}", f.name(), e);
387 return Err(e);
388 }
389 }
390 }
391 }
392
393 for f in self.global_filters.read().await.iter() {
394 if f.filter_type().is_post() || f.filter_type().is_both() {
395 log::trace!("Applying global post-filter: {}", f.name());
396 match f.post_filter(request.clone(), proxy_resp).await {
397 Ok(resp) => proxy_resp = resp,
398 Err(e) => {
399 log::error!("Global post-filter '{}' failed: {}", f.name(), e);
400 return Err(e);
401 }
402 }
403 }
404 }
405
406 proxy_resp = match self.security_chain.read().await.apply_post(request.clone(), proxy_resp).await {
408 Ok(resp) => {
409 log::trace!("Security post-auth passed for {} {}", method, path);
410 resp
411 },
412 Err(e) => {
413 log::warn!("Security post-auth failed for {} {}: {}", method, path, e);
414 return Err(e);
415 }
416 };
417
418 let overall_elapsed = overall_start.elapsed();
420 let internal_elapsed = overall_elapsed.saturating_sub(upstream_elapsed);
421
422 log::debug!(
423 "[timing] {} {} -> {} | total={:?} upstream={:?} internal={:?}",
424 request.method,
425 request.path,
426 proxy_resp.status,
427 overall_elapsed,
428 upstream_elapsed,
429 internal_elapsed
430 );
431
432 Ok(proxy_resp)
433 }
434}
435
436#[derive(Debug, Clone, Copy, PartialEq, Eq)]
438pub enum FilterType {
439 Pre,
441 Post,
443 Both,
445}
446
447impl FilterType {
448 pub fn is_pre(&self) -> bool {
450 matches!(self, FilterType::Pre | FilterType::Both)
451 }
452
453 pub fn is_post(&self) -> bool {
455 matches!(self, FilterType::Post | FilterType::Both)
456 }
457
458 pub fn is_both(&self) -> bool {
460 matches!(self, FilterType::Both)
461 }
462}
463
464#[async_trait::async_trait]
466pub trait Filter: fmt::Debug + Send + Sync {
467 fn filter_type(&self) -> FilterType;
469
470 fn name(&self) -> &str;
472
473 async fn pre_filter(&self, request: ProxyRequest) -> Result<ProxyRequest, ProxyError> {
475 Ok(request)
477 }
478
479 async fn post_filter(&self, _request: ProxyRequest, response: ProxyResponse) -> Result<ProxyResponse, ProxyError> {
481 Ok(response)
483 }
484}
485
486#[derive(Debug, Clone)]
488pub struct Route {
489 pub id: String,
491 pub target_base_url: String,
493 pub path_pattern: String,
495 pub filters: Option<Vec<Arc<dyn Filter>>>,
497}
498
499#[async_trait::async_trait]
501pub trait Router: fmt::Debug + Send + Sync {
502 async fn route(&self, request: &ProxyRequest) -> Result<Route, ProxyError>;
504
505 async fn get_routes(&self) -> Vec<Route>;
507
508 async fn add_route(&self, route: Route) -> Result<(), ProxyError>;
510
511 async fn remove_route(&self, route_id: &str) -> Result<(), ProxyError>;
513}