1pub mod config;
5pub mod metrics;
6pub mod tls;
7
8use std::convert::Infallible;
9use std::net::SocketAddr;
10use std::path::Path;
11use std::pin::Pin;
12use std::sync::atomic::{AtomicU64, Ordering};
13use std::sync::{Arc, RwLock};
14use std::task::{Context, Poll};
15use std::time::{Instant, SystemTime};
16
17use http_body_util::combinators::BoxBody;
18use http_body_util::{BodyExt, Full};
19use hyper::body::{Body, Bytes, Frame, Incoming};
20use hyper::service::service_fn;
21use hyper::{HeaderMap, Request, Response, Uri};
22use hyper_util::client::legacy::connect::HttpConnector;
23use hyper_util::client::legacy::Client;
24use hyper_util::rt::{TokioExecutor, TokioIo};
25use hyper_util::server::conn::auto;
26use tokio::net::TcpListener;
27use tokio_rustls::TlsAcceptor;
28
29use crate::metrics::{Metrics, Outcome};
30use tls::TlsCertSource;
31use tracing::{debug, error, info, warn};
32
33use waf_core::{
34 ClientIpResolver, Config, FailMode, IpSource, Normalized, RateLimitState, RequestContext,
35 ResilienceConfig, StateStore, WafModule,
36};
37use waf_detection::{
38 crs::CrsModule,
39 evasion::EvasionModule,
40 graphql::GraphqlModule, grpc::GrpcModule, header_injection::HeaderInjectionModule, ldap::LdapModule,
41 lfi_rfi::LfiRfiModule,
42 mail::MailModule, nosql::NosqlModule, path_traversal::PathTraversalModule,
43 rate_limit::RateLimitModule,
44 rce::RceModule, request_smuggling::RequestSmugglingModule, scanner::ScannerModule,
45 sqli::SqliModule, ssi::SsiModule, ssrf::SsrfModule, ssti::SstiModule, xss::XssModule,
46 xxe::XxeModule, ContentPrefilter,
47};
48use waf_normalizer::Normalizer;
49use waf_pipeline::{NoopLogger, Pipeline, PipelineVerdict};
50use waf_wasm::{WasmModule, WasmOptions};
51
52pub type HyperBoxBody = BoxBody<Bytes, hyper::Error>;
53
54pub type ModuleFactory =
61 dyn Fn() -> Result<Vec<Box<dyn WafModule>>, Box<dyn std::error::Error + Send + Sync>>
62 + Send
63 + Sync;
64
65const HOP_BY_HOP: &[&str] = &[
67 "connection",
68 "host", "keep-alive",
70 "proxy-authenticate",
71 "proxy-authorization",
72 "te",
73 "trailers",
74 "transfer-encoding",
75 "upgrade",
76];
77
78static REQUEST_COUNTER: AtomicU64 = AtomicU64::new(0);
79
80fn next_request_id() -> String {
81 let n = REQUEST_COUNTER.fetch_add(1, Ordering::Relaxed);
82 format!("req-{n:016x}")
83}
84
85pub fn full_body(data: impl Into<Bytes>) -> HyperBoxBody {
86 Full::new(data.into())
87 .map_err(|never| match never {})
88 .boxed()
89}
90
91struct FramedBody {
96 data: Option<Bytes>,
97 trailers: Option<HeaderMap>,
98}
99
100impl Body for FramedBody {
101 type Data = Bytes;
102 type Error = Infallible;
103
104 fn poll_frame(
105 mut self: Pin<&mut Self>,
106 _cx: &mut Context<'_>,
107 ) -> Poll<Option<Result<Frame<Bytes>, Infallible>>> {
108 if let Some(d) = self.data.take() {
109 return Poll::Ready(Some(Ok(Frame::data(d))));
110 }
111 if let Some(t) = self.trailers.take() {
112 return Poll::Ready(Some(Ok(Frame::trailers(t))));
113 }
114 Poll::Ready(None)
115 }
116}
117
118fn body_with_trailers(data: Bytes, trailers: Option<HeaderMap>) -> HyperBoxBody {
121 match trailers {
122 None => full_body(data),
123 Some(t) => FramedBody { data: Some(data), trailers: Some(t) }
124 .map_err(|never| match never {})
125 .boxed(),
126 }
127}
128
129async fn collect_with_trailers<B>(body: B) -> Result<(Bytes, Option<HeaderMap>), B::Error>
133where
134 B: Body<Data = Bytes>,
135{
136 let collected = body.collect().await?;
137 let trailers = collected.trailers().cloned();
138 Ok((collected.to_bytes(), trailers))
139}
140
141fn is_grpc_request(parts: &hyper::http::request::Parts) -> bool {
145 parts
146 .headers
147 .get(hyper::header::CONTENT_TYPE)
148 .and_then(|v| v.to_str().ok())
149 .map(|ct| ct.trim_start().starts_with("application/grpc"))
150 .unwrap_or(false)
151}
152
153fn parse_cookies(headers: &[(String, String)]) -> Vec<(String, String)> {
154 headers
155 .iter()
156 .filter(|(name, _)| name.eq_ignore_ascii_case("cookie"))
157 .flat_map(|(_, value)| {
158 value.split(';').filter_map(|pair| {
159 let mut parts = pair.splitn(2, '=');
160 let key = parts.next()?.trim().to_string();
161 let val = parts.next().unwrap_or("").trim().to_string();
162 Some((key, val))
163 })
164 })
165 .collect()
166}
167
168fn build_context(
169 parts: &hyper::http::request::Parts,
170 body: &Bytes,
171 client_addr: SocketAddr,
172 ip_resolver: &ClientIpResolver,
173) -> RequestContext {
174 let path = parts.uri.path().to_string();
175 let query = parts.uri.query().map(str::to_string);
176 let method = parts.method.to_string();
177 let http_version = format!("{:?}", parts.version);
178
179 let headers: Vec<(String, String)> = parts
180 .headers
181 .iter()
182 .filter_map(|(name, value)| {
183 value.to_str().ok().map(|v| (name.to_string(), v.to_string()))
184 })
185 .collect();
186
187 let cookies = parse_cookies(&headers);
188
189 let normalized = Normalized::default();
190
191 let request_id = next_request_id();
196 let resolved = ip_resolver.resolve(client_addr.ip(), &headers);
197 match resolved.source {
198 IpSource::FallbackMissingHeader | IpSource::FallbackMalformed => warn!(
199 request_id = %request_id,
200 peer = %client_addr.ip(),
201 source = ?resolved.source,
202 "client-IP resolution fell back to peer address"
203 ),
204 IpSource::DirectPeer | IpSource::TrustedHeader => {}
205 }
206
207 RequestContext {
208 client_ip: resolved.ip,
209 request_id,
210 timestamp: SystemTime::now(),
211 method,
212 path: path.clone(),
213 raw_path: path,
214 query,
215 http_version,
216 headers,
217 cookies,
218 body: body.clone(),
219 normalized,
220 score: 0,
221 score_contributions: vec![],
222 }
223}
224
225struct Reloadable {
229 backend: String,
230 normalizer: Normalizer,
231 pipeline: Pipeline,
232 prefilter: ContentPrefilter,
236 ip_resolver: ClientIpResolver,
237 resilience: ResilienceConfig,
238}
239
240struct StaticState {
247 client: Client<HttpConnector, HyperBoxBody>,
248 grpc_client: Client<HttpConnector, HyperBoxBody>,
252 listen_addr: SocketAddr,
253 rl_state: RateLimitState,
254 current: RwLock<Arc<Reloadable>>,
255 mode: HandlerMode,
256 tls_acceptor: Option<TlsAcceptor>,
260 metrics: Arc<Metrics>,
263 module_factory: Option<Arc<ModuleFactory>>,
268}
269
270#[derive(Clone, Copy)]
277enum HandlerMode {
278 Inspect,
279 Passthrough,
280}
281
282impl StaticState {
283 fn current(&self) -> Arc<Reloadable> {
288 self.current
289 .read()
290 .unwrap_or_else(|poisoned| poisoned.into_inner())
291 .clone()
292 }
293}
294
295#[derive(Clone)]
299pub struct Reloader(Arc<StaticState>);
300
301impl Reloader {
302 pub fn reload_from(&self, path: &Path) -> Result<(), config::LoadError> {
306 let new_cfg = match config::load(path) {
307 Ok(c) => c,
308 Err(e) => {
309 error!(error = %e, "config reload failed; keeping current configuration");
310 return Err(e);
311 }
312 };
313
314 if new_cfg.proxy.listen != self.0.listen_addr {
316 warn!(
317 current = %self.0.listen_addr,
318 requested = %new_cfg.proxy.listen,
319 "proxy.listen change requires a restart; keeping the current bind address"
320 );
321 }
322
323 let extra = match &self.0.module_factory {
330 Some(factory) => match factory() {
331 Ok(modules) => modules,
332 Err(e) => {
333 error!(error = %e, "module factory failed on reload; keeping current configuration");
334 return Err(config::LoadError::ModuleFactory(e.to_string()));
335 }
336 },
337 None => Vec::new(),
338 };
339
340 let new_reloadable = build_reloadable(&new_cfg, self.0.rl_state.clone(), extra);
343
344 *self
348 .0
349 .current
350 .write()
351 .unwrap_or_else(|poisoned| poisoned.into_inner()) = Arc::new(new_reloadable);
352 info!("configuration reloaded");
353 Ok(())
354 }
355}
356
357fn upstream_error_response(
362 ctx: &RequestContext,
363 resilience: &ResilienceConfig,
364 detail: &str,
365) -> Response<HyperBoxBody> {
366 let (status, body) = match resilience.on_upstream_error {
367 FailMode::FailClosed => (502, "Bad Gateway"),
368 FailMode::FailOpen => (503, "Service Unavailable"),
369 };
370 warn!(
371 request_id = %ctx.request_id,
372 client_ip = %ctx.client_ip,
373 status = status,
374 policy = ?resilience.on_upstream_error,
375 detail = detail,
376 "upstream error: applying on_upstream_error policy"
377 );
378 Response::builder().status(status).body(full_body(body)).unwrap()
379}
380
381fn contributions_json(ctx: &RequestContext) -> String {
387 serde_json::to_string(&ctx.score_contributions).unwrap_or_else(|_| "[]".to_string())
388}
389
390fn deny_response(
393 ctx: &RequestContext,
394 verdict: PipelineVerdict,
395) -> Option<(Response<HyperBoxBody>, Outcome)> {
396 match verdict {
397 PipelineVerdict::Allow => None,
398 PipelineVerdict::Block { rule_id, reason } => {
399 warn!(
400 request_id = %ctx.request_id,
401 rule_id = %rule_id,
402 reason = %reason,
403 score = ctx.score,
404 score_contributions = %contributions_json(ctx),
405 "request blocked"
406 );
407 Some((
408 Response::builder()
409 .status(403)
410 .body(full_body("Forbidden"))
411 .unwrap(),
412 Outcome::Blocked,
413 ))
414 }
415 PipelineVerdict::Reject { rule_id, reason, status, retry_after } => {
416 warn!(
417 request_id = %ctx.request_id,
418 rule_id = %rule_id,
419 reason = %reason,
420 status = status,
421 score_contributions = %contributions_json(ctx),
422 "request rejected"
423 );
424 let (body, outcome) = match status {
427 429 => ("Too Many Requests", Outcome::RateLimited),
428 400 => ("Bad Request", Outcome::BadRequest),
429 _ => ("Rejected", Outcome::BadRequest),
430 };
431 let mut builder = Response::builder().status(status);
432 if let Some(secs) = retry_after {
433 builder = builder.header("retry-after", secs.to_string());
434 }
435 Some((builder.body(full_body(body)).unwrap(), outcome))
436 }
437 }
438}
439
440async fn try_forward(
441 req: Request<Incoming>,
442 state: &StaticState,
443 client_addr: SocketAddr,
444) -> Result<(Response<HyperBoxBody>, Outcome), Box<dyn std::error::Error + Send + Sync>> {
445 let rel = state.current();
448
449 let (parts, body) = req.into_parts();
450 let (body_bytes, req_trailers) = collect_with_trailers(body).await?;
453
454 let mut ctx = build_context(&parts, &body_bytes, client_addr, &rel.ip_resolver);
455
456 let connection_verdict = rel.pipeline.run_connection(&mut ctx);
459 if let Some(denied) = deny_response(&ctx, connection_verdict) {
460 return Ok(denied);
461 }
462
463 let normalized_ok = match rel.normalizer.normalize(&mut ctx) {
467 Ok(()) => true,
468 Err(e) => match rel.resilience.on_parser_limit {
469 FailMode::FailClosed => {
470 warn!(
471 request_id = %ctx.request_id,
472 error = %e,
473 policy = ?FailMode::FailClosed,
474 "normalization failed: rejecting (on_parser_limit)"
475 );
476 return Ok((
477 Response::builder()
478 .status(400)
479 .body(full_body("Bad Request"))
480 .unwrap(),
481 Outcome::BadRequest,
482 ));
483 }
484 FailMode::FailOpen => {
485 warn!(
486 request_id = %ctx.request_id,
487 error = %e,
488 policy = ?FailMode::FailOpen,
489 "normalization failed: forwarding UNINSPECTED (on_parser_limit)"
490 );
491 false
492 }
493 },
494 };
495
496 let path_and_query = parts
497 .uri
498 .path_and_query()
499 .map(|pq| pq.as_str())
500 .unwrap_or("/")
501 .to_string();
502
503 info!(
504 request_id = %ctx.request_id,
505 method = %ctx.method,
506 path = %path_and_query,
507 client_ip = %ctx.client_ip,
508 "→ request"
509 );
510
511 if normalized_ok {
514 let inspect = rel.prefilter.is_candidate(&ctx);
520 let inspection_verdict = rel.pipeline.run_inspection_gated(&mut ctx, inspect);
521 if let Some(denied) = deny_response(&ctx, inspection_verdict) {
522 return Ok(denied);
523 }
524 }
525
526 forward_to_backend(state, &rel, &parts, &path_and_query, body_bytes, req_trailers, client_addr, &ctx).await
527}
528
529#[allow(clippy::too_many_arguments)]
538async fn forward_to_backend(
539 state: &StaticState,
540 rel: &Reloadable,
541 parts: &hyper::http::request::Parts,
542 path_and_query: &str,
543 body_bytes: Bytes,
544 req_trailers: Option<HeaderMap>,
545 client_addr: SocketAddr,
546 ctx: &RequestContext,
547) -> Result<(Response<HyperBoxBody>, Outcome), Box<dyn std::error::Error + Send + Sync>> {
548 let backend_uri: Uri = format!("{}{}", rel.backend, path_and_query).parse()?;
549 let is_grpc = is_grpc_request(parts);
550
551 let mut builder = Request::builder()
552 .method(parts.method.clone())
553 .uri(backend_uri);
554
555 for (name, value) in &parts.headers {
556 if !HOP_BY_HOP.contains(&name.as_str()) {
557 builder = builder.header(name, value);
558 }
559 }
560 builder = builder.header("x-forwarded-for", client_addr.ip().to_string());
563 builder = builder.header("x-request-id", ctx.request_id.as_str());
564 if is_grpc {
567 builder = builder.header("te", "trailers");
568 }
569
570 let (client, fwd_body) = if is_grpc {
573 (&state.grpc_client, body_with_trailers(body_bytes, req_trailers))
574 } else {
575 (&state.client, full_body(body_bytes))
576 };
577 let fwd_req = builder.body(fwd_body)?;
578
579 let upstream = tokio::time::timeout(rel.resilience.upstream_timeout(), async {
583 let resp = client.request(fwd_req).await?;
584 let (resp_parts, resp_body) = resp.into_parts();
585 let (resp_bytes, resp_trailers) = collect_with_trailers(resp_body).await?;
588 Ok::<_, Box<dyn std::error::Error + Send + Sync>>((resp_parts, resp_bytes, resp_trailers))
589 })
590 .await;
591
592 let (resp_parts, resp_bytes, resp_trailers) = match upstream {
593 Ok(Ok(triple)) => triple,
594 Ok(Err(e)) => {
595 return Ok((
596 upstream_error_response(ctx, &rel.resilience, &e.to_string()),
597 Outcome::UpstreamError,
598 ))
599 }
600 Err(_elapsed) => {
601 return Ok((
602 upstream_error_response(ctx, &rel.resilience, "upstream timeout"),
603 Outcome::UpstreamError,
604 ))
605 }
606 };
607
608 info!(
609 request_id = %ctx.request_id,
610 status = %resp_parts.status,
611 score = ctx.score,
612 "← response"
613 );
614
615 Ok((
616 Response::from_parts(resp_parts, body_with_trailers(resp_bytes, resp_trailers)),
617 Outcome::Allowed,
618 ))
619}
620
621async fn try_passthrough(
628 req: Request<Incoming>,
629 state: &StaticState,
630 client_addr: SocketAddr,
631) -> Result<(Response<HyperBoxBody>, Outcome), Box<dyn std::error::Error + Send + Sync>> {
632 let rel = state.current();
633 let (parts, body) = req.into_parts();
634 let (body_bytes, req_trailers) = collect_with_trailers(body).await?;
635 let ctx = build_context(&parts, &body_bytes, client_addr, &rel.ip_resolver);
636 let path_and_query = parts
637 .uri
638 .path_and_query()
639 .map(|pq| pq.as_str())
640 .unwrap_or("/")
641 .to_string();
642 forward_to_backend(state, &rel, &parts, &path_and_query, body_bytes, req_trailers, client_addr, &ctx).await
643}
644
645async fn handle(
646 req: Request<Incoming>,
647 state: Arc<StaticState>,
648 client_addr: SocketAddr,
649) -> Result<Response<HyperBoxBody>, Infallible> {
650 let start = Instant::now();
653 let result = match state.mode {
654 HandlerMode::Inspect => try_forward(req, &state, client_addr).await,
655 HandlerMode::Passthrough => try_passthrough(req, &state, client_addr).await,
656 };
657 let (resp, outcome) = match result {
661 Ok((resp, outcome)) => (resp, outcome),
662 Err(e) => {
663 error!(error = %e, client_ip = %client_addr.ip(), "forwarding error");
664 let resp = Response::builder()
665 .status(502)
666 .body(full_body("Bad Gateway"))
667 .unwrap();
668 (resp, Outcome::InternalError)
669 }
670 };
671 state.metrics.record(outcome, start.elapsed());
672 Ok(resp)
673}
674
675pub struct Proxy {
676 listener: TcpListener,
677 state: Arc<StaticState>,
678 metrics_listener: Option<TcpListener>,
682}
683
684fn build_modules(config: &Config, rl_state: &RateLimitState) -> Vec<Box<dyn WafModule>> {
687 let mut modules: Vec<Box<dyn WafModule>> = vec![Box::new(NoopLogger)];
688 if config.modules.request_smuggling.enabled {
691 modules.push(Box::new(RequestSmugglingModule::new()));
692 }
693 if config.rate_limit.enabled {
694 modules.push(Box::new(RateLimitModule::with_state(rl_state.clone())));
695 }
696 if config.modules.sqli.enabled {
697 modules.push(Box::new(SqliModule::new()));
698 }
699 if config.modules.xss.enabled {
700 modules.push(Box::new(XssModule::new()));
701 }
702 if config.modules.path_traversal.enabled {
703 modules.push(Box::new(PathTraversalModule::new()));
704 }
705 if config.modules.rce.enabled {
706 modules.push(Box::new(RceModule::new()));
707 }
708 if config.modules.lfi_rfi.enabled {
709 modules.push(Box::new(LfiRfiModule::new()));
710 }
711 if config.modules.ssrf.enabled {
712 modules.push(Box::new(SsrfModule::new()));
713 }
714 if config.modules.ldap.enabled {
715 modules.push(Box::new(LdapModule::new()));
716 }
717 if config.modules.nosql.enabled {
718 modules.push(Box::new(NosqlModule::new()));
719 }
720 if config.modules.mail.enabled {
721 modules.push(Box::new(MailModule::new()));
722 }
723 if config.modules.ssti.enabled {
724 modules.push(Box::new(SstiModule::new()));
725 }
726 if config.modules.scanner.enabled {
727 modules.push(Box::new(ScannerModule::new()));
728 }
729 if config.modules.ssi.enabled {
730 modules.push(Box::new(SsiModule::new()));
731 }
732 if config.modules.xxe.enabled {
733 modules.push(Box::new(XxeModule::new()));
734 }
735 if config.modules.header_injection.enabled {
736 modules.push(Box::new(HeaderInjectionModule::new()));
737 }
738 if config.modules.evasion.enabled {
739 modules.push(Box::new(EvasionModule::new()));
740 }
741 if config.modules.graphql.enabled {
742 modules.push(Box::new(GraphqlModule::new()));
743 }
744 if config.modules.grpc.enabled {
745 modules.push(Box::new(GrpcModule::new()));
746 }
747 if config.modules.crs.enabled {
748 modules.push(Box::new(load_crs_module(&config.modules.crs.files)));
749 }
750 if config.modules.wasm.enabled {
751 for plugin in &config.modules.wasm.plugins {
752 if let Some(m) = load_wasm_plugin(plugin, &config.modules.wasm) {
753 modules.push(Box::new(m));
754 }
755 }
756 }
757 modules
758}
759
760fn load_wasm_plugin(
766 plugin: &waf_core::WasmPluginConfig,
767 cfg: &waf_core::WasmConfig,
768) -> Option<WasmModule> {
769 let bytes = match std::fs::read(&plugin.path) {
770 Ok(b) => b,
771 Err(e) => {
772 error!(file = %plugin.path, error = %e, "WASM: cannot read plugin (skipped)");
773 return None;
774 }
775 };
776 let name = plugin_name(&plugin.path);
777 let opts = WasmOptions {
778 pool_size: cfg.pool_size,
779 fuel_per_request: cfg.fuel_per_request,
780 max_memory_bytes: cfg.max_memory_bytes,
781 checkout_timeout: std::time::Duration::from_millis(cfg.checkout_timeout_ms),
782 };
783 let config_bytes = plugin.config.as_deref().unwrap_or("").as_bytes();
784 match WasmModule::from_bytes(&name, &bytes, config_bytes, &opts) {
785 Ok((module, report)) => {
786 info!(plugin = %name, "{}", report.summary());
789 Some(module)
790 }
791 Err(e) => {
792 error!(file = %plugin.path, error = %e, "WASM: plugin failed to load (skipped)");
793 None
794 }
795 }
796}
797
798fn plugin_name(path: &str) -> String {
800 std::path::Path::new(path)
801 .file_stem()
802 .and_then(|s| s.to_str())
803 .unwrap_or("plugin")
804 .to_string()
805}
806
807fn load_crs_module(files: &[String]) -> CrsModule {
814 let mut combined = String::new();
815 for path in files {
816 match std::fs::read_to_string(path) {
817 Ok(text) => {
818 combined.push_str(&text);
819 combined.push('\n');
820 }
821 Err(e) => error!(file = %path, error = %e, "CRS import: cannot read file (skipped)"),
822 }
823 }
824 let module = CrsModule::from_source(&combined);
825 info!(files = files.len(), "{}", module.report());
826 if !module.skipped().is_empty() {
827 warn!(
828 skipped = module.skipped().len(),
829 "CRS import: some rules fall outside the supported subset (see debug logs for reasons)"
830 );
831 for s in module.skipped() {
832 debug!(id = ?s.id, line = s.line_no, reason = %s.reason, "CRS import: rule skipped");
833 }
834 }
835 module
836}
837
838fn build_reloadable(
843 config: &Config,
844 rl_state: RateLimitState,
845 extra: Vec<Box<dyn WafModule>>,
846) -> Reloadable {
847 let mut modules = build_modules(config, &rl_state);
848 modules.extend(extra);
849 let pipeline = Pipeline::new(config, modules);
850
851 if config.waf.paranoia_level > waf_detection::HIGHEST_RULE_PARANOIA {
854 warn!(
855 paranoia_level = config.waf.paranoia_level,
856 highest_rule_paranoia = waf_detection::HIGHEST_RULE_PARANOIA,
857 "paranoia_level exceeds the highest existing rule paranoia: no additional rules are activated"
858 );
859 }
860 let ip_resolver = ClientIpResolver::from_config(&config.network);
861 if ip_resolver.trusted_count() < config.network.trusted_proxies.len() {
862 warn!(
863 configured = config.network.trusted_proxies.len(),
864 valid = ip_resolver.trusted_count(),
865 "some trusted_proxies CIDR entries were invalid and skipped"
866 );
867 }
868
869 Reloadable {
870 backend: config.proxy.backend.trim_end_matches('/').to_string(),
871 normalizer: Normalizer::new(&config.limits),
872 pipeline,
873 prefilter: ContentPrefilter::new(config.waf.paranoia_level),
875 ip_resolver,
876 resilience: config.resilience,
877 }
878}
879
880impl Proxy {
881 pub async fn bind(config: &Config) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
885 Self::builder(config).build().await
886 }
887
888 pub fn builder(config: &Config) -> ProxyBuilder<'_> {
893 ProxyBuilder {
894 config,
895 modules: Vec::new(),
896 state_store: None,
897 cert_source: None,
898 module_factory: None,
899 mode: HandlerMode::Inspect,
900 }
901 }
902
903 #[doc(hidden)]
909 pub async fn bind_with_modules(
910 config: &Config,
911 extra: Vec<Box<dyn WafModule>>,
912 ) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
913 Self::bind_inner(config, extra, HandlerMode::Inspect, None, None, None).await
914 }
915
916 #[doc(hidden)]
921 pub async fn bind_passthrough(
922 config: &Config,
923 ) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
924 Self::bind_inner(config, Vec::new(), HandlerMode::Passthrough, None, None, None).await
925 }
926
927 async fn bind_inner(
928 config: &Config,
929 extra: Vec<Box<dyn WafModule>>,
930 mode: HandlerMode,
931 state_store: Option<RateLimitState>,
932 cert_source: Option<Arc<dyn TlsCertSource>>,
933 module_factory: Option<Arc<ModuleFactory>>,
934 ) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
935 let listener = TcpListener::bind(config.proxy.listen).await?;
936 let listen_addr = listener.local_addr()?;
937 let client: Client<HttpConnector, HyperBoxBody> =
938 Client::builder(TokioExecutor::new()).build(HttpConnector::new());
939 let grpc_client: Client<HttpConnector, HyperBoxBody> =
941 Client::builder(TokioExecutor::new()).http2_only(true).build(HttpConnector::new());
942
943 let tls_acceptor = tls::acceptor_from_source(&config.tls, cert_source)
947 .map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
948 if tls_acceptor.is_some() {
949 info!(listen = %listen_addr, alpn = ?config.tls.alpn, "TLS termination enabled");
950 }
951
952 let rl_state = state_store
957 .unwrap_or_else(|| RateLimitState::in_memory(config.rate_limit.max_tracked_keys));
958
959 let mut extra_total = extra;
964 if let Some(factory) = &module_factory {
965 extra_total.extend(factory()?);
966 }
967 let reloadable = build_reloadable(config, rl_state.clone(), extra_total);
968
969 let metrics = Arc::new(Metrics::new());
973 let metrics_listener = if config.metrics.enabled {
974 let l = TcpListener::bind(config.metrics.listen).await?;
975 info!(listen = %l.local_addr()?, "metrics endpoint enabled (/metrics)");
976 Some(l)
977 } else {
978 None
979 };
980
981 Ok(Self {
982 listener,
983 state: Arc::new(StaticState {
984 client,
985 grpc_client,
986 listen_addr,
987 rl_state,
988 current: RwLock::new(Arc::new(reloadable)),
989 mode,
990 tls_acceptor,
991 metrics,
992 module_factory,
993 }),
994 metrics_listener,
995 })
996 }
997
998 pub fn reloader(&self) -> Reloader {
1002 Reloader(Arc::clone(&self.state))
1003 }
1004
1005 pub fn local_addr(&self) -> std::io::Result<SocketAddr> {
1006 self.listener.local_addr()
1007 }
1008
1009 pub fn metrics_addr(&self) -> Option<SocketAddr> {
1011 self.metrics_listener.as_ref().and_then(|l| l.local_addr().ok())
1012 }
1013
1014 pub async fn run(self) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
1015 if let Some(metrics_listener) = self.metrics_listener {
1018 let metrics = Arc::clone(&self.state.metrics);
1019 tokio::spawn(serve_metrics(metrics_listener, metrics));
1020 }
1021 loop {
1022 let (stream, client_addr) = self.listener.accept().await?;
1023 let state = Arc::clone(&self.state);
1024
1025 tokio::spawn(async move {
1026 match state.tls_acceptor.clone() {
1031 Some(acceptor) => match acceptor.accept(stream).await {
1032 Ok(tls_stream) => {
1033 serve_connection(TokioIo::new(tls_stream), state, client_addr).await;
1034 }
1035 Err(e) => {
1036 warn!(error = %e, client_ip = %client_addr.ip(), "TLS handshake error");
1037 }
1038 },
1039 None => {
1040 serve_connection(TokioIo::new(stream), state, client_addr).await;
1041 }
1042 }
1043 });
1044 }
1045 }
1046}
1047
1048pub struct ProxyBuilder<'a> {
1054 config: &'a Config,
1055 modules: Vec<Box<dyn WafModule>>,
1056 state_store: Option<RateLimitState>,
1057 cert_source: Option<Arc<dyn TlsCertSource>>,
1058 module_factory: Option<Arc<ModuleFactory>>,
1059 mode: HandlerMode,
1060}
1061
1062impl<'a> ProxyBuilder<'a> {
1063 pub fn modules(mut self, modules: Vec<Box<dyn WafModule>>) -> Self {
1066 self.modules = modules;
1067 self
1068 }
1069
1070 pub fn add_module(mut self, module: Box<dyn WafModule>) -> Self {
1072 self.modules.push(module);
1073 self
1074 }
1075
1076 pub fn state_store(mut self, store: Arc<dyn StateStore>) -> Self {
1080 self.state_store = Some(RateLimitState::with_store(store));
1081 self
1082 }
1083
1084 pub fn cert_source(mut self, source: Arc<dyn TlsCertSource>) -> Self {
1089 self.cert_source = Some(source);
1090 self
1091 }
1092
1093 pub fn module_factory<F>(mut self, factory: F) -> Self
1101 where
1102 F: Fn() -> Result<Vec<Box<dyn WafModule>>, Box<dyn std::error::Error + Send + Sync>>
1103 + Send
1104 + Sync
1105 + 'static,
1106 {
1107 self.module_factory = Some(Arc::new(factory));
1108 self
1109 }
1110
1111 pub async fn build(self) -> Result<Proxy, Box<dyn std::error::Error + Send + Sync>> {
1113 Proxy::bind_inner(
1114 self.config,
1115 self.modules,
1116 self.mode,
1117 self.state_store,
1118 self.cert_source,
1119 self.module_factory,
1120 )
1121 .await
1122 }
1123}
1124
1125async fn serve_connection<I>(io: I, state: Arc<StaticState>, client_addr: SocketAddr)
1129where
1130 I: hyper::rt::Read + hyper::rt::Write + Unpin + Send + 'static,
1131{
1132 let svc = service_fn(move |req| {
1133 let state = Arc::clone(&state);
1134 handle(req, state, client_addr)
1135 });
1136 if let Err(e) = auto::Builder::new(TokioExecutor::new())
1137 .serve_connection(io, svc)
1138 .await
1139 {
1140 warn!(error = %e, client_ip = %client_addr.ip(), "connection error");
1141 }
1142}
1143
1144async fn serve_metrics(listener: TcpListener, metrics: Arc<Metrics>) {
1148 loop {
1149 let Ok((stream, _)) = listener.accept().await else { continue };
1150 let metrics = Arc::clone(&metrics);
1151 tokio::spawn(async move {
1152 let svc = service_fn(move |req: Request<Incoming>| {
1153 let metrics = Arc::clone(&metrics);
1154 async move { Ok::<_, Infallible>(metrics_response(&req, &metrics)) }
1155 });
1156 let _ = hyper::server::conn::http1::Builder::new()
1157 .serve_connection(TokioIo::new(stream), svc)
1158 .await;
1159 });
1160 }
1161}
1162
1163fn metrics_response(req: &Request<Incoming>, metrics: &Metrics) -> Response<HyperBoxBody> {
1165 if req.method() == hyper::Method::GET && req.uri().path() == "/metrics" {
1166 Response::builder()
1167 .status(200)
1168 .header("content-type", "text/plain; version=0.0.4; charset=utf-8")
1169 .body(full_body(metrics.render()))
1170 .unwrap()
1171 } else {
1172 Response::builder().status(404).body(full_body("Not Found")).unwrap()
1173 }
1174}
1175
1176#[cfg(test)]
1177mod tests {
1178 use super::*;
1179 use waf_core::WafMode;
1180
1181 #[test]
1182 fn hop_by_hop_includes_connection_and_host() {
1183 assert!(HOP_BY_HOP.contains(&"connection"));
1184 assert!(HOP_BY_HOP.contains(&"host"));
1185 assert!(HOP_BY_HOP.contains(&"transfer-encoding"));
1186 }
1187
1188 #[test]
1189 fn hop_by_hop_excludes_regular_headers() {
1190 assert!(!HOP_BY_HOP.contains(&"content-type"));
1191 assert!(!HOP_BY_HOP.contains(&"authorization"));
1192 assert!(!HOP_BY_HOP.contains(&"x-custom-header"));
1193 }
1194
1195 #[test]
1196 fn config_parses_from_toml() {
1197 let raw = r#"
1198[proxy]
1199listen = "127.0.0.1:8080"
1200backend = "http://localhost:3000"
1201
1202[waf]
1203mode = "detection-only"
1204block_threshold = 10
1205"#;
1206 let config: Config = toml::from_str(raw).unwrap();
1207 assert_eq!(config.proxy.backend, "http://localhost:3000");
1208 assert_eq!(config.waf.mode, WafMode::DetectionOnly);
1209 assert_eq!(config.waf.block_threshold, 10);
1210 }
1211
1212 #[test]
1213 fn config_uses_default_block_threshold_when_omitted() {
1214 let raw = r#"
1215[proxy]
1216listen = "127.0.0.1:8080"
1217backend = "http://localhost:3000"
1218
1219[waf]
1220mode = "detection-only"
1221"#;
1222 let config: Config = toml::from_str(raw).unwrap();
1223 assert_eq!(config.waf.block_threshold, 5);
1224 }
1225
1226 #[test]
1227 fn config_parses_network_section() {
1228 let raw = r#"
1229[proxy]
1230listen = "127.0.0.1:8080"
1231backend = "http://localhost:3000"
1232
1233[waf]
1234mode = "blocking"
1235
1236[network]
1237trusted_proxies = ["10.0.0.0/8", "::1"]
1238client_ip_header = "X-Forwarded-For"
1239trusted_hops = 2
1240"#;
1241 let config: Config = toml::from_str(raw).unwrap();
1242 assert_eq!(config.network.trusted_proxies, vec!["10.0.0.0/8", "::1"]);
1243 assert_eq!(config.network.client_ip_header, "X-Forwarded-For");
1244 assert_eq!(config.network.trusted_hops, 2);
1245 }
1246
1247 #[test]
1248 fn config_network_defaults_to_failsafe_when_absent() {
1249 let raw = r#"
1250[proxy]
1251listen = "127.0.0.1:8080"
1252backend = "http://localhost:3000"
1253
1254[waf]
1255mode = "detection-only"
1256"#;
1257 let config: Config = toml::from_str(raw).unwrap();
1258 assert!(config.network.trusted_proxies.is_empty());
1259 assert_eq!(config.network.trusted_hops, 1);
1260 assert_eq!(config.network.client_ip_header, "x-forwarded-for".to_string());
1261 }
1262
1263 #[test]
1264 fn config_rejects_unknown_mode() {
1265 let raw = r#"
1266[proxy]
1267listen = "127.0.0.1:8080"
1268backend = "http://localhost:3000"
1269
1270[waf]
1271mode = "unknown-mode"
1272"#;
1273 assert!(toml::from_str::<Config>(raw).is_err());
1274 }
1275
1276 #[test]
1277 fn parse_cookies_splits_on_semicolon() {
1278 let headers = vec![("cookie".to_string(), "session=abc; user=123".to_string())];
1279 let cookies = parse_cookies(&headers);
1280 assert_eq!(cookies.len(), 2);
1281 assert!(cookies.contains(&("session".to_string(), "abc".to_string())));
1282 assert!(cookies.contains(&("user".to_string(), "123".to_string())));
1283 }
1284
1285 #[test]
1286 fn parse_cookies_handles_missing_value() {
1287 let headers = vec![("cookie".to_string(), "flag=; token=xyz".to_string())];
1288 let cookies = parse_cookies(&headers);
1289 assert!(cookies.contains(&("flag".to_string(), "".to_string())));
1290 assert!(cookies.contains(&("token".to_string(), "xyz".to_string())));
1291 }
1292
1293 #[test]
1294 fn parse_cookies_handles_empty_header_list() {
1295 assert!(parse_cookies(&[]).is_empty());
1296 }
1297
1298 #[test]
1299 fn request_id_is_unique_per_call() {
1300 let id1 = next_request_id();
1301 let id2 = next_request_id();
1302 assert_ne!(id1, id2);
1303 assert!(id1.starts_with("req-"));
1304 }
1305}