1mod 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#[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#[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#[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 #[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
206pub 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 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 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 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
295pub 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
314pub type GatewayMiddlewareFuture = LocalBoxFuture<'static, HttpResponse>;
316
317#[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
329pub 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#[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 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
596pub 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 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
684pub 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#[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}