From 94c7980197833c6e0d3418d2e6d7fef79b1fcaeb Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 10:54:14 +0300 Subject: [PATCH 1/9] feat(guards)!: scope guards and reject in the request's protocol - Every guard (maintenance, rate limits, JWT, ext_authz, the auth decider) takes a scope: the traffic it covers (transcoded, endpoints, grpc, fallback, all), optionally narrowed by path globs and methods. Guards are mounted per traffic class, so a request runs only the guards of its own class; native gRPC and the fallback can now be guarded too. Defaults keep the previous coverage. - New concurrency limit guard (`concurrency.max_in_flight`): excess requests get 503 UNAVAILABLE with Retry-After at once; a request holds its slot until its response body ends. - All guards reject through one path: a REST client gets the google.rpc.Status JSON body with the mapped HTTP status, a gRPC or gRPC-Web client a trailers-only status with the same code and the guard's headers as metadata. - `ProxyServer::with_auth_decider_scope` chooses the decider's traffic. - A scope with no traffic, an invalid glob or method fails the build. BREAKING CHANGE: guard rejections use the google.rpc.Status body ({"error","code","message","details"}) instead of {"error","message"}; maintenance answers that JSON body instead of plain text; an unmatched path under maintenance is 404, since the fallback class is outside the default scope. Closes #120 --- Cargo.toml | 3 +- README.md | 103 ++++++++- src/auth/authz.rs | 39 ++-- src/auth/forward/tests.rs | 1 + src/auth/mod.rs | 21 +- src/auth/tests.rs | 13 ++ src/config.rs | 99 ++++++++ src/config/tests.rs | 4 +- src/embed.rs | 19 +- src/guard.rs | 475 ++++++++++++++++++++++++++++++++++++++ src/guard/concurrency.rs | 107 +++++++++ src/guard/grpc.rs | 105 +++++++++ src/guard/tests.rs | 320 +++++++++++++++++++++++++ src/lib.rs | 323 +++++++++++++------------- src/service.rs | 167 +++++++++++--- src/service/tests.rs | 14 +- src/shield/mod.rs | 14 +- src/shield/tests.rs | 2 + src/tests.rs | 34 --- src/transcode/error.rs | 6 + src/upstream.rs | 11 +- tests/edge.rs | 129 ++++++++++- tests/embedded.rs | 1 + 23 files changed, 1715 insertions(+), 295 deletions(-) create mode 100644 src/guard.rs create mode 100644 src/guard/concurrency.rs create mode 100644 src/guard/grpc.rs create mode 100644 src/guard/tests.rs diff --git a/Cargo.toml b/Cargo.toml index 14d39df..41f0bc7 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -35,7 +35,8 @@ doc = false # 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" +# `util`: the guarded gRPC and fallback paths are boxed services. +tower = { version = "0.5", features = ["util"] } tower-http = { version = "0.7", features = ["cors", "trace"] } # Foundational HTTP types, used directly by the framework-agnostic embedding # hooks (src/hooks.rs) so an embedder never names `axum`. Already in the tree diff --git a/README.md b/README.md index a7555a6..896a6bc 100644 --- a/README.md +++ b/README.md @@ -26,7 +26,9 @@ Works with **any** gRPC service via proto descriptor files. No code generation, - **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/*`) +- **Scoped guards**: maintenance, rate limits, a concurrency limit, JWT, ext_authz and the auth decider each cover the traffic you name (transcoded, the proxy's own endpoints, native gRPC, the fallback), narrowed by path and method; a rejection answers in the request's protocol, a `google.rpc.Status` JSON body for REST and a gRPC status for gRPC (see [Guards and scopes](#guards-and-scopes)) - **Maintenance mode** returning 503 with a configurable exempt-path list +- **Concurrency limit**: requests past `max_in_flight` are shed at once with 503 instead of queueing; a stream holds its slot until its body ends - **Health endpoints** `/health/live`, `/health/ready` (upstream gRPC health probe), `/health/startup` - **Prometheus metrics** at `/metrics` - **CORS** with a configurable origin allow-list, exposed headers and preflight cache, applied to gRPC-Web pass-through too @@ -133,6 +135,16 @@ runtime: maintenance: enabled: false message: "Service is under maintenance. Please try again later." + # Which traffic it turns away (see "Guards and scopes"). Default: + # [transcoded, endpoints]. + # scope: { traffic: [all] } + +# Optional: request concurrency limit. Requests past max_in_flight get 503 +# UNAVAILABLE with `Retry-After: 1` at once; a request holds its slot until its +# response body ends. Default scope: [transcoded, endpoints, grpc]. +concurrency: + max_in_flight: 512 + # scope: { traffic: [grpc], paths: ["/acme.v1.Orders/*"] } # Optional: server-streaming response behavior. # Streaming RPCs return NDJSON by default; clients sending @@ -177,6 +189,8 @@ response_headers: # the rest run before auth so anonymous floods are shed cheaply. shield: enabled: true + # Which traffic the rules apply to. Default: [transcoded, endpoints]. + # scope: { traffic: [transcoded, endpoints, grpc] } # CIDR ranges of trusted proxies/LBs. X-Forwarded-For is honored only from # these peers; set this behind a load balancer for correct per-client limits. trusted_proxies: ["10.0.0.0/8"] @@ -206,6 +220,10 @@ shield: # JWT auth auth: mode: "jwt" + # Which traffic needs a token. Default: [transcoded, endpoints]; the + # forward-auth endpoint is never behind it. `auth.authz` takes a `scope` too + # (default: [transcoded]). + # scope: { traffic: [transcoded, grpc] } jwt: jwks_uri: "https://idp.example.com/.well-known/jwks.json" # OR a static key: public_key_pem_file: "/etc/proxy/idp-ed25519.pub.pem" @@ -340,6 +358,64 @@ there is no boundary burst on top of this lag. See the `shield:` block under [Configuration](#configuration) for the full schema. +## Guards and scopes + +A guard is a check that may turn a request away: maintenance mode, the +concurrency limit, the rate limits, JWT auth, ext_authz and the auth decider. +Each one covers the traffic its `scope` names, and nothing else: + +| Traffic | What it is | +|---------|------------| +| `transcoded` | REST calls the proxy transcodes to gRPC | +| `endpoints` | the proxy's own endpoints: health, metrics, OpenAPI, OIDC, extra routes, forward-auth | +| `grpc` | native gRPC and gRPC-Web calls passed through to the upstream | +| `fallback` | requests no route answers, handed to `ProxyService::with_fallback` | +| `all` | every class above | + +```yaml +scope: + traffic: [transcoded, grpc] + paths: ["/v1/orders/**", "/acme.v1.Orders/*"] # optional, globs + methods: ["POST"] # optional +``` + +`paths` and `methods` narrow the guard within its traffic; `*` stays within a +path segment and `**` spans segments. A native gRPC call's path is +`/./`. A scope whose `traffic` is empty, a relative +or invalid glob, or an invalid method stops the proxy at startup. + +| Guard | Configured by | Default traffic | +|-------|---------------|-----------------| +| maintenance | `maintenance.scope` | `transcoded`, `endpoints` | +| concurrency limit | `concurrency.scope` | `transcoded`, `endpoints`, `grpc` | +| rate limits | `shield.scope` | `transcoded`, `endpoints` | +| JWT | `auth.scope` | `transcoded`, `endpoints` | +| ext_authz | `auth.authz.scope` | `transcoded` | +| auth decider | `ProxyServer::with_auth_decider_scope` | `transcoded` | + +The forward-auth endpoint answers for JWT, ext_authz and the decider, so it is +never behind them, whatever their scope. A request runs the guards of its own +class only, in this order: maintenance, concurrency, rate limits keyed before +auth, JWT, rate limits keyed by verified claims, ext_authz, the decider. A +guard outside the request's class costs it nothing. + +**Rejections in the request's protocol.** A REST client gets the +`google.rpc.Status` JSON body of [Error responses](#error-responses) with the +mapped HTTP status and an empty `details`: + +```json +{ "error": "RESOURCE_EXHAUSTED", "code": 8, "message": "rate limit exceeded", "details": [] } +``` + +A gRPC or gRPC-Web client gets a trailers-only response with the same code +(`UNAUTHENTICATED`, `PERMISSION_DENIED`, `RESOURCE_EXHAUSTED`, `UNAVAILABLE`) +and message, and the guard's headers (`Retry-After`, `RateLimit-*`, +`WWW-Authenticate`, `Location`) as metadata. A gRPC-Web rejection keeps the +proxy's CORS policy. An ext_authz denial carries the Check's own status code; +a decider's `Deny` maps its HTTP status back to a code by the `google.rpc.Code` +table, and its `Redirect` reaches a gRPC client as `UNAUTHENTICATED` with +`Location` in the metadata, since a gRPC client cannot follow it. + ## Error responses A failed gRPC call becomes a JSON body with the status of the gRPC → HTTP @@ -714,9 +790,11 @@ A plain TCP connection passes `stream.connect_info()` directly 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. +`ProxyService::with_fallback(my_axum_app)`. By default no guard covers those +requests or native gRPC ones, and they reach your service untouched; name +`grpc` or `fallback` in a guard's scope to put them behind it (see +[Guards and scopes](#guards-and-scopes)). CORS and tracing stay the fallback's +own. gRPC-Web requests pass through unchanged as well, so the upstream answers them in that protocol: wrap your services in tonic-web's layer @@ -793,20 +871,21 @@ ProxyServer::from_config(config) The hooks are: -- **`with_auth_decider`** — an in-process forward-auth / PDP decision, run inline - on every proxied request and exposed at `/verify` (path configurable via - `with_verify_path`). -- **`with_token_verifier`** — replaces the built-in JWT signature check +- **`with_auth_decider`**: an in-process forward-auth / PDP decision, run inline + on transcoded requests (other traffic with `with_auth_decider_scope`, see + [Guards and scopes](#guards-and-scopes)) and exposed at `/verify` (path + configurable via `with_verify_path`). +- **`with_token_verifier`**: replaces the built-in JWT signature check (see [JWT verification](#jwt-verification)) while keeping the route policies, the roles claim, and the claim→header forwarding. -- **`with_oidc_backend`** — backs the stateless OIDC surface (discovery, JWKS, +- **`with_oidc_backend`**: backs the stateless OIDC surface (discovery, JWKS, userinfo) with your key/client metadata; supersedes the config-driven static discovery. -- **`with_extra_routes`** — registers extra stateless routes through a +- **`with_extra_routes`**: registers extra stateless routes through a framework-agnostic adapter (request parts in, response parts out). -- **`with_error_details`** — chooses which transcoded routes return the +- **`with_error_details`**: chooses which transcoded routes return the upstream's `google.rpc.Status` details (see [Error responses](#error-responses)). -- **`with_denied_response_headers`** — keeps upstream response metadata keys +- **`with_denied_response_headers`**: keeps upstream response metadata keys off the HTTP responses (see [Upstream controls](#upstream-controls)). ## JWT verification @@ -975,6 +1054,8 @@ Client (HTTP/JSON) │ ├─────────────────┤ │ │ │ Maintenance │ │ 503 gate (exempt paths) │ ├─────────────────┤ │ +│ │ Concurrency │ │ in-flight limit (503) +│ ├─────────────────┤ │ │ │ Shield │ │ rate limiting (429) │ ├─────────────────┤ │ │ │ Auth (JWT) │ │ validate + policies (401/403) diff --git a/src/auth/authz.rs b/src/auth/authz.rs index d2e9fed..aed18ce 100644 --- a/src/auth/authz.rs +++ b/src/auth/authz.rs @@ -14,7 +14,6 @@ use axum::http::header::{HeaderName, HeaderValue, HOST}; use axum::http::{HeaderMap, StatusCode, Uri}; use axum::middleware::Next; use axum::response::{IntoResponse, Response}; -use axum::Json; use tonic::transport::Channel; use envoy_types::pb::envoy::config::core::v3::HeaderValueOption; @@ -25,8 +24,8 @@ use envoy_types::pb::envoy::service::auth::v3::{ AttributeContext, CheckRequest, CheckResponse, DeniedHttpResponse, }; -use super::forbidden; use crate::config::AuthzConfig; +use crate::guard::{http_to_grpc_code, mark_rejection, reject}; /// A configured ext_authz client. pub struct Authz { @@ -82,7 +81,10 @@ pub async fn middleware( } Err(status) => { tracing::warn!(error = %status, "authz check failed; failing closed"); - service_unavailable("authorization service unavailable") + reject( + tonic::Code::Unavailable, + "authorization service unavailable", + ) } } } @@ -148,13 +150,30 @@ fn evaluate(resp: CheckResponse) -> Decision { } else { let response = match resp.http_response { Some(EnvoyHttpResponse::DeniedResponse(denied)) - | Some(EnvoyHttpResponse::ErrorResponse(denied)) => denied_to_response(denied), - _ => forbidden("forbidden by authorization policy"), + | Some(EnvoyHttpResponse::ErrorResponse(denied)) => { + let response = denied_to_response(denied); + // A gRPC caller gets the Check's own status, the decision in + // gRPC terms; a Check that sets no code maps the HTTP one. + let code = match resp.status.as_ref().map(|s| s.code) { + Some(code) if code != 0 => tonic::Code::from_i32(code), + _ => http_to_grpc_code(response.status()), + }; + let message = resp + .status + .map(|s| s.message) + .filter(|m| !m.is_empty()) + .unwrap_or_else(|| DENIED.to_string()); + mark_rejection(response, code, message) + } + _ => reject(tonic::Code::PermissionDenied, DENIED), }; Decision::Deny(response) } } +/// The message of a denial that carries none of its own. +const DENIED: &str = "forbidden by authorization policy"; + /// Append authz-supplied headers, preserving multiple values for the same name /// (e.g. several `Set-Cookie`) instead of overwriting all but the last. fn apply_headers(dst: &mut HeaderMap, headers: Vec<(HeaderName, HeaderValue)>) { @@ -187,14 +206,6 @@ fn denied_to_response(denied: DeniedHttpResponse) -> Response { (status, headers, denied.body).into_response() } -fn service_unavailable(message: &str) -> Response { - ( - StatusCode::SERVICE_UNAVAILABLE, - Json(serde_json::json!({ "error": "UNAVAILABLE", "message": message })), - ) - .into_response() -} - #[cfg(test)] mod tests { use super::*; @@ -332,6 +343,7 @@ mod tests { endpoint: "http://127.0.0.1:1".into(), timeout_ms: 100, failure_mode_allow: false, + scope: None, }) .unwrap() .unwrap(); @@ -360,6 +372,7 @@ mod tests { endpoint: "http://127.0.0.1:1".into(), timeout_ms: 100, failure_mode_allow: true, + scope: None, }) .unwrap() .unwrap(); diff --git a/src/auth/forward/tests.rs b/src/auth/forward/tests.rs index 006a7aa..557c3dc 100644 --- a/src/auth/forward/tests.rs +++ b/src/auth/forward/tests.rs @@ -70,6 +70,7 @@ fn forward_auth(pem_path: std::path::PathBuf, login_url: Option) -> Arc< applications_path: None, }), authz: None, + scope: None, }; let auth = Auth::build(&config, None).unwrap().unwrap(); ForwardAuth::build(&config, auth).unwrap() diff --git a/src/auth/mod.rs b/src/auth/mod.rs index 02e906d..b2725ee 100644 --- a/src/auth/mod.rs +++ b/src/auth/mod.rs @@ -32,10 +32,9 @@ use std::sync::Arc; use axum::extract::State; use axum::http::header::{HeaderName, HeaderValue}; -use axum::http::{HeaderMap, StatusCode}; +use axum::http::HeaderMap; use axum::middleware::Next; -use axum::response::{IntoResponse, Response}; -use axum::Json; +use axum::response::Response; use serde_json::Value; use crate::config::{default_roles_claim, AuthConfig}; @@ -356,18 +355,10 @@ fn inject_claim_headers( } } -fn unauthorized(message: &str) -> Response { - ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({ "error": "UNAUTHENTICATED", "message": message })), - ) - .into_response() +fn unauthorized(message: &'static str) -> Response { + crate::guard::reject(tonic::Code::Unauthenticated, message) } -fn forbidden(message: &str) -> Response { - ( - StatusCode::FORBIDDEN, - Json(serde_json::json!({ "error": "PERMISSION_DENIED", "message": message })), - ) - .into_response() +fn forbidden(message: &'static str) -> Response { + crate::guard::reject(tonic::Code::PermissionDenied, message) } diff --git a/src/auth/tests.rs b/src/auth/tests.rs index e5fb9b6..2c86e7f 100644 --- a/src/auth/tests.rs +++ b/src/auth/tests.rs @@ -5,6 +5,7 @@ use super::*; use crate::config::{AuthConfig, ForwardAuthConfig, JwtConfig, RoutePolicyConfig}; use axum::http::Request as HttpRequest; +use axum::http::StatusCode; use tower::ServiceExt; #[test] @@ -122,6 +123,7 @@ fn auth_with_stub(roles: &[&str], jwt: Option) -> Arc { jwt, forward_auth: Some(secure_policy(roles)), authz: None, + scope: None, }; let stub = StubVerifier { accepts: "good-token", @@ -230,6 +232,7 @@ async fn injected_verifier_is_called_on_every_request() { jwt: Some(jwt_claims_only()), forward_auth: Some(secure_policy(&[])), authz: None, + scope: None, }; let auth = Auth::build(&cfg, Some(verifier.clone())).unwrap().unwrap(); for _ in 0..3 { @@ -259,6 +262,7 @@ fn injected_verifier_supersedes_a_configured_key_source() { }), forward_auth: None, authz: None, + scope: None, }; let stub = StubVerifier { accepts: "good-token", @@ -274,6 +278,7 @@ fn no_auth_when_mode_is_not_jwt() { jwt: None, forward_auth: None, authz: None, + scope: None, }; assert!(Auth::build(&cfg, None).unwrap().is_none()); } @@ -288,6 +293,7 @@ fn jwt_mode_without_a_verifier_is_rejected() { jwt: Some(jwt_claims_only()), forward_auth: None, authz: None, + scope: None, }; let Err(err) = Auth::build(&cfg, None) else { panic!("a jwt config with no verifier must not build"); @@ -349,6 +355,7 @@ mod builtin { }), forward_auth: Some(secure_policy(roles)), authz: None, + scope: None, }; Auth::build(&cfg, None).unwrap().unwrap() } @@ -397,6 +404,7 @@ mod builtin { ..secure_policy(&[]) }), authz: None, + scope: None, }; let auth = Auth::build(&cfg, None).unwrap().unwrap(); let resp = app(auth) @@ -553,6 +561,7 @@ mod builtin { }), forward_auth: Some(secure_policy(&[])), authz: None, + scope: None, }; let auth = Auth::build(&cfg, None).unwrap().unwrap(); let app = app(auth.clone()); @@ -614,6 +623,7 @@ mod builtin { }), forward_auth: Some(secure_policy(&[])), authz: None, + scope: None, }; let auth = Auth::build(&cfg, None).unwrap().unwrap(); let app = app(auth.clone()); @@ -643,6 +653,7 @@ mod builtin { }), forward_auth: None, authz: None, + scope: None, }; let Err(err) = Auth::build(&jwks(59), None) else { panic!("a JWKS age below 60 s must be rejected"); @@ -662,6 +673,7 @@ mod builtin { }), forward_auth: None, authz: None, + scope: None, }; let Err(err) = Auth::build(&cfg, None) else { panic!("a zero-sized cache must be rejected"); @@ -678,6 +690,7 @@ mod builtin { jwt: Some(jwt_claims_only()), forward_auth: None, authz: None, + scope: None, }; let Err(err) = Auth::build(&cfg, None) else { panic!("a built-in verifier with no key source must not build"); diff --git a/src/config.rs b/src/config.rs index c033528..98c59a7 100644 --- a/src/config.rs +++ b/src/config.rs @@ -85,6 +85,10 @@ pub struct ProxyConfig { /// Server-streaming response behavior. #[serde(default)] pub streaming: StreamingConfig, + + /// Request concurrency limit. + #[serde(default)] + pub concurrency: Option, } fn default_forwarded_headers() -> Vec { @@ -230,6 +234,7 @@ pub(crate) const KNOWN_TOP_LEVEL_KEYS: &[&str] = &[ "error_details", "response_headers", "runtime", + "concurrency", ]; /// Every `streaming:` key: the [`StreamingConfig`] fields plus the ones @@ -456,6 +461,89 @@ pub struct AuthConfig { /// AuthZ integration (optional gRPC call). #[serde(default)] pub authz: Option, + + /// The traffic JWT authentication covers. Default: `transcoded` and + /// `endpoints`; the forward-auth endpoint is never behind it. + #[serde(default)] + pub scope: Option, +} + +/// The traffic a guard covers, and optionally which paths and methods of it. +/// +/// ```yaml +/// scope: +/// traffic: [transcoded, grpc] # or [all] +/// paths: ["/v1/**", "/acme.v1.Orders/*"] +/// methods: ["POST"] +/// ``` +/// +/// Omitted keys keep the guard's own default traffic and cover every path and +/// method of it. +#[derive(Debug, Clone, Default, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ScopeConfig { + /// The classes of traffic covered; `all` names every class. + #[serde(default)] + pub traffic: Option>, + /// Path globs (`*` within a segment, `**` across segments) the guard is + /// narrowed to; empty covers every path. A native gRPC call's path is + /// `/./`. + #[serde(default)] + pub paths: Vec, + /// Methods the guard is narrowed to; empty covers every method. + #[serde(default)] + pub methods: Vec, +} + +impl ScopeConfig { + /// A scope covering every path and method of `traffic`. + /// + /// # Examples + /// + /// ``` + /// use structured_proxy::config::{ScopeConfig, Traffic}; + /// + /// let scope = ScopeConfig::traffic([Traffic::Transcoded, Traffic::Grpc]); + /// assert!(scope.paths.is_empty()); + /// ``` + pub fn traffic(traffic: impl IntoIterator) -> Self { + Self { + traffic: Some(traffic.into_iter().collect()), + paths: Vec::new(), + methods: Vec::new(), + } + } +} + +/// A class of traffic the proxy tells apart before any guard runs. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum Traffic { + /// REST calls the proxy transcodes to gRPC. + Transcoded, + /// The proxy's own endpoints: health, metrics, OpenAPI, OIDC, extra + /// routes and forward-auth. + Endpoints, + /// Native gRPC and gRPC-Web calls passed through to the upstream. + Grpc, + /// Requests no route answers, handed to the fallback. + Fallback, + /// Every class above. + All, +} + +/// Request concurrency limit: requests past `max_in_flight` are answered +/// `UNAVAILABLE` (503) at once instead of queueing. +#[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ConcurrencyConfig { + /// Most requests in flight at once, counted until each response body + /// ends, so a stream holds its slot for its whole life. At least 1. + pub max_in_flight: usize, + /// The traffic the limit covers, one budget shared by all of it. + /// Default: `transcoded`, `endpoints` and `grpc`. + #[serde(default)] + pub scope: Option, } fn default_auth_mode() -> String { @@ -616,6 +704,9 @@ pub struct AuthzConfig { /// request through instead of denying. Defaults to false (fail closed). #[serde(default)] pub failure_mode_allow: bool, + /// The traffic the check covers. Default: `transcoded`. + #[serde(default)] + pub scope: Option, } fn default_authz_timeout_ms() -> u64 { @@ -667,6 +758,9 @@ pub struct ShieldConfig { /// trust forwarding headers; set this behind a load balancer. #[serde(default)] pub trusted_proxies: Vec, + /// The traffic the rules apply to. Default: `transcoded` and `endpoints`. + #[serde(default)] + pub scope: Option, } /// A named limit tier: a sustained rate plus an instantaneous burst capacity. @@ -976,6 +1070,10 @@ pub struct MaintenanceConfig { pub exempt_paths: Vec, #[serde(default = "default_maintenance_message")] pub message: String, + /// The traffic maintenance mode turns away. Default: `transcoded` and + /// `endpoints`. + #[serde(default)] + pub scope: Option, } fn default_exempt_paths() -> Vec { @@ -997,6 +1095,7 @@ impl Default for MaintenanceConfig { enabled: false, exempt_paths: default_exempt_paths(), message: default_maintenance_message(), + scope: None, } } } diff --git a/src/config/tests.rs b/src/config/tests.rs index bfa7bf5..b3c2201 100644 --- a/src/config/tests.rs +++ b/src/config/tests.rs @@ -498,6 +498,7 @@ fn known_top_level_keys_cover_every_proxy_config_field() { metrics_classes: _, forwarded_headers: _, streaming: _, + concurrency: _, } = config; for key in [ "upstream", @@ -520,10 +521,11 @@ fn known_top_level_keys_cover_every_proxy_config_field() { "error_details", "response_headers", "runtime", + "concurrency", ] { assert!(KNOWN_TOP_LEVEL_KEYS.contains(&key), "{key}"); } - assert_eq!(KNOWN_TOP_LEVEL_KEYS.len(), 20); + assert_eq!(KNOWN_TOP_LEVEL_KEYS.len(), 21); } #[test] diff --git a/src/embed.rs b/src/embed.rs index f52ebc6..074fdf2 100644 --- a/src/embed.rs +++ b/src/embed.rs @@ -19,6 +19,7 @@ use axum::response::{IntoResponse, Response}; use axum::routing::{get, on, MethodFilter, MethodRouter}; use axum::{Json, Router}; +use crate::guard::{http_to_grpc_code, mark_rejection}; use crate::hooks::{AuthDecider, Decision, ExtraRoute, OidcBackend, RequestParts, RouteRequest}; /// Cap on the body an extra-route handler will buffer (16 MiB). Extra routes are @@ -69,9 +70,21 @@ pub(crate) async fn auth_decider_gate( strip_then_insert(dst, &inject_headers); next.run(request).await } - Decision::Deny { status, body } => deny_response(status, body), - // Inline (browser-facing) path: drive a real redirect. - Decision::Redirect { location } => redirect_response(StatusCode::FOUND, &location), + // The decider owns the HTTP answer; a gRPC caller gets the code its + // status maps to. + Decision::Deny { status, body } => mark_rejection( + deny_response(status, body), + http_to_grpc_code(status), + "denied by the auth decider", + ), + // Inline (browser-facing) path: drive a real redirect. A gRPC client + // cannot follow one; it is told to authenticate, `Location` in its + // metadata. + Decision::Redirect { location } => mark_rejection( + redirect_response(StatusCode::FOUND, &location), + tonic::Code::Unauthenticated, + "authentication required", + ), } } diff --git a/src/guard.rs b/src/guard.rs new file mode 100644 index 0000000..0b5c988 --- /dev/null +++ b/src/guard.rs @@ -0,0 +1,475 @@ +//! Guards: the checks that may turn a request away before it reaches a route, +//! the upstream or the fallback (maintenance, concurrency, rate limits, JWT, +//! ext_authz, the auth decider). Each covers the traffic its scope names, and +//! each rejects through [`reject`] or [`mark_rejection`], so a rejection +//! reaches a REST client as its HTTP answer and a gRPC client as a gRPC +//! status ([`GrpcRejections`]). + +mod concurrency; +mod grpc; + +use alloc::borrow::Cow; +use alloc::string::String; +use alloc::sync::Arc; +use alloc::vec::Vec; +use core::convert::Infallible; +use core::task::{Context, Poll}; + +use axum::extract::{Request, State}; +use axum::middleware::{from_fn_with_state, Next}; +use axum::response::{IntoResponse, Response}; +use axum::{Json, Router}; +use futures::future::Either; +// no-std: a segment-wise glob matcher over `&str` (globset needs std). +use globset::{GlobBuilder, GlobSet, GlobSetBuilder}; +use http::{Method, StatusCode}; +use tower::util::BoxCloneSyncService; +use tower::{Layer, Service}; + +use crate::config::{ScopeConfig, Traffic}; +use crate::hooks::AuthDecider; +use crate::transcode::error::{grpc_to_http_status, guard_error_body}; + +pub(crate) use concurrency::Concurrency; +pub(crate) use grpc::GrpcRejections; + +/// What a guard's rejection means in gRPC terms, attached to its response so a +/// gRPC request is answered with a status instead of the HTTP body. +#[derive(Clone, Debug)] +pub(crate) struct Rejection { + pub(crate) code: tonic::Code, + pub(crate) message: Cow<'static, str>, +} + +/// A guard's rejection: the `google.rpc.Status` JSON body the transcoder's +/// own errors use, with the HTTP status `code` maps to, carrying the +/// [`Rejection`] for a gRPC request. +pub(crate) fn reject(code: tonic::Code, message: impl Into>) -> Response { + let message = message.into(); + let body = guard_error_body(code, &message); + mark_rejection( + (grpc_to_http_status(code), Json(body)).into_response(), + code, + message, + ) +} + +/// Mark a guard's own HTTP answer (a decider's body, an ext_authz denial, a +/// redirect) with the gRPC status a gRPC request gets instead. +pub(crate) fn mark_rejection( + mut response: Response, + code: tonic::Code, + message: impl Into>, +) -> Response { + response.extensions_mut().insert(Rejection { + code, + message: message.into(), + }); + response +} + +/// The gRPC code of an HTTP status, by the HTTP mapping of +/// `google/rpc/code.proto`: the inverse of the transcoder's own mapping, so a +/// code survives the round trip. +pub(crate) fn http_to_grpc_code(status: StatusCode) -> tonic::Code { + match status.as_u16() { + 400 => tonic::Code::InvalidArgument, + 401 => tonic::Code::Unauthenticated, + 403 => tonic::Code::PermissionDenied, + 404 => tonic::Code::NotFound, + 409 => tonic::Code::Aborted, + 429 => tonic::Code::ResourceExhausted, + 499 => tonic::Code::Cancelled, + 501 => tonic::Code::Unimplemented, + 503 => tonic::Code::Unavailable, + 504 => tonic::Code::DeadlineExceeded, + 500..=599 => tonic::Code::Internal, + _ => tonic::Code::Unknown, + } +} + +/// A class of traffic a guard stack is built for. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum Class { + Transcoded, + Endpoints, + /// The forward-auth endpoint: an endpoint that answers the JWT and decider + /// gates, so it is never behind them. + Verify, + Grpc, + Fallback, +} + +const TRANSCODED: u8 = 1; +const ENDPOINTS: u8 = 1 << 1; +const GRPC: u8 = 1 << 2; +const FALLBACK: u8 = 1 << 3; + +impl Class { + fn bit(self) -> u8 { + match self { + Self::Transcoded => TRANSCODED, + Self::Endpoints | Self::Verify => ENDPOINTS, + Self::Grpc => GRPC, + Self::Fallback => FALLBACK, + } + } + + /// Whether the authenticating guards (JWT, ext_authz, the decider) may + /// run on this class. + fn authenticates(self) -> bool { + self != Self::Verify + } +} + +fn traffic_bits(traffic: Traffic) -> u8 { + match traffic { + Traffic::Transcoded => TRANSCODED, + Traffic::Endpoints => ENDPOINTS, + Traffic::Grpc => GRPC, + Traffic::Fallback => FALLBACK, + Traffic::All => TRANSCODED | ENDPOINTS | GRPC | FALLBACK, + } +} + +/// A compiled [`ScopeConfig`]: the classes a guard is mounted on, and the +/// paths and methods it is narrowed to within them. +#[derive(Debug)] +pub(crate) struct Scope { + classes: u8, + paths: Option, + methods: Option>, +} + +impl Scope { + /// Compile `config` for the guard named `what`, with `default` traffic + /// when the config names none. + /// + /// # Errors + /// A scope that covers no traffic, a path glob that is relative or does + /// not compile, or a method that is not an HTTP method token. + pub(crate) fn compile( + config: Option<&ScopeConfig>, + default: &[Traffic], + what: &str, + ) -> Result, String> { + let traffic = config.and_then(|c| c.traffic.as_deref()).unwrap_or(default); + let classes = traffic.iter().fold(0, |acc, t| acc | traffic_bits(*t)); + if classes == 0 { + return Err(format!("{what}.scope.traffic covers no traffic")); + } + let patterns = config.map(|c| c.paths.as_slice()).unwrap_or_default(); + let paths = if patterns.is_empty() { + None + } else { + let mut set = GlobSetBuilder::new(); + for pattern in patterns { + // Paths always start with `/`; a relative pattern is a typo + // that would silently never match. + if !pattern.starts_with('/') { + return Err(format!( + "{what}.scope.paths entry {pattern:?} must start with '/'" + )); + } + let glob = GlobBuilder::new(pattern) + .literal_separator(true) + .build() + .map_err(|e| format!("{what}.scope.paths entry {pattern:?} is invalid: {e}"))?; + set.add(glob); + } + Some( + set.build() + .map_err(|e| format!("{what}.scope.paths: {e}"))?, + ) + }; + let names = config.map(|c| c.methods.as_slice()).unwrap_or_default(); + let methods = if names.is_empty() { + None + } else { + Some( + names + .iter() + .map(|name| { + Method::from_bytes(name.to_ascii_uppercase().as_bytes()).map_err(|_| { + format!("{what}.scope.methods entry {name:?} is not a method") + }) + }) + .collect::, _>>()?, + ) + }; + Ok(Arc::new(Self { + classes, + paths, + methods, + })) + } + + fn covers(&self, class: Class) -> bool { + self.classes & class.bit() != 0 + } + + fn narrows(&self) -> bool { + self.paths.is_some() || self.methods.is_some() + } + + /// Whether the guard applies to `request`, within a class it covers. + fn matches(&self, request: &http::Request) -> bool { + self.methods + .as_ref() + .is_none_or(|methods| methods.contains(request.method())) + && self + .paths + .as_ref() + .is_none_or(|paths| paths.is_match(request.uri().path())) + } +} + +/// A guard layer `L` applied only to the requests its scope's paths and +/// methods select; the others go straight to the inner service. +#[derive(Clone)] +struct Scoped { + layer: L, + scope: Arc, +} + +impl, S: Clone> Layer for Scoped { + type Service = ScopedService; + + fn layer(&self, inner: S) -> Self::Service { + ScopedService { + guarded: self.layer.layer(inner.clone()), + plain: inner, + scope: self.scope.clone(), + } + } +} + +#[derive(Clone)] +struct ScopedService { + guarded: G, + plain: S, + scope: Arc, +} + +impl Service for ScopedService +where + G: Service, + S: Service, +{ + type Response = Response; + type Error = Infallible; + type Future = Either; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + // Guard middleware and routes are always ready; readiness further in + // is waited for per request. + Poll::Ready(Ok(())) + } + + fn call(&mut self, request: Request) -> Self::Future { + if self.scope.matches(&request) { + Either::Left(self.guarded.call(request)) + } else { + Either::Right(self.plain.call(request)) + } + } +} + +/// A gRPC-path or fallback service with its guards around it. +pub(crate) type BoxedService = BoxCloneSyncService; + +/// The guards the proxy runs, each with its scope. +#[derive(Default)] +pub(crate) struct Guards { + pub(crate) maintenance: Option<(Arc, Arc)>, + pub(crate) concurrency: Option<(Arc, Arc)>, + pub(crate) shield: Option<(Arc, Arc)>, + pub(crate) auth: Option<(Arc, Arc)>, + pub(crate) authz: Option<(Arc, Arc)>, + pub(crate) decider: Option<(Arc, Arc)>, +} + +/// Put the guards that cover `$class` around `$target`, in pipeline order +/// (outermost first): maintenance, concurrency, rate limits before auth, JWT, +/// rate limits after auth (keyed by verified claims), ext_authz, the decider. +/// `$wrap` applies one layer to the target type at hand; a macro, since each +/// guard's layer is its own type. +macro_rules! guard_stack { + ($guards:expr, $target:expr, $class:expr, $wrap:ident) => {{ + let guards: &Guards = $guards; + let class: Class = $class; + let mut target = $target; + // Built inside out: the first layer added runs last. + if class.authenticates() { + if let Some((decider, scope)) = &guards.decider { + target = $wrap!( + target, + class, + scope, + from_fn_with_state(decider.clone(), crate::embed::auth_decider_gate) + ); + } + if let Some((authz, scope)) = &guards.authz { + target = $wrap!( + target, + class, + scope, + from_fn_with_state(authz.clone(), crate::auth::authz::middleware) + ); + } + } + if class.authenticates() { + // Keyed by claims only the JWT gate verifies. + if let Some((shield, scope)) = &guards.shield { + target = $wrap!( + target, + class, + scope, + from_fn_with_state(shield.clone(), crate::shield::post_auth_middleware) + ); + } + if let Some((auth, scope)) = &guards.auth { + target = $wrap!( + target, + class, + scope, + from_fn_with_state(auth.clone(), crate::auth::middleware) + ); + } + } + if let Some((shield, scope)) = &guards.shield { + target = $wrap!( + target, + class, + scope, + from_fn_with_state(shield.clone(), crate::shield::pre_auth_middleware) + ); + } + if let Some((concurrency, scope)) = &guards.concurrency { + target = $wrap!( + target, + class, + scope, + from_fn_with_state(concurrency.clone(), concurrency::middleware) + ); + } + if let Some((maintenance, scope)) = &guards.maintenance { + target = $wrap!( + target, + class, + scope, + from_fn_with_state(maintenance.clone(), maintenance_middleware) + ); + } + target + }}; +} + +/// One guard layer on a router, when its scope covers the class. +macro_rules! wrap_router { + ($target:expr, $class:expr, $scope:expr, $layer:expr) => {{ + let scope: &Arc = $scope; + if !scope.covers($class) { + $target + } else if scope.narrows() { + $target.layer(Scoped { + layer: $layer, + scope: scope.clone(), + }) + } else { + $target.layer($layer) + } + }}; +} + +/// One guard layer on a boxed service, when its scope covers the class. +macro_rules! wrap_service { + ($target:expr, $class:expr, $scope:expr, $layer:expr) => {{ + let scope: &Arc = $scope; + let target: BoxedService = $target; + if !scope.covers($class) { + target + } else if scope.narrows() { + BoxedService::new( + Scoped { + layer: $layer, + scope: scope.clone(), + } + .layer(target), + ) + } else { + BoxedService::new($layer.layer(target)) + } + }}; +} + +impl Guards { + /// Whether any guard covers `class`. + pub(crate) fn cover(&self, class: Class) -> bool { + fn on(guard: &Option<(T, Arc)>, class: Class) -> bool { + guard.as_ref().is_some_and(|(_, scope)| scope.covers(class)) + } + on(&self.maintenance, class) + || on(&self.concurrency, class) + || on(&self.shield, class) + || (class.authenticates() + && (on(&self.auth, class) || on(&self.authz, class) || on(&self.decider, class))) + } + + /// `router`, all of class `class`, behind the guards that cover it. + pub(crate) fn router(&self, router: Router, class: Class) -> Router + where + S: Clone + Send + Sync + 'static, + { + guard_stack!(self, router, class, wrap_router) + } + + /// `service`, of class `class`, behind the guards that cover it. + pub(crate) fn service(&self, service: BoxedService, class: Class) -> BoxedService { + guard_stack!(self, service, class, wrap_service) + } +} + +/// Maintenance mode: every request outside the exempt paths gets a `503`. +#[derive(Debug)] +pub(crate) struct Maintenance { + pub(crate) exempt: Vec, + pub(crate) message: String, +} + +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). + pub(crate) 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: `UNAVAILABLE` with a five-minute +/// `Retry-After` for every request outside the exempt paths. +async fn maintenance_middleware( + State(maintenance): State>, + request: Request, + next: Next, +) -> Response { + if maintenance.exempts(request.uri().path()) { + return next.run(request).await; + } + let mut response = reject(tonic::Code::Unavailable, maintenance.message.clone()); + response.headers_mut().insert( + http::header::RETRY_AFTER, + http::HeaderValue::from_static("300"), + ); + response +} + +#[cfg(test)] +mod tests; diff --git a/src/guard/concurrency.rs b/src/guard/concurrency.rs new file mode 100644 index 0000000..d4a5bbc --- /dev/null +++ b/src/guard/concurrency.rs @@ -0,0 +1,107 @@ +//! The concurrency limit: at most `max_in_flight` requests at once, the rest +//! turned away at once rather than queued, since a queue behind a saturated +//! upstream only adds latency to what will time out anyway. + +use alloc::format; +use alloc::string::String; +use alloc::sync::Arc; +use core::pin::Pin; +use core::task::{Context, Poll}; + +use axum::body::Body; +use axum::extract::{Request, State}; +use axum::middleware::Next; +use axum::response::Response; +use bytes::Bytes; +use http_body::{Frame, SizeHint}; +use pin_project_lite::pin_project; +// no-std: an `AtomicUsize` slot counter with a drop guard (only try-acquire is used). +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; + +use super::reject; +use crate::config::ConcurrencyConfig; + +/// The in-flight slots. +#[derive(Debug)] +pub(crate) struct Concurrency { + slots: Arc, +} + +impl Concurrency { + /// The limit `config` sets. + /// + /// # Errors + /// A `max_in_flight` of zero, which would turn every request away, or one + /// past what a semaphore can count. + pub(crate) fn build(config: &ConcurrencyConfig) -> Result, String> { + if config.max_in_flight == 0 || config.max_in_flight > Semaphore::MAX_PERMITS { + return Err(format!( + "concurrency.max_in_flight must be between 1 and {}", + Semaphore::MAX_PERMITS + )); + } + Ok(Arc::new(Self { + slots: Arc::new(Semaphore::new(config.max_in_flight)), + })) + } +} + +/// Take a slot for the request and hold it until its response body ends, so +/// a stream counts for its whole life; with none free, `UNAVAILABLE` and a +/// one-second `Retry-After`. +pub(super) async fn middleware( + State(concurrency): State>, + request: Request, + next: Next, +) -> Response { + let Ok(slot) = concurrency.slots.clone().try_acquire_owned() else { + let mut response = reject(tonic::Code::Unavailable, "too many requests in flight"); + response.headers_mut().insert( + http::header::RETRY_AFTER, + http::HeaderValue::from_static("1"), + ); + return response; + }; + next.run(request).await.map(|body| { + Body::new(Holding { + body, + slot: Some(slot), + }) + }) +} + +pin_project! { + /// A response body that frees its request's slot when it ends or is + /// dropped. + struct Holding { + #[pin] + body: Body, + slot: Option, + } +} + +impl http_body::Body for Holding { + type Data = Bytes; + type Error = axum::Error; + + fn poll_frame( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, axum::Error>>> { + let this = self.project(); + let frame = this.body.poll_frame(cx); + if let Poll::Ready(None | Some(Err(_))) = frame { + // Ended: the slot is free before the connection moves on. + this.slot.take(); + } + frame + } + + fn is_end_stream(&self) -> bool { + self.body.is_end_stream() + } + + fn size_hint(&self) -> SizeHint { + self.body.size_hint() + } +} diff --git a/src/guard/grpc.rs b/src/guard/grpc.rs new file mode 100644 index 0000000..2fa696b --- /dev/null +++ b/src/guard/grpc.rs @@ -0,0 +1,105 @@ +//! Guard rejections on the gRPC path, answered in gRPC. + +use alloc::string::ToString; +use core::convert::Infallible; +use core::future::Future; +use core::pin::Pin; +use core::task::{ready, Context, Poll}; + +use axum::extract::Request; +use axum::response::Response; +use pin_project_lite::pin_project; +use tower::Service; + +use super::{http_to_grpc_code, BoxedService, Rejection}; +use crate::upstream::{grpc_protocol, trailers_only, GrpcProtocol}; + +/// The guarded gRPC path of one protocol: a guard's HTTP rejection becomes a +/// trailers-only status in `protocol` (gRPC PROTOCOL-HTTP2, "Responses"), the +/// only answer a gRPC client reads. The upstream's own answers pass as they +/// are. +#[derive(Clone)] +pub(crate) struct GrpcRejections { + pub(crate) inner: BoxedService, + pub(crate) protocol: GrpcProtocol, +} + +impl Service for GrpcRejections { + type Response = Response; + type Error = Infallible; + type Future = RejectionFuture; + + #[inline] + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, request: Request) -> Self::Future { + RejectionFuture { + future: self.inner.call(request), + protocol: self.protocol, + } + } +} + +pin_project! { + /// The response future of [`GrpcRejections`]. + pub(crate) struct RejectionFuture { + #[pin] + future: >::Future, + protocol: GrpcProtocol, + } +} + +impl Future for RejectionFuture { + type Output = Result; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + let this = self.project(); + let response = ready!(this.future.poll(cx))?; + Poll::Ready(Ok(in_grpc(response, *this.protocol))) + } +} + +/// Headers of an HTTP rejection that describe the body, not the rejection, +/// and so do not become gRPC metadata. +const BODY_HEADERS: [http::HeaderName; 3] = [ + http::header::CONTENT_TYPE, + http::header::CONTENT_LENGTH, + http::header::TRANSFER_ENCODING, +]; + +/// `response` as a gRPC client reads it: a guard's rejection (marked, or any +/// answer without a gRPC content type) as a trailers-only status, carrying the +/// guard's headers (`Retry-After`, `RateLimit-*`, `WWW-Authenticate`, +/// `Location`) as metadata. +fn in_grpc(response: Response, protocol: GrpcProtocol) -> Response { + let (code, message) = match response.extensions().get::() { + Some(rejection) => (rejection.code, rejection.message.to_string()), + None if grpc_protocol(response.headers()).is_some() => return response, + None => ( + http_to_grpc_code(response.status()), + response + .status() + .canonical_reason() + .unwrap_or_default() + .to_string(), + ), + }; + let (parts, _) = response.into_parts(); + let mut answer = trailers_only(tonic::Status::new(code, message), protocol); + let headers = answer.headers_mut(); + let mut previous = None; + for (name, value) in parts.headers { + // `HeaderMap::into_iter` names a header once, before all its values. + if let Some(name) = name { + previous = Some(name); + } + let Some(name) = &previous else { continue }; + if BODY_HEADERS.contains(name) { + continue; + } + headers.append(name.clone(), value); + } + answer +} diff --git a/src/guard/tests.rs b/src/guard/tests.rs new file mode 100644 index 0000000..02d69e3 --- /dev/null +++ b/src/guard/tests.rs @@ -0,0 +1,320 @@ +use super::*; +use crate::upstream::GrpcProtocol; + +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")); +} + +fn request(method: &str, path: &str) -> http::Request<()> { + http::Request::builder() + .method(method) + .uri(path) + .body(()) + .unwrap() +} + +#[test] +fn a_scope_without_traffic_takes_the_default() { + let scope = Scope::compile(None, &[Traffic::Transcoded, Traffic::Grpc], "shield").unwrap(); + assert!(scope.covers(Class::Transcoded)); + assert!(scope.covers(Class::Grpc)); + assert!(!scope.covers(Class::Endpoints)); + assert!(!scope.covers(Class::Fallback)); + assert!(!scope.narrows()); +} + +#[test] +fn all_traffic_covers_every_class() { + let config = ScopeConfig::traffic([Traffic::All]); + let scope = Scope::compile(Some(&config), &[], "shield").unwrap(); + for class in [ + Class::Transcoded, + Class::Endpoints, + Class::Verify, + Class::Grpc, + Class::Fallback, + ] { + assert!(scope.covers(class), "{class:?}"); + } +} + +#[test] +fn a_scope_that_covers_no_traffic_is_an_error() { + // An empty list would mount the guard nowhere; a typo, not a choice. + let config = ScopeConfig::traffic([]); + let err = Scope::compile(Some(&config), &[Traffic::Transcoded], "auth").unwrap_err(); + assert!(err.contains("auth.scope.traffic"), "{err}"); +} + +#[test] +fn a_relative_or_invalid_path_glob_is_an_error() { + let relative = ScopeConfig { + paths: vec!["v1/**".into()], + ..ScopeConfig::default() + }; + let err = Scope::compile(Some(&relative), &[Traffic::Transcoded], "shield").unwrap_err(); + assert!(err.contains("must start with '/'"), "{err}"); + + let invalid = ScopeConfig { + paths: vec!["/v1/[".into()], + ..ScopeConfig::default() + }; + let err = Scope::compile(Some(&invalid), &[Traffic::Transcoded], "shield").unwrap_err(); + assert!(err.contains("is invalid"), "{err}"); +} + +#[test] +fn an_invalid_method_is_an_error() { + let config = ScopeConfig { + methods: vec!["GE T".into()], + ..ScopeConfig::default() + }; + let err = Scope::compile(Some(&config), &[Traffic::Transcoded], "shield").unwrap_err(); + assert!(err.contains("is not a method"), "{err}"); +} + +#[test] +fn paths_and_methods_narrow_the_requests_a_scope_matches() { + let config = ScopeConfig { + paths: vec!["/v1/orders/**".into()], + methods: vec!["post".into()], + ..ScopeConfig::default() + }; + let scope = Scope::compile(Some(&config), &[Traffic::Transcoded], "shield").unwrap(); + assert!(scope.narrows()); + assert!(scope.matches(&request("POST", "/v1/orders/42"))); + // Methods compare as tokens, case aside. + assert!(!scope.matches(&request("GET", "/v1/orders/42"))); + // `*` stays within a segment and `**` needs its own: `/v1/ordersx` is not + // below `/v1/orders`. + assert!(!scope.matches(&request("POST", "/v1/ordersx"))); + assert!(!scope.matches(&request("POST", "/v1/users/1"))); +} + +#[test] +fn http_statuses_map_back_to_their_grpc_codes() { + // The inverse of google/rpc/code.proto's HTTP mapping, so a code a guard + // answers with over REST is the one a gRPC client gets. + for (status, code) in [ + (400, tonic::Code::InvalidArgument), + (401, tonic::Code::Unauthenticated), + (403, tonic::Code::PermissionDenied), + (404, tonic::Code::NotFound), + (409, tonic::Code::Aborted), + (429, tonic::Code::ResourceExhausted), + (499, tonic::Code::Cancelled), + (500, tonic::Code::Internal), + (501, tonic::Code::Unimplemented), + (502, tonic::Code::Internal), + (503, tonic::Code::Unavailable), + (504, tonic::Code::DeadlineExceeded), + (302, tonic::Code::Unknown), + ] { + let status = StatusCode::from_u16(status).unwrap(); + assert_eq!(http_to_grpc_code(status), code, "{status}"); + } +} + +#[tokio::test] +async fn a_rejection_carries_the_status_json_body() { + let response = reject(tonic::Code::ResourceExhausted, "rate limit exceeded"); + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + let rejection = response.extensions().get::().unwrap(); + assert_eq!(rejection.code, tonic::Code::ResourceExhausted); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + let body: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!( + body, + serde_json::json!({ + "error": "RESOURCE_EXHAUSTED", + "code": 8, + "message": "rate limit exceeded", + "details": [], + }) + ); +} + +#[tokio::test] +async fn the_concurrency_limit_holds_a_slot_until_the_response_body_ends() { + use tower::ServiceExt; + let concurrency = Concurrency::build(&crate::config::ConcurrencyConfig { + max_in_flight: 1, + scope: None, + }) + .unwrap(); + // `/stream` answers with a body that ends when its sender is dropped. + let (sender, receiver) = tokio::sync::mpsc::channel::>(1); + let receiver = Arc::new(std::sync::Mutex::new(Some(receiver))); + let app: Router = Router::new() + .route( + "/stream", + axum::routing::get(move || { + let receiver = receiver.lock().unwrap().take().unwrap(); + async move { + axum::body::Body::from_stream(tokio_stream::wrappers::ReceiverStream::new( + receiver, + )) + } + }), + ) + .route("/x", axum::routing::get(|| async { "x" })) + .layer(from_fn_with_state(concurrency, concurrency::middleware)); + let get = |path: &str| { + http::Request::get(path) + .body(axum::body::Body::empty()) + .unwrap() + }; + + let streaming = app.clone().oneshot(get("/stream")).await.unwrap(); + assert_eq!(streaming.status(), StatusCode::OK); + + // The stream's headers are sent, its body is not: the slot is still taken. + let shed = app.clone().oneshot(get("/x")).await.unwrap(); + assert_eq!(shed.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(shed.headers()["retry-after"], "1"); + let rejection = shed.extensions().get::().unwrap(); + assert_eq!(rejection.code, tonic::Code::Unavailable); + + // Ending the body frees the slot. + sender + .send(Ok(bytes::Bytes::from_static(b"tail"))) + .await + .unwrap(); + drop(sender); + let body = axum::body::to_bytes(streaming.into_body(), usize::MAX) + .await + .unwrap(); + assert_eq!(&body[..], b"tail"); + let after = app.oneshot(get("/x")).await.unwrap(); + assert_eq!(after.status(), StatusCode::OK); +} + +#[test] +fn a_concurrency_limit_of_zero_is_an_error() { + let err = Concurrency::build(&crate::config::ConcurrencyConfig { + max_in_flight: 0, + scope: None, + }) + .unwrap_err(); + assert!(err.contains("max_in_flight"), "{err}"); +} + +/// The gRPC-path service answering every request with `response`. +fn answering(response: fn() -> Response, protocol: GrpcProtocol) -> GrpcRejections { + GrpcRejections { + inner: BoxedService::new(tower::service_fn(move |_: Request| async move { + Ok::<_, Infallible>(response()) + })), + protocol, + } +} + +async fn call(mut service: GrpcRejections) -> Response { + use tower::ServiceExt; + let request = Request::new(axum::body::Body::empty()); + service.ready().await.unwrap().call(request).await.unwrap() +} + +#[tokio::test] +async fn a_rejection_reaches_grpc_as_a_trailers_only_status_with_its_headers() { + let service = answering( + || { + let mut response = reject(tonic::Code::ResourceExhausted, "rate limit exceeded"); + response + .headers_mut() + .insert("retry-after", http::HeaderValue::from_static("7")); + response + }, + GrpcProtocol::Grpc, + ); + let response = call(service).await; + assert_eq!(response.status(), StatusCode::OK); + let headers = response.headers(); + assert_eq!(headers["content-type"], "application/grpc"); + assert_eq!(headers["grpc-status"], "8"); + assert_eq!(headers["grpc-message"], "rate%20limit%20exceeded"); + assert_eq!(headers["retry-after"], "7"); + // The JSON body's length does not describe the empty gRPC answer. + assert!(!headers.contains_key("content-length")); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + assert!(body.is_empty()); +} + +#[tokio::test] +async fn a_rejection_on_grpc_web_takes_the_grpc_web_content_type() { + let service = answering( + || reject(tonic::Code::Unauthenticated, "authentication required"), + GrpcProtocol::WebText, + ); + let response = call(service).await; + let headers = response.headers(); + assert_eq!(headers["content-type"], "application/grpc-web-text+proto"); + assert_eq!(headers["grpc-status"], "16"); +} + +#[tokio::test] +async fn an_unmarked_http_answer_maps_its_status() { + // A guard answer the proxy did not write (an embedder's decider body) + // still reaches a gRPC client as a status. + let service = answering( + || (StatusCode::FORBIDDEN, "nope").into_response(), + GrpcProtocol::Grpc, + ); + let response = call(service).await; + assert_eq!(response.headers()["grpc-status"], "7"); +} + +#[tokio::test] +async fn an_upstream_grpc_answer_passes_unchanged() { + let service = answering( + || { + let mut response = Response::new(axum::body::Body::from("frame")); + response.headers_mut().insert( + http::header::CONTENT_TYPE, + http::HeaderValue::from_static("application/grpc"), + ); + response + }, + GrpcProtocol::Grpc, + ); + let response = call(service).await; + assert!(!response.headers().contains_key("grpc-status")); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + assert_eq!(&body[..], b"frame"); +} diff --git a/src/lib.rs b/src/lib.rs index 7163e02..0a388cf 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -53,10 +53,13 @@ compile_error!( (or neither, and inject a verifier with `ProxyServer::with_token_verifier`)" ); +extern crate alloc; + pub mod auth; pub mod config; mod cors; mod embed; +mod guard; pub mod hooks; pub mod oidc; pub mod openapi; @@ -73,9 +76,8 @@ pub use auth::crypto::install_default_crypto_provider; pub use service::{serve, ConnectionInfo, ProxyService}; use axum::extract::State; -use axum::http::{Request, StatusCode}; -use axum::middleware::Next; -use axum::response::{IntoResponse, Response}; +use axum::http::StatusCode; +use axum::response::IntoResponse; use axum::routing::get; use axum::{Json, Router}; use prost_reflect::DescriptorPool; @@ -85,7 +87,7 @@ use tower_http::trace::TraceLayer; use std::sync::Arc; -use config::{DescriptorSource, ProxyConfig}; +use config::{DescriptorSource, ProxyConfig, ScopeConfig}; use hooks::{AuthDecider, ExtraRoute, OidcBackend, TokenVerifier}; use upstream::Upstream; @@ -101,13 +103,6 @@ pub(crate) struct ProxyState { 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. pub struct ProxyServer { config: ProxyConfig, @@ -115,6 +110,8 @@ pub struct ProxyServer { descriptor_pool: Option, /// Optional in-process forward-auth/PDP gate (embedded Tier-2 hook). auth_decider: Option>, + /// The traffic the injected decider gates; `transcoded` when unset. + auth_decider_scope: Option, /// Optional stateless OIDC surface backing (embedded Tier-2 hook). oidc_backend: Option>, /// Embedder-supplied extra stateless routes (embedded Tier-2 hook). @@ -141,6 +138,7 @@ impl ProxyServer { config, descriptor_pool: None, auth_decider: None, + auth_decider_scope: None, oidc_backend: None, extra_routes: Vec::new(), verify_path: None, @@ -203,14 +201,40 @@ impl ProxyServer { /// Inject an in-process forward-auth / PDP decision (embedded Tier-2 hook). /// - /// The decider gates every proxied request inline and also backs the - /// `/verify` forward-auth endpoint. Its signature is `axum`-free (see - /// [`hooks::AuthDecider`]), so the embedder never names an HTTP framework. + /// The decider gates the transcoded requests inline (other traffic with + /// [`with_auth_decider_scope`](Self::with_auth_decider_scope)) and also + /// backs the `/verify` forward-auth endpoint. Its signature is `axum`-free + /// (see [`hooks::AuthDecider`]), so the embedder never names an HTTP + /// framework. pub fn with_auth_decider(mut self, decider: Arc) -> Self { self.auth_decider = Some(decider); self } + /// Choose the traffic the injected [`AuthDecider`] gates, `transcoded` + /// by default; the `/verify` endpoint is never behind it, since it answers + /// for the decider. + /// + /// # Examples + /// + /// ``` + /// use structured_proxy::config::{ScopeConfig, Traffic}; + /// use structured_proxy::ProxyServer; + /// + /// # fn build() -> anyhow::Result<()> { + /// // Gate native gRPC calls as well as the transcoded ones. + /// let server = ProxyServer::from_yaml_str("service:\n name: demo\n")? + /// .with_auth_decider_scope(ScopeConfig::traffic([Traffic::Transcoded, Traffic::Grpc])); + /// # let _ = server; + /// # Ok(()) + /// # } + /// # build().unwrap(); + /// ``` + pub fn with_auth_decider_scope(mut self, scope: ScopeConfig) -> Self { + self.auth_decider_scope = Some(scope); + self + } + /// Back the stateless OIDC surface (discovery, JWKS, userinfo) with the /// embedder's key/client metadata (embedded Tier-2 hook). /// @@ -443,7 +467,7 @@ impl ProxyServer { /// No valid upstream address, or a configuration [`service`](Self::service) /// rejects. pub fn router(&self) -> anyhow::Result { - let (router, _) = self.routes(self.upstream()?)?; + let (router, _, _) = self.routes(self.upstream()?)?; Ok(router) } @@ -477,16 +501,19 @@ impl ProxyServer { /// # build().unwrap(); /// ``` pub fn service(&self, upstream: U) -> anyhow::Result> { - let (routes, cors) = self.routes(upstream.clone())?; + let (routes, cors, guards) = 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)) + Ok(ProxyService::new(upstream, routes, grpc_web_cors, guards)) } - /// 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)> { + /// Build the axum router with all endpoints, calling `upstream`, the CORS + /// policy it answers under, and the guards, for the traffic outside it. + fn routes( + &self, + upstream: U, + ) -> anyhow::Result<(Router, CorsLayer, Arc)> { // Enforce cross-field invariants on the embedded path too, where the // config is built directly instead of through `from_yaml_str`. self.config.validate()?; @@ -578,36 +605,20 @@ impl ProxyServer { let cors = self.build_cors()?; // Build transcoding routes from descriptor pool. - let mut transcode_routes = + let transcode_routes = transcode::routes_with_options(&pool, &self.config.aliases, &self.transcode); - // External authorization (Envoy ext_authz) gates only the proxied API - // routes, never health / metrics / discovery. It runs inside the auth - // layer below, so the Check call sees the identity headers the JWT - // middleware injected. - let authz = match self.config.auth.as_ref().and_then(|a| a.authz.as_ref()) { - Some(cfg) => auth::authz::Authz::build(cfg) - .map_err(|e| anyhow::anyhow!("invalid authz config: {e}"))?, + // JWT auth, if configured (auth.mode == "jwt"). + let auth = match &self.config.auth { + Some(cfg) => auth::Auth::build(cfg, self.token_verifier.clone()) + .map_err(|e| anyhow::anyhow!("invalid auth config: {e}"))?, None => None, }; - - // Order matters: in axum the LAST-added layer is outermost and runs - // FIRST. We want `authz -> AuthDecider -> handler`, so add the decider - // layer first (inner) and the authz layer second (outer). That way, when - // both are configured, ext_authz runs first and the in-process decider - // sees any headers the authz Check injected. - if let Some(decider) = &self.auth_decider { - transcode_routes = transcode_routes.layer(axum::middleware::from_fn_with_state( - decider.clone(), - embed::auth_decider_gate, - )); - } - if let Some(authz) = authz { - transcode_routes = transcode_routes.layer(axum::middleware::from_fn_with_state( - authz, - auth::authz::middleware, - )); - } + // Forward-auth verification endpoint, sharing the built Auth. + let forward_auth = auth.as_ref().and_then(|built| { + auth::forward::ForwardAuth::build(self.config.auth.as_ref()?, built.clone()) + }); + let guards = self.guards(auth, maintenance_exempt)?; // Health routes. Paths are configurable; the whole group is skippable. let health_routes = if self.config.health.enabled { @@ -703,105 +714,132 @@ impl ProxyServer { }, }; - // Rate limiting (Shield), if configured and enabled. - let shield = match &self.config.shield { - Some(cfg) => shield::Shield::build(cfg) - .map_err(|e| anyhow::anyhow!("invalid shield config: {e}"))?, - None => None, - }; - - // JWT auth, if configured (auth.mode == "jwt"). - let auth = match &self.config.auth { - Some(cfg) => auth::Auth::build(cfg, self.token_verifier.clone()) - .map_err(|e| anyhow::anyhow!("invalid auth config: {e}"))?, - None => None, - }; - - let mut router = Router::new() + let endpoints = Router::new() .merge(health_routes) .merge(metrics_routes) .merge(openapi_routes) .merge(oidc_routes) - .merge(embed::extra_routes_router(&self.extra_routes)) - .merge(transcode_routes); - // CORS is applied as the outermost layer below, so it wraps the auth and - // rate-limit enforcement: a short-circuited 401/429/503 still carries CORS - // headers, and preflight OPTIONS is answered before auth can reject it. - - // Forward-auth verification endpoint, sharing the built Auth. Mounted - // after the auth layer below so the endpoint itself is not gated by the - // JWT middleware (it answers the gate, it isn't behind it). - let forward_auth = auth.as_ref().and_then(|built| { - auth::forward::ForwardAuth::build(self.config.auth.as_ref()?, built.clone()) - }); - - // Duplicate-route collisions (including the verify path) were already - // rejected up front, before any router was built. - - // Two-phase rate limiting around auth. The post-auth phase (rules keyed - // by a validated JWT claim) is layered first so it sits *inside* auth and - // sees the verified claims; the pre-auth phase (IP / header keys) is - // layered after auth below so it runs *first* and sheds anonymous floods - // before any signature verification. - if let Some(shield) = &shield { - router = router.layer(axum::middleware::from_fn_with_state( - shield.clone(), - shield::post_auth_middleware, - )); - } - - if let Some(auth) = auth { - router = router.layer(axum::middleware::from_fn_with_state(auth, auth::middleware)); - } + .merge(embed::extra_routes_router(&self.extra_routes)); // Forward-auth `/verify` endpoint. An injected AuthDecider owns it when // present (in-process PDP); otherwise the config-driven JWT ForwardAuth - // backs it. Mounted after the auth layer so it is not itself JWT-gated. - if let Some(decider) = &self.auth_decider { - // Collision / shape of this path was already validated above. + // backs it. Its path was validated above. + let verify = if let Some(decider) = &self.auth_decider { let decider = decider.clone(); - let path = self.decider_verify_path(); - router = router.route( - &path, + Router::new().route( + &self.decider_verify_path(), axum::routing::any(move |req: axum::extract::Request| { let decider = decider.clone(); async move { embed::verify_via_decider(decider, req).await } }), - ); + ) } else if let Some(forward_auth) = &forward_auth { - router = router.merge(forward_auth.routes()); - } + forward_auth.routes() + } else { + Router::new() + }; - // Pre-auth phase, added before maintenance so maintenance wraps it (outer - // layers run first): a request rejected by the maintenance gate must not - // be charged against its rate-limit budget. Placed after the auth layer - // so it runs before auth, and after the verify route so that endpoint is - // rate-limited too (but not JWT-gated). - if let Some(shield) = &shield { - router = router.layer(axum::middleware::from_fn_with_state( - shield.clone(), - shield::pre_auth_middleware, - )); - } + // Each class of traffic behind the guards that cover it, so a request + // runs only the guards of its own class. A path no route answers + // reaches the plain 404 (or the fallback's own guards). + let router = Router::new() + .merge(guards.router(transcode_routes, guard::Class::Transcoded)) + .merge(guards.router(endpoints, guard::Class::Endpoints)) + .merge(guards.router(verify, guard::Class::Verify)) + .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.clone()).with_state(state); + + Ok((router, cors, Arc::new(guards))) + } + /// The guards the configuration and the hooks turn on, each with its + /// scope; `maintenance_exempt` lists the paths maintenance mode leaves + /// reachable. + /// + /// # Errors + /// + /// A malformed shield, authz or concurrency section, or a scope that + /// covers no traffic, names an invalid path glob or an invalid method. + fn guards( + &self, + auth: Option>, + maintenance_exempt: Vec, + ) -> anyhow::Result { + use config::Traffic::{Endpoints, Grpc, Transcoded}; + let scope = |config: Option<&ScopeConfig>, default: &[config::Traffic], what: &str| { + guard::Scope::compile(config, default, what).map_err(anyhow::Error::msg) + }; + let mut guards = guard::Guards::default(); // 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, + let maintenance = &self.config.maintenance; + if maintenance.enabled { + guards.maintenance = Some(( + Arc::new(guard::Maintenance { + exempt: maintenance_exempt, + message: maintenance.message.clone(), + }), + scope( + maintenance.scope.as_ref(), + &[Transcoded, Endpoints], + "maintenance", + )?, )); } - 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.clone()).with_state(state); - - Ok((router, cors)) + if let Some(cfg) = &self.config.concurrency { + guards.concurrency = Some(( + guard::Concurrency::build(cfg).map_err(anyhow::Error::msg)?, + scope( + cfg.scope.as_ref(), + &[Transcoded, Endpoints, Grpc], + "concurrency", + )?, + )); + } + if let Some(cfg) = &self.config.shield { + if let Some(shield) = shield::Shield::build(cfg) + .map_err(|e| anyhow::anyhow!("invalid shield config: {e}"))? + { + guards.shield = Some(( + shield, + scope(cfg.scope.as_ref(), &[Transcoded, Endpoints], "shield")?, + )); + } + } + let auth_config = self.config.auth.as_ref(); + if let Some(auth) = auth { + guards.auth = Some(( + auth, + scope( + auth_config.and_then(|a| a.scope.as_ref()), + &[Transcoded, Endpoints], + "auth", + )?, + )); + } + if let Some(cfg) = auth_config.and_then(|a| a.authz.as_ref()) { + if let Some(authz) = auth::authz::Authz::build(cfg) + .map_err(|e| anyhow::anyhow!("invalid authz config: {e}"))? + { + guards.authz = Some(( + authz, + scope(cfg.scope.as_ref(), &[Transcoded], "auth.authz")?, + )); + } + } + if let Some(decider) = &self.auth_decider { + guards.decider = Some(( + decider.clone(), + scope( + self.auth_decider_scope.as_ref(), + &[Transcoded], + "auth_decider", + )?, + )); + } + Ok(guards) } fn build_openapi_routes(&self, pool: &DescriptorPool) -> Router @@ -952,39 +990,6 @@ 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(maintenance): State>, - request: Request, - next: Next, -) -> Response { - if maintenance.exempts(request.uri().path()) { - return next.run(request).await; - } - ( - StatusCode::SERVICE_UNAVAILABLE, - [("retry-after", "300")], - maintenance.message.clone(), - ) - .into_response() -} - /// A [`ProxyState`] for tests whose routers never call the upstream: a lazy /// channel to a port nothing listens on. #[cfg(test)] diff --git a/src/service.rs b/src/service.rs index 8a44832..4b681f8 100644 --- a/src/service.rs +++ b/src/service.rs @@ -16,9 +16,11 @@ 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::{Layer, Service, ServiceExt}; use tower_http::cors::{Cors, CorsLayer}; +use crate::guard::{BoxedService, Class, GrpcRejections, Guards}; + /// 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. @@ -41,7 +43,9 @@ use crate::upstream::{ /// 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. +/// [`with_fallback`](Self::with_fallback). Native gRPC and the fallback pass +/// only the guards whose scope names them (`grpc`, `fallback`); a guard's +/// rejection of a gRPC call is a gRPC status. /// /// 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 @@ -63,18 +67,40 @@ use crate::upstream::{ /// # } /// # build().unwrap(); /// ``` -#[derive(Clone, Debug)] +#[derive(Clone)] 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 gRPC path behind the guards that cover it, one per protocol; + /// `None` when no guard does, so gRPC calls pay nothing for guards. + guarded: Option, + /// The guards, for the fallback an embedder sets later. + guards: Arc, /// The connection the requests arrive on, set per connection by the /// server. connection: Option, } +impl std::fmt::Debug for ProxyService { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ProxyService") + .field("guarded_grpc", &self.guarded.is_some()) + .field("connection", &self.connection) + .finish_non_exhaustive() + } +} + +/// The guarded gRPC path, each protocol's stack built once. +#[derive(Clone)] +struct GuardedGrpc { + grpc: BoxedService, + web: BoxedService, + web_text: BoxedService, +} + /// gRPC-Web pass-through, one CORS-wrapped path per encoding. #[derive(Clone, Debug)] struct GrpcWebCors { @@ -106,6 +132,23 @@ impl Service> for Forward { } } +/// The guarded path's end: the guards run on axum's request type. +impl Service for Forward { + type Response = http::Response; + type Error = Infallible; + type Future = PassThrough; + + #[inline] + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, request: axum::extract::Request) -> Self::Future { + let request = request.map(tonic::body::Body::new); + 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. /// @@ -165,28 +208,55 @@ impl ConnectionInfo { 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, + pub(crate) fn new( + upstream: U, + routes: axum::Router, + grpc_web_cors: Option, + guards: Arc, + ) -> Self { + let forward = |protocol| Forward { + upstream: upstream.clone(), + protocol, + }; + let guarded = guards.cover(Class::Grpc).then(|| { + // CORS outermost, so a rejected gRPC-Web call still carries it + // and a preflight is answered before any guard. + let stack = |protocol| { + let rejections = GrpcRejections { + inner: guards.service(BoxedService::new(forward(protocol)), Class::Grpc), + protocol, + }; + match (&grpc_web_cors, protocol) { + (Some(cors), GrpcProtocol::Web | GrpcProtocol::WebText) => { + BoxedService::new(cors.layer(rejections)) + } + _ => BoxedService::new(rejections), + } }; - GrpcWebCors { - web: cors.layer(forward(GrpcProtocol::Web)), - web_text: cors.layer(forward(GrpcProtocol::WebText)), + GuardedGrpc { + grpc: stack(GrpcProtocol::Grpc), + web: stack(GrpcProtocol::Web), + web_text: stack(GrpcProtocol::WebText), } }); + let grpc_web = grpc_web_cors.map(|cors| GrpcWebCors { + web: cors.layer(forward(GrpcProtocol::Web)), + web_text: cors.layer(forward(GrpcProtocol::WebText)), + }); Self { upstream, routes, grpc_web, + guarded, + guards, 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. + /// tower service. Only the guards whose scope names `fallback` traffic see + /// them; CORS and tracing are the fallback's own. /// /// 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 @@ -198,7 +268,13 @@ impl ProxyService { F::Response: IntoResponse, F::Future: Send + 'static, { - self.routes = self.routes.fallback_service(fallback); + self.routes = if self.guards.cover(Class::Fallback) { + let fallback = BoxedService::new(fallback.map_response(IntoResponse::into_response)); + self.routes + .fallback_service(self.guards.service(fallback, Class::Fallback)) + } else { + self.routes.fallback_service(fallback) + }; self } @@ -238,24 +314,29 @@ impl ProxyService { upstream: self.upstream.clone(), routes: self.routes.clone(), grpc_web: self.grpc_web.clone(), + guarded: self.guarded.clone(), + guards: self.guards.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), - })) +/// The connection `request` came on: `connection`, the one given to +/// [`ProxyService::for_connection`], else the peer an outer axum server +/// recorded as `ConnectInfo`, so the upstream sees the same client the proxy's +/// middleware does. A free function, so the caller keeps its other fields. +fn connection_of( + connection: Option<&ConnectionInfo>, + request: &http::Request, +) -> Option { + if let Some(connection) = 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 @@ -278,13 +359,34 @@ where // 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(|| { + let call = grpc_protocol(request.headers()); + let protocol = call.or_else(|| { is_grpc_web_preflight(request.method(), request.headers()).then_some(GrpcProtocol::Web) }); - let inner = if let Some(protocol) = protocol { + let inner = if let (Some(protocol), Some(guarded)) = (call, &mut self.guarded) { + // The guards read the peer as axum's `ConnectInfo`, the upstream + // as tonic records it. A preflight is not a call and skips them. + if let Some(connection) = connection_of(self.connection.as_ref(), &request) { + let extensions = request.extensions_mut(); + if let Some(remote) = connection.remote_addr() { + extensions.insert(ConnectInfo(remote)); + } + connection.into_tonic_extensions(extensions); + } + let service = match protocol { + GrpcProtocol::Grpc => &mut guarded.grpc, + GrpcProtocol::Web => &mut guarded.web, + GrpcProtocol::WebText => &mut guarded.web_text, + }; + // Every service on the guarded path is ready at once: the guards + // are middleware and `Forward` waits for the upstream per request. + Inner::Guarded { + future: service.call(request.map(axum::body::Body::new)), + } + } else 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) { + if let Some(connection) = connection_of(self.connection.as_ref(), &request) { connection.into_tonic_extensions(request.extensions_mut()); } let request = request.map(tonic::body::Body::new); @@ -313,7 +415,7 @@ where extensions.insert(ConnectInfo(remote)); } extensions.insert(connection.clone()); - } else if let Some(connection) = self.connection_of(&request) { + } else if let Some(connection) = connection_of(None, &request) { request.extensions_mut().insert(connection); } Inner::Routes { @@ -347,6 +449,10 @@ pin_project! { #[pin] future: tower_http::cors::ResponseFuture>, }, + Guarded { + #[pin] + future: >::Future, + }, } } @@ -359,6 +465,7 @@ impl Future for ResponseFuture { InnerProj::Routes { future } => future.poll(cx), InnerProj::Grpc { call } => call.poll(cx), InnerProj::GrpcWeb { future } => future.poll(cx), + InnerProj::Guarded { future } => future.poll(cx), } } } diff --git a/src/service/tests.rs b/src/service/tests.rs index 4c18c45..20989fc 100644 --- a/src/service/tests.rs +++ b/src/service/tests.rs @@ -61,7 +61,7 @@ impl Service> for Recorder { /// 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) + ProxyService::new(upstream, routes, None, Arc::default()) } fn grpc_request(path: &str) -> http::Request { @@ -287,8 +287,8 @@ fn peer_routes() -> axum::Router { #[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 proxy = ProxyService::new(Recorder::default(), peer_routes(), None, Arc::default()) + .for_connection(connection()); let response = proxy .oneshot(http::Request::get("/peer").body(Body::empty()).unwrap()) .await @@ -298,8 +298,8 @@ async fn an_http_request_carries_its_peer_for_the_middleware() { #[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 proxy = ProxyService::new(Recorder::default(), peer_routes(), None, Arc::default()) + .for_connection(connection()); let response = proxy .oneshot( http::Request::get("/connection") @@ -313,7 +313,7 @@ async fn an_http_request_carries_its_connection_for_the_transcoder() { #[tokio::test] async fn an_http_request_without_a_connection_carries_none() { - let proxy = ProxyService::new(Recorder::default(), peer_routes(), None); + let proxy = ProxyService::new(Recorder::default(), peer_routes(), None, Arc::default()); let response = proxy .oneshot( http::Request::get("/connection") @@ -341,7 +341,7 @@ async fn a_grpc_request_on_an_axum_server_carries_its_peer_to_the_upstream() { #[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 proxy = ProxyService::new(Recorder::default(), peer_routes(), None, Arc::default()); let mut request = http::Request::get("/connection") .body(Body::empty()) .unwrap(); diff --git a/src/shield/mod.rs b/src/shield/mod.rs index 3ff33bf..c714138 100644 --- a/src/shield/mod.rs +++ b/src/shield/mod.rs @@ -24,10 +24,9 @@ use std::sync::Arc; use std::time::Duration; use axum::extract::Request; -use axum::http::{HeaderMap, StatusCode}; +use axum::http::HeaderMap; use axum::middleware::Next; -use axum::response::{IntoResponse, Response}; -use axum::Json; +use axum::response::Response; use crate::config::ShieldConfig; use gcra::Verdict; @@ -530,14 +529,7 @@ fn attach_rate_headers(headers: &mut HeaderMap, limit: u64, remaining: u64, verd /// A `429` response carrying the rate-limit headers plus `Retry-After`. fn too_many_requests(limit: u64, remaining: u64, verdict: &Verdict) -> Response { - let mut response = ( - StatusCode::TOO_MANY_REQUESTS, - Json(serde_json::json!({ - "error": "RESOURCE_EXHAUSTED", - "message": "rate limit exceeded", - })), - ) - .into_response(); + let mut response = crate::guard::reject(tonic::Code::ResourceExhausted, "rate limit exceeded"); let headers = response.headers_mut(); attach_rate_headers(headers, limit, remaining, verdict); if let Ok(v) = secs_ceil(verdict.retry_after).to_string().parse() { diff --git a/src/shield/tests.rs b/src/shield/tests.rs index 9236bdc..91fec37 100644 --- a/src/shield/tests.rs +++ b/src/shield/tests.rs @@ -33,6 +33,7 @@ fn config(profiles: Vec<(&str, &str, Option)>, rules: Vec) limit_service: None, sync: None, trusted_proxies: Vec::new(), + scope: None, } } @@ -373,6 +374,7 @@ mod two_phase { }), forward_auth: None, authz: None, + scope: None, }; Auth::build(&cfg, None).unwrap().unwrap() } diff --git a/src/tests.rs b/src/tests.rs index 4186a2c..d881ab5 100644 --- a/src/tests.rs +++ b/src/tests.rs @@ -27,40 +27,6 @@ upstream: 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 diff --git a/src/transcode/error.rs b/src/transcode/error.rs index 3fb8bf9..11ff761 100644 --- a/src/transcode/error.rs +++ b/src/transcode/error.rs @@ -112,6 +112,12 @@ pub(crate) fn malformed_response(details: Option<&StatusDetails>) -> Response { (StatusCode::INTERNAL_SERVER_ERROR, Json(body)).into_response() } +/// The error body of a request a guard turned away: the shape of the +/// transcoder's own errors, with empty `details`. +pub(crate) fn guard_error_body(code: tonic::Code, message: &str) -> Value { + body(code, message, Some(RenderedDetails::default())) +} + /// The JSON error body for a failed call, shared by the unary response and the /// terminal frame of a stream so a client parses one shape everywhere. /// diff --git a/src/upstream.rs b/src/upstream.rs index d9075c8..423f78d 100644 --- a/src/upstream.rs +++ b/src/upstream.rs @@ -155,8 +155,15 @@ impl Future for PassThrough { /// 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(); + trailers_only(tonic::Status::from_error(error), protocol) +} + +/// `status` as a trailers-only gRPC answer in `protocol`. +pub(crate) fn trailers_only( + status: tonic::Status, + protocol: GrpcProtocol, +) -> http::Response { + let mut response: http::Response = status.into_http(); if protocol != GrpcProtocol::Grpc { response.headers_mut().insert( http::header::CONTENT_TYPE, diff --git a/tests/edge.rs b/tests/edge.rs index 44dfe75..069a260 100644 --- a/tests/edge.rs +++ b/tests/edge.rs @@ -156,6 +156,16 @@ async fn listen( upstream: common::Upstream, extra_yaml: &str, fallback: Option, +) -> SocketAddr { + listen_configured(upstream, extra_yaml, fallback, |server| server).await +} + +/// [`listen`], with `configure` applied to the server before it is built. +async fn listen_configured( + upstream: common::Upstream, + extra_yaml: &str, + fallback: Option, + configure: impl FnOnce(ProxyServer) -> ProxyServer, ) -> SocketAddr { let pool = pool(); let service = Edge { pool: pool.clone() }; @@ -164,11 +174,13 @@ async fn listen( 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 server = configure( + 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); @@ -176,9 +188,11 @@ async fn listen( tokio::spawn(structured_proxy::serve(listener, proxy)); } common::Upstream::InProcess => { - let server = ProxyServer::from_yaml_str(extra_yaml) - .unwrap() - .with_descriptors(pool); + let server = configure( + ProxyServer::from_yaml_str(extra_yaml) + .unwrap() + .with_descriptors(pool), + ); let mut proxy = server .service(tonic::service::Routes::new(service)) .unwrap(); @@ -321,6 +335,105 @@ async fn maintenance_gates_the_routes_but_not_the_fallback_or_native_grpc() { let seen = grpc_echo(addr, "native").await.unwrap(); assert_eq!(field(&seen, "name"), "native"); } + +// --- scoped guards --------------------------------------------------------------- + +async fn a_guard_scoped_to_grpc_answers_native_calls_with_a_status() { + // A gRPC client reads only a status: the guard's 503 reaches it as + // UNAVAILABLE with the guard's message, while the REST route it does not + // cover stays open. + let yaml = "maintenance:\n enabled: true\n message: \"back soon\"\n scope:\n traffic: [grpc]\n"; + let addr = listen(UPSTREAM, yaml, None).await; + let status = grpc_echo(addr, "native").await.unwrap_err(); + assert_eq!(status.code(), tonic::Code::Unavailable); + assert_eq!(status.message(), "back soon"); + let (status, body, _) = http1_get(addr, "/v1/echo/rest").await; + assert_eq!(status, 200, "{body}"); +} + +async fn a_guard_scoped_to_the_fallback_covers_it_and_nothing_else() { + let fallback = axum::Router::new().route( + "/static/index.html", + axum::routing::get(|| async { "static" }), + ); + let yaml = "maintenance:\n enabled: true\n scope:\n traffic: [fallback]\n"; + let addr = listen(UPSTREAM, yaml, Some(fallback)).await; + let (status, body, _) = http1_get(addr, "/static/index.html").await; + assert_eq!(status, 503, "{body}"); + let body: Value = serde_json::from_str(&body).unwrap(); + assert_eq!(body["error"], "UNAVAILABLE"); + let (status, body, _) = http1_get(addr, "/v1/echo/rest").await; + assert_eq!(status, 200, "{body}"); +} + +async fn a_guard_narrowed_by_path_leaves_the_other_paths_alone() { + let yaml = "maintenance:\n enabled: true\n scope:\n paths: [\"/v1/echo/closed\"]\n"; + let addr = listen(UPSTREAM, yaml, None).await; + let (status, _, _) = http1_get(addr, "/v1/echo/closed").await; + assert_eq!(status, 503); + let (status, body, _) = http1_get(addr, "/v1/echo/open").await; + assert_eq!(status, 200, "{body}"); +} + +async fn the_auth_decider_scoped_to_grpc_gates_native_calls_by_the_client_address() { + // The decider sees the gRPC client's address, and its 403 reaches the + // client as PERMISSION_DENIED. + let decider = std::sync::Arc::new(PeerDecider::default()); + let scope = structured_proxy::config::ScopeConfig::traffic([ + structured_proxy::config::Traffic::Grpc, + ]); + let addr = listen_configured(UPSTREAM, "", None, { + let decider = decider.clone(); + move |server| { + server + .with_auth_decider(decider) + .with_auth_decider_scope(scope) + } + }) + .await; + let status = grpc_echo(addr, "native").await.unwrap_err(); + assert_eq!(status.code(), tonic::Code::PermissionDenied); + let peer = decider.peer.lock().unwrap().unwrap(); + assert!(peer.ip().is_loopback(), "{peer}"); + // Scoped to gRPC: the transcoded route is not behind it. + let (status, body, _) = http1_get(addr, "/v1/echo/rest").await; + assert_eq!(status, 200, "{body}"); +} +} + +/// Denies every request with `403`, recording the peer it saw. +#[derive(Default)] +struct PeerDecider { + peer: std::sync::Mutex>, +} + +#[async_trait::async_trait] +impl structured_proxy::hooks::AuthDecider for PeerDecider { + async fn decide( + &self, + req: &structured_proxy::hooks::RequestParts<'_>, + ) -> structured_proxy::hooks::Decision { + *self.peer.lock().unwrap() = Some(req.peer); + structured_proxy::hooks::Decision::Deny { + status: StatusCode::FORBIDDEN, + body: bytes::Bytes::from_static(br#"{"error":"denied"}"#), + } + } +} + +#[test] +fn a_scope_that_covers_no_traffic_fails_the_build() { + let server = + ProxyServer::from_yaml_str("maintenance:\n enabled: true\n scope:\n traffic: []\n") + .unwrap() + .with_descriptors(pool()); + let Err(err) = server.service(tonic::service::Routes::default()) else { + panic!("a scope with no traffic must be refused"); + }; + assert!( + err.to_string().contains("maintenance.scope.traffic"), + "{err}" + ); } // --- deadlines --------------------------------------------------------------- diff --git a/tests/embedded.rs b/tests/embedded.rs index 788d38b..020b108 100644 --- a/tests/embedded.rs +++ b/tests/embedded.rs @@ -45,6 +45,7 @@ fn embedded_config_is_constructible() { // the config via from_file / from_yaml_str, where the default list applies). forwarded_headers: vec!["authorization".into()], streaming: Default::default(), + concurrency: None, }; // The server accepts a programmatically-built config (the embedded path). let _server = ProxyServer::from_config(config); From eab80c511bd2eb80bab2e4a04b67a3c9ae9532dc Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 10:58:45 +0300 Subject: [PATCH 2/9] ci: build and test the musl targets on every change The release packages ship static musl binaries, which were built only when a release was cut: a dependency that breaks the target would pass CI and fail the release. Build the packaged binary and run the tests on x86_64 and aarch64 musl in CI. --- .github/workflows/ci.yml | 40 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e8811bf..8716891 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -134,6 +134,46 @@ jobs: if: matrix.backend.name == 'rust_crypto' run: cargo publish --dry-run --features cli + # The release packages ship static musl binaries; build and test that target + # on every change, so a dependency that breaks it fails here, not at release. + musl: + name: musl (${{ matrix.target }}) + runs-on: ${{ matrix.runs-on }} + strategy: + fail-fast: false + matrix: + include: + - target: x86_64-unknown-linux-musl + runs-on: ubuntu-latest + - target: aarch64-unknown-linux-musl + runs-on: ubuntu-24.04-arm + steps: + - uses: actions/checkout@v7 + with: + persist-credentials: false + - uses: dtolnay/rust-toolchain@stable + with: + targets: ${{ matrix.target }} + - name: Install musl tools + run: | + sudo apt-get update + sudo apt-get install -y musl-tools + # Third-party action: pinned to a commit SHA (dependabot keeps it fresh). + - uses: Swatinem/rust-cache@e18b497796c12c097a38f9edb9d0641fb99eee32 # v2 + with: + key: ${{ matrix.target }} + # Third-party action: pinned to a commit SHA (dependabot keeps it fresh). + - uses: taiki-e/install-action@94c31af3204a9f15ab40b35ad084410b905bbc73 # v2 + with: + tool: cargo-nextest + + # The build the release workflow packages. + - name: Build + run: cargo build --release --target ${{ matrix.target }} --features cli --bin structured-proxy + + - name: Test + run: cargo nextest run --target ${{ matrix.target }} --features cli + security-audit: name: Security Audit runs-on: ubuntu-latest From fdf717262ab7631077a58cc0d4f3ae6f60a26ad9 Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 10:59:54 +0300 Subject: [PATCH 3/9] refactor(guards): use std paths, the crate targets std only --- src/guard.rs | 11 ++++------- src/guard/concurrency.rs | 9 +++------ src/guard/grpc.rs | 9 ++++----- src/lib.rs | 2 -- 4 files changed, 11 insertions(+), 20 deletions(-) diff --git a/src/guard.rs b/src/guard.rs index 0b5c988..2b99808 100644 --- a/src/guard.rs +++ b/src/guard.rs @@ -8,19 +8,16 @@ mod concurrency; mod grpc; -use alloc::borrow::Cow; -use alloc::string::String; -use alloc::sync::Arc; -use alloc::vec::Vec; -use core::convert::Infallible; -use core::task::{Context, Poll}; +use std::borrow::Cow; +use std::convert::Infallible; +use std::sync::Arc; +use std::task::{Context, Poll}; use axum::extract::{Request, State}; use axum::middleware::{from_fn_with_state, Next}; use axum::response::{IntoResponse, Response}; use axum::{Json, Router}; use futures::future::Either; -// no-std: a segment-wise glob matcher over `&str` (globset needs std). use globset::{GlobBuilder, GlobSet, GlobSetBuilder}; use http::{Method, StatusCode}; use tower::util::BoxCloneSyncService; diff --git a/src/guard/concurrency.rs b/src/guard/concurrency.rs index d4a5bbc..a46bd3d 100644 --- a/src/guard/concurrency.rs +++ b/src/guard/concurrency.rs @@ -2,11 +2,9 @@ //! turned away at once rather than queued, since a queue behind a saturated //! upstream only adds latency to what will time out anyway. -use alloc::format; -use alloc::string::String; -use alloc::sync::Arc; -use core::pin::Pin; -use core::task::{Context, Poll}; +use std::pin::Pin; +use std::sync::Arc; +use std::task::{Context, Poll}; use axum::body::Body; use axum::extract::{Request, State}; @@ -15,7 +13,6 @@ use axum::response::Response; use bytes::Bytes; use http_body::{Frame, SizeHint}; use pin_project_lite::pin_project; -// no-std: an `AtomicUsize` slot counter with a drop guard (only try-acquire is used). use tokio::sync::{OwnedSemaphorePermit, Semaphore}; use super::reject; diff --git a/src/guard/grpc.rs b/src/guard/grpc.rs index 2fa696b..daedebf 100644 --- a/src/guard/grpc.rs +++ b/src/guard/grpc.rs @@ -1,10 +1,9 @@ //! Guard rejections on the gRPC path, answered in gRPC. -use alloc::string::ToString; -use core::convert::Infallible; -use core::future::Future; -use core::pin::Pin; -use core::task::{ready, Context, Poll}; +use std::convert::Infallible; +use std::future::Future; +use std::pin::Pin; +use std::task::{ready, Context, Poll}; use axum::extract::Request; use axum::response::Response; diff --git a/src/lib.rs b/src/lib.rs index 0a388cf..380f581 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -53,8 +53,6 @@ compile_error!( (or neither, and inject a verifier with `ProxyServer::with_token_verifier`)" ); -extern crate alloc; - pub mod auth; pub mod config; mod cors; From 7c5b54d2b70d5ee19d4e9ea9b628989f3c519c59 Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 11:09:59 +0300 Subject: [PATCH 4/9] feat(serve)!: built-in TLS, mTLS and a connection limit - `serve_with(listener, service, ServeOptions)`: TLS termination with any rustls config (ALPN h2 + http/1.1 filled in), the handshake in each connection's task under a 10 s limit, and `max_connections`, which waits for a free slot before accepting, so excess clients wait in the listen backlog. `serve` is `serve_with` with no options. - `listen.tls` (cert_file, key_file, client_ca_file, client_auth) and `listen.max_connections` configure `ProxyServer::serve`; `ProxyServer::serve_options` hands the same to an embedder's listener. A verified client certificate reaches an in-process tonic upstream as `Request::peer_certs`. - Connections run on hyper-util instead of `axum::serve`; axum no longer needs its `http2` feature. - README: TLS and connection limits; the TLS crypto section covers the listener too. BREAKING CHANGE: `ListenConfig` gains `max_connections` and `tls` and rejects unknown keys; `structured_proxy::service::serve` moves to `structured_proxy::serve`. Part of #118 --- Cargo.toml | 8 +- README.md | 65 ++++++++-- src/config.rs | 51 +++++++- src/lib.rs | 47 ++++++- src/serve.rs | 222 ++++++++++++++++++++++++++++++++ src/serve/tests.rs | 121 +++++++++++++++++ src/service.rs | 64 ++------- src/tls.rs | 65 +++++++++- src/tls/testdata/client-ca.pem | 12 ++ src/tls/testdata/client.key.pem | 5 + src/tls/testdata/client.pem | 12 ++ src/tls/testdata/generate.sh | 19 ++- tests/embedded.rs | 2 + tests/tls.rs | 174 ++++++++++++++++++++++--- 14 files changed, 763 insertions(+), 104 deletions(-) create mode 100644 src/serve.rs create mode 100644 src/serve/tests.rs create mode 100644 src/tls/testdata/client-ca.pem create mode 100644 src/tls/testdata/client.key.pem create mode 100644 src/tls/testdata/client.pem diff --git a/Cargo.toml b/Cargo.toml index 41f0bc7..00810d3 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -32,9 +32,8 @@ required-features = ["cli"] doc = false [dependencies] -# 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"] } +# HTTP framework of the proxy's routes; `serve` runs connections on hyper-util. +axum = { version = "0.8", features = ["macros"] } # `util`: the guarded gRPC and fallback paths are boxed services. tower = { version = "0.5", features = ["util"] } tower-http = { version = "0.7", features = ["cors", "trace"] } @@ -49,6 +48,9 @@ bytes = "1" http-body = "1" # The response futures of `ProxyService`, named without boxing each one. pin-project-lite = "0.2" +# `serve_with`'s connections: HTTP/1.1 and HTTP/2 on one port, over TCP or +# TLS. Already in the tree through axum and tonic. +hyper-util = { version = "0.1", features = ["server-auto", "service", "tokio"] } # gRPC client (to upstream service). `tls-connect-info` gives a rustls server # stream the connection record (client certificates) a tonic handler reads, so diff --git a/README.md b/README.md index 896a6bc..4916971 100644 --- a/README.md +++ b/README.md @@ -21,7 +21,7 @@ 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 +- **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 the built-in TLS listener (mTLS with a client CA, a cap on open connections) or your own TLS, 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 @@ -69,6 +69,17 @@ log line (`RUST_LOG=info`) states the count and where it came from. # my-service.yaml listen: http: "0.0.0.0:8080" + # Optional: most connections served at once; past it the next one waits in + # the listen backlog until a connection closes. Unset: no limit. + # max_connections: 10000 + # Optional: TLS on the listener (REST and gRPC share the port; ALPN offers + # h2 and http/1.1). With client_ca_file, client certificates are verified + # (mTLS) and reach an in-process tonic upstream as Request::peer_certs. + # tls: + # cert_file: "/etc/proxy/tls.crt" # PEM chain, leaf first + # key_file: "/etc/proxy/tls.key" # PEM private key + # client_ca_file: "/etc/proxy/ca.crt" + # client_auth: required # or `optional` # The gRPC service behind the proxy. Required by the standalone binary; an # embedder with an in-process upstream leaves it out. @@ -744,9 +755,36 @@ structured_proxy::serve(listener, proxy).await?; 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. +### TLS and connection limits + +`listen.tls` and `listen.max_connections` configure the listener of +`ProxyServer::serve`. An embedder with a listener of its own gets the same +from `ProxyServer::serve_options`, or builds `ServeOptions` in code, and runs +`structured_proxy::serve_with`: + +```rust +use structured_proxy::{ProxyServer, ServeOptions}; + +# async fn run(tls: rustls::ServerConfig, grpc: tonic::service::Routes) -> anyhow::Result<()> { +let proxy = ProxyServer::from_file(std::path::Path::new("my-service.yaml"))?.service(grpc)?; +let listener = tokio::net::TcpListener::bind("0.0.0.0:8443").await?; +// Any rustls config: your own certificate resolver, client verifier, ... +let options = ServeOptions::new().tls(tls).max_connections(10_000); +structured_proxy::serve_with(listener, proxy, options).await?; +# Ok(()) +# } +``` + +The TLS handshake runs in each connection's task with a ten-second limit, so +a slow or stalled client holds up no one else. A client certificate the +listener verified reaches a tonic handler in process as `Request::peer_certs`. +TLS needs a rustls crypto provider: the one a crypto backend feature brings, +or the one your process installed (see [TLS crypto](#tls-crypto)). + ### Behind your own TLS -`serve` speaks cleartext. For TLS, run the service on your own acceptor: one +For a server of your own (another TLS stack, a Unix socket), 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 @@ -998,16 +1036,18 @@ async-trait = "0.1" serde_json = "1" ``` -which links no JWT or TLS crypto (see [Outbound TLS](#outbound-tls)), and +which links no JWT or TLS crypto (see [TLS crypto](#tls-crypto)), and supplies the backend from its own binary. With no verifier injected and no backend feature, an `auth.mode: "jwt"` config is rejected at startup with that instruction, rather than silently accepting tokens. -## Outbound TLS +## TLS crypto -The proxy's own HTTP calls (JWKS fetches, the rate-limit service) use rustls -with Mozilla's root store bundled from `webpki-roots`, so no system CA bundle is +All of the proxy's TLS is rustls: the listener (`listen.tls`, see +[TLS and connection limits](#tls-and-connection-limits)) and its own outbound +HTTP calls (JWKS fetches, the rate-limit service). Outbound calls trust +Mozilla's root store bundled from `webpki-roots`, so no system CA bundle is needed. The rustls crypto provider is, in order: 1. the one your process installed with @@ -1016,13 +1056,12 @@ needed. The rustls crypto provider is, in order: 2. aws-lc, with the `aws_lc_rs` feature; 3. the pure-Rust RustCrypto provider (`rustls-rustcrypto`), with `rust_crypto`. -Neither `ring` nor aws-lc is linked unless you ask: the default build contains -no C crypto, which CI checks. A `default-features = false` build links no TLS -crypto provider at all, so a crate that only transcodes pulls in neither `rsa` -nor `rustls-rustcrypto`. If such a build configures a JWKS endpoint or the -rate-limit service, install a provider before building the proxy (the client -needs one even for an `http://` endpoint); otherwise startup fails with an error -that says so. +The default build is pure Rust; aws-lc (C) comes with `aws_lc_rs`. A +`default-features = false` build brings no provider, so a crate that only +transcodes stays free of crypto dependencies. Such a build that terminates TLS, +or configures a JWKS endpoint or the rate-limit service, installs a provider +before building the proxy (the outbound client needs one even for an +`http://` endpoint); otherwise startup fails with an error that says so. The RustCrypto provider verifies RSA server signatures with `rsa`, under the same RUSTSEC-2023-0071 note as the `rust_crypto` JWT backend: only public-key diff --git a/src/config.rs b/src/config.rs index 98c59a7..87e585b 100644 --- a/src/config.rs +++ b/src/config.rs @@ -364,12 +364,57 @@ where Ok(yaml_sources.into_iter().map(Into::into).collect()) } -/// Listen address configuration. +/// The listener [`ProxyServer::serve`](crate::ProxyServer::serve) runs. #[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] pub struct ListenConfig { - /// HTTP listen address (default: "0.0.0.0:8080"). + /// Listen address (default: "0.0.0.0:8080"). #[serde(default = "default_http_listen")] pub http: String, + /// Most connections served at once; past it the next connection is + /// accepted when one closes. Unset: no limit. + #[serde(default)] + pub max_connections: Option, + /// TLS on the listener, mTLS with `client_ca_file`. Unset: cleartext. + #[serde(default)] + pub tls: Option, +} + +/// TLS for the listener. +/// +/// ```yaml +/// tls: +/// cert_file: /etc/proxy/tls.crt # PEM chain, leaf first +/// key_file: /etc/proxy/tls.key # PEM private key +/// client_ca_file: /etc/proxy/ca.crt # optional: verify client certificates +/// client_auth: required # or `optional` +/// ``` +#[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ListenTlsConfig { + /// The server certificate chain, PEM, leaf first. + pub cert_file: PathBuf, + /// The server private key, PEM (PKCS#8, PKCS#1 or SEC1). + pub key_file: PathBuf, + /// CA certificates, PEM, that client certificates are verified against. + /// Unset: clients present none. + #[serde(default)] + pub client_ca_file: Option, + /// Whether a client must present a certificate once `client_ca_file` + /// is set. + #[serde(default)] + pub client_auth: ClientAuth, +} + +/// Whether a TLS client must present a certificate. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ClientAuth { + /// A client without a valid certificate is refused in the handshake. + #[default] + Required, + /// A certificate is verified when presented; a client may present none. + Optional, } fn default_http_listen() -> String { @@ -380,6 +425,8 @@ impl Default for ListenConfig { fn default() -> Self { Self { http: default_http_listen(), + max_connections: None, + tls: None, } } } diff --git a/src/lib.rs b/src/lib.rs index 380f581..17ce749 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -61,6 +61,7 @@ mod guard; pub mod hooks; pub mod oidc; pub mod openapi; +mod serve; pub mod service; pub mod shield; mod tls; @@ -71,7 +72,8 @@ pub mod upstream; /// [`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}; +pub use serve::{serve, serve_with, ServeOptions}; +pub use service::{ConnectionInfo, ProxyService}; use axum::extract::State; use axum::http::StatusCode; @@ -949,21 +951,52 @@ impl ProxyServer { }) } + /// The [`ServeOptions`] of `listen:`: TLS from `listen.tls` (mTLS with its + /// `client_ca_file`) and the `listen.max_connections` cap, for + /// [`serve_with`] on a listener of your own. + /// + /// # Errors + /// + /// A `max_connections` of zero, or TLS files that cannot be loaded, or no + /// rustls crypto provider for TLS. + pub fn serve_options(&self) -> anyhow::Result { + let listen = &self.config.listen; + let mut options = ServeOptions::new(); + if let Some(max) = listen.max_connections { + anyhow::ensure!(max > 0, "listen.max_connections must be at least 1"); + options = options.max_connections(max); + } + if let Some(tls) = &listen.tls { + options = options.tls( + tls::server_config(tls).map_err(|e| anyhow::anyhow!("invalid listen.tls: {e}"))?, + ); + } + Ok(options) + } + /// 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`]). + /// configured upstream address: REST and native gRPC on one port, with + /// the TLS and connection limit of `listen:` (see + /// [`serve_options`](Self::serve_options) and [`serve_with`]). /// /// # Errors /// - /// What [`upstream`](Self::upstream) and [`service`](Self::service) - /// reject, an invalid listen address, or a listener that fails. + /// What [`upstream`](Self::upstream), [`service`](Self::service) and + /// [`serve_options`](Self::serve_options) reject, an invalid listen + /// address, or a listener that fails. pub async fn serve(&self) -> anyhow::Result<()> { let service = self.service(self.upstream()?)?; + let options = self.serve_options()?; 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); - serve(listener, service).await?; + tracing::info!( + tls = self.config.listen.tls.is_some(), + "{} listening on {}", + self.config.service.name, + addr + ); + serve_with(listener, service, options).await?; Ok(()) } } diff --git a/src/serve.rs b/src/serve.rs new file mode 100644 index 0000000..1052e0d --- /dev/null +++ b/src/serve.rs @@ -0,0 +1,222 @@ +//! Running a [`ProxyService`] on a TCP listener: HTTP/1.1 and HTTP/2 on one +//! port, optionally behind TLS, with an optional cap on open connections. + +use std::sync::Arc; +use std::time::Duration; + +use hyper_util::rt::{TokioExecutor, TokioIo}; +use hyper_util::server::conn::auto::Builder; +use hyper_util::service::TowerToHyperService; +use tokio::io::{AsyncRead, AsyncWrite}; +use tokio::net::TcpListener; +use tokio::sync::Semaphore; +use tonic::transport::server::Connected; + +use crate::service::{ConnectionInfo, ProxyService}; +use crate::upstream::Upstream; + +/// How long a client has to finish its TLS handshake before the connection is +/// dropped, so a stalled client holds neither a task nor a connection slot. +const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10); + +/// How [`serve_with`] runs its listener. +/// +/// # Examples +/// +/// ``` +/// use structured_proxy::ServeOptions; +/// +/// let options = ServeOptions::new().max_connections(10_000); +/// # let _ = options; +/// ``` +#[derive(Clone, Debug, Default)] +pub struct ServeOptions { + max_connections: Option, + tls: Option>, +} + +impl ServeOptions { + /// Cleartext, with no limit on connections. + pub fn new() -> Self { + Self::default() + } + + /// Serve at most `max` connections at once: past it, the next connection + /// is accepted when one closes, and waits in the listen backlog until then. + /// + /// # Panics + /// + /// `max` is zero, which would never accept a connection. + #[must_use] + pub fn max_connections(mut self, max: usize) -> Self { + assert!(max > 0, "max_connections must be at least 1"); + self.max_connections = Some(max.min(Semaphore::MAX_PERMITS)); + self + } + + /// Terminate TLS with `config`. Set its `alpn_protocols` to `h2` and + /// `http/1.1` for gRPC clients to get HTTP/2; an empty list is filled in + /// with those two. A client certificate the config verifies reaches a + /// tonic upstream as `Request::peer_certs`. + #[must_use] + pub fn tls(mut self, mut config: rustls::ServerConfig) -> Self { + if config.alpn_protocols.is_empty() { + config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()]; + } + self.tls = Some(Arc::new(config)); + self + } +} + +/// 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. [`serve_with`] adds TLS and a connection limit. +/// +/// # Errors +/// +/// None so far: a failed accept is logged and retried, as a full file +/// descriptor table recovers when connections close. +/// +/// # 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: TcpListener, + service: ProxyService, +) -> std::io::Result<()> { + serve_with(listener, service, ServeOptions::new()).await +} + +/// [`serve`] with `options`: TLS termination (mTLS when the config verifies +/// client certificates) and a cap on open connections. Each connection's +/// service gets its peer, and behind TLS its client certificates, through +/// [`ProxyService::for_connection`]. The TLS handshake runs in the +/// connection's own task, so a slow client does not hold up the others. +/// +/// # Errors +/// +/// See [`serve`]. +/// +/// # Examples +/// +/// ```no_run +/// use structured_proxy::{ProxyServer, ServeOptions}; +/// +/// # async fn run(tls: rustls::ServerConfig) -> anyhow::Result<()> { +/// let service = 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 options = ServeOptions::new().tls(tls).max_connections(10_000); +/// structured_proxy::serve_with(listener, service, options).await?; +/// # Ok(()) +/// # } +/// ``` +pub async fn serve_with( + listener: TcpListener, + service: ProxyService, + options: ServeOptions, +) -> std::io::Result<()> { + let slots = options + .max_connections + .map(|max| Arc::new(Semaphore::new(max))); + let acceptor = options.tls.map(tokio_rustls::TlsAcceptor::from); + loop { + // The slot is taken before the accept, so a full server leaves new + // connections in the kernel's backlog instead of accepting and + // dropping them. + let slot = match &slots { + Some(slots) => Some( + slots + .clone() + .acquire_owned() + .await + .expect("the connection semaphore is never closed"), + ), + None => None, + }; + let tcp = match listener.accept().await { + Ok((tcp, _)) => tcp, + Err(error) => { + accept_failed(error).await; + continue; + } + }; + let service = service.clone(); + let acceptor = acceptor.clone(); + tokio::spawn(async move { + // Held for the connection's life. + let _slot = slot; + // Small gRPC frames and REST answers are latency-bound. + if let Err(error) = tcp.set_nodelay(true) { + tracing::debug!(%error, "cannot set TCP_NODELAY"); + } + match acceptor { + None => { + let service = service.for_connection(tcp.connect_info()); + serve_connection(tcp, service).await; + } + Some(acceptor) => { + let stream = + match tokio::time::timeout(TLS_HANDSHAKE_TIMEOUT, acceptor.accept(tcp)) + .await + { + Ok(Ok(stream)) => stream, + Ok(Err(error)) => { + tracing::debug!(%error, "TLS handshake failed"); + return; + } + Err(_) => { + tracing::debug!("TLS handshake timed out"); + return; + } + }; + let service = + service.for_connection(ConnectionInfo::tls(stream.connect_info())); + serve_connection(stream, service).await; + } + } + }); + } +} + +/// HTTP/1.1 or HTTP/2, whichever the client speaks, on one connection. +async fn serve_connection(io: I, service: ProxyService) +where + U: Upstream, + I: AsyncRead + AsyncWrite + Unpin + Send + 'static, +{ + let served = Builder::new(TokioExecutor::new()) + .serve_connection_with_upgrades(TokioIo::new(io), TowerToHyperService::new(service)) + .await; + if let Err(error) = served { + tracing::debug!(%error, "connection ended"); + } +} + +/// A failed accept: a connection the client gave up on is skipped; anything +/// else (a full file descriptor table) is logged and retried a second later, +/// the pause that lets open connections close. +async fn accept_failed(error: std::io::Error) { + use std::io::ErrorKind; + if matches!( + error.kind(), + ErrorKind::ConnectionRefused | ErrorKind::ConnectionAborted | ErrorKind::ConnectionReset + ) { + return; + } + tracing::error!(%error, "accepting a connection failed; retrying in 1s"); + tokio::time::sleep(Duration::from_secs(1)).await; +} + +#[cfg(test)] +mod tests; diff --git a/src/serve/tests.rs b/src/serve/tests.rs new file mode 100644 index 0000000..7ffd278 --- /dev/null +++ b/src/serve/tests.rs @@ -0,0 +1,121 @@ +use super::*; + +use std::net::SocketAddr; +use std::time::Duration; + +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpStream; + +/// Serve the proxy's health routes with `options` on a local port. +async fn listen(options: ServeOptions) -> SocketAddr { + let service = crate::ProxyServer::from_yaml_str("service:\n name: demo\n") + .unwrap() + .service(tonic::service::Routes::default()) + .unwrap(); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(serve_with(listener, service, options)); + addr +} + +/// Send a keep-alive `GET /health/live` on `stream`. +async fn request(stream: &mut TcpStream) { + stream + .write_all(b"GET /health/live HTTP/1.1\r\nHost: localhost\r\n\r\n") + .await + .unwrap(); +} + +/// Read until the end of a response head; returns its status line. +async fn status_line(stream: &mut TcpStream) -> String { + let mut head = Vec::new(); + let mut byte = [0; 1]; + while !head.ends_with(b"\r\n\r\n") { + let read = stream.read(&mut byte).await.unwrap(); + assert_eq!(read, 1, "the connection closed before a response"); + head.push(byte[0]); + } + let head = String::from_utf8(head).unwrap(); + head.lines().next().unwrap().to_owned() +} + +#[tokio::test] +async fn a_connection_past_the_limit_is_served_once_another_closes() { + let addr = listen(ServeOptions::new().max_connections(1)).await; + + let mut first = TcpStream::connect(addr).await.unwrap(); + request(&mut first).await; + assert_eq!(status_line(&mut first).await, "HTTP/1.1 200 OK"); + + // The kernel completes the second connection, but the proxy does not + // accept it while the first stays open. + let mut second = TcpStream::connect(addr).await.unwrap(); + request(&mut second).await; + let waiting = tokio::time::timeout(Duration::from_millis(300), status_line(&mut second)).await; + assert!(waiting.is_err(), "served past the connection limit"); + + drop(first); + let served = tokio::time::timeout(Duration::from_secs(5), status_line(&mut second)) + .await + .expect("the freed slot serves the waiting connection"); + assert_eq!(served, "HTTP/1.1 200 OK"); +} + +#[tokio::test] +async fn without_a_limit_connections_are_served_together() { + let addr = listen(ServeOptions::new()).await; + let mut first = TcpStream::connect(addr).await.unwrap(); + let mut second = TcpStream::connect(addr).await.unwrap(); + request(&mut first).await; + request(&mut second).await; + assert_eq!(status_line(&mut second).await, "HTTP/1.1 200 OK"); + assert_eq!(status_line(&mut first).await, "HTTP/1.1 200 OK"); +} + +#[test] +#[should_panic(expected = "max_connections must be at least 1")] +fn a_limit_of_zero_is_refused() { + let _ = ServeOptions::new().max_connections(0); +} + +#[test] +fn a_zero_limit_in_the_config_is_an_error() { + let server = crate::ProxyServer::from_yaml_str("listen:\n max_connections: 0\n").unwrap(); + let err = server.serve_options().unwrap_err(); + assert!(err.to_string().contains("listen.max_connections"), "{err}"); +} + +/// A process crypto provider, for builds without a crypto backend feature. +fn install_provider() { + if rustls::crypto::CryptoProvider::get_default().is_none() { + rustls::crypto::CryptoProvider::install_default(rustls_rustcrypto::provider()) + .expect("each test runs in its own process"); + } +} + +#[test] +fn tls_is_set_up_from_the_config_files() { + install_provider(); + let dir = concat!(env!("CARGO_MANIFEST_DIR"), "/src/tls/testdata"); + let yaml = format!( + "listen:\n tls:\n cert_file: {dir}/ecdsa.pem\n key_file: {dir}/ecdsa.key.pem\n client_ca_file: {dir}/client-ca.pem\n" + ); + let options = crate::ProxyServer::from_yaml_str(&yaml) + .unwrap() + .serve_options() + .unwrap(); + let tls = options.tls.expect("listen.tls configures TLS"); + // gRPC clients negotiate HTTP/2 on the same port as HTTP/1.1 ones. + assert_eq!(tls.alpn_protocols, [b"h2".to_vec(), b"http/1.1".to_vec()]); +} + +#[test] +fn a_missing_tls_file_names_itself() { + install_provider(); + let yaml = "listen:\n tls:\n cert_file: /nonexistent/tls.crt\n key_file: /nonexistent/tls.key\n"; + let err = crate::ProxyServer::from_yaml_str(yaml) + .unwrap() + .serve_options() + .unwrap_err(); + assert!(err.to_string().contains("/nonexistent/tls.crt"), "{err}"); +} diff --git a/src/service.rs b/src/service.rs index 4b681f8..1f28bec 100644 --- a/src/service.rs +++ b/src/service.rs @@ -11,7 +11,6 @@ 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; @@ -47,11 +46,13 @@ use crate::upstream::{ /// only the guards whose scope names them (`grpc`, `fallback`); a guard's /// rejection of a gRPC call is a gRPC status. /// -/// 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). +/// Serve it with [`serve`](crate::serve) or +/// [`serve_with`](crate::serve_with) (TLS, a connection limit), or hand it to +/// any server that takes a tower service of `http` types: 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 /// @@ -286,7 +287,8 @@ impl ProxyService { /// /// 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. + /// reports for the stream. [`serve_with`](crate::serve_with) does this + /// itself. /// /// # Examples /// @@ -470,53 +472,5 @@ impl Future for ResponseFuture { } } -/// 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/tls.rs b/src/tls.rs index a44c228..d737c70 100644 --- a/src/tls.rs +++ b/src/tls.rs @@ -1,5 +1,5 @@ -//! The rustls client configuration for the proxy's own outbound HTTPS calls -//! (JWKS fetches, the rate-limit service). +//! The rustls configurations: the client of the proxy's own outbound HTTPS +//! calls (JWKS fetches, the rate-limit service) and the server of its listener. use std::sync::Arc; @@ -36,7 +36,64 @@ pub(crate) fn client_config_with( .with_no_client_auth()) } -/// The crypto provider for outbound TLS: the one the process `installed`, else +/// The rustls server config `config` describes: its certificate and key, and +/// with `client_ca_file` a verifier for client certificates. ALPN offers `h2` +/// and `http/1.1`, so gRPC clients get HTTP/2 on the same port as REST ones. +/// +/// # Errors +/// +/// No crypto provider (see [`select_provider`]), a file that cannot be read or +/// holds no certificate / key, or a certificate rustls refuses. +pub(crate) fn server_config( + config: &crate::config::ListenTlsConfig, +) -> Result { + use rustls::pki_types::pem::PemObject; + use rustls::pki_types::{CertificateDer, PrivateKeyDer}; + + let provider = select_provider(CryptoProvider::get_default(), builtin_provider)?; + let certs = |path: &std::path::Path| { + let certs = CertificateDer::pem_file_iter(path) + .and_then(|certs| certs.collect::, _>>()) + .map_err(|e| format!("{}: {e}", path.display()))?; + if certs.is_empty() { + return Err(format!("{}: no PEM certificate", path.display())); + } + Ok(certs) + }; + let chain = certs(&config.cert_file)?; + let key = PrivateKeyDer::from_pem_file(&config.key_file) + .map_err(|e| format!("{}: {e}", config.key_file.display()))?; + let builder = rustls::ServerConfig::builder_with_provider(provider.clone()) + .with_safe_default_protocol_versions() + .map_err(|e| format!("the rustls crypto provider cannot do TLS 1.2 or 1.3: {e}"))?; + let builder = match &config.client_ca_file { + None => builder.with_no_client_auth(), + Some(path) => { + let mut roots = rustls::RootCertStore::empty(); + for cert in certs(path)? { + roots + .add(cert) + .map_err(|e| format!("{}: {e}", path.display()))?; + } + let verifier = + rustls::server::WebPkiClientVerifier::builder_with_provider(roots.into(), provider); + let verifier = match config.client_auth { + crate::config::ClientAuth::Required => verifier, + crate::config::ClientAuth::Optional => verifier.allow_unauthenticated(), + } + .build() + .map_err(|e| format!("{}: {e}", path.display()))?; + builder.with_client_cert_verifier(verifier) + } + }; + let mut server = builder + .with_single_cert(chain, key) + .map_err(|e| format!("{}: {e}", config.cert_file.display()))?; + server.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()]; + Ok(server) +} + +/// The crypto provider for the proxy's TLS: the one the process `installed`, else /// the `builtin` one this crate's crypto backend brings. /// /// An installed provider is the application's explicit choice and wins, the @@ -56,7 +113,7 @@ fn select_provider( return Ok(Arc::clone(installed)); } builtin().map(Arc::new).ok_or_else(|| { - "the outbound HTTP client needs a rustls crypto provider: enable the \ + "TLS (the listener, the outbound HTTP client) needs a rustls crypto provider: enable the \ `rust_crypto` or `aws_lc_rs` feature, or install one with \ `rustls::crypto::CryptoProvider::install_default` before building the proxy" .to_string() diff --git a/src/tls/testdata/client-ca.pem b/src/tls/testdata/client-ca.pem new file mode 100644 index 0000000..d303887 --- /dev/null +++ b/src/tls/testdata/client-ca.pem @@ -0,0 +1,12 @@ +-----BEGIN CERTIFICATE----- +MIIBujCCAWGgAwIBAgIUE0lQ1GFJJJyC90ofG7a6U1GaCK8wCgYIKoZIzj0EAwIw +KjEoMCYGA1UEAwwfc3RydWN0dXJlZC1wcm94eSB0ZXN0IGNsaWVudCBDQTAgFw0y +NjA5MjgwODA1MDlaGA8yMTI2MDkwNDA4MDUwOVowKjEoMCYGA1UEAwwfc3RydWN0 +dXJlZC1wcm94eSB0ZXN0IGNsaWVudCBDQTBZMBMGByqGSM49AgEGCCqGSM49AwEH +A0IABGjFplBZ9BTU8G9MEWrF3pFqzoTm5TWfxjvKPO7znbNREmzcxhWjSSNPy6B/ +6E6uNasumN4qUNhT1CSbzbkygrujYzBhMB0GA1UdDgQWBBQ4JeejxcGIULEsa0Ww +fV/ynjCm5zAfBgNVHSMEGDAWgBQ4JeejxcGIULEsa0WwfV/ynjCm5zAPBgNVHRMB +Af8EBTADAQH/MA4GA1UdDwEB/wQEAwIBBjAKBggqhkjOPQQDAgNHADBEAiAGSZ3Y +cragoxwNQAqvwdwYNL1n1mEi8KdhdPJf2UNo+gIgcuKsDxlILQAfJcY97K0rsae1 +BUVPEFgJ9QJ3Mtertak= +-----END CERTIFICATE----- diff --git a/src/tls/testdata/client.key.pem b/src/tls/testdata/client.key.pem new file mode 100644 index 0000000..183c8a8 --- /dev/null +++ b/src/tls/testdata/client.key.pem @@ -0,0 +1,5 @@ +-----BEGIN PRIVATE KEY----- +MIGHAgEAMBMGByqGSM49AgEGCCqGSM49AwEHBG0wawIBAQQgEiiQ8fAGTFkW9jx3 +ilXEwFCGNzk7ZmKQeqoKN0gK17+hRANCAAT4zHfvPKkRGCYPncSNHz0gxo01e7p+ +Kl9hkAHfwg8M27O8OCgPA9yuAW+Zv7rsYoYvE7QtTdmTNKF513smFBUD +-----END PRIVATE KEY----- diff --git a/src/tls/testdata/client.pem b/src/tls/testdata/client.pem new file mode 100644 index 0000000..89b0401 --- /dev/null +++ b/src/tls/testdata/client.pem @@ -0,0 +1,12 @@ +-----BEGIN CERTIFICATE----- +MIIBszCCAVmgAwIBAgIUYeMC8FFE9IKIEeXRhbzuUMsws8QwCgYIKoZIzj0EAwIw +KjEoMCYGA1UEAwwfc3RydWN0dXJlZC1wcm94eSB0ZXN0IGNsaWVudCBDQTAgFw0y +NjA5MjgwODA1MDlaGA8yMTI2MDkwNDA4MDUwOVowFjEUMBIGA1UEAwwLdGVzdCBj +bGllbnQwWTATBgcqhkjOPQIBBggqhkjOPQMBBwNCAAT4zHfvPKkRGCYPncSNHz0g +xo01e7p+Kl9hkAHfwg8M27O8OCgPA9yuAW+Zv7rsYoYvE7QtTdmTNKF513smFBUD +o28wbTAJBgNVHRMEAjAAMAsGA1UdDwQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcD +AjAdBgNVHQ4EFgQUXw56ugGH2jG3CyK3VsdUDy9TIF4wHwYDVR0jBBgwFoAUOCXn +o8XBiFCxLGtFsH1f8p4wpucwCgYIKoZIzj0EAwIDSAAwRQIhAJOQxfuz64mSXykR +31qhC0XY+wLoO0SDLvUCXzIxlRwvAiAPRUao4Iz6eTSVVzvWqwWoefNJ0RThnnNL +TG75dmzsRA== +-----END CERTIFICATE----- diff --git a/src/tls/testdata/generate.sh b/src/tls/testdata/generate.sh index bb6b727..62ab9e0 100644 --- a/src/tls/testdata/generate.sh +++ b/src/tls/testdata/generate.sh @@ -1,6 +1,7 @@ #!/bin/sh # Test PKI for src/tls/tests.rs: a CA, an ECDSA and an RSA leaf for -# `localhost`, and an unrelated CA. Valid for 100 years. +# `localhost`, and an unrelated CA; for mTLS, a client CA and a client leaf +# (tests/serve.rs). Valid for 100 years. set -eu out="$1" cd "$out" @@ -42,4 +43,18 @@ leaf ecdsa openssl genpkey -algorithm RSA -pkeyopt rsa_keygen_bits:2048 -out rsa.key.pem leaf rsa -rm -f ca.srl ca.key.pem +# A client certificate, which a verifier of client certificates accepts only +# with the clientAuth extended key usage (RFC 5280 4.2.1.12). +ca client-ca "structured-proxy test client CA" +openssl ecparam -name prime256v1 -genkey -noout -out client.key.tmp +openssl pkcs8 -topk8 -nocrypt -in client.key.tmp -out client.key.pem +rm client.key.tmp +printf '%s\n' "basicConstraints=CA:FALSE +keyUsage=digitalSignature +extendedKeyUsage=clientAuth" > client.ext +openssl req -new -key client.key.pem -subj "/CN=test client" -out client.csr +openssl x509 -req -in client.csr -CA client-ca.pem -CAkey client-ca.key.pem -CAcreateserial \ + -days "$days" -extfile client.ext -out client.pem +rm client.csr client.ext + +rm -f ca.srl ca.key.pem client-ca.srl client-ca.key.pem diff --git a/tests/embedded.rs b/tests/embedded.rs index 020b108..36b823e 100644 --- a/tests/embedded.rs +++ b/tests/embedded.rs @@ -24,6 +24,8 @@ fn embedded_config_is_constructible() { }], listen: ListenConfig { http: "0.0.0.0:8080".into(), + max_connections: None, + tls: None, }, service: ServiceConfig { name: "embedded-test".into(), diff --git a/tests/tls.rs b/tests/tls.rs index 9b2fcd8..a4b2864 100644 --- a/tests/tls.rs +++ b/tests/tls.rs @@ -1,7 +1,8 @@ -//! 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. +//! The proxy behind TLS, an embedder's own (a rustls acceptor and hyper's +//! HTTP/1.1 + HTTP/2 connection, with the proxy as the service) or its +//! built-in listener (`listen.tls`, mTLS with a client CA). 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; @@ -236,14 +237,82 @@ async fn listen() -> SocketAddr { addr } +/// The directory of the test PKI. +const TESTDATA: &str = concat!(env!("CARGO_MANIFEST_DIR"), "/src/tls/testdata"); + +/// Serve the proxy in front of `Who` with the proxy's own TLS listener, +/// configured by the `listen:` of `listen_yaml` (its `tls:` and its limits). +async fn listen_builtin(listen_yaml: &str) -> SocketAddr { + // Builds without a crypto backend feature take the process's provider. + if rustls::crypto::CryptoProvider::get_default().is_none() { + rustls::crypto::CryptoProvider::install_default(rustls_rustcrypto::provider()) + .expect("each test runs in its own process"); + } + let pool = pool(); + let server = ProxyServer::from_yaml_str(&format!("listen:\n{listen_yaml}")) + .unwrap() + .with_descriptors(pool.clone()); + let proxy = server + .service(tonic::service::Routes::new(Who { pool })) + .unwrap(); + let options = server.serve_options().unwrap(); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(structured_proxy::serve_with(listener, proxy, options)); + addr +} + +/// `listen.tls` for the test server certificate, verifying client +/// certificates against the client CA with `client_auth` when given. +fn tls_yaml(client_auth: Option<&str>) -> String { + let mut yaml = format!( + " tls:\n cert_file: {TESTDATA}/ecdsa.pem\n key_file: {TESTDATA}/ecdsa.key.pem\n" + ); + if let Some(client_auth) = client_auth { + yaml.push_str(&format!( + " client_ca_file: {TESTDATA}/client-ca.pem\n client_auth: {client_auth}\n" + )); + } + yaml +} + // --- clients -------------------------------------------------------------------- +/// The certificate a test client presents. +#[derive(Clone, Copy)] +enum Identity { + /// None. + Anonymous, + /// The server's own leaf, which a verifier that checks usage refuses. + ServerLeaf, + /// A client leaf of the client CA. + Client, +} + +impl Identity { + fn cert(self) -> Option<(Vec>, PrivateKeyDer<'static>)> { + let client = || { + let chain = CertificateDer::pem_file_iter(format!("{TESTDATA}/client.pem")) + .unwrap() + .collect::>() + .unwrap(); + let key = PrivateKeyDer::from_pem_file(format!("{TESTDATA}/client.key.pem")).unwrap(); + (chain, key) + }; + match self { + Self::Anonymous => None, + Self::ServerLeaf => Some((chain(), key())), + Self::Client => Some(client()), + } + } +} + /// A TLS connection to `addr` trusting the test CA, offering `alpn`, and -/// presenting the test certificate when `with_cert`. +/// presenting `identity`'s certificate. async fn connect( addr: SocketAddr, alpn: &[u8], - with_cert: bool, + identity: Identity, ) -> tokio_rustls::client::TlsStream { let mut roots = rustls::RootCertStore::empty(); for cert in CertificateDer::pem_slice_iter(CA.as_bytes()) { @@ -253,10 +322,9 @@ async fn connect( .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() + let mut config = match identity.cert() { + Some((chain, key)) => builder.with_client_auth_cert(chain, key).unwrap(), + None => builder.with_no_client_auth(), }; config.alpn_protocols = vec![alpn.to_vec()]; let tcp = tokio::net::TcpStream::connect(addr).await.unwrap(); @@ -268,8 +336,8 @@ async fn connect( /// `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; +async fn rest_me(addr: SocketAddr, identity: Identity) -> (u16, Value, SocketAddr) { + let mut tls = connect(addr, b"http/1.1", identity).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 @@ -290,12 +358,12 @@ async fn rest_me(addr: SocketAddr, with_cert: bool) -> (u16, Value, SocketAddr) /// 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) { +async fn grpc_me(addr: SocketAddr, identity: Identity) -> (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; + let tls = connect(addr, b"h2", identity).await; sender .send(tls.get_ref().0.local_addr().unwrap()) .await @@ -333,7 +401,7 @@ fn leaf_der() -> Vec { #[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; + let (status, seen, client) = rest_me(addr, Identity::Anonymous).await; assert_eq!(status, 200, "{seen}"); assert_eq!(seen["peer"], client.to_string()); // No certificate presented, none reported. @@ -343,7 +411,7 @@ async fn rest_over_tls_reaches_the_upstream_with_the_client_address() { #[tokio::test] async fn native_grpc_shares_the_tls_port() { let addr = listen().await; - let (seen, client) = grpc_me(addr, false).await; + let (seen, client) = grpc_me(addr, Identity::Anonymous).await; let peer = match seen.get_field_by_name("peer").as_deref() { Some(PbValue::String(peer)) => peer.clone(), _ => String::new(), @@ -356,7 +424,7 @@ 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; + let (status, seen, _) = rest_me(addr, Identity::ServerLeaf).await; assert_eq!(status, 200, "{seen}"); let cert = base64::engine::general_purpose::STANDARD .decode(seen["cert"].as_str().unwrap()) @@ -367,10 +435,80 @@ async fn a_client_certificate_reaches_a_transcoded_call() { #[tokio::test] async fn a_client_certificate_reaches_a_native_call() { let addr = listen().await; - let (seen, _) = grpc_me(addr, true).await; + let (seen, _) = grpc_me(addr, Identity::ServerLeaf).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()); } + +// --- the proxy's own TLS listener ----------------------------------------------- + +fn seen_cert(seen: &DynamicMessage) -> Vec { + match seen.get_field_by_name("cert").as_deref() { + Some(PbValue::Bytes(cert)) => cert.to_vec(), + _ => Vec::new(), + } +} + +fn client_leaf_der() -> Vec { + Identity::Client.cert().unwrap().0[0].as_ref().to_vec() +} + +#[tokio::test] +async fn the_builtin_tls_listener_serves_rest_and_native_grpc() { + let addr = listen_builtin(&tls_yaml(None)).await; + let (status, seen, client) = rest_me(addr, Identity::Anonymous).await; + assert_eq!(status, 200, "{seen}"); + assert_eq!(seen["peer"], client.to_string()); + let (seen, client) = grpc_me(addr, Identity::Anonymous).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 builtin_mtls_passes_a_verified_client_certificate_to_the_upstream() { + let addr = listen_builtin(&tls_yaml(Some("required"))).await; + let (status, seen, _) = rest_me(addr, Identity::Client).await; + assert_eq!(status, 200, "{seen}"); + let cert = base64::engine::general_purpose::STANDARD + .decode(seen["cert"].as_str().unwrap()) + .unwrap(); + assert_eq!(cert, client_leaf_der()); + let (seen, _) = grpc_me(addr, Identity::Client).await; + assert_eq!(seen_cert(&seen), client_leaf_der()); +} + +/// Whether a request on a TLS connection presenting `identity` gets any +/// answer: a refused client certificate ends the connection instead (in TLS +/// 1.3 after the client's side of the handshake completed). +async fn answered(addr: SocketAddr, identity: Identity) -> bool { + let mut tls = connect(addr, b"http/1.1", identity).await; + let written = tls + .write_all(b"GET /v1/me HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n") + .await; + let mut response = Vec::new(); + let read = tls.read_to_end(&mut response).await; + written.is_ok() && read.is_ok() && response.starts_with(b"HTTP/1.1 200") +} + +#[tokio::test] +async fn builtin_mtls_required_refuses_a_client_without_a_valid_certificate() { + let addr = listen_builtin(&tls_yaml(Some("required"))).await; + assert!(!answered(addr, Identity::Anonymous).await); + // Signed by a CA the listener does not trust for clients. + assert!(!answered(addr, Identity::ServerLeaf).await); + assert!(answered(addr, Identity::Client).await); +} + +#[tokio::test] +async fn builtin_mtls_optional_serves_a_client_without_a_certificate() { + let addr = listen_builtin(&tls_yaml(Some("optional"))).await; + let (status, seen, _) = rest_me(addr, Identity::Anonymous).await; + assert_eq!(status, 200, "{seen}"); + assert_eq!(seen["cert"], ""); +} From 59579d5579fa4c57b9a5d7a2cfe49c17c1e1c71d Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 11:13:35 +0300 Subject: [PATCH 5/9] docs(readme): make the README current and easier to read Group the feature list by purpose, rewrite the densest paragraphs (rate limiting, error details, header forwarding, raw bodies, pass-through) as short sentences and lists, replace the transcoding-only architecture diagram with the request flow of the edge, and fix the stale crate version in the injected-verifier example. --- README.md | 570 ++++++++++++++++++++++++++---------------------------- 1 file changed, 278 insertions(+), 292 deletions(-) diff --git a/README.md b/README.md index 4916971..4ae0c56 100644 --- a/README.md +++ b/README.md @@ -6,42 +6,66 @@ [![downloads](https://img.shields.io/crates/d/structured-proxy.svg)](https://crates.io/crates/structured-proxy) [![license](https://img.shields.io/crates/l/structured-proxy.svg)](https://github.com/structured-world/structured-proxy/blob/main/LICENSE) -Universal, config-driven gRPC→REST transcoding proxy. One binary, different YAML configs, different products. - -Works with **any** gRPC service via proto descriptor files. No code generation, no custom handlers, just configuration. +A gRPC→REST transcoding proxy and edge for any gRPC service. Point it at your +proto descriptors and it serves REST/JSON next to native gRPC on one port, +with auth, rate limits and the rest of an edge configured in YAML. Run it as a +standalone binary, or embed it in your own Rust service in front of your own +tonic services. ## Features -- **Dynamic REST routes** from proto descriptors using `google.api.http` annotations -- **Full request mapping**: path params, query parameters (typed + repeated + nested), and a JSON or form `body` (`*` / named field / none), decoded straight into the request message (see [Request mapping](#request-mapping)) -- **`response_body`** to return a single response subfield, and **`additional_bindings`** for multiple routes per RPC -- **`custom` rules**: any HTTP method (`HEAD`, `OPTIONS`, extension methods), or `kind: "*"` for every method -- **Upstream-controlled HTTP answers**: response metadata becomes response headers, `x-http-code` sets the status, and `google.api.HttpBody` carries a raw body and content type in either direction, so OAuth 2.0 / OIDC endpoints, redirects and file downloads work as gRPC (see [Upstream controls](#upstream-controls)) -- **Auto-generated OpenAPI** documentation from proto messages, served at `/openapi.json` -- **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 the built-in TLS listener (mTLS with a client CA, a cap on open connections) or your own TLS, 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/*`) -- **Scoped guards**: maintenance, rate limits, a concurrency limit, JWT, ext_authz and the auth decider each cover the traffic you name (transcoded, the proxy's own endpoints, native gRPC, the fallback), narrowed by path and method; a rejection answers in the request's protocol, a `google.rpc.Status` JSON body for REST and a gRPC status for gRPC (see [Guards and scopes](#guards-and-scopes)) -- **Maintenance mode** returning 503 with a configurable exempt-path list -- **Concurrency limit**: requests past `max_in_flight` are shed at once with 503 instead of queueing; a stream holds its slot until its body ends -- **Health endpoints** `/health/live`, `/health/ready` (upstream gRPC health probe), `/health/startup` -- **Prometheus metrics** at `/metrics` -- **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 -- **Forward-auth**: a verification endpoint (`/auth/verify`) for a fronting proxy (nginx `auth_request`, Traefik `forwardAuth`) to delegate auth, returning the verified identity as headers -- **External AuthZ**: gate proxied requests through an Envoy ext_authz gRPC server (`envoy.service.auth.v3.Authorization/Check`), interoperating with OPA and any ext_authz server, with fail-open/closed control -- **Zero code changes** between services: same binary, different config +**Transcoding** + +- REST routes from the `google.api.http` annotations in your protos: path, + query and JSON or form body mapped onto the request message + ([Request mapping](#request-mapping)), `response_body`, + `additional_bindings`, and `custom` rules for any HTTP method +- Server streaming as NDJSON, or Server-Sent Events when the client asks for + `text/event-stream` +- Errors as `google.rpc.Status` JSON with the standard HTTP status mapping and + typed error details ([Error responses](#error-responses)) +- The upstream decides status, headers and raw bodies, so OAuth 2.0 / OIDC + endpoints, redirects and file downloads can be written as gRPC + ([Upstream controls](#upstream-controls)) +- OpenAPI generated from the protos, served at `/openapi.json` +- Header forwarding, W3C trace context and client deadlines carried across + REST↔gRPC; path aliases + +**Edge** + +- REST and native gRPC (and gRPC-Web) on one port, HTTP/1.1 and HTTP/2 +- Built-in TLS and mTLS, and a cap on open connections; or your own TLS, with + the client's address and certificate reaching the upstream +- Guards you scope to the traffic they cover (transcoded calls, the proxy's + own endpoints, native gRPC, the fallback), rejecting in the protocol of the + request ([Guards and scopes](#guards-and-scopes)): + - rate limits keyed by IP, header or JWT claim, optionally shared across + instances through Redis ([Rate limiting](#rate-limiting)) + - a limit on requests in flight + - JWT auth with per-route policies, or your own token verifier + - Envoy ext_authz (OPA and any other ext_authz server) and an in-process + auth decider + - maintenance mode +- CORS, applied to gRPC-Web calls too +- Health probes (`/health/live`, `/health/ready` against the upstream's gRPC + health, `/health/startup`) and Prometheus metrics at `/metrics` +- Forward-auth endpoint for nginx `auth_request` / Traefik `forwardAuth`, and + OIDC discovery with a JWKS endpoint + +**Embedding** + +- Your own tonic services as the upstream, called in process + ([Library Usage](#library-usage)) +- Requests the proxy does not serve go to your own fallback service +- Hooks for auth decisions, token verification, OIDC and extra routes, with no + HTTP framework in your code ## Non-goals -- **Session / BFF management** (cookie-based login, server-side token storage, refresh flows) and **stateful OIDC** (`authorize` / `token` with auth codes / PKCE state). The **default build** is a stateless transcoding data plane with stateless auth primitives; session lifecycle is a separate, stateful concern. Put a dedicated BFF (e.g. `oauth2-proxy`, Pomerium) in front, or drive auth through the stateless forward-auth / external-authz hooks below. (A stateful surface behind an opt-in, default-off `bff` Cargo feature is planned; it does not affect the default data-plane build.) +- **Sessions and stateful OIDC** (cookie login, server-side tokens, refresh, + `authorize` / `token` with PKCE state). The proxy is a stateless data plane: + put a BFF such as `oauth2-proxy` or Pomerium in front, or decide auth + through the forward-auth and ext_authz hooks. ## Quick Start @@ -58,10 +82,10 @@ Prebuilt static Linux binaries and deb/rpm packages are attached to each GitHub release. To embed the proxy in your own service instead, add the library with `cargo add structured-proxy` (see [Library Usage](#library-usage)). -The binary runs the proxy on a multi-thread async runtime. `runtime.worker_threads` -in the config file sets how many worker threads, and so CPU cores, it keeps busy; -unset, `TOKIO_WORKER_THREADS` or the available parallelism decides. The startup -log line (`RUST_LOG=info`) states the count and where it came from. +`runtime.worker_threads` sets how many worker threads, and so CPU cores, the +binary uses. Unset, `TOKIO_WORKER_THREADS` decides, else the number of CPUs +available to the process. The startup log (`RUST_LOG=info`) shows the count +and where it came from. ## Configuration @@ -279,8 +303,8 @@ protoc --descriptor_set_out=my-service.descriptor.bin --include_imports *.proto ## Request mapping -The gRPC request message is built from the three sources `google.api.http` -names, in one pass and without an intermediate JSON tree: +The gRPC request message is built from the path, the query string and the +body, as the route's `google.api.http` rule says: - **Precedence:** a path parameter wins over the body, and the body over the query string. A query parameter only fills a field the body did not send; @@ -304,67 +328,57 @@ answered with `INVALID_ARGUMENT` (400) before the upstream is called. ## Rate limiting -Shield is an embedded, config-driven limiter designed for a data plane: every -decision is made in-process by a GCRA shaper, so it adds no blocking latency to -the request path. GCRA (a token-bucket equivalent storing one timestamp per key) -lets legitimate bursts through up to a configured `burst` while throttling -sustained abuse to the `rate`, with no fixed-window boundary burst. - -**Keying and phases.** A rule keys on the client IP, a header value (API key), -or a validated JWT claim. The phase is derived from the key, not configured: an -IP/header rule needs no verified identity so it runs *before* auth (a fast, -purely local check that sheds anonymous floods before any signature verification, -and short-circuits so blocked clients never reach the auth layer); a `jwt_claim` -rule needs the verified principal so it runs *after* auth. A key falls back to -the client IP when its value is absent within its own phase, so a limit can't be -dodged by omitting a header. Note the fallback is phase-local: an anonymous -request under a `jwt_claim` rule keys by IP in the post-auth phase, but is *not* -shed pre-auth. For anonymous flood protection, add a separate pre-auth IP (or -header) rule covering the same paths; a path may match one rule per phase and -each is enforced independently (defense in depth). - -**Limit sources.** A key's `{rate, burst}` resolves in order: the JWT itself -(a `ratelimit_tier` claim naming a profile, or explicit `ratelimit_rpm` / -`ratelimit_burst`), then an external service (cached and refreshed in the -background, never blocking), then the rule's pinned profile, then the default. -JWT-based resolution only applies to `jwt_claim` rules, since only they run with -verified claims available; setting `jwt_limits` has no effect on an IP/header -rule, which runs pre-auth (use a `jwt_claim` key if you want the token's tier to -drive the limit). Tier-name indirection lets you retune the numbers in config -without re-issuing tokens or changing the service. - -**Response headers.** Every metered response carries the +Shield decides every request in the proxy's own process with GCRA, so a limit +never waits on the network. GCRA lets a client burst up to `burst` requests +and then holds it to `rate`, without the spikes a fixed window allows at its +edges. + +**Keys and when they run.** A rule keys on the client IP, a header (an API +key) or a verified JWT claim, and the key decides when the rule runs: + +- IP and header rules run before auth, so a flood is shed before any token + signature is checked. +- `jwt_claim` rules run after auth, on the verified claim. + +A request that lacks the key's value (no header, no token) is keyed by its IP, +so leaving the header out does not escape the limit. That fallback happens in +the rule's own phase: an anonymous request under a `jwt_claim` rule is limited +by IP only after auth. To shed anonymous floods early, add an IP rule for the +same paths; a path matches at most one rule in each phase, and both apply. + +**Which limit applies.** A key's rate and burst come from, in order: the token +(a `ratelimit_tier` claim naming a profile, or `ratelimit_rpm` / +`ratelimit_burst` claims), the external limit service (cached and refreshed in +the background), the rule's profile, the default profile. Limits from the +token apply to `jwt_claim` rules only, the only ones that see verified claims. +A tier name in the token lets you retune the numbers in config without +reissuing tokens. + +**Response headers.** Every limited response carries the [draft-ietf-httpapi-ratelimit-headers](https://datatracker.ietf.org/doc/draft-ietf-httpapi-ratelimit-headers/) -fields: `RateLimit-Limit` (the tier's per-window quota), `RateLimit-Remaining` -(requests still admissible now), and `RateLimit-Reset` (whole seconds until the -limiter drains toward full). A rejected request returns `429` with `Retry-After` -(whole seconds until a retry would conform). Clients should back off for -`Retry-After` seconds on a `429`, and may pace themselves using `RateLimit-*` on -allowed responses. Behind a browser, these are exposed via CORS. +fields: `RateLimit-Limit` (the quota per window), `RateLimit-Remaining` +(requests allowed right now) and `RateLimit-Reset` (seconds until the budget +refills). A rejected request gets `429` with `Retry-After`, the seconds to wait +before retrying. CORS exposes all of them to browser scripts. **Deployment modes.** -- **Local (default).** No shared store. Each instance enforces the limit - independently, so the fleet-wide effect is roughly `N × rate` for `N` - instances. Zero dependencies, lowest latency. Set per-instance limits with - that multiplier in mind. -- **Reconciled (`sync` + `redis` feature).** Instances asynchronously push their - deltas to a shared store and pull the aggregate on an interval, converging on - an approximate fleet-wide limit. The configured `rate` is then the *fleet* - budget. The request path still never blocks on the store; if the store is - unreachable, instances degrade to local limiting rather than failing requests. - -**Sizing the overshoot.** In reconciled mode the aggregate lags by up to one -`sync.interval_ms`. Within that lag each of the other instances can admit its -local `burst` plus a `rate` fraction of the window before the estimate catches -up (the fleet gate caps sustained fleet volume at `rate`, but `burst` is a -per-instance allowance the gate does not pre-reserve). With the interval in the -same time unit as the window, the worst-case fleet overshoot is about -`(N - 1) × (burst + rate × (interval / window))` requests. For example, -`burst = 100`, `rate = 1000/min`, `interval = 500 ms`, `N = 4` gives -`3 × (100 + 1000 × (0.5 / 60)) ≈ 325` extra requests. Smaller `burst` and shorter -intervals tighten the bound. The global view uses a sliding-window counter, so -there is no boundary burst on top of this lag. +- **Local (default).** Each instance enforces the limit on its own, so `N` + instances admit roughly `N × rate` together. Set per-instance limits with + that in mind. +- **Shared (`sync`, with the `redis` feature).** Instances push their counts + to Redis in the background and read the fleet's total on an interval, so + `rate` becomes the budget of the whole fleet, approximately. A request never + waits on Redis; while Redis is down, each instance limits locally. + +**How far the fleet can overshoot.** The shared total is up to one +`sync.interval_ms` old. Within that time every other instance can still admit +its own `burst` plus the share of `rate` that falls in the interval, so the +fleet can exceed its budget by about +`(N - 1) × (burst + rate × interval / window)` requests. With `burst = 100`, +`rate = 1000/min`, a 500 ms interval and 4 instances that is +`3 × (100 + 1000 × 0.5 / 60) ≈ 325` requests. A smaller `burst` and a shorter +interval shrink it. See the `shield:` block under [Configuration](#configuration) for the full schema. @@ -474,16 +488,15 @@ mapping (`INVALID_ARGUMENT` → 400, `NOT_FOUND` → 404, ...): request that cannot be mapped onto the RPC (`INVALID_ARGUMENT`, 400), an upstream that is not reachable (`UNAVAILABLE`, 503), a response that cannot be serialized (`INTERNAL`, 500). Their `details` is empty. -- A broken upstream error status is never passed on in part or reinterpreted: - a trailer that is not a `google.rpc.Status` or disagrees with `grpc-status` / - `grpc-message`, a type URL without a `/` or whose last segment is not a - protobuf full name, or a detail of a known type whose bytes do not decode or - whose value has no valid JSON form (a `Duration` beyond its range), turns the - whole error into +- A malformed upstream error is replaced as a whole, never passed on in part + or guessed at. The client gets `{"error": "INTERNAL", "code": 13, "message": "upstream returned a malformed error status", "details": []}` - (500, or the terminal frame of a started stream). The cause is logged by the - proxy and not sent to the client. With details switched off for a route the - trailer is not read, so this does not apply there. + (500, or the terminal frame of a started stream), and the proxy logs the + cause. Malformed means: a details trailer that is not a `google.rpc.Status` + or contradicts `grpc-status` / `grpc-message`, a type URL that is not + `…/`, or a detail of a known type that does not decode + or has no JSON form (a `Duration` out of range). Routes with details + switched off never read the trailer. **Opaque-detail extension.** ProtoJSON cannot represent an `Any` whose type is unknown to the writer, so a detail whose type is in neither descriptor set has @@ -522,51 +535,43 @@ of ProtoJSON or `google.rpc.Status`: - Only an unknown type goes there. A detail of a known type that fails to decode is a broken upstream status (see above), never an opaque entry. -**Errors in server-streaming responses.** An upstream that rejects the call -outright (a gRPC trailers-only response, with no response headers or messages) -gets the mapped HTTP status and the body above. Once the upstream has accepted -the call, the proxy answers `200` and starts the stream right away, so that -headers and SSE keep-alives are not held back waiting for the first message. -Any later failure, including one that arrives before the first message, is then -delivered as a terminal frame whose payload is exactly that body, after which -the stream ends and no further data follows: - -- **NDJSON**: the last line, framed by an extra - `"@type": "type.googleapis.com/google.rpc.Status"` next to the error body. A - data line is the ProtoJSON of a response message, which has a top-level - `@type` only when the RPC streams `google.protobuf.Any`, while `Struct`, - `Value` and `ListValue` messages can carry any key at all. For those RPCs no - in-band marker is collision-free: set `streaming.ndjson_envelope: true` (or - `ProxyServer::with_ndjson_envelope(true)`) and every line is wrapped instead, - `{"result": }` for data and `{"error": }` for the - terminal error, the grpc-gateway stream shape. The envelope changes data - lines too, so it is off by default. -- **SSE**: one event with type `stream-error` (listen with - `addEventListener("stream-error", ...)`), distinct from the `EventSource` - `onerror` that fires on transport failures. The event type is the framing, - so the event data is exactly the error body, without the NDJSON marker. - -The same applies to a message the proxy cannot serialize mid-stream: the -stream ends with an `INTERNAL` terminal frame. +**Errors in server-streaming responses.** If the upstream rejects the call +outright, the client gets the mapped HTTP status and the body above. Once the +upstream accepts the call, the proxy answers `200` and starts the stream at +once, without waiting for the first message. A failure after that, or a +message the proxy cannot serialize, ends the stream with one terminal frame +carrying that body: + +- **NDJSON**: the last line, the error body with + `"@type": "type.googleapis.com/google.rpc.Status"` added. An RPC that streams + `google.protobuf.Any`, `Struct`, `Value` or `ListValue` can send data lines + that look the same; for those set `streaming.ndjson_envelope: true` (or + `ProxyServer::with_ndjson_envelope(true)`), and every data line becomes + `{"result": }` and the error `{"error": }`, the + grpc-gateway shape. It is off by default because it changes the data lines + too. +- **SSE**: an event of type `stream-error` (listen with + `addEventListener("stream-error", ...)`) whose data is the error body. It is + not `EventSource.onerror`, which fires on transport failures. This is the HTTP/JSON transcoding format. It is not the Connect protocol's error format, and it is not an OAuth 2.0 token endpoint error body (RFC 6749 §5.2): an upstream that needs one answers successfully with that body instead (see [Upstream controls](#upstream-controls)). -**Switching details off.** In the config file, `error_details:` (see -[Configuration](#configuration)) is read by the standalone binary and by -`ProxyServer::from_yaml_str` / `ProxyServer::from_file`; it is not part of -`ProxyConfig`, so `ProxyConfig::from_yaml_str` alone ignores it. Both log a -warning for a top-level or `streaming:` key no setting reads, so a misspelled -`error_detail:` or `ndjson_envelop:` shows up at startup instead of silently -leaving the default in force. An embedding -service can choose in code with `ProxyServer::with_error_details`. Overrides are -checked in the order they are added; for each switch (`enabled`, `opaque`) the -first rule whose pattern matches the mounted route and that sets the switch -decides, otherwise the global value; `*` stays within one path segment (a path -parameter counts as one) and `**` spans segments. A config rule that sets -neither switch is rejected: +**Switching details off.** `error_details:` in the config file (see +[Configuration](#configuration)) is read by the binary and by +`ProxyServer::from_yaml_str` / `ProxyServer::from_file`. It is not part of +`ProxyConfig`, so `ProxyConfig::from_yaml_str` alone ignores it. Both warn at +startup about a top-level or `streaming:` key they do not know, so a typo such +as `error_detail:` does not silently keep the default. In code, use +`ProxyServer::with_error_details`. + +Route rules are checked in order: for each switch (`enabled`, `opaque`), the +first rule that matches the route and sets that switch decides; with none, the +global value holds. `*` matches within one path segment (a path parameter +counts as one), `**` across segments. A config rule that sets neither switch +is an error: ```rust use structured_proxy::transcode::error::ErrorDetailsPolicy; @@ -592,31 +597,30 @@ methods other than the five standard ones. The upstream RPC decides all of these; the proxy only carries them, as Envoy's `grpc_json_transcoder` and grpc-gateway do, so the same service works behind any of them. -**Request headers → request metadata.** Each header named in -`forwarded_headers` reaches the upstream as request metadata byte for byte, -every value in the order the client sent it, so a check that depends on how -often a header was sent (RFC 9449 §4.3 rejects a request with two `DPoP` -headers) sees the same request behind the proxy. A `-bin` header keeps the -base64 it arrived with, one metadata value per comma-separated part. A value -gRPC metadata cannot carry (empty, or outside visible ASCII and space, such as -a tab or obs-text, or not canonical base64 under a `-bin` key) is refused with -`INVALID_ARGUMENT` (400) naming the header: gRPC lets a receiver drop such a -value, which would change what the upstream counts. A `forwarded_headers` name -must be a gRPC metadata key (letters, digits, `_`, `-`, `.`), or the proxy -does not start. W3C trace-context is the exception, listed or -not: the upstream always gets exactly one valid `traceparent` (the client's -first, or a fresh one when it is missing or malformed), and every `tracestate` -line only with the client's own trace. - -**Response metadata → response headers.** The upstream's response metadata is -its HTTP response headers. Every ASCII entry becomes a header, in order, with -repeated values as repeated fields; a key sent in both the initial metadata and -the trailers keeps both values. This covers a successful unary call (initial -metadata and trailers), a failed call (its trailers-only metadata, or the -response headers and trailers of a call that failed after sending headers, so a -`401` carries its `WWW-Authenticate`), and the initial metadata of a server-streaming -call (its trailers arrive after the headers are sent and are not forwarded). -Never forwarded: +**Request headers → request metadata.** Each header listed in +`forwarded_headers` reaches the upstream as metadata: + +- byte for byte, every value, in the client's order, so a check that counts + headers (RFC 9449 §4.3 rejects two `DPoP` headers) sees the request as sent; +- a `-bin` header keeps its base64, one metadata value per comma-separated + part; +- a value gRPC metadata cannot carry (empty, a character outside visible ASCII + and space, or bad base64 under a `-bin` key) is refused with + `INVALID_ARGUMENT` (400) naming the header, since gRPC lets a receiver + silently drop such a value; +- a listed name that is not a gRPC metadata key (letters, digits, `_`, `-`, + `.`) stops the proxy at startup. + +Trace context goes through whether listed or not: the upstream gets exactly +one valid `traceparent` (the client's first, or a new one when it is missing +or malformed), and `tracestate` only together with the client's own trace. + +**Response metadata → response headers.** The upstream's response metadata +becomes HTTP response headers: every ASCII entry, in order, repeated values as +repeated headers, a key in both initial metadata and trailers with both +values. That covers a unary call, succeeded or failed (so a `401` carries its +`WWW-Authenticate`), and the initial metadata of a server stream; a stream's +trailers arrive after its headers are sent. Not forwarded: - gRPC's own keys: `grpc-*`, binary `-bin` keys and `content-type` (the proxy sets it for the body it writes); @@ -638,31 +642,30 @@ response headers plus the exposed ones, so a browser client that must read a 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 -integer from 200 to 599. Anything else (a value that is not three digits, out -of range, or given twice) turns the answer into +metadata `x-http-code` (grpc-gateway's convention) sets the HTTP status, one +integer from 200 to 599. A value that is not three digits, is out of range or +comes twice turns the whole answer into `{"error": "INTERNAL", "code": 13, "message": "upstream returned a malformed response", "details": []}` -(500), with nothing else of the upstream's answer. `204`, `205` and `304` are -sent without a body or `Content-Type` (RFC 9110 §15.3.5, §15.3.6, §15.4.5). -Errors keep the -`google.rpc.Code` mapping: a protocol-specific error body is a successful -answer with `x-http-code` and that body. Server-streaming calls ignore the key. - -**Raw bodies with `google.api.HttpBody`.** An RPC whose response type is -`google.api.HttpBody`, or whose `response_body` names a field of that type, -answers with `content_type` as `Content-Type` (none when empty) and `data` as -the raw body. An RPC whose request type is `HttpBody` with `body: "*"`, or -whose `body` names a field of that type, receives the raw request body and its -full `Content-Type` value there. With a named field, the other fields still -come from the path and query (a query key naming the body field is ignored); -with `body: "*"`, the query binds nothing, since every field comes from the -body. A server-streaming `HttpBody` writes each message's `data` as it -arrives, with `Content-Type` from the first message; as a raw body has no -in-band error frame, a failure after the first message aborts the transfer so -the client does not take a partial body for a complete one. An `HttpBody` -content type that is not a valid header value is a malformed response (500). -`google/api/httpbody.proto` is always resolvable for error details, like the -`google/rpc` types. +(500). `204`, `205` and `304` go out without a body or `Content-Type` +(RFC 9110 §15.3.5, §15.3.6, §15.4.5). Failed calls keep the `google.rpc.Code` +mapping, so a protocol's own error body is sent as a successful call with +`x-http-code` and that body. Server-streaming calls ignore the key. + +**Raw bodies with `google.api.HttpBody`.** + +- An RPC that returns `google.api.HttpBody` (or whose `response_body` names a + field of that type) answers with `data` as the body and `content_type` as + `Content-Type` (none when empty; a value that is not a valid header is a + malformed response, 500). +- An RPC that takes `HttpBody` with `body: "*"` (or whose `body` names a field + of that type) receives the raw request body and its full `Content-Type`. + With a named field, the other fields still come from the path and query; a + query key naming the body field is ignored. With `body: "*"` every field + comes from the body, so the query binds nothing. +- A server-streaming `HttpBody` writes each message's `data` as it arrives, + with `Content-Type` from the first message. A raw body has no way to carry + an error, so a failure after the first message aborts the transfer, and the + client does not take a partial body for a complete one. An RFC 6749 token endpoint, for example: @@ -678,28 +681,26 @@ client gets exactly that `400`. An authorization endpoint answers `x-http-code: 302` with `location` and an empty `HttpBody`; a JWKS endpoint returns `application/jwk-set+json` (RFC 7517 §8.5). -**`custom` rules.** `HttpRule.custom` (`{kind, path}`) binds any method token: +**`custom` rules.** `HttpRule.custom` (`{kind, path}`) binds any method: `kind: "HEAD"`, `kind: "OPTIONS"`, an extension method such as `PROPFIND` (case-sensitive, RFC 9110 §9.1), or `kind: "*"` for every method, as -`google/api/http.proto` defines. A forward-auth sub-request (nginx -`auth_request`, Traefik `forwardAuth`) arrives with the original request's -method, so a `*` rule answers it whatever that method is. `custom` works in -`additional_bindings` too. A `*` rule takes its path for every method, so -another binding on that path is rejected at startup. OpenAPI lists a `*` rule -under every operation, and cannot describe an extension method. Only a real -CORS preflight (an `OPTIONS` request with both `Origin` and -`Access-Control-Request-Method`) is -answered by the CORS layer; any other `OPTIONS` request reaches its route. +`google/api/http.proto` defines; in `additional_bindings` too. + +- A `*` rule suits a forward-auth sub-request (nginx `auth_request`, Traefik + `forwardAuth`), which arrives with the original request's method. +- A `*` rule owns its path for every method, so another binding on the same + path is an error at startup. +- OpenAPI lists a `*` rule under every operation and leaves extension methods + out. +- Only a real CORS preflight (`OPTIONS` with both `Origin` and + `Access-Control-Request-Method`) is answered by CORS; any other `OPTIONS` + request reaches its route. ## Library Usage -`cargo add structured-proxy` adds the library alone: the binary and its -command-line dependencies (`clap`, `tracing-subscriber`) sit behind the `cli` -feature, which is off by default, so a service that embeds the proxy compiles -none of them. The library starts no runtime and installs no logger of its own; -it runs on the -embedder's tokio runtime and logs through `tracing` to whatever subscriber the -embedder sets up. +`cargo add structured-proxy` adds the library; the binary and its +command-line dependencies come with the `cli` feature. The library runs on +your tokio runtime and logs through `tracing` to the subscriber you set up. ```rust use std::path::Path; @@ -724,12 +725,11 @@ 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. +A gRPC service that embeds the proxy to add REST hands the proxy its own +services instead of an address. Transcoded calls then reach them in process, +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; @@ -826,33 +826,35 @@ loop { 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)`. By default no guard covers those -requests or native gRPC ones, and they reach your service untouched; name -`grpc` or `fallback` in a guard's scope to put them behind it (see -[Guards and scopes](#guards-and-scopes)). CORS and tracing stay the fallback's -own. - -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. +### What passes through + +A request no route matches gets `404`, or goes to your own service with +`ProxyService::with_fallback(my_axum_app)`. Native gRPC calls and the fallback +pass only the guards whose scope names `grpc` or `fallback` (see +[Guards and scopes](#guards-and-scopes)); by default none does. The fallback +keeps its own CORS and tracing. + +gRPC-Web calls pass through as they are, so the upstream answers them: wrap +your services in tonic-web's layer +(`tower::ServiceBuilder::new().layer(tonic_web::GrpcWebLayer::new()).service(grpc)`) +for binary and text gRPC-Web alike. + +- Browsers get the proxy's CORS policy on these calls, the same one their + preflight got (`cors.grpc_web`, on by default). `grpc-status`, + `grpc-message` and `grpc-status-details-bin` are always exposed to them. +- A preflight for a gRPC-Web call (one announcing `x-grpc-web`) goes where + the call goes: the proxy answers it, the fallback never sees it. With + `cors.grpc_web: false` it reaches the upstream, whose own CORS policy then + covers the preflight and the call. +- When the upstream cannot take a call at all, the proxy answers in the + request's own protocol. + +**Deadlines.** A call waits at most five seconds for the upstream's response +headers, or less when the client's `grpc-timeout` says so; then the client +gets `504` `DEADLINE_EXCEEDED`, with an upstream in process or remote. The +client's `grpc-timeout` travels to the upstream; the five-second default does +not, so an upstream that applies `grpc-timeout` to the whole call does not cut +a long server stream short. ### Merging into an axum application @@ -876,11 +878,9 @@ async fn main() -> anyhow::Result<()> { ### Embedding hooks (axum-free) -Inject *stateless* service-specific logic without naming an HTTP framework in -your own crate: implement the hook traits with foundational types (`http`, -`bytes`, `serde_json`) plus `async-trait` (the traits are `#[async_trait]`), -none of which is an HTTP framework. `cargo tree -i axum` in your crate then -shows `axum` solely under `structured-proxy`. +Hooks plug your own stateless logic into the proxy. Their traits use `http`, +`bytes` and `serde_json` types and `#[async_trait]`, so your crate implements +them without depending on axum itself. ```rust use std::sync::Arc; @@ -1020,27 +1020,23 @@ ProxyServer::from_config(config) # } ``` -Injection also resolves a problem the features cannot: Cargo unifies features -across the whole dependency graph, so `rust_crypto` / `aws_lc_rs` is a property -of the *resolution*, not of a binary. Two crates in one workspace that link this -one and want different backends cannot both get their way: the resolution enables -both features, and the tie-break above picks `aws_lc_rs` for everyone. A consumer -that injects its own verifier is not in that argument at all: it takes +Injection also gets around Cargo feature unification. Features apply to the +whole dependency graph, so when two crates in one workspace ask for different +backends, both are enabled and the tie-break above picks `aws_lc_rs` for +everyone. A crate that injects its own verifier takes no backend at all: ```toml [dependencies] -structured-proxy = { version = "4", default-features = false } +structured-proxy = { version = "6", default-features = false } # What the verifier above is written with: the trait is `#[async_trait]`, and # claims cross it as `serde_json::Value`. Neither is re-exported. async-trait = "0.1" serde_json = "1" ``` -which links no JWT or TLS crypto (see [TLS crypto](#tls-crypto)), and -supplies the backend from its own binary. With no -verifier injected and no backend feature, an `auth.mode: "jwt"` config is -rejected at startup with that instruction, rather than silently accepting -tokens. +and brings its own crypto (see [TLS crypto](#tls-crypto)). With neither an +injected verifier nor a backend feature, an `auth.mode: "jwt"` config fails at +startup with a message saying so. ## TLS crypto @@ -1070,46 +1066,36 @@ CRL and name-constraint advisories are listed in `deny.toml` with why they do no apply: the provider reads only algorithm identifiers from it, and rustls verifies certificates with its own patched `rustls-webpki`. -## How It Works +## How it works -1. Load the proto descriptor from a pre-compiled descriptor file -2. Parse `google.api.http` annotations → generate REST routes -3. Incoming HTTP request → transcode to gRPC (path params + query params + JSON body → protobuf) -4. Forward to the upstream gRPC service -5. Response protobuf → transcode to JSON -6. Serve the OpenAPI spec at `/openapi.json` - -## Architecture +At startup the proxy reads your proto descriptors and turns every +`google.api.http` rule into a REST route. Each request is then sorted once: ``` -Client (HTTP/JSON) - │ - ▼ -┌──────────────────────┐ -│ structured-proxy │ -│ │ -│ ┌─────────────────┐ │ -│ │ CORS │ │ -│ ├─────────────────┤ │ -│ │ Maintenance │ │ 503 gate (exempt paths) -│ ├─────────────────┤ │ -│ │ Concurrency │ │ in-flight limit (503) -│ ├─────────────────┤ │ -│ │ Shield │ │ rate limiting (429) -│ ├─────────────────┤ │ -│ │ Auth (JWT) │ │ validate + policies (401/403) -│ ├─────────────────┤ │ -│ │ Transcoder │ │ REST → gRPC -│ │ (prost-reflect) │ │ JSON → Protobuf -│ ├─────────────────┤ │ -│ │ OpenAPI gen │ │ /openapi.json -│ └─────────────────┘ │ -└─────────┬─────────────┘ - │ gRPC - ▼ - Upstream Service + REST, gRPC and gRPC-Web clients (HTTP/1.1, HTTP/2, optional TLS) + │ + ┌──────────────▼──────────────┐ + │ listener │ TLS / mTLS, connection limit + └──────────────┬──────────────┘ + ┌────────────────────────┼────────────────────────┐ + ▼ ▼ ▼ + a REST route gRPC / gRPC-Web no route matches + (transcoded call or content type + own endpoint) + │ │ │ + CORS, guards guards in scope guards in scope + │ │ │ + transcoder passed through your fallback + JSON ↔ protobuf unchanged (or 404) + │ │ + └───────────┬────────────┘ + ▼ + upstream: a remote gRPC server, or your tonic services in process ``` +Guards run in this order: maintenance, concurrency limit, rate limits keyed +before auth, JWT, rate limits keyed by claims, ext_authz, the auth decider. +
## Support the Project From 3d0057a27b3a7e976ea0d5491964b452e54d3b45 Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 11:18:32 +0300 Subject: [PATCH 6/9] feat(grpc-web): translate gRPC-Web for a gRPC-only upstream - `grpc_web.translate: true` converts gRPC-Web calls (binary and text, HTTP/1.1 too) to gRPC through tonic-web, for an upstream without a gRPC-Web layer; CORS and the guards stay around it. - The translated call goes out as HTTP/2 and without a size hint: tonic-web reports the base64 length of a text body for the decoded bytes, which a remote HTTP/2 upstream refused as a protocol error (covered by the remote text translation test). - Translation without the proxy's gRPC-Web CORS is a config error: an upstream that speaks only gRPC cannot answer the browser's preflight. - gRPC paths that need guards or translation are built as boxed stacks per protocol; the others still pass through with nothing in between. Part of #118 --- Cargo.toml | 8 +- README.md | 14 +++- src/config.rs | 24 ++++++ src/config/tests.rs | 4 +- src/lib.rs | 8 +- src/service.rs | 173 +++++++++++++++++++++++++++++++------------ src/service/tests.rs | 38 ++++++++-- tests/edge.rs | 91 ++++++++++++++++++++++- tests/embedded.rs | 1 + 9 files changed, 295 insertions(+), 66 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 00810d3..fed0eb7 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -60,6 +60,8 @@ hyper-util = { version = "0.1", features = ["server-auto", "service", "tokio"] } tonic = { version = "0.14", features = ["tls-connect-info"] } tokio-rustls = { version = "0.26", default-features = false } tonic-health = "0.14" +# `grpc_web.translate`: gRPC-Web to gRPC for an upstream that speaks only gRPC. +tonic-web = "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 # bodies can render typed details even when the product descriptors do not @@ -184,12 +186,6 @@ 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 4ae0c56..cf54cb9 100644 --- a/README.md +++ b/README.md @@ -33,7 +33,8 @@ tonic services. **Edge** -- REST and native gRPC (and gRPC-Web) on one port, HTTP/1.1 and HTTP/2 +- REST and native gRPC on one port, HTTP/1.1 and HTTP/2; gRPC-Web passed + through, or translated to gRPC for an upstream that speaks only gRPC - Built-in TLS and mTLS, and a cap on open connections; or your own TLS, with the client's address and certificate reaching the upstream - Guards you scope to the traffic they cover (transcoded calls, the proxy's @@ -139,6 +140,12 @@ cors: # itself: its preflights then reach the upstream too. grpc_web: true +# Optional: convert gRPC-Web calls to gRPC for an upstream that speaks only +# gRPC (binary and text gRPC-Web, over HTTP/1.1 too). Off: gRPC-Web passes +# through for the upstream to answer. Needs cors.grpc_web. +grpc_web: + translate: false + # Optional: path aliases (rewrite before routing) aliases: - from: "/api/v1/*" @@ -837,7 +844,10 @@ keeps its own CORS and tracing. gRPC-Web calls pass through as they are, so the upstream answers them: wrap your services in tonic-web's layer (`tower::ServiceBuilder::new().layer(tonic_web::GrpcWebLayer::new()).service(grpc)`) -for binary and text gRPC-Web alike. +for binary and text gRPC-Web alike. For an upstream that speaks only gRPC, +set `grpc_web.translate: true` and the proxy converts gRPC-Web calls to gRPC +and the answers back; that needs `cors.grpc_web`, since such an upstream +cannot answer a browser's preflight. - Browsers get the proxy's CORS policy on these calls, the same one their preflight got (`cors.grpc_web`, on by default). `grpc-status`, diff --git a/src/config.rs b/src/config.rs index 87e585b..7d5c208 100644 --- a/src/config.rs +++ b/src/config.rs @@ -89,6 +89,23 @@ pub struct ProxyConfig { /// Request concurrency limit. #[serde(default)] pub concurrency: Option, + + /// How gRPC-Web calls reach the upstream. + #[serde(default)] + pub grpc_web: GrpcWebConfig, +} + +/// How gRPC-Web calls reach the upstream. +#[derive(Debug, Clone, Default, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct GrpcWebConfig { + /// Translate gRPC-Web (binary and text, over HTTP/1.1 too) to gRPC for an + /// upstream that speaks only gRPC. Off: gRPC-Web passes through, for an + /// upstream that translates it itself. Needs the proxy's CORS policy on + /// gRPC-Web (`cors.grpc_web`), since the upstream then cannot answer a + /// browser's preflight. + #[serde(default)] + pub translate: bool, } fn default_forwarded_headers() -> Vec { @@ -235,6 +252,7 @@ pub(crate) const KNOWN_TOP_LEVEL_KEYS: &[&str] = &[ "response_headers", "runtime", "concurrency", + "grpc_web", ]; /// Every `streaming:` key: the [`StreamingConfig`] fields plus the ones @@ -1244,6 +1262,12 @@ impl ProxyConfig { if self.streaming.sse_keep_alive_secs == 0 { anyhow::bail!("streaming.sse_keep_alive_secs must be greater than 0"); } + if self.grpc_web.translate && !self.cors.grpc_web { + anyhow::bail!( + "grpc_web.translate needs cors.grpc_web: an upstream that speaks only \ + gRPC cannot answer a browser's preflight for a gRPC-Web call" + ); + } self.validate_edge_paths()?; Ok(()) } diff --git a/src/config/tests.rs b/src/config/tests.rs index b3c2201..2d705c2 100644 --- a/src/config/tests.rs +++ b/src/config/tests.rs @@ -499,6 +499,7 @@ fn known_top_level_keys_cover_every_proxy_config_field() { forwarded_headers: _, streaming: _, concurrency: _, + grpc_web: _, } = config; for key in [ "upstream", @@ -522,10 +523,11 @@ fn known_top_level_keys_cover_every_proxy_config_field() { "response_headers", "runtime", "concurrency", + "grpc_web", ] { assert!(KNOWN_TOP_LEVEL_KEYS.contains(&key), "{key}"); } - assert_eq!(KNOWN_TOP_LEVEL_KEYS.len(), 21); + assert_eq!(KNOWN_TOP_LEVEL_KEYS.len(), 22); } #[test] diff --git a/src/lib.rs b/src/lib.rs index 17ce749..3aedd5c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -505,7 +505,13 @@ impl ProxyServer { // 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, guards)) + Ok(ProxyService::new( + upstream, + routes, + grpc_web_cors, + guards, + self.config.grpc_web.translate, + )) } /// Build the axum router with all endpoints, calling `upstream`, the CORS diff --git a/src/service.rs b/src/service.rs index 1f28bec..d36283e 100644 --- a/src/service.rs +++ b/src/service.rs @@ -34,9 +34,10 @@ use crate::upstream::{ /// [`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 +/// 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 +/// or the proxy's with `grpc_web.translate` for an upstream that speaks only +/// gRPC. 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, @@ -75,9 +76,9 @@ pub struct ProxyService { /// 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 gRPC path behind the guards that cover it, one per protocol; - /// `None` when no guard does, so gRPC calls pay nothing for guards. - guarded: Option, + /// The gRPC paths behind guards or through gRPC-Web translation; a + /// protocol that needs neither passes through with nothing in between. + boxed: BoxedGrpc, /// The guards, for the fallback an embedder sets later. guards: Arc, /// The connection the requests arrive on, set per connection by the @@ -88,18 +89,81 @@ pub struct ProxyService { impl std::fmt::Debug for ProxyService { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ProxyService") - .field("guarded_grpc", &self.guarded.is_some()) + .field("boxed_grpc", &self.boxed.grpc.is_some()) + .field("boxed_grpc_web", &self.boxed.web.is_some()) .field("connection", &self.connection) .finish_non_exhaustive() } } -/// The guarded gRPC path, each protocol's stack built once. -#[derive(Clone)] -struct GuardedGrpc { - grpc: BoxedService, - web: BoxedService, - web_text: BoxedService, +/// The gRPC paths that need more than a pass-through (guards, gRPC-Web +/// translation), each protocol's stack built once; `None` for a protocol that +/// passes through as it is. +#[derive(Clone, Default)] +struct BoxedGrpc { + grpc: Option, + web: Option, + web_text: Option, +} + +/// Hands a gRPC-Web call, translated to gRPC by tonic-web, to an upstream +/// that speaks only gRPC. +#[derive(Clone, Debug)] +struct Translated { + upstream: U, +} + +impl Service> for Translated { + type Response = http::Response; + type Error = Infallible; + type Future = PassThrough; + + #[inline] + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, request: http::Request) -> Self::Future { + let (mut parts, body) = request.into_parts(); + // A browser's call arrives over HTTP/1.1; as gRPC it is an HTTP/2 + // request (gRPC PROTOCOL-HTTP2). + parts.version = http::Version::HTTP_2; + // tonic-web reports the size of the base64 text body for the decoded + // one, and an HTTP/2 client would announce that as its length; the + // gRPC body goes without one. + let body = tonic::body::Body::new(Unsized { body }); + PassThrough::new( + self.upstream.clone(), + http::Request::from_parts(parts, body), + GrpcProtocol::Grpc, + ) + } +} + +pin_project! { + /// `body` without its size hint. + struct Unsized { + #[pin] + body: tonic::body::Body, + } +} + +impl http_body::Body for Unsized { + type Data = Bytes; + type Error = tonic::Status; + + #[inline] + fn poll_frame( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, tonic::Status>>> { + self.project().body.poll_frame(cx) + } + + #[inline] + fn is_end_stream(&self) -> bool { + self.body.is_end_stream() + } } /// gRPC-Web pass-through, one CORS-wrapped path per encoding. @@ -208,38 +272,52 @@ impl ConnectionInfo { impl ProxyService { /// The service over `upstream` and `routes`; gRPC-Web answers carry - /// `grpc_web_cors` when set. + /// `grpc_web_cors` when set, and gRPC-Web calls are translated to gRPC + /// when `translate_grpc_web`. pub(crate) fn new( upstream: U, routes: axum::Router, grpc_web_cors: Option, guards: Arc, + translate_grpc_web: bool, ) -> Self { let forward = |protocol| Forward { upstream: upstream.clone(), protocol, }; - let guarded = guards.cover(Class::Grpc).then(|| { - // CORS outermost, so a rejected gRPC-Web call still carries it - // and a preflight is answered before any guard. - let stack = |protocol| { - let rejections = GrpcRejections { - inner: guards.service(BoxedService::new(forward(protocol)), Class::Grpc), - protocol, - }; - match (&grpc_web_cors, protocol) { - (Some(cors), GrpcProtocol::Web | GrpcProtocol::WebText) => { - BoxedService::new(cors.layer(rejections)) - } - _ => BoxedService::new(rejections), - } + let guard_grpc = guards.cover(Class::Grpc); + let stack = |protocol: GrpcProtocol| { + let web = protocol != GrpcProtocol::Grpc; + let mut service = if web && translate_grpc_web { + let translator = tonic_web::GrpcWebLayer::new().layer(Translated { + upstream: upstream.clone(), + }); + BoxedService::new(ServiceExt::::map_response( + translator, + |response| response.map(axum::body::Body::new), + )) + } else { + BoxedService::new(forward(protocol)) }; - GuardedGrpc { - grpc: stack(GrpcProtocol::Grpc), - web: stack(GrpcProtocol::Web), - web_text: stack(GrpcProtocol::WebText), + if guard_grpc { + service = BoxedService::new(GrpcRejections { + inner: guards.service(service, Class::Grpc), + protocol, + }); } - }); + // CORS outermost, so a rejected gRPC-Web call still carries it. + match &grpc_web_cors { + Some(cors) if web => BoxedService::new(cors.layer(service)), + _ => service, + } + }; + let boxed_web = guard_grpc || translate_grpc_web; + let boxed = BoxedGrpc { + grpc: guard_grpc.then(|| stack(GrpcProtocol::Grpc)), + web: boxed_web.then(|| stack(GrpcProtocol::Web)), + web_text: boxed_web.then(|| stack(GrpcProtocol::WebText)), + }; + // The pass-through of gRPC-Web, and the answer to its preflights. let grpc_web = grpc_web_cors.map(|cors| GrpcWebCors { web: cors.layer(forward(GrpcProtocol::Web)), web_text: cors.layer(forward(GrpcProtocol::WebText)), @@ -248,7 +326,7 @@ impl ProxyService { upstream, routes, grpc_web, - guarded, + boxed, guards, connection: None, } @@ -316,7 +394,7 @@ impl ProxyService { upstream: self.upstream.clone(), routes: self.routes.clone(), grpc_web: self.grpc_web.clone(), - guarded: self.guarded.clone(), + boxed: self.boxed.clone(), guards: self.guards.clone(), connection: Some(connection.into()), } @@ -365,9 +443,16 @@ where let protocol = call.or_else(|| { is_grpc_web_preflight(request.method(), request.headers()).then_some(GrpcProtocol::Web) }); - let inner = if let (Some(protocol), Some(guarded)) = (call, &mut self.guarded) { + // A preflight is not a call: it skips the guards and the translation. + let boxed = match call { + Some(GrpcProtocol::Grpc) => self.boxed.grpc.as_mut(), + Some(GrpcProtocol::Web) => self.boxed.web.as_mut(), + Some(GrpcProtocol::WebText) => self.boxed.web_text.as_mut(), + None => None, + }; + let inner = if let Some(service) = boxed { // The guards read the peer as axum's `ConnectInfo`, the upstream - // as tonic records it. A preflight is not a call and skips them. + // as tonic records it. if let Some(connection) = connection_of(self.connection.as_ref(), &request) { let extensions = request.extensions_mut(); if let Some(remote) = connection.remote_addr() { @@ -375,14 +460,10 @@ where } connection.into_tonic_extensions(extensions); } - let service = match protocol { - GrpcProtocol::Grpc => &mut guarded.grpc, - GrpcProtocol::Web => &mut guarded.web, - GrpcProtocol::WebText => &mut guarded.web_text, - }; - // Every service on the guarded path is ready at once: the guards - // are middleware and `Forward` waits for the upstream per request. - Inner::Guarded { + // Every service on these paths is ready at once: guards and the + // translator are middleware, and `Forward` waits for the upstream + // per request. + Inner::Boxed { future: service.call(request.map(axum::body::Body::new)), } } else if let Some(protocol) = protocol { @@ -451,7 +532,7 @@ pin_project! { #[pin] future: tower_http::cors::ResponseFuture>, }, - Guarded { + Boxed { #[pin] future: >::Future, }, @@ -467,7 +548,7 @@ impl Future for ResponseFuture { InnerProj::Routes { future } => future.poll(cx), InnerProj::Grpc { call } => call.poll(cx), InnerProj::GrpcWeb { future } => future.poll(cx), - InnerProj::Guarded { future } => future.poll(cx), + InnerProj::Boxed { future } => future.poll(cx), } } } diff --git a/src/service/tests.rs b/src/service/tests.rs index 20989fc..e06ba62 100644 --- a/src/service/tests.rs +++ b/src/service/tests.rs @@ -61,7 +61,7 @@ impl Service> for Recorder { /// 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, Arc::default()) + ProxyService::new(upstream, routes, None, Arc::default(), false) } fn grpc_request(path: &str) -> http::Request { @@ -287,8 +287,14 @@ fn peer_routes() -> axum::Router { #[tokio::test] async fn an_http_request_carries_its_peer_for_the_middleware() { - let proxy = ProxyService::new(Recorder::default(), peer_routes(), None, Arc::default()) - .for_connection(connection()); + let proxy = ProxyService::new( + Recorder::default(), + peer_routes(), + None, + Arc::default(), + false, + ) + .for_connection(connection()); let response = proxy .oneshot(http::Request::get("/peer").body(Body::empty()).unwrap()) .await @@ -298,8 +304,14 @@ async fn an_http_request_carries_its_peer_for_the_middleware() { #[tokio::test] async fn an_http_request_carries_its_connection_for_the_transcoder() { - let proxy = ProxyService::new(Recorder::default(), peer_routes(), None, Arc::default()) - .for_connection(connection()); + let proxy = ProxyService::new( + Recorder::default(), + peer_routes(), + None, + Arc::default(), + false, + ) + .for_connection(connection()); let response = proxy .oneshot( http::Request::get("/connection") @@ -313,7 +325,13 @@ async fn an_http_request_carries_its_connection_for_the_transcoder() { #[tokio::test] async fn an_http_request_without_a_connection_carries_none() { - let proxy = ProxyService::new(Recorder::default(), peer_routes(), None, Arc::default()); + let proxy = ProxyService::new( + Recorder::default(), + peer_routes(), + None, + Arc::default(), + false, + ); let response = proxy .oneshot( http::Request::get("/connection") @@ -341,7 +359,13 @@ async fn a_grpc_request_on_an_axum_server_carries_its_peer_to_the_upstream() { #[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, Arc::default()); + let proxy = ProxyService::new( + Recorder::default(), + peer_routes(), + None, + Arc::default(), + false, + ); let mut request = http::Request::get("/connection") .body(Body::empty()) .unwrap(); diff --git a/tests/edge.rs b/tests/edge.rs index 069a260..cd400f4 100644 --- a/tests/edge.rs +++ b/tests/edge.rs @@ -399,6 +399,75 @@ async fn the_auth_decider_scoped_to_grpc_gates_native_calls_by_the_client_addres let (status, body, _) = http1_get(addr, "/v1/echo/rest").await; assert_eq!(status, 200, "{body}"); } + +// --- gRPC-Web translation ------------------------------------------------------------ + +async fn binary_grpc_web_is_translated_for_an_upstream_that_speaks_only_grpc() { + // The upstream has no gRPC-Web layer; the proxy converts the call, over + // HTTP/1.1 here, to gRPC and the answer back. + let app = translating_proxy(UPSTREAM, "").await; + let (content_type, body) = + grpc_web_echo_via(app, "application/grpc-web+proto", request_frame("web")).await; + assert_eq!(content_type, "application/grpc-web+proto"); + assert_eq!(seen_name(&body), "web"); +} + +async fn text_grpc_web_is_translated_for_an_upstream_that_speaks_only_grpc() { + use base64::Engine as _; + let engine = base64::engine::general_purpose::STANDARD; + let app = translating_proxy(UPSTREAM, "").await; + let (content_type, body) = grpc_web_echo_via( + app, + "application/grpc-web-text+proto", + engine.encode(request_frame("text")).into_bytes(), + ) + .await; + assert_eq!(content_type, "application/grpc-web-text+proto"); + let decoded: Vec = body + .chunks(4) + .flat_map(|group| engine.decode(group).unwrap()) + .collect(); + assert_eq!(seen_name(&decoded), "text"); +} + +async fn a_guard_rejects_a_translated_call_in_grpc_web() { + let yaml = "maintenance:\n enabled: true\n scope:\n traffic: [grpc]\n"; + let request = http::Request::post("/test.v1.Edge/Echo") + .header("content-type", "application/grpc-web+proto") + .header("x-grpc-web", "1") + .body(Body::from(request_frame("web"))) + .unwrap(); + let app = translating_proxy(UPSTREAM, yaml).await; + let response = tower::ServiceExt::oneshot(app, request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response.headers()["content-type"], + "application/grpc-web+proto" + ); + assert_eq!(response.headers()["grpc-status"], "14"); +} +} + +/// The proxy in front of the `Edge` service with no gRPC-Web layer of its +/// own, translating gRPC-Web, with `extra_yaml`. +async fn translating_proxy(upstream: common::Upstream, extra_yaml: &str) -> common::App { + let pool = pool(); + common::app(upstream, Edge { pool: pool.clone() }, |yaml| { + ProxyServer::from_yaml_str(&format!("{yaml}grpc_web:\n translate: true\n{extra_yaml}")) + .unwrap() + .with_descriptors(pool) + }) + .await +} + +#[test] +fn translation_needs_the_proxys_grpc_web_cors() { + // The upstream cannot answer the preflight of a call it cannot read. + let err = + ProxyServer::from_yaml_str("grpc_web:\n translate: true\ncors:\n grpc_web: false\n") + .err() + .expect("translation without the proxy's gRPC-Web CORS is refused"); + assert!(err.to_string().contains("grpc_web.translate"), "{err}"); } /// Denies every request with `403`, recording the peer it saw. @@ -804,6 +873,15 @@ fn seen_name(body: &[u8]) -> String { /// 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) { + grpc_web_echo_via(grpc_web_proxy(), content_type, body).await +} + +/// [`grpc_web_echo`] through `app`. +async fn grpc_web_echo_via( + app: common::App, + 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) @@ -811,10 +889,17 @@ async fn grpc_web_echo(content_type: &str, body: Vec) -> (String, bytes::Byt .header("x-grpc-web", "1") .body(Body::from(body)) .unwrap(); - let response = tower::ServiceExt::oneshot(grpc_web_proxy(), request) - .await - .unwrap(); + let response = tower::ServiceExt::oneshot(app, request).await.unwrap(); assert_eq!(response.status(), StatusCode::OK); + // A trailers-only answer is a failed call, whatever its HTTP status. + assert!( + response + .headers() + .get("grpc-status") + .is_none_or(|status| status == "0"), + "{:?}", + response.headers() + ); let content_type = response.headers()["content-type"] .to_str() .unwrap() diff --git a/tests/embedded.rs b/tests/embedded.rs index 36b823e..c0c73cf 100644 --- a/tests/embedded.rs +++ b/tests/embedded.rs @@ -48,6 +48,7 @@ fn embedded_config_is_constructible() { forwarded_headers: vec!["authorization".into()], streaming: Default::default(), concurrency: None, + grpc_web: Default::default(), }; // The server accepts a programmatically-built config (the embedded path). let _server = ProxyServer::from_config(config); From bb9c7ae355b52ff62c6c1612fea1ac9ba408dbd1 Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 11:24:21 +0300 Subject: [PATCH 7/9] feat(embed)!: capability builder and a selection of transcoded RPCs - `ProxyServer::new()` starts with every capability off: native gRPC goes to the upstream, everything else to the fallback. One builder method per config section (`with_listen`, `with_health`, `with_metrics`, `with_openapi`, `with_cors`, `with_maintenance`, `with_concurrency_limit`, `with_rate_limits`, `with_auth`, ...) sets the same field the YAML does, so code and file describe one proxy. - `config::from_yaml` reads a single section, for the sections that are built through serde. - `transcode.only` / `with_transcoded_rpcs` narrows transcoding to named services and methods; routes, the collision check and the OpenAPI spec follow the same selection, and a name the descriptors do not hold fails the build. - `ProxyConfig` implements `Default` (the empty file). BREAKING CHANGE: `transcode::route_paths` and `openapi::generate` take the RPC selection; `ProxyConfig` gains the `transcode` field. Closes #118 --- README.md | 45 ++++++++++ src/config.rs | 74 ++++++++++++++++ src/config/tests.rs | 4 +- src/lib.rs | 184 ++++++++++++++++++++++++++++++++++++++-- src/openapi.rs | 19 ++++- src/openapi/tests.rs | 35 +++++++- src/transcode/mod.rs | 42 ++++++--- src/transcode/select.rs | 88 +++++++++++++++++++ src/transcode/tests.rs | 2 +- tests/edge.rs | 97 +++++++++++++++++++++ tests/embedded.rs | 1 + 11 files changed, 564 insertions(+), 27 deletions(-) create mode 100644 src/transcode/select.rs diff --git a/README.md b/README.md index cf54cb9..f8b0931 100644 --- a/README.md +++ b/README.md @@ -57,6 +57,8 @@ tonic services. - Your own tonic services as the upstream, called in process ([Library Usage](#library-usage)) +- A builder that starts as a plain pass-through and turns on only the + capabilities you ask for ([Building the proxy in code](#building-the-proxy-in-code)) - Requests the proxy does not serve go to your own fallback service - Hooks for auth decisions, token verification, OIDC and extra routes, with no HTTP framework in your code @@ -146,6 +148,11 @@ cors: grpc_web: translate: false +# Optional: transcode only these services and methods; every annotated RPC +# by default. A name the descriptors do not hold stops the proxy at startup. +transcode: + only: ["my.package.v1.MyService", "my.package.v1.Admin/GetStatus"] + # Optional: path aliases (rewrite before routing) aliases: - from: "/api/v1/*" @@ -762,6 +769,44 @@ structured_proxy::serve(listener, proxy).await?; 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. +### Building the proxy in code + +`ProxyServer::new()` starts with every capability off: native gRPC goes to the +upstream, every other request to your fallback. Turn on only what you need; +each method sets the config section of the same name, so a proxy built in code +and one read from YAML behave the same: + +```rust +use structured_proxy::config::{self, ConcurrencyConfig, ScopeConfig, Traffic}; +use structured_proxy::ProxyServer; + +# fn build(pool: prost_reflect::DescriptorPool, grpc: tonic::service::Routes) -> anyhow::Result<()> { +let service = ProxyServer::new() + // REST for two RPCs of your API. + .with_descriptors(pool) + .with_transcoded_rpcs(["acme.v1.Orders/GetOrder", "acme.v1.Orders/ListOrders"]) + // At most 1000 native gRPC calls in flight. + .with_concurrency_limit(ConcurrencyConfig { + max_in_flight: 1000, + scope: Some(ScopeConfig::traffic([Traffic::Grpc])), + }) + // Sections with many options come from YAML, the file's own syntax. + .with_rate_limits(config::from_yaml( + "enabled: true\nprofiles:\n anon: { rate: \"600/min\" }\nrules:\n - pattern: \"/**\"\n key: { type: ip }\n profile: anon\n", + )?) + .service(grpc)?; +# let _ = service; +# Ok(()) +# } +``` + +The methods are `with_upstream_address`, `with_listen`, `with_descriptors`, +`with_transcoded_rpcs`, `with_aliases`, `with_forwarded_headers`, +`with_health`, `with_metrics`, `with_openapi`, `with_oidc_discovery`, +`with_cors`, `with_grpc_web_translation`, `with_streaming`, +`with_maintenance`, `with_concurrency_limit`, `with_rate_limits` and +`with_auth`, next to the hooks below. + ### TLS and connection limits `listen.tls` and `listen.max_connections` configure the listener of diff --git a/src/config.rs b/src/config.rs index 7d5c208..fae50e5 100644 --- a/src/config.rs +++ b/src/config.rs @@ -93,6 +93,79 @@ pub struct ProxyConfig { /// How gRPC-Web calls reach the upstream. #[serde(default)] pub grpc_web: GrpcWebConfig, + + /// Which annotated RPCs are transcoded. + #[serde(default)] + pub transcode: TranscodeConfig, +} + +impl Default for ProxyConfig { + /// The configuration of an empty YAML file: health and metrics endpoints + /// on, everything else off. + fn default() -> Self { + Self { + upstream: None, + descriptors: Vec::new(), + listen: ListenConfig::default(), + service: ServiceConfig::default(), + aliases: Vec::new(), + openapi: None, + auth: None, + shield: None, + oidc_discovery: None, + health: HealthConfig::default(), + metrics: MetricsConfig::default(), + maintenance: MaintenanceConfig::default(), + cors: CorsConfig::default(), + logging: LoggingConfig::default(), + metrics_classes: Vec::new(), + forwarded_headers: default_forwarded_headers(), + streaming: StreamingConfig::default(), + concurrency: None, + grpc_web: GrpcWebConfig::default(), + transcode: TranscodeConfig::default(), + } + } +} + +/// Which annotated RPCs are transcoded. +/// +/// ```yaml +/// transcode: +/// only: ["acme.v1.Orders", "acme.v1.Users/GetUser"] +/// ``` +#[derive(Debug, Clone, Default, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct TranscodeConfig { + /// Services (`package.Service`) and single methods + /// (`package.Service/Method`) whose `google.api.http` rules become REST + /// routes; empty transcodes every annotated RPC. A name the descriptors do + /// not hold stops the proxy at startup. + #[serde(default)] + pub only: Vec, +} + +/// One config section from YAML, the way the config file reads it: for the +/// sections an embedder hands to a builder method, such as +/// [`ProxyServer::with_rate_limits`](crate::ProxyServer::with_rate_limits). +/// +/// # Errors +/// +/// YAML that is not that section. +/// +/// # Examples +/// +/// ``` +/// use structured_proxy::config::{self, ShieldConfig}; +/// +/// let shield: ShieldConfig = config::from_yaml( +/// "enabled: true\nprofiles:\n anon: { rate: \"60/min\" }\nrules:\n - pattern: \"/**\"\n key: { type: ip }\n profile: anon\n", +/// ) +/// .unwrap(); +/// assert!(shield.enabled); +/// ``` +pub fn from_yaml(yaml: &str) -> anyhow::Result { + Ok(serde_yaml::from_str(yaml)?) } /// How gRPC-Web calls reach the upstream. @@ -253,6 +326,7 @@ pub(crate) const KNOWN_TOP_LEVEL_KEYS: &[&str] = &[ "runtime", "concurrency", "grpc_web", + "transcode", ]; /// Every `streaming:` key: the [`StreamingConfig`] fields plus the ones diff --git a/src/config/tests.rs b/src/config/tests.rs index 2d705c2..109b338 100644 --- a/src/config/tests.rs +++ b/src/config/tests.rs @@ -500,6 +500,7 @@ fn known_top_level_keys_cover_every_proxy_config_field() { streaming: _, concurrency: _, grpc_web: _, + transcode: _, } = config; for key in [ "upstream", @@ -524,10 +525,11 @@ fn known_top_level_keys_cover_every_proxy_config_field() { "runtime", "concurrency", "grpc_web", + "transcode", ] { assert!(KNOWN_TOP_LEVEL_KEYS.contains(&key), "{key}"); } - assert_eq!(KNOWN_TOP_LEVEL_KEYS.len(), 22); + assert_eq!(KNOWN_TOP_LEVEL_KEYS.len(), 23); } #[test] diff --git a/src/lib.rs b/src/lib.rs index 3aedd5c..e8edf9d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -124,6 +124,13 @@ pub struct ProxyServer { transcode: transcode::TranscodeOptions, } +impl Default for ProxyServer { + /// [`ProxyServer::new`]: every capability off. + fn default() -> Self { + Self::new() + } +} + impl ProxyServer { /// Create from YAML config file. pub fn from_config(config: ProxyConfig) -> Self { @@ -193,6 +200,150 @@ impl ProxyServer { &self.config } + /// A proxy with every capability off: native gRPC and gRPC-Web pass to + /// the upstream, every other request to the fallback (`404` by default). + /// Turn on what you need with the methods below; each sets the config + /// section of the same name, so code and YAML describe one proxy. + /// + /// # Examples + /// + /// ``` + /// use structured_proxy::config::{ConcurrencyConfig, ScopeConfig, Traffic}; + /// use structured_proxy::ProxyServer; + /// + /// # fn build() -> anyhow::Result<()> { + /// // Your gRPC API, with at most 1000 calls in flight and nothing else. + /// let grpc = tonic::service::Routes::default(); + /// let service = ProxyServer::new() + /// .with_concurrency_limit(ConcurrencyConfig { + /// max_in_flight: 1000, + /// scope: Some(ScopeConfig::traffic([Traffic::Grpc])), + /// }) + /// .service(grpc)?; + /// # let _ = service; + /// # Ok(()) + /// # } + /// # build().unwrap(); + /// ``` + pub fn new() -> Self { + let mut config = ProxyConfig::default(); + config.health.enabled = false; + config.metrics.enabled = false; + Self::from_config(config) + } + + /// The remote gRPC upstream (`upstream.default`), for + /// [`upstream`](Self::upstream) and [`serve`](Self::serve). + pub fn with_upstream_address(mut self, address: impl Into) -> Self { + self.config.upstream = Some(config::UpstreamConfig { + default: address.into(), + }); + self + } + + /// The listener of [`serve`](Self::serve): address, TLS, connection + /// limit (`listen:`). + pub fn with_listen(mut self, listen: config::ListenConfig) -> Self { + self.config.listen = listen; + self + } + + /// Transcode only these services (`package.Service`) and methods + /// (`package.Service/Method`) of the descriptors (`transcode.only`); every + /// annotated RPC by default. A name the descriptors do not hold fails the + /// build. + pub fn with_transcoded_rpcs( + mut self, + names: impl IntoIterator>, + ) -> Self { + self.config.transcode.only = names.into_iter().map(Into::into).collect(); + self + } + + /// Path aliases rewritten before routing (`aliases:`). + pub fn with_aliases(mut self, aliases: impl IntoIterator) -> Self { + self.config.aliases = aliases.into_iter().collect(); + self + } + + /// The request headers transcoded calls forward as gRPC metadata + /// (`forwarded_headers:`), replacing the default list. + pub fn with_forwarded_headers( + mut self, + names: impl IntoIterator>, + ) -> Self { + self.config.forwarded_headers = names.into_iter().map(Into::into).collect(); + self + } + + /// Health probe endpoints (`health:`). + pub fn with_health(mut self, health: config::HealthConfig) -> Self { + self.config.health = health; + self + } + + /// The Prometheus metrics endpoint (`metrics:`). + pub fn with_metrics(mut self, metrics: config::MetricsConfig) -> Self { + self.config.metrics = metrics; + self + } + + /// The OpenAPI spec and docs endpoints (`openapi:`). + pub fn with_openapi(mut self, openapi: config::OpenApiConfig) -> Self { + self.config.openapi = Some(openapi); + self + } + + /// Static OIDC discovery and JWKS endpoints (`oidc_discovery:`). + pub fn with_oidc_discovery(mut self, oidc: config::OidcDiscoveryConfig) -> Self { + self.config.oidc_discovery = Some(oidc); + self + } + + /// The CORS policy (`cors:`). + pub fn with_cors(mut self, cors: config::CorsConfig) -> Self { + self.config.cors = cors; + self + } + + /// Translate gRPC-Web to gRPC for an upstream that speaks only gRPC + /// (`grpc_web.translate`). + pub fn with_grpc_web_translation(mut self, translate: bool) -> Self { + self.config.grpc_web.translate = translate; + self + } + + /// Server-streaming behavior (`streaming:`). + pub fn with_streaming(mut self, streaming: config::StreamingConfig) -> Self { + self.config.streaming = streaming; + self + } + + /// Maintenance mode (`maintenance:`). + pub fn with_maintenance(mut self, maintenance: config::MaintenanceConfig) -> Self { + self.config.maintenance = maintenance; + self + } + + /// The limit on requests in flight (`concurrency:`). + pub fn with_concurrency_limit(mut self, concurrency: config::ConcurrencyConfig) -> Self { + self.config.concurrency = Some(concurrency); + self + } + + /// Rate limits (`shield:`); build the section with [`config::from_yaml`]. + pub fn with_rate_limits(mut self, shield: config::ShieldConfig) -> Self { + self.config.shield = Some(shield); + self + } + + /// JWT auth, route policies, forward-auth and ext_authz (`auth:`); build + /// the section with [`config::from_yaml`]. + pub fn with_auth(mut self, auth: config::AuthConfig) -> Self { + self.config.auth = Some(auth); + self + } + /// Create with an embedded descriptor pool (for sid-proxy backward compat). pub fn with_descriptors(mut self, pool: DescriptorPool) -> Self { self.descriptor_pool = Some(pool); @@ -398,7 +549,11 @@ impl ProxyServer { /// surface (injected backend or config-driven static discovery), embedder /// extra routes, and the transcoded REST routes. All built-in surfaces here /// are `GET`. - fn reserved_routes(&self, pool: &DescriptorPool) -> anyhow::Result> { + fn reserved_routes( + &self, + pool: &DescriptorPool, + selection: &transcode::RpcSelection, + ) -> anyhow::Result> { let mut routes = Vec::new(); let mut get = |path: String| routes.push(("GET".to_string(), path)); if self.config.health.enabled { @@ -433,7 +588,11 @@ impl ProxyServer { for route in &self.extra_routes { routes.push((route.method.as_str().to_string(), route.path.clone())); } - routes.extend(transcode::route_paths(pool, &self.config.aliases)); + routes.extend(transcode::route_paths( + pool, + &self.config.aliases, + selection, + )); Ok(routes) } @@ -524,6 +683,8 @@ impl ProxyServer { // config is built directly instead of through `from_yaml_str`. self.config.validate()?; let pool = self.load_descriptors()?; + let selection = transcode::RpcSelection::new(&pool, &self.config.transcode.only) + .map_err(|e| anyhow::anyhow!("invalid transcode.only: {e}"))?; let service_name = self.config.service.name.clone(); @@ -538,7 +699,7 @@ impl ProxyServer { // routes with different methods are legal (they merge), so only a // repeated (method, path) — or any overlap with the verify endpoint, // which answers ALL methods (`*`) — is a real conflict. - let mut mounted = self.reserved_routes(&pool)?; + let mut mounted = self.reserved_routes(&pool, &selection)?; if let Some(vp) = &verify_path { mounted.push(("*".to_string(), vp.clone())); } @@ -611,8 +772,11 @@ impl ProxyServer { let cors = self.build_cors()?; // Build transcoding routes from descriptor pool. - let transcode_routes = - transcode::routes_with_options(&pool, &self.config.aliases, &self.transcode); + let transcode_routes = transcode::routes_with_options( + &pool, + &self.config.aliases, + &self.transcode.clone().with_selection(selection.clone()), + ); // JWT auth, if configured (auth.mode == "jwt"). let auth = match &self.config.auth { @@ -704,7 +868,7 @@ impl ProxyServer { }; // OpenAPI + docs routes (if enabled). - let openapi_routes = self.build_openapi_routes(&pool); + let openapi_routes = self.build_openapi_routes(&pool, &selection); // OIDC routes (public, like the health endpoints). An injected // OidcBackend supersedes the config-driven static discovery: the proxy @@ -848,7 +1012,11 @@ impl ProxyServer { Ok(guards) } - fn build_openapi_routes(&self, pool: &DescriptorPool) -> Router + fn build_openapi_routes( + &self, + pool: &DescriptorPool, + selection: &transcode::RpcSelection, + ) -> Router where S: Clone + Send + Sync + 'static, { @@ -857,7 +1025,7 @@ impl ProxyServer { _ => return Router::new(), }; - let spec = openapi::generate(pool, openapi_config, &self.config.aliases); + let spec = openapi::generate(pool, openapi_config, &self.config.aliases, selection); let spec_json = serde_json::to_string_pretty(&spec).unwrap_or_default(); let openapi_path = openapi_config.path.clone(); let docs_path = openapi_config.docs_path.clone(); diff --git a/src/openapi.rs b/src/openapi.rs index 4b8b951..a3722a8 100644 --- a/src/openapi.rs +++ b/src/openapi.rs @@ -14,6 +14,7 @@ use crate::config::{AliasConfig, OpenApiConfig}; use crate::transcode::httpbody; use crate::transcode::request::BodyMapping; use crate::transcode::rule::{self, HttpBinding, RouteMethod}; +use crate::transcode::RpcSelection; /// The operations an OpenAPI 3.0 path item can hold, in the order a `*` rule /// lists them. @@ -21,8 +22,14 @@ const OPERATIONS: [&str; 8] = [ "get", "put", "post", "delete", "options", "head", "patch", "trace", ]; -/// Generate OpenAPI 3.0 JSON spec from a descriptor pool. -pub fn generate(pool: &DescriptorPool, config: &OpenApiConfig, aliases: &[AliasConfig]) -> Value { +/// Generate OpenAPI 3.0 JSON spec from a descriptor pool: the RPCs +/// `selection` transcodes, and their services as tags. +pub fn generate( + pool: &DescriptorPool, + config: &OpenApiConfig, + aliases: &[AliasConfig], + selection: &RpcSelection, +) -> Value { let title = config.title.as_deref().unwrap_or("API"); let version = config.version.as_deref().unwrap_or("1.0.0"); @@ -33,6 +40,9 @@ pub fn generate(pool: &DescriptorPool, config: &OpenApiConfig, aliases: &[AliasC let http_ext = pool.get_extension_by_name("google.api.http"); for service in pool.services() { + if !selection.is_all() && !service.methods().any(|m| selection.selects(&m)) { + continue; + } let service_name = service.name().to_string(); let service_full = service.full_name().to_string(); @@ -48,8 +58,9 @@ pub fn generate(pool: &DescriptorPool, config: &OpenApiConfig, aliases: &[AliasC continue; }; for method in service.methods() { - if method.is_client_streaming() { - continue; // No REST mapping for client-streaming. + // No REST mapping for client-streaming. + if method.is_client_streaming() || !selection.selects(&method) { + continue; } for binding in rule::http_bindings(&method, http_ext) { diff --git a/src/openapi/tests.rs b/src/openapi/tests.rs index 57abf22..0f0eb74 100644 --- a/src/openapi/tests.rs +++ b/src/openapi/tests.rs @@ -53,7 +53,7 @@ fn config() -> OpenApiConfig { #[test] fn test_generate_empty_pool() { let pool = DescriptorPool::new(); - let spec = generate(&pool, &config(), &[]); + let spec = generate(&pool, &config(), &[], &RpcSelection::default()); assert_eq!(spec["openapi"], "3.0.3"); assert_eq!(spec["info"]["title"], "Test API"); @@ -203,11 +203,17 @@ impl protox::file::FileResolver for TestProtos { } fn spec(aliases: &[AliasConfig]) -> Value { + spec_of(aliases, &[]) +} + +/// [`spec`] documenting only the RPCs `only` selects. +fn spec_of(aliases: &[AliasConfig], only: &[&str]) -> Value { let pool = protox::Compiler::with_file_resolver(TestProtos) .open_file("test/v1/api.proto") .unwrap() .descriptor_pool(); - generate(&pool, &config(), aliases) + let selection = RpcSelection::new(&pool, only).unwrap(); + generate(&pool, &config(), aliases, &selection) } #[test] @@ -371,3 +377,28 @@ fn aliases_get_their_own_operation_ids() { assert!(aliased.is_string()); assert_ne!(aliased, "Api.Create"); } + +#[test] +fn a_selection_documents_only_the_rpcs_it_transcodes() { + // The spec describes what is mounted: an RPC left out of `transcode.only` + // has no route, so it has no operation either. + let spec = spec_of(&[], &["test.v1.Api/Find"]); + let paths = spec["paths"].as_object().unwrap(); + assert_eq!(paths.keys().collect::>(), ["/v1/find"]); + assert_eq!(spec["tags"][0]["name"], "Api"); + // A whole service selects all of its RPCs. + let spec = spec_of(&[], &["test.v1.Api"]); + assert!(spec["paths"]["/v1/items"]["post"].is_object()); +} + +#[test] +fn a_selection_naming_nothing_in_the_descriptors_is_an_error() { + let pool = protox::Compiler::with_file_resolver(TestProtos) + .open_file("test/v1/api.proto") + .unwrap() + .descriptor_pool(); + for name in ["test.v1.Nope", "test.v1.Api/Nope", "Api"] { + let err = RpcSelection::new(&pool, [name]).unwrap_err(); + assert!(err.contains(name), "{err}"); + } +} diff --git a/src/transcode/mod.rs b/src/transcode/mod.rs index 1453077..87b27a8 100644 --- a/src/transcode/mod.rs +++ b/src/transcode/mod.rs @@ -17,6 +17,9 @@ pub mod metadata; pub mod request; pub(crate) mod response; pub(crate) mod rule; +mod select; + +pub use select::RpcSelection; use axum::body::{Body, Bytes}; use axum::extract::{Path, RawQuery, State}; @@ -174,9 +177,17 @@ pub struct TranscodeOptions { pub(crate) error_details: ErrorDetailsPolicy, pub(crate) ndjson_envelope: bool, pub(crate) denied_response_headers: Arc<[HeaderName]>, + pub(crate) selection: RpcSelection, } impl TranscodeOptions { + /// Transcode only the RPCs `selection` names (every annotated one by + /// default). + pub fn with_selection(mut self, selection: RpcSelection) -> Self { + self.selection = selection; + self + } + /// Which routes return `google.rpc.Status` details in their error bodies /// (all of them by default). pub fn with_error_details(mut self, policy: ErrorDetailsPolicy) -> Self { @@ -230,7 +241,7 @@ pub fn routes_with_options( aliases: &[AliasConfig], options: &TranscodeOptions, ) -> Router { - let bindings = route_bindings(pool, aliases); + let bindings = route_bindings(pool, aliases, &options.selection); if bindings.is_empty() { tracing::warn!("No HTTP-annotated RPCs found in proto descriptors"); return Router::new(); @@ -392,9 +403,13 @@ struct RouteBinding { /// [`routes`] (to build handlers) and [`route_paths`] (to enumerate paths for /// collision checks) consume this, so the mounted set and the enumerated set /// cannot drift apart. -fn route_bindings(pool: &DescriptorPool, aliases: &[AliasConfig]) -> Vec { +fn route_bindings( + pool: &DescriptorPool, + aliases: &[AliasConfig], + selection: &RpcSelection, +) -> Vec { let mut bindings = Vec::new(); - for entry in extract_routes(pool) { + for entry in extract_routes(pool, selection) { for alias in aliases { if let Some(suffix) = entry.http_path.strip_prefix(&alias.to) { if alias.from.ends_with("/{path}") { @@ -414,9 +429,10 @@ fn route_bindings(pool: &DescriptorPool, aliases: &[AliasConfig]) -> Vec Vec Vec<(String, String)> { - route_bindings(pool, aliases) +pub fn route_paths( + pool: &DescriptorPool, + aliases: &[AliasConfig], + selection: &RpcSelection, +) -> Vec<(String, String)> { + route_bindings(pool, aliases, selection) .into_iter() .map(|b| (b.entry.http_method.as_str().to_string(), b.axum_path)) .collect() @@ -1089,9 +1109,9 @@ fn request_content_type(headers: &HeaderMap) -> Result { } } -/// Extract the route entries of every HTTP binding of every unary and -/// server-streaming RPC. Client-streaming RPCs have no HTTP mapping. -fn extract_routes(pool: &DescriptorPool) -> Vec { +/// Extract the route entries of every HTTP binding of every selected unary +/// and server-streaming RPC. Client-streaming RPCs have no HTTP mapping. +fn extract_routes(pool: &DescriptorPool, selection: &RpcSelection) -> Vec { let http_ext = match pool.get_extension_by_name("google.api.http") { Some(ext) => ext, None => { @@ -1104,7 +1124,7 @@ fn extract_routes(pool: &DescriptorPool) -> Vec { for service in pool.services() { for method in service.methods() { - if method.is_client_streaming() { + if method.is_client_streaming() || !selection.selects(&method) { continue; } diff --git a/src/transcode/select.rs b/src/transcode/select.rs new file mode 100644 index 0000000..86def57 --- /dev/null +++ b/src/transcode/select.rs @@ -0,0 +1,88 @@ +//! Which annotated RPCs become REST routes. + +use std::sync::Arc; + +use prost_reflect::{DescriptorPool, MethodDescriptor}; + +/// The RPCs transcoded to REST: every one with a `google.api.http` rule, or +/// only the services and methods named. +/// +/// # Examples +/// +/// ``` +/// use structured_proxy::transcode::RpcSelection; +/// +/// let pool = prost_reflect::DescriptorPool::new(); +/// // An empty list selects every annotated RPC. +/// let all = RpcSelection::new(&pool, Vec::::new()).unwrap(); +/// # let _ = all; +/// // A name the descriptors do not hold is an error. +/// assert!(RpcSelection::new(&pool, ["acme.v1.Orders"]).is_err()); +/// ``` +#[derive(Clone, Debug, Default)] +pub struct RpcSelection { + names: Arc<[Name]>, +} + +/// A selected service, or one method of it. +#[derive(Debug)] +struct Name { + service: String, + method: Option, +} + +impl RpcSelection { + /// `names` out of `pool`: `package.Service` for every RPC of a service, + /// `package.Service/Method` for one; none selects every annotated RPC. + /// + /// # Errors + /// + /// A name that is not a service, or a method of a service, of `pool`: a + /// typo would otherwise leave its routes silently unmounted. + pub fn new( + pool: &DescriptorPool, + names: impl IntoIterator>, + ) -> Result { + let names = names + .into_iter() + .map(|name| { + let name = name.as_ref(); + let (service, method) = match name.split_once('/') { + Some((service, method)) => (service, Some(method)), + None => (name, None), + }; + let known = pool.get_service_by_name(service).is_some_and(|found| { + method.is_none_or(|method| found.methods().any(|m| m.name() == method)) + }); + if !known { + return Err(format!( + "{name:?} is not a service or method of the proto descriptors" + )); + } + Ok(Name { + service: service.to_owned(), + method: method.map(str::to_owned), + }) + }) + .collect::, _>>()?; + Ok(Self { + names: names.into(), + }) + } + + /// Whether every annotated RPC is transcoded. + pub(crate) fn is_all(&self) -> bool { + self.names.is_empty() + } + + /// Whether `method` is transcoded. + pub(crate) fn selects(&self, method: &MethodDescriptor) -> bool { + self.is_all() || { + let service = method.parent_service(); + self.names.iter().any(|name| { + name.service == service.full_name() + && name.method.as_deref().is_none_or(|m| m == method.name()) + }) + } + } +} diff --git a/src/transcode/tests.rs b/src/transcode/tests.rs index b4f370c..4595e66 100644 --- a/src/transcode/tests.rs +++ b/src/transcode/tests.rs @@ -79,7 +79,7 @@ service S { "#, ); let alias: AliasConfig = serde_yaml::from_str("from: /api/{path}\nto: /v1").unwrap(); - let paths = route_paths(&pool, &[alias]); + let paths = route_paths(&pool, &[alias], &RpcSelection::default()); for expected in [ "/v1/files/{*path}", "/api/files/{*path}", diff --git a/tests/edge.rs b/tests/edge.rs index cd400f4..3d7bcfb 100644 --- a/tests/edge.rs +++ b/tests/edge.rs @@ -460,6 +460,103 @@ async fn translating_proxy(upstream: common::Upstream, extra_yaml: &str) -> comm .await } +// --- the capability builder ------------------------------------------------------ + +/// Send a native gRPC `Echo` through `app`; returns the gRPC status code. +async fn grpc_echo_status(app: common::App) -> String { + let request = http::Request::post("/test.v1.Edge/Echo") + .version(http::Version::HTTP_2) + .header("content-type", "application/grpc") + .header("te", "trailers") + .body(Body::from(request_frame("native"))) + .unwrap(); + let response = tower::ServiceExt::oneshot(app, request).await.unwrap(); + let (parts, body) = response.into_parts(); + if let Some(status) = parts.headers.get("grpc-status") { + return status.to_str().unwrap().to_owned(); + } + let collected = http_body_util::BodyExt::collect(body).await.unwrap(); + collected.trailers().unwrap()["grpc-status"] + .to_str() + .unwrap() + .to_owned() +} + +#[tokio::test] +async fn a_new_proxy_passes_everything_through() { + // Nothing on: no transcoding, no health or metrics endpoints; native gRPC + // reaches the upstream and REST the fallback. + let fallback = axum::Router::new().fallback(|| async { (StatusCode::IM_A_TEAPOT, "yours") }); + let service = ProxyServer::new() + .service(tonic::service::Routes::new(Edge { pool: pool() })) + .unwrap() + .with_fallback(fallback); + let app = common::App::new(service); + for path in ["/v1/echo/a", "/health", "/metrics"] { + let (status, body) = + common::send(&app, http::Request::get(path).body(Body::empty()).unwrap()).await; + assert_eq!(status, StatusCode::IM_A_TEAPOT, "{path}"); + assert_eq!(body, "yours"); + } + assert_eq!(grpc_echo_status(app).await, "0"); +} + +#[tokio::test] +async fn only_the_selected_rpcs_are_transcoded() { + // `Hang` never answers: were it transcoded, its route would hang. + let service = ProxyServer::new() + .with_descriptors(pool()) + .with_transcoded_rpcs(["test.v1.Edge/Echo"]) + .service(tonic::service::Routes::new(Edge { pool: pool() })) + .unwrap(); + let app = common::App::new(service); + let (status, seen) = get(&app, "/v1/echo/a", &[]).await; + assert_eq!(status, StatusCode::OK, "{seen}"); + let (status, _) = common::send( + &app, + http::Request::get("/v1/hang").body(Body::empty()).unwrap(), + ) + .await; + assert_eq!(status, StatusCode::NOT_FOUND); +} + +#[test] +fn a_selection_the_descriptors_do_not_hold_fails_the_build() { + let server = ProxyServer::new() + .with_descriptors(pool()) + .with_transcoded_rpcs(["test.v1.Edge/Missing"]); + let Err(err) = server.service(tonic::service::Routes::default()) else { + panic!("an unknown RPC must be refused"); + }; + assert!(err.to_string().contains("transcode.only"), "{err}"); +} + +#[tokio::test] +async fn a_builder_section_behaves_as_its_yaml() { + let maintenance = + structured_proxy::config::from_yaml("enabled: true\nmessage: down\n").unwrap(); + let service = ProxyServer::new() + .with_descriptors(pool()) + .with_maintenance(maintenance) + .service(tonic::service::Routes::new(Edge { pool: pool() })) + .unwrap(); + let (status, body) = get(&common::App::new(service), "/v1/echo/a", &[]).await; + assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(body["message"], "down"); +} + +#[tokio::test] +async fn a_new_proxy_reaches_a_remote_upstream_by_address() { + let url = common::serve(Edge { pool: pool() }).await; + let server = ProxyServer::new() + .with_descriptors(pool()) + .with_upstream_address(url); + let app = common::App::new(server.service(server.upstream().unwrap()).unwrap()); + let (status, seen) = get(&app, "/v1/echo/remote", &[]).await; + assert_eq!(status, StatusCode::OK, "{seen}"); + assert_eq!(seen["name"], "remote"); +} + #[test] fn translation_needs_the_proxys_grpc_web_cors() { // The upstream cannot answer the preflight of a call it cannot read. diff --git a/tests/embedded.rs b/tests/embedded.rs index c0c73cf..8677c27 100644 --- a/tests/embedded.rs +++ b/tests/embedded.rs @@ -49,6 +49,7 @@ fn embedded_config_is_constructible() { streaming: Default::default(), concurrency: None, grpc_web: Default::default(), + transcode: Default::default(), }; // The server accepts a programmatically-built config (the embedded path). let _server = ProxyServer::from_config(config); From 365990e2baaadef4909cac0fae7bd85c5167fd20 Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 20:27:26 +0300 Subject: [PATCH 8/9] fix(edge): close idle connections and refuse silent config - An HTTP/2 connection with no stream open held its max_connections slot forever. A connection with no request in flight for listen.idle_timeout_secs (default 60, 0 = never) is now shut down gracefully (GOAWAY on HTTP/2); a request counts until its response body ends, so streams keep their connection. - hyper drops its default HTTP/1.1 header read timeout without a timer, so a client trickling headers held its connection forever. The connection now runs with a timer and listen.header_read_timeout_secs (default 30); the TLS handshake timeout is listen.tls. handshake_timeout_secs (default 10). ServeOptions exposes all three. - listen.tls.client_auth without client_ca_file gave a listener that verified no client; it is now an error, and a CA alone means required. - A scope method of "*" or one no route answers (a typo) matched no request and left the guard covering nothing; both fail the build, while methods of custom rules and extra routes, and any method for fallback traffic, stay allowed. - The concurrency limit no longer covers the proxy's own endpoints by default: a saturated proxy refused its health probes, getting busy instances restarted. - The response-body wrappers of the concurrency guard and the idle tracking share one type. Regression tests: an_idle_http2_connection_gives_up_its_slot, a_client_that_trickles_its_headers_is_disconnected, client_auth_without_a_client_ca_is_an_error, a_catch_all_method_is_an_error, a_method_no_route_answers_is_an_error, a_saturated_proxy_still_answers_its_health_probes. --- README.md | 30 +++++++--- src/config.rs | 41 +++++++++++-- src/guard.rs | 52 ++++++++++++++--- src/guard/concurrency.rs | 53 ++--------------- src/guard/tests.rs | 54 ++++++++++++++--- src/held.rs | 52 +++++++++++++++++ src/lib.rs | 54 ++++++++++++----- src/serve.rs | 121 +++++++++++++++++++++++++++++++++----- src/serve/idle.rs | 123 +++++++++++++++++++++++++++++++++++++++ src/serve/idle/tests.rs | 56 ++++++++++++++++++ src/serve/tests.rs | 102 ++++++++++++++++++++++++++++++++ src/tls.rs | 12 +++- tests/edge.rs | 21 +++++++ tests/embedded.rs | 2 + tests/tls.rs | 11 ++++ 15 files changed, 675 insertions(+), 109 deletions(-) create mode 100644 src/held.rs create mode 100644 src/serve/idle.rs create mode 100644 src/serve/idle/tests.rs diff --git a/README.md b/README.md index f8b0931..286ad61 100644 --- a/README.md +++ b/README.md @@ -99,6 +99,12 @@ listen: # Optional: most connections served at once; past it the next one waits in # the listen backlog until a connection closes. Unset: no limit. # max_connections: 10000 + # Seconds a connection may go without a request in flight before it is + # closed (HTTP/2 gets a GOAWAY), so idle clients do not hold connection + # slots; 0 keeps idle connections open. A stream keeps its connection. + idle_timeout_secs: 60 + # Seconds an HTTP/1.1 client has to send a request's headers. + header_read_timeout_secs: 30 # Optional: TLS on the listener (REST and gRPC share the port; ALPN offers # h2 and http/1.1). With client_ca_file, client certificates are verified # (mTLS) and reach an in-process tonic upstream as Request::peer_certs. @@ -106,7 +112,8 @@ listen: # cert_file: "/etc/proxy/tls.crt" # PEM chain, leaf first # key_file: "/etc/proxy/tls.key" # PEM private key # client_ca_file: "/etc/proxy/ca.crt" - # client_auth: required # or `optional` + # client_auth: required # or `optional`; needs client_ca_file + # handshake_timeout_secs: 10 # The gRPC service behind the proxy. Required by the standalone binary; an # embedder with an in-process upstream leaves it out. @@ -190,7 +197,8 @@ maintenance: # Optional: request concurrency limit. Requests past max_in_flight get 503 # UNAVAILABLE with `Retry-After: 1` at once; a request holds its slot until its -# response body ends. Default scope: [transcoded, endpoints, grpc]. +# response body ends. Default scope: [transcoded, grpc], so health probes and +# metrics still answer while the proxy is saturated. concurrency: max_in_flight: 512 # scope: { traffic: [grpc], paths: ["/acme.v1.Orders/*"] } @@ -420,13 +428,16 @@ scope: `paths` and `methods` narrow the guard within its traffic; `*` stays within a path segment and `**` spans segments. A native gRPC call's path is -`/./`. A scope whose `traffic` is empty, a relative -or invalid glob, or an invalid method stops the proxy at startup. +`/./`. A method is a standard one or one a route +answers (a `custom` rule, an extra route); a scope covering `fallback` takes +any method. A scope whose `traffic` is empty, a relative or invalid glob, `*` +as a method (leave `methods` out to cover every method) or a method no route +answers stops the proxy at startup. | Guard | Configured by | Default traffic | |-------|---------------|-----------------| | maintenance | `maintenance.scope` | `transcoded`, `endpoints` | -| concurrency limit | `concurrency.scope` | `transcoded`, `endpoints`, `grpc` | +| concurrency limit | `concurrency.scope` | `transcoded`, `grpc` | | rate limits | `shield.scope` | `transcoded`, `endpoints` | | JWT | `auth.scope` | `transcoded`, `endpoints` | | ext_authz | `auth.authz.scope` | `transcoded` | @@ -827,8 +838,13 @@ structured_proxy::serve_with(listener, proxy, options).await?; # } ``` -The TLS handshake runs in each connection's task with a ten-second limit, so -a slow or stalled client holds up no one else. A client certificate the +The TLS handshake runs in each connection's task under +`listen.tls.handshake_timeout_secs`, so a slow or stalled client holds up no +one else. A connection with no request in flight for `listen.idle_timeout_secs` +is closed (HTTP/2 gets a GOAWAY, and a gRPC client reconnects when it next +calls), so idle clients cannot keep `max_connections` slots. The same settings +are `ServeOptions::idle_timeout`, `header_read_timeout` and +`tls_handshake_timeout` in code. A client certificate the listener verified reaches a tonic handler in process as `Request::peer_certs`. TLS needs a rustls crypto provider: the one a crypto backend feature brings, or the one your process installed (see [TLS crypto](#tls-crypto)). diff --git a/src/config.rs b/src/config.rs index fae50e5..f0aa15c 100644 --- a/src/config.rs +++ b/src/config.rs @@ -470,6 +470,27 @@ pub struct ListenConfig { /// TLS on the listener, mTLS with `client_ca_file`. Unset: cleartext. #[serde(default)] pub tls: Option, + /// Seconds a connection may go without a request in flight before it is + /// closed (HTTP/2 gets a GOAWAY), so idle clients do not hold + /// `max_connections` slots. 0 keeps idle connections open. Default: 60. + #[serde(default = "default_idle_timeout_secs")] + pub idle_timeout_secs: u64, + /// Seconds an HTTP/1.1 client has to send a request's headers. At least + /// 1. Default: 30. + #[serde(default = "default_header_read_timeout_secs")] + pub header_read_timeout_secs: u64, +} + +fn default_idle_timeout_secs() -> u64 { + 60 +} + +fn default_header_read_timeout_secs() -> u64 { + 30 +} + +fn default_tls_handshake_timeout_secs() -> u64 { + 10 } /// TLS for the listener. @@ -492,10 +513,15 @@ pub struct ListenTlsConfig { /// Unset: clients present none. #[serde(default)] pub client_ca_file: Option, - /// Whether a client must present a certificate once `client_ca_file` - /// is set. + /// Whether a client must present a certificate; `required` when + /// `client_ca_file` is set and this is not. Setting it without + /// `client_ca_file` is an error: nothing could verify the certificate. #[serde(default)] - pub client_auth: ClientAuth, + pub client_auth: Option, + /// Seconds a client has to finish the TLS handshake. At least 1. + /// Default: 10. + #[serde(default = "default_tls_handshake_timeout_secs")] + pub handshake_timeout_secs: u64, } /// Whether a TLS client must present a certificate. @@ -519,6 +545,8 @@ impl Default for ListenConfig { http: default_http_listen(), max_connections: None, tls: None, + idle_timeout_secs: default_idle_timeout_secs(), + header_read_timeout_secs: default_header_read_timeout_secs(), } } } @@ -629,7 +657,9 @@ pub struct ScopeConfig { /// `/./`. #[serde(default)] pub paths: Vec, - /// Methods the guard is narrowed to; empty covers every method. + /// Methods the guard is narrowed to; empty covers every method. Each is a + /// standard method or one a route answers (a `custom` rule, an extra + /// route); a scope covering the fallback takes any method token. #[serde(default)] pub methods: Vec, } @@ -680,7 +710,8 @@ pub struct ConcurrencyConfig { /// ends, so a stream holds its slot for its whole life. At least 1. pub max_in_flight: usize, /// The traffic the limit covers, one budget shared by all of it. - /// Default: `transcoded`, `endpoints` and `grpc`. + /// Default: `transcoded` and `grpc`; the proxy's own endpoints stay out, + /// so health probes answer while the proxy is saturated. #[serde(default)] pub scope: Option, } diff --git a/src/guard.rs b/src/guard.rs index 2b99808..47df2fb 100644 --- a/src/guard.rs +++ b/src/guard.rs @@ -140,15 +140,18 @@ pub(crate) struct Scope { impl Scope { /// Compile `config` for the guard named `what`, with `default` traffic - /// when the config names none. + /// when the config names none; `routed` are the methods the proxy's + /// routes answer beyond the standard ones (`custom` rules, extra routes). /// /// # Errors /// A scope that covers no traffic, a path glob that is relative or does - /// not compile, or a method that is not an HTTP method token. + /// not compile, `*` or a method no request can carry (see + /// [`scope_method`]). pub(crate) fn compile( config: Option<&ScopeConfig>, default: &[Traffic], what: &str, + routed: &[Method], ) -> Result, String> { let traffic = config.and_then(|c| c.traffic.as_deref()).unwrap_or(default); let classes = traffic.iter().fold(0, |acc, t| acc | traffic_bits(*t)); @@ -186,11 +189,7 @@ impl Scope { Some( names .iter() - .map(|name| { - Method::from_bytes(name.to_ascii_uppercase().as_bytes()).map_err(|_| { - format!("{what}.scope.methods entry {name:?} is not a method") - }) - }) + .map(|name| scope_method(name, what, classes & FALLBACK != 0, routed)) .collect::, _>>()?, ) }; @@ -221,6 +220,45 @@ impl Scope { } } +/// The methods of RFC 9110 §9.3 and PATCH (RFC 5789). +const STANDARD_METHODS: [Method; 9] = [ + Method::GET, + Method::HEAD, + Method::POST, + Method::PUT, + Method::DELETE, + Method::CONNECT, + Method::OPTIONS, + Method::TRACE, + Method::PATCH, +]; + +/// One `scope.methods` entry of the guard `what`: a method some request can +/// carry, or an error. A method no route answers would match no request and +/// leave the guard covering nothing, unless the scope covers the fallback, +/// whose methods are the embedder's. +fn scope_method( + name: &str, + what: &str, + covers_fallback: bool, + routed: &[Method], +) -> Result { + if name == "*" { + return Err(format!( + "{what}.scope.methods entry \"*\" is not a method: leave methods out to cover every one" + )); + } + let method = Method::from_bytes(name.to_ascii_uppercase().as_bytes()) + .map_err(|_| format!("{what}.scope.methods entry {name:?} is not a method"))?; + if covers_fallback || STANDARD_METHODS.contains(&method) || routed.contains(&method) { + Ok(method) + } else { + Err(format!( + "{what}.scope.methods entry {name:?} is neither a standard method nor one a route answers" + )) + } +} + /// A guard layer `L` applied only to the requests its scope's paths and /// methods select; the others go straight to the inner service. #[derive(Clone)] diff --git a/src/guard/concurrency.rs b/src/guard/concurrency.rs index a46bd3d..fde5ce4 100644 --- a/src/guard/concurrency.rs +++ b/src/guard/concurrency.rs @@ -2,18 +2,12 @@ //! turned away at once rather than queued, since a queue behind a saturated //! upstream only adds latency to what will time out anyway. -use std::pin::Pin; use std::sync::Arc; -use std::task::{Context, Poll}; -use axum::body::Body; use axum::extract::{Request, State}; use axum::middleware::Next; use axum::response::Response; -use bytes::Bytes; -use http_body::{Frame, SizeHint}; -use pin_project_lite::pin_project; -use tokio::sync::{OwnedSemaphorePermit, Semaphore}; +use tokio::sync::Semaphore; use super::reject; use crate::config::ConcurrencyConfig; @@ -59,46 +53,7 @@ pub(super) async fn middleware( ); return response; }; - next.run(request).await.map(|body| { - Body::new(Holding { - body, - slot: Some(slot), - }) - }) -} - -pin_project! { - /// A response body that frees its request's slot when it ends or is - /// dropped. - struct Holding { - #[pin] - body: Body, - slot: Option, - } -} - -impl http_body::Body for Holding { - type Data = Bytes; - type Error = axum::Error; - - fn poll_frame( - self: Pin<&mut Self>, - cx: &mut Context<'_>, - ) -> Poll, axum::Error>>> { - let this = self.project(); - let frame = this.body.poll_frame(cx); - if let Poll::Ready(None | Some(Err(_))) = frame { - // Ended: the slot is free before the connection moves on. - this.slot.take(); - } - frame - } - - fn is_end_stream(&self) -> bool { - self.body.is_end_stream() - } - - fn size_hint(&self) -> SizeHint { - self.body.size_hint() - } + next.run(request) + .await + .map(|body| crate::held::until_end(body, slot)) } diff --git a/src/guard/tests.rs b/src/guard/tests.rs index 02d69e3..4e685e5 100644 --- a/src/guard/tests.rs +++ b/src/guard/tests.rs @@ -45,7 +45,7 @@ fn request(method: &str, path: &str) -> http::Request<()> { #[test] fn a_scope_without_traffic_takes_the_default() { - let scope = Scope::compile(None, &[Traffic::Transcoded, Traffic::Grpc], "shield").unwrap(); + let scope = Scope::compile(None, &[Traffic::Transcoded, Traffic::Grpc], "shield", &[]).unwrap(); assert!(scope.covers(Class::Transcoded)); assert!(scope.covers(Class::Grpc)); assert!(!scope.covers(Class::Endpoints)); @@ -56,7 +56,7 @@ fn a_scope_without_traffic_takes_the_default() { #[test] fn all_traffic_covers_every_class() { let config = ScopeConfig::traffic([Traffic::All]); - let scope = Scope::compile(Some(&config), &[], "shield").unwrap(); + let scope = Scope::compile(Some(&config), &[], "shield", &[]).unwrap(); for class in [ Class::Transcoded, Class::Endpoints, @@ -72,7 +72,7 @@ fn all_traffic_covers_every_class() { fn a_scope_that_covers_no_traffic_is_an_error() { // An empty list would mount the guard nowhere; a typo, not a choice. let config = ScopeConfig::traffic([]); - let err = Scope::compile(Some(&config), &[Traffic::Transcoded], "auth").unwrap_err(); + let err = Scope::compile(Some(&config), &[Traffic::Transcoded], "auth", &[]).unwrap_err(); assert!(err.contains("auth.scope.traffic"), "{err}"); } @@ -82,14 +82,14 @@ fn a_relative_or_invalid_path_glob_is_an_error() { paths: vec!["v1/**".into()], ..ScopeConfig::default() }; - let err = Scope::compile(Some(&relative), &[Traffic::Transcoded], "shield").unwrap_err(); + let err = Scope::compile(Some(&relative), &[Traffic::Transcoded], "shield", &[]).unwrap_err(); assert!(err.contains("must start with '/'"), "{err}"); let invalid = ScopeConfig { paths: vec!["/v1/[".into()], ..ScopeConfig::default() }; - let err = Scope::compile(Some(&invalid), &[Traffic::Transcoded], "shield").unwrap_err(); + let err = Scope::compile(Some(&invalid), &[Traffic::Transcoded], "shield", &[]).unwrap_err(); assert!(err.contains("is invalid"), "{err}"); } @@ -99,10 +99,50 @@ fn an_invalid_method_is_an_error() { methods: vec!["GE T".into()], ..ScopeConfig::default() }; - let err = Scope::compile(Some(&config), &[Traffic::Transcoded], "shield").unwrap_err(); + let err = Scope::compile(Some(&config), &[Traffic::Transcoded], "shield", &[]).unwrap_err(); assert!(err.contains("is not a method"), "{err}"); } +fn methods(names: &[&str], traffic: Traffic) -> ScopeConfig { + ScopeConfig { + methods: names.iter().map(|&name| name.into()).collect(), + ..ScopeConfig::traffic([traffic]) + } +} + +#[test] +fn a_catch_all_method_is_an_error() { + // `*` means every method in route policies; here it would match no + // request and leave the guard covering nothing. + let config = methods(&["*"], Traffic::Transcoded); + let err = Scope::compile(Some(&config), &[], "auth", &[]).unwrap_err(); + assert!(err.contains("leave methods out"), "{err}"); +} + +#[test] +fn a_method_no_route_answers_is_an_error() { + // A typo matches no request, and the guard would silently cover nothing. + let config = methods(&["PSOT"], Traffic::Transcoded); + let err = Scope::compile(Some(&config), &[], "shield", &[]).unwrap_err(); + assert!(err.contains("\"PSOT\""), "{err}"); +} + +#[test] +fn an_extension_method_a_route_answers_is_accepted() { + let config = methods(&["PROPFIND"], Traffic::Transcoded); + let routed = [Method::from_bytes(b"PROPFIND").unwrap()]; + let scope = Scope::compile(Some(&config), &[], "shield", &routed).unwrap(); + assert!(scope.matches(&request("PROPFIND", "/dav"))); +} + +#[test] +fn any_method_token_is_accepted_for_the_fallback() { + // The fallback is the embedder's own service: which methods it answers + // is not the proxy's to know. + let config = methods(&["PROPFIND"], Traffic::Fallback); + assert!(Scope::compile(Some(&config), &[], "shield", &[]).is_ok()); +} + #[test] fn paths_and_methods_narrow_the_requests_a_scope_matches() { let config = ScopeConfig { @@ -110,7 +150,7 @@ fn paths_and_methods_narrow_the_requests_a_scope_matches() { methods: vec!["post".into()], ..ScopeConfig::default() }; - let scope = Scope::compile(Some(&config), &[Traffic::Transcoded], "shield").unwrap(); + let scope = Scope::compile(Some(&config), &[Traffic::Transcoded], "shield", &[]).unwrap(); assert!(scope.narrows()); assert!(scope.matches(&request("POST", "/v1/orders/42"))); // Methods compare as tokens, case aside. diff --git a/src/held.rs b/src/held.rs new file mode 100644 index 0000000..55e5328 --- /dev/null +++ b/src/held.rs @@ -0,0 +1,52 @@ +//! A response body that keeps a value alive until it ends. + +use std::pin::Pin; +use std::task::{Context, Poll}; + +use axum::body::Body; +use bytes::Bytes; +use http_body::{Frame, SizeHint}; +use pin_project_lite::pin_project; + +/// `body`, holding `value` until the body ends or is dropped: a slot, a +/// request counted as in flight. +pub(crate) fn until_end(body: Body, value: T) -> Body { + Body::new(Held { + body, + value: Some(value), + }) +} + +pin_project! { + struct Held { + #[pin] + body: Body, + value: Option, + } +} + +impl http_body::Body for Held { + type Data = Bytes; + type Error = axum::Error; + + fn poll_frame( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, axum::Error>>> { + let this = self.project(); + let frame = this.body.poll_frame(cx); + if let Poll::Ready(None | Some(Err(_))) = frame { + // Ended: let go before the connection moves on. + this.value.take(); + } + frame + } + + fn is_end_stream(&self) -> bool { + self.body.is_end_stream() + } + + fn size_hint(&self) -> SizeHint { + self.body.size_hint() + } +} diff --git a/src/lib.rs b/src/lib.rs index e8edf9d..1d9757b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -58,6 +58,7 @@ pub mod config; mod cors; mod embed; mod guard; +mod held; pub mod hooks; pub mod oidc; pub mod openapi; @@ -788,7 +789,7 @@ impl ProxyServer { let forward_auth = auth.as_ref().and_then(|built| { auth::forward::ForwardAuth::build(self.config.auth.as_ref()?, built.clone()) }); - let guards = self.guards(auth, maintenance_exempt)?; + let guards = self.guards(auth, maintenance_exempt, &mounted)?; // Health routes. Paths are configurable; the whole group is skippable. let health_routes = if self.config.health.enabled { @@ -926,7 +927,7 @@ impl ProxyServer { /// The guards the configuration and the hooks turn on, each with its /// scope; `maintenance_exempt` lists the paths maintenance mode leaves - /// reachable. + /// reachable, `mounted` the `(method, path)` of every route. /// /// # Errors /// @@ -936,10 +937,17 @@ impl ProxyServer { &self, auth: Option>, maintenance_exempt: Vec, + mounted: &[(String, String)], ) -> anyhow::Result { use config::Traffic::{Endpoints, Grpc, Transcoded}; + // `*` marks a route that answers every method, not a method. + let routed: Vec = mounted + .iter() + .filter(|(method, _)| method != "*") + .filter_map(|(method, _)| http::Method::from_bytes(method.as_bytes()).ok()) + .collect(); let scope = |config: Option<&ScopeConfig>, default: &[config::Traffic], what: &str| { - guard::Scope::compile(config, default, what).map_err(anyhow::Error::msg) + guard::Scope::compile(config, default, what, &routed).map_err(anyhow::Error::msg) }; let mut guards = guard::Guards::default(); // Mounted only while maintenance is on, so normal traffic pays nothing @@ -961,11 +969,9 @@ impl ProxyServer { if let Some(cfg) = &self.config.concurrency { guards.concurrency = Some(( guard::Concurrency::build(cfg).map_err(anyhow::Error::msg)?, - scope( - cfg.scope.as_ref(), - &[Transcoded, Endpoints, Grpc], - "concurrency", - )?, + // The proxy's own endpoints stay out by default: a health + // probe refused under load gets a busy instance restarted. + scope(cfg.scope.as_ref(), &[Transcoded, Grpc], "concurrency")?, )); } if let Some(cfg) = &self.config.shield { @@ -1126,24 +1132,42 @@ impl ProxyServer { } /// The [`ServeOptions`] of `listen:`: TLS from `listen.tls` (mTLS with its - /// `client_ca_file`) and the `listen.max_connections` cap, for - /// [`serve_with`] on a listener of your own. + /// `client_ca_file`), the `listen.max_connections` cap and the connection + /// timeouts, for [`serve_with`] on a listener of your own. /// /// # Errors /// - /// A `max_connections` of zero, or TLS files that cannot be loaded, or no - /// rustls crypto provider for TLS. + /// A `max_connections`, `header_read_timeout_secs` or + /// `tls.handshake_timeout_secs` of zero, TLS files that cannot be loaded, + /// or no rustls crypto provider for TLS. pub fn serve_options(&self) -> anyhow::Result { + use std::time::Duration; let listen = &self.config.listen; - let mut options = ServeOptions::new(); + anyhow::ensure!( + listen.header_read_timeout_secs > 0, + "listen.header_read_timeout_secs must be at least 1" + ); + let mut options = ServeOptions::new() + .idle_timeout( + (listen.idle_timeout_secs > 0) + .then(|| Duration::from_secs(listen.idle_timeout_secs)), + ) + .header_read_timeout(Duration::from_secs(listen.header_read_timeout_secs)); if let Some(max) = listen.max_connections { anyhow::ensure!(max > 0, "listen.max_connections must be at least 1"); options = options.max_connections(max); } if let Some(tls) = &listen.tls { - options = options.tls( - tls::server_config(tls).map_err(|e| anyhow::anyhow!("invalid listen.tls: {e}"))?, + anyhow::ensure!( + tls.handshake_timeout_secs > 0, + "listen.tls.handshake_timeout_secs must be at least 1" ); + options = options + .tls( + tls::server_config(tls) + .map_err(|e| anyhow::anyhow!("invalid listen.tls: {e}"))?, + ) + .tls_handshake_timeout(Duration::from_secs(tls.handshake_timeout_secs)); } Ok(options) } diff --git a/src/serve.rs b/src/serve.rs index 1052e0d..3bac898 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -4,7 +4,7 @@ use std::sync::Arc; use std::time::Duration; -use hyper_util::rt::{TokioExecutor, TokioIo}; +use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer}; use hyper_util::server::conn::auto::Builder; use hyper_util::service::TowerToHyperService; use tokio::io::{AsyncRead, AsyncWrite}; @@ -15,32 +15,78 @@ use tonic::transport::server::Connected; use crate::service::{ConnectionInfo, ProxyService}; use crate::upstream::Upstream; -/// How long a client has to finish its TLS handshake before the connection is -/// dropped, so a stalled client holds neither a task nor a connection slot. -const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10); +mod idle; /// How [`serve_with`] runs its listener. /// /// # Examples /// /// ``` +/// use std::time::Duration; /// use structured_proxy::ServeOptions; /// -/// let options = ServeOptions::new().max_connections(10_000); +/// let options = ServeOptions::new() +/// .max_connections(10_000) +/// .idle_timeout(Some(Duration::from_secs(120))); /// # let _ = options; /// ``` -#[derive(Clone, Debug, Default)] +#[derive(Clone, Debug)] pub struct ServeOptions { max_connections: Option, tls: Option>, + idle_timeout: Option, + header_read_timeout: Duration, + tls_handshake_timeout: Duration, +} + +impl Default for ServeOptions { + fn default() -> Self { + Self { + max_connections: None, + tls: None, + idle_timeout: Some(Duration::from_secs(60)), + header_read_timeout: Duration::from_secs(30), + tls_handshake_timeout: Duration::from_secs(10), + } + } } impl ServeOptions { - /// Cleartext, with no limit on connections. + /// Cleartext, with no limit on connections; a connection idle for 60 s is + /// closed, a client gets 30 s to send the headers of an HTTP/1.1 request + /// and 10 s to finish a TLS handshake. pub fn new() -> Self { Self::default() } + /// Close a connection that has had no request in flight for `timeout` + /// (gracefully: HTTP/2 gets a GOAWAY); `None` keeps idle connections + /// open. A request counts until its response body ends, so a stream keeps + /// its connection. Without it an idle client holds a + /// [`max_connections`](Self::max_connections) slot for as long as it likes. + #[must_use] + pub fn idle_timeout(mut self, timeout: Option) -> Self { + self.idle_timeout = timeout; + self + } + + /// How long an HTTP/1.1 client has to send a request's headers, so a + /// client that trickles them in holds no connection for long. + #[must_use] + pub fn header_read_timeout(mut self, timeout: Duration) -> Self { + self.header_read_timeout = timeout; + self + } + + /// How long a client has to finish its TLS handshake before the + /// connection is dropped, so a stalled client holds neither a task nor a + /// connection slot. + #[must_use] + pub fn tls_handshake_timeout(mut self, timeout: Duration) -> Self { + self.tls_handshake_timeout = timeout; + self + } + /// Serve at most `max` connections at once: past it, the next connection /// is accepted when one closes, and waits in the listen backlog until then. /// @@ -130,6 +176,11 @@ pub async fn serve_with( .max_connections .map(|max| Arc::new(Semaphore::new(max))); let acceptor = options.tls.map(tokio_rustls::TlsAcceptor::from); + let handshake_timeout = options.tls_handshake_timeout; + let limits = ConnectionLimits { + idle_timeout: options.idle_timeout, + header_read_timeout: options.header_read_timeout, + }; loop { // The slot is taken before the accept, so a full server leaves new // connections in the kernel's backlog instead of accepting and @@ -163,13 +214,11 @@ pub async fn serve_with( match acceptor { None => { let service = service.for_connection(tcp.connect_info()); - serve_connection(tcp, service).await; + serve_connection(tcp, service, limits).await; } Some(acceptor) => { let stream = - match tokio::time::timeout(TLS_HANDSHAKE_TIMEOUT, acceptor.accept(tcp)) - .await - { + match tokio::time::timeout(handshake_timeout, acceptor.accept(tcp)).await { Ok(Ok(stream)) => stream, Ok(Err(error)) => { tracing::debug!(%error, "TLS handshake failed"); @@ -182,22 +231,62 @@ pub async fn serve_with( }; let service = service.for_connection(ConnectionInfo::tls(stream.connect_info())); - serve_connection(stream, service).await; + serve_connection(stream, service, limits).await; } } }); } } +/// The timeouts every connection is served under. +#[derive(Clone, Copy, Debug)] +struct ConnectionLimits { + idle_timeout: Option, + header_read_timeout: Duration, +} + /// HTTP/1.1 or HTTP/2, whichever the client speaks, on one connection. -async fn serve_connection(io: I, service: ProxyService) +async fn serve_connection(io: I, service: ProxyService, limits: ConnectionLimits) where U: Upstream, I: AsyncRead + AsyncWrite + Unpin + Send + 'static, { - let served = Builder::new(TokioExecutor::new()) - .serve_connection_with_upgrades(TokioIo::new(io), TowerToHyperService::new(service)) - .await; + let mut builder = Builder::new(TokioExecutor::new()); + // hyper times nothing without a timer: its own default header read + // timeout is dropped with a warning. + builder + .http1() + .timer(TokioTimer::new()) + .header_read_timeout(limits.header_read_timeout); + builder.http2().timer(TokioTimer::new()); + let served = match limits.idle_timeout { + None => { + builder + .serve_connection_with_upgrades(TokioIo::new(io), TowerToHyperService::new(service)) + .await + } + Some(timeout) => { + let activity = Arc::new(idle::Activity::default()); + let service = idle::Tracked { + inner: service, + activity: Arc::clone(&activity), + }; + let connection = builder.serve_connection_with_upgrades( + TokioIo::new(io), + TowerToHyperService::new(service), + ); + tokio::pin!(connection); + tokio::select! { + served = connection.as_mut() => served, + () = activity.idle_for(timeout) => { + // HTTP/2 gets a GOAWAY, HTTP/1.1 closes after the request + // it is reading, if any. + connection.as_mut().graceful_shutdown(); + connection.await + } + } + } + }; if let Err(error) = served { tracing::debug!(%error, "connection ended"); } diff --git a/src/serve/idle.rs b/src/serve/idle.rs new file mode 100644 index 0000000..96e4931 --- /dev/null +++ b/src/serve/idle.rs @@ -0,0 +1,123 @@ +//! When a connection has had no request in flight for long enough to close. + +use std::convert::Infallible; +use std::future::Future; +use std::pin::Pin; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Arc; +use std::task::{ready, Context, Poll}; +use std::time::Duration; + +use axum::body::Body; +use pin_project_lite::pin_project; +use tokio::sync::Notify; +use tower::Service; + +/// The requests in flight on one connection. +#[derive(Debug, Default)] +pub(super) struct Activity { + active: AtomicUsize, + /// Woken when the count leaves or reaches zero. + changed: Notify, +} + +impl Activity { + /// A request starts; it ends when the returned guard drops. + fn enter(self: &Arc) -> InFlight { + if self.active.fetch_add(1, Ordering::AcqRel) == 0 { + self.changed.notify_one(); + } + InFlight(Arc::clone(self)) + } + + /// Resolves once no request has been in flight for `timeout`. + pub(super) async fn idle_for(&self, timeout: Duration) { + loop { + if self.active.load(Ordering::Acquire) == 0 { + tokio::select! { + () = tokio::time::sleep(timeout) => { + if self.active.load(Ordering::Acquire) == 0 { + return; + } + } + // A request started (and maybe ended): measure again. + () = self.changed.notified() => {} + } + } else { + // `notify_one` keeps a permit when nobody waits, so an end + // between the load and this await is not missed. + self.changed.notified().await; + } + } + } +} + +/// One request in flight. +#[derive(Debug)] +struct InFlight(Arc); + +impl Drop for InFlight { + fn drop(&mut self) { + if self.0.active.fetch_sub(1, Ordering::AcqRel) == 1 { + self.0.changed.notify_one(); + } + } +} + +/// `inner`, counting each request in flight on `activity` until its response +/// body ends. +#[derive(Clone, Debug)] +pub(super) struct Tracked { + pub(super) inner: S, + pub(super) activity: Arc, +} + +impl Service for Tracked +where + S: Service, Error = Infallible>, +{ + type Response = http::Response; + type Error = Infallible; + type Future = TrackedFuture; + + #[inline] + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, request: R) -> Self::Future { + TrackedFuture { + in_flight: Some(self.activity.enter()), + future: self.inner.call(request), + } + } +} + +pin_project! { + /// The response future of [`Tracked`]. + pub(super) struct TrackedFuture { + #[pin] + future: F, + in_flight: Option, + } +} + +impl Future for TrackedFuture +where + F: Future, Infallible>>, +{ + type Output = Result, Infallible>; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + let this = self.project(); + let response = ready!(this.future.poll(cx))?; + Poll::Ready(Ok(match this.in_flight.take() { + // The request stays in flight until its response body ends. + Some(in_flight) => response.map(|body| crate::held::until_end(body, in_flight)), + None => response, + })) + } +} + +#[cfg(test)] +mod tests; diff --git a/src/serve/idle/tests.rs b/src/serve/idle/tests.rs new file mode 100644 index 0000000..348ca55 --- /dev/null +++ b/src/serve/idle/tests.rs @@ -0,0 +1,56 @@ +use super::*; + +use std::time::Duration; + +/// Whether `future` has resolved, without waiting. +async fn ready(future: Pin<&mut F>) -> bool { + tokio::select! { + biased; + _ = future => true, + () = std::future::ready(()) => false, + } +} + +#[tokio::test(start_paused = true)] +async fn a_connection_with_nothing_in_flight_is_idle_after_the_timeout() { + let activity = Arc::new(Activity::default()); + let idle = activity.idle_for(Duration::from_secs(1)); + tokio::pin!(idle); + assert!(!ready(idle.as_mut()).await); + tokio::time::advance(Duration::from_millis(1001)).await; + assert!(ready(idle.as_mut()).await); +} + +#[tokio::test(start_paused = true)] +async fn a_request_in_flight_keeps_its_connection() { + // A long stream must not be cut because it outlasts the idle timeout. + let activity = Arc::new(Activity::default()); + let in_flight = activity.enter(); + let idle = activity.idle_for(Duration::from_secs(1)); + tokio::pin!(idle); + tokio::time::advance(Duration::from_secs(10)).await; + assert!(!ready(idle.as_mut()).await); + + // The timeout counts from the end of the request. + drop(in_flight); + assert!(!ready(idle.as_mut()).await); + tokio::time::advance(Duration::from_millis(500)).await; + assert!(!ready(idle.as_mut()).await); + tokio::time::advance(Duration::from_millis(501)).await; + assert!(ready(idle.as_mut()).await); +} + +#[tokio::test(start_paused = true)] +async fn a_request_that_comes_and_goes_restarts_the_timeout() { + let activity = Arc::new(Activity::default()); + let idle = activity.idle_for(Duration::from_secs(1)); + tokio::pin!(idle); + assert!(!ready(idle.as_mut()).await); + tokio::time::advance(Duration::from_millis(800)).await; + drop(activity.enter()); + assert!(!ready(idle.as_mut()).await); + tokio::time::advance(Duration::from_millis(800)).await; + assert!(!ready(idle.as_mut()).await); + tokio::time::advance(Duration::from_millis(201)).await; + assert!(ready(idle.as_mut()).await); +} diff --git a/src/serve/tests.rs b/src/serve/tests.rs index 7ffd278..8c9cc0b 100644 --- a/src/serve/tests.rs +++ b/src/serve/tests.rs @@ -61,6 +61,64 @@ async fn a_connection_past_the_limit_is_served_once_another_closes() { assert_eq!(served, "HTTP/1.1 200 OK"); } +#[tokio::test] +async fn an_idle_http2_connection_gives_up_its_slot() { + // A gRPC channel keeps its HTTP/2 connection open between calls; with + // one slot, that alone would lock every other client out. + let options = ServeOptions::new() + .max_connections(1) + .idle_timeout(Some(Duration::from_millis(300))); + let addr = listen(options).await; + let channel = tonic::transport::Channel::from_shared(format!("http://{addr}")) + .unwrap() + .connect() + .await + .unwrap(); + let mut client = tonic_health::pb::health_client::HealthClient::new(channel); + // No upstream service: the call completes with UNIMPLEMENTED, and the + // connection stays open with no stream on it. + let status = client + .check(tonic_health::pb::HealthCheckRequest::default()) + .await + .unwrap_err(); + assert_eq!(status.code(), tonic::Code::Unimplemented); + + let mut other = TcpStream::connect(addr).await.unwrap(); + request(&mut other).await; + let served = tokio::time::timeout(Duration::from_secs(5), status_line(&mut other)) + .await + .expect("the idle connection's slot is freed"); + assert_eq!(served, "HTTP/1.1 200 OK"); + drop(client); +} + +#[tokio::test] +async fn a_client_that_trickles_its_headers_is_disconnected() { + // Headers that never complete would otherwise hold the connection (and + // its slot) forever. + let options = ServeOptions::new() + .idle_timeout(None) + .header_read_timeout(Duration::from_millis(300)); + let addr = listen(options).await; + let mut stream = TcpStream::connect(addr).await.unwrap(); + stream + .write_all(b"GET /health/live HTTP/1.1\r\nHost: local") + .await + .unwrap(); + let mut rest = Vec::new(); + let closed = tokio::time::timeout(Duration::from_secs(5), stream.read_to_end(&mut rest)) + .await + .expect("the server closes a connection whose headers never end"); + // End of stream, or a reset: either way the server let go. + assert!( + closed.as_ref().map_or_else( + |e| e.kind() == std::io::ErrorKind::ConnectionReset, + |_| true + ), + "{closed:?}" + ); +} + #[tokio::test] async fn without_a_limit_connections_are_served_together() { let addr = listen(ServeOptions::new()).await; @@ -78,6 +136,32 @@ fn a_limit_of_zero_is_refused() { let _ = ServeOptions::new().max_connections(0); } +#[test] +fn connection_timeouts_come_from_the_config() { + let server = crate::ProxyServer::from_yaml_str( + "listen:\n idle_timeout_secs: 0\n header_read_timeout_secs: 5\n", + ) + .unwrap(); + let options = server.serve_options().unwrap(); + // 0 keeps idle connections open. + assert_eq!(options.idle_timeout, None); + assert_eq!(options.header_read_timeout, Duration::from_secs(5)); + let defaults = crate::ProxyServer::new().serve_options().unwrap(); + assert_eq!(defaults.idle_timeout, Some(Duration::from_secs(60))); + assert_eq!(defaults.header_read_timeout, Duration::from_secs(30)); +} + +#[test] +fn a_zero_header_read_timeout_is_an_error() { + let server = + crate::ProxyServer::from_yaml_str("listen:\n header_read_timeout_secs: 0\n").unwrap(); + let err = server.serve_options().unwrap_err(); + assert!( + err.to_string().contains("header_read_timeout_secs"), + "{err}" + ); +} + #[test] fn a_zero_limit_in_the_config_is_an_error() { let server = crate::ProxyServer::from_yaml_str("listen:\n max_connections: 0\n").unwrap(); @@ -109,6 +193,24 @@ fn tls_is_set_up_from_the_config_files() { assert_eq!(tls.alpn_protocols, [b"h2".to_vec(), b"http/1.1".to_vec()]); } +#[test] +fn client_auth_without_a_client_ca_is_an_error() { + // `client_auth: required` alone would otherwise give a listener that + // verifies no client, although the config asks for mTLS. + install_provider(); + let dir = concat!(env!("CARGO_MANIFEST_DIR"), "/src/tls/testdata"); + for client_auth in ["required", "optional"] { + let yaml = format!( + "listen:\n tls:\n cert_file: {dir}/ecdsa.pem\n key_file: {dir}/ecdsa.key.pem\n client_auth: {client_auth}\n" + ); + let err = crate::ProxyServer::from_yaml_str(&yaml) + .unwrap() + .serve_options() + .unwrap_err(); + assert!(err.to_string().contains("client_ca_file"), "{err}"); + } +} + #[test] fn a_missing_tls_file_names_itself() { install_provider(); diff --git a/src/tls.rs b/src/tls.rs index d737c70..8e5abd0 100644 --- a/src/tls.rs +++ b/src/tls.rs @@ -42,14 +42,20 @@ pub(crate) fn client_config_with( /// /// # Errors /// -/// No crypto provider (see [`select_provider`]), a file that cannot be read or -/// holds no certificate / key, or a certificate rustls refuses. +/// `client_auth` without `client_ca_file`, no crypto provider (see +/// [`select_provider`]), a file that cannot be read or holds no certificate / +/// key, or a certificate rustls refuses. pub(crate) fn server_config( config: &crate::config::ListenTlsConfig, ) -> Result { use rustls::pki_types::pem::PemObject; use rustls::pki_types::{CertificateDer, PrivateKeyDer}; + // A client policy with nothing to verify against would leave the listener + // accepting every client while the config reads as mTLS. + if config.client_auth.is_some() && config.client_ca_file.is_none() { + return Err("client_auth needs client_ca_file to verify the certificates against".into()); + } let provider = select_provider(CryptoProvider::get_default(), builtin_provider)?; let certs = |path: &std::path::Path| { let certs = CertificateDer::pem_file_iter(path) @@ -77,7 +83,7 @@ pub(crate) fn server_config( } let verifier = rustls::server::WebPkiClientVerifier::builder_with_provider(roots.into(), provider); - let verifier = match config.client_auth { + let verifier = match config.client_auth.unwrap_or_default() { crate::config::ClientAuth::Required => verifier, crate::config::ClientAuth::Optional => verifier.allow_unauthenticated(), } diff --git a/tests/edge.rs b/tests/edge.rs index 3d7bcfb..e6d3dbf 100644 --- a/tests/edge.rs +++ b/tests/edge.rs @@ -567,6 +567,27 @@ fn translation_needs_the_proxys_grpc_web_cors() { assert!(err.to_string().contains("grpc_web.translate"), "{err}"); } +#[tokio::test] +async fn a_saturated_proxy_still_answers_its_health_probes() { + // A liveness probe that fails under load gets a busy pod restarted, + // which moves its load onto the others; the limit covers the traffic + // that loads the upstream, not the probes. + let addr = listen( + common::Upstream::InProcess, + "concurrency:\n max_in_flight: 1\n", + None, + ) + .await; + // `Hang` holds the only slot until its deadline. + let hang = tokio::spawn(http1_get(addr, "/v1/hang")); + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + let (status, body, _) = http1_get(addr, "/health/live").await; + assert_eq!(status, 200, "{body}"); + let (status, body, _) = http1_get(addr, "/v1/echo/a").await; + assert_eq!(status, 503, "{body}"); + hang.abort(); +} + /// Denies every request with `403`, recording the peer it saw. #[derive(Default)] struct PeerDecider { diff --git a/tests/embedded.rs b/tests/embedded.rs index 8677c27..c1d4ba8 100644 --- a/tests/embedded.rs +++ b/tests/embedded.rs @@ -26,6 +26,8 @@ fn embedded_config_is_constructible() { http: "0.0.0.0:8080".into(), max_connections: None, tls: None, + idle_timeout_secs: 60, + header_read_timeout_secs: 30, }, service: ServiceConfig { name: "embedded-test".into(), diff --git a/tests/tls.rs b/tests/tls.rs index a4b2864..e5fe9b9 100644 --- a/tests/tls.rs +++ b/tests/tls.rs @@ -505,6 +505,17 @@ async fn builtin_mtls_required_refuses_a_client_without_a_valid_certificate() { assert!(answered(addr, Identity::Client).await); } +#[tokio::test] +async fn a_client_ca_alone_requires_a_client_certificate() { + let yaml = format!( + "{} client_ca_file: {TESTDATA}/client-ca.pem\n", + tls_yaml(None) + ); + let addr = listen_builtin(&yaml).await; + assert!(!answered(addr, Identity::Anonymous).await); + assert!(answered(addr, Identity::Client).await); +} + #[tokio::test] async fn builtin_mtls_optional_serves_a_client_without_a_certificate() { let addr = listen_builtin(&tls_yaml(Some("optional"))).await; From 0fc8dd47c3b9c9febf10aad4cdaec7651ed16666 Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Mon, 28 Sep 2026 20:58:35 +0300 Subject: [PATCH 9/9] fix(edge): ready the scoped branch and keep upgraded slots - A guard narrowed by path or method called the service it did not guard without polling it first, so a fallback that needs poll_ready (a concurrency limit) panicked. The chosen branch is now called as a oneshot on its own clone, as axum's Route does. - An HTTP upgrade (a WebSocket in a fallback) hands the socket out of hyper and ends the connection future, which dropped the max_connections slot while the socket stayed open. The slot now lives with the connection's IO and is released when the socket closes. - README: the connection slot, the three timeouts and upgrades, stated precisely. Regression tests: a_narrowed_guard_readies_the_branch_it_calls, an_upgraded_connection_keeps_its_slot_until_it_closes. --- Cargo.toml | 4 ++- README.md | 23 +++++++++----- src/guard.rs | 20 ++++++------ src/guard/tests.rs | 30 ++++++++++++++++++ src/serve.rs | 76 ++++++++++++++++++++++++++++++++++++++++++---- src/serve/tests.rs | 55 +++++++++++++++++++++++++++++++++ 6 files changed, 185 insertions(+), 23 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index fed0eb7..70c6bee 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -176,8 +176,10 @@ cli = ["dep:clap", "dep:tracing-subscriber", "redis"] [dev-dependencies] # `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"] } +tower = { version = "0.5", features = ["util", "limit"] } http-body-util = "0.1" +# `hyper::upgrade::on`: a fallback that upgrades its connection (src/serve/tests.rs). +hyper = "1" # Trailer frames for the hand-written upstream response in # tests/upstream_controls.rs (tonic's server API cannot set success trailers). http-body = "1" diff --git a/README.md b/README.md index 286ad61..fdcb346 100644 --- a/README.md +++ b/README.md @@ -838,13 +838,22 @@ structured_proxy::serve_with(listener, proxy, options).await?; # } ``` -The TLS handshake runs in each connection's task under -`listen.tls.handshake_timeout_secs`, so a slow or stalled client holds up no -one else. A connection with no request in flight for `listen.idle_timeout_secs` -is closed (HTTP/2 gets a GOAWAY, and a gRPC client reconnects when it next -calls), so idle clients cannot keep `max_connections` slots. The same settings -are `ServeOptions::idle_timeout`, `header_read_timeout` and -`tls_handshake_timeout` in code. A client certificate the +Every connection takes a `max_connections` slot from the moment it is +accepted until its socket closes; past the limit, new connections wait in the +listen backlog. Three timeouts keep slots from being held for nothing: + +- the TLS handshake runs in the connection's own task, so it never stalls the + accept loop, and must finish within `listen.tls.handshake_timeout_secs`; +- an HTTP/1.1 request's headers must arrive within + `listen.header_read_timeout_secs`; +- a connection with no request in flight for `listen.idle_timeout_secs` is + closed (HTTP/2 gets a GOAWAY, and a gRPC client reconnects when it next + calls). + +A connection a fallback upgrades (a WebSocket) leaves HTTP and these +timeouts, and keeps its slot until it closes. In code the same settings are +`ServeOptions::idle_timeout`, `header_read_timeout` and +`tls_handshake_timeout`. A client certificate the listener verified reaches a tonic handler in process as `Request::peer_certs`. TLS needs a rustls crypto provider: the one a crypto backend feature brings, or the one your process installed (see [TLS crypto](#tls-crypto)). diff --git a/src/guard.rs b/src/guard.rs index 47df2fb..8a44ff7 100644 --- a/src/guard.rs +++ b/src/guard.rs @@ -20,8 +20,8 @@ use axum::{Json, Router}; use futures::future::Either; use globset::{GlobBuilder, GlobSet, GlobSetBuilder}; use http::{Method, StatusCode}; -use tower::util::BoxCloneSyncService; -use tower::{Layer, Service}; +use tower::util::{BoxCloneSyncService, Oneshot}; +use tower::{Layer, Service, ServiceExt}; use crate::config::{ScopeConfig, Traffic}; use crate::hooks::AuthDecider; @@ -288,24 +288,26 @@ struct ScopedService { impl Service for ScopedService where - G: Service, - S: Service, + G: Service + Clone, + S: Service + Clone, { type Response = Response; type Error = Infallible; - type Future = Either; + type Future = Either, Oneshot>; fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { - // Guard middleware and routes are always ready; readiness further in - // is waited for per request. + // Which branch serves a request is known only from the request, so + // readiness is waited for per request, on the chosen branch's clone + // (as axum's `Route` does): polling both here would hold a + // reservation of the other one, such as a concurrency permit. Poll::Ready(Ok(())) } fn call(&mut self, request: Request) -> Self::Future { if self.scope.matches(&request) { - Either::Left(self.guarded.call(request)) + Either::Left(self.guarded.clone().oneshot(request)) } else { - Either::Right(self.plain.call(request)) + Either::Right(self.plain.clone().oneshot(request)) } } } diff --git a/src/guard/tests.rs b/src/guard/tests.rs index 4e685e5..6b96bf3 100644 --- a/src/guard/tests.rs +++ b/src/guard/tests.rs @@ -161,6 +161,36 @@ fn paths_and_methods_narrow_the_requests_a_scope_matches() { assert!(!scope.matches(&request("POST", "/v1/users/1"))); } +#[tokio::test] +async fn a_narrowed_guard_readies_the_branch_it_calls() { + // An embedder's fallback may need `poll_ready` before `call` (a + // concurrency limit panics without it); a request outside the scope's + // paths reaches it through the plain branch. + use tower::ServiceExt; + let fallback = tower::limit::ConcurrencyLimit::new( + tower::service_fn(|_: Request| async { + Ok::<_, Infallible>(Response::new(axum::body::Body::empty())) + }), + 1, + ); + let config = ScopeConfig { + paths: vec!["/guarded".into()], + ..ScopeConfig::traffic([Traffic::Fallback]) + }; + let scoped = Scoped { + layer: tower::layer::layer_fn(|inner| inner), + scope: Scope::compile(Some(&config), &[], "maintenance", &[]).unwrap(), + } + .layer(fallback); + for path in ["/other", "/guarded", "/other"] { + let request = http::Request::get(path) + .body(axum::body::Body::empty()) + .unwrap(); + let response = scoped.clone().oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK, "{path}"); + } +} + #[test] fn http_statuses_map_back_to_their_grpc_codes() { // The inverse of google/rpc/code.proto's HTTP mapping, so a code a guard diff --git a/src/serve.rs b/src/serve.rs index 3bac898..728658b 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -1,15 +1,18 @@ //! Running a [`ProxyService`] on a TCP listener: HTTP/1.1 and HTTP/2 on one //! port, optionally behind TLS, with an optional cap on open connections. +use std::pin::Pin; use std::sync::Arc; +use std::task::{Context, Poll}; use std::time::Duration; use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer}; use hyper_util::server::conn::auto::Builder; use hyper_util::service::TowerToHyperService; -use tokio::io::{AsyncRead, AsyncWrite}; +use pin_project_lite::pin_project; +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; use tokio::net::TcpListener; -use tokio::sync::Semaphore; +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; use tonic::transport::server::Connected; use crate::service::{ConnectionInfo, ProxyService}; @@ -64,6 +67,9 @@ impl ServeOptions { /// open. A request counts until its response body ends, so a stream keeps /// its connection. Without it an idle client holds a /// [`max_connections`](Self::max_connections) slot for as long as it likes. + /// A connection upgraded to another protocol (a WebSocket in a fallback) + /// leaves HTTP and this timeout with it; it keeps its slot until it + /// closes. #[must_use] pub fn idle_timeout(mut self, timeout: Option) -> Self { self.idle_timeout = timeout; @@ -205,8 +211,6 @@ pub async fn serve_with( let service = service.clone(); let acceptor = acceptor.clone(); tokio::spawn(async move { - // Held for the connection's life. - let _slot = slot; // Small gRPC frames and REST answers are latency-bound. if let Err(error) = tcp.set_nodelay(true) { tracing::debug!(%error, "cannot set TCP_NODELAY"); @@ -214,9 +218,11 @@ pub async fn serve_with( match acceptor { None => { let service = service.for_connection(tcp.connect_info()); - serve_connection(tcp, service, limits).await; + serve_connection(SlotIo { io: tcp, slot }, service, limits).await; } Some(acceptor) => { + // The slot is held through the handshake, then by the + // stream. let stream = match tokio::time::timeout(handshake_timeout, acceptor.accept(tcp)).await { Ok(Ok(stream)) => stream, @@ -231,13 +237,71 @@ pub async fn serve_with( }; let service = service.for_connection(ConnectionInfo::tls(stream.connect_info())); - serve_connection(stream, service, limits).await; + serve_connection(SlotIo { io: stream, slot }, service, limits).await; } } }); } } +pin_project! { + /// A connection's IO holding its `max_connections` slot for as long as + /// the socket is open. hyper hands the IO of an upgraded connection (a + /// WebSocket in a fallback) to the upgrading service and finishes the + /// connection future, so the slot has to live with the IO, not the task. + struct SlotIo { + #[pin] + io: I, + slot: Option, + } +} + +impl AsyncRead for SlotIo { + #[inline] + fn poll_read( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + self.project().io.poll_read(cx, buf) + } +} + +impl AsyncWrite for SlotIo { + #[inline] + fn poll_write( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + self.project().io.poll_write(cx, buf) + } + + #[inline] + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.project().io.poll_flush(cx) + } + + #[inline] + fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.project().io.poll_shutdown(cx) + } + + #[inline] + fn poll_write_vectored( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + bufs: &[std::io::IoSlice<'_>], + ) -> Poll> { + self.project().io.poll_write_vectored(cx, bufs) + } + + #[inline] + fn is_write_vectored(&self) -> bool { + self.io.is_write_vectored() + } +} + /// The timeouts every connection is served under. #[derive(Clone, Copy, Debug)] struct ConnectionLimits { diff --git a/src/serve/tests.rs b/src/serve/tests.rs index 8c9cc0b..4c5587b 100644 --- a/src/serve/tests.rs +++ b/src/serve/tests.rs @@ -39,6 +39,61 @@ async fn status_line(stream: &mut TcpStream) -> String { head.lines().next().unwrap().to_owned() } +#[tokio::test] +async fn an_upgraded_connection_keeps_its_slot_until_it_closes() { + // A fallback that upgrades (a WebSocket) takes the socket out of hyper; + // the connection still counts against the limit while it is open. + let fallback = axum::Router::new().route( + "/upgrade", + axum::routing::get(|request: axum::extract::Request| async move { + let upgrade = hyper::upgrade::on(request); + tokio::spawn(async move { + let mut upgraded = hyper_util::rt::TokioIo::new(upgrade.await.unwrap()); + // Hold the socket until the client closes it; a reset ends it + // as well as an end of stream. + let mut rest = Vec::new(); + upgraded.read_to_end(&mut rest).await.unwrap_or_default(); + }); + http::Response::builder() + .status(http::StatusCode::SWITCHING_PROTOCOLS) + .header("connection", "upgrade") + .header("upgrade", "test") + .body(axum::body::Body::empty()) + .unwrap() + }), + ); + let service = crate::ProxyServer::from_yaml_str("service:\n name: demo\n") + .unwrap() + .service(tonic::service::Routes::default()) + .unwrap() + .with_fallback(fallback); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let options = ServeOptions::new().max_connections(1).idle_timeout(None); + tokio::spawn(serve_with(listener, service, options)); + + let mut upgraded = TcpStream::connect(addr).await.unwrap(); + upgraded + .write_all(b"GET /upgrade HTTP/1.1\r\nHost: localhost\r\nConnection: upgrade\r\nUpgrade: test\r\n\r\n") + .await + .unwrap(); + assert_eq!( + status_line(&mut upgraded).await, + "HTTP/1.1 101 Switching Protocols" + ); + + let mut other = TcpStream::connect(addr).await.unwrap(); + request(&mut other).await; + let waiting = tokio::time::timeout(Duration::from_millis(300), status_line(&mut other)).await; + assert!(waiting.is_err(), "served past the connection limit"); + + drop(upgraded); + let served = tokio::time::timeout(Duration::from_secs(5), status_line(&mut other)) + .await + .expect("closing the upgraded connection frees its slot"); + assert_eq!(served, "HTTP/1.1 200 OK"); +} + #[tokio::test] async fn a_connection_past_the_limit_is_served_once_another_closes() { let addr = listen(ServeOptions::new().max_connections(1)).await;