diff --git a/Cargo.toml b/Cargo.toml index b85ef94..14d39df 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -32,8 +32,9 @@ required-features = ["cli"] doc = false [dependencies] -# HTTP framework -axum = { version = "0.8", features = ["macros"] } +# HTTP framework. `http2`: `serve` answers HTTP/1.1 and HTTP/2 on one port, +# so native gRPC clients share the listener with REST ones. +axum = { version = "0.8", features = ["macros", "http2"] } tower = "0.5" tower-http = { version = "0.7", features = ["cors", "trace"] } # Foundational HTTP types, used directly by the framework-agnostic embedding @@ -42,9 +43,19 @@ tower-http = { version = "0.7", features = ["cors", "trace"] } # explicit. http = "1" bytes = "1" - -# gRPC client (to upstream service) -tonic = "0.14" +# The body bound of an upstream's responses and of requests `ProxyService` +# accepts; already in the tree through hyper, axum and tonic. +http-body = "1" +# The response futures of `ProxyService`, named without boxing each one. +pin-project-lite = "0.2" + +# gRPC client (to upstream service). `tls-connect-info` gives a rustls server +# stream the connection record (client certificates) a tonic handler reads, so +# an in-process upstream behind the proxy gets it; tonic exports the type only +# with its TLS backends, so it is named through that stream (`tokio-rustls`, +# the same crate tonic uses). Neither links a crypto provider. +tonic = { version = "0.14", features = ["tls-connect-info"] } +tokio-rustls = { version = "0.26", default-features = false } tonic-health = "0.14" # Canonical google.rpc.Status / error_details descriptors (FILE_DESCRIPTOR_SET) # and the Status message used to decode `grpc-status-details-bin`, so REST error @@ -158,7 +169,8 @@ redis = ["dep:redis"] cli = ["dep:clap", "dep:tracing-subscriber", "redis"] [dev-dependencies] -tokio = { version = "1", features = ["macros", "rt-multi-thread"] } +# `test-util`: deadline tests run on a paused clock instead of waiting. +tokio = { version = "1", features = ["macros", "rt-multi-thread", "test-util"] } tower = { version = "0.5", features = ["util"] } http-body-util = "0.1" # Trailer frames for the hand-written upstream response in @@ -169,6 +181,12 @@ http-body = "1" # features would pull in aws-lc. tokio-rustls = { version = "0.26", default-features = false, features = ["tls12"] } rustls-rustcrypto = "0.0.2-alpha" +# The embedder-owned TLS server of tests/tls.rs (HTTP/1.1 and HTTP/2 on one +# connection type), and the TLS stream a tonic client dials through. +hyper-util = { version = "0.1", features = ["server-auto", "service", "tokio"] } +# An upstream that speaks gRPC-Web, the way an embedder gives it that +# protocol (tests/edge.rs). +tonic-web = "0.14" # benches/jwt_verify.rs. Plots and rayon are left out: numbers are enough. criterion = { version = "0.8", default-features = false, features = ["async_tokio", "cargo_bench_support"] } # The embedding-hooks integration test (tests/hooks.rs) writes hook impls using diff --git a/README.md b/README.md index 233324f..a7555a6 100644 --- a/README.md +++ b/README.md @@ -21,13 +21,15 @@ Works with **any** gRPC service via proto descriptor files. No code generation, - **Server-streaming** RPC → NDJSON by default, or Server-Sent Events via `Accept: text/event-stream` negotiation - **gRPC → HTTP status mapping** following the standard `google.rpc.Code` table - **Typed error details**: the upstream's `google.rpc.Status` details (`ErrorInfo`, `BadRequest`, `RetryInfo`, ...) reach the HTTP client as ProtoJSON, switchable globally and per route (see [Error responses](#error-responses)) +- **One port for REST and native gRPC**: HTTP/1.1 and HTTP/2 on the same listener, gRPC and gRPC-Web requests pass through to the upstream unchanged (an upstream that speaks gRPC-Web, e.g. behind tonic-web, answers it); behind your own TLS too, with the client's address and certificate reaching the upstream +- **In-process upstream** for embedders: transcoded calls reach your own tonic services with no socket or loopback hop (see [Library Usage](#library-usage)) - **Header forwarding** from HTTP requests to gRPC metadata (configurable allow-list) - **Context propagation**: W3C trace-context (`traceparent` forwarded or synthesized) and client deadlines (`grpc-timeout`) carried across the REST↔gRPC boundary - **Path aliasing** for route remapping (e.g. `/oauth2/*` → `/v1/oauth2/*`) - **Maintenance mode** returning 503 with a configurable exempt-path list - **Health endpoints** `/health/live`, `/health/ready` (upstream gRPC health probe), `/health/startup` - **Prometheus metrics** at `/metrics` -- **CORS** with a configurable origin allow-list +- **CORS** with a configurable origin allow-list, exposed headers and preflight cache, applied to gRPC-Web pass-through too - **Rate limiting (Shield)**: local GCRA shaper (no blocking latency) keyed by client IP, header, or validated JWT claim; named limit tiers as config data; optional async cross-instance reconciliation for an approximate fleet-wide limit (requires both the `redis` feature and a configured `sync` block) - **JWT auth**: validate `Bearer` tokens via an Ed25519 PEM key or JWKS auto-discovery, enforce per-route `require_auth` / `required_roles`, and forward claims as headers — or hand the signature check to your own verifier (a validated / FIPS module, an HSM) without changing anything else - **OIDC discovery**: serve `/.well-known/openid-configuration` and a JWKS endpoint (Ed25519) built from config, to front an identity provider @@ -66,6 +68,8 @@ log line (`RUST_LOG=info`) states the count and where it came from. listen: http: "0.0.0.0:8080" +# The gRPC service behind the proxy. Required by the standalone binary; an +# embedder with an in-process upstream leaves it out. upstream: default: "http://127.0.0.1:50051" @@ -79,10 +83,24 @@ service: cors: # Empty list = permissive CORS (dev mode, reflects any Origin). - # A non-empty list allows those exact origins; there is no "*" wildcard - # (browsers never send `Origin: *`, so listing "*" would block everything). + # A non-empty list allows those exact origins, with credentials; the + # preflight echoes the methods and headers the browser asks for. There is + # no "*" wildcard (browsers never send `Origin: *`, so listing "*" would + # block everything). origins: [] # e.g. origins: ["https://app.example.com", "https://admin.example.com"] + # Response headers a browser script may read, on top of grpc-status, + # grpc-message, grpc-status-details-bin and the rate-limit headers: + # typically upstream metadata forwarded as a header. + expose_headers: [] + # e.g. expose_headers: ["x-request-id"] + # How long a browser caches a preflight answer (seconds). Unset: the + # browser's default. + # max_age_secs: 600 + # Apply this policy to gRPC-Web calls passed through to the upstream and to + # their preflights. Turn off only when the upstream sets CORS on gRPC-Web + # itself: its preflights then reach the upstream too. + grpc_web: true # Optional: path aliases (rewrite before routing) aliases: @@ -530,7 +548,7 @@ details are malformed is dropped along with it (see [Error responses](#error-responses)). Browsers read only [CORS-safelisted](https://fetch.spec.whatwg.org/#cors-safelisted-response-header-name) response headers plus the exposed ones, so a browser client that must read a -forwarded header needs a CORS setup that exposes it. +forwarded header needs it listed in `cors.expose_headers`. **Status from `x-http-code`.** On a successful unary call, the response metadata `x-http-code` (grpc-gateway's convention) sets the HTTP status: one @@ -613,7 +631,117 @@ async fn main() -> anyhow::Result<()> { } ``` -Or build the axum `Router` yourself for custom serving / embedding: +`serve` answers HTTP/1.1 and HTTP/2 on one port: REST requests go to the +proxy's routes, and requests with a gRPC or gRPC-Web content type go to the +upstream as they arrived, so native gRPC clients can use the same address. + +### Your own gRPC services as the upstream + +A gRPC service that embeds the proxy to add REST (a forward-auth decision +service, an API that also speaks gRPC) hands its own services to the proxy +instead of an address. Transcoded calls then reach them in process: no +socket, no loopback connection, no second HTTP/2 round, and they pass through +the service's whole tonic stack (interceptors, layers) like a native gRPC +call. `Request::remote_addr` in a handler gives the HTTP client's address. + +```rust +use structured_proxy::ProxyServer; + +# async fn run() -> anyhow::Result<()> { +// Your services, exactly as you would give them to tonic's server. +let grpc = tonic::service::Routes::default(); // .add_service(MyServer::new(...)) + +// No `upstream:` in the config: the upstream is `grpc`. +let proxy = ProxyServer::from_file(std::path::Path::new("my-service.yaml"))? + .service(grpc)?; + +// REST and native gRPC on one port. +let listener = tokio::net::TcpListener::bind("0.0.0.0:8080").await?; +structured_proxy::serve(listener, proxy).await?; +# Ok(()) +# } +``` + +`ProxyServer::service` takes any gRPC tower service +(`structured_proxy::upstream::Upstream`): `tonic::service::Routes`, a remote +`tonic::transport::Channel` (what `ProxyServer::upstream` builds from the +config), or anything else that speaks gRPC over `http` types. The result is a +tower service, so it can also run on a server of your own. + +### Behind your own TLS + +`serve` speaks cleartext. For TLS, run the service on your own acceptor: one +`ProxyService::for_connection` call per accepted connection tells the proxy +who is on the other end, so its middleware sees the client's address and a +tonic handler in process reads it with `Request::remote_addr`, and the client +certificate with `Request::peer_certs` (mTLS), for native and transcoded calls +alike. Advertise `h2` next to `http/1.1` in ALPN so gRPC clients get HTTP/2. + +```rust +use std::sync::Arc; +use hyper_util::rt::{TokioExecutor, TokioIo}; +use hyper_util::server::conn::auto::Builder; +use hyper_util::service::TowerToHyperService; +use structured_proxy::{ConnectionInfo, ProxyServer}; +use tonic::transport::server::Connected; + +# async fn run(mut tls: rustls::ServerConfig, grpc: tonic::service::Routes) -> anyhow::Result<()> { +let proxy = ProxyServer::from_file(std::path::Path::new("my-service.yaml"))?.service(grpc)?; +tls.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()]; +let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(tls)); +let listener = tokio::net::TcpListener::bind("0.0.0.0:8443").await?; +loop { + let (tcp, _) = listener.accept().await?; + let (acceptor, proxy) = (acceptor.clone(), proxy.clone()); + tokio::spawn(async move { + // The handshake runs in the connection's task, so a slow client + // does not hold up the others. + let Ok(stream) = acceptor.accept(tcp).await else { return }; + let service = proxy.for_connection(ConnectionInfo::tls(stream.connect_info())); + let served = Builder::new(TokioExecutor::new()) + .serve_connection(TokioIo::new(stream), TowerToHyperService::new(service)) + .await; + if let Err(error) = served { + tracing::debug!(%error, "connection ended"); + } + }); +} +# } +``` + +A plain TCP connection passes `stream.connect_info()` directly +(`ConnectionInfo` converts from tonic's `TcpConnectInfo`). + +What the proxy does not serve is not its business: a request no route matches +gets `404`, or goes to a service of yours with +`ProxyService::with_fallback(my_axum_app)`. The proxy's middleware (CORS, +maintenance, rate limits, auth) sees neither those requests nor native gRPC +ones; they reach your service untouched. + +gRPC-Web requests pass through unchanged as well, so the upstream answers +them in that protocol: wrap your services in tonic-web's layer +(`tower::ServiceBuilder::new().layer(tonic_web::GrpcWebLayer::new()).service(grpc)`), +binary and text gRPC-Web alike. When the upstream cannot take a call at all, +the proxy's own error answer keeps the request's protocol. Browsers get the +proxy's CORS policy on these calls, the same one their preflight got +(`cors.grpc_web`, on by default); a gRPC-Web client reads `grpc-status`, +`grpc-message` and `grpc-status-details-bin`, which are always exposed. A +browser's preflight for a gRPC-Web call (one announcing `x-grpc-web`) goes +where the call goes: the proxy answers it under that policy, a fallback never +sees it, and with `cors.grpc_web: false` it reaches the upstream, whose own +CORS policy then covers preflight and call alike. + +**Deadlines.** Every call waits at most five seconds for the upstream's +response headers, or less when the client's `grpc-timeout` says so; after that +the client gets `504` `DEADLINE_EXCEEDED`. The proxy enforces this itself, in +process and remote alike. The client's `grpc-timeout` travels to the upstream; +the five-second default does not, so an upstream that applies `grpc-timeout` +to a whole call does not cut a long server stream short. + +### Merging into an axum application + +`ProxyServer::router` returns the proxy's HTTP routes in front of the +configured upstream address, to serve or to merge into your own axum `Router`: ```rust use std::path::Path; diff --git a/src/config.rs b/src/config.rs index 8da5d7e..c033528 100644 --- a/src/config.rs +++ b/src/config.rs @@ -15,8 +15,12 @@ use std::path::PathBuf; /// `#[non_exhaustive]` instead, since those are deserialized, not hand-built. #[derive(Debug, Clone, Deserialize)] pub struct ProxyConfig { - /// Upstream gRPC service(s). - pub upstream: UpstreamConfig, + /// The remote gRPC upstream. Required by the standalone proxy and by + /// [`ProxyServer::upstream`](crate::ProxyServer::upstream); an embedder + /// whose upstream is in process + /// ([`ProxyServer::service`](crate::ProxyServer::service)) leaves it out. + #[serde(default)] + pub upstream: Option, /// Proto descriptor sources. #[serde(default, deserialize_with = "deserialize_descriptor_sources")] @@ -998,12 +1002,39 @@ impl Default for MaintenanceConfig { } /// CORS configuration. -#[derive(Debug, Clone, Default, Deserialize)] +#[derive(Debug, Clone, Deserialize)] #[non_exhaustive] pub struct CorsConfig { /// Allowed origins. Empty = permissive (dev mode). #[serde(default)] pub origins: Vec, + /// Response headers a browser script may read on top of the ones the + /// proxy always exposes (`grpc-status`, `grpc-message`, + /// `grpc-status-details-bin` and the rate-limit headers): typically + /// upstream metadata forwarded as a header, such as `x-request-id`. + #[serde(default)] + pub expose_headers: Vec, + /// How long a browser may cache a preflight answer, in seconds. Unset, + /// the browser's own default applies. + #[serde(default)] + pub max_age_secs: Option, + /// Apply this policy to gRPC-Web calls that pass through to the upstream + /// and to their preflights, so a browser gets one policy for both (on by + /// default). Off only for an upstream that sets CORS on gRPC-Web itself: + /// its preflights then reach the upstream as well. + #[serde(default = "default_true")] + pub grpc_web: bool, +} + +impl Default for CorsConfig { + fn default() -> Self { + Self { + origins: Vec::new(), + expose_headers: Vec::new(), + max_age_secs: None, + grpc_web: true, + } + } } /// Logging configuration. diff --git a/src/config/tests.rs b/src/config/tests.rs index a3d8db0..bfa7bf5 100644 --- a/src/config/tests.rs +++ b/src/config/tests.rs @@ -7,7 +7,7 @@ upstream: default: "grpc://localhost:4180" "#; let config: ProxyConfig = serde_yaml::from_str(yaml).unwrap(); - assert_eq!(config.upstream.default, "grpc://localhost:4180"); + assert_eq!(config.upstream.unwrap().default, "grpc://localhost:4180"); assert_eq!(config.listen.http, "0.0.0.0:8080"); assert_eq!(config.service.name, "structured-proxy"); assert_eq!(config.streaming.sse_keep_alive_secs, 15); @@ -16,6 +16,44 @@ upstream: assert!(config.shield.is_none()); } +#[test] +fn cors_defaults_cover_grpc_web_and_nothing_extra() { + // Absent `cors:` and an empty one agree: permissive origins, gRPC-Web + // included, no extra exposed headers, the browser's own preflight cache. + for yaml in ["service:\n name: demo\n", "cors: {}\n"] { + let cors = serde_yaml::from_str::(yaml).unwrap().cors; + assert!(cors.origins.is_empty(), "{yaml}"); + assert!(cors.expose_headers.is_empty(), "{yaml}"); + assert_eq!(cors.max_age_secs, None, "{yaml}"); + assert!(cors.grpc_web, "{yaml}"); + } +} + +#[test] +fn cors_settings_are_read() { + let yaml = "cors:\n origins: [\"https://app.example\"]\n expose_headers: [\"x-request-id\"]\n max_age_secs: 600\n grpc_web: false\n"; + let cors = serde_yaml::from_str::(yaml).unwrap().cors; + assert_eq!(cors.origins, ["https://app.example"]); + assert_eq!(cors.expose_headers, ["x-request-id"]); + assert_eq!(cors.max_age_secs, Some(600)); + assert!(!cors.grpc_web); +} + +#[test] +fn upstream_is_optional() { + // An embedder whose upstream is in process names no address. + let config: ProxyConfig = serde_yaml::from_str("service:\n name: demo\n").unwrap(); + assert!(config.upstream.is_none()); +} + +#[test] +fn upstream_without_its_address_is_rejected() { + // A present `upstream:` block must say where: an empty one is a mistake, + // not a request for an in-process upstream. + let err = serde_yaml::from_str::("upstream: {}\n").unwrap_err(); + assert!(err.to_string().contains("default"), "{err}"); +} + #[test] fn jwks_max_age_defaults_and_overrides() { let yaml = "upstream:\n default: \"grpc://x:1\"\nauth:\n mode: jwt\n jwt:\n jwks_uri: \"https://idp/jwks\"\n"; @@ -231,7 +269,10 @@ forwarded_headers: - "x-request-id" "#; let config: ProxyConfig = serde_yaml::from_str(yaml).unwrap(); - assert_eq!(config.upstream.default, "grpc://sid-identity:4180"); + assert_eq!( + config.upstream.as_ref().unwrap().default, + "grpc://sid-identity:4180" + ); assert_eq!(config.listen.http, "0.0.0.0:9090"); assert_eq!(config.service.name, "sid-proxy"); assert_eq!(config.aliases.len(), 1); diff --git a/src/embed.rs b/src/embed.rs index 35b4b97..f52ebc6 100644 --- a/src/embed.rs +++ b/src/embed.rs @@ -20,7 +20,6 @@ use axum::routing::{get, on, MethodFilter, MethodRouter}; use axum::{Json, Router}; use crate::hooks::{AuthDecider, Decision, ExtraRoute, OidcBackend, RequestParts, RouteRequest}; -use crate::ProxyState; /// Cap on the body an extra-route handler will buffer (16 MiB). Extra routes are /// a stateless escape hatch, not a bulk-upload path; a bounded buffer keeps a @@ -115,7 +114,10 @@ pub(crate) async fn verify_via_decider( } /// Routes for the stateless OIDC surface supplied by an [`OidcBackend`]. -pub(crate) fn oidc_backend_routes(backend: Arc) -> Router { +pub(crate) fn oidc_backend_routes(backend: Arc) -> Router +where + S: Clone + Send + Sync + 'static, +{ let mut router = Router::new(); // Static metadata documents (openid-configuration, provider-specific docs). @@ -180,10 +182,13 @@ pub(crate) fn oidc_backend_routes(backend: Arc) -> Router Router { +pub(crate) fn extra_routes_router(routes: &[ExtraRoute]) -> Router +where + S: Clone + Send + Sync + 'static, +{ use std::collections::HashMap; - let mut by_path: HashMap> = HashMap::new(); + let mut by_path: HashMap> = HashMap::new(); for route in routes { let Ok(filter) = MethodFilter::try_from(route.method.clone()) else { tracing::warn!( diff --git a/src/lib.rs b/src/lib.rs index 0705ff6..7163e02 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -60,14 +60,17 @@ mod embed; pub mod hooks; pub mod oidc; pub mod openapi; +pub mod service; pub mod shield; mod tls; pub mod transcode; +pub mod upstream; /// Settle the process-wide JWT crypto provider. See /// [`install_default_crypto_provider`] for when a call is needed. #[cfg(feature = "builtin_jwt")] pub use auth::crypto::install_default_crypto_provider; +pub use service::{serve, ConnectionInfo, ProxyService}; use axum::extract::State; use axum::http::{Request, StatusCode}; @@ -77,37 +80,32 @@ use axum::routing::get; use axum::{Json, Router}; use prost_reflect::DescriptorPool; use std::net::SocketAddr; -use tower_http::cors::{AllowOrigin, CorsLayer}; +use tower_http::cors::{AllowHeaders, AllowMethods, AllowOrigin, CorsLayer}; use tower_http::trace::TraceLayer; use std::sync::Arc; use config::{DescriptorSource, ProxyConfig}; use hooks::{AuthDecider, ExtraRoute, OidcBackend, TokenVerifier}; +use upstream::Upstream; -/// Shared state for all proxy handlers. +/// What the proxy's handlers share. Every request extracts its own clone, so +/// it holds the upstream handle and reference-counted settings only. #[derive(Clone, Debug)] -pub struct ProxyState { - /// Service name from config. - pub service_name: String, - /// gRPC upstream address. - pub grpc_upstream: String, - /// Lazy gRPC channel to upstream service. - pub grpc_channel: tonic::transport::Channel, - /// Maintenance mode active. - pub maintenance_mode: bool, - /// Maintenance exempt path patterns. - pub maintenance_exempt: Vec, - /// Maintenance message. - pub maintenance_message: String, +pub(crate) struct ProxyState { + /// The gRPC service transcoded calls and readiness probes go to. + pub(crate) upstream: U, /// Headers to forward from HTTP to gRPC. - pub forwarded_headers: Vec, - /// Metrics namespace (derived from service name). - pub metrics_namespace: String, - /// Path class patterns for metrics. - pub metrics_classes: Vec, + pub(crate) forwarded_headers: Arc<[String]>, /// SSE keep-alive interval (seconds) for server-streaming responses. - pub sse_keep_alive_secs: u64, + pub(crate) sse_keep_alive_secs: u64, +} + +/// Maintenance mode: every request outside the exempt paths gets a `503`. +#[derive(Debug)] +struct Maintenance { + exempt: Vec, + message: String, } /// Universal proxy server. @@ -415,22 +413,86 @@ impl ProxyServer { Ok(routes) } - /// Build the axum router with all endpoints. + /// A lazy channel to the configured upstream address (`upstream.default`), + /// the upstream of the standalone proxy. It connects on first use, giving + /// up after five seconds. + /// + /// # Errors + /// + /// No upstream address is configured, or it is not a valid URI. + pub fn upstream(&self) -> anyhow::Result { + let Some(upstream) = &self.config.upstream else { + anyhow::bail!("no gRPC upstream address is configured (upstream.default)"); + }; + Ok( + tonic::transport::Channel::from_shared(upstream.default.clone()) + .map_err(|e| anyhow::anyhow!("invalid gRPC upstream URL: {e}"))? + .connect_timeout(std::time::Duration::from_secs(5)) + .connect_lazy(), + ) + } + + /// The proxy's HTTP routes in front of the configured upstream address + /// (see [`upstream`](Self::upstream)), as an axum `Router` to serve or to + /// merge into an axum application. It answers HTTP only; use + /// [`service`](Self::service) for native gRPC on the same listener, or for + /// an upstream in process. + /// + /// # Errors + /// + /// No valid upstream address, or a configuration [`service`](Self::service) + /// rejects. pub fn router(&self) -> anyhow::Result { + let (router, _) = self.routes(self.upstream()?)?; + Ok(router) + } + + /// The whole proxy as one tower service in front of `upstream`: native + /// gRPC requests reach `upstream` unchanged, every other request the + /// proxy's routes, whose transcoded calls go to `upstream` too. See + /// [`ProxyService`]. + /// + /// `upstream` is any gRPC service ([`upstream::Upstream`]): the embedder's + /// own tonic services in process, with no socket between them and the + /// proxy, or a remote [`Channel`](tonic::transport::Channel) such as + /// [`upstream`](Self::upstream). + /// + /// # Errors + /// + /// An invalid configuration (see [`ProxyConfig::validate`]), descriptors + /// that cannot be loaded, a route mounted twice, or a malformed auth, + /// authz, shield or OIDC section. + /// + /// # Examples + /// + /// ``` + /// use structured_proxy::ProxyServer; + /// + /// # fn build() -> anyhow::Result<()> { + /// let grpc = tonic::service::Routes::default(); // add your services here + /// let service = ProxyServer::from_yaml_str("service:\n name: demo\n")?.service(grpc)?; + /// # let _ = service; + /// # Ok(()) + /// # } + /// # build().unwrap(); + /// ``` + pub fn service(&self, upstream: U) -> anyhow::Result> { + let (routes, cors) = self.routes(upstream.clone())?; + // The routes answer a browser's preflight for gRPC-Web too, so its + // call carries the same policy unless the upstream sets its own. + let grpc_web_cors = self.config.cors.grpc_web.then_some(cors); + Ok(ProxyService::new(upstream, routes, grpc_web_cors)) + } + + /// Build the axum router with all endpoints, calling `upstream`, and the + /// CORS policy it answers under. + fn routes(&self, upstream: U) -> anyhow::Result<(Router, CorsLayer)> { // Enforce cross-field invariants on the embedded path too, where the // config is built directly instead of through `from_yaml_str`. self.config.validate()?; let pool = self.load_descriptors()?; - let grpc_upstream = self.config.upstream.default.clone(); - let grpc_channel = tonic::transport::Channel::from_shared(grpc_upstream.clone()) - .map_err(|e| anyhow::anyhow!("invalid gRPC upstream URL: {}", e))? - .connect_timeout(std::time::Duration::from_secs(5)) - .timeout(std::time::Duration::from_secs(5)) - .connect_lazy(); - let service_name = self.config.service.name.clone(); - let metrics_namespace = service_name.replace('-', "_"); // The verify path that is actually mounted (branch-correct), if any. let verify_path = self.mounted_verify_path(); @@ -508,19 +570,12 @@ impl ProxyServer { } let state = ProxyState { - service_name: service_name.clone(), - grpc_upstream, - grpc_channel, - maintenance_mode: self.config.maintenance.enabled, - maintenance_exempt, - maintenance_message: self.config.maintenance.message.clone(), - forwarded_headers: self.config.forwarded_headers.clone(), - metrics_namespace, - metrics_classes: self.config.metrics_classes.clone(), + upstream, + forwarded_headers: self.config.forwarded_headers.as_slice().into(), sse_keep_alive_secs: self.config.streaming.sse_keep_alive_secs, }; - let cors = self.build_cors(); + let cors = self.build_cors()?; // Build transcoding routes from descriptor pool. let mut transcode_routes = @@ -574,15 +629,18 @@ impl ProxyServer { .route(&health.live_path, get(|| async { StatusCode::OK })) .route( &health.ready_path, - get(|State(state): State| async move { + get(|State(state): State>| async move { let mut client = - tonic_health::pb::health_client::HealthClient::new(state.grpc_channel); - match client - .check(tonic_health::pb::HealthCheckRequest { - service: String::new(), - }) + tonic_health::pb::health_client::HealthClient::new(state.upstream); + let check = client.check(tonic_health::pb::HealthCheckRequest { + service: String::new(), + }); + // An upstream that does not answer in time is not ready. + match tokio::time::timeout(transcode::UPSTREAM_DEADLINE, check) .await - { + .unwrap_or_else(|_| { + Err(tonic::Status::deadline_exceeded("health check timed out")) + }) { Ok(resp) => { let status = resp.into_inner().status; if status @@ -726,20 +784,30 @@ impl ProxyServer { )); } - let router = router - .layer(axum::middleware::from_fn_with_state( - state.clone(), + // Mounted only while maintenance is on, so normal traffic pays nothing + // for it. + if self.config.maintenance.enabled { + let maintenance = Arc::new(Maintenance { + exempt: maintenance_exempt, + message: self.config.maintenance.message.clone(), + }); + router = router.layer(axum::middleware::from_fn_with_state( + maintenance, maintenance_middleware, - )) - .layer(TraceLayer::new_for_http()); + )); + } + let router = router.layer(TraceLayer::new_for_http()); // Outermost: wraps every enforcement layer so short-circuited // responses keep CORS headers, and answers preflight before auth. - let router = cors::layer(router, cors).with_state(state); + let router = cors::layer(router, cors.clone()).with_state(state); - Ok(router) + Ok((router, cors)) } - fn build_openapi_routes(&self, pool: &DescriptorPool) -> Router { + fn build_openapi_routes(&self, pool: &DescriptorPool) -> Router + where + S: Clone + Send + Sync + 'static, + { let openapi_config = match &self.config.openapi { Some(cfg) if cfg.enabled => cfg, _ => return Router::new(), @@ -784,44 +852,82 @@ impl ProxyServer { ) } - fn build_cors(&self) -> CorsLayer { - if self.config.cors.origins.is_empty() { + /// The CORS policy of `cors:`. + /// + /// # Errors + /// + /// An origin that is not a header value, or an `expose_headers` entry that + /// is not a header name: dropping either would quietly narrow the policy. + fn build_cors(&self) -> anyhow::Result { + let config = &self.config.cors; + // Checked in both modes, so a typo fails the same way whether or not + // origins are set. + let configured = config + .expose_headers + .iter() + .map(|name| { + http::HeaderName::from_bytes(name.as_bytes()).map_err(|_| { + anyhow::anyhow!("cors.expose_headers entry {name:?} is not a header name") + }) + }) + .collect::>>()?; + let layer = if config.origins.is_empty() { tracing::warn!("CORS origins not set — using permissive CORS (dev mode)"); + // Exposes every header already. CorsLayer::permissive() } else { - let origins: Vec<_> = self - .config - .cors + let origins = config .origins .iter() - .filter_map(|o| o.parse().ok()) - .collect(); + .map(|origin| { + http::HeaderValue::from_str(origin) + .map_err(|_| anyhow::anyhow!("cors.origins entry {origin:?} is not valid")) + }) + .collect::>>()?; + let exposed = [ + // What a gRPC-Web client reads its status and details from. + http::HeaderName::from_static("grpc-status"), + http::HeaderName::from_static("grpc-message"), + http::HeaderName::from_static("grpc-status-details-bin"), + // Let browser clients read the rate-limit budget and back off. + http::HeaderName::from_static("ratelimit-limit"), + http::HeaderName::from_static("ratelimit-remaining"), + http::HeaderName::from_static("ratelimit-reset"), + http::HeaderName::from_static("retry-after"), + ] + .into_iter() + .chain(configured); + // With credentials the Fetch standard (§3.2.5) forbids `*` for + // methods and headers, so the preflight's own request is echoed + // back instead: what the browser asked for, from an allowed origin. CorsLayer::new() .allow_origin(AllowOrigin::list(origins)) - .allow_methods(tower_http::cors::Any) - .allow_headers(tower_http::cors::Any) + .allow_methods(AllowMethods::mirror_request()) + .allow_headers(AllowHeaders::mirror_request()) .allow_credentials(true) - .expose_headers([ - "grpc-status".parse().unwrap(), - "grpc-message".parse().unwrap(), - // Let browser clients read the rate-limit budget and back off. - "ratelimit-limit".parse().unwrap(), - "ratelimit-remaining".parse().unwrap(), - "ratelimit-reset".parse().unwrap(), - "retry-after".parse().unwrap(), - ]) - } + .expose_headers(exposed.collect::>()) + }; + Ok(match config.max_age_secs { + Some(secs) => layer.max_age(std::time::Duration::from_secs(secs)), + None => layer, + }) } - /// Start serving on configured address. + /// Serve the proxy on the configured listen address in front of the + /// configured upstream address: REST and native gRPC on one port (see + /// [`service`](Self::service) and [`serve`]). + /// + /// # Errors + /// + /// What [`upstream`](Self::upstream) and [`service`](Self::service) + /// reject, an invalid listen address, or a listener that fails. pub async fn serve(&self) -> anyhow::Result<()> { - let router = self.router()?; - let app = router.into_make_service_with_connect_info::(); + let service = self.service(self.upstream()?)?; let addr: SocketAddr = self.config.listen.http.parse()?; let listener = tokio::net::TcpListener::bind(addr).await?; tracing::info!("{} listening on {}", self.config.service.name, addr); - axum::serve(listener, app).await?; + serve(listener, service).await?; Ok(()) } } @@ -846,126 +952,51 @@ fn normalize_route_shape(path: &str) -> String { .join("/") } +impl Maintenance { + /// Whether `path` stays reachable: an exact exempt path, or `prefix` and + /// what lies below it for a `prefix/**` one (a sibling that only shares + /// the prefix, `/healthz` for `/health/**`, does not). + fn exempts(&self, path: &str) -> bool { + self.exempt + .iter() + .any(|pattern| match pattern.strip_suffix("/**") { + Some(prefix) => path + .strip_prefix(prefix) + .is_some_and(|rest| rest.is_empty() || rest.starts_with('/')), + None => path == pattern, + }) + } +} + /// Maintenance mode middleware. async fn maintenance_middleware( - State(state): State, + State(maintenance): State>, request: Request, next: Next, ) -> Response { - if state.maintenance_mode { - let path = request.uri().path(); - let exempt = state.maintenance_exempt.iter().any(|pattern| { - if pattern.ends_with("/**") { - let prefix = &pattern[..pattern.len() - 3]; - path.starts_with(prefix) - } else { - path == pattern - } - }); - if !exempt { - return ( - StatusCode::SERVICE_UNAVAILABLE, - [("retry-after", "300")], - state.maintenance_message.clone(), - ) - .into_response(); - } + if maintenance.exempts(request.uri().path()) { + return next.run(request).await; } - next.run(request).await + ( + StatusCode::SERVICE_UNAVAILABLE, + [("retry-after", "300")], + maintenance.message.clone(), + ) + .into_response() } -/// Create a lazy gRPC channel for testing (connects to nowhere). +/// A [`ProxyState`] for tests whose routers never call the upstream: a lazy +/// channel to a port nothing listens on. #[cfg(test)] -pub(crate) fn test_channel() -> tonic::transport::Channel { - tonic::transport::Channel::from_static("http://127.0.0.1:1") - .connect_timeout(std::time::Duration::from_millis(100)) - .connect_lazy() -} - -/// A minimal [`ProxyState`] for tests that only need a state to satisfy a -/// `Router` (the hook routers do not read it). -#[cfg(test)] -pub(crate) fn test_state() -> ProxyState { +pub(crate) fn test_state() -> ProxyState { ProxyState { - service_name: "test".into(), - grpc_upstream: "http://127.0.0.1:1".into(), - grpc_channel: test_channel(), - maintenance_mode: false, - maintenance_exempt: vec![], - maintenance_message: String::new(), - forwarded_headers: vec![], - metrics_namespace: "test".into(), - metrics_classes: vec![], + upstream: tonic::transport::Channel::from_static("http://127.0.0.1:1") + .connect_timeout(std::time::Duration::from_millis(100)) + .connect_lazy(), + forwarded_headers: Arc::from([]), sse_keep_alive_secs: 15, } } #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn normalize_route_shape_collapses_param_names() { - // Same shape, different param names → same key. - assert_eq!( - normalize_route_shape("/v1/x/{profile_id}"), - normalize_route_shape("/v1/x/{id}") - ); - // Wildcard vs named capture stay distinct; literals are untouched. - assert_eq!(normalize_route_shape("/a/{p}/b"), "/a/{}/b"); - assert_eq!(normalize_route_shape("/a/{*rest}"), "/a/{*}"); - assert_ne!( - normalize_route_shape("/a/{p}"), - normalize_route_shape("/a/b") - ); - } - - #[test] - fn test_minimal_config_server() { - let yaml = r#" -upstream: - default: "http://127.0.0.1:50051" -"#; - let config: ProxyConfig = serde_yaml::from_str(yaml).unwrap(); - let server = ProxyServer::from_config(config); - assert!(server.descriptor_pool.is_none()); - } - - #[tokio::test] - async fn test_maintenance_exempt_matching() { - let state = ProxyState { - service_name: "test".into(), - grpc_upstream: "http://localhost:50051".into(), - grpc_channel: test_channel(), - maintenance_mode: true, - maintenance_exempt: vec![ - "/health/**".into(), - "/.well-known/**".into(), - "/metrics".into(), - ], - maintenance_message: "Down".into(), - forwarded_headers: vec![], - metrics_namespace: "test".into(), - metrics_classes: vec![], - sse_keep_alive_secs: 15, - }; - - let check = |path: &str| -> bool { - state.maintenance_exempt.iter().any(|pattern| { - if pattern.ends_with("/**") { - let prefix = &pattern[..pattern.len() - 3]; - path.starts_with(prefix) - } else { - path == pattern - } - }) - }; - - assert!(check("/health")); - assert!(check("/health/ready")); - assert!(check("/.well-known/openid-configuration")); - assert!(check("/metrics")); - assert!(!check("/v1/auth/login")); - assert!(!check("/oauth2/token")); - } -} +mod tests; diff --git a/src/main.rs b/src/main.rs index d187532..c7ea721 100644 --- a/src/main.rs +++ b/src/main.rs @@ -41,10 +41,11 @@ fn main() -> anyhow::Result<()> { let (rt, source) = file.runtime.build().context("starting the async runtime")?; let config = server.config(); + let upstream = config.upstream.as_ref().map_or("", |u| u.default.as_str()); tracing::info!( service = %config.service.name, listen = %config.listen.http, - upstream = %config.upstream.default, + upstream = %upstream, descriptors = config.descriptors.len(), worker_threads = rt.metrics().num_workers(), worker_threads_from = %source, diff --git a/src/service.rs b/src/service.rs new file mode 100644 index 0000000..8a44832 --- /dev/null +++ b/src/service.rs @@ -0,0 +1,415 @@ +//! The proxy as one tower service: native gRPC to the upstream, everything +//! else to the proxy's routes, and what no route answers to a fallback. + +use std::convert::Infallible; +use std::future::Future; +use std::net::SocketAddr; +use std::pin::Pin; +use std::sync::Arc; +use std::task::{Context, Poll}; + +use axum::extract::connect_info::ConnectInfo; +use axum::response::IntoResponse; +use axum::routing::future::RouteFuture; +use axum::serve::IncomingStream; +use bytes::Bytes; +use pin_project_lite::pin_project; +use rustls::pki_types::CertificateDer; +use tonic::transport::server::{Connected, TcpConnectInfo}; +use tower::{Layer, Service}; +use tower_http::cors::{Cors, CorsLayer}; + +/// tonic's `TlsConnectInfo`, the record its TLS server puts on +/// a request and `Request::peer_certs` reads. tonic exports the name only with +/// a TLS backend feature, so it is reached through the stream it describes. +pub type TlsConnectInfo = + as Connected>::ConnectInfo; + +use crate::upstream::{ + grpc_protocol, is_grpc_web_preflight, BoxError, GrpcProtocol, PassThrough, Upstream, +}; + +/// The proxy as a tower service, built by +/// [`ProxyServer::service`](crate::ProxyServer::service). +/// +/// A request with a gRPC or gRPC-Web content type goes to the upstream as it +/// arrived, so one listener carries REST and native gRPC; gRPC-Web is the +/// upstream's to translate (tonic-web's `GrpcWebLayer` around its services), +/// the proxy passes protocols through rather than converting them. A gRPC-Web +/// answer gets the proxy's CORS policy (`cors.grpc_web`), since the proxy +/// answers the browser's preflight for it. Every other request goes to the +/// proxy's routes (transcoded RPCs, health, metrics, OpenAPI, OIDC, +/// forward-auth, extra routes) behind the proxy's middleware. A request no +/// route matches is answered `404`, or handed to the service set with +/// [`with_fallback`](Self::with_fallback), untouched by that middleware. +/// +/// Serve it with [`serve`], or hand it to any server that takes a tower +/// service of `http` types: your own TLS, a Unix socket, an existing hyper or +/// axum server. Native gRPC needs HTTP/2 on that server (ALPN `h2` next to +/// `http/1.1` behind TLS). Such a server tells the proxy which connection a +/// request came on with [`for_connection`](Self::for_connection). +/// +/// # Examples +/// +/// ``` +/// use structured_proxy::ProxyServer; +/// +/// # fn build() -> anyhow::Result<()> { +/// // The embedder's own tonic services, called in process. +/// let grpc = tonic::service::Routes::default(); +/// let service = ProxyServer::from_yaml_str("service:\n name: demo\n")?.service(grpc)?; +/// # let _ = service; +/// # Ok(()) +/// # } +/// # build().unwrap(); +/// ``` +#[derive(Clone, Debug)] +pub struct ProxyService { + upstream: U, + routes: axum::Router, + /// Binary and text gRPC-Web to the upstream under the CORS policy, each + /// built once; `None` when the upstream owns CORS for gRPC-Web. + grpc_web: Option>, + /// The connection the requests arrive on, set per connection by the + /// server. + connection: Option, +} + +/// gRPC-Web pass-through, one CORS-wrapped path per encoding. +#[derive(Clone, Debug)] +struct GrpcWebCors { + web: Cors>, + web_text: Cors>, +} + +/// Hands a request to the upstream in a known protocol, as a tower service so +/// a layer can wrap it. +#[derive(Clone, Debug)] +struct Forward { + upstream: U, + protocol: GrpcProtocol, +} + +impl Service> for Forward { + type Response = http::Response; + type Error = Infallible; + type Future = PassThrough; + + #[inline] + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + // The upstream's readiness is waited for per request, on its clone. + Poll::Ready(Ok(())) + } + + fn call(&mut self, request: http::Request) -> Self::Future { + PassThrough::new(self.upstream.clone(), request, self.protocol) + } +} + +/// The connection requests arrive on, in the form a tonic server records it: +/// the TCP ends, and behind TLS the client's certificate chain. +/// +/// Built from what [`Connected::connect_info`] returns for a +/// `tokio::net::TcpStream` (`ConnectionInfo::from`) or a +/// `tokio_rustls::server::TlsStream` ([`ConnectionInfo::tls`]), so a +/// tonic handler behind the proxy reads it as it would behind tonic's own +/// server: `Request::remote_addr`, and `Request::peer_certs` for mTLS. +/// +/// [`Connected::connect_info`]: tonic::transport::server::Connected::connect_info +#[derive(Clone, Debug)] +pub struct ConnectionInfo { + tcp: TcpConnectInfo, + tls: Option, +} + +impl ConnectionInfo { + /// The client's address. + pub fn remote_addr(&self) -> Option { + self.tcp.remote_addr + } + + /// The certificate chain the client presented over TLS, if it did. + pub fn peer_certs(&self) -> Option>>> { + self.tls.as_ref().and_then(|tls| tls.peer_certs()) + } + + /// Put the connection where tonic reads it: [`TcpConnectInfo`] for + /// `Request::remote_addr` and, behind TLS, [`TlsConnectInfo`] for + /// `Request::peer_certs`. + pub(crate) fn into_tonic_extensions(self, extensions: &mut http::Extensions) { + extensions.insert(self.tcp); + if let Some(tls) = self.tls { + extensions.insert(tls); + } + } +} + +impl From for ConnectionInfo { + fn from(tcp: TcpConnectInfo) -> Self { + Self { tcp, tls: None } + } +} + +impl ConnectionInfo { + /// A TLS connection, from what `connect_info` reports for a + /// `tokio_rustls::server::TlsStream`: the TCP ends and the + /// client's certificate chain. + pub fn tls(tls: TlsConnectInfo) -> Self { + Self { + tcp: tls.get_ref().clone(), + tls: Some(tls), + } + } +} + +impl ProxyService { + /// The service over `upstream` and `routes`; gRPC-Web answers carry + /// `grpc_web_cors` when set. + pub(crate) fn new(upstream: U, routes: axum::Router, grpc_web_cors: Option) -> Self { + let grpc_web = grpc_web_cors.map(|cors| { + let forward = |protocol| Forward { + upstream: upstream.clone(), + protocol, + }; + GrpcWebCors { + web: cors.layer(forward(GrpcProtocol::Web)), + web_text: cors.layer(forward(GrpcProtocol::WebText)), + } + }); + Self { + upstream, + routes, + grpc_web, + connection: None, + } + } + + /// Hand the requests no route matches to `fallback` instead of answering + /// `404`: an embedder's own REST routes, a static site, anything that is a + /// tower service. The proxy's middleware does not see them. + /// + /// A request whose path a route answers but not with its method stays with + /// the proxy (`405`), as does every gRPC request and every browser + /// preflight for a gRPC-Web call, which follows the call it announces. + #[must_use] + pub fn with_fallback(mut self, fallback: F) -> Self + where + F: Service + Clone + Send + Sync + 'static, + F::Response: IntoResponse, + F::Future: Send + 'static, + { + self.routes = self.routes.fallback_service(fallback); + self + } + + /// This service for the requests of one connection: the proxy's + /// middleware sees its client's address (rate limits by IP, the auth + /// decider), and a tonic upstream in process reads it with + /// `Request::remote_addr`, and its TLS client certificates with + /// `Request::peer_certs`, for native and transcoded calls alike. + /// + /// A server of your own calls it once per accepted connection, with what + /// tonic's [`Connected`] trait + /// reports for the stream. [`serve`] does this itself. + /// + /// # Examples + /// + /// ```no_run + /// use structured_proxy::ProxyServer; + /// use tonic::transport::server::Connected; + /// + /// # async fn run() -> anyhow::Result<()> { + /// let proxy = ProxyServer::from_yaml_str("service:\n name: demo\n")? + /// .service(tonic::service::Routes::default())?; + /// let listener = tokio::net::TcpListener::bind("0.0.0.0:8443").await?; + /// let (tcp, _) = listener.accept().await?; + /// // After a TLS handshake, the `TlsStream` reports the client's certificates too. + /// let service = proxy.for_connection(tcp.connect_info()); + /// # let _ = service; + /// # Ok(()) + /// # } + /// ``` + /// + /// Without it, a request's peer comes from the `ConnectInfo` an outer + /// axum server recorded, when there is one. + #[must_use] + pub fn for_connection(&self, connection: impl Into) -> Self { + Self { + upstream: self.upstream.clone(), + routes: self.routes.clone(), + grpc_web: self.grpc_web.clone(), + connection: Some(connection.into()), + } + } + + /// The connection `request` came on: the one given to + /// [`for_connection`](Self::for_connection), else the peer an outer axum + /// server recorded as `ConnectInfo`, so the upstream sees the same client + /// the proxy's middleware does. + fn connection_of(&self, request: &http::Request) -> Option { + if let Some(connection) = &self.connection { + return Some(connection.clone()); + } + let ConnectInfo(remote) = request.extensions().get::>()?; + Some(ConnectionInfo::from(TcpConnectInfo { + local_addr: None, + remote_addr: Some(*remote), + })) + } +} + +impl Service> for ProxyService +where + U: Upstream, + B: http_body::Body + Send + 'static, + B::Error: Into, +{ + type Response = http::Response; + type Error = Infallible; + type Future = ResponseFuture; + + #[inline] + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + // Readiness of the upstream is waited for per request, on its clone. + Poll::Ready(Ok(())) + } + + fn call(&mut self, mut request: http::Request) -> Self::Future { + // A gRPC-Web preflight goes where the call it announces goes, so both + // get one CORS policy: the proxy's, or the upstream's when it owns + // CORS. A fallback in between would answer it with neither. + let protocol = grpc_protocol(request.headers()).or_else(|| { + is_grpc_web_preflight(request.method(), request.headers()).then_some(GrpcProtocol::Web) + }); + let inner = if let Some(protocol) = protocol { + // A native call carries its connection the way tonic's server + // hands it to a handler. + if let Some(connection) = self.connection_of(&request) { + connection.into_tonic_extensions(request.extensions_mut()); + } + let request = request.map(tonic::body::Body::new); + // A browser only speaks gRPC-Web, and the routes answered its + // preflight: its call gets the same CORS policy. + let cors = match (&mut self.grpc_web, protocol) { + (Some(cors), GrpcProtocol::Web) => Some(&mut cors.web), + (Some(cors), GrpcProtocol::WebText) => Some(&mut cors.web_text), + _ => None, + }; + match cors { + // `Forward` is always ready, and so is CORS around it. + Some(cors) => Inner::GrpcWeb { + future: cors.call(request), + }, + None => Inner::Grpc { + call: PassThrough::new(self.upstream.clone(), request, protocol), + }, + } + } else { + // The middleware reads the peer as axum's `ConnectInfo`; the + // transcoder passes the whole connection on to the upstream. + if let Some(connection) = &self.connection { + let extensions = request.extensions_mut(); + if let Some(remote) = connection.remote_addr() { + extensions.insert(ConnectInfo(remote)); + } + extensions.insert(connection.clone()); + } else if let Some(connection) = self.connection_of(&request) { + request.extensions_mut().insert(connection); + } + Inner::Routes { + future: self.routes.call(request), + } + }; + ResponseFuture { inner } + } +} + +pin_project! { + /// The response future of [`ProxyService`]. + pub struct ResponseFuture { + #[pin] + inner: Inner, + } +} + +pin_project! { + #[project = InnerProj] + enum Inner { + Routes { + #[pin] + future: RouteFuture, + }, + Grpc { + #[pin] + call: PassThrough, + }, + GrpcWeb { + #[pin] + future: tower_http::cors::ResponseFuture>, + }, + } +} + +impl Future for ResponseFuture { + type Output = Result, Infallible>; + + #[inline] + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + match self.project().inner.project() { + InnerProj::Routes { future } => future.poll(cx), + InnerProj::Grpc { call } => call.poll(cx), + InnerProj::GrpcWeb { future } => future.poll(cx), + } + } +} + +/// Serve `service` on `listener` until the listener fails: cleartext HTTP/1.1 +/// and HTTP/2 on the same port, so REST clients and native gRPC clients share +/// it. Each connection's service gets its peer through +/// [`ProxyService::for_connection`]. For TLS, run the service on a server of +/// your own (see [`ProxyService`]). +/// +/// # Errors +/// +/// The listener's own I/O failure. +/// +/// # Examples +/// +/// ```no_run +/// use structured_proxy::ProxyServer; +/// +/// # async fn run() -> anyhow::Result<()> { +/// let grpc = tonic::service::Routes::default(); +/// let service = ProxyServer::from_yaml_str("service:\n name: demo\n")?.service(grpc)?; +/// let listener = tokio::net::TcpListener::bind("0.0.0.0:8080").await?; +/// structured_proxy::serve(listener, service).await?; +/// # Ok(()) +/// # } +/// ``` +pub async fn serve( + listener: tokio::net::TcpListener, + service: ProxyService, +) -> std::io::Result<()> { + axum::serve(listener, PerConnection(service)).await +} + +/// Makes the [`ProxyService`] of each accepted connection. +struct PerConnection(ProxyService); + +impl Service> for PerConnection { + type Response = ProxyService; + type Error = Infallible; + type Future = std::future::Ready, Infallible>>; + + #[inline] + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, stream: IncomingStream<'_, tokio::net::TcpListener>) -> Self::Future { + std::future::ready(Ok(self.0.for_connection(stream.io().connect_info()))) + } +} + +#[cfg(test)] +mod tests; diff --git a/src/service/tests.rs b/src/service/tests.rs new file mode 100644 index 0000000..4c18c45 --- /dev/null +++ b/src/service/tests.rs @@ -0,0 +1,382 @@ +use super::*; + +use std::sync::{Arc, Mutex}; + +use axum::body::Body; +use axum::routing::get; +use tower::ServiceExt; + +/// What a [`Recorder`] upstream saw of the last request it was called with. +#[derive(Debug, Default)] +struct Seen { + path: Option, + body: Option, + remote: Option, + local: Option, +} + +/// An upstream that records its request and answers `200` with an +/// `x-upstream` header, or fails to become ready. +#[derive(Clone, Default)] +struct Recorder { + seen: Arc>, + unready: bool, +} + +impl Service> for Recorder { + type Response = http::Response; + type Error = BoxError; + type Future = Pin> + Send>>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + if self.unready { + return Poll::Ready(Err(Box::new(tonic::Status::unavailable("upstream down")))); + } + Poll::Ready(Ok(())) + } + + fn call(&mut self, request: http::Request) -> Self::Future { + let seen = self.seen.clone(); + Box::pin(async move { + let (parts, body) = request.into_parts(); + let body = http_body_util::BodyExt::collect(body) + .await + .map_err(Into::::into)? + .to_bytes(); + let connection = parts.extensions.get::(); + *seen.lock().unwrap() = Seen { + path: Some(parts.uri.path().to_owned()), + body: Some(body), + remote: connection.and_then(|c| c.remote_addr), + local: connection.and_then(|c| c.local_addr), + }; + Ok(http::Response::builder() + .header("x-upstream", "1") + .body(tonic::body::Body::empty()) + .unwrap()) + }) + } +} + +/// A proxy with one HTTP route, `GET /route`, in front of `upstream`. +fn service(upstream: Recorder) -> ProxyService { + let routes = axum::Router::new().route("/route", get(|| async { "route" })); + ProxyService::new(upstream, routes, None) +} + +fn grpc_request(path: &str) -> http::Request { + http::Request::post(path) + .header("content-type", "application/grpc") + .body(Body::from("frame")) + .unwrap() +} + +async fn body_text(response: http::Response) -> String { + let bytes = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + String::from_utf8(bytes.to_vec()).unwrap() +} + +#[tokio::test] +async fn a_grpc_request_reaches_the_upstream_unchanged() { + let upstream = Recorder::default(); + let response = service(upstream.clone()) + .oneshot(grpc_request("/pkg.Svc/Method")) + .await + .unwrap(); + assert_eq!(response.headers()["x-upstream"], "1"); + let seen = upstream.seen.lock().unwrap(); + assert_eq!(seen.path.as_deref(), Some("/pkg.Svc/Method")); + assert_eq!(seen.body.as_deref(), Some(&b"frame"[..])); +} + +#[tokio::test] +async fn a_grpc_request_on_a_route_path_still_goes_to_the_upstream() { + // The content type decides first, so a REST route (even a catch-all) can + // never take a native gRPC call. + let upstream = Recorder::default(); + let response = service(upstream.clone()) + .oneshot(grpc_request("/route")) + .await + .unwrap(); + assert_eq!(response.headers()["x-upstream"], "1"); + assert_eq!( + upstream.seen.lock().unwrap().path.as_deref(), + Some("/route") + ); +} + +#[tokio::test] +async fn an_http_request_reaches_the_routes_not_the_upstream() { + let upstream = Recorder::default(); + let response = service(upstream.clone()) + .oneshot(http::Request::get("/route").body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(response.status(), http::StatusCode::OK); + assert_eq!(body_text(response).await, "route"); + assert!(upstream.seen.lock().unwrap().path.is_none()); +} + +#[tokio::test] +async fn an_unmatched_request_is_404_without_a_fallback() { + let upstream = Recorder::default(); + let response = service(upstream.clone()) + .oneshot( + http::Request::get("/elsewhere") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), http::StatusCode::NOT_FOUND); + assert!(upstream.seen.lock().unwrap().path.is_none()); +} + +#[tokio::test] +async fn an_unmatched_request_goes_to_the_fallback() { + let fallback = axum::Router::new().route("/elsewhere", get(|| async { "fallback" })); + let response = service(Recorder::default()) + .with_fallback(fallback) + .oneshot( + http::Request::get("/elsewhere") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), http::StatusCode::OK); + assert_eq!(body_text(response).await, "fallback"); +} + +#[tokio::test] +async fn a_route_path_with_another_method_stays_with_the_proxy() { + // The path is the proxy's: a method it does not bind is its 405, not a + // request for the fallback. + let fallback = axum::Router::new().route("/route", axum::routing::post(|| async { "x" })); + let response = service(Recorder::default()) + .with_fallback(fallback) + .oneshot(http::Request::post("/route").body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(response.status(), http::StatusCode::METHOD_NOT_ALLOWED); +} + +#[tokio::test] +async fn an_upstream_that_cannot_take_the_call_answers_a_grpc_status() { + // The gRPC client gets UNAVAILABLE in a trailers-only response, not a + // dropped connection. + let upstream = Recorder { + unready: true, + ..Recorder::default() + }; + let response = service(upstream) + .oneshot(grpc_request("/pkg.Svc/Method")) + .await + .unwrap(); + assert_eq!(response.status(), http::StatusCode::OK); + assert_eq!(response.headers()["grpc-status"], "14"); + assert_eq!(response.headers()["content-type"], "application/grpc"); +} + +/// The content type the proxy answers a failed pass-through call with, for a +/// request of `content_type`. +async fn failure_content_type(content_type: &'static str) -> String { + let upstream = Recorder { + unready: true, + ..Recorder::default() + }; + let request = http::Request::post("/pkg.Svc/Method") + .header("content-type", content_type) + .body(Body::from("frame")) + .unwrap(); + let response = service(upstream).oneshot(request).await.unwrap(); + assert_eq!(response.headers()["grpc-status"], "14", "{content_type}"); + response.headers()["content-type"] + .to_str() + .unwrap() + .to_owned() +} + +#[tokio::test] +async fn a_failed_call_answers_in_the_protocol_of_the_request() { + // A gRPC-Web client reads a status only from a gRPC-Web response (gRPC + // PROTOCOL-WEB): the failure keeps the protocol and its encoding. + assert_eq!( + failure_content_type("application/grpc").await, + "application/grpc" + ); + assert_eq!( + failure_content_type("application/grpc+proto").await, + "application/grpc" + ); + assert_eq!( + failure_content_type("application/grpc-web").await, + "application/grpc-web+proto" + ); + assert_eq!( + failure_content_type("application/grpc-web+proto").await, + "application/grpc-web+proto" + ); + assert_eq!( + failure_content_type("application/grpc-web-text").await, + "application/grpc-web-text+proto" + ); + assert_eq!( + failure_content_type("application/GRPC-WEB-TEXT+proto").await, + "application/grpc-web-text+proto" + ); +} + +fn connection() -> TcpConnectInfo { + TcpConnectInfo { + local_addr: Some("10.0.0.1:8080".parse().unwrap()), + remote_addr: Some("192.0.2.7:40000".parse().unwrap()), + } +} + +#[tokio::test] +async fn a_grpc_request_carries_its_connection_to_the_upstream() { + // A tonic handler behind the proxy reads the client's address with + // `Request::remote_addr`, as it would behind tonic's own server. + let upstream = Recorder::default(); + let proxy = service(upstream.clone()).for_connection(connection()); + proxy + .oneshot(grpc_request("/pkg.Svc/Method")) + .await + .unwrap(); + let seen = upstream.seen.lock().unwrap(); + assert_eq!(seen.remote, Some("192.0.2.7:40000".parse().unwrap())); + assert_eq!(seen.local, Some("10.0.0.1:8080".parse().unwrap())); +} + +#[tokio::test] +async fn a_grpc_request_without_a_connection_carries_none() { + // A service no server told about its connection invents no peer. + let upstream = Recorder::default(); + service(upstream.clone()) + .oneshot(grpc_request("/pkg.Svc/Method")) + .await + .unwrap(); + let seen = upstream.seen.lock().unwrap(); + assert_eq!(seen.remote, None); + assert_eq!(seen.local, None); +} + +/// Routes answering with the peer the middleware sees (`/peer`) and with the +/// connection the transcoder gets (`/connection`). +fn peer_routes() -> axum::Router { + axum::Router::new() + .route( + "/peer", + get(|ConnectInfo(peer): ConnectInfo| async move { peer.to_string() }), + ) + .route( + "/connection", + get( + |connection: Option>| async move { + match connection { + Some(axum::Extension(c)) => format!("{:?}", c.remote_addr()), + None => "none".to_owned(), + } + }, + ), + ) +} + +#[tokio::test] +async fn an_http_request_carries_its_peer_for_the_middleware() { + let proxy = + ProxyService::new(Recorder::default(), peer_routes(), None).for_connection(connection()); + let response = proxy + .oneshot(http::Request::get("/peer").body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(body_text(response).await, "192.0.2.7:40000"); +} + +#[tokio::test] +async fn an_http_request_carries_its_connection_for_the_transcoder() { + let proxy = + ProxyService::new(Recorder::default(), peer_routes(), None).for_connection(connection()); + let response = proxy + .oneshot( + http::Request::get("/connection") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(body_text(response).await, "Some(192.0.2.7:40000)"); +} + +#[tokio::test] +async fn an_http_request_without_a_connection_carries_none() { + let proxy = ProxyService::new(Recorder::default(), peer_routes(), None); + let response = proxy + .oneshot( + http::Request::get("/connection") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(body_text(response).await, "none"); +} + +#[tokio::test] +async fn a_grpc_request_on_an_axum_server_carries_its_peer_to_the_upstream() { + // Hosted by an embedder's axum server that set `ConnectInfo`, with no + // `for_connection`: the upstream sees the peer the middleware sees. + let upstream = Recorder::default(); + let mut request = grpc_request("/pkg.Svc/Method"); + let peer: SocketAddr = "198.51.100.1:5000".parse().unwrap(); + request.extensions_mut().insert(ConnectInfo(peer)); + service(upstream.clone()).oneshot(request).await.unwrap(); + let seen = upstream.seen.lock().unwrap(); + assert_eq!(seen.remote, Some(peer)); + assert_eq!(seen.local, None); +} + +#[tokio::test] +async fn an_http_request_on_an_axum_server_carries_its_peer_for_the_transcoder() { + let proxy = ProxyService::new(Recorder::default(), peer_routes(), None); + let mut request = http::Request::get("/connection") + .body(Body::empty()) + .unwrap(); + request.extensions_mut().insert(ConnectInfo::( + "198.51.100.1:5000".parse().unwrap(), + )); + let response = proxy.oneshot(request).await.unwrap(); + assert_eq!(body_text(response).await, "Some(198.51.100.1:5000)"); +} + +#[tokio::test] +async fn the_connection_given_to_the_service_wins_over_an_outer_peer() { + // `for_connection` is the server's own statement of the connection. + let upstream = Recorder::default(); + let mut request = grpc_request("/pkg.Svc/Method"); + request.extensions_mut().insert(ConnectInfo::( + "198.51.100.1:5000".parse().unwrap(), + )); + service(upstream.clone()) + .for_connection(connection()) + .oneshot(request) + .await + .unwrap(); + assert_eq!( + upstream.seen.lock().unwrap().remote, + Some("192.0.2.7:40000".parse().unwrap()) + ); +} + +#[test] +fn a_tcp_connection_has_no_certificates() { + let connection = ConnectionInfo::from(connection()); + assert_eq!( + connection.remote_addr(), + Some("192.0.2.7:40000".parse().unwrap()) + ); + assert!(connection.peer_certs().is_none()); +} diff --git a/src/tests.rs b/src/tests.rs new file mode 100644 index 0000000..4186a2c --- /dev/null +++ b/src/tests.rs @@ -0,0 +1,85 @@ +use super::*; + +#[test] +fn normalize_route_shape_collapses_param_names() { + // Same shape, different param names → same key. + assert_eq!( + normalize_route_shape("/v1/x/{profile_id}"), + normalize_route_shape("/v1/x/{id}") + ); + // Wildcard vs named capture stay distinct; literals are untouched. + assert_eq!(normalize_route_shape("/a/{p}/b"), "/a/{}/b"); + assert_eq!(normalize_route_shape("/a/{*rest}"), "/a/{*}"); + assert_ne!( + normalize_route_shape("/a/{p}"), + normalize_route_shape("/a/b") + ); +} + +#[test] +fn test_minimal_config_server() { + let yaml = r#" +upstream: + default: "http://127.0.0.1:50051" +"#; + let config: ProxyConfig = serde_yaml::from_str(yaml).unwrap(); + let server = ProxyServer::from_config(config); + assert!(server.descriptor_pool.is_none()); +} + +fn maintenance() -> Maintenance { + Maintenance { + exempt: vec![ + "/health/**".into(), + "/.well-known/**".into(), + "/metrics".into(), + ], + message: "Down".into(), + } +} + +#[test] +fn maintenance_exempts_exact_paths_and_subtrees() { + let maintenance = maintenance(); + assert!(maintenance.exempts("/health")); + assert!(maintenance.exempts("/health/ready")); + assert!(maintenance.exempts("/.well-known/openid-configuration")); + assert!(maintenance.exempts("/metrics")); + assert!(!maintenance.exempts("/v1/auth/login")); + assert!(!maintenance.exempts("/oauth2/token")); + // An exact path covers nothing below it. + assert!(!maintenance.exempts("/metrics/extra")); +} + +#[test] +fn maintenance_subtree_stops_at_a_segment_boundary() { + // `/health/**` is the `/health` subtree: a sibling path that only shares + // the prefix (`/healthz`, `/health-admin`) stays behind the 503. + let maintenance = maintenance(); + assert!(!maintenance.exempts("/healthz")); + assert!(!maintenance.exempts("/health-admin/drop")); + assert!(!maintenance.exempts("/.well-knownx")); +} + +#[test] +fn no_configured_upstream_is_an_error_naming_the_key() { + // An embedder with an in-process upstream needs no address; asking for + // the remote channel without one is a startup error, not a panic. + let server = ProxyServer::from_yaml_str("service:\n name: demo\n").unwrap(); + let err = server.upstream().unwrap_err(); + assert!(err.to_string().contains("upstream.default"), "{err}"); + let Err(err) = server.router() else { + panic!("router() needs an upstream address"); + }; + assert!(err.to_string().contains("upstream.default"), "{err}"); +} + +#[test] +fn an_invalid_upstream_address_is_an_error() { + let server = ProxyServer::from_yaml_str("upstream:\n default: \"not a uri\"\n").unwrap(); + let err = server.upstream().unwrap_err(); + assert!( + err.to_string().contains("invalid gRPC upstream URL"), + "{err}" + ); +} diff --git a/src/transcode/metadata.rs b/src/transcode/metadata.rs index c6fd2fe..48665ef 100644 --- a/src/transcode/metadata.rs +++ b/src/transcode/metadata.rs @@ -312,9 +312,11 @@ fn hex(bytes: &[u8]) -> String { /// Apply a client-supplied deadline to the upstream gRPC call. /// /// Reads the gRPC-standard `grpc-timeout` header (``, unit one of -/// `H`/`M`/`S`/`m`/`u`/`n`) and sets it as the request timeout. Absent or -/// malformed values leave the channel default in place. Returns the deadline -/// that was applied, if any. +/// `H`/`M`/`S`/`m`/`u`/`n`) and sets it as the request timeout, which travels +/// to the upstream as its `grpc-timeout`. Absent or malformed values set +/// nothing, leaving the proxy's own +/// [`UPSTREAM_DEADLINE`](super::UPSTREAM_DEADLINE). Returns the deadline that +/// was applied, if any. pub fn apply_request_deadline( request: &mut tonic::Request, headers: &HeaderMap, @@ -332,7 +334,7 @@ pub fn apply_request_deadline( /// Units: `H` hours, `M` minutes, `S` seconds, `m` milliseconds, `u` /// microseconds, `n` nanoseconds. Per the gRPC wire spec the value is at most 8 /// digits. Returns `None` on a malformed value, an over-long digit run, or a -/// zero duration (which would expire the call immediately, so the channel +/// zero duration (which would expire the call immediately, so the proxy's /// default is used instead). fn parse_grpc_timeout(value: &str) -> Option { let value = value.trim(); diff --git a/src/transcode/mod.rs b/src/transcode/mod.rs index f7cde46..1453077 100644 --- a/src/transcode/mod.rs +++ b/src/transcode/mod.rs @@ -25,6 +25,7 @@ use axum::http::{HeaderMap, HeaderName, HeaderValue, Method, StatusCode}; use axum::response::sse::{Event, KeepAlive, Sse}; use axum::response::{IntoResponse, Response}; use axum::routing::{MethodFilter, MethodRouter}; +use axum::Extension; use axum::Router; use futures::{StreamExt, TryStreamExt}; use prost_reflect::{ @@ -33,30 +34,41 @@ use prost_reflect::{ }; use std::collections::HashMap; use std::sync::Arc; +use std::time::Duration; use tonic::client::Grpc; use tonic::metadata::MetadataMap; use crate::config::AliasConfig; +use crate::service::ConnectionInfo; +use crate::upstream::Upstream; use error::{ErrorDetailsPolicy, StatusDetails}; use response::UpstreamHeaders; use rule::RouteMethod; +/// How long the upstream may take to answer a call with its response headers +/// when the client's `grpc-timeout` asks for no shorter deadline. +pub const UPSTREAM_DEADLINE: Duration = Duration::from_secs(5); + /// Trait for state types that support REST→gRPC transcoding. /// /// Implement this for your application's state type to use `transcode::routes()`. -/// Provides the minimal interface needed by transcode handlers. +/// Provides the minimal interface needed by transcode handlers. Each request +/// works on its own clone of the state, so a clone should be cheap. pub trait TranscodeState: Clone + Send + Sync + 'static { - /// Lazy gRPC channel to upstream service. - fn grpc_channel(&self) -> tonic::transport::Channel; + /// The gRPC service transcoded calls go to. + type Upstream: Upstream; + /// The upstream, taken out of this request's clone of the state. + fn into_upstream(self) -> Self::Upstream; /// Headers to forward from HTTP to gRPC metadata. fn forwarded_headers(&self) -> &[String]; /// SSE keep-alive interval (seconds) for server-streaming responses. fn sse_keep_alive_secs(&self) -> u64; } -impl TranscodeState for crate::ProxyState { - fn grpc_channel(&self) -> tonic::transport::Channel { - self.grpc_channel.clone() +impl TranscodeState for crate::ProxyState { + type Upstream = U; + fn into_upstream(self) -> U { + self.upstream } fn forwarded_headers(&self) -> &[String] { &self.forwarded_headers @@ -282,7 +294,14 @@ macro_rules! endpoint { headers: HeaderMap, path_params: Path, raw_query: RawQuery, - body: Bytes| handle(state, headers, path_params, raw_query, body, entry) + connection: Option>, + body: Bytes| { + let client = Client { + headers, + connection: connection.map(|Extension(connection)| connection), + }; + handle(state, client, path_params, raw_query, body, entry) + } }}; } @@ -512,85 +531,137 @@ fn accept_range_selects_sse(range: &str) -> bool { true } +/// Who sent a transcoded request: its headers, and the connection it came +/// on when the server recorded one. +struct Client { + headers: HeaderMap, + connection: Option, +} + /// Serve one request on a transcoded route. async fn handle( State(proxy_state): State, - headers: HeaderMap, + client: Client, Path(path_params): Path, RawQuery(raw_query): RawQuery, body: Bytes, entry: Arc, ) -> Response { + let keep_alive_secs = proxy_state.sse_keep_alive_secs(); + let Client { + headers, + connection, + } = client; let prepared = prepare( - &proxy_state, + proxy_state, &headers, + connection, &path_params, raw_query.as_deref(), body, &entry, - ) - .await; - let (client, request) = match prepared { - Ok(prepared) => prepared, + ); + let call = match prepared { + Ok(call) => call, Err(rejection) => return rejection.into_response(&entry), }; if entry.streaming { - let keep_alive_secs = proxy_state.sse_keep_alive_secs(); - streaming_call(client, request, entry, wants_sse(&headers), keep_alive_secs).await + streaming_call(call, entry, wants_sse(&headers), keep_alive_secs).await } else { - unary_call(client, request, &entry).await + unary_call(call, &entry).await } } -/// Why a request ends before the upstream is called. -enum Rejection { - /// It cannot be mapped onto the RPC (`INVALID_ARGUMENT`, 400). - Unmappable(String), - /// The upstream channel is not ready (`UNAVAILABLE`, 503). - NotReady(String), -} +/// Why a request ends before the upstream is called: it cannot be mapped onto +/// the RPC (`INVALID_ARGUMENT`, 400). +struct Unmappable(String); -impl Rejection { +impl Unmappable { /// The answer, in the error body the upstream's own errors get on the route. fn into_response(self, entry: &RouteEntry) -> Response { - let status = match self { - Self::Unmappable(message) => tonic::Status::invalid_argument(message), - Self::NotReady(message) => tonic::Status::unavailable(message), - }; + let status = tonic::Status::invalid_argument(self.0); error::status_to_response_with_details(&status, entry.error_details.as_deref()) } } -/// Map the request onto the RPC's input message and get a client whose -/// channel is ready. -async fn prepare( - proxy_state: &S, +/// A call ready to be made: the upstream, the request, and how long the +/// upstream may take to answer it. +struct Call { + upstream: U, + request: tonic::Request, + deadline: Duration, +} + +impl Call { + /// Wait for the upstream to take the call, start it, and wait for its + /// response headers, all within the one deadline: an upstream under + /// backpressure that never frees a slot answers `DEADLINE_EXCEEDED` like + /// one that never answers. Every upstream gets the same deadline here, in + /// process or remote, rather than whatever its transport enforces. + async fn open( + self, + entry: &RouteEntry, + ) -> Result>, tonic::Status> { + let Self { + upstream, + request, + deadline, + } = self; + let call = async move { + let mut client = Grpc::new(upstream); + if let Err(e) = client.ready().await { + let e: crate::upstream::BoxError = e.into(); + return Err(tonic::Status::unavailable(format!( + "gRPC upstream not ready: {e}" + ))); + } + client + .server_streaming(request, entry.grpc_path.clone(), entry.codec()) + .await + }; + match tokio::time::timeout(deadline, call).await { + Ok(result) => result, + Err(_) => Err(tonic::Status::deadline_exceeded( + "upstream did not answer within the deadline", + )), + } + } +} + +/// Map the request onto the RPC's input message. +fn prepare( + proxy_state: S, headers: &HeaderMap, + connection: Option, path_params: &PathParams, raw_query: Option<&str>, body: Bytes, entry: &RouteEntry, -) -> Result< - ( - Grpc, - tonic::Request, - ), - Rejection, -> { +) -> Result, Unmappable> { let request_metadata = metadata::try_http_headers_to_grpc_metadata(headers, proxy_state.forwarded_headers()) - .map_err(|e| Rejection::Unmappable(e.to_string()))?; - let message = decode_request(entry, headers, path_params, raw_query, body) - .map_err(Rejection::Unmappable)?; + .map_err(|e| Unmappable(e.to_string()))?; + let message = + decode_request(entry, headers, path_params, raw_query, body).map_err(Unmappable)?; let mut request = tonic::Request::new(message); *request.metadata_mut() = request_metadata; - metadata::apply_request_deadline(&mut request, headers); - - let mut client = Grpc::new(proxy_state.grpc_channel()); - if let Err(e) = client.ready().await { - return Err(Rejection::NotReady(format!("gRPC upstream not ready: {e}"))); + // An upstream in process reads the HTTP client's address and TLS + // certificates with `Request::remote_addr` / `peer_certs`; a remote one + // never sees request extensions. + if let Some(connection) = connection { + connection.into_tonic_extensions(request.extensions_mut()); } - Ok((client, request)) + // Only the client's own deadline travels upstream: a default one would + // cut a long server stream short on an upstream that applies + // `grpc-timeout` to the whole call. + let deadline = metadata::apply_request_deadline(&mut request, headers) + .map_or(UPSTREAM_DEADLINE, |client| client.min(UPSTREAM_DEADLINE)); + + Ok(Call { + upstream: proxy_state.into_upstream(), + request, + deadline, + }) } /// A successful unary answer with its initial metadata and trailers kept @@ -605,15 +676,11 @@ struct UnaryAnswer { /// initial metadata, so a key sent in both keeps only its trailer value; the /// HTTP response carries both, so the call is made as a one-message stream /// instead, reading exactly what `Grpc::unary` reads. -async fn call_unary( - client: &mut Grpc, - request: tonic::Request, +async fn call_unary( + call: Call, entry: &RouteEntry, ) -> Result { - let response = client - .server_streaming(request, entry.grpc_path.clone(), entry.codec()) - .await?; - let (initial, mut stream, _) = response.into_parts(); + let (initial, mut stream, _) = call.open(entry).await?.into_parts(); let message = match stream.message().await { Ok(Some(message)) => message, Ok(None) => { @@ -654,12 +721,8 @@ fn with_initial(mut status: tonic::Status, initial: MetadataMap) -> tonic::Statu } /// Serve a unary RPC. -async fn unary_call( - mut client: Grpc, - request: tonic::Request, - entry: &RouteEntry, -) -> Response { - match call_unary(&mut client, request, entry).await { +async fn unary_call(call: Call, entry: &RouteEntry) -> Response { + match call_unary(call, entry).await { Ok(answer) => unary_success(entry, answer), Err(status) => upstream_error(status, entry), } @@ -754,19 +817,15 @@ fn upstream_error(mut status: tonic::Status, entry: &RouteEntry) -> Response { /// stream is the concatenated `data` of its messages. The upstream's initial /// metadata becomes response headers; trailers arrive after the headers are /// sent and are not forwarded. -async fn streaming_call( - mut client: Grpc, - request: tonic::Request, +async fn streaming_call( + call: Call, entry: Arc, use_sse: bool, keep_alive_secs: u64, ) -> Response { - let response = match client - .server_streaming(request, entry.grpc_path.clone(), entry.codec()) - .await - { + let response = match call.open(&entry).await { Ok(response) => response, - // Only a trailers-only rejection lands here. + // A trailers-only rejection, or no answer within the deadline. Err(status) => return upstream_error(status, &entry), }; let (initial, stream, _) = response.into_parts(); diff --git a/src/upstream.rs b/src/upstream.rs new file mode 100644 index 0000000..d9075c8 --- /dev/null +++ b/src/upstream.rs @@ -0,0 +1,242 @@ +//! The gRPC service the proxy calls. +//! +//! An upstream is any tower service that speaks gRPC over `http` types: a +//! remote [`tonic::transport::Channel`], or an embedder's own services in +//! process, such as [`tonic::service::Routes`]. The transcoder calls it for +//! every transcoded request, and [`ProxyService`](crate::ProxyService) hands it +//! native gRPC requests unchanged. + +use std::convert::Infallible; +use std::future::Future; +use std::pin::Pin; +use std::task::{ready, Context, Poll}; + +use bytes::Bytes; +use pin_project_lite::pin_project; +use tonic::client::GrpcService; + +/// The error type a gRPC service hands back through tonic. +pub type BoxError = Box; + +/// A gRPC service the proxy can call: a remote +/// [`Channel`](tonic::transport::Channel), [`tonic::service::Routes`], or any +/// other `tower::Service>` answering with a +/// gRPC response. Implemented for every such service; there is nothing to +/// implement by hand. +/// +/// An in-process upstream sees a transcoded call exactly as it sees a native +/// gRPC request: through its whole stack, interceptors and layers included. +/// +/// # Examples +/// +/// ``` +/// use structured_proxy::upstream::Upstream; +/// +/// fn accepts(_upstream: U) {} +/// +/// // An empty set of tonic services answers every call with UNIMPLEMENTED. +/// accepts(tonic::service::Routes::default()); +/// ``` +pub trait Upstream: + GrpcService< + tonic::body::Body, + ResponseBody: http_body::Body + Send> + Send + 'static, + Error: Into, + Future: Send, + > + Clone + + Send + + Sync + + 'static +{ +} + +impl Upstream for T where + T: GrpcService< + tonic::body::Body, + ResponseBody: http_body::Body + Send> + + Send + + 'static, + Error: Into, + Future: Send, + > + Clone + + Send + + Sync + + 'static +{ +} + +pin_project! { + /// One native gRPC request on its way through an upstream: waits for the + /// upstream to be ready, calls it, and answers with its response. A + /// failure of the upstream itself (a remote one that cannot be reached) + /// becomes a trailers-only gRPC error, so the client gets a status rather + /// than a broken connection. + /// + /// No proxy timer runs here and `grpc-timeout` travels unchanged: the + /// caller is a gRPC client, which enforces its own deadline (gRPC + /// PROTOCOL-HTTP2, "Timeout") by cancelling the stream, and that drops + /// this future wherever it waits, readiness included. A transcoded call is + /// different: there the proxy is the gRPC client and bounds the call itself. + #[project = PassThroughProj] + #[project_replace = PassThroughReplace] + pub(crate) enum PassThrough { + Ready { + upstream: U, + request: http::Request, + protocol: GrpcProtocol, + }, + Call { + #[pin] + future: U::Future, + protocol: GrpcProtocol, + }, + Done, + } +} + +impl PassThrough { + pub(crate) fn new( + upstream: U, + request: http::Request, + protocol: GrpcProtocol, + ) -> Self { + Self::Ready { + upstream, + request, + protocol, + } + } +} + +impl Future for PassThrough { + type Output = Result, Infallible>; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + loop { + match self.as_mut().project() { + PassThroughProj::Ready { + upstream, protocol, .. + } => { + let protocol = *protocol; + if let Err(e) = ready!(upstream.poll_ready(cx)) { + self.set(Self::Done); + return Poll::Ready(Ok(failure(e.into(), protocol))); + } + let PassThroughReplace::Ready { + mut upstream, + request, + .. + } = self.as_mut().project_replace(Self::Done) + else { + unreachable!("the state was just matched as Ready"); + }; + // The returned future owns what it needs; the service + // handle is dropped, as tower's `Oneshot` does. + let future = upstream.call(request); + self.set(Self::Call { future, protocol }); + } + PassThroughProj::Call { future, protocol } => { + let protocol = *protocol; + let result = ready!(future.poll(cx)); + self.set(Self::Done); + return Poll::Ready(Ok(match result { + Ok(response) => response.map(axum::body::Body::new), + Err(e) => failure(e.into(), protocol), + })); + } + PassThroughProj::Done => panic!("PassThrough polled after completion"), + } + } + } +} + +/// The trailers-only answer to an upstream that failed to take the call, in +/// the request's protocol: gRPC-Web carries a trailers-only status in its +/// headers too, but its client reads it only under a gRPC-Web content type +/// (gRPC PROTOCOL-WEB). +fn failure(error: BoxError, protocol: GrpcProtocol) -> http::Response { + let mut response: http::Response = + tonic::Status::from_error(error).into_http(); + if protocol != GrpcProtocol::Grpc { + response.headers_mut().insert( + http::header::CONTENT_TYPE, + http::HeaderValue::from_static(protocol.content_type()), + ); + } + response +} + +/// The gRPC protocol a request speaks. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum GrpcProtocol { + /// gRPC over HTTP/2 (gRPC PROTOCOL-HTTP2). + Grpc, + /// Binary gRPC-Web (gRPC PROTOCOL-WEB). + Web, + /// Base64 gRPC-Web, `-text` (gRPC PROTOCOL-WEB). + WebText, +} + +impl GrpcProtocol { + /// The content type of an answer the proxy writes itself: the protocol + /// with gRPC's default protobuf codec. + fn content_type(self) -> &'static str { + match self { + Self::Grpc => "application/grpc", + Self::Web => "application/grpc-web+proto", + Self::WebText => "application/grpc-web-text+proto", + } + } +} + +/// The gRPC protocol `headers` announce, if any: a media type of +/// `application/grpc` or `application/grpc+` (gRPC PROTOCOL-HTTP2, +/// "Content-Type"), or `application/grpc-web[-text][+]` (gRPC +/// PROTOCOL-WEB). Media types compare without case (RFC 9110 §8.3.1). +pub(crate) fn grpc_protocol(headers: &http::HeaderMap) -> Option { + let value = headers.get(http::header::CONTENT_TYPE)?.as_bytes(); + let media = match value.iter().position(|&b| b == b';') { + Some(end) => &value[..end], + None => value, + } + .trim_ascii(); + let rest = strip_prefix_ignore_case(media, b"application/grpc")?; + match rest { + [] | [b'+', ..] => Some(GrpcProtocol::Grpc), + _ => web_protocol(strip_prefix_ignore_case(rest, b"-web")?), + } +} + +/// Whether a request is a browser's CORS preflight for a gRPC-Web call: an +/// `OPTIONS` with an `Origin` (Fetch §3.2.2) whose +/// `Access-Control-Request-Headers` names `x-grpc-web`, the header gRPC-Web +/// clients send with every call (gRPC PROTOCOL-WEB). It carries no gRPC +/// content type, so only these headers tell it apart from a REST preflight. +pub(crate) fn is_grpc_web_preflight(method: &http::Method, headers: &http::HeaderMap) -> bool { + method == http::Method::OPTIONS + && headers.contains_key(http::header::ORIGIN) + && headers + .get_all(http::header::ACCESS_CONTROL_REQUEST_HEADERS) + .iter() + .flat_map(|value| value.as_bytes().split(|&b| b == b',')) + .any(|name| name.trim_ascii().eq_ignore_ascii_case(b"x-grpc-web")) +} + +/// What follows `application/grpc-web`: nothing or `+` for binary, +/// `-text` with an optional `+` for base64. +fn web_protocol(rest: &[u8]) -> Option { + let (protocol, rest) = match strip_prefix_ignore_case(rest, b"-text") { + Some(rest) => (GrpcProtocol::WebText, rest), + None => (GrpcProtocol::Web, rest), + }; + matches!(rest, [] | [b'+', ..]).then_some(protocol) +} + +/// `bytes` after `prefix`, compared without ASCII case. +fn strip_prefix_ignore_case<'a>(bytes: &'a [u8], prefix: &[u8]) -> Option<&'a [u8]> { + let (head, rest) = bytes.split_at_checked(prefix.len())?; + head.eq_ignore_ascii_case(prefix).then_some(rest) +} + +#[cfg(test)] +mod tests; diff --git a/src/upstream/tests.rs b/src/upstream/tests.rs new file mode 100644 index 0000000..39fd18d --- /dev/null +++ b/src/upstream/tests.rs @@ -0,0 +1,154 @@ +use super::*; + +fn with_content_type(value: &'static str) -> http::HeaderMap { + let mut headers = http::HeaderMap::new(); + headers.insert( + http::header::CONTENT_TYPE, + http::HeaderValue::from_static(value), + ); + headers +} + +fn protocol(value: &'static str) -> Option { + grpc_protocol(&with_content_type(value)) +} + +#[test] +fn grpc_media_types_are_grpc() { + // gRPC PROTOCOL-HTTP2: `application/grpc` with an optional `+codec`. + for value in [ + "application/grpc", + "application/grpc+proto", + "application/grpc+json", + "Application/GRPC", + "application/grpc; charset=utf-8", + " application/grpc ", + ] { + assert_eq!(protocol(value), Some(GrpcProtocol::Grpc), "{value}"); + } +} + +#[test] +fn grpc_web_media_types_are_grpc_web() { + // gRPC PROTOCOL-WEB: the binary and the base64 text form, each with a + // codec, told apart so a failure answers in the same one. + for value in [ + "application/grpc-web", + "application/grpc-web+proto", + "APPLICATION/GRPC-WEB", + ] { + assert_eq!(protocol(value), Some(GrpcProtocol::Web), "{value}"); + } + for value in [ + "application/grpc-web-text", + "application/grpc-web-text+proto", + "application/GRPC-WEB-TEXT", + ] { + assert_eq!(protocol(value), Some(GrpcProtocol::WebText), "{value}"); + } +} + +#[test] +fn other_media_types_are_not_grpc() { + // A type that only shares the prefix is someone else's: it must reach the + // proxy's routes, never the upstream. + for value in [ + "application/json", + "application/grpcx", + "application/grpc-webx", + "application/grpc-web-textual", + "application/grpc-json", + "application/gr", + "text/grpc", + "", + ] { + assert_eq!(protocol(value), None, "{value:?}"); + } +} + +#[test] +fn a_request_without_a_content_type_is_not_grpc() { + assert_eq!(grpc_protocol(&http::HeaderMap::new()), None); +} + +fn preflight(headers: &[(&'static str, &'static str)]) -> http::HeaderMap { + let mut map = http::HeaderMap::new(); + for (name, value) in headers { + map.append(*name, http::HeaderValue::from_static(value)); + } + map +} + +#[test] +fn a_preflight_announcing_x_grpc_web_is_a_grpc_web_preflight() { + for requested in [ + "x-grpc-web", + "content-type,x-grpc-web,x-user-agent", + "Content-Type, X-Grpc-Web", + ] { + let headers = preflight(&[ + ("origin", "https://app.example"), + ("access-control-request-headers", requested), + ]); + assert!( + is_grpc_web_preflight(&http::Method::OPTIONS, &headers), + "{requested}" + ); + } + // Split over several header lines, as a list may be. + let headers = preflight(&[ + ("origin", "https://app.example"), + ("access-control-request-headers", "content-type"), + ("access-control-request-headers", "x-grpc-web"), + ]); + assert!(is_grpc_web_preflight(&http::Method::OPTIONS, &headers)); +} + +#[test] +fn other_requests_are_not_grpc_web_preflights() { + let announcing = preflight(&[ + ("origin", "https://app.example"), + ("access-control-request-headers", "x-grpc-web"), + ]); + // Not an OPTIONS request. + assert!(!is_grpc_web_preflight(&http::Method::POST, &announcing)); + // No Origin: not a CORS request at all. + let no_origin = preflight(&[("access-control-request-headers", "x-grpc-web")]); + assert!(!is_grpc_web_preflight(&http::Method::OPTIONS, &no_origin)); + // A REST preflight, or a header that only shares the prefix. + for requested in ["content-type,authorization", "x-grpc-webx", "x-grpc"] { + let headers = preflight(&[ + ("origin", "https://app.example"), + ("access-control-request-headers", requested), + ]); + assert!( + !is_grpc_web_preflight(&http::Method::OPTIONS, &headers), + "{requested}" + ); + } +} + +#[test] +fn an_upstream_failure_is_a_trailers_only_grpc_status() { + // A remote upstream that cannot be reached answers the gRPC client with a + // status in the headers, not a broken connection. + let status = tonic::Status::unavailable("connection refused"); + let response = failure(Box::new(status), GrpcProtocol::Grpc); + assert_eq!(response.status(), http::StatusCode::OK); + assert_eq!(response.headers()["content-type"], "application/grpc"); + assert_eq!(response.headers()["grpc-status"], "14"); + assert_eq!(response.headers()["grpc-message"], "connection%20refused"); +} + +#[test] +fn an_upstream_failure_keeps_the_grpc_web_protocol() { + for (protocol, content_type) in [ + (GrpcProtocol::Web, "application/grpc-web+proto"), + (GrpcProtocol::WebText, "application/grpc-web-text+proto"), + ] { + let status = tonic::Status::unavailable("down"); + let response = failure(Box::new(status), protocol); + assert_eq!(response.headers()["content-type"], content_type); + assert_eq!(response.headers()["grpc-status"], "14"); + } +} diff --git a/tests/common/mod.rs b/tests/common/mod.rs index f050d98..c61ec4f 100644 --- a/tests/common/mod.rs +++ b/tests/common/mod.rs @@ -1,104 +1,53 @@ //! Harness for proxy tests against a real tonic upstream: compiles a test -//! `.proto` with `google.api.http` routes in memory, serves a gRPC service on a -//! random local port, and drives the proxy router built by `ProxyServer`. +//! `.proto` with `google.api.http` routes in memory, runs the gRPC service +//! either on a random local port (a remote upstream) or in process, and drives +//! the proxy service built by `ProxyServer`. + +use std::convert::Infallible; use axum::body::Body; use http::StatusCode; use prost_reflect::DescriptorPool; -use structured_proxy::config::ProxyConfig; use structured_proxy::transcode::error::ErrorDetailsPolicy; use structured_proxy::ProxyServer; +use tower::util::BoxCloneSyncService; use tower::ServiceExt; -/// Minimal `google/api/http.proto`: the fields the transcoder reads. -const HTTP_PROTO: &str = r#" -syntax = "proto3"; -package google.api; -message HttpRule { - string selector = 1; - oneof pattern { - string get = 2; - string put = 3; - string post = 4; - string delete = 5; - string patch = 6; - CustomHttpPattern custom = 8; - } - string body = 7; - string response_body = 12; - repeated HttpRule additional_bindings = 11; -} -message CustomHttpPattern { - string kind = 1; - string path = 2; -} -"#; - -/// `google/api/httpbody.proto`. -const HTTPBODY_PROTO: &str = r#" -syntax = "proto3"; -package google.api; -import "google/protobuf/any.proto"; -message HttpBody { - string content_type = 1; - bytes data = 2; - repeated google.protobuf.Any extensions = 3; -} -"#; - -const ANNOTATIONS_PROTO: &str = r#" -syntax = "proto3"; -package google.api; -import "google/api/http.proto"; -import "google/protobuf/descriptor.proto"; -extend google.protobuf.MethodOptions { - HttpRule http = 72295728; -} -"#; - -/// Serves the google.api sources and one test file from memory, and the -/// `google/protobuf` files from protox's bundled Google files. -struct TestProtos { - name: &'static str, - source: &'static str, -} +mod protos; -impl protox::file::FileResolver for TestProtos { - fn open_file(&self, name: &str) -> Result { - let source = match name { - "google/api/http.proto" => HTTP_PROTO, - "google/api/annotations.proto" => ANNOTATIONS_PROTO, - "google/api/httpbody.proto" => HTTPBODY_PROTO, - _ if name == self.name => self.source, - _ => return protox::file::GoogleFileResolver::new().open_file(name), - }; - protox::file::File::from_source(name, source) - } -} +pub use protos::compile; -/// Compile the test file `name` (which may import -/// `google/api/annotations.proto`) into a descriptor pool. -pub fn compile(name: &'static str, source: &'static str) -> DescriptorPool { - protox::Compiler::with_file_resolver(TestProtos { name, source }) - .open_file(name) - .expect("test protos compile") - .descriptor_pool() +/// The bounds of a tonic service a test runs as its upstream. +pub trait TestService: + tower::Service< + http::Request, + Response = http::Response, + Error = Infallible, + Future: Send + 'static, + > + tonic::server::NamedService + + Clone + + Send + + Sync + + 'static +{ } -/// Serve `service` on a random local port; returns its `http://` URL. -pub async fn serve(service: S) -> String -where +impl TestService for S where S: tower::Service< http::Request, Response = http::Response, - Error = std::convert::Infallible, + Error = Infallible, + Future: Send + 'static, > + tonic::server::NamedService + Clone + Send + Sync - + 'static, - S::Future: Send + 'static, + + 'static { +} + +/// Serve `service` on a random local port; returns its `http://` URL. +pub async fn serve(service: S) -> String { let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let addr = listener.local_addr().unwrap(); let incoming = futures::stream::unfold(listener, |listener| async move { @@ -113,24 +62,62 @@ where format!("http://{addr}") } -/// The proxy router for `pool` in front of `upstream`, returning error details -/// as `error_details` decides. -pub fn proxy( - upstream: &str, +/// Where the upstream of a proxy under test runs. +#[derive(Clone, Copy, Debug)] +pub enum Upstream { + /// A tonic server on a local port, reached over HTTP/2. + Remote, + /// The same service called in process, with no socket in between. + InProcess, +} + +/// A proxy under test, whatever its upstream. +pub type App = BoxCloneSyncService, http::Response, Infallible>; + +/// The proxy `configure` builds, in front of `service` running as `upstream` +/// says. `configure` gets the YAML naming the upstream address: an +/// `upstream:` block for a remote upstream, nothing for one in process. +pub async fn app( + upstream: Upstream, + service: S, + configure: impl FnOnce(&str) -> ProxyServer, +) -> App { + match upstream { + Upstream::Remote => { + let url = serve(service).await; + let server = configure(&format!("upstream:\n default: \"{url}\"\n")); + App::new(server.service(server.upstream().unwrap()).unwrap()) + } + Upstream::InProcess => { + let server = configure(""); + App::new( + server + .service(tonic::service::Routes::new(service)) + .unwrap(), + ) + } + } +} + +/// The proxy for `pool` in front of `service`, returning error details as +/// `error_details` decides. +pub async fn proxy( + upstream: Upstream, + service: S, pool: DescriptorPool, error_details: ErrorDetailsPolicy, -) -> axum::Router { - let config = - ProxyConfig::from_yaml_str(&format!("upstream:\n default: \"{upstream}\"\n")).unwrap(); - ProxyServer::from_config(config) - .with_descriptors(pool) - .with_error_details(error_details) - .router() - .unwrap() +) -> App { + app(upstream, service, |yaml| { + ProxyServer::from_yaml_str(yaml) + .unwrap() + .with_descriptors(pool) + .with_error_details(error_details) + }) + .await } /// Send `request` through `app`; returns the status and the body as text. -pub async fn send(app: &axum::Router, request: http::Request) -> (StatusCode, String) { +pub async fn send(app: &App, request: http::Request) -> (StatusCode, String) { let resp = app.clone().oneshot(request).await.unwrap(); let status = resp.status(); let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX) @@ -138,3 +125,21 @@ pub async fn send(app: &axum::Router, request: http::Request) -> (StatusCo .unwrap(); (status, String::from_utf8(bytes.to_vec()).unwrap()) } + +/// Declare each test twice, in a `remote` and an `in_process` module, so every +/// case runs against both kinds of upstream. A body names the one it runs +/// against as `UPSTREAM`. +macro_rules! upstream_tests { + ($($(#[$meta:meta])* async fn $name:ident() $body:block)*) => { + mod remote { + use super::*; + const UPSTREAM: common::Upstream = common::Upstream::Remote; + $($(#[$meta])* #[tokio::test] async fn $name() $body)* + } + mod in_process { + use super::*; + const UPSTREAM: common::Upstream = common::Upstream::InProcess; + $($(#[$meta])* #[tokio::test] async fn $name() $body)* + } + }; +} diff --git a/tests/common/protos.rs b/tests/common/protos.rs new file mode 100644 index 0000000..0c6ca37 --- /dev/null +++ b/tests/common/protos.rs @@ -0,0 +1,79 @@ +//! Compiles a test `.proto` with `google.api.http` routes in memory, with no +//! protoc binary. + +use prost_reflect::DescriptorPool; + +/// Minimal `google/api/http.proto`: the fields the transcoder reads. +const HTTP_PROTO: &str = r#" +syntax = "proto3"; +package google.api; +message HttpRule { + string selector = 1; + oneof pattern { + string get = 2; + string put = 3; + string post = 4; + string delete = 5; + string patch = 6; + CustomHttpPattern custom = 8; + } + string body = 7; + string response_body = 12; + repeated HttpRule additional_bindings = 11; +} +message CustomHttpPattern { + string kind = 1; + string path = 2; +} +"#; + +/// `google/api/httpbody.proto`. +const HTTPBODY_PROTO: &str = r#" +syntax = "proto3"; +package google.api; +import "google/protobuf/any.proto"; +message HttpBody { + string content_type = 1; + bytes data = 2; + repeated google.protobuf.Any extensions = 3; +} +"#; + +const ANNOTATIONS_PROTO: &str = r#" +syntax = "proto3"; +package google.api; +import "google/api/http.proto"; +import "google/protobuf/descriptor.proto"; +extend google.protobuf.MethodOptions { + HttpRule http = 72295728; +} +"#; + +/// Serves the google.api sources and one test file from memory, and the +/// `google/protobuf` files from protox's bundled Google files. +struct TestProtos { + name: &'static str, + source: &'static str, +} + +impl protox::file::FileResolver for TestProtos { + fn open_file(&self, name: &str) -> Result { + let source = match name { + "google/api/http.proto" => HTTP_PROTO, + "google/api/annotations.proto" => ANNOTATIONS_PROTO, + "google/api/httpbody.proto" => HTTPBODY_PROTO, + _ if name == self.name => self.source, + _ => return protox::file::GoogleFileResolver::new().open_file(name), + }; + protox::file::File::from_source(name, source) + } +} + +/// Compile the test file `name` (which may import +/// `google/api/annotations.proto`) into a descriptor pool. +pub fn compile(name: &'static str, source: &'static str) -> DescriptorPool { + protox::Compiler::with_file_resolver(TestProtos { name, source }) + .open_file(name) + .expect("test protos compile") + .descriptor_pool() +} diff --git a/tests/edge.rs b/tests/edge.rs new file mode 100644 index 0000000..44dfe75 --- /dev/null +++ b/tests/edge.rs @@ -0,0 +1,766 @@ +//! The proxy as the whole edge of a service: deadlines and trace context on +//! the way to the upstream, native gRPC and REST on one listener, the fallback +//! for what no route answers, and the client's address reaching an upstream in +//! process. Cases that hold for any upstream run against a remote and an +//! in-process one. + +#[macro_use] +mod common; + +use std::convert::Infallible; +use std::net::SocketAddr; +use std::pin::Pin; +use std::task::{Context, Poll}; + +use axum::body::Body; +use http::StatusCode; +use prost_reflect::{DescriptorPool, DynamicMessage, Value as PbValue}; +use serde_json::Value; +use structured_proxy::transcode::codec::DynamicCodec; +use structured_proxy::ProxyServer; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; + +const EDGE_PROTO: &str = r#" +syntax = "proto3"; +package test.v1; +import "google/api/annotations.proto"; + +message Req { + string name = 1; +} +// What the upstream saw of the call. +message Seen { + string name = 1; + string traceparent = 2; + string grpc_timeout = 3; + string peer = 4; +} + +service Edge { + rpc Echo(Req) returns (Seen) { + option (google.api.http) = { get: "/v1/echo/{name}" }; + } + // Never answers. + rpc Hang(Req) returns (Seen) { + option (google.api.http) = { get: "/v1/hang" }; + } +} +"#; + +fn pool() -> DescriptorPool { + common::compile("test/v1/edge.proto", EDGE_PROTO) +} + +// --- upstream --------------------------------------------------------------- + +type UnaryFuture = Pin< + Box< + dyn std::future::Future, tonic::Status>> + + Send, + >, +>; + +/// `Echo` answers with what it saw; `Hang` never answers. +#[derive(Clone)] +struct Handler { + pool: DescriptorPool, + rpc: String, +} + +impl tonic::server::UnaryService for Handler { + type Response = DynamicMessage; + type Future = UnaryFuture; + + fn call(&mut self, request: tonic::Request) -> Self::Future { + if self.rpc == "Hang" { + return Box::pin(std::future::pending()); + } + let metadata = |key: &str| { + request + .metadata() + .get(key) + .and_then(|v| v.to_str().ok()) + .unwrap_or_default() + .to_owned() + }; + let mut seen = DynamicMessage::new(self.pool.get_message_by_name("test.v1.Seen").unwrap()); + if let Some(PbValue::String(name)) = request.get_ref().get_field_by_name("name").as_deref() + { + seen.set_field_by_name("name", PbValue::String(name.clone())); + } + seen.set_field_by_name("traceparent", PbValue::String(metadata("traceparent"))); + seen.set_field_by_name("grpc_timeout", PbValue::String(metadata("grpc-timeout"))); + let peer = request + .remote_addr() + .map(|addr| addr.to_string()) + .unwrap_or_default(); + seen.set_field_by_name("peer", PbValue::String(peer)); + Box::pin(std::future::ready(Ok(tonic::Response::new(seen)))) + } +} + +/// The `test.v1.Edge` gRPC service, dispatching by method path. +#[derive(Clone)] +struct Edge { + pool: DescriptorPool, +} + +impl tonic::server::NamedService for Edge { + const NAME: &'static str = "test.v1.Edge"; +} + +impl tower::Service> for Edge { + type Response = http::Response; + type Error = Infallible; + type Future = + Pin> + Send>>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, req: http::Request) -> Self::Future { + let pool = self.pool.clone(); + Box::pin(async move { + let rpc = req + .uri() + .path() + .strip_prefix("/test.v1.Edge/") + .unwrap() + .to_owned(); + let input = pool.get_message_by_name("test.v1.Req").unwrap(); + let mut grpc = tonic::server::Grpc::new(DynamicCodec::new(input)); + Ok(grpc.unary(Handler { pool, rpc }, req).await) + }) + } +} + +// --- harness ---------------------------------------------------------------- + +/// The proxy in front of the `Edge` service. +async fn proxy(upstream: common::Upstream) -> common::App { + let pool = pool(); + common::proxy( + upstream, + Edge { pool: pool.clone() }, + pool, + Default::default(), + ) + .await +} + +/// Serve the proxy (built from `extra_yaml`, with `fallback` for what no route +/// answers) on a local port in front of the `Edge` service; returns the port's +/// address. +async fn listen( + upstream: common::Upstream, + extra_yaml: &str, + fallback: Option, +) -> SocketAddr { + let pool = pool(); + let service = Edge { pool: pool.clone() }; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + match upstream { + common::Upstream::Remote => { + let url = common::serve(service).await; + let server = ProxyServer::from_yaml_str(&format!( + "upstream:\n default: \"{url}\"\n{extra_yaml}" + )) + .unwrap() + .with_descriptors(pool); + let mut proxy = server.service(server.upstream().unwrap()).unwrap(); + if let Some(fallback) = fallback { + proxy = proxy.with_fallback(fallback); + } + tokio::spawn(structured_proxy::serve(listener, proxy)); + } + common::Upstream::InProcess => { + let server = ProxyServer::from_yaml_str(extra_yaml) + .unwrap() + .with_descriptors(pool); + let mut proxy = server + .service(tonic::service::Routes::new(service)) + .unwrap(); + if let Some(fallback) = fallback { + proxy = proxy.with_fallback(fallback); + } + tokio::spawn(structured_proxy::serve(listener, proxy)); + } + } + addr +} + +async fn get(app: &common::App, path: &str, headers: &[(&str, &str)]) -> (StatusCode, Value) { + let mut request = http::Request::get(path); + for (name, value) in headers { + request = request.header(*name, *value); + } + let (status, body) = common::send(app, request.body(Body::empty()).unwrap()).await; + (status, serde_json::from_str(&body).unwrap()) +} + +/// A plain HTTP/1.1 `GET` over a fresh connection; returns the status, the +/// body, and the client's own address. +async fn http1_get(addr: SocketAddr, path: &str) -> (u16, String, SocketAddr) { + let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap(); + let local = stream.local_addr().unwrap(); + let request = format!("GET {path} HTTP/1.1\r\nHost: {addr}\r\nConnection: close\r\n\r\n"); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut response = String::new(); + stream.read_to_string(&mut response).await.unwrap(); + let status = response[9..12].parse().unwrap(); + let body = response + .split_once("\r\n\r\n") + .map(|(_, body)| body.to_owned()) + .unwrap_or_default(); + (status, body, local) +} + +/// A native gRPC `Echo` call over HTTP/2 to `addr`. +async fn grpc_echo(addr: SocketAddr, name: &str) -> Result { + let pool = pool(); + let channel = tonic::transport::Channel::from_shared(format!("http://{addr}")) + .unwrap() + .connect() + .await + .unwrap(); + let mut grpc = tonic::client::Grpc::new(channel); + grpc.ready().await.unwrap(); + let mut req = DynamicMessage::new(pool.get_message_by_name("test.v1.Req").unwrap()); + req.set_field_by_name("name", PbValue::String(name.into())); + let codec = DynamicCodec::new(pool.get_message_by_name("test.v1.Seen").unwrap()); + grpc.unary( + tonic::Request::new(req), + http::uri::PathAndQuery::from_static("/test.v1.Edge/Echo"), + codec, + ) + .await + .map(tonic::Response::into_inner) +} + +fn field(msg: &DynamicMessage, name: &str) -> String { + match msg.get_field_by_name(name).as_deref() { + Some(PbValue::String(s)) => s.clone(), + _ => String::new(), + } +} + +const TRACEPARENT: &str = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"; + +upstream_tests! { +// --- context propagation ----------------------------------------------------- + +async fn a_client_traceparent_reaches_the_upstream() { + // W3C Trace Context §3.2: the upstream joins the client's trace. + let app = proxy(UPSTREAM).await; + let (status, seen) = get(&app, "/v1/echo/a", &[("traceparent", TRACEPARENT)]).await; + assert_eq!(status, StatusCode::OK, "{seen}"); + assert_eq!(seen["traceparent"], TRACEPARENT); +} + +async fn a_missing_traceparent_is_synthesized_for_the_upstream() { + let app = proxy(UPSTREAM).await; + let (status, seen) = get(&app, "/v1/echo/a", &[]).await; + assert_eq!(status, StatusCode::OK, "{seen}"); + let traceparent = seen["traceparent"].as_str().unwrap(); + assert!(traceparent.starts_with("00-"), "{traceparent}"); + assert_eq!(traceparent.len(), 55, "{traceparent}"); + assert_ne!(traceparent, TRACEPARENT); +} + +async fn a_client_deadline_reaches_the_upstream() { + // The upstream learns how long the client waits, whatever carries it. + let app = proxy(UPSTREAM).await; + let (status, seen) = get(&app, "/v1/echo/a", &[("grpc-timeout", "3S")]).await; + assert_eq!(status, StatusCode::OK, "{seen}"); + assert!(!seen["grpcTimeout"].as_str().unwrap().is_empty(), "{seen}"); +} + +async fn the_default_deadline_does_not_reach_the_upstream() { + // A default deadline sent upstream would end a long server stream on an + // upstream that applies `grpc-timeout` to the whole call. + let app = proxy(UPSTREAM).await; + let (status, seen) = get(&app, "/v1/echo/a", &[]).await; + assert_eq!(status, StatusCode::OK, "{seen}"); + assert_eq!(seen["grpcTimeout"], ""); +} + +// --- one listener -------------------------------------------------------------- + +async fn rest_and_native_grpc_share_one_listener() { + // HTTP/1.1 REST and HTTP/2 gRPC clients reach the same port: the REST call + // is transcoded, the gRPC call reaches the upstream as it was sent. + let addr = listen(UPSTREAM, "", None).await; + let (status, body, _) = http1_get(addr, "/v1/echo/rest").await; + assert_eq!(status, 200, "{body}"); + assert_eq!(serde_json::from_str::(&body).unwrap()["name"], "rest"); + let seen = grpc_echo(addr, "native").await.unwrap(); + assert_eq!(field(&seen, "name"), "native"); +} + +async fn an_unmatched_request_is_404_without_a_fallback() { + let addr = listen(UPSTREAM, "", None).await; + let (status, _, _) = http1_get(addr, "/static/index.html").await; + assert_eq!(status, 404); +} + +async fn maintenance_gates_the_routes_but_not_the_fallback_or_native_grpc() { + // What the proxy does not serve is passed through untouched: its + // maintenance gate answers 503 on its own routes only. + let fallback = axum::Router::new().route( + "/static/index.html", + axum::routing::get(|| async { "static" }), + ); + let addr = listen(UPSTREAM, "maintenance:\n enabled: true\n", Some(fallback)).await; + let (status, _, _) = http1_get(addr, "/v1/echo/rest").await; + assert_eq!(status, 503); + let (status, body, _) = http1_get(addr, "/static/index.html").await; + assert_eq!(status, 200); + assert_eq!(body, "static"); + let seen = grpc_echo(addr, "native").await.unwrap(); + assert_eq!(field(&seen, "name"), "native"); +} +} + +// --- deadlines --------------------------------------------------------------- + +#[tokio::test] +async fn remote_default_deadline_ends_a_call_the_upstream_never_answers() { + // Nothing on the remote side limits the call (no client `grpc-timeout`), + // so the proxy's own deadline is what ends it: 504 DEADLINE_EXCEEDED. + let app = proxy(common::Upstream::Remote).await; + let started = std::time::Instant::now(); + let (status, body) = get(&app, "/v1/hang", &[]).await; + assert_eq!(status, StatusCode::GATEWAY_TIMEOUT, "{body}"); + assert_eq!(body["error"], "DEADLINE_EXCEEDED"); + assert!(started.elapsed() >= structured_proxy::transcode::UPSTREAM_DEADLINE); +} + +#[tokio::test(start_paused = true)] +async fn in_process_default_deadline_ends_a_call_the_upstream_never_answers() { + // An upstream in process has no transport to time the call out, so the + // proxy does. + let app = proxy(common::Upstream::InProcess).await; + let started = tokio::time::Instant::now(); + let (status, body) = get(&app, "/v1/hang", &[]).await; + assert_eq!(status, StatusCode::GATEWAY_TIMEOUT, "{body}"); + assert_eq!(body["error"], "DEADLINE_EXCEEDED"); + assert_eq!( + started.elapsed(), + structured_proxy::transcode::UPSTREAM_DEADLINE + ); +} + +#[tokio::test(start_paused = true)] +async fn in_process_client_deadline_shorter_than_the_default_ends_the_call() { + // A remote tonic server enforces `grpc-timeout` itself and races the + // proxy to the answer; in process only the proxy's clock runs. + let app = proxy(common::Upstream::InProcess).await; + let started = tokio::time::Instant::now(); + let (status, body) = get(&app, "/v1/hang", &[("grpc-timeout", "250m")]).await; + assert_eq!(status, StatusCode::GATEWAY_TIMEOUT, "{body}"); + assert_eq!(started.elapsed(), std::time::Duration::from_millis(250)); +} + +#[tokio::test(start_paused = true)] +async fn in_process_client_deadline_longer_than_the_default_is_capped() { + let app = proxy(common::Upstream::InProcess).await; + let started = tokio::time::Instant::now(); + let (status, _) = get(&app, "/v1/hang", &[("grpc-timeout", "1M")]).await; + assert_eq!(status, StatusCode::GATEWAY_TIMEOUT); + assert_eq!( + started.elapsed(), + structured_proxy::transcode::UPSTREAM_DEADLINE + ); +} + +/// An upstream under backpressure that never frees a slot: `poll_ready` stays +/// pending, as behind a saturated concurrency limit. +#[derive(Clone)] +struct NeverReady; + +impl tower::Service> for NeverReady { + type Response = http::Response; + type Error = Infallible; + type Future = std::future::Pending>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Pending + } + + fn call(&mut self, _req: http::Request) -> Self::Future { + unreachable!("never ready, never called") + } +} + +#[tokio::test(start_paused = true)] +async fn waiting_for_a_saturated_upstream_counts_against_the_deadline() { + // Readiness is part of the call: an upstream that never takes the call + // must not hold the request past its deadline. + let server = ProxyServer::from_yaml_str("") + .unwrap() + .with_descriptors(pool()); + let app = common::App::new(server.service(NeverReady).unwrap()); + let started = tokio::time::Instant::now(); + let (status, body) = get(&app, "/v1/echo/a", &[("grpc-timeout", "250m")]).await; + assert_eq!(status, StatusCode::GATEWAY_TIMEOUT, "{body}"); + assert_eq!(body["error"], "DEADLINE_EXCEEDED"); + assert_eq!(started.elapsed(), std::time::Duration::from_millis(250)); +} + +// --- gRPC-Web -------------------------------------------------------------------- + +/// The proxy in front of the `Edge` service in process, made to speak gRPC-Web +/// the way an embedder does it: tonic-web's layer around its services. +fn grpc_web_proxy() -> common::App { + grpc_web_proxy_with("") +} + +/// [`grpc_web_proxy`] configured by `yaml`. +fn grpc_web_proxy_with(yaml: &str) -> common::App { + let pool = pool(); + let upstream = tower::ServiceBuilder::new() + .layer(tonic_web::GrpcWebLayer::new()) + .service(tonic::service::Routes::new(Edge { pool: pool.clone() })); + let server = ProxyServer::from_yaml_str(yaml) + .unwrap() + .with_descriptors(pool); + common::App::new(server.service(upstream).unwrap()) +} + +const ORIGIN: &str = "https://app.example"; + +/// A browser's cross-origin gRPC-Web call from [`ORIGIN`] through `app`; +/// returns the response headers. +async fn browser_grpc_web_call(app: common::App) -> http::HeaderMap { + let request = http::Request::post("/test.v1.Edge/Echo") + .header("origin", ORIGIN) + .header("content-type", "application/grpc-web+proto") + .header("x-grpc-web", "1") + .body(Body::from(request_frame("browser"))) + .unwrap(); + let response = tower::ServiceExt::oneshot(app, request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + response.headers().clone() +} + +/// The browser's preflight for that call; returns the response headers. +async fn browser_grpc_web_preflight(app: common::App) -> http::HeaderMap { + let request = http::Request::builder() + .method("OPTIONS") + .uri("/test.v1.Edge/Echo") + .header("origin", ORIGIN) + .header("access-control-request-method", "POST") + .header( + "access-control-request-headers", + "content-type,x-grpc-web,x-user-agent", + ) + .body(Body::empty()) + .unwrap(); + tower::ServiceExt::oneshot(app, request) + .await + .unwrap() + .headers() + .clone() +} + +fn exposed(headers: &http::HeaderMap) -> Vec { + headers + .get_all("access-control-expose-headers") + .iter() + .flat_map(|v| v.to_str().unwrap().split(',')) + .map(|name| name.trim().to_ascii_lowercase()) + .collect() +} + +const CORS_YAML: &str = "cors:\n origins: [\"https://app.example\"]\n"; + +#[tokio::test] +async fn a_browser_grpc_web_call_gets_the_cors_policy_its_preflight_got() { + // The proxy answers the preflight under its CORS policy, so the call must + // carry the same policy: without it the browser discards the response. + let preflight = browser_grpc_web_preflight(grpc_web_proxy_with(CORS_YAML)).await; + assert_eq!(preflight["access-control-allow-origin"], ORIGIN); + let call = browser_grpc_web_call(grpc_web_proxy_with(CORS_YAML)).await; + assert_eq!(call["access-control-allow-origin"], ORIGIN); + assert_eq!(call["access-control-allow-credentials"], "true"); + // A gRPC-Web client reads the status and its details from these. + let exposed = exposed(&call); + for name in ["grpc-status", "grpc-message", "grpc-status-details-bin"] { + assert!(exposed.iter().any(|e| e == name), "{name}: {exposed:?}"); + } +} + +#[tokio::test] +async fn a_grpc_web_call_from_an_unlisted_origin_gets_no_allowance() { + let request = http::Request::post("/test.v1.Edge/Echo") + .header("origin", "https://evil.example") + .header("content-type", "application/grpc-web+proto") + .body(Body::from(request_frame("x"))) + .unwrap(); + let response = tower::ServiceExt::oneshot(grpc_web_proxy_with(CORS_YAML), request) + .await + .unwrap(); + assert!(response + .headers() + .get("access-control-allow-origin") + .is_none()); +} + +#[tokio::test] +async fn configured_expose_headers_and_max_age_reach_the_browser() { + // Upstream metadata a browser must read is exposed by name, and the + // preflight tells the browser how long to cache it. + let yaml = "cors:\n origins: [\"https://app.example\"]\n expose_headers: [\"x-request-id\"]\n max_age_secs: 600\n"; + let call = browser_grpc_web_call(grpc_web_proxy_with(yaml)).await; + assert!(exposed(&call).iter().any(|e| e == "x-request-id")); + let preflight = browser_grpc_web_preflight(grpc_web_proxy_with(yaml)).await; + assert_eq!(preflight["access-control-max-age"], "600"); +} + +#[tokio::test] +async fn cors_can_be_left_to_an_upstream_that_does_it_itself() { + // `grpc_web: false`: the upstream's own CORS layer answers, and the proxy + // adds nothing that could clash with it. + let yaml = "cors:\n origins: [\"https://app.example\"]\n grpc_web: false\n"; + let call = browser_grpc_web_call(grpc_web_proxy_with(yaml)).await; + assert!(call.get("access-control-allow-origin").is_none()); +} + +#[tokio::test] +async fn a_grpc_web_preflight_is_answered_by_the_proxy_despite_a_fallback() { + // The fallback takes what no route matches, but a gRPC-Web preflight + // belongs to the call it announces, which carries the proxy's policy. + let pool = pool(); + let upstream = tower::ServiceBuilder::new() + .layer(tonic_web::GrpcWebLayer::new()) + .service(tonic::service::Routes::new(Edge { pool: pool.clone() })); + let fallback = axum::Router::new().fallback(|| async { (StatusCode::IM_A_TEAPOT, "fallback") }); + let service = ProxyServer::from_yaml_str(CORS_YAML) + .unwrap() + .with_descriptors(pool) + .service(upstream) + .unwrap() + .with_fallback(fallback); + let preflight = browser_grpc_web_preflight(common::App::new(service)).await; + assert_eq!(preflight["access-control-allow-origin"], ORIGIN); +} + +/// An upstream that speaks gRPC-Web and sets its own CORS policy, allowing +/// only [`UPSTREAM_ORIGIN`]. +fn upstream_with_its_own_cors() -> common::App { + let pool = pool(); + let upstream = tower::ServiceBuilder::new() + .layer( + tower_http::cors::CorsLayer::new() + .allow_origin(http::HeaderValue::from_static(UPSTREAM_ORIGIN)) + .allow_methods([http::Method::POST]) + .allow_headers(tower_http::cors::Any), + ) + .layer(tonic_web::GrpcWebLayer::new()) + .service(tonic::service::Routes::new(Edge { pool: pool.clone() })); + let service = ProxyServer::from_yaml_str( + "cors:\n origins: [\"https://app.example\"]\n grpc_web: false\n", + ) + .unwrap() + .with_descriptors(pool) + .service(upstream) + .unwrap(); + common::App::new(service) +} + +const UPSTREAM_ORIGIN: &str = "https://upstream-policy.example"; + +#[tokio::test] +async fn a_grpc_web_preflight_goes_to_an_upstream_that_owns_cors() { + // `grpc_web: false`: the preflight must get the policy the call will get, + // the upstream's, not the proxy's. + let request = http::Request::builder() + .method("OPTIONS") + .uri("/test.v1.Edge/Echo") + .header("origin", UPSTREAM_ORIGIN) + .header("access-control-request-method", "POST") + .header("access-control-request-headers", "content-type,x-grpc-web") + .body(Body::empty()) + .unwrap(); + let response = tower::ServiceExt::oneshot(upstream_with_its_own_cors(), request) + .await + .unwrap(); + assert_eq!( + response.headers()["access-control-allow-origin"], + UPSTREAM_ORIGIN + ); +} + +#[tokio::test] +async fn a_rest_preflight_stays_with_the_proxy_when_the_upstream_owns_grpc_web_cors() { + // Only gRPC-Web preflights follow their call upstream: a REST route's + // preflight keeps the proxy's policy. + let request = http::Request::builder() + .method("OPTIONS") + .uri("/v1/echo/rest") + .header("origin", ORIGIN) + .header("access-control-request-method", "GET") + .body(Body::empty()) + .unwrap(); + let response = tower::ServiceExt::oneshot(upstream_with_its_own_cors(), request) + .await + .unwrap(); + assert_eq!(response.headers()["access-control-allow-origin"], ORIGIN); +} + +#[tokio::test] +async fn a_listed_origin_gets_its_cors_allowance_on_a_rest_route() { + // A named origin list with credentials must not use `*` for methods or + // headers (Fetch §3.2.5): such a policy used to stop the proxy at startup. + let app = grpc_web_proxy_with(CORS_YAML); + let request = http::Request::get("/v1/echo/rest") + .header("origin", ORIGIN) + .body(Body::empty()) + .unwrap(); + let response = tower::ServiceExt::oneshot(app.clone(), request) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!(response.headers()["access-control-allow-origin"], ORIGIN); + // The preflight echoes what the browser asked for. + let preflight = http::Request::builder() + .method("OPTIONS") + .uri("/v1/echo/rest") + .header("origin", ORIGIN) + .header("access-control-request-method", "GET") + .header("access-control-request-headers", "authorization") + .body(Body::empty()) + .unwrap(); + let response = tower::ServiceExt::oneshot(app, preflight).await.unwrap(); + assert_eq!(response.headers()["access-control-allow-methods"], "GET"); + assert_eq!( + response.headers()["access-control-allow-headers"], + "authorization" + ); +} + +#[test] +fn an_origin_that_is_not_a_header_value_is_a_config_error() { + // Dropping it would quietly narrow the policy. + let server = ProxyServer::from_yaml_str("cors:\n origins: [\"https://a.example\\n\"]\n") + .unwrap() + .with_descriptors(pool()); + let Err(err) = server.service(tonic::service::Routes::default()) else { + panic!("an invalid origin must be refused"); + }; + assert!(err.to_string().contains("cors.origins"), "{err}"); +} + +#[test] +fn an_expose_header_that_is_not_a_header_name_is_a_config_error() { + let server = ProxyServer::from_yaml_str("cors:\n expose_headers: [\"not a header\"]\n") + .unwrap() + .with_descriptors(pool()); + let Err(err) = server.service(tonic::service::Routes::default()) else { + panic!("an invalid header name must be refused"); + }; + assert!(err.to_string().contains("not a header"), "{err}"); +} + +/// One gRPC message frame (flag 0, big-endian length) holding `Req{name}`. +fn request_frame(name: &str) -> Vec { + let pool = pool(); + let mut req = DynamicMessage::new(pool.get_message_by_name("test.v1.Req").unwrap()); + req.set_field_by_name("name", PbValue::String(name.into())); + let payload = prost::Message::encode_to_vec(&req); + let mut frame = vec![0]; + frame.extend_from_slice(&u32::try_from(payload.len()).unwrap().to_be_bytes()); + frame.extend_from_slice(&payload); + frame +} + +/// The `name` of the `Seen` message in the first frame of a gRPC-Web body. +fn seen_name(body: &[u8]) -> String { + assert_eq!(body[0], 0, "a message frame comes first"); + let len = u32::from_be_bytes(body[1..5].try_into().unwrap()) as usize; + let seen = DynamicMessage::decode( + pool().get_message_by_name("test.v1.Seen").unwrap(), + &body[5..5 + len], + ) + .unwrap(); + field(&seen, "name") +} + +/// Send a gRPC-Web `Echo` with `body` as `content_type`; returns the response +/// content type and body. +async fn grpc_web_echo(content_type: &str, body: Vec) -> (String, bytes::Bytes) { + // A gRPC-Web client names the encoding it reads back in `Accept`. + let request = http::Request::post("/test.v1.Edge/Echo") + .header("content-type", content_type) + .header("accept", content_type) + .header("x-grpc-web", "1") + .body(Body::from(body)) + .unwrap(); + let response = tower::ServiceExt::oneshot(grpc_web_proxy(), request) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let content_type = response.headers()["content-type"] + .to_str() + .unwrap() + .to_owned(); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + (content_type, body) +} + +#[tokio::test] +async fn binary_grpc_web_reaches_an_upstream_that_speaks_it() { + // gRPC-Web passes through unchanged; the upstream's own gRPC-Web layer + // answers it. + let (content_type, body) = + grpc_web_echo("application/grpc-web+proto", request_frame("web")).await; + assert_eq!(content_type, "application/grpc-web+proto"); + assert_eq!(seen_name(&body), "web"); +} + +#[tokio::test] +async fn text_grpc_web_reaches_an_upstream_that_speaks_it() { + use base64::Engine as _; + let engine = base64::engine::general_purpose::STANDARD; + let (content_type, body) = grpc_web_echo( + "application/grpc-web-text+proto", + engine.encode(request_frame("text")).into_bytes(), + ) + .await; + assert_eq!(content_type, "application/grpc-web-text+proto"); + // Each frame is its own base64 run, padded at its end (gRPC + // PROTOCOL-WEB), so the body decodes group by group. + let decoded: Vec = body + .chunks(4) + .flat_map(|group| engine.decode(group).unwrap()) + .collect(); + assert_eq!(seen_name(&decoded), "text"); +} + +// --- the client's address ------------------------------------------------------- + +#[tokio::test] +async fn an_in_process_upstream_sees_the_client_address_of_a_rest_call() { + // `Request::remote_addr` gives the HTTP client, not a loopback hop: a + // forward-auth decision can key on it. + let addr = listen(common::Upstream::InProcess, "", None).await; + let (status, body, client) = http1_get(addr, "/v1/echo/rest").await; + assert_eq!(status, 200, "{body}"); + let seen: Value = serde_json::from_str(&body).unwrap(); + assert_eq!(seen["peer"], client.to_string()); +} + +#[tokio::test] +async fn an_in_process_upstream_sees_the_client_address_of_a_native_call() { + // As behind tonic's own server: the gRPC client's address, on loopback + // here, with the port its connection came from. + let addr = listen(common::Upstream::InProcess, "", None).await; + let seen = grpc_echo(addr, "native").await.unwrap(); + let peer: SocketAddr = field(&seen, "peer").parse().unwrap(); + assert!(peer.ip().is_loopback(), "{peer}"); + assert_ne!(peer.port(), addr.port()); +} diff --git a/tests/embedded.rs b/tests/embedded.rs index 06d2f1f..788d38b 100644 --- a/tests/embedded.rs +++ b/tests/embedded.rs @@ -16,9 +16,9 @@ use structured_proxy::ProxyServer; fn embedded_config_is_constructible() { static DESCRIPTOR_BYTES: &[u8] = &[]; let config = ProxyConfig { - upstream: UpstreamConfig { + upstream: Some(UpstreamConfig { default: "http://127.0.0.1:50051".into(), - }, + }), descriptors: vec![DescriptorSource::Embedded { bytes: DESCRIPTOR_BYTES, }], diff --git a/tests/error_details.rs b/tests/error_details.rs index d6c2215..5075aac 100644 --- a/tests/error_details.rs +++ b/tests/error_details.rs @@ -4,8 +4,10 @@ //! gRPC server that fails with `tonic_types` details, so every case below goes //! over the actual `grpc-status-details-bin` trailer: unary errors, a stream //! refused before any response header, and a stream that fails after its first -//! message, in both NDJSON and SSE. +//! message, in both NDJSON and SSE. Every case runs against a remote and an +//! in-process upstream. +#[macro_use] mod common; use std::convert::Infallible; @@ -293,15 +295,26 @@ impl tower::Service> for Things { // --- proxy harness ---------------------------------------------------------- -/// A proxy router in front of a fresh upstream, returning error details as +/// A proxy in front of a fresh upstream, returning error details as /// `error_details` decides. -async fn proxy(error_details: ErrorDetailsPolicy) -> axum::Router { +async fn proxy(upstream: common::Upstream, error_details: ErrorDetailsPolicy) -> common::App { let pool = pool(); - let upstream = common::serve(Things { pool: pool.clone() }).await; - common::proxy(&upstream, pool, error_details) + common::proxy(upstream, Things { pool: pool.clone() }, pool, error_details).await } -async fn get(app: &axum::Router, path: &str, accept: Option<&str>) -> (StatusCode, String) { +/// A proxy created from a YAML document (the upstream plus `extra_yaml`), the +/// way the standalone binary reads its config file. +async fn proxy_from_yaml(upstream: common::Upstream, extra_yaml: &str) -> common::App { + let pool = pool(); + common::app(upstream, Things { pool: pool.clone() }, |upstream_yaml| { + structured_proxy::ProxyServer::from_yaml_str(&format!("{upstream_yaml}{extra_yaml}")) + .unwrap() + .with_descriptors(pool) + }) + .await +} + +async fn get(app: &common::App, path: &str, accept: Option<&str>) -> (StatusCode, String) { let mut req = http::Request::get(path); if let Some(accept) = accept { req = req.header("accept", accept); @@ -309,19 +322,19 @@ async fn get(app: &axum::Router, path: &str, accept: Option<&str>) -> (StatusCod common::send(app, req.body(Body::empty()).unwrap()).await } -async fn get_json(app: &axum::Router, path: &str) -> (StatusCode, Value) { +async fn get_json(app: &common::App, path: &str) -> (StatusCode, Value) { let (status, body) = get(app, path, None).await; (status, serde_json::from_str(&body).unwrap()) } +upstream_tests! { // --- unary ------------------------------------------------------------------ -#[tokio::test] async fn unary_error_carries_error_info_and_bad_request() { // The acceptance case: typed details arrive as ProtoJSON `Any`s next to // the existing fields, with the HTTP status of the gRPC → HTTP mapping, // and DebugInfo stays behind. - let app = proxy(ErrorDetailsPolicy::default()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::default()).await; let (status, body) = get_json(&app, "/v1/things/rich").await; assert_eq!(status, StatusCode::BAD_REQUEST); assert_eq!( @@ -337,9 +350,8 @@ async fn unary_error_carries_error_info_and_bad_request() { assert!(!text.contains("register.rs") && !text.contains("users_email_key")); } -#[tokio::test] async fn unary_error_without_trailer_has_empty_details() { - let app = proxy(ErrorDetailsPolicy::default()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::default()).await; let (status, body) = get_json(&app, "/v1/things/missing").await; assert_eq!(status, StatusCode::NOT_FOUND); assert_eq!( @@ -348,7 +360,6 @@ async fn unary_error_without_trailer_has_empty_details() { ); } -#[tokio::test] async fn product_unknown_and_well_known_details_are_told_apart() { // Three different renderings side by side: a product message expands to // its fields and a well-known type with a special JSON form sits under @@ -357,6 +368,7 @@ async fn product_unknown_and_well_known_details_are_told_apart() { // base64 of the original bytes), never into `details`, on the route that // switches the extension on. let app = proxy( + UPSTREAM, ErrorDetailsPolicy::default() .opaque_route("/v1/things/*", true) .unwrap(), @@ -385,11 +397,11 @@ async fn product_unknown_and_well_known_details_are_told_apart() { // --- per-route switch ------------------------------------------------------- -#[tokio::test] async fn route_rule_switches_details_off_for_one_route() { // Only the matched route loses `details` (the key is absent, not empty); // its HTTP status and the other routes are unaffected. let app = proxy( + UPSTREAM, ErrorDetailsPolicy::default() .route("/v1/quiet/*", false) .unwrap(), @@ -405,11 +417,11 @@ async fn route_rule_switches_details_off_for_one_route() { assert_eq!(loud["details"], rich_details()); } -#[tokio::test] async fn global_switch_off_with_a_sub_route_back_on() { // Global off, `/v1/things/**` back on: the sub-route (including its // streaming routes) keeps details, everything else drops them. let app = proxy( + UPSTREAM, ErrorDetailsPolicy::disabled() .route("/v1/things/**", true) .unwrap(), @@ -423,11 +435,10 @@ async fn global_switch_off_with_a_sub_route_back_on() { assert_eq!(denied["details"][0]["reason"], "NOT_OWNER"); } -#[tokio::test] async fn global_switch_off_removes_details_from_stream_frames_too() { // The switch covers the in-stream terminal frame as well: a route with // details off ends its stream with the bare error body. - let app = proxy(ErrorDetailsPolicy::disabled()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::disabled()).await; let (status, body) = get(&app, "/v1/things/x/watch", None).await; assert_eq!(status, StatusCode::OK, "{body}"); let last: Value = serde_json::from_str(body.lines().last().unwrap()).unwrap(); @@ -444,13 +455,12 @@ async fn global_switch_off_removes_details_from_stream_frames_too() { // --- broken upstream status --------------------------------------------------- -#[tokio::test] async fn unary_error_with_a_corrupt_known_detail_becomes_a_safe_internal() { // The type resolves but its bytes do not decode: a broken upstream // response. Before headers the proxy still owns the status, so the client // gets a generic 500 INTERNAL, not the upstream's NOT_FOUND with the detail // dropped, passed on as base64, or otherwise reinterpreted. - let app = proxy(ErrorDetailsPolicy::default()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::default()).await; let (status, body) = get(&app, "/v1/things/corrupt", None).await; assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); assert_eq!( @@ -460,11 +470,10 @@ async fn unary_error_with_a_corrupt_known_detail_becomes_a_safe_internal() { assert!(!body.contains("CgVh") && !body.contains("gone"), "{body}"); } -#[tokio::test] async fn stream_error_with_a_corrupt_known_detail_ends_with_a_safe_internal_frame() { // After the first message the 200 is sent, so the same failure becomes the // terminal frame instead, in both formats. - let app = proxy(ErrorDetailsPolicy::default()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::default()).await; let (status, body) = get(&app, "/v1/things/corrupt/watch", None).await; assert_eq!(status, StatusCode::OK); let mut expected = malformed_upstream_status_body(); @@ -493,25 +502,11 @@ async fn stream_error_with_a_corrupt_known_detail_ends_with_a_safe_internal_fram // --- YAML settings ------------------------------------------------------------ -/// A proxy created from a YAML document (upstream plus `extra_yaml`), the way -/// the standalone binary reads its config file. -async fn proxy_from_yaml(extra_yaml: &str) -> axum::Router { - let pool = pool(); - let upstream = common::serve(Things { pool: pool.clone() }).await; - structured_proxy::ProxyServer::from_yaml_str(&format!( - "upstream:\n default: \"{upstream}\"\n{extra_yaml}" - )) - .unwrap() - .with_descriptors(pool) - .router() - .unwrap() -} - -#[tokio::test] async fn yaml_switches_route_details_off_and_envelopes_ndjson() { // Both settings come from the config file: the quiet route loses its // details, the others keep them, and NDJSON lines are enveloped. let app = proxy_from_yaml( + UPSTREAM, "error_details:\n routes:\n - pattern: \"/v1/quiet/*\"\n enabled: false\nstreaming:\n ndjson_envelope: true\n", ) .await; @@ -540,11 +535,11 @@ async fn yaml_switches_route_details_off_and_envelopes_ndjson() { ); } -#[tokio::test] async fn yaml_switches_opaque_details_globally_and_per_route() { // `opaque: true` turns the extension on everywhere and a rule setting only // `opaque: false` turns it back off for one route, leaving its details on. let app = proxy_from_yaml( + UPSTREAM, "error_details:\n opaque: true\n routes:\n - pattern: \"/v1/quiet/*\"\n opaque: false\n", ) .await; @@ -555,39 +550,13 @@ async fn yaml_switches_opaque_details_globally_and_per_route() { assert!(quiet.get("opaqueDetails").is_none(), "{quiet}"); } -#[test] -fn yaml_route_rule_without_a_switch_is_rejected() { - // A rule that sets neither `enabled` nor `opaque` would match routes and - // change nothing, which is a mistake rather than an intent. - let err = structured_proxy::ProxyServer::from_yaml_str( - "upstream:\n default: \"http://127.0.0.1:1\"\nerror_details:\n routes:\n - pattern: \"/v1/**\"\n", - ) - .err() - .expect("a rule without a switch must be rejected"); - assert!( - err.to_string().contains("neither enabled nor opaque"), - "{err}" - ); -} - -#[test] -fn yaml_with_an_invalid_error_details_pattern_is_rejected() { - let err = structured_proxy::ProxyServer::from_yaml_str( - "upstream:\n default: \"http://127.0.0.1:1\"\nerror_details:\n routes:\n - pattern: \"v1/**\"\n enabled: false\n", - ) - .err() - .expect("a relative pattern must be rejected"); - assert!(err.to_string().contains("must start with '/'"), "{err}"); -} - // --- errors the proxy raises itself ------------------------------------------ -#[tokio::test] async fn unmappable_request_gets_the_shared_error_body() { // A request the proxy rejects before calling the upstream answers in the // same body as an upstream error on that route, so a client parses one // shape: here INVALID_ARGUMENT with empty details. - let app = proxy(ErrorDetailsPolicy::default()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::default()).await; let (status, body) = get_json(&app, "/v1/things/rich?count=many").await; assert_eq!(status, StatusCode::BAD_REQUEST); assert_eq!(body["error"], "INVALID_ARGUMENT"); @@ -596,37 +565,12 @@ async fn unmappable_request_gets_the_shared_error_body() { assert!(body["message"].is_string()); } -#[tokio::test] -async fn unreachable_upstream_gets_the_shared_error_body() { - // Nothing listens on the upstream port: 503 UNAVAILABLE in the shared - // body. With details switched off for the route, the key is absent here - // too. - let app = common::proxy( - "http://127.0.0.1:1", - pool(), - ErrorDetailsPolicy::default() - .route("/v1/quiet/*", false) - .unwrap(), - ); - let (status, body) = get_json(&app, "/v1/things/rich").await; - assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE); - assert_eq!(body["error"], "UNAVAILABLE"); - assert_eq!(body["code"], 14); - assert_eq!(body["details"], json!([])); - - let (status, quiet) = get_json(&app, "/v1/quiet/rich").await; - assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE); - assert_eq!(quiet["code"], 14); - assert!(quiet.get("details").is_none(), "{quiet}"); -} - // --- streaming -------------------------------------------------------------- -#[tokio::test] async fn stream_refused_before_headers_maps_like_a_unary_error() { // No message was sent yet, so the proxy still owns the HTTP status: it is // mapped (PERMISSION_DENIED → 403) and the body is the unary error body. - let app = proxy(ErrorDetailsPolicy::default()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::default()).await; let (status, body) = get_json(&app, "/v1/things/x/denied").await; assert_eq!(status, StatusCode::FORBIDDEN); assert_eq!( @@ -644,13 +588,12 @@ async fn stream_refused_before_headers_maps_like_a_unary_error() { ); } -#[tokio::test] async fn ndjson_stream_failing_after_first_message_ends_with_detailed_error_line() { // The 200 and the first message are already on the wire when the upstream // fails, so the status cannot change: the error arrives as exactly one // final NDJSON line holding the same body a unary error would have, marked // by `@type: google.rpc.Status` so it is not mistaken for a data line. - let app = proxy(ErrorDetailsPolicy::default()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::default()).await; let (status, body) = get(&app, "/v1/things/x/watch", None).await; assert_eq!(status, StatusCode::OK); let lines: Vec = body @@ -672,11 +615,10 @@ async fn ndjson_stream_failing_after_first_message_ends_with_detailed_error_line ); } -#[tokio::test] async fn sse_stream_failing_after_first_message_ends_with_detailed_stream_error_event() { // Same failure over SSE: one data event, then exactly one `stream-error` // event with the full error body, and nothing after it. - let app = proxy(ErrorDetailsPolicy::default()).await; + let app = proxy(UPSTREAM, ErrorDetailsPolicy::default()).await; let (status, body) = get(&app, "/v1/things/x/watch", Some("text/event-stream")).await; assert_eq!(status, StatusCode::OK); let events: Vec<(Option<&str>, Value)> = body @@ -711,3 +653,57 @@ async fn sse_stream_failing_after_first_message_ends_with_detailed_stream_error_ ] ); } +} + +#[test] +fn yaml_route_rule_without_a_switch_is_rejected() { + // A rule that sets neither `enabled` nor `opaque` would match routes and + // change nothing, which is a mistake rather than an intent. + let err = structured_proxy::ProxyServer::from_yaml_str( + "upstream:\n default: \"http://127.0.0.1:1\"\nerror_details:\n routes:\n - pattern: \"/v1/**\"\n", + ) + .err() + .expect("a rule without a switch must be rejected"); + assert!( + err.to_string().contains("neither enabled nor opaque"), + "{err}" + ); +} + +#[test] +fn yaml_with_an_invalid_error_details_pattern_is_rejected() { + let err = structured_proxy::ProxyServer::from_yaml_str( + "upstream:\n default: \"http://127.0.0.1:1\"\nerror_details:\n routes:\n - pattern: \"v1/**\"\n enabled: false\n", + ) + .err() + .expect("a relative pattern must be rejected"); + assert!(err.to_string().contains("must start with '/'"), "{err}"); +} + +#[tokio::test] +async fn unreachable_upstream_gets_the_shared_error_body() { + // Nothing listens on the remote upstream's port: 503 UNAVAILABLE in the + // shared body. With details switched off for the route, the key is absent + // here too. + let server = structured_proxy::ProxyServer::from_yaml_str( + "upstream:\n default: \"http://127.0.0.1:1\"\n", + ) + .unwrap() + .with_descriptors(pool()) + .with_error_details( + ErrorDetailsPolicy::default() + .route("/v1/quiet/*", false) + .unwrap(), + ); + let app = common::App::new(server.service(server.upstream().unwrap()).unwrap()); + let (status, body) = get_json(&app, "/v1/things/rich").await; + assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(body["error"], "UNAVAILABLE"); + assert_eq!(body["code"], 14); + assert_eq!(body["details"], json!([])); + + let (status, quiet) = get_json(&app, "/v1/quiet/rich").await; + assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(quiet["code"], 14); + assert!(quiet.get("details").is_none(), "{quiet}"); +} diff --git a/tests/forwarded_headers.rs b/tests/forwarded_headers.rs index aba36f6..60bd4dc 100644 --- a/tests/forwarded_headers.rs +++ b/tests/forwarded_headers.rs @@ -1,6 +1,7 @@ //! Forwarded request headers reach a real tonic upstream as the client sent -//! them: every value, in order, over the actual HTTP/2 stream. +//! them: every value, in order, over the actual HTTP/2 stream and in process. +#[macro_use] mod common; use std::convert::Infallible; @@ -89,16 +90,21 @@ impl tower::Service> for Seen { } } -async fn proxy() -> axum::Router { - let upstream = common::serve(Seen { pool: pool() }).await; - common::proxy(&upstream, pool(), ErrorDetailsPolicy::default()) +async fn proxy(upstream: common::Upstream) -> common::App { + common::proxy( + upstream, + Seen { pool: pool() }, + pool(), + ErrorDetailsPolicy::default(), + ) + .await } -#[tokio::test] +upstream_tests! { async fn every_value_of_a_repeated_header_reaches_the_upstream_in_order() { // RFC 9449 §4.3 has the server reject a request with two DPoP headers; // behind the proxy it can only do that if both arrive. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::get("/v1/seen") .header("dpop", "proof-a") .header("dpop", "proof-b") @@ -112,12 +118,11 @@ async fn every_value_of_a_repeated_header_reaches_the_upstream_in_order() { ); } -#[tokio::test] async fn a_value_grpc_metadata_cannot_carry_is_refused() { // gRPC lets a receiver drop an ASCII metadata value outside %x20-%x7E, which // would change how many DPoP headers the upstream counts; the request is // refused instead of reaching it altered. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::get("/v1/seen") .header("dpop", "proof-a") .header("dpop", HeaderValue::from_bytes(b"caf\xe9").unwrap()) @@ -130,6 +135,21 @@ async fn a_value_grpc_metadata_cannot_carry_is_refused() { assert!(body["message"].as_str().unwrap().contains("dpop"), "{body}"); } +async fn a_single_value_is_unchanged() { + let app = proxy(UPSTREAM).await; + let request = http::Request::get("/v1/seen") + .header("dpop", "proof-a") + .body(Body::empty()) + .unwrap(); + let (status, body) = common::send(&app, request).await; + assert_eq!(status, StatusCode::OK, "{body}"); + assert_eq!( + serde_json::from_str::(&body).unwrap()["name"], + "proof-a" + ); +} +} + #[tokio::test] async fn a_forwarded_name_grpc_cannot_carry_is_a_config_error() { // `+` is a valid HTTP field-name character but not a gRPC key one, so an @@ -146,18 +166,3 @@ async fn a_forwarded_name_grpc_cannot_carry_is_a_config_error() { }; assert!(err.to_string().contains("x+proof"), "{err}"); } - -#[tokio::test] -async fn a_single_value_is_unchanged() { - let app = proxy().await; - let request = http::Request::get("/v1/seen") - .header("dpop", "proof-a") - .body(Body::empty()) - .unwrap(); - let (status, body) = common::send(&app, request).await; - assert_eq!(status, StatusCode::OK, "{body}"); - assert_eq!( - serde_json::from_str::(&body).unwrap()["name"], - "proof-a" - ); -} diff --git a/tests/request_mapping.rs b/tests/request_mapping.rs index 22c7428..67e906b 100644 --- a/tests/request_mapping.rs +++ b/tests/request_mapping.rs @@ -2,8 +2,10 @@ //! and reaches the upstream as that message. //! //! The upstream answers with the request it received, so each test sees what -//! actually reached the service. +//! actually reached the service. Every case runs against a remote and an +//! in-process upstream. +#[macro_use] mod common; use std::convert::Infallible; @@ -87,20 +89,19 @@ impl tower::Service> for Items { } } -async fn proxy() -> axum::Router { +async fn proxy(upstream: common::Upstream) -> common::App { let pool: DescriptorPool = common::compile("test/v1/items.proto", ITEMS_PROTO); let item = pool.get_message_by_name("test.v1.Item").unwrap(); - let upstream = common::serve(Items { item }).await; - common::proxy(&upstream, pool, Default::default()) + common::proxy(upstream, Items { item }, pool, Default::default()).await } -async fn get(app: &axum::Router, uri: &str) -> (StatusCode, Value) { +async fn get(app: &common::App, uri: &str) -> (StatusCode, Value) { let (status, body) = common::send(app, http::Request::get(uri).body(Body::empty()).unwrap()).await; (status, serde_json::from_str(&body).unwrap()) } -async fn post_form(app: &axum::Router, uri: &str, form: &'static str) -> (StatusCode, Value) { +async fn post_form(app: &common::App, uri: &str, form: &'static str) -> (StatusCode, Value) { let request = http::Request::post(uri) .header("content-type", "application/x-www-form-urlencoded") .body(Body::from(form)) @@ -109,11 +110,11 @@ async fn post_form(app: &axum::Router, uri: &str, form: &'static str) -> (Status (status, serde_json::from_str(&body).unwrap()) } -#[tokio::test] +upstream_tests! { async fn query_and_form_keys_bind_by_proto_or_json_name() { // ProtoJSON reads a field under either name; a query or form key does too, // rather than being dropped as unknown. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, body) = get(&app, "/v1/items/a?displayName=Ann&max_items=3").await; assert_eq!(status, StatusCode::OK, "{body}"); assert_eq!( @@ -128,11 +129,10 @@ async fn query_and_form_keys_bind_by_proto_or_json_name() { ); } -#[tokio::test] async fn well_known_type_fields_bind_one_by_one() { // A Timestamp, Duration or wrapper can be sent whole or field by field; // none of the client's values is lost. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, body) = get( &app, "/v1/items/a?at.seconds=5&at.nanos=7&ttl=1.5s¬e.value=hi", @@ -149,11 +149,10 @@ async fn well_known_type_fields_bind_one_by_one() { assert_eq!(body["note"], "hello"); } -#[tokio::test] async fn an_invalid_well_known_value_is_rejected_before_the_upstream() { // Field by field a client could write what no JSON form holds: nanos past // 999999999, or seconds and nanos of opposite sign. - let app = proxy().await; + let app = proxy(UPSTREAM).await; for uri in [ "/v1/items/a?at.nanos=2000000000", "/v1/items/a?ttl.seconds=5&ttl.nanos=-1", @@ -164,3 +163,4 @@ async fn an_invalid_well_known_value_is_rejected_before_the_upstream() { assert_eq!(body["code"], 3, "{uri}"); } } +} diff --git a/tests/streaming_request.rs b/tests/streaming_request.rs index befbbca..17d23e1 100644 --- a/tests/streaming_request.rs +++ b/tests/streaming_request.rs @@ -2,8 +2,10 @@ //! like unary routes: path parameters, query parameters and the `body` rule. //! //! The upstream echoes the request it received as the only stream message, so -//! each test sees what actually reached the service. +//! each test sees what actually reached the service. Every case runs against +//! a remote and an in-process upstream. +#[macro_use] mod common; use std::convert::Infallible; @@ -84,11 +86,10 @@ impl tower::Service> for Things { } } -async fn proxy() -> axum::Router { +async fn proxy(upstream: common::Upstream) -> common::App { let pool: DescriptorPool = common::compile("test/v1/things.proto", THINGS_PROTO); let item = pool.get_message_by_name("test.v1.Item").unwrap(); - let upstream = common::serve(Things { item }).await; - common::proxy(&upstream, pool, Default::default()) + common::proxy(upstream, Things { item }, pool, Default::default()).await } /// The NDJSON lines of a streaming response, parsed. @@ -98,11 +99,11 @@ fn ndjson(body: &str) -> Vec { .collect() } -#[tokio::test] +upstream_tests! { async fn get_stream_binds_path_and_query_parameters() { // `{name}` comes from the path and `count` from the query string; before // the fix the upstream received an empty request. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, body) = common::send( &app, http::Request::get("/v1/things/alpha/watch?count=7") @@ -114,10 +115,9 @@ async fn get_stream_binds_path_and_query_parameters() { assert_eq!(ndjson(&body), vec![json!({"name": "alpha", "count": "7"})]); } -#[tokio::test] async fn sse_stream_binds_the_request_too() { // The request mapping does not depend on the negotiated stream format. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, body) = common::send( &app, http::Request::get("/v1/things/alpha/watch?count=7") @@ -135,9 +135,8 @@ async fn sse_stream_binds_the_request_too() { assert_eq!(data, vec![json!({"name": "alpha", "count": "7"})]); } -#[tokio::test] async fn post_stream_maps_the_whole_body() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, body) = common::send( &app, http::Request::post("/v1/things:watch") @@ -150,11 +149,10 @@ async fn post_stream_maps_the_whole_body() { assert_eq!(ndjson(&body), vec![json!({"name": "beta", "count": "3"})]); } -#[tokio::test] async fn post_stream_with_malformed_body_is_rejected_before_the_upstream() { // A body that is not JSON is the client's error, answered with 400 like on // a unary route, instead of opening a stream with an empty request. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, body) = common::send( &app, http::Request::post("/v1/things:watch") @@ -170,10 +168,9 @@ async fn post_stream_with_malformed_body_is_rejected_before_the_upstream() { assert_eq!(error["details"], json!([])); } -#[tokio::test] async fn get_stream_with_ill_typed_query_is_rejected_before_the_upstream() { // `count` is an int64: a non-numeric value cannot build the request. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, body) = common::send( &app, http::Request::get("/v1/things/alpha/watch?count=many") @@ -187,3 +184,4 @@ async fn get_stream_with_ill_typed_query_is_rejected_before_the_upstream() { assert_eq!(error["code"], 3); assert_eq!(error["details"], json!([])); } +} diff --git a/tests/tls.rs b/tests/tls.rs new file mode 100644 index 0000000..9b2fcd8 --- /dev/null +++ b/tests/tls.rs @@ -0,0 +1,376 @@ +//! An embedder serves the proxy behind its own TLS: a rustls acceptor and +//! hyper's HTTP/1.1 + HTTP/2 connection, with the proxy as the service. REST +//! and native gRPC share the TLS port, and a tonic upstream in process reads +//! the client's address and TLS certificate as behind tonic's own server. + +#[path = "common/protos.rs"] +mod protos; + +use std::convert::Infallible; +use std::future::{ready, Ready}; +use std::net::SocketAddr; +use std::pin::Pin; +use std::sync::Arc; +use std::task::{Context, Poll}; + +use base64::Engine as _; +use hyper_util::rt::{TokioExecutor, TokioIo}; +use prost_reflect::{DescriptorPool, DynamicMessage, Value as PbValue}; +use rustls::client::danger::HandshakeSignatureValid; +use rustls::pki_types::pem::PemObject; +use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName, UnixTime}; +use rustls::server::danger::{ClientCertVerified, ClientCertVerifier}; +use rustls::{DigitallySignedStruct, DistinguishedName, SignatureScheme}; +use serde_json::Value; +use structured_proxy::service::TlsConnectInfo; +use structured_proxy::transcode::codec::DynamicCodec; +use structured_proxy::{ConnectionInfo, ProxyServer}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tonic::transport::server::Connected; + +const CA: &str = include_str!("../src/tls/testdata/ca.pem"); +const CERT: &str = include_str!("../src/tls/testdata/ecdsa.pem"); +const KEY: &str = include_str!("../src/tls/testdata/ecdsa.key.pem"); + +const WHO_PROTO: &str = r#" +syntax = "proto3"; +package test.v1; +import "google/api/annotations.proto"; + +message Req {} +// The connection the upstream saw the call on. +message Seen { + string peer = 1; + bytes cert = 2; +} + +service Who { + rpc Me(Req) returns (Seen) { + option (google.api.http) = { get: "/v1/me" }; + } +} +"#; + +fn pool() -> DescriptorPool { + protos::compile("test/v1/who.proto", WHO_PROTO) +} + +// --- upstream --------------------------------------------------------------- + +/// Answers with the caller's address and the first certificate it presented. +#[derive(Clone)] +struct Me { + pool: DescriptorPool, +} + +impl tonic::server::UnaryService for Me { + type Response = DynamicMessage; + type Future = Ready, tonic::Status>>; + + fn call(&mut self, request: tonic::Request) -> Self::Future { + let peer = request + .remote_addr() + .map(|addr| addr.to_string()) + .unwrap_or_default(); + // What `Request::peer_certs` reads, with tonic's TLS feature on. + let cert = request + .extensions() + .get::() + .and_then(|tls| tls.peer_certs()) + .and_then(|chain| chain.first().map(|c| c.as_ref().to_vec())) + .unwrap_or_default(); + let mut seen = DynamicMessage::new(self.pool.get_message_by_name("test.v1.Seen").unwrap()); + seen.set_field_by_name("peer", PbValue::String(peer)); + seen.set_field_by_name("cert", PbValue::Bytes(cert.into())); + ready(Ok(tonic::Response::new(seen))) + } +} + +#[derive(Clone)] +struct Who { + pool: DescriptorPool, +} + +impl tonic::server::NamedService for Who { + const NAME: &'static str = "test.v1.Who"; +} + +impl tower::Service> for Who { + type Response = http::Response; + type Error = Infallible; + type Future = + Pin> + Send>>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, req: http::Request) -> Self::Future { + let pool = self.pool.clone(); + Box::pin(async move { + let input = pool.get_message_by_name("test.v1.Req").unwrap(); + let mut grpc = tonic::server::Grpc::new(DynamicCodec::new(input)); + Ok(grpc.unary(Me { pool }, req).await) + }) + } +} + +// --- the embedder's TLS server ------------------------------------------------ + +fn provider() -> Arc { + Arc::new(rustls_rustcrypto::provider()) +} + +fn chain() -> Vec> { + CertificateDer::pem_slice_iter(CERT.as_bytes()) + .collect::>() + .unwrap() +} + +fn key() -> PrivateKeyDer<'static> { + PrivateKeyDer::from_pem_slice(KEY.as_bytes()).unwrap() +} + +/// Asks for a client certificate and takes any that proves its key: these +/// tests check that the certificate reaches the upstream, not how an embedder +/// decides to trust it. +#[derive(Debug)] +struct AnyClientCert(Arc); + +impl ClientCertVerifier for AnyClientCert { + fn client_auth_mandatory(&self) -> bool { + false + } + + fn root_hint_subjects(&self) -> &[DistinguishedName] { + &[] + } + + fn verify_client_cert( + &self, + _end_entity: &CertificateDer<'_>, + _intermediates: &[CertificateDer<'_>], + _now: UnixTime, + ) -> Result { + Ok(ClientCertVerified::assertion()) + } + + fn verify_tls12_signature( + &self, + message: &[u8], + cert: &CertificateDer<'_>, + dss: &DigitallySignedStruct, + ) -> Result { + rustls::crypto::verify_tls12_signature( + message, + cert, + dss, + &self.0.signature_verification_algorithms, + ) + } + + fn verify_tls13_signature( + &self, + message: &[u8], + cert: &CertificateDer<'_>, + dss: &DigitallySignedStruct, + ) -> Result { + rustls::crypto::verify_tls13_signature( + message, + cert, + dss, + &self.0.signature_verification_algorithms, + ) + } + + fn supported_verify_schemes(&self) -> Vec { + self.0.signature_verification_algorithms.supported_schemes() + } +} + +/// Serve the proxy in front of `Who`, in process, on a local TLS port, the +/// way an embedder with its own TLS does: accept, handshake, tell the proxy +/// the connection, serve HTTP/1.1 or HTTP/2 as the client speaks. +async fn listen() -> SocketAddr { + let pool = pool(); + let proxy = ProxyServer::from_yaml_str("") + .unwrap() + .with_descriptors(pool.clone()) + .service(tonic::service::Routes::new(Who { pool })) + .unwrap(); + let mut config = rustls::ServerConfig::builder_with_provider(provider()) + .with_safe_default_protocol_versions() + .unwrap() + .with_client_cert_verifier(Arc::new(AnyClientCert(provider()))) + .with_single_cert(chain(), key()) + .unwrap(); + // gRPC clients negotiate HTTP/2, browsers and curl may take HTTP/1.1. + config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()]; + let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(config)); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + loop { + let (tcp, _) = listener.accept().await.unwrap(); + let acceptor = acceptor.clone(); + let proxy = proxy.clone(); + tokio::spawn(async move { + let Ok(tls) = acceptor.accept(tcp).await else { + return; + }; + let service = proxy.for_connection(ConnectionInfo::tls(tls.connect_info())); + let served = hyper_util::server::conn::auto::Builder::new(TokioExecutor::new()) + .serve_connection( + TokioIo::new(tls), + hyper_util::service::TowerToHyperService::new(service), + ) + .await; + // A client dropping its connection ends it with an error; the + // cases check what the client received. + if let Err(error) = served { + eprintln!("connection ended: {error}"); + } + }); + } + }); + addr +} + +// --- clients -------------------------------------------------------------------- + +/// A TLS connection to `addr` trusting the test CA, offering `alpn`, and +/// presenting the test certificate when `with_cert`. +async fn connect( + addr: SocketAddr, + alpn: &[u8], + with_cert: bool, +) -> tokio_rustls::client::TlsStream { + let mut roots = rustls::RootCertStore::empty(); + for cert in CertificateDer::pem_slice_iter(CA.as_bytes()) { + roots.add(cert.unwrap()).unwrap(); + } + let builder = rustls::ClientConfig::builder_with_provider(provider()) + .with_safe_default_protocol_versions() + .unwrap() + .with_root_certificates(roots); + let mut config = if with_cert { + builder.with_client_auth_cert(chain(), key()).unwrap() + } else { + builder.with_no_client_auth() + }; + config.alpn_protocols = vec![alpn.to_vec()]; + let tcp = tokio::net::TcpStream::connect(addr).await.unwrap(); + tokio_rustls::TlsConnector::from(Arc::new(config)) + .connect(ServerName::try_from("localhost").unwrap(), tcp) + .await + .unwrap() +} + +/// `GET /v1/me` over HTTP/1.1 and TLS; returns the status, the JSON body and +/// the client's own address. +async fn rest_me(addr: SocketAddr, with_cert: bool) -> (u16, Value, SocketAddr) { + let mut tls = connect(addr, b"http/1.1", with_cert).await; + let client = tls.get_ref().0.local_addr().unwrap(); + tls.write_all(b"GET /v1/me HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n") + .await + .unwrap(); + let mut response = Vec::new(); + match tls.read_to_end(&mut response).await { + Ok(_) => {} + // A server that closes without close_notify still delivered the + // response. + Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => {} + Err(e) => panic!("reading the response: {e}"), + } + let response = String::from_utf8(response).unwrap(); + let status = response[9..12].parse().unwrap(); + let body = response.split_once("\r\n\r\n").unwrap().1; + (status, serde_json::from_str(body).unwrap(), client) +} + +/// A native `Me` call over HTTP/2 and TLS; returns what the upstream saw and +/// the client's own address. +async fn grpc_me(addr: SocketAddr, with_cert: bool) -> (DynamicMessage, SocketAddr) { + let (sender, mut client_addr) = tokio::sync::mpsc::channel(1); + let connector = tower::service_fn(move |_: http::Uri| { + let sender = sender.clone(); + async move { + let tls = connect(addr, b"h2", with_cert).await; + sender + .send(tls.get_ref().0.local_addr().unwrap()) + .await + .unwrap(); + Ok::<_, std::io::Error>(TokioIo::new(tls)) + } + }); + let channel = tonic::transport::Endpoint::from_static("http://localhost") + .connect_with_connector(connector) + .await + .unwrap(); + let mut grpc = tonic::client::Grpc::new(channel); + grpc.ready().await.unwrap(); + let pool = pool(); + let req = DynamicMessage::new(pool.get_message_by_name("test.v1.Req").unwrap()); + let codec = DynamicCodec::new(pool.get_message_by_name("test.v1.Seen").unwrap()); + let seen = grpc + .unary( + tonic::Request::new(req), + http::uri::PathAndQuery::from_static("/test.v1.Who/Me"), + codec, + ) + .await + .unwrap() + .into_inner(); + (seen, client_addr.recv().await.unwrap()) +} + +fn leaf_der() -> Vec { + chain()[0].as_ref().to_vec() +} + +// --- cases -------------------------------------------------------------------- + +#[tokio::test] +async fn rest_over_tls_reaches_the_upstream_with_the_client_address() { + let addr = listen().await; + let (status, seen, client) = rest_me(addr, false).await; + assert_eq!(status, 200, "{seen}"); + assert_eq!(seen["peer"], client.to_string()); + // No certificate presented, none reported. + assert_eq!(seen["cert"], ""); +} + +#[tokio::test] +async fn native_grpc_shares_the_tls_port() { + let addr = listen().await; + let (seen, client) = grpc_me(addr, false).await; + let peer = match seen.get_field_by_name("peer").as_deref() { + Some(PbValue::String(peer)) => peer.clone(), + _ => String::new(), + }; + assert_eq!(peer, client.to_string()); +} + +#[tokio::test] +async fn a_client_certificate_reaches_a_transcoded_call() { + // mTLS: the upstream authorizes on the certificate as behind tonic's own + // TLS server, although the call came in as REST. + let addr = listen().await; + let (status, seen, _) = rest_me(addr, true).await; + assert_eq!(status, 200, "{seen}"); + let cert = base64::engine::general_purpose::STANDARD + .decode(seen["cert"].as_str().unwrap()) + .unwrap(); + assert_eq!(cert, leaf_der()); +} + +#[tokio::test] +async fn a_client_certificate_reaches_a_native_call() { + let addr = listen().await; + let (seen, _) = grpc_me(addr, true).await; + let cert = match seen.get_field_by_name("cert").as_deref() { + Some(PbValue::Bytes(cert)) => cert.to_vec(), + _ => Vec::new(), + }; + assert_eq!(cert, leaf_der()); +} diff --git a/tests/upstream_controls.rs b/tests/upstream_controls.rs index 558c789..cdbb87f 100644 --- a/tests/upstream_controls.rs +++ b/tests/upstream_controls.rs @@ -4,9 +4,11 @@ //! any method. //! //! Runs the proxy (through its public `ProxyServer`) in front of a real tonic -//! gRPC server, so metadata, trailers and trailers-only errors go over the -//! actual HTTP/2 stream. +//! gRPC service, reached over an actual HTTP/2 stream and in process, so every +//! case holds for both kinds of upstream: metadata, trailers and trailers-only +//! errors included. +#[macro_use] mod common; use std::convert::Infallible; @@ -439,32 +441,43 @@ impl tower::Service> for Controls { // --- proxy harness ---------------------------------------------------------- -/// The proxy router in front of a fresh upstream, built by `configure`. -async fn proxy_with(configure: impl FnOnce(ProxyServer) -> ProxyServer) -> axum::Router { +/// The proxy in front of a fresh upstream, built by `configure`. +async fn proxy_with( + upstream: common::Upstream, + configure: impl FnOnce(ProxyServer) -> ProxyServer, +) -> common::App { let pool = pool(); - let upstream = common::serve(Controls { pool: pool.clone() }).await; - let server = ProxyServer::from_yaml_str(&format!("upstream:\n default: \"{upstream}\"\n")) - .unwrap() - .with_descriptors(pool); - configure(server).router().unwrap() + common::app(upstream, Controls { pool: pool.clone() }, |yaml| { + configure( + ProxyServer::from_yaml_str(yaml) + .unwrap() + .with_descriptors(pool), + ) + }) + .await } -/// The proxy router in front of a fresh upstream, with default settings. -async fn proxy() -> axum::Router { +/// The proxy in front of a fresh upstream, with default settings. +async fn proxy(upstream: common::Upstream) -> common::App { let pool = pool(); - let upstream = common::serve(Controls { pool: pool.clone() }).await; - common::proxy(&upstream, pool, ErrorDetailsPolicy::default()) + common::proxy( + upstream, + Controls { pool: pool.clone() }, + pool, + ErrorDetailsPolicy::default(), + ) + .await } /// Send `request`; returns the status, the headers and the raw body. -async fn send(app: &axum::Router, request: http::Request) -> (StatusCode, HeaderMap, Bytes) { +async fn send(app: &common::App, request: http::Request) -> (StatusCode, HeaderMap, Bytes) { let resp = app.clone().oneshot(request).await.unwrap(); let (parts, body) = resp.into_parts(); let body = axum::body::to_bytes(body, usize::MAX).await.unwrap(); (parts.status, parts.headers, body) } -async fn call(app: &axum::Router, method: Method, path: &str) -> (StatusCode, HeaderMap, Bytes) { +async fn call(app: &common::App, method: Method, path: &str) -> (StatusCode, HeaderMap, Bytes) { let request = http::Request::builder() .method(method) .uri(path) @@ -473,6 +486,18 @@ async fn call(app: &axum::Router, method: Method, path: &str) -> (StatusCode, He send(app, request).await } +async fn post_form( + app: &common::App, + path: &str, + form: &'static str, +) -> (StatusCode, HeaderMap, Bytes) { + let request = http::Request::post(path) + .header("content-type", "application/x-www-form-urlencoded") + .body(Body::from(form)) + .unwrap(); + send(app, request).await +} + fn values<'a>(headers: &'a HeaderMap, name: &str) -> Vec<&'a str> { headers .get_all(name) @@ -485,11 +510,11 @@ fn json_body(body: &Bytes) -> Value { serde_json::from_slice(body).unwrap() } +upstream_tests! { // --- response metadata → headers --------------------------------------------- -#[tokio::test] async fn unary_metadata_becomes_response_headers() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/things/meta").await; assert_eq!(status, StatusCode::OK); assert_eq!(json_body(&body), json!({"name": "meta"})); @@ -505,9 +530,8 @@ async fn unary_metadata_becomes_response_headers() { assert_eq!(values(&headers, "content-type"), ["application/json"]); } -#[tokio::test] async fn plain_answer_adds_no_upstream_headers() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/things/plain").await; assert_eq!(status, StatusCode::OK); assert_eq!(json_body(&body), json!({"name": "plain"})); @@ -515,11 +539,10 @@ async fn plain_answer_adds_no_upstream_headers() { assert!(headers.get("x-http-code").is_none()); } -#[tokio::test] async fn trailers_of_a_successful_call_become_headers_after_initial_metadata() { // A key in both the initial metadata and the trailers keeps both values, // initial first; gRPC's trailer keys stay behind. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/trailers").await; assert_eq!(status, StatusCode::OK); assert_eq!(json_body(&body), json!({"name": "trailed"})); @@ -528,9 +551,8 @@ async fn trailers_of_a_successful_call_become_headers_after_initial_metadata() { assert!(headers.get("grpc-status").is_none()); } -#[tokio::test] async fn deny_list_from_the_builder_drops_its_keys() { - let app = proxy_with(|server| { + let app = proxy_with(UPSTREAM, |server| { server.with_denied_response_headers([HeaderName::from_static("x-debug")]) }) .await; @@ -539,27 +561,25 @@ async fn deny_list_from_the_builder_drops_its_keys() { assert_eq!(values(&headers, "cache-control"), ["no-store"]); } -#[tokio::test] async fn deny_list_from_yaml_drops_its_keys() { let pool = pool(); - let upstream = common::serve(Controls { pool: pool.clone() }).await; - let app = ProxyServer::from_yaml_str(&format!( - "upstream:\n default: \"{upstream}\"\nresponse_headers:\n deny: [\"X-Debug\", \"set-cookie\"]\n" - )) - .unwrap() - .with_descriptors(pool) - .router() - .unwrap(); + let app = common::app(UPSTREAM, Controls { pool: pool.clone() }, |yaml| { + ProxyServer::from_yaml_str(&format!( + "{yaml}response_headers:\n deny: [\"X-Debug\", \"set-cookie\"]\n" + )) + .unwrap() + .with_descriptors(pool) + }) + .await; let (_, headers, _) = call(&app, Method::GET, "/v1/things/meta").await; assert!(headers.get("x-debug").is_none()); assert!(headers.get("set-cookie").is_none()); assert_eq!(values(&headers, "cache-control"), ["no-store"]); } -#[tokio::test] async fn trailers_only_error_carries_its_metadata() { // RFC 6750 §3: the 401 carries the upstream's `WWW-Authenticate`. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/things/unauth").await; assert_eq!(status, StatusCode::UNAUTHORIZED); assert_eq!( @@ -570,12 +590,11 @@ async fn trailers_only_error_carries_its_metadata() { assert_eq!(values(&headers, "content-type"), ["application/json"]); } -#[tokio::test] async fn unary_error_after_headers_carries_initial_metadata_and_trailers() { // The upstream sent response headers, then failed in its trailers: the // error keeps both, initial metadata first, as tonic's own unary call // does. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/late-error").await; assert_eq!(status, StatusCode::UNAUTHORIZED); assert_eq!(values(&headers, "x-initial"), ["1"]); @@ -586,11 +605,10 @@ async fn unary_error_after_headers_carries_initial_metadata_and_trailers() { assert_eq!(json_body(&body)["message"], "expired"); } -#[tokio::test] async fn http_body_stream_failing_before_its_first_message_keeps_initial_metadata() { // Headers were sent before the failure, so they belong to the error // response, next to the failure's own metadata. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, _) = call(&app, Method::GET, "/v1/files/late-denied").await; assert_eq!(status, StatusCode::UNAUTHORIZED); assert_eq!(values(&headers, "x-stream"), ["1"]); @@ -598,11 +616,10 @@ async fn http_body_stream_failing_before_its_first_message_keeps_initial_metadat assert!(values(&headers, "www-authenticate")[0].starts_with("Bearer error=")); } -#[tokio::test] async fn malformed_error_status_carries_no_upstream_metadata() { // The answer is the generic INTERNAL, so nothing of the broken error, // headers included, reaches the client. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/things/corrupt-unauth").await; assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); assert_eq!(json_body(&body)["error"], "INTERNAL"); @@ -611,9 +628,8 @@ async fn malformed_error_status_carries_no_upstream_metadata() { // --- x-http-code ------------------------------------------------------------ -#[tokio::test] async fn http_code_sets_the_status_of_a_successful_call() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/things/created").await; assert_eq!(status, StatusCode::CREATED); assert_eq!(values(&headers, "location"), ["/v1/things/created"]); @@ -621,19 +637,17 @@ async fn http_code_sets_the_status_of_a_successful_call() { assert_eq!(json_body(&body), json!({"name": "created"})); } -#[tokio::test] async fn http_code_204_answers_without_content() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/things/no-content").await; assert_eq!(status, StatusCode::NO_CONTENT); assert!(body.is_empty()); assert!(headers.get("content-type").is_none()); } -#[tokio::test] async fn invalid_http_code_is_a_malformed_upstream_internal() { // Never a partial response: no other upstream header rides along. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/things/bad-code").await; assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); assert_eq!( @@ -649,10 +663,9 @@ async fn invalid_http_code_is_a_malformed_upstream_internal() { assert!(headers.get("x-http-code").is_none()); } -#[tokio::test] async fn redirect_with_location_and_an_empty_body() { // RFC 6749 §4.1.2: the authorization endpoint answers 302 + Location. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/authorize").await; assert_eq!(status, StatusCode::FOUND); assert_eq!( @@ -665,23 +678,10 @@ async fn redirect_with_location_and_an_empty_body() { // --- google.api.HttpBody ---------------------------------------------------- -async fn post_form( - app: &axum::Router, - path: &str, - form: &'static str, -) -> (StatusCode, HeaderMap, Bytes) { - let request = http::Request::post(path) - .header("content-type", "application/x-www-form-urlencoded") - .body(Body::from(form)) - .unwrap(); - send(app, request).await -} - -#[tokio::test] async fn token_error_is_the_rfc_6749_body_with_400_and_no_store() { // RFC 6749 §5.2: 400, the upstream's own JSON body, `Cache-Control: // no-store` (§5.1). Not the transcoder's error body. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = post_form(&app, "/v1/token", "grant_type=refresh_token").await; assert_eq!(status, StatusCode::BAD_REQUEST); assert_eq!(values(&headers, "cache-control"), ["no-store"]); @@ -696,9 +696,8 @@ async fn token_error_is_the_rfc_6749_body_with_400_and_no_store() { ); } -#[tokio::test] async fn token_success_is_the_raw_json_with_no_store() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = post_form(&app, "/v1/token", "grant_type=authorization_code").await; assert_eq!(status, StatusCode::OK); @@ -706,10 +705,9 @@ async fn token_success_is_the_raw_json_with_no_store() { assert_eq!(&body[..], br#"{"access_token":"at","token_type":"Bearer"}"#); } -#[tokio::test] async fn http_body_response_keeps_the_content_type_and_bytes() { // RFC 7517 §8.5: a JWK Set is served as application/jwk-set+json. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/jwks").await; assert_eq!(status, StatusCode::OK); assert_eq!( @@ -719,11 +717,10 @@ async fn http_body_response_keeps_the_content_type_and_bytes() { assert_eq!(&body[..], br#"{"keys":[]}"#); } -#[tokio::test] async fn http_body_request_receives_the_raw_body_and_content_type() { // Bytes that are neither JSON nor UTF-8 arrive untouched, with the full // Content-Type value (parameters included). - let app = proxy().await; + let app = proxy(UPSTREAM).await; let raw: &'static [u8] = b"\x89PNG\r\n\x1a\n\x00\xff"; let request = http::Request::post("/v1/echo") .header("content-type", "image/png; name=x") @@ -735,12 +732,11 @@ async fn http_body_request_receives_the_raw_body_and_content_type() { assert_eq!(&body[..], raw); } -#[tokio::test] async fn query_on_a_whole_message_http_body_route_is_ignored() { // With `body: "*"` on an HttpBody input every field comes from the body // (google/api/http.proto: no HTTP parameters with `*`), so a query that // names an HttpBody field is ignored instead of failing the request. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::post("/v1/echo?extensions=x&content_type=y&data=z") .header("content-type", "text/plain") .body(Body::from("raw")) @@ -751,9 +747,8 @@ async fn query_on_a_whole_message_http_body_route_is_ignored() { assert_eq!(&body[..], b"raw"); } -#[tokio::test] async fn http_body_request_without_a_body_is_empty() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::post("/v1/echo").body(Body::empty()).unwrap(); let (status, headers, body) = send(&app, request).await; assert_eq!(status, StatusCode::OK); @@ -761,9 +756,8 @@ async fn http_body_request_without_a_body_is_empty() { assert!(headers.get("content-type").is_none()); } -#[tokio::test] async fn http_body_field_receives_the_raw_body_next_to_path_fields() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::put("/v1/uploads/report") .header("content-type", "text/csv") .body(Body::from("a,b")) @@ -776,12 +770,11 @@ async fn http_body_field_receives_the_raw_body_next_to_path_fields() { ); } -#[tokio::test] async fn query_key_naming_the_raw_body_field_does_not_break_the_upload() { // The body binds `file`, and the body wins over the query: a `file` // query parameter (or one under it) is ignored rather than bound into // the HttpBody field before the raw body replaces it. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::put("/v1/uploads/report?file=x&file.content_type=y") .header("content-type", "text/csv") .body(Body::from("a,b")) @@ -791,11 +784,10 @@ async fn query_key_naming_the_raw_body_field_does_not_break_the_upload() { assert_eq!(json_body(&body), json!({"name": "report|text/csv|a,b"})); } -#[tokio::test] async fn http_body_request_with_a_non_ascii_content_type_is_rejected() { // HttpBody.content_type is a proto string; bytes that are not visible // ASCII are refused before the upstream is called. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::post("/v1/echo") .header( "content-type", @@ -810,9 +802,8 @@ async fn http_body_request_with_a_non_ascii_content_type_is_rejected() { // --- server streaming ------------------------------------------------------- -#[tokio::test] async fn streaming_initial_metadata_becomes_headers() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/things/x/watch").await; assert_eq!(status, StatusCode::OK); assert_eq!(values(&headers, "x-stream"), ["1"]); @@ -825,11 +816,10 @@ async fn streaming_initial_metadata_becomes_headers() { assert_eq!(lines, [json!({"name": "one"}), json!({"name": "two"})]); } -#[tokio::test] async fn sse_keeps_its_own_cache_control_over_the_upstream_one() { // What the proxy writes describes the body it writes: SSE must not be // cached, whatever the upstream asked for. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::get("/v1/things/x/watch") .header("accept", "text/event-stream") .body(Body::empty()) @@ -841,18 +831,16 @@ async fn sse_keeps_its_own_cache_control_over_the_upstream_one() { assert_eq!(values(&headers, "x-stream"), ["1"]); } -#[tokio::test] async fn refused_stream_carries_its_metadata() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, _) = call(&app, Method::GET, "/v1/things/denied/watch").await; assert_eq!(status, StatusCode::UNAUTHORIZED); assert!(values(&headers, "www-authenticate")[0].starts_with("Bearer error=\"invalid_token\"")); } -#[tokio::test] async fn streaming_http_body_is_chunked_raw_data() { // Content type from the first message; every message's data in order. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/files/report").await; assert_eq!(status, StatusCode::OK); assert_eq!(values(&headers, "content-type"), ["text/csv"]); @@ -860,9 +848,8 @@ async fn streaming_http_body_is_chunked_raw_data() { assert_eq!(&body[..], b"a,b\n1,2\n3,4\n"); } -#[tokio::test] async fn streaming_http_body_ignores_sse_negotiation() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::get("/v1/files/report") .header("accept", "text/event-stream") .body(Body::empty()) @@ -872,20 +859,18 @@ async fn streaming_http_body_ignores_sse_negotiation() { assert_eq!(&body[..], b"a,b\n1,2\n3,4\n"); } -#[tokio::test] async fn empty_streaming_http_body_is_an_empty_ok() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::GET, "/v1/files/empty").await; assert_eq!(status, StatusCode::OK); assert!(body.is_empty()); assert!(headers.get("content-type").is_none()); } -#[tokio::test] async fn streaming_http_body_failing_mid_stream_aborts_the_body() { // A raw body has no in-band error frame: the transfer is cut short so the // client cannot take the partial file for a complete one. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::get("/v1/files/broken") .body(Body::empty()) .unwrap(); @@ -898,9 +883,8 @@ async fn streaming_http_body_failing_mid_stream_aborts_the_body() { // --- custom rules ------------------------------------------------------------ -#[tokio::test] async fn custom_head_rule_routes_head_to_the_rpc() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let (status, headers, body) = call(&app, Method::HEAD, "/v1/probe/disk").await; assert_eq!(status, StatusCode::OK); assert_eq!(values(&headers, "x-probe"), ["disk"]); @@ -910,10 +894,9 @@ async fn custom_head_rule_routes_head_to_the_rpc() { assert_eq!(status, StatusCode::METHOD_NOT_ALLOWED); } -#[tokio::test] async fn star_rule_routes_every_method_to_the_rpc() { // A forward-auth sub-request arrives with the original request's method. - let app = proxy().await; + let app = proxy(UPSTREAM).await; for method in [ Method::GET, Method::POST, @@ -930,9 +913,8 @@ async fn star_rule_routes_every_method_to_the_rpc() { } } -#[tokio::test] async fn extension_method_rule_routes_next_to_a_standard_one() { - let app = proxy().await; + let app = proxy(UPSTREAM).await; let propfind = Method::from_bytes(b"PROPFIND").unwrap(); let (status, _, body) = call(&app, propfind, "/v1/dav").await; assert_eq!(status, StatusCode::OK); @@ -949,11 +931,10 @@ async fn extension_method_rule_routes_next_to_a_standard_one() { assert_eq!(status, StatusCode::METHOD_NOT_ALLOWED); } -#[tokio::test] async fn unbound_method_is_405_before_its_body_is_read() { // The method decides first: a method nobody binds on the path is 405 // even with a body over the extractor limit, which is never buffered. - let app = proxy().await; + let app = proxy(UPSTREAM).await; let request = http::Request::builder() .method(Method::DELETE) .uri("/v1/dav") @@ -963,6 +944,7 @@ async fn unbound_method_is_405_before_its_body_is_read() { assert_eq!(status, StatusCode::METHOD_NOT_ALLOWED); assert_eq!(values(&headers, "allow"), ["PROPFIND, GET, HEAD"]); } +} #[tokio::test] async fn star_rule_collides_with_another_method_on_its_path() {