Skip to main content

gateway/
lib.rs

1//! Health-aware gateway routing and an Actix/Reqwest reverse proxy.
2
3mod transcode;
4
5pub use transcode::{
6    grpc_status_to_http, transcode, HttpBinding, HttpVerb, TranscodeError, Transcoder,
7    TranscoderBuilder,
8};
9
10use actix_web::{
11    http::{header, StatusCode},
12    web, App, HttpRequest, HttpResponse, HttpServer,
13};
14use futures::{future::LocalBoxFuture, StreamExt};
15use serde::{Deserialize, Serialize};
16use std::{
17    cmp::Reverse,
18    collections::{HashMap, HashSet},
19    future::Future,
20    io,
21    net::{SocketAddr, TcpListener},
22    path::Path,
23    rc::Rc,
24    sync::{
25        atomic::{AtomicBool, AtomicUsize, Ordering},
26        Arc,
27    },
28    time::Duration,
29};
30
31fn default_address() -> SocketAddr {
32    "127.0.0.1:8080".parse().expect("static socket address")
33}
34
35fn default_workers() -> usize {
36    1
37}
38
39fn default_timeout_ms() -> u64 {
40    30_000
41}
42
43fn default_request_body_limit() -> usize {
44    10 * 1024 * 1024
45}
46
47fn default_response_body_limit() -> usize {
48    50 * 1024 * 1024
49}
50
51/// File-loadable configuration for the HTTP gateway runtime.
52#[derive(Debug, Clone, Deserialize, Serialize)]
53#[serde(default, deny_unknown_fields)]
54pub struct GatewayConfig {
55    pub address: SocketAddr,
56    pub workers: usize,
57    pub request_timeout_ms: u64,
58    pub shutdown_timeout_ms: u64,
59    pub request_body_limit: usize,
60    pub response_body_limit: usize,
61    pub routes: Vec<GatewayRoute>,
62}
63
64impl Default for GatewayConfig {
65    fn default() -> Self {
66        Self {
67            address: default_address(),
68            workers: default_workers(),
69            request_timeout_ms: default_timeout_ms(),
70            shutdown_timeout_ms: default_timeout_ms(),
71            request_body_limit: default_request_body_limit(),
72            response_body_limit: default_response_body_limit(),
73            routes: Vec::new(),
74        }
75    }
76}
77
78impl GatewayConfig {
79    pub fn load(path: impl AsRef<Path>) -> Result<Self, GatewayConfigError> {
80        let config: Self = rust_zero_core::load_config(path).map_err(GatewayConfigError::Load)?;
81        config.validate()?;
82        Ok(config)
83    }
84
85    pub fn validate(&self) -> Result<(), GatewayConfigError> {
86        if self.workers == 0 {
87            return Err(GatewayConfigError::Invalid(
88                "workers must be greater than zero",
89            ));
90        }
91        if self.request_timeout_ms == 0 {
92            return Err(GatewayConfigError::Invalid(
93                "request_timeout_ms must be greater than zero",
94            ));
95        }
96        if self.shutdown_timeout_ms == 0 {
97            return Err(GatewayConfigError::Invalid(
98                "shutdown_timeout_ms must be greater than zero",
99            ));
100        }
101        if self.request_body_limit == 0 || self.response_body_limit == 0 {
102            return Err(GatewayConfigError::Invalid(
103                "gateway body limits must be greater than zero",
104            ));
105        }
106        if self.routes.is_empty() {
107            return Err(GatewayConfigError::Invalid(
108                "at least one gateway route is required",
109            ));
110        }
111        for route in &self.routes {
112            let normalized =
113                normalize_prefix(route.prefix.clone()).map_err(GatewayConfigError::Route)?;
114            if normalized != route.prefix {
115                return Err(GatewayConfigError::Invalid(
116                    "gateway route prefixes must not have a trailing slash",
117                ));
118            }
119            if route.upstreams.is_empty() {
120                return Err(GatewayConfigError::Route(GatewayError::EmptyUpstreams(
121                    route.prefix.clone(),
122                )));
123            }
124            validate_middleware_names(&route.middleware).map_err(GatewayConfigError::Middleware)?;
125            for upstream in &route.upstreams {
126                let url = reqwest::Url::parse(upstream).map_err(|_| {
127                    GatewayConfigError::Invalid("gateway upstream must be a valid URL")
128                })?;
129                if !matches!(url.scheme(), "http" | "https") || url.host_str().is_none() {
130                    return Err(GatewayConfigError::Invalid(
131                        "gateway upstream must use HTTP or HTTPS and include a host",
132                    ));
133                }
134            }
135        }
136        GatewayRouter::new(self.routes.clone()).map_err(GatewayConfigError::Route)?;
137        Ok(())
138    }
139}
140
141/// Configuration loading or validation failure.
142#[derive(Debug)]
143pub enum GatewayConfigError {
144    Load(rust_zero_core::ConfigError),
145    Route(GatewayError),
146    Middleware(String),
147    Invalid(&'static str),
148}
149
150impl std::fmt::Display for GatewayConfigError {
151    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
152        match self {
153            Self::Load(error) => write!(formatter, "failed to load gateway configuration: {error}"),
154            Self::Route(error) => write!(formatter, "invalid gateway route: {error}"),
155            Self::Middleware(error) => write!(formatter, "invalid gateway middleware: {error}"),
156            Self::Invalid(message) => formatter.write_str(message),
157        }
158    }
159}
160
161impl std::error::Error for GatewayConfigError {
162    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
163        match self {
164            Self::Load(error) => Some(error),
165            Self::Route(error) => Some(error),
166            Self::Middleware(_) | Self::Invalid(_) => None,
167        }
168    }
169}
170
171/// A configured HTTP path prefix and its upstream endpoints.
172#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
173#[serde(deny_unknown_fields)]
174pub struct GatewayRoute {
175    pub prefix: String,
176    pub upstreams: Vec<String>,
177    /// Ordered application middleware applied to requests using this upstream pool.
178    #[serde(default)]
179    pub middleware: Vec<String>,
180}
181
182impl GatewayRoute {
183    pub fn new(prefix: impl Into<String>, upstreams: Vec<String>) -> Result<Self, GatewayError> {
184        let prefix = normalize_prefix(prefix.into())?;
185        if upstreams.is_empty() {
186            return Err(GatewayError::EmptyUpstreams(prefix));
187        }
188        if upstreams.iter().any(|upstream| upstream.is_empty()) {
189            return Err(GatewayError::EmptyUpstream);
190        }
191
192        Ok(Self {
193            prefix,
194            upstreams,
195            middleware: Vec::new(),
196        })
197    }
198}
199
200struct RoutePool {
201    route: GatewayRoute,
202    next_upstream: AtomicUsize,
203    healthy: Vec<AtomicBool>,
204}
205
206/// Selects the most specific route and distributes requests across its upstreams.
207pub struct GatewayRouter {
208    routes: Vec<RoutePool>,
209}
210
211impl GatewayRouter {
212    pub fn new(routes: impl IntoIterator<Item = GatewayRoute>) -> Result<Self, GatewayError> {
213        let mut routes: Vec<_> = routes
214            .into_iter()
215            .map(|route| RoutePool {
216                healthy: route
217                    .upstreams
218                    .iter()
219                    .map(|_| AtomicBool::new(true))
220                    .collect(),
221                route,
222                next_upstream: AtomicUsize::new(0),
223            })
224            .collect();
225
226        routes.sort_unstable_by_key(|route| Reverse(route.route.prefix.len()));
227        for routes_with_same_prefix in routes.windows(2) {
228            if routes_with_same_prefix[0].route.prefix == routes_with_same_prefix[1].route.prefix {
229                return Err(GatewayError::DuplicatePrefix(
230                    routes_with_same_prefix[0].route.prefix.clone(),
231                ));
232            }
233        }
234
235        Ok(Self { routes })
236    }
237
238    /// Selects an upstream for a request path, using round robin within its matched route.
239    pub fn select(&self, path: &str) -> Option<&str> {
240        self.routes
241            .iter()
242            .find(|route| matches_prefix(path, &route.route.prefix))
243            .and_then(select_healthy)
244    }
245
246    /// Builds the full upstream URL for a path and optional query string.
247    pub fn select_target(&self, path_and_query: &str) -> Option<String> {
248        self.select_target_with_middleware(path_and_query)
249            .map(|(target, _)| target)
250    }
251
252    fn select_target_with_middleware(&self, path_and_query: &str) -> Option<(String, &[String])> {
253        let path = path_and_query.split('?').next().unwrap_or(path_and_query);
254        let route = self
255            .routes
256            .iter()
257            .find(|route| matches_prefix(path, &route.route.prefix))?;
258        select_healthy(route).map(|upstream| {
259            let target = format!(
260                "{}{}",
261                upstream.trim_end_matches('/'),
262                if path_and_query.starts_with('/') {
263                    path_and_query.to_owned()
264                } else {
265                    format!("/{path_and_query}")
266                }
267            );
268            (target, route.route.middleware.as_slice())
269        })
270    }
271
272    /// Includes or excludes an upstream from selection, for use by active health checks.
273    pub fn set_upstream_health(
274        &self,
275        prefix: &str,
276        upstream: &str,
277        healthy: bool,
278    ) -> Result<(), GatewayError> {
279        let route = self
280            .routes
281            .iter()
282            .find(|route| route.route.prefix == prefix)
283            .ok_or_else(|| GatewayError::UnknownPrefix(prefix.to_owned()))?;
284        let index = route
285            .route
286            .upstreams
287            .iter()
288            .position(|candidate| candidate == upstream)
289            .ok_or_else(|| GatewayError::UnknownUpstream(upstream.to_owned()))?;
290        route.healthy[index].store(healthy, Ordering::Release);
291        Ok(())
292    }
293}
294
295/// An outbound request passed through a configured gateway middleware chain.
296pub struct GatewayMiddlewareRequest {
297    request: reqwest::Request,
298}
299
300impl GatewayMiddlewareRequest {
301    pub fn request(&self) -> &reqwest::Request {
302        &self.request
303    }
304
305    pub fn request_mut(&mut self) -> &mut reqwest::Request {
306        &mut self.request
307    }
308
309    pub fn into_request(self) -> reqwest::Request {
310        self.request
311    }
312}
313
314/// The boxed future returned by application-defined upstream middleware.
315pub type GatewayMiddlewareFuture = LocalBoxFuture<'static, HttpResponse>;
316
317/// The remainder of an upstream middleware chain and its network dispatch.
318#[derive(Clone)]
319pub struct GatewayMiddlewareNext {
320    inner: Rc<dyn Fn(GatewayMiddlewareRequest) -> GatewayMiddlewareFuture>,
321}
322
323impl GatewayMiddlewareNext {
324    pub fn call(&self, request: GatewayMiddlewareRequest) -> GatewayMiddlewareFuture {
325        (self.inner)(request)
326    }
327}
328
329/// Type-erased application policy registered by name on [`GatewayProxy`] or [`GatewayServer`].
330pub trait GatewayMiddleware: Send + Sync + 'static {
331    fn call(
332        &self,
333        request: GatewayMiddlewareRequest,
334        next: GatewayMiddlewareNext,
335    ) -> GatewayMiddlewareFuture;
336}
337
338impl<F, Fut> GatewayMiddleware for F
339where
340    F: Fn(GatewayMiddlewareRequest, GatewayMiddlewareNext) -> Fut + Send + Sync + 'static,
341    Fut: Future<Output = HttpResponse> + 'static,
342{
343    fn call(
344        &self,
345        request: GatewayMiddlewareRequest,
346        next: GatewayMiddlewareNext,
347    ) -> GatewayMiddlewareFuture {
348        Box::pin((self)(request, next))
349    }
350}
351
352fn select_healthy(route: &RoutePool) -> Option<&str> {
353    let len = route.route.upstreams.len();
354    let start = route.next_upstream.fetch_add(1, Ordering::Relaxed) % len;
355    (0..len).find_map(|offset| {
356        let index = (start + offset) % len;
357        route.healthy[index]
358            .load(Ordering::Acquire)
359            .then(|| route.route.upstreams[index].as_str())
360    })
361}
362
363/// An HTTP reverse proxy backed by a [`GatewayRouter`].
364#[derive(Clone)]
365pub struct GatewayProxy {
366    router: Arc<GatewayRouter>,
367    client: reqwest::Client,
368    request_body_limit: usize,
369    response_body_limit: usize,
370    timeout: Duration,
371    middleware: Arc<HashMap<String, Arc<dyn GatewayMiddleware>>>,
372}
373
374impl GatewayProxy {
375    pub fn new(router: GatewayRouter) -> Self {
376        Self {
377            router: Arc::new(router),
378            client: reqwest::Client::new(),
379            request_body_limit: 10 * 1024 * 1024,
380            response_body_limit: 50 * 1024 * 1024,
381            timeout: Duration::from_secs(30),
382            middleware: Arc::new(HashMap::new()),
383        }
384    }
385
386    pub fn with_client(mut self, client: reqwest::Client) -> Self {
387        self.client = client;
388        self
389    }
390
391    pub fn with_request_body_limit(mut self, bytes: usize) -> Self {
392        assert!(bytes > 0, "request body limit must be greater than zero");
393        self.request_body_limit = bytes;
394        self
395    }
396
397    pub fn with_response_body_limit(mut self, bytes: usize) -> Self {
398        assert!(bytes > 0, "response body limit must be greater than zero");
399        self.response_body_limit = bytes;
400        self
401    }
402
403    pub fn with_timeout(mut self, timeout: Duration) -> Self {
404        assert!(
405            !timeout.is_zero(),
406            "gateway timeout must be greater than zero"
407        );
408        self.timeout = timeout;
409        self
410    }
411
412    /// Registers a named policy referenced by one or more configured upstream pools.
413    pub fn with_upstream_middleware<M>(
414        mut self,
415        name: impl Into<String>,
416        middleware: M,
417    ) -> Result<Self, GatewayConfigError>
418    where
419        M: GatewayMiddleware,
420    {
421        let name = name.into();
422        if name.trim().is_empty() {
423            return Err(GatewayConfigError::Middleware(
424                "middleware name must not be empty".to_owned(),
425            ));
426        }
427        let registry = Arc::make_mut(&mut self.middleware);
428        if registry
429            .insert(name.clone(), Arc::new(middleware))
430            .is_some()
431        {
432            return Err(GatewayConfigError::Middleware(format!(
433                "middleware '{name}' is already registered"
434            )));
435        }
436        Ok(self)
437    }
438
439    fn validate_middleware(&self) -> Result<(), GatewayConfigError> {
440        for route in &self.router.routes {
441            for name in &route.route.middleware {
442                if !self.middleware.contains_key(name) {
443                    return Err(GatewayConfigError::Middleware(format!(
444                        "middleware '{name}' is not registered"
445                    )));
446                }
447            }
448        }
449        Ok(())
450    }
451
452    pub fn router(&self) -> &GatewayRouter {
453        &self.router
454    }
455
456    pub async fn forward(&self, request: HttpRequest, mut payload: web::Payload) -> HttpResponse {
457        let path_and_query = request
458            .uri()
459            .path_and_query()
460            .map_or_else(|| request.path(), |value| value.as_str());
461        let Some((target, middleware)) = self.router.select_target_with_middleware(path_and_query)
462        else {
463            return HttpResponse::NotFound().body("no healthy gateway upstream");
464        };
465        let middleware: Arc<[String]> = middleware.to_vec().into();
466
467        let mut request_body = web::BytesMut::new();
468        while let Some(chunk) = payload.next().await {
469            let chunk = match chunk {
470                Ok(chunk) => chunk,
471                Err(_) => return HttpResponse::BadRequest().body("invalid request body"),
472            };
473            if request_body.len().saturating_add(chunk.len()) > self.request_body_limit {
474                return HttpResponse::PayloadTooLarge().body("gateway request body limit exceeded");
475            }
476            request_body.extend_from_slice(&chunk);
477        }
478
479        let method = match reqwest::Method::from_bytes(request.method().as_str().as_bytes()) {
480            Ok(method) => method,
481            Err(_) => return HttpResponse::BadRequest().body("invalid request method"),
482        };
483        let mut upstream = self
484            .client
485            .request(method, target)
486            .timeout(self.timeout)
487            .body(request_body.freeze());
488        for (name, value) in request.headers() {
489            if is_hop_by_hop(name.as_str())
490                || name == header::HOST
491                || name == header::CONTENT_LENGTH
492            {
493                continue;
494            }
495            upstream = upstream.header(name.as_str(), value.as_bytes());
496        }
497        if let Some(peer) = request.peer_addr() {
498            upstream = upstream.header("x-forwarded-for", peer.ip().to_string());
499        }
500        upstream = upstream
501            .header("x-forwarded-proto", request.connection_info().scheme())
502            .header("x-forwarded-host", request.connection_info().host());
503        let upstream = match upstream.build() {
504            Ok(request) => GatewayMiddlewareRequest { request },
505            Err(_) => return HttpResponse::BadGateway().body("invalid gateway upstream request"),
506        };
507
508        let client = self.client.clone();
509        let response_body_limit = self.response_body_limit;
510        let terminal = GatewayMiddlewareNext {
511            inner: Rc::new(move |request| {
512                let client = client.clone();
513                Box::pin(async move {
514                    dispatch_upstream(client, request.into_request(), response_body_limit).await
515                })
516            }),
517        };
518        let registry = Arc::clone(&self.middleware);
519        let chain = middleware.iter().rev().fold(terminal, |next, name| {
520            let Some(policy) = registry.get(name).cloned() else {
521                return GatewayMiddlewareNext {
522                    inner: Rc::new(move |_| {
523                        Box::pin(async {
524                            HttpResponse::InternalServerError()
525                                .body("gateway middleware is not registered")
526                        })
527                    }),
528                };
529            };
530            GatewayMiddlewareNext {
531                inner: Rc::new(move |request| policy.call(request, next.clone())),
532            }
533        });
534        chain.call(upstream).await
535    }
536}
537
538async fn dispatch_upstream(
539    client: reqwest::Client,
540    request: reqwest::Request,
541    response_body_limit: usize,
542) -> HttpResponse {
543    let response = match client.execute(request).await {
544        Ok(response) => response,
545        Err(error) if error.is_timeout() => {
546            return HttpResponse::GatewayTimeout().body("gateway upstream timed out");
547        }
548        Err(_) => return HttpResponse::BadGateway().body("gateway upstream unavailable"),
549    };
550    if response
551        .content_length()
552        .is_some_and(|length| length > response_body_limit as u64)
553    {
554        return HttpResponse::BadGateway().body("gateway response body limit exceeded");
555    }
556
557    let status =
558        StatusCode::from_u16(response.status().as_u16()).unwrap_or(StatusCode::BAD_GATEWAY);
559    let headers: Vec<_> = response
560        .headers()
561        .iter()
562        .filter(|(name, _)| {
563            !is_hop_by_hop(name.as_str()) && *name != reqwest::header::CONTENT_LENGTH
564        })
565        .map(|(name, value)| (name.as_str().to_owned(), value.as_bytes().to_vec()))
566        .collect();
567    let mut downstream = HttpResponse::build(status);
568    for (name, value) in headers {
569        if let (Ok(name), Ok(value)) = (
570            header::HeaderName::try_from(name),
571            header::HeaderValue::from_bytes(&value),
572        ) {
573            downstream.insert_header((name, value));
574        }
575    }
576    let limit = response_body_limit;
577    let mut received = 0usize;
578    let stream = response.bytes_stream().map(move |chunk| match chunk {
579        Ok(chunk) => {
580            received = received.saturating_add(chunk.len());
581            if received > limit {
582                Err(actix_web::error::ErrorBadGateway(
583                    "gateway response body limit exceeded",
584                ))
585            } else {
586                Ok(chunk)
587            }
588        }
589        Err(_) => Err(actix_web::error::ErrorBadGateway(
590            "invalid upstream response",
591        )),
592    });
593    downstream.streaming(stream)
594}
595
596/// Configuration-driven Actix gateway with bounded graceful draining.
597pub struct GatewayServer {
598    config: GatewayConfig,
599    proxy: GatewayProxy,
600}
601
602impl GatewayServer {
603    pub fn new(config: GatewayConfig) -> Result<Self, GatewayConfigError> {
604        config.validate()?;
605        let router =
606            GatewayRouter::new(config.routes.clone()).map_err(GatewayConfigError::Route)?;
607        let proxy = GatewayProxy::new(router)
608            .with_request_body_limit(config.request_body_limit)
609            .with_response_body_limit(config.response_body_limit)
610            .with_timeout(Duration::from_millis(config.request_timeout_ms));
611        Ok(Self { config, proxy })
612    }
613
614    /// Registers a named policy referenced by configured upstream pools.
615    pub fn with_upstream_middleware<M>(
616        mut self,
617        name: impl Into<String>,
618        middleware: M,
619    ) -> Result<Self, GatewayConfigError>
620    where
621        M: GatewayMiddleware,
622    {
623        self.proxy = self.proxy.with_upstream_middleware(name, middleware)?;
624        Ok(self)
625    }
626
627    pub fn run(&self) -> io::Result<actix_web::dev::Server> {
628        self.proxy
629            .validate_middleware()
630            .map_err(|error| io::Error::new(io::ErrorKind::InvalidInput, error))?;
631        let listener = TcpListener::bind(self.config.address)?;
632        self.run_on(listener)
633    }
634
635    pub fn run_on(&self, listener: TcpListener) -> io::Result<actix_web::dev::Server> {
636        self.proxy
637            .validate_middleware()
638            .map_err(|error| io::Error::new(io::ErrorKind::InvalidInput, error))?;
639        let proxy = self.proxy.clone();
640        let workers = self.config.workers;
641        let shutdown_seconds = self.config.shutdown_timeout_ms.div_ceil(1_000);
642        HttpServer::new(move || {
643            App::new()
644                .app_data(web::Data::new(proxy.clone()))
645                .default_service(web::to(crate::proxy))
646        })
647        .workers(workers)
648        .shutdown_timeout(shutdown_seconds)
649        .listen(listener)
650        .map(HttpServer::run)
651    }
652
653    pub async fn serve_until<F>(&self, shutdown: F) -> io::Result<()>
654    where
655        F: Future<Output = ()>,
656    {
657        drain_on_signal(self.run()?, shutdown).await
658    }
659
660    pub async fn serve_on_until<F>(&self, listener: TcpListener, shutdown: F) -> io::Result<()>
661    where
662        F: Future<Output = ()>,
663    {
664        drain_on_signal(self.run_on(listener)?, shutdown).await
665    }
666}
667
668async fn drain_on_signal<F>(server: actix_web::dev::Server, shutdown: F) -> io::Result<()>
669where
670    F: Future<Output = ()>,
671{
672    use futures::future::{select, Either};
673
674    let handle = server.handle();
675    match select(Box::pin(server), Box::pin(shutdown)).await {
676        Either::Left((result, _)) => result,
677        Either::Right(((), server)) => {
678            let (_, result) = futures::future::join(handle.stop(true), server).await;
679            result
680        }
681    }
682}
683
684/// Actix handler for mounting a [`GatewayProxy`] stored in `web::Data`.
685pub async fn proxy(
686    gateway: web::Data<GatewayProxy>,
687    request: HttpRequest,
688    payload: web::Payload,
689) -> HttpResponse {
690    gateway.forward(request, payload).await
691}
692
693fn is_hop_by_hop(name: &str) -> bool {
694    matches!(
695        name.to_ascii_lowercase().as_str(),
696        "connection"
697            | "keep-alive"
698            | "proxy-authenticate"
699            | "proxy-authorization"
700            | "te"
701            | "trailer"
702            | "transfer-encoding"
703            | "upgrade"
704    )
705}
706
707/// Errors produced by invalid gateway routing configuration.
708#[derive(Debug, Clone, PartialEq, Eq)]
709pub enum GatewayError {
710    InvalidPrefix(String),
711    EmptyUpstreams(String),
712    EmptyUpstream,
713    DuplicatePrefix(String),
714    UnknownPrefix(String),
715    UnknownUpstream(String),
716}
717
718impl std::fmt::Display for GatewayError {
719    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
720        match self {
721            Self::InvalidPrefix(prefix) => {
722                write!(formatter, "gateway prefix must start with '/': {prefix}")
723            }
724            Self::EmptyUpstreams(prefix) => {
725                write!(formatter, "gateway route {prefix} has no upstreams")
726            }
727            Self::EmptyUpstream => formatter.write_str("gateway upstream cannot be empty"),
728            Self::DuplicatePrefix(prefix) => {
729                write!(formatter, "duplicate gateway prefix: {prefix}")
730            }
731            Self::UnknownPrefix(prefix) => write!(formatter, "unknown gateway prefix: {prefix}"),
732            Self::UnknownUpstream(upstream) => {
733                write!(formatter, "unknown gateway upstream: {upstream}")
734            }
735        }
736    }
737}
738
739impl std::error::Error for GatewayError {}
740
741fn normalize_prefix(mut prefix: String) -> Result<String, GatewayError> {
742    if !prefix.starts_with('/') {
743        return Err(GatewayError::InvalidPrefix(prefix));
744    }
745    while prefix.len() > 1 && prefix.ends_with('/') {
746        prefix.pop();
747    }
748    Ok(prefix)
749}
750
751fn matches_prefix(path: &str, prefix: &str) -> bool {
752    prefix == "/"
753        || path == prefix
754        || path
755            .strip_prefix(prefix)
756            .is_some_and(|suffix| suffix.starts_with('/'))
757}
758
759fn validate_middleware_names(names: &[String]) -> Result<(), String> {
760    let mut unique = HashSet::new();
761    for name in names {
762        if name.trim().is_empty() {
763            return Err("middleware name must not be empty".to_owned());
764        }
765        if !unique.insert(name) {
766            return Err(format!("duplicate middleware name: {name}"));
767        }
768    }
769    Ok(())
770}
771
772#[cfg(test)]
773mod tests {
774    use super::*;
775    use actix_web::{test as actix_test, App, HttpServer};
776    use rust_zero_core::{parse_config, ConfigFormat};
777
778    fn route(prefix: &str, upstreams: &[&str]) -> GatewayRoute {
779        GatewayRoute::new(
780            prefix,
781            upstreams
782                .iter()
783                .map(|upstream| (*upstream).to_owned())
784                .collect(),
785        )
786        .unwrap()
787    }
788
789    #[test]
790    fn parses_and_validates_gateway_configuration() {
791        let config: GatewayConfig = parse_config(
792            r#"
793address = "127.0.0.1:9100"
794workers = 2
795request_timeout_ms = 1500
796shutdown_timeout_ms = 4000
797request_body_limit = 1024
798response_body_limit = 2048
799
800[[routes]]
801prefix = "/api"
802upstreams = ["http://api-a:8080", "https://api-b:8443"]
803middleware = ["sign", "audit"]
804"#,
805            ConfigFormat::Toml,
806        )
807        .unwrap();
808
809        config.validate().unwrap();
810        assert_eq!(config.workers, 2);
811        assert_eq!(config.routes[0].prefix, "/api");
812        assert_eq!(config.routes[0].middleware, ["sign", "audit"]);
813
814        let mut invalid = config;
815        invalid.routes[0].upstreams[0] = "ftp://api-a".to_owned();
816        assert!(invalid.validate().unwrap_err().to_string().contains("HTTP"));
817    }
818
819    #[test]
820    fn rejects_invalid_or_unregistered_middleware_names() {
821        let mut configured = route("/api", &["http://api"]);
822        configured.middleware = vec!["audit".to_owned(), "audit".to_owned()];
823        let config = GatewayConfig {
824            routes: vec![configured],
825            ..GatewayConfig::default()
826        };
827        assert!(config
828            .validate()
829            .unwrap_err()
830            .to_string()
831            .contains("duplicate"));
832
833        let mut configured = route("/api", &["http://api"]);
834        configured.middleware = vec!["audit".to_owned()];
835        let proxy = GatewayProxy::new(GatewayRouter::new([configured]).unwrap());
836        assert!(proxy
837            .validate_middleware()
838            .unwrap_err()
839            .to_string()
840            .contains("not registered"));
841    }
842
843    #[actix_web::test]
844    async fn applies_named_upstream_middleware_in_order_and_can_short_circuit() {
845        let mut configured = route("/api", &["http://unused.invalid"]);
846        configured.middleware = vec!["decorate".to_owned(), "authorize".to_owned()];
847        let gateway = GatewayProxy::new(GatewayRouter::new([configured]).unwrap())
848            .with_upstream_middleware(
849                "decorate",
850                |mut request: GatewayMiddlewareRequest, next: GatewayMiddlewareNext| async move {
851                    request
852                        .request_mut()
853                        .headers_mut()
854                        .insert("x-gateway-policy", "decorated".parse().unwrap());
855                    let mut response = next.call(request).await;
856                    response.headers_mut().insert(
857                        header::HeaderName::from_static("x-policy-response"),
858                        header::HeaderValue::from_static("wrapped"),
859                    );
860                    response
861                },
862            )
863            .unwrap()
864            .with_upstream_middleware(
865                "authorize",
866                |request: GatewayMiddlewareRequest, _next: GatewayMiddlewareNext| async move {
867                    assert_eq!(request.request().headers()["x-gateway-policy"], "decorated");
868                    HttpResponse::Accepted().body("short-circuited")
869                },
870            )
871            .unwrap();
872        gateway.validate_middleware().unwrap();
873
874        let app = actix_test::init_service(
875            App::new()
876                .app_data(web::Data::new(gateway))
877                .default_service(web::to(proxy)),
878        )
879        .await;
880        let response = actix_test::call_service(
881            &app,
882            actix_test::TestRequest::get()
883                .uri("/api/items")
884                .to_request(),
885        )
886        .await;
887
888        assert_eq!(response.status(), StatusCode::ACCEPTED);
889        assert_eq!(
890            response.headers().get("x-policy-response").unwrap(),
891            "wrapped"
892        );
893        assert_eq!(actix_test::read_body(response).await, "short-circuited");
894    }
895
896    #[test]
897    fn selects_the_most_specific_matching_prefix() {
898        let router = GatewayRouter::new([
899            route("/", &["http://home"]),
900            route("/api", &["http://api"]),
901            route("/api/admin", &["http://admin"]),
902        ])
903        .unwrap();
904
905        assert_eq!(router.select("/api/admin/users"), Some("http://admin"));
906        assert_eq!(router.select("/api/users"), Some("http://api"));
907        assert_eq!(router.select("/apix"), Some("http://home"));
908    }
909
910    #[test]
911    fn rotates_through_matched_route_upstreams() {
912        let router = GatewayRouter::new([route("/api", &["http://one", "http://two"])]).unwrap();
913
914        assert_eq!(router.select("/api/items"), Some("http://one"));
915        assert_eq!(router.select("/api/items"), Some("http://two"));
916        assert_eq!(router.select("/api/items"), Some("http://one"));
917    }
918
919    #[test]
920    fn skips_unhealthy_upstreams_and_builds_targets() {
921        let router = GatewayRouter::new([route("/api", &["http://one/", "http://two"])]).unwrap();
922        router
923            .set_upstream_health("/api", "http://one/", false)
924            .unwrap();
925
926        assert_eq!(router.select("/api/items"), Some("http://two"));
927        assert_eq!(
928            router.select_target("/api/items?page=2"),
929            Some("http://two/api/items?page=2".to_owned())
930        );
931        router
932            .set_upstream_health("/api", "http://two", false)
933            .unwrap();
934        assert_eq!(router.select("/api/items"), None);
935    }
936
937    #[test]
938    fn identifies_hop_by_hop_headers() {
939        assert!(is_hop_by_hop("Connection"));
940        assert!(is_hop_by_hop("transfer-encoding"));
941        assert!(!is_hop_by_hop("content-type"));
942    }
943
944    #[actix_web::test]
945    async fn forwards_requests_and_upstream_responses() {
946        let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
947        let address = listener.local_addr().unwrap();
948        let server = HttpServer::new(|| {
949            App::new().default_service(web::to(
950                |request: HttpRequest, body: web::Bytes| async move {
951                    HttpResponse::Created()
952                        .insert_header(("x-upstream", "yes"))
953                        .body(format!(
954                            "{} {} {}",
955                            request.method(),
956                            request.uri(),
957                            String::from_utf8_lossy(&body)
958                        ))
959                },
960            ))
961        })
962        .listen(listener)
963        .unwrap()
964        .run();
965        let server_handle = server.handle();
966        actix_web::rt::spawn(server);
967
968        let gateway = GatewayProxy::new(
969            GatewayRouter::new([route("/api", &[&format!("http://{address}")])]).unwrap(),
970        );
971        let app = actix_test::init_service(
972            App::new()
973                .app_data(web::Data::new(gateway))
974                .default_service(web::to(proxy)),
975        )
976        .await;
977        let response = actix_test::call_service(
978            &app,
979            actix_test::TestRequest::post()
980                .uri("/api/hello?x=1")
981                .set_payload("world")
982                .to_request(),
983        )
984        .await;
985
986        assert_eq!(response.status(), StatusCode::CREATED);
987        assert_eq!(response.headers().get("x-upstream").unwrap(), "yes");
988        assert_eq!(
989            actix_test::read_body(response).await,
990            "POST /api/hello?x=1 world"
991        );
992        server_handle.stop(true).await;
993    }
994
995    #[actix_web::test]
996    async fn streams_upstream_response_without_buffering_it() {
997        let upstream_listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
998        let upstream_address = upstream_listener.local_addr().unwrap();
999        let upstream = HttpServer::new(|| {
1000            App::new().default_service(web::to(|| async {
1001                let chunks = futures::stream::unfold(0, |step| async move {
1002                    match step {
1003                        0 => Some((
1004                            Ok::<_, actix_web::Error>(web::Bytes::from_static(b"first")),
1005                            1,
1006                        )),
1007                        1 => {
1008                            tokio::time::sleep(Duration::from_millis(250)).await;
1009                            Some((Ok(web::Bytes::from_static(b"second")), 2))
1010                        }
1011                        _ => None,
1012                    }
1013                });
1014                HttpResponse::Ok().streaming(chunks)
1015            }))
1016        })
1017        .listen(upstream_listener)
1018        .unwrap()
1019        .run();
1020        let upstream_handle = upstream.handle();
1021        actix_web::rt::spawn(upstream);
1022
1023        let gateway_listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
1024        let gateway_address = gateway_listener.local_addr().unwrap();
1025        let gateway = GatewayServer::new(GatewayConfig {
1026            routes: vec![route("/", &[&format!("http://{upstream_address}")])],
1027            ..GatewayConfig::default()
1028        })
1029        .unwrap()
1030        .run_on(gateway_listener)
1031        .unwrap();
1032        let gateway_handle = gateway.handle();
1033        actix_web::rt::spawn(gateway);
1034
1035        let started = tokio::time::Instant::now();
1036        let response = reqwest::get(format!("http://{gateway_address}/events"))
1037            .await
1038            .unwrap();
1039        let mut body = response.bytes_stream();
1040        assert_eq!(body.next().await.unwrap().unwrap(), "first");
1041        assert!(started.elapsed() < Duration::from_millis(200));
1042        assert_eq!(body.next().await.unwrap().unwrap(), "second");
1043
1044        gateway_handle.stop(true).await;
1045        upstream_handle.stop(true).await;
1046    }
1047
1048    #[actix_web::test]
1049    async fn shutdown_gracefully_drains_an_inflight_proxy_request() {
1050        let request_started = Arc::new(tokio::sync::Notify::new());
1051        let release_request = Arc::new(tokio::sync::Notify::new());
1052        let upstream_listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
1053        let upstream_address = upstream_listener.local_addr().unwrap();
1054        let upstream = HttpServer::new({
1055            let request_started = request_started.clone();
1056            let release_request = release_request.clone();
1057            move || {
1058                App::new()
1059                    .app_data(web::Data::new((
1060                        request_started.clone(),
1061                        release_request.clone(),
1062                    )))
1063                    .default_service(web::to(
1064                        |signals: web::Data<(
1065                            Arc<tokio::sync::Notify>,
1066                            Arc<tokio::sync::Notify>,
1067                        )>| async move {
1068                            signals.0.notify_one();
1069                            signals.1.notified().await;
1070                            HttpResponse::Ok().body("drained")
1071                        },
1072                    ))
1073            }
1074        })
1075        .listen(upstream_listener)
1076        .unwrap()
1077        .run();
1078        let upstream_handle = upstream.handle();
1079        actix_web::rt::spawn(upstream);
1080
1081        let gateway_listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
1082        let gateway_address = gateway_listener.local_addr().unwrap();
1083        let server = GatewayServer::new(GatewayConfig {
1084            routes: vec![route("/", &[&format!("http://{upstream_address}")])],
1085            shutdown_timeout_ms: 2_000,
1086            ..GatewayConfig::default()
1087        })
1088        .unwrap();
1089        let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel();
1090        let serving = actix_web::rt::spawn(async move {
1091            server
1092                .serve_on_until(gateway_listener, async {
1093                    let _ = shutdown_rx.await;
1094                })
1095                .await
1096        });
1097
1098        let request = actix_web::rt::spawn(async move {
1099            reqwest::get(format!("http://{gateway_address}/slow"))
1100                .await
1101                .unwrap()
1102                .text()
1103                .await
1104                .unwrap()
1105        });
1106        request_started.notified().await;
1107        shutdown_tx.send(()).unwrap();
1108        tokio::time::sleep(Duration::from_millis(50)).await;
1109        assert!(!serving.is_finished());
1110        release_request.notify_one();
1111
1112        assert_eq!(request.await.unwrap(), "drained");
1113        serving.await.unwrap().unwrap();
1114        upstream_handle.stop(true).await;
1115    }
1116
1117    #[test]
1118    fn rejects_invalid_routes() {
1119        assert_eq!(
1120            GatewayRoute::new("api", vec!["http://api".to_owned()]).unwrap_err(),
1121            GatewayError::InvalidPrefix("api".to_owned())
1122        );
1123        assert_eq!(
1124            GatewayRoute::new("/api", Vec::new()).unwrap_err(),
1125            GatewayError::EmptyUpstreams("/api".to_owned())
1126        );
1127    }
1128}