From 3e11a300c5e89a3b3844e1d235bba79af3c1a75e Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Sun, 27 Sep 2026 03:46:18 +0300 Subject: [PATCH 1/3] feat(transcode): upstream-controlled HTTP answers - Forward upstream response metadata as HTTP response headers: initial metadata and trailers of a successful call, the metadata of a trailers-only error, and the initial metadata of a server stream. grpc-*, -bin, content-type, hop-by-hop/framing keys and x-http-code are never forwarded; an operator deny-list (response_headers.deny, ProxyServer::with_denied_response_headers) drops more - Set the status of a successful unary call from x-http-code (200-599); an invalid value becomes the malformed-upstream INTERNAL (500) - Carry google.api.HttpBody raw bodies in both directions, including server-streaming chunks; make httpbody.proto resolvable for details - Parse HttpRule.custom: any method token, and kind "*" for every method - Answer only real CORS preflights (OPTIONS with Access-Control-Request-Method) in the CORS layer; every other OPTIONS reaches its route. The CORS layer used to answer any OPTIONS with 200, which let an unauthenticated OPTIONS pass the forward-auth endpoint - OpenAPI reads bindings from the same parser as routing (additional_bindings, custom rules, unique operationIds) and maps body/query by the body rule; HttpBody is raw content - Mount server-streaming bindings on any method and under aliases - Serialize JSON responses without an intermediate value tree Closes #92 --- Cargo.toml | 3 + README.md | 105 +++- src/auth/verifier.rs | 2 +- src/config.rs | 30 +- src/config/tests.rs | 64 +- src/cors.rs | 69 +++ src/cors/tests.rs | 88 +++ src/lib.rs | 38 +- src/openapi.rs | 416 ++++++------- src/openapi/tests.rs | 295 +++++++++ src/transcode/error.rs | 64 +- src/transcode/error/tests.rs | 74 +++ src/transcode/httpbody.rs | 131 ++++ src/transcode/httpbody/tests.rs | 212 +++++++ src/transcode/mod.rs | 1021 +++++++++++++++++++------------ src/transcode/response.rs | 191 ++++++ src/transcode/response/tests.rs | 262 ++++++++ src/transcode/rule.rs | 140 +++++ src/transcode/rule/tests.rs | 168 +++++ src/transcode/tests.rs | 88 +-- tests/common/mod.rs | 22 +- tests/hooks.rs | 45 ++ tests/upstream_controls.rs | 887 +++++++++++++++++++++++++++ 23 files changed, 3677 insertions(+), 738 deletions(-) create mode 100644 src/cors.rs create mode 100644 src/cors/tests.rs create mode 100644 src/openapi/tests.rs create mode 100644 src/transcode/httpbody.rs create mode 100644 src/transcode/httpbody/tests.rs create mode 100644 src/transcode/response.rs create mode 100644 src/transcode/response/tests.rs create mode 100644 src/transcode/rule.rs create mode 100644 src/transcode/rule/tests.rs create mode 100644 tests/upstream_controls.rs diff --git a/Cargo.toml b/Cargo.toml index 797eec1..9a06466 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -141,6 +141,9 @@ redis = ["dep:redis"] tokio = { version = "1", features = ["macros", "rt-multi-thread"] } tower = { version = "0.5", features = ["util"] } http-body-util = "0.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" # The embedding-hooks integration test (tests/hooks.rs) writes hook impls using # only these (the same crates a real embedder uses — none is an HTTP framework) # and drives the resulting router via axum + tower for assertions. diff --git a/README.md b/README.md index e8cbc3e..18279d8 100644 --- a/README.md +++ b/README.md @@ -15,6 +15,8 @@ Works with **any** gRPC service via proto descriptor files. No code generation, - **Dynamic REST routes** from proto descriptors using `google.api.http` annotations - **Full request mapping**: path params, query parameters (typed + repeated + nested), and `body` (`*` / named field / none) - **`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 @@ -126,6 +128,11 @@ error_details: - pattern: "/v1/partner/**" opaque: true +# Optional: upstream response metadata keys kept off the HTTP response (see +# "Upstream controls"). Every other application key is forwarded as a header. +response_headers: + deny: ["x-debug-trace"] + # Rate limiting (Shield) # # Every decision is made locally with a GCRA shaper (no blocking latency). @@ -387,7 +394,9 @@ The same applies to a message the proxy cannot serialize mid-stream: the stream ends with an `INTERNAL` terminal frame. 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). +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 @@ -418,6 +427,93 @@ Ok(ProxyServer::from_config(config).with_error_details(policy)) # } ``` +## Upstream controls + +HTTP protocols served as gRPC (an OAuth 2.0 / OpenID Connect provider, a +forward-auth endpoint, a file download) need more than a JSON body with `200`: +a status of their choosing, response headers, bodies that are not JSON, and +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. + +**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, 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: + +- gRPC's own keys: `grpc-*`, binary `-bin` keys and `content-type` (the proxy + sets it for the body it writes); +- hop-by-hop and framing fields, which describe the upstream connection: + `connection`, `keep-alive`, `proxy-connection`, `te`, `trailer`, + `transfer-encoding`, `upgrade`, `content-length`; +- `x-http-code` (below); +- anything the operator denies: `response_headers.deny` in the config file or + `ProxyServer::with_denied_response_headers`, e.g. to keep internal debugging + headers off a public edge. There is no allow-list: a header the upstream sets + is meant for its HTTP clients. + +A header the proxy writes for the body itself wins over the same upstream key +(an SSE stream stays `Cache-Control: no-cache`). The metadata of an error whose +details are malformed is dropped along with it (see +[Error responses](#error-responses)). Browsers read only +[CORS-safelisted](https://fetch.spec.whatwg.org/#cors-safelisted-response-header-name) +response headers plus the exposed ones, so a browser client that must read a +forwarded header needs a CORS setup that exposes it. + +**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 +`{"error": "INTERNAL", "code": 13, "message": "upstream returned a malformed response", "details": []}` +(500), with nothing else of the upstream's answer. `204` and `304` are sent +without a body or `Content-Type` (RFC 9110 §15.3.5, §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; the other fields still come from the path and +query. 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. + +An RFC 6749 token endpoint, for example: + +```proto +rpc Token(TokenRequest) returns (google.api.HttpBody) { + option (google.api.http) = { post: "/oauth2/token" body: "*" }; +} +``` + +answers a bad grant with `x-http-code: 400`, `cache-control: no-store` and an +`HttpBody` of `application/json` holding `{"error": "invalid_grant"}`; the +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: +`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 `Access-Control-Request-Method`) is +answered by the CORS layer; any other `OPTIONS` request reaches its route. + ## Library Usage ```rust @@ -426,8 +522,9 @@ use structured_proxy::ProxyServer; #[tokio::main] async fn main() -> anyhow::Result<()> { - // Reads the whole config file, including `error_details` and - // `streaming.ndjson_envelope`, which live outside `ProxyConfig`. + // Reads the whole config file, including `error_details`, + // `streaming.ndjson_envelope` and `response_headers`, which live outside + // `ProxyConfig`. let server = ProxyServer::from_file(Path::new("my-service.yaml"))?; // Run the proxy on the configured listen address. @@ -501,6 +598,8 @@ The hooks are: framework-agnostic adapter (request parts in, response parts out). - **`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 + off the HTTP responses (see [Upstream controls](#upstream-controls)). ## JWT verification diff --git a/src/auth/verifier.rs b/src/auth/verifier.rs index 2685acf..e21fc95 100644 --- a/src/auth/verifier.rs +++ b/src/auth/verifier.rs @@ -1,4 +1,4 @@ -//! The built-in [`TokenVerifier`]: keys from config, verification via +//! The built-in [`TokenVerifier`](crate::hooks::TokenVerifier): keys from config, verification via //! `jsonwebtoken`. //! //! Compiled only with the `builtin_jwt` feature (implied by `rust_crypto` / diff --git a/src/config.rs b/src/config.rs index fc3f4c0..ad9a2d3 100644 --- a/src/config.rs +++ b/src/config.rs @@ -138,6 +138,8 @@ impl Default for StreamingConfig { /// opaque: true /// streaming: /// ndjson_envelope: true +/// response_headers: +/// deny: ["x-debug-trace"] /// ``` /// /// They live outside [`ProxyConfig`] so that embedders who build it as a @@ -148,6 +150,18 @@ pub(crate) struct TranscodeFileConfig { error_details: Option, #[serde(default)] streaming: StreamingFileConfig, + #[serde(default)] + response_headers: Option, +} + +/// `response_headers:`. A typo here would silently let an internal header +/// reach clients, so unknown keys are rejected. +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] +struct ResponseHeadersFileConfig { + /// Upstream response metadata keys kept off the HTTP response. + #[serde(default)] + deny: Vec, } /// `error_details:`. A typo here would silently change what clients learn @@ -209,6 +223,7 @@ pub(crate) const KNOWN_TOP_LEVEL_KEYS: &[&str] = &[ "forwarded_headers", "streaming", "error_details", + "response_headers", ]; /// Every `streaming:` key: the [`StreamingConfig`] fields plus the ones @@ -249,11 +264,24 @@ impl TranscodeFileConfig { /// # Errors /// /// An `error_details` route pattern that is relative or not a valid glob, - /// or a route rule that sets neither `enabled` nor `opaque`. + /// a route rule that sets neither `enabled` nor `opaque`, or a + /// `response_headers.deny` entry that is not a header name. pub(crate) fn options(&self) -> Result { use crate::transcode::error::ErrorDetailsPolicy; let mut options = crate::transcode::TranscodeOptions::default() .with_ndjson_envelope(self.streaming.ndjson_envelope); + if let Some(cfg) = &self.response_headers { + let deny = cfg + .deny + .iter() + .map(|name| { + http::HeaderName::from_bytes(name.as_bytes()).map_err(|_| { + format!("response_headers.deny entry {name:?} is not a header name") + }) + }) + .collect::, _>>()?; + options = options.with_denied_response_headers(deny); + } if let Some(cfg) = &self.error_details { let base = if cfg.enabled { ErrorDetailsPolicy::default() diff --git a/src/config/tests.rs b/src/config/tests.rs index 8cdf264..41fe502 100644 --- a/src/config/tests.rs +++ b/src/config/tests.rs @@ -437,10 +437,72 @@ fn known_top_level_keys_cover_every_proxy_config_field() { "forwarded_headers", "streaming", "error_details", + "response_headers", ] { assert!(KNOWN_TOP_LEVEL_KEYS.contains(&key), "{key}"); } - assert_eq!(KNOWN_TOP_LEVEL_KEYS.len(), 18); + assert_eq!(KNOWN_TOP_LEVEL_KEYS.len(), 19); +} + +#[test] +fn response_headers_key_is_known() { + // The deny-list lives outside ProxyConfig; its key must not be reported + // as a typo, while a misspelling of it is. + let yaml = r#" +upstream: + default: "grpc://x:1" +response_headers: + deny: ["x-debug-trace"] +response_header: + deny: ["x-debug-trace"] +"#; + assert_eq!( + unknown_config_keys(yaml), + vec!["response_header".to_string()] + ); +} + +#[test] +fn transcode_settings_read_the_response_header_deny_list() { + // Names are normalized to lowercase header names, in order. + let options = transcode_options( + "upstream:\n default: \"grpc://x:1\"\nresponse_headers:\n deny: [\"X-Debug-Trace\", \"x-backend\"]\n", + ) + .unwrap(); + assert_eq!( + options + .denied_response_headers + .iter() + .map(|name| name.as_str()) + .collect::>(), + ["x-debug-trace", "x-backend"] + ); +} + +#[test] +fn transcode_settings_without_response_headers_deny_nothing() { + let options = transcode_options("upstream:\n default: \"grpc://x:1\"\n").unwrap(); + assert!(options.denied_response_headers.is_empty()); +} + +#[test] +fn transcode_settings_reject_an_invalid_deny_entry() { + // A space is not allowed in a header name; the entry could never match. + let err = transcode_options( + "upstream:\n default: \"grpc://x:1\"\nresponse_headers:\n deny: [\"x debug\"]\n", + ) + .unwrap_err(); + assert!(err.contains("not a header name"), "{err}"); +} + +#[test] +fn transcode_settings_reject_unknown_response_headers_key() { + // `denied` for `deny` would otherwise let the header through silently. + let err = transcode_options( + "upstream:\n default: \"grpc://x:1\"\nresponse_headers:\n denied: [\"x-debug-trace\"]\n", + ) + .unwrap_err(); + assert!(err.contains("denied"), "{err}"); } #[test] diff --git a/src/cors.rs b/src/cors.rs new file mode 100644 index 0000000..13bf480 --- /dev/null +++ b/src/cors.rs @@ -0,0 +1,69 @@ +//! CORS around the router that answers only real preflight requests. +//! +//! tower-http's `CorsLayer` answers every `OPTIONS` request as a preflight, so +//! no route could serve `OPTIONS`: not a `custom` rule, not a `*` rule, not the +//! forward-auth endpoint, whose sub-request carries the original request's +//! method and would be allowed through by that `200`. The Fetch standard (§3.2.2, +//! CORS-preflight request) makes a preflight an `OPTIONS` request with +//! `Access-Control-Request-Method`; any other `OPTIONS` is an ordinary request. Such a request passes the CORS layer under a stand-in +//! method, so it gets the response headers of an ordinary CORS request, and has +//! its method restored before anything else sees it. + +use axum::extract::Request; +use axum::http::header::ACCESS_CONTROL_REQUEST_METHOD; +use axum::http::Method; +use axum::middleware::{self, Next}; +use axum::response::Response; +use axum::Router; +use tower_http::cors::CorsLayer; + +/// Marks a request whose method the outer layer replaced with [`stand_in`]. +#[derive(Clone, Copy)] +struct OrdinaryOptions; + +/// The method an ordinary `OPTIONS` request carries through the CORS layer +/// (short enough for `http` to store inline, without allocating). A client +/// sending it itself gets no special treatment: without the marker it is left +/// as it is and matches no route. +fn stand_in() -> Method { + Method::from_bytes(b"X-SP-OPTIONS").expect("a valid method token") +} + +/// `router` wrapped in `cors`, with only real preflights answered by it. +pub(crate) fn layer(router: Router, cors: CorsLayer) -> Router +where + S: Clone + Send + Sync + 'static, +{ + router + .layer(middleware::from_fn(restore_options)) + .layer(cors) + .layer(middleware::from_fn(disguise_options)) +} + +/// Outermost: hide an ordinary `OPTIONS` from the CORS layer. +async fn disguise_options(mut request: Request, next: Next) -> Response { + if request.method() == Method::OPTIONS + && !request + .headers() + .contains_key(ACCESS_CONTROL_REQUEST_METHOD) + { + *request.method_mut() = stand_in(); + request.extensions_mut().insert(OrdinaryOptions); + } + next.run(request).await +} + +/// Right inside the CORS layer: give the request its method back. +async fn restore_options(mut request: Request, next: Next) -> Response { + if request + .extensions_mut() + .remove::() + .is_some() + { + *request.method_mut() = Method::OPTIONS; + } + next.run(request).await +} + +#[cfg(test)] +mod tests; diff --git a/src/cors/tests.rs b/src/cors/tests.rs new file mode 100644 index 0000000..22698b9 --- /dev/null +++ b/src/cors/tests.rs @@ -0,0 +1,88 @@ +use super::*; +use axum::body::Body; +use axum::http::StatusCode; +use axum::routing::any; +use tower::ServiceExt; + +/// A router whose only route answers every method with the method it saw. +fn app() -> Router { + let router = Router::new().route( + "/x", + any(|method: Method| async move { method.as_str().to_owned() }), + ); + layer(router, CorsLayer::permissive()) +} + +async fn send(request: Request) -> (StatusCode, axum::http::HeaderMap, String) { + let response = app().oneshot(request).await.unwrap(); + let (parts, body) = response.into_parts(); + let body = axum::body::to_bytes(body, usize::MAX).await.unwrap(); + ( + parts.status, + parts.headers, + String::from_utf8(body.to_vec()).unwrap(), + ) +} + +fn options(headers: &[(&str, &str)]) -> Request { + let mut builder = Request::builder().method(Method::OPTIONS).uri("/x"); + for (name, value) in headers { + builder = builder.header(*name, *value); + } + builder.body(Body::empty()).unwrap() +} + +#[tokio::test] +async fn ordinary_options_reaches_the_route_as_options() { + let (status, headers, body) = send(options(&[("origin", "https://a.example")])).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(body, "OPTIONS"); + // It is still a CORS response, with the headers of an ordinary request. + assert_eq!(headers["access-control-allow-origin"], "*"); + assert!(!headers.contains_key("access-control-allow-methods")); +} + +#[tokio::test] +async fn options_without_origin_reaches_the_route() { + let (status, _, body) = send(options(&[])).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(body, "OPTIONS"); +} + +#[tokio::test] +async fn preflight_is_answered_by_the_cors_layer() { + let (status, headers, body) = send(options(&[ + ("origin", "https://a.example"), + ("access-control-request-method", "PUT"), + ])) + .await; + assert_eq!(status, StatusCode::OK); + assert!(body.is_empty(), "{body}"); + assert!(headers.contains_key("access-control-allow-methods")); +} + +#[tokio::test] +async fn other_methods_are_untouched() { + let request = Request::builder() + .method(Method::DELETE) + .uri("/x") + .header("origin", "https://a.example") + .body(Body::empty()) + .unwrap(); + let (status, headers, body) = send(request).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(body, "DELETE"); + assert_eq!(headers["access-control-allow-origin"], "*"); +} + +#[tokio::test] +async fn stand_in_method_sent_by_a_client_is_not_turned_into_options() { + let request = Request::builder() + .method(stand_in()) + .uri("/x") + .body(Body::empty()) + .unwrap(); + let (status, _, body) = send(request).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(body, "X-SP-OPTIONS"); +} diff --git a/src/lib.rs b/src/lib.rs index d7d02d9..77ab14d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -45,6 +45,7 @@ compile_error!( pub mod auth; pub mod config; +mod cors; mod embed; pub mod hooks; pub mod oidc; @@ -141,19 +142,22 @@ impl ProxyServer { } /// Create from a YAML document: the [`ProxyConfig`] plus the transcoding - /// settings it does not hold (`error_details:` and - /// `streaming.ndjson_envelope`), applied as [`with_error_details`] and - /// [`with_ndjson_envelope`] would. A top-level or `streaming:` key no - /// setting reads is logged as a warning. + /// settings it does not hold (`error_details:`, + /// `streaming.ndjson_envelope` and `response_headers:`), applied as + /// [`with_error_details`], [`with_ndjson_envelope`] and + /// [`with_denied_response_headers`] would. A top-level or `streaming:` key + /// no setting reads is logged as a warning. /// /// # Errors /// /// Invalid YAML, a [`ProxyConfig`] that fails - /// [`validate`](ProxyConfig::validate), or an `error_details` route pattern - /// that is relative or not a valid glob. + /// [`validate`](ProxyConfig::validate), an `error_details` route pattern + /// that is relative or not a valid glob, or a `response_headers.deny` + /// entry that is not a header name. /// /// [`with_error_details`]: Self::with_error_details /// [`with_ndjson_envelope`]: Self::with_ndjson_envelope + /// [`with_denied_response_headers`]: Self::with_denied_response_headers pub fn from_yaml_str(yaml: &str) -> anyhow::Result { let config = ProxyConfig::from_yaml_str(yaml)?; for key in config::unknown_config_keys(yaml) { @@ -162,7 +166,7 @@ impl ProxyServer { let settings: config::TranscodeFileConfig = serde_yaml::from_str(yaml)?; let options = settings .options() - .map_err(|e| anyhow::anyhow!("invalid error_details config: {e}"))?; + .map_err(|e| anyhow::anyhow!("invalid transcoding config: {e}"))?; let mut server = Self::from_config(config); server.transcode = options; Ok(server) @@ -262,6 +266,17 @@ impl ProxyServer { self } + /// Keep these upstream response metadata keys off the HTTP responses of + /// the transcoded routes; see + /// [`transcode::TranscodeOptions::with_denied_response_headers`]. + pub fn with_denied_response_headers( + mut self, + names: impl IntoIterator, + ) -> Self { + self.transcode = self.transcode.with_denied_response_headers(names); + self + } + /// Load descriptor pool from configured sources. /// /// Multiple descriptor files are merged into a single pool, @@ -692,11 +707,10 @@ impl ProxyServer { state.clone(), maintenance_middleware, )) - .layer(TraceLayer::new_for_http()) - // Outermost: wraps every enforcement layer so short-circuited - // responses keep CORS headers, and answers preflight before auth. - .layer(cors) - .with_state(state); + .layer(TraceLayer::new_for_http()); + // Outermost: wraps every enforcement layer so short-circuited + // responses keep CORS headers, and answers preflight before auth. + let router = cors::layer(router, cors).with_state(state); Ok(router) } diff --git a/src/openapi.rs b/src/openapi.rs index c10e50d..745615b 100644 --- a/src/openapi.rs +++ b/src/openapi.rs @@ -4,10 +4,22 @@ //! to produce a complete OpenAPI 3.0 JSON spec at runtime. //! No codegen, no build step — same descriptor pool used for transcoding. +use std::collections::HashSet; + +use axum::http::Method; use prost_reflect::{DescriptorPool, FieldDescriptor, Kind, MessageDescriptor, MethodDescriptor}; use serde_json::{json, Map, Value}; use crate::config::{AliasConfig, OpenApiConfig}; +use crate::transcode::httpbody; +use crate::transcode::request::BodyMapping; +use crate::transcode::rule::{self, HttpBinding, RouteMethod}; + +/// The operations an OpenAPI 3.0 path item can hold, in the order a `*` rule +/// lists them. +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 { @@ -17,6 +29,8 @@ pub fn generate(pool: &DescriptorPool, config: &OpenApiConfig, aliases: &[AliasC let mut paths = Map::new(); let mut schemas = Map::new(); let mut tags = Vec::new(); + let mut operation_ids = HashSet::new(); + let http_ext = pool.get_extension_by_name("google.api.http"); for service in pool.services() { let service_name = service.name().to_string(); @@ -30,38 +44,41 @@ pub fn generate(pool: &DescriptorPool, config: &OpenApiConfig, aliases: &[AliasC } tags.push(tag); + let Some(http_ext) = &http_ext else { + continue; + }; for method in service.methods() { if method.is_client_streaming() { continue; // No REST mapping for client-streaming. } - if let Some((http_method, http_path)) = extract_http_rule(&method, pool) { - let operation = build_operation( - &method, - &service_name, - &http_method, - &http_path, - pool, - &mut schemas, - ); - - // Main path. - add_path_operation(&mut paths, &http_path, &http_method, operation.clone()); - - // Aliases. - for alias in aliases { - if let Some(suffix) = http_path.strip_prefix(&alias.to) { - if alias.from.ends_with("/{path}") { - let prefix = alias.from.trim_end_matches("/{path}"); - let alias_path = format!("{}{}", prefix, suffix); - add_path_operation( - &mut paths, - &alias_path, - &http_method, - operation.clone(), - ); + for binding in rule::http_bindings(&method, http_ext) { + let operations = operation_methods(&binding.method); + let expanded = operations.len() > 1; + if operations.is_empty() { + continue; + } + let operation = build_operation(&method, &service_name, &binding, &mut schemas); + for http_method in operations { + let mut id = format!("{service_name}.{}", method.name()); + if expanded { + id = format!("{id}_{http_method}"); + } + + let mut targets = vec![binding.path.clone()]; + for alias in aliases { + if let Some(suffix) = binding.path.strip_prefix(&alias.to) { + if alias.from.ends_with("/{path}") { + let prefix = alias.from.trim_end_matches("/{path}"); + targets.push(format!("{prefix}{suffix}")); + } } } + for path in targets { + let mut operation = operation.clone(); + operation["operationId"] = json!(unique_id(&mut operation_ids, &id)); + add_path_operation(&mut paths, &path, http_method, operation); + } } } } @@ -115,6 +132,40 @@ pub fn docs_html(openapi_path: &str, title: &str) -> String { ) } +/// The OpenAPI operations a binding appears under: every one for a `*` rule, +/// the method's own for a method OpenAPI 3.0 has an operation for, and none for +/// any other method token (OpenAPI 3.0 cannot describe it). +fn operation_methods(method: &RouteMethod) -> Vec<&'static str> { + let RouteMethod::One(method) = method else { + return OPERATIONS.to_vec(); + }; + let operation = match *method { + Method::GET => "get", + Method::PUT => "put", + Method::POST => "post", + Method::DELETE => "delete", + Method::OPTIONS => "options", + Method::HEAD => "head", + Method::PATCH => "patch", + Method::TRACE => "trace", + _ => return Vec::new(), + }; + vec![operation] +} + +/// `base`, or `base_2`, `base_3`, ... when taken: OpenAPI requires every +/// `operationId` to be unique, and one RPC can be mounted several times +/// (additional bindings, aliases, a `*` rule). +fn unique_id(ids: &mut HashSet, base: &str) -> String { + let mut id = base.to_owned(); + let mut n = 1; + while !ids.insert(id.clone()) { + n += 1; + id = format!("{base}_{n}"); + } + id +} + fn add_path_operation(paths: &mut Map, path: &str, method: &str, operation: Value) { let path_item = paths.entry(path.to_string()).or_insert_with(|| json!({})); if let Some(obj) = path_item.as_object_mut() { @@ -122,12 +173,16 @@ fn add_path_operation(paths: &mut Map, path: &str, method: &str, } } +/// Content of a raw `google.api.HttpBody`: any media type, as bytes. +fn raw_content() -> Value { + json!({ "*/*": { "schema": { "type": "string", "format": "binary" } } }) +} + +/// The operation for one binding of `method`, without its `operationId`. fn build_operation( method: &MethodDescriptor, service_name: &str, - http_method: &str, - http_path: &str, - pool: &DescriptorPool, + binding: &HttpBinding, schemas: &mut Map, ) -> Value { let method_name = method.name().to_string(); @@ -135,15 +190,10 @@ fn build_operation( let input = method.input(); let output = method.output(); - let is_streaming = method.is_server_streaming(); - // Description from proto comments. - let description = get_comments(&full_name, pool).unwrap_or_default(); - - let operation_id = format!("{}.{}", service_name, method_name); + let description = get_comments(&full_name, method.parent_pool()).unwrap_or_default(); let mut op = json!({ - "operationId": operation_id, "tags": [service_name], "summary": method_name, }); @@ -153,82 +203,100 @@ fn build_operation( } // Path parameters. - let path_params = extract_path_params(http_path); - if !path_params.is_empty() { - let params: Vec = path_params - .iter() - .map(|name| { - let mut param = json!({ - "name": name, - "in": "path", - "required": true, - "schema": { "type": "string" }, - }); - - // Try to get type from input message field. - if let Some(field) = input.get_field_by_name(name) { - param["schema"] = field_to_schema(&field); - } - - param - }) - .collect(); - op["parameters"] = json!(params); - } - - // Request body (for POST/PUT/PATCH/DELETE with body fields). - if http_method != "get" { - let has_body_fields = input - .fields() - .any(|f| !path_params.contains(&f.name().to_string())); - - if has_body_fields { - let schema_name = input.name().to_string(); - let body_schema = message_to_schema(&input, &path_params, schemas); - - schemas.insert(schema_name.clone(), body_schema); - - op["requestBody"] = json!({ + let path_params = extract_path_params(&binding.path); + let mut params: Vec = path_params + .iter() + .map(|name| { + let mut param = json!({ + "name": name, + "in": "path", "required": true, - "content": { - "application/json": { - "schema": { - "$ref": format!("#/components/schemas/{}", schema_name), + "schema": { "type": "string" }, + }); + // Try to get type from input message field. + if let Some(field) = input.get_field_by_name(name) { + param["schema"] = field_to_schema(&field); + } + param + }) + .collect(); + + // The body rule decides which fields travel in the body; every field + // bound by neither the path nor the body is a query parameter. + match &binding.body { + BodyMapping::None => {} + BodyMapping::Root => { + let has_body_fields = input + .fields() + .any(|f| !path_params.contains(&f.name().to_string())); + if httpbody::is_http_body(&input) { + op["requestBody"] = json!({ "required": true, "content": raw_content() }); + } else if has_body_fields { + let schema_name = input.name().to_string(); + let body_schema = message_to_schema(&input, &path_params, schemas); + schemas.insert(schema_name.clone(), body_schema); + op["requestBody"] = json!({ + "required": true, + "content": { + "application/json": { + "schema": { "$ref": format!("#/components/schemas/{}", schema_name) }, }, }, - }, - }); + }); + } } - } else { - // GET: non-path fields become query parameters. - let query_params: Vec = input - .fields() - .filter(|f| !path_params.contains(&f.name().to_string())) - .map(|field| { - json!({ - "name": field.name(), - "in": "query", - "required": false, - "schema": field_to_schema(&field), - }) - }) - .collect(); - - if !query_params.is_empty() { - let existing = op - .get("parameters") - .and_then(|v| v.as_array()) - .cloned() - .unwrap_or_default(); - let mut all_params = existing; - all_params.extend(query_params); - op["parameters"] = json!(all_params); + BodyMapping::Field(name) => { + if let Some(field) = input.get_field_by_name(name) { + let content = if httpbody::http_body_field(&input, name).is_some() { + raw_content() + } else { + register_nested(&field, schemas); + json!({ "application/json": { "schema": field_to_schema(&field) } }) + }; + op["requestBody"] = json!({ "required": true, "content": content }); + } } } + let body_field = match &binding.body { + BodyMapping::Field(name) => Some(name.as_str()), + BodyMapping::None | BodyMapping::Root => None, + }; + // With `body: "*"` the whole message is the body: no query parameters. + if binding.body != BodyMapping::Root { + params.extend( + input + .fields() + .filter(|f| { + !path_params.contains(&f.name().to_string()) && body_field != Some(f.name()) + }) + .map(|field| { + json!({ + "name": field.name(), + "in": "query", + "required": false, + "schema": field_to_schema(&field), + }) + }), + ); + } + if !params.is_empty() { + op["parameters"] = json!(params); + } // Response. - if is_streaming { - op["responses"] = json!({ + let raw_response = match &binding.response_body { + None => httpbody::is_http_body(&output), + Some(path) => httpbody::http_body_path(&output, path).is_some(), + }; + op["responses"] = if raw_response { + let description = if method.is_server_streaming() { + "Server-streaming raw body (the data of every message, concatenated)" + } else { + "Success" + }; + json!({ "200": { "description": description, "content": raw_content() } }) + } else if method.is_server_streaming() { + json!({ "200": { "description": "Server-streaming response (NDJSON)", "content": { @@ -237,31 +305,24 @@ fn build_operation( }, }, }, - }); + }) } else if output.full_name() == "google.protobuf.Empty" { - op["responses"] = json!({ - "200": { - "description": "Success (empty response)", - }, - }); + json!({ "200": { "description": "Success (empty response)" } }) } else { let schema_name = output.name().to_string(); let response_schema = message_to_schema(&output, &[], schemas); schemas.insert(schema_name.clone(), response_schema); - - op["responses"] = json!({ + json!({ "200": { "description": "Success", "content": { "application/json": { - "schema": { - "$ref": format!("#/components/schemas/{}", schema_name), - }, + "schema": { "$ref": format!("#/components/schemas/{}", schema_name) }, }, }, }, - }); - } + }) + }; // Common error responses. if let Some(responses) = op.get_mut("responses").and_then(|r| r.as_object_mut()) { @@ -287,6 +348,16 @@ fn build_operation( op } +/// Register the schema of a message-typed field so its `$ref` resolves. +fn register_nested(field: &FieldDescriptor, schemas: &mut Map) { + if let Kind::Message(nested) = field.kind() { + if !is_well_known(&nested) && !schemas.contains_key(nested.name()) { + let nested_schema = message_to_schema(&nested, &[], schemas); + schemas.insert(nested.name().to_string(), nested_schema); + } + } +} + /// Generate a JSON Schema for a protobuf message, excluding path parameter fields. fn message_to_schema( msg: &MessageDescriptor, @@ -320,12 +391,7 @@ fn message_to_schema( if exclude_fields.contains(&field.name().to_string()) { continue; } - if let Kind::Message(nested) = field.kind() { - if !is_well_known(&nested) && !schemas.contains_key(nested.name()) { - let nested_schema = message_to_schema(&nested, &[], schemas); - schemas.insert(nested.name().to_string(), nested_schema); - } - } + register_nested(&field, schemas); } schema @@ -428,7 +494,8 @@ fn well_known_schema(msg: &MessageDescriptor) -> Value { } } -/// Extract `{param}` names from a path like `/v1/profiles/{profile_id}/devices`. +/// Extract the field names of the `{param}` / `{param=template}` captures of a +/// path like `/v1/profiles/{profile_id}/devices`. fn extract_path_params(path: &str) -> Vec { let mut params = Vec::new(); let mut in_brace = false; @@ -442,8 +509,10 @@ fn extract_path_params(path: &str) -> Vec { } '}' => { in_brace = false; - if !current.is_empty() { - params.push(current.clone()); + // `{name=shelves/*}` binds the field `name`. + let name = current.split('=').next().unwrap_or_default(); + if !name.is_empty() { + params.push(name.to_string()); } } _ if in_brace => current.push(ch), @@ -454,37 +523,6 @@ fn extract_path_params(path: &str) -> Vec { params } -/// Extract HTTP method and path from google.api.http annotation. -fn extract_http_rule(method: &MethodDescriptor, pool: &DescriptorPool) -> Option<(String, String)> { - let http_ext = pool.get_extension_by_name("google.api.http")?; - let options = method.options(); - - if !options.has_extension(&http_ext) { - return None; - } - - let http_rule = options.get_extension(&http_ext); - if let prost_reflect::Value::Message(rule_msg) = http_rule.into_owned() { - for (method_name, _) in [ - ("get", "get"), - ("post", "post"), - ("put", "put"), - ("delete", "delete"), - ("patch", "patch"), - ] { - if let Some(val) = rule_msg.get_field_by_name(method_name) { - if let prost_reflect::Value::String(path) = val.into_owned() { - if !path.is_empty() { - return Some((method_name.to_string(), path)); - } - } - } - } - } - - None -} - /// Get proto source comments for a given fully-qualified name. fn get_comments(_full_name: &str, _pool: &DescriptorPool) -> Option { // prost-reflect doesn't expose source code info comments easily. @@ -494,70 +532,4 @@ fn get_comments(_full_name: &str, _pool: &DescriptorPool) -> Option { } #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_extract_path_params() { - assert_eq!( - extract_path_params("/v1/profiles/{profile_id}"), - vec!["profile_id"] - ); - assert_eq!( - extract_path_params("/v1/profiles/{profile_id}/devices/{device_id}"), - vec!["profile_id", "device_id"] - ); - assert!(extract_path_params("/v1/auth/login").is_empty()); - } - - #[test] - fn test_docs_html_contains_scalar() { - let html = docs_html("/openapi.json", "Test API"); - assert!(html.contains("@scalar/api-reference")); - assert!(html.contains("/openapi.json")); - assert!(html.contains("Test API")); - } - - #[test] - fn test_well_known_schemas() { - // Verify well-known type mappings are correct. - let pool = DescriptorPool::global(); - if let Some(ts) = pool.get_message_by_name("google.protobuf.Timestamp") { - let schema = well_known_schema(&ts); - assert_eq!(schema["type"], "string"); - assert_eq!(schema["format"], "date-time"); - } - } - - #[test] - fn test_generate_empty_pool() { - let pool = DescriptorPool::new(); - let config = OpenApiConfig { - enabled: true, - path: "/openapi.json".into(), - docs_path: "/docs".into(), - title: Some("Test API".into()), - version: Some("0.1.0".into()), - }; - let spec = generate(&pool, &config, &[]); - - assert_eq!(spec["openapi"], "3.0.3"); - assert_eq!(spec["info"]["title"], "Test API"); - assert_eq!(spec["info"]["version"], "0.1.0"); - assert!(spec["paths"].as_object().unwrap().is_empty()); - } - - #[test] - fn test_field_to_schema_primitives() { - // Test via JSON output structure. - let schema = json!({ "type": "string" }); - assert_eq!(schema["type"], "string"); - - let int_schema = json!({ "type": "integer", "format": "int32" }); - assert_eq!(int_schema["format"], "int32"); - - let i64_schema = json!({ "type": "string", "format": "int64", "description": "64-bit integer (string-encoded)" }); - assert_eq!(i64_schema["type"], "string"); - assert_eq!(i64_schema["format"], "int64"); - } -} +mod tests; diff --git a/src/openapi/tests.rs b/src/openapi/tests.rs new file mode 100644 index 0000000..1ebc56e --- /dev/null +++ b/src/openapi/tests.rs @@ -0,0 +1,295 @@ +use super::*; + +#[test] +fn test_extract_path_params() { + assert_eq!( + extract_path_params("/v1/profiles/{profile_id}"), + vec!["profile_id"] + ); + assert_eq!( + extract_path_params("/v1/profiles/{profile_id}/devices/{device_id}"), + vec!["profile_id", "device_id"] + ); + assert!(extract_path_params("/v1/auth/login").is_empty()); +} + +#[test] +fn path_param_with_a_field_template_names_the_field() { + // `{name=shelves/*}` binds the field `name`, not a field called + // `name=shelves/*`. + assert_eq!(extract_path_params("/v1/{name=shelves/*}"), vec!["name"]); + assert_eq!(extract_path_params("/v1/files/{path=**}"), vec!["path"]); +} + +#[test] +fn test_docs_html_contains_scalar() { + let html = docs_html("/openapi.json", "Test API"); + assert!(html.contains("@scalar/api-reference")); + assert!(html.contains("/openapi.json")); + assert!(html.contains("Test API")); +} + +#[test] +fn test_well_known_schemas() { + // Verify well-known type mappings are correct. + let pool = DescriptorPool::global(); + if let Some(ts) = pool.get_message_by_name("google.protobuf.Timestamp") { + let schema = well_known_schema(&ts); + assert_eq!(schema["type"], "string"); + assert_eq!(schema["format"], "date-time"); + } +} + +fn config() -> OpenApiConfig { + OpenApiConfig { + enabled: true, + path: "/openapi.json".into(), + docs_path: "/docs".into(), + title: Some("Test API".into()), + version: Some("0.1.0".into()), + } +} + +#[test] +fn test_generate_empty_pool() { + let pool = DescriptorPool::new(); + let spec = generate(&pool, &config(), &[]); + + assert_eq!(spec["openapi"], "3.0.3"); + assert_eq!(spec["info"]["title"], "Test API"); + assert_eq!(spec["info"]["version"], "0.1.0"); + assert!(spec["paths"].as_object().unwrap().is_empty()); +} + +#[test] +fn test_field_to_schema_primitives() { + // Test via JSON output structure. + let schema = json!({ "type": "string" }); + assert_eq!(schema["type"], "string"); + + let int_schema = json!({ "type": "integer", "format": "int32" }); + assert_eq!(int_schema["format"], "int32"); + + let i64_schema = json!({ "type": "string", "format": "int64", "description": "64-bit integer (string-encoded)" }); + assert_eq!(i64_schema["type"], "string"); + assert_eq!(i64_schema["format"], "int64"); +} + +// --- specs generated from annotated descriptors -------------------------------- + +const HTTP_PROTO: &str = r#" +syntax = "proto3"; +package google.api; +message HttpRule { + string selector = 1; + oneof pattern { + string get = 2; + string put = 3; + string post = 4; + string delete = 5; + string patch = 6; + CustomHttpPattern custom = 8; + } + string body = 7; + string response_body = 12; + repeated HttpRule additional_bindings = 11; +} +message CustomHttpPattern { + string kind = 1; + string path = 2; +} +"#; + +const ANNOTATIONS_PROTO: &str = r#" +syntax = "proto3"; +package google.api; +import "google/api/http.proto"; +import "google/protobuf/descriptor.proto"; +extend google.protobuf.MethodOptions { + HttpRule http = 72295728; +} +"#; + +const HTTPBODY_PROTO: &str = r#" +syntax = "proto3"; +package google.api; +import "google/protobuf/any.proto"; +message HttpBody { + string content_type = 1; + bytes data = 2; + repeated google.protobuf.Any extensions = 3; +} +"#; + +const API_PROTO: &str = r#" +syntax = "proto3"; +package test.v1; +import "google/api/annotations.proto"; +import "google/api/httpbody.proto"; + +message Item { + string name = 1; + int32 count = 2; + Note note = 3; +} +message Note { + string text = 1; +} +message Upload { + string name = 1; + google.api.HttpBody file = 2; +} + +service Api { + rpc Head(Item) returns (Item) { + option (google.api.http) = { custom: { kind: "HEAD", path: "/v1/items/{name}" } }; + } + rpc Verify(Item) returns (Item) { + option (google.api.http) = { custom: { kind: "*", path: "/v1/verify" } }; + } + rpc Propfind(Item) returns (Item) { + option (google.api.http) = { custom: { kind: "PROPFIND", path: "/v1/dav" } }; + } + rpc Create(Item) returns (Item) { + option (google.api.http) = { + post: "/v1/items" + body: "*" + additional_bindings { put: "/v1/items/{name}" body: "note" } + additional_bindings { post: "/v1/items:touch" } + }; + } + rpc Put(Upload) returns (google.api.HttpBody) { + option (google.api.http) = { put: "/v1/uploads/{name}" body: "file" }; + } + rpc Raw(google.api.HttpBody) returns (google.api.HttpBody) { + option (google.api.http) = { post: "/v1/raw" body: "*" }; + } +} +"#; + +struct TestProtos; + +impl protox::file::FileResolver for TestProtos { + fn open_file(&self, name: &str) -> Result { + let source = match name { + "google/api/http.proto" => HTTP_PROTO, + "google/api/annotations.proto" => ANNOTATIONS_PROTO, + "google/api/httpbody.proto" => HTTPBODY_PROTO, + "test/v1/api.proto" => API_PROTO, + _ => return protox::file::GoogleFileResolver::new().open_file(name), + }; + protox::file::File::from_source(name, source) + } +} + +fn spec(aliases: &[AliasConfig]) -> Value { + let pool = protox::Compiler::with_file_resolver(TestProtos) + .open_file("test/v1/api.proto") + .unwrap() + .descriptor_pool(); + generate(&pool, &config(), aliases) +} + +#[test] +fn custom_head_rule_is_a_head_operation() { + let spec = spec(&[]); + let item = &spec["paths"]["/v1/items/{name}"]; + assert_eq!(item["head"]["operationId"], "Api.Head"); + assert_eq!(item["head"]["parameters"][0]["in"], "path"); +} + +#[test] +fn star_rule_is_listed_under_every_method_with_distinct_ids() { + let spec = spec(&[]); + let verify = spec["paths"]["/v1/verify"].as_object().unwrap(); + let mut methods: Vec<&str> = verify.keys().map(String::as_str).collect(); + methods.sort_unstable(); + let mut expected = OPERATIONS.to_vec(); + expected.sort_unstable(); + assert_eq!(methods, expected); + let ids: HashSet<&str> = verify + .values() + .map(|op| op["operationId"].as_str().unwrap()) + .collect(); + assert_eq!(ids.len(), OPERATIONS.len()); + assert!(ids.contains("Api.Verify_head")); +} + +#[test] +fn extension_method_has_no_openapi_operation() { + // OpenAPI 3.0 has no slot for PROPFIND. + assert!(spec(&[])["paths"].get("/v1/dav").is_none()); +} + +#[test] +fn additional_bindings_are_listed() { + let spec = spec(&[]); + assert_eq!( + spec["paths"]["/v1/items"]["post"]["operationId"], + "Api.Create" + ); + assert_eq!( + spec["paths"]["/v1/items/{name}"]["put"]["operationId"], + "Api.Create_2" + ); + assert_eq!( + spec["paths"]["/v1/items:touch"]["post"]["operationId"], + "Api.Create_3" + ); +} + +#[test] +fn body_rule_decides_body_and_query_fields() { + let spec = spec(&[]); + // `body: "*"`: the whole message is the body, nothing in the query. + let create = &spec["paths"]["/v1/items"]["post"]; + assert_eq!( + create["requestBody"]["content"]["application/json"]["schema"]["$ref"], + "#/components/schemas/Item" + ); + assert!(create.get("parameters").is_none(), "{create}"); + + // `body: "note"`: that field is the body, the rest is path or query. + let put = &spec["paths"]["/v1/items/{name}"]["put"]; + assert_eq!( + put["requestBody"]["content"]["application/json"]["schema"]["$ref"], + "#/components/schemas/Note" + ); + let params: Vec<(&str, &str)> = put["parameters"] + .as_array() + .unwrap() + .iter() + .map(|p| (p["name"].as_str().unwrap(), p["in"].as_str().unwrap())) + .collect(); + assert_eq!(params, [("name", "path"), ("count", "query")]); + + // No body rule on a POST: every field is a query parameter. + let touch = &spec["paths"]["/v1/items:touch"]["post"]; + assert!(touch.get("requestBody").is_none(), "{touch}"); + assert_eq!(touch["parameters"].as_array().unwrap().len(), 3); +} + +#[test] +fn http_body_request_and_response_are_raw_content() { + let spec = spec(&[]); + let raw = json!({ "*/*": { "schema": { "type": "string", "format": "binary" } } }); + let put = &spec["paths"]["/v1/uploads/{name}"]["put"]; + assert_eq!(put["requestBody"]["content"], raw); + assert_eq!(put["responses"]["200"]["content"], raw); + let root = &spec["paths"]["/v1/raw"]["post"]; + assert_eq!(root["requestBody"]["content"], raw); + assert_eq!(root["responses"]["200"]["content"], raw); +} + +#[test] +fn aliases_get_their_own_operation_ids() { + let alias: AliasConfig = serde_yaml::from_str("from: /api/{path}\nto: /v1").unwrap(); + let spec = spec(&[alias]); + assert_eq!( + spec["paths"]["/v1/items"]["post"]["operationId"], + "Api.Create" + ); + let aliased = &spec["paths"]["/api/items"]["post"]["operationId"]; + assert!(aliased.is_string()); + assert_ne!(aliased, "Api.Create"); +} diff --git a/src/transcode/error.rs b/src/transcode/error.rs index cb10d9b..3fb8bf9 100644 --- a/src/transcode/error.rs +++ b/src/transcode/error.rs @@ -76,8 +76,40 @@ pub fn status_to_response_with_details( status: &tonic::Status, details: Option<&StatusDetails>, ) -> Response { - let (code, body) = render(status, details); - (grpc_to_http_status(code), Json(body)).into_response() + render_response(status, details).0 +} + +/// [`status_to_response_with_details`], plus whether the response reports the +/// upstream's own error: `false` when its details could not be rendered +/// faithfully and the response is the generic `INTERNAL` instead, which then +/// carries nothing else from the upstream either. +pub(crate) fn render_response( + status: &tonic::Status, + details: Option<&StatusDetails>, +) -> (Response, bool) { + let (code, body, faithful) = render(status, details); + ( + (grpc_to_http_status(code), Json(body)).into_response(), + faithful, + ) +} + +/// Message of the `INTERNAL` a client gets instead of a successful upstream +/// answer the proxy cannot turn into a faithful HTTP response (an invalid +/// `x-http-code`, an `HttpBody` content type that is not a header value). +const MALFORMED_RESPONSE_MESSAGE: &str = "upstream returned a malformed response"; + +/// The `INTERNAL` (500) answering a successful upstream call whose response +/// cannot be passed on faithfully, in the route's error body. Like a malformed +/// error status, nothing of the upstream's answer reaches the client; the +/// caller logs the cause. +pub(crate) fn malformed_response(details: Option<&StatusDetails>) -> Response { + let body = body( + tonic::Code::Internal, + MALFORMED_RESPONSE_MESSAGE, + details.map(|_| RenderedDetails::default()), + ); + (StatusCode::INTERNAL_SERVER_ERROR, Json(body)).into_response() } /// The JSON error body for a failed call, shared by the unary response and the @@ -95,16 +127,22 @@ pub fn error_body(status: &tonic::Status, details: Option<&StatusDetails>) -> Va render(status, details).1 } -/// The error body and the gRPC code it reports: the upstream's own, or -/// `INTERNAL` when its details cannot be rendered faithfully. -fn render(status: &tonic::Status, details: Option<&StatusDetails>) -> (tonic::Code, Value) { +/// The error body, the gRPC code it reports and whether that is the +/// upstream's own error: `false` when its details cannot be rendered +/// faithfully and the body is the generic `INTERNAL`. +fn render(status: &tonic::Status, details: Option<&StatusDetails>) -> (tonic::Code, Value, bool) { let Some(details) = details else { - return (status.code(), body(status.code(), status.message(), None)); + return ( + status.code(), + body(status.code(), status.message(), None), + true, + ); }; match details.render(status) { Ok(rendered) => ( status.code(), body(status.code(), status.message(), Some(rendered)), + true, ), Err(MalformedStatus) => ( tonic::Code::Internal, @@ -113,6 +151,7 @@ fn render(status: &tonic::Status, details: Option<&StatusDetails>) -> (tonic::Co MALFORMED_STATUS_MESSAGE, Some(RenderedDetails::default()), ), + false, ), } } @@ -526,6 +565,15 @@ impl StatusDetails { canonical .decode_file_descriptor_set(tonic_types::pb::FILE_DESCRIPTOR_SET) .expect("tonic-types ships a valid google.rpc descriptor set"); + // The global pool is process-wide; another crate may have added it. + if canonical + .get_message_by_name(super::httpbody::HTTP_BODY) + .is_none() + { + canonical + .add_file_descriptor_proto(super::httpbody::file_descriptor()) + .expect("google/api/httpbody.proto is a valid descriptor"); + } let mut pool = product.clone(); complete_with_canonical(&mut pool, &canonical); Self { @@ -535,7 +583,9 @@ impl StatusDetails { } /// Whether a detail whose type no descriptor describes goes to - /// `opaqueDetails` (see [`opaque_entry`]) instead of being withheld. Its + /// `opaqueDetails` (`{"index", "typeUrl", "bytes"}` entries: its position + /// among the forwarded details, its type URL and the standard base64 of + /// its bytes) instead of being withheld. Its /// bytes cannot be inspected, so a DebugInfo it carries in a field would /// reach the client: switch it on only for upstreams trusted not to nest /// one in types the proxy has no descriptor for. diff --git a/src/transcode/error/tests.rs b/src/transcode/error/tests.rs index b82b0a0..465509b 100644 --- a/src/transcode/error/tests.rs +++ b/src/transcode/error/tests.rs @@ -1149,3 +1149,77 @@ fn policy_rejects_invalid_glob() { let err = policy(true, &[("/v1/[admin", false)]).unwrap_err(); assert!(err.contains("invalid glob pattern"), "{err}"); } + +// --- google.api.HttpBody, malformed responses --------------------------------- + +#[test] +fn http_body_detail_renders_without_product_descriptors() { + // google/api/httpbody.proto is canonical, like the google.rpc types: a + // detail of that type renders even when the product pool lacks it. + // content_type = "text/plain" (field 1), data = "hi" (field 2). + let http_body = b"\x0a\x0atext/plain\x12\x02hi".to_vec(); + let status = status_with_raw_details(&[("type.googleapis.com/google.api.HttpBody", http_body)]); + let body = error_body(&status, Some(&canonical_only())); + assert_eq!( + body["details"], + json!([{ + "@type": "type.googleapis.com/google.api.HttpBody", + "contentType": "text/plain", + "data": "aGk=" + }]) + ); +} + +#[test] +fn render_response_reports_whether_the_error_is_the_upstreams() { + let (response, faithful) = render_response(&rich_status(), Some(&canonical_only())); + assert!(faithful); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + + // Without details the trailer is not read, so the error is always faithful. + let broken = tonic::Status::with_details( + tonic::Code::NotFound, + "gone", + bytes::Bytes::from_static(b"\xff\xff"), + ); + let (response, faithful) = render_response(&broken, None); + assert!(faithful); + assert_eq!(response.status(), StatusCode::NOT_FOUND); + + // With details on, a trailer that is not a google.rpc.Status is replaced. + let (response, faithful) = render_response(&broken, Some(&canonical_only())); + assert!(!faithful); + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); +} + +async fn json_of(response: Response) -> serde_json::Value { + let bytes = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + serde_json::from_slice(&bytes).unwrap() +} + +#[tokio::test] +async fn malformed_response_is_an_internal_in_the_route_error_body() { + let response = malformed_response(Some(&canonical_only())); + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + assert_eq!( + json_of(response).await, + json!({ + "error": "INTERNAL", + "message": "upstream returned a malformed response", + "code": 13, + "details": [] + }) + ); + // With details off for the route, the key is absent as on every error. + let response = malformed_response(None); + assert_eq!( + json_of(response).await, + json!({ + "error": "INTERNAL", + "message": "upstream returned a malformed response", + "code": 13 + }) + ); +} diff --git a/src/transcode/httpbody.rs b/src/transcode/httpbody.rs new file mode 100644 index 0000000..7adf1e0 --- /dev/null +++ b/src/transcode/httpbody.rs @@ -0,0 +1,131 @@ +//! `google.api.HttpBody`: an RPC that takes or returns one carries the raw HTTP +//! body and its `Content-Type` instead of JSON, as in `google/api/httpbody.proto`. + +use bytes::Bytes; +use prost_reflect::{DynamicMessage, FieldDescriptor, Kind, MessageDescriptor, Value}; + +/// Full name of `google.api.HttpBody`. +pub(crate) const HTTP_BODY: &str = "google.api.HttpBody"; + +/// Whether `desc` is `google.api.HttpBody` with the `content_type` (string) and +/// `data` (bytes) fields the transcoder reads and writes. A product revision of +/// the type that lacks either is transcoded as an ordinary message. +pub(crate) fn is_http_body(desc: &MessageDescriptor) -> bool { + let field_is = |name: &str, kind: Kind| { + desc.get_field_by_name(name) + .is_some_and(|field| field.kind() == kind && !field.is_list()) + }; + desc.full_name() == HTTP_BODY + && field_is("content_type", Kind::String) + && field_is("data", Kind::Bytes) +} + +/// The singular message field `name` of `desc`, and its type, when that type +/// is an HttpBody. +pub(crate) fn http_body_field( + desc: &MessageDescriptor, + name: &str, +) -> Option<(FieldDescriptor, MessageDescriptor)> { + let field = desc.get_field_by_name(name)?; + match field.kind() { + Kind::Message(inner) if !field.is_list() && is_http_body(&inner) => Some((field, inner)), + _ => None, + } +} + +/// The chain of singular message fields a dotted `response_body` path names in +/// `desc`, when it ends at an HttpBody. +pub(crate) fn http_body_path(desc: &MessageDescriptor, path: &str) -> Option> { + let mut fields = Vec::new(); + let mut current = desc.clone(); + for segment in path.split('.') { + let field = current.get_field_by_name(segment)?; + let Kind::Message(inner) = field.kind() else { + return None; + }; + if field.is_list() { + return None; + } + fields.push(field); + current = inner; + } + is_http_body(¤t).then_some(fields) +} + +/// The raw body an HttpBody carries. +#[derive(Debug, Default)] +pub(crate) struct RawBody { + /// `content_type`; empty when the upstream left it unset. + pub(crate) content_type: String, + pub(crate) data: Bytes, +} + +/// Move the HttpBody at `path` (the message itself when empty) out of `msg`. +/// An unset field on the way is an empty body, as ProtoJSON renders an unset +/// message field as absent. +pub(crate) fn take(mut msg: DynamicMessage, path: &[FieldDescriptor]) -> RawBody { + for field in path { + match msg.take_field(field) { + Some(Value::Message(inner)) => msg = inner, + _ => return RawBody::default(), + } + } + let content_type = match msg.take_field_by_name("content_type") { + Some(Value::String(content_type)) => content_type, + _ => String::new(), + }; + let data = match msg.take_field_by_name("data") { + Some(Value::Bytes(data)) => data, + _ => Bytes::new(), + }; + RawBody { content_type, data } +} + +/// Set the `content_type` and `data` of the HttpBody `msg`. +pub(crate) fn fill(msg: &mut DynamicMessage, content_type: String, data: Bytes) { + msg.set_field_by_name("content_type", Value::String(content_type)); + msg.set_field_by_name("data", Value::Bytes(data)); +} + +/// `google/api/httpbody.proto`, so error details of this type render even when +/// the product descriptors do not import it. +pub(crate) fn file_descriptor() -> prost_reflect::prost_types::FileDescriptorProto { + use prost_reflect::prost_types::field_descriptor_proto::{Label, Type}; + use prost_reflect::prost_types::{DescriptorProto, FieldDescriptorProto, FileDescriptorProto}; + + let field = |name: &str, number: i32, label: Label, ty: Type, type_name: Option<&str>| { + FieldDescriptorProto { + name: Some(name.to_owned()), + number: Some(number), + label: Some(label as i32), + r#type: Some(ty as i32), + type_name: type_name.map(str::to_owned), + ..Default::default() + } + }; + FileDescriptorProto { + name: Some("google/api/httpbody.proto".to_owned()), + package: Some("google.api".to_owned()), + dependency: vec!["google/protobuf/any.proto".to_owned()], + message_type: vec![DescriptorProto { + name: Some("HttpBody".to_owned()), + field: vec![ + field("content_type", 1, Label::Optional, Type::String, None), + field("data", 2, Label::Optional, Type::Bytes, None), + field( + "extensions", + 3, + Label::Repeated, + Type::Message, + Some(".google.protobuf.Any"), + ), + ], + ..Default::default() + }], + syntax: Some("proto3".to_owned()), + ..Default::default() + } +} + +#[cfg(test)] +mod tests; diff --git a/src/transcode/httpbody/tests.rs b/src/transcode/httpbody/tests.rs new file mode 100644 index 0000000..709d057 --- /dev/null +++ b/src/transcode/httpbody/tests.rs @@ -0,0 +1,212 @@ +use super::*; +use prost_reflect::DescriptorPool; + +/// A pool with the canonical HttpBody and `test.Wrapper { HttpBody body = 1; +/// repeated HttpBody many = 2; string name = 3; Inner inner = 4; }`, +/// `test.Inner { HttpBody body = 1; }`, and a product type named like HttpBody +/// but missing `data`. +fn pool() -> DescriptorPool { + use prost_reflect::prost_types::field_descriptor_proto::{Label, Type}; + use prost_reflect::prost_types::{DescriptorProto, FieldDescriptorProto, FileDescriptorProto}; + + let field = |name: &str, number: i32, label: Label, ty: Type, type_name: Option<&str>| { + FieldDescriptorProto { + name: Some(name.to_owned()), + number: Some(number), + label: Some(label as i32), + r#type: Some(ty as i32), + type_name: type_name.map(str::to_owned), + ..Default::default() + } + }; + let mut pool = DescriptorPool::global(); + pool.add_file_descriptor_proto(file_descriptor()).unwrap(); + pool.add_file_descriptor_proto(FileDescriptorProto { + name: Some("test/wrapper.proto".to_owned()), + package: Some("test".to_owned()), + dependency: vec!["google/api/httpbody.proto".to_owned()], + message_type: vec![ + DescriptorProto { + name: Some("Wrapper".to_owned()), + field: vec![ + field( + "body", + 1, + Label::Optional, + Type::Message, + Some(".google.api.HttpBody"), + ), + field( + "many", + 2, + Label::Repeated, + Type::Message, + Some(".google.api.HttpBody"), + ), + field("name", 3, Label::Optional, Type::String, None), + field( + "inner", + 4, + Label::Optional, + Type::Message, + Some(".test.Inner"), + ), + ], + ..Default::default() + }, + DescriptorProto { + name: Some("Inner".to_owned()), + field: vec![field( + "body", + 1, + Label::Optional, + Type::Message, + Some(".google.api.HttpBody"), + )], + ..Default::default() + }, + ], + syntax: Some("proto3".to_owned()), + ..Default::default() + }) + .unwrap(); + pool +} + +fn message(pool: &DescriptorPool, name: &str) -> MessageDescriptor { + pool.get_message_by_name(name).unwrap() +} + +#[test] +fn canonical_http_body_is_recognized() { + let pool = pool(); + assert!(is_http_body(&message(&pool, HTTP_BODY))); + assert!(!is_http_body(&message(&pool, "test.Wrapper"))); +} + +#[test] +fn http_body_named_type_without_its_fields_is_an_ordinary_message() { + // A product revision of google.api.HttpBody that lacks `data` cannot be + // filled or read as a raw body, so it is transcoded as JSON instead of + // panicking on the missing field. + use prost_reflect::prost_types::field_descriptor_proto::{Label, Type}; + use prost_reflect::prost_types::{DescriptorProto, FieldDescriptorProto, FileDescriptorProto}; + let mut pool = DescriptorPool::new(); + pool.add_file_descriptor_proto(FileDescriptorProto { + name: Some("google/api/httpbody.proto".to_owned()), + package: Some("google.api".to_owned()), + message_type: vec![DescriptorProto { + name: Some("HttpBody".to_owned()), + field: vec![FieldDescriptorProto { + name: Some("content_type".to_owned()), + number: Some(1), + label: Some(Label::Optional as i32), + r#type: Some(Type::String as i32), + ..Default::default() + }], + ..Default::default() + }], + syntax: Some("proto3".to_owned()), + ..Default::default() + }) + .unwrap(); + assert!(!is_http_body(&message(&pool, HTTP_BODY))); +} + +#[test] +fn http_body_field_finds_singular_http_body_fields_only() { + let pool = pool(); + let wrapper = message(&pool, "test.Wrapper"); + let (field, desc) = http_body_field(&wrapper, "body").unwrap(); + assert_eq!(field.name(), "body"); + assert_eq!(desc.full_name(), HTTP_BODY); + // Repeated, scalar and missing fields are not an HttpBody body target. + assert!(http_body_field(&wrapper, "many").is_none()); + assert!(http_body_field(&wrapper, "name").is_none()); + assert!(http_body_field(&wrapper, "missing").is_none()); +} + +#[test] +fn http_body_path_resolves_nested_fields() { + let pool = pool(); + let wrapper = message(&pool, "test.Wrapper"); + let path = http_body_path(&wrapper, "inner.body").unwrap(); + assert_eq!( + path.iter().map(|f| f.name()).collect::>(), + ["inner", "body"] + ); + assert_eq!(http_body_path(&wrapper, "body").unwrap().len(), 1); + // Paths ending elsewhere, through a scalar, a repeated field or a missing + // name are not HttpBody responses. + assert!(http_body_path(&wrapper, "inner").is_none()); + assert!(http_body_path(&wrapper, "name").is_none()); + assert!(http_body_path(&wrapper, "many").is_none()); + assert!(http_body_path(&wrapper, "inner.missing").is_none()); +} + +#[test] +fn fill_then_take_round_trips_without_copying_the_data() { + let pool = pool(); + let mut msg = DynamicMessage::new(message(&pool, HTTP_BODY)); + let data = Bytes::from_static(b"\x89PNG raw"); + fill(&mut msg, "image/png".to_owned(), data.clone()); + let raw = take(msg, &[]); + assert_eq!(raw.content_type, "image/png"); + assert_eq!(raw.data, data); + // Same allocation: the bytes were moved, not copied. + assert_eq!(raw.data.as_ptr(), data.as_ptr()); +} + +#[test] +fn take_walks_a_field_path() { + let pool = pool(); + let mut body = DynamicMessage::new(message(&pool, HTTP_BODY)); + fill( + &mut body, + "text/plain".to_owned(), + Bytes::from_static(b"hi"), + ); + let mut inner = DynamicMessage::new(message(&pool, "test.Inner")); + inner.set_field_by_name("body", Value::Message(body)); + let wrapper_desc = message(&pool, "test.Wrapper"); + let mut wrapper = DynamicMessage::new(wrapper_desc.clone()); + wrapper.set_field_by_name("inner", Value::Message(inner)); + let raw = take( + wrapper, + &http_body_path(&wrapper_desc, "inner.body").unwrap(), + ); + assert_eq!(raw.content_type, "text/plain"); + assert_eq!(&raw.data[..], b"hi"); +} + +#[test] +fn take_of_an_unset_field_is_an_empty_body() { + let pool = pool(); + let wrapper_desc = message(&pool, "test.Wrapper"); + let raw = take( + DynamicMessage::new(wrapper_desc.clone()), + &http_body_path(&wrapper_desc, "inner.body").unwrap(), + ); + assert!(raw.content_type.is_empty()); + assert!(raw.data.is_empty()); +} + +#[test] +fn canonical_descriptor_matches_google_api_httpbody() { + let pool = pool(); + let desc = message(&pool, HTTP_BODY); + let fields: Vec<(u32, String, bool)> = desc + .fields() + .map(|f| (f.number(), f.name().to_owned(), f.is_list())) + .collect(); + assert_eq!( + fields, + [ + (1, "content_type".to_owned(), false), + (2, "data".to_owned(), false), + (3, "extensions".to_owned(), true) + ] + ); + assert_eq!(desc.get_field(2).unwrap().json_name(), "data"); + assert_eq!(desc.get_field(1).unwrap().json_name(), "contentType"); +} diff --git a/src/transcode/mod.rs b/src/transcode/mod.rs index af02c70..c69267b 100644 --- a/src/transcode/mod.rs +++ b/src/transcode/mod.rs @@ -2,28 +2,44 @@ //! //! Reads `google.api.http` annotations from proto service descriptors //! and builds axum routes that proxy JSON/form requests to gRPC upstream. +//! The upstream decides the HTTP answer beyond the JSON body where it needs +//! to: its response metadata becomes response headers, `x-http-code` sets the +//! status of a successful unary call, and `google.api.HttpBody` carries a raw +//! body in either direction. //! //! Generic: works with ANY proto descriptor set. No product-specific code. pub mod body; pub mod codec; pub mod error; +pub(crate) mod httpbody; pub mod metadata; pub mod request; +pub(crate) mod response; +pub(crate) mod rule; +use axum::body::{Body, Bytes}; use axum::extract::{Path, RawQuery, State}; -use axum::http::{HeaderMap, StatusCode}; +use axum::http::header::{ALLOW, CONTENT_TYPE}; +use axum::http::{HeaderMap, HeaderName, HeaderValue, Method, StatusCode}; use axum::response::sse::{Event, KeepAlive, Sse}; use axum::response::{IntoResponse, Response}; -use axum::routing::{delete, get, patch, post, put, MethodRouter}; -use axum::{Json, Router}; -use futures::StreamExt; -use prost_reflect::{DescriptorPool, DynamicMessage, MethodDescriptor, SerializeOptions}; +use axum::routing::{MethodFilter, MethodRouter}; +use axum::Router; +use futures::{StreamExt, TryStreamExt}; +use prost_reflect::{ + DescriptorPool, DynamicMessage, FieldDescriptor, MessageDescriptor, MethodDescriptor, + SerializeOptions, +}; +use std::collections::HashMap; use std::sync::Arc; use tonic::client::Grpc; +use tonic::metadata::MetadataMap; use crate::config::AliasConfig; use error::{ErrorDetailsPolicy, StatusDetails}; +use response::UpstreamHeaders; +use rule::RouteMethod; /// Trait for state types that support REST→gRPC transcoding. /// @@ -50,22 +66,49 @@ impl TranscodeState for crate::ProxyState { } } +/// Path parameters of a matched route. +type PathParams = HashMap; + +/// How the HTTP request body reaches the RPC's input message. +#[derive(Debug, Clone)] +enum RequestBody { + /// JSON or a form, mapped as the rule's `body` says. + Parsed(request::BodyMapping), + /// Raw bytes and `Content-Type` into the input message, a `google.api.HttpBody`. + RawRoot, + /// Raw bytes and `Content-Type` into the HttpBody field (of that type) `body` + /// names; the other fields still come from path and query. + RawField(FieldDescriptor, MessageDescriptor), +} + +/// What the HTTP response body is made of. +#[derive(Debug, Clone)] +enum ResponseShape { + /// ProtoJSON of the response message, or of its `response_body` subfield. + Json(Option), + /// The raw body of a `google.api.HttpBody`: the response message itself + /// (empty path) or the field chain `response_body` names. + HttpBody(Vec), +} + /// Route entry extracted from proto HTTP annotations. #[derive(Debug, Clone)] struct RouteEntry { /// HTTP path pattern (e.g., "/v1/auth/opaque/login/start"). http_path: String, - /// HTTP method (GET, POST, PUT, PATCH, DELETE). - http_method: HttpMethod, + /// The HTTP method(s) the binding answers. + http_method: RouteMethod, /// gRPC path (e.g., "/sid.v1.AuthService/OpaqueLoginStart"), parsed once at /// route-build time so each request clones a cheap `Bytes` refcount. grpc_path: axum::http::uri::PathAndQuery, /// Method descriptor for input/output message resolution. method: MethodDescriptor, + /// Server-streaming RPC (NDJSON / SSE, or chunked HttpBody). + streaming: bool, /// How the request body maps onto the gRPC request message. - body: request::BodyMapping, - /// Optional response subfield to return as the HTTP body (`response_body`). - response_body: Option, + request_body: RequestBody, + /// What the HTTP response body is made of. + response: ResponseShape, /// Renderer for the status details of this route's errors; `None` when the /// error-details policy switches them off for the route. error_details: Option>, @@ -74,6 +117,23 @@ struct RouteEntry { /// [`codec::has_required_fields`] of the response type, computed once so a /// request does not walk the descriptor. response_has_required: bool, + /// Response metadata keys the operator keeps off the HTTP response, on top + /// of the ones that never go there. + denied_headers: Arc<[HeaderName]>, +} + +impl RouteEntry { + fn codec(&self) -> codec::DynamicCodec { + codec::DynamicCodec::with_required_check(self.method.output(), self.response_has_required) + } + + /// The raw body of `message`, on a route that answers with an HttpBody. + fn raw_body(&self, message: DynamicMessage) -> Option { + match &self.response { + ResponseShape::HttpBody(path) => Some(httpbody::take(message, path)), + ResponseShape::Json(_) => None, + } + } } /// How [`routes_with_options`] builds the transcoded routes. @@ -81,18 +141,21 @@ struct RouteEntry { /// # Examples /// /// ``` +/// use axum::http::HeaderName; /// use structured_proxy::transcode::error::ErrorDetailsPolicy; /// use structured_proxy::transcode::TranscodeOptions; /// /// let options = TranscodeOptions::default() /// .with_error_details(ErrorDetailsPolicy::default().route("/v1/admin/**", false).unwrap()) -/// .with_ndjson_envelope(true); +/// .with_ndjson_envelope(true) +/// .with_denied_response_headers([HeaderName::from_static("x-debug-trace")]); /// # let _ = options; /// ``` #[derive(Debug, Clone, Default)] pub struct TranscodeOptions { pub(crate) error_details: ErrorDetailsPolicy, pub(crate) ndjson_envelope: bool, + pub(crate) denied_response_headers: Arc<[HeaderName]>, } impl TranscodeOptions { @@ -114,27 +177,18 @@ impl TranscodeOptions { self.ndjson_envelope = enabled; self } -} - -#[derive(Debug, Clone, Copy)] -enum HttpMethod { - Get, - Post, - Put, - Patch, - Delete, -} -impl HttpMethod { - /// The uppercase HTTP method token (e.g. `"GET"`). - fn as_str(self) -> &'static str { - match self { - HttpMethod::Get => "GET", - HttpMethod::Post => "POST", - HttpMethod::Put => "PUT", - HttpMethod::Patch => "PATCH", - HttpMethod::Delete => "DELETE", - } + /// Response metadata keys that never become HTTP response headers, on top + /// of gRPC's own keys, hop-by-hop fields and `x-http-code`, which never do. + /// Replaces any list set before. Use it to keep internal headers (debug + /// traces, backend names) off a public edge; an upstream key not listed + /// here reaches the client. + pub fn with_denied_response_headers( + mut self, + names: impl IntoIterator, + ) -> Self { + self.denied_response_headers = names.into_iter().collect(); + self } } @@ -149,6 +203,10 @@ pub fn routes(pool: &DescriptorPool, aliases: &[AliasConfig]) } /// [`routes`], built as `options` describe. +/// +/// A second binding for a method and path already taken (or any binding on a +/// path a `custom` `*` rule takes, which answers every method) is skipped with +/// an error in the log. pub fn routes_with_options( pool: &DescriptorPool, aliases: &[AliasConfig], @@ -166,7 +224,10 @@ pub fn routes_with_options( // each shared by every route that uses it and built only when one does. // The second is a copy of the first: the descriptor pool inside is shared. let mut status_details: [Option>; 2] = [None, None]; - let mut router: Router = Router::new(); + // Every binding of one path goes into the same method router, in binding + // order. + let mut paths: Vec = Vec::new(); + let mut path_index: HashMap = HashMap::new(); for mut binding in bindings { let policy = &options.error_details; if policy.enabled_for(&binding.axum_path) { @@ -184,64 +245,131 @@ pub fn routes_with_options( binding.entry.error_details = status_details[slot].clone(); } binding.entry.ndjson_envelope = options.ndjson_envelope; - let method = binding.entry.http_method; - let entry = Arc::new(binding.entry); - let method_router: MethodRouter = if binding.streaming { - let handler = move |proxy_state: State, - headers: HeaderMap, - path_params: Path>, - raw_query: RawQuery, - body: axum::body::Bytes| { - streaming_handler(proxy_state, headers, path_params, raw_query, body, entry) - }; - match method { - HttpMethod::Get => get(handler), - HttpMethod::Post => post(handler), - // route_bindings only yields GET/POST streaming bindings. - _ => unreachable!("streaming routes are GET/POST only"), - } - } else { - let handler = move |proxy_state: State, - headers: HeaderMap, - path_params: Path>, - raw_query: RawQuery, - body: axum::body::Bytes| { - transcode_handler(proxy_state, headers, path_params, raw_query, body, entry) - }; - match method { - HttpMethod::Get => get(handler), - HttpMethod::Post => post(handler), - HttpMethod::Put => put(handler), - HttpMethod::Patch => patch(handler), - HttpMethod::Delete => delete(handler), + binding.entry.denied_headers = options.denied_response_headers.clone(); + let index = match path_index.get(&binding.axum_path) { + Some(&index) => index, + None => { + path_index.insert(binding.axum_path.clone(), paths.len()); + paths.push(PathRoutes { + path: binding.axum_path, + methods: Vec::new(), + }); + paths.len() - 1 } }; - router = router.route(&binding.axum_path, method_router); + paths[index].add(Arc::new(binding.entry)); } + let mut router: Router = Router::new(); + for path in &paths { + router = router.route(&path.path, path.method_router()); + } router } -/// One transcode route to mount: the RPC entry that serves it, the axum path to -/// register it at, and whether it is the server-streaming variant. +/// The axum handler serving the route entry `$entry`, for the state type `S` +/// in scope. A macro because the closure's handler type cannot be named. +macro_rules! endpoint { + ($entry:expr) => {{ + let entry: Arc = $entry; + move |state: State, + headers: HeaderMap, + path_params: Path, + raw_query: RawQuery, + body: Bytes| handle(state, headers, path_params, raw_query, body, entry) + }}; +} + +/// The bindings mounted at one axum path. +struct PathRoutes { + path: String, + methods: Vec>, +} + +impl PathRoutes { + /// Add `entry` unless its method is already answered on this path. + fn add(&mut self, entry: Arc) { + let taken = self.methods.iter().any(|existing| { + existing.http_method == entry.http_method + || existing.http_method == RouteMethod::Any + || entry.http_method == RouteMethod::Any + }); + if taken { + tracing::error!( + method = entry.http_method.as_str(), + path = %self.path, + rpc = %entry.grpc_path, + "HTTP method and path already bound to another RPC; skipping this binding" + ); + return; + } + self.methods.push(entry); + } + + /// One method router for every binding of the path. Methods axum routes by + /// itself are registered directly; any other token (a `custom` rule such + /// as `PROPFIND`) is dispatched by a fallback that answers `405` with the + /// full `Allow` list (RFC 9110 §15.5.6) for a method nobody binds. + fn method_router(&self) -> MethodRouter { + let mut router = MethodRouter::new(); + let mut extension: Vec<(Method, Arc)> = Vec::new(); + let mut allow: Vec<&str> = Vec::new(); + for entry in &self.methods { + match &entry.http_method { + // `add` keeps a `*` binding alone on its path. + RouteMethod::Any => return axum::routing::any(endpoint!(entry.clone())), + RouteMethod::One(method) => { + allow.push(method.as_str()); + match MethodFilter::try_from(method.clone()) { + Ok(filter) => router = router.on(filter, endpoint!(entry.clone())), + Err(_) => extension.push((method.clone(), entry.clone())), + } + } + } + } + if extension.is_empty() { + return router; + } + // A GET route answers HEAD too. + if allow.contains(&"GET") && !allow.contains(&"HEAD") { + allow.push("HEAD"); + } + let allow = HeaderValue::from_str(&allow.join(", ")) + .expect("method tokens are valid header value characters"); + let extension: Arc<[(Method, Arc)]> = extension.into(); + router.fallback( + move |method: Method, + state: State, + headers: HeaderMap, + path_params: Path, + raw_query: RawQuery, + body: Bytes| async move { + match extension.iter().find(|(bound, _)| *bound == method) { + Some((_, entry)) => { + handle(state, headers, path_params, raw_query, body, entry.clone()).await + } + None => (StatusCode::METHOD_NOT_ALLOWED, [(ALLOW, allow)]).into_response(), + } + }, + ) + } +} + +/// One transcode route to mount: the RPC entry that serves it and the axum path +/// to register it at. struct RouteBinding { entry: RouteEntry, axum_path: String, - streaming: bool, } -/// The single source of truth for what [`routes`] mounts: unary RPCs, their -/// config aliases, and server-streaming RPCs. Both [`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. +/// The single source of truth for what [`routes`] mounts: every binding of +/// every unary and server-streaming RPC, plus its config aliases. Both +/// [`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 { let mut bindings = Vec::new(); for entry in extract_routes(pool) { - bindings.push(RouteBinding { - axum_path: proto_path_to_axum(&entry.http_path), - entry: entry.clone(), - streaming: false, - }); for alias in aliases { if let Some(suffix) = entry.http_path.strip_prefix(&alias.to) { if alias.from.ends_with("/{path}") { @@ -249,32 +377,28 @@ fn route_bindings(pool: &DescriptorPool, aliases: &[AliasConfig]) -> Vec Vec<(String, String)> { route_bindings(pool, aliases) .into_iter() @@ -290,12 +414,50 @@ fn response_serialize_options() -> SerializeOptions { .stringify_64_bit_integers(true) } +/// Serialize a message straight to JSON bytes, without an intermediate tree. +fn message_to_json_bytes( + msg: &DynamicMessage, + opts: &SerializeOptions, +) -> Result, serde_json::Error> { + let mut buf = Vec::with_capacity(128); + msg.serialize_with_options(&mut serde_json::Serializer::new(&mut buf), opts)?; + Ok(buf) +} + /// Serialize one streamed gRPC message to a compact JSON string. fn message_to_json_string(msg: &DynamicMessage, opts: &SerializeOptions) -> Result { - let value = msg - .serialize_with_options(serde_json::value::Serializer, opts) - .map_err(|e| e.to_string())?; - serde_json::to_string(&value).map_err(|e| e.to_string()) + let buf = message_to_json_bytes(msg, opts).map_err(|e| e.to_string())?; + // SAFETY: serde_json's serializer writes only valid UTF-8, the same + // guarantee `serde_json::to_string` relies on. + Ok(unsafe { String::from_utf8_unchecked(buf) }) +} + +/// The unary response body as JSON: the whole message, or the subfield +/// `response_body` names (JSON `null` when the path does not exist). +fn json_body( + msg: &DynamicMessage, + response_body: Option<&str>, +) -> Result, serde_json::Error> { + let opts = response_serialize_options(); + let Some(path) = response_body else { + return message_to_json_bytes(msg, &opts); + }; + // Walk the tree by moving each subtree out, so nothing is copied. + let mut value = Some(msg.serialize_with_options(serde_json::value::Serializer, &opts)?); + for segment in path.split('.') { + value = match value { + Some(serde_json::Value::Object(mut fields)) => fields.remove(segment), + _ => None, + }; + } + let value = value.unwrap_or_else(|| { + tracing::warn!( + response_body = %path, + "configured response_body path not found in response; returning null" + ); + serde_json::Value::Null + }); + serde_json::to_vec(&value) } /// Whether the client negotiated a Server-Sent Events response via `Accept`. @@ -334,79 +496,299 @@ fn accept_range_selects_sse(range: &str) -> bool { true } -/// Handler for server-streaming RPCs. -/// -/// Returns Server-Sent Events when the client sends `Accept: text/event-stream`, -/// otherwise newline-delimited JSON (NDJSON). In both formats a gRPC error -/// mid-stream is delivered as an explicit terminal frame before the stream is -/// closed cleanly, rather than truncating the HTTP body. -async fn streaming_handler( +/// Serve one request on a transcoded route. +async fn handle( State(proxy_state): State, headers: HeaderMap, - Path(path_params): Path>, + Path(path_params): Path, RawQuery(raw_query): RawQuery, - body_bytes: axum::body::Bytes, - entry: std::sync::Arc, + body: Bytes, + entry: Arc, ) -> Response { - let channel = proxy_state.grpc_channel(); - - let request_msg = match decode_request( - &entry, + let prepared = prepare( + &proxy_state, &headers, &path_params, raw_query.as_deref(), - &body_bytes, - ) { - Ok(msg) => msg, - Err(message) => return bad_request(&entry, message), + body, + &entry, + ) + .await; + let (client, request) = match prepared { + Ok(prepared) => prepared, + Err(rejection) => return rejection.into_response(&entry), }; + if entry.streaming { + let keep_alive_secs = proxy_state.sse_keep_alive_secs(); + streaming_call(client, request, entry, wants_sse(&headers), keep_alive_secs).await + } else { + unary_call(client, request, &entry).await + } +} - let grpc_metadata = - metadata::http_headers_to_grpc_metadata(&headers, proxy_state.forwarded_headers()); - let mut grpc_request = tonic::Request::new(request_msg); - *grpc_request.metadata_mut() = grpc_metadata; - metadata::apply_request_deadline(&mut grpc_request, &headers); +/// Why a request ends before the upstream is called. +enum Rejection { + /// It cannot be mapped onto the RPC (`INVALID_ARGUMENT`, 400). + Unmappable(String), + /// The upstream channel is not ready (`UNAVAILABLE`, 503). + NotReady(String), +} - let output_desc = entry.method.output(); - let grpc_codec = - codec::DynamicCodec::with_required_check(output_desc.clone(), entry.response_has_required); - let grpc_path = entry.grpc_path.clone(); +impl Rejection { + /// The answer, in the error body the upstream's own errors get on the route. + fn into_response(self, entry: &RouteEntry) -> Response { + let status = match self { + Self::Unmappable(message) => tonic::Status::invalid_argument(message), + Self::NotReady(message) => tonic::Status::unavailable(message), + }; + error::status_to_response_with_details(&status, entry.error_details.as_deref()) + } +} - let mut grpc_client = Grpc::new(channel); - if let Err(e) = grpc_client.ready().await { - let status = tonic::Status::unavailable(format!("gRPC upstream not ready: {e}")); - return error::status_to_response_with_details(&status, entry.error_details.as_deref()); +/// Map the request onto the RPC's input message and get a client whose +/// channel is ready. +async fn prepare( + proxy_state: &S, + headers: &HeaderMap, + path_params: &PathParams, + raw_query: Option<&str>, + body: Bytes, + entry: &RouteEntry, +) -> Result< + ( + Grpc, + tonic::Request, + ), + Rejection, +> { + let message = decode_request(entry, headers, path_params, raw_query, body) + .map_err(Rejection::Unmappable)?; + let mut request = tonic::Request::new(message); + *request.metadata_mut() = + metadata::http_headers_to_grpc_metadata(headers, proxy_state.forwarded_headers()); + metadata::apply_request_deadline(&mut request, headers); + + let mut client = Grpc::new(proxy_state.grpc_channel()); + if let Err(e) = client.ready().await { + return Err(Rejection::NotReady(format!("gRPC upstream not ready: {e}"))); } + Ok((client, request)) +} - let use_sse = wants_sse(&headers); +/// A successful unary answer with its initial metadata and trailers kept +/// apart. +struct UnaryAnswer { + initial: MetadataMap, + message: DynamicMessage, + trailers: Option, +} - match grpc_client - .server_streaming(grpc_request, grpc_path, grpc_codec) - .await - { - // Only a trailers-only rejection lands in `Err`. Once the upstream - // accepted the call, the response starts at once rather than waiting - // for the first item, so headers and SSE keep-alives are not held back; - // an error that comes before the first message is a terminal frame. - Ok(response) => { - let stream = response.into_inner(); - // The terminal frame renders like the unary error body. The - // closure takes over this request's route handle, so the stream - // keeps it alive without another refcount. - let envelope = entry.ndjson_envelope; - let render_error = move |status: &tonic::Status| { - error::error_body(status, entry.error_details.as_deref()) - }; - if use_sse { - sse_response(stream, render_error, proxy_state.sse_keep_alive_secs()) - } else { - ndjson_response(stream, render_error, envelope) +/// Call a unary RPC. tonic's `Grpc::unary` merges the trailers over the +/// initial metadata, so a key sent in both keeps only its trailer value; the +/// HTTP response carries both, so the call is made as a one-message stream +/// instead, reading exactly what `Grpc::unary` reads. +async fn call_unary( + client: &mut Grpc, + request: tonic::Request, + entry: &RouteEntry, +) -> Result { + let response = client + .server_streaming(request, entry.grpc_path.clone(), entry.codec()) + .await?; + let (initial, mut stream, _) = response.into_parts(); + let message = stream + .message() + .await? + .ok_or_else(|| tonic::Status::internal("Missing response message."))?; + let trailers = stream.trailers().await?; + Ok(UnaryAnswer { + initial, + message, + trailers, + }) +} + +/// Serve a unary RPC. +async fn unary_call( + mut client: Grpc, + request: tonic::Request, + entry: &RouteEntry, +) -> Response { + match call_unary(&mut client, request, entry).await { + Ok(answer) => unary_success(entry, answer), + Err(status) => upstream_error(status, entry), + } +} + +/// The HTTP response to a successful unary call: the status `x-http-code` +/// sets (200 otherwise), the forwarded response metadata, and the body. +fn unary_success(entry: &RouteEntry, answer: UnaryAnswer) -> Response { + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(answer.initial, &entry.denied_headers); + if let Some(trailers) = answer.trailers { + upstream.absorb(trailers, &entry.denied_headers); + } + let status = match upstream.status() { + Ok(status) => status.unwrap_or(StatusCode::OK), + Err(response::InvalidHttpCode) => { + tracing::error!( + rpc = %entry.grpc_path, + "upstream set x-http-code to something other than one integer in 200-599" + ); + return error::malformed_response(entry.error_details.as_deref()); + } + }; + let (content_type, body) = match &entry.response { + ResponseShape::HttpBody(path) => { + let raw = httpbody::take(answer.message, path); + match content_type_header(&raw.content_type) { + Ok(content_type) => (content_type, Body::from(raw.data)), + Err(()) => return invalid_content_type(entry), } } - Err(status) => { - error::status_to_response_with_details(&status, entry.error_details.as_deref()) + ResponseShape::Json(response_body) => { + match json_body(&answer.message, response_body.as_deref()) { + Ok(json) => ( + Some(HeaderValue::from_static("application/json")), + Body::from(json), + ), + Err(e) => { + tracing::error!("Failed to serialize gRPC response: {e}"); + return error::status_to_response_with_details( + &tonic::Status::internal("failed to serialize response"), + entry.error_details.as_deref(), + ); + } + } } + }; + response::build(status, upstream.into_headers(), content_type, body) +} + +/// The `Content-Type` an HttpBody asks for: none when it left the field empty. +fn content_type_header(content_type: &str) -> Result, ()> { + if content_type.is_empty() { + return Ok(None); } + HeaderValue::from_str(content_type) + .map(Some) + .map_err(|_| ()) +} + +/// The answer to an HttpBody whose content type cannot be a header value. +fn invalid_content_type(entry: &RouteEntry) -> Response { + tracing::error!( + rpc = %entry.grpc_path, + "upstream HttpBody content_type is not a valid header value" + ); + error::malformed_response(entry.error_details.as_deref()) +} + +/// The HTTP response to a failed call, carrying the failure's metadata as +/// headers unless its details were malformed and the answer is the generic +/// `INTERNAL`. Only the failure's own metadata is used (a trailers-only +/// response, or the trailers ending the call): a failure the proxy's client +/// raises itself, such as an undecodable message, has none, so nothing of an +/// answer the proxy rejected reaches the client. +fn upstream_error(mut status: tonic::Status, entry: &RouteEntry) -> Response { + let (response, faithful) = error::render_response(&status, entry.error_details.as_deref()); + let metadata = std::mem::take(status.metadata_mut()); + if !faithful || metadata.is_empty() { + return response; + } + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(metadata, &entry.denied_headers); + response::with_upstream_headers(response, upstream.into_headers()) +} + +/// Serve a server-streaming RPC. +/// +/// A JSON stream is Server-Sent Events when the client sends +/// `Accept: text/event-stream`, otherwise newline-delimited JSON (NDJSON); a +/// gRPC error mid-stream is delivered as an explicit terminal frame before the +/// stream is closed cleanly, rather than truncating the HTTP body. An HttpBody +/// stream is the concatenated `data` of its messages. The upstream's initial +/// metadata becomes response headers; trailers arrive after the headers are +/// sent and are not forwarded. +async fn streaming_call( + mut client: Grpc, + request: tonic::Request, + entry: Arc, + use_sse: bool, + keep_alive_secs: u64, +) -> Response { + let response = match client + .server_streaming(request, entry.grpc_path.clone(), entry.codec()) + .await + { + Ok(response) => response, + // Only a trailers-only rejection lands here. + Err(status) => return upstream_error(status, &entry), + }; + let (initial, stream, _) = response.into_parts(); + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(initial, &entry.denied_headers); + let headers = upstream.into_headers(); + + if matches!(entry.response, ResponseShape::HttpBody(_)) { + return http_body_stream(stream, entry, headers).await; + } + // Once the upstream accepted the call, the response starts at once rather + // than waiting for the first item, so headers and SSE keep-alives are not + // held back; an error that comes before the first message is a terminal + // frame. The terminal frame renders like the unary error body. The + // closure takes over this request's route handle, so the stream keeps it + // alive without another refcount. + let envelope = entry.ndjson_envelope; + let render_error = + move |status: &tonic::Status| error::error_body(status, entry.error_details.as_deref()); + let response = if use_sse { + sse_response(stream, render_error, keep_alive_secs) + } else { + ndjson_response(stream, render_error, envelope) + }; + response::with_upstream_headers(response, headers) +} + +/// A server-streaming HttpBody response: `Content-Type` from the first +/// message, so the headers wait for it, then every message's `data` as it +/// arrives. An error before the first message is an ordinary error response; +/// after it, the raw body has no in-band error frame, so the body is aborted +/// and the client sees a truncated transfer instead of a clean end. +async fn http_body_stream( + mut stream: tonic::Streaming, + entry: Arc, + headers: HeaderMap, +) -> Response { + let first = match stream.message().await { + Ok(Some(message)) => entry.raw_body(message).unwrap_or_default(), + Ok(None) => httpbody::RawBody::default(), + Err(status) => return upstream_error(status, &entry), + }; + let content_type = match content_type_header(&first.content_type) { + Ok(content_type) => content_type, + Err(()) => return invalid_content_type(&entry), + }; + let rest = stream.map(move |item| match item { + Ok(message) => Ok(entry.raw_body(message).unwrap_or_default().data), + Err(status) => { + tracing::error!( + rpc = %entry.grpc_path, + code = ?status.code(), + "server-streaming HttpBody failed after the response started; aborting the body" + ); + Err(std::io::Error::other(status)) + } + }); + let chunks = futures::stream::once(futures::future::ready(Ok(first.data))) + .chain(rest) + .try_filter(|chunk| futures::future::ready(!chunk.is_empty())); + response::build( + StatusCode::OK, + headers, + content_type, + Body::from_stream(chunks), + ) } /// One frame of a streaming response: a serialized message, or the error body @@ -495,16 +877,16 @@ where } }; line.push('\n'); - Ok::(axum::body::Bytes::from(line)) + Ok::(Bytes::from(line)) }); - let body = axum::body::Body::from_stream(byte_stream); + let body = Body::from_stream(byte_stream); // Body framing (chunked on HTTP/1.1, DATA frames on HTTP/2) is chosen by // hyper from the protocol version; setting transfer-encoding by hand would // be redundant on HTTP/1.1 and illegal on HTTP/2. Response::builder() .status(StatusCode::OK) - .header("content-type", "application/x-ndjson") + .header(CONTENT_TYPE, "application/x-ndjson") .body(body) .unwrap_or_else(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response()) } @@ -533,21 +915,29 @@ where .into_response() } +/// Request body mapping of the raw-body routes: nothing parsed. +static NO_PARSED_BODY: request::BodyMapping = request::BodyMapping::None; + /// Map the HTTP request onto the RPC's input message: path parameters, query -/// parameters and the route's `body` rule. Unary and server-streaming routes -/// share it, so both bind a request the same way. The error is the message of -/// the 400 the caller answers with, before the upstream is called. +/// parameters and the route's `body` rule, or the raw body for an HttpBody. +/// Unary and server-streaming routes share it, so both bind a request the same +/// way. The error is the message of the 400 the caller answers with, before +/// the upstream is called. fn decode_request( entry: &RouteEntry, headers: &HeaderMap, - path_params: &std::collections::HashMap, + path_params: &PathParams, raw_query: Option<&str>, - body_bytes: &[u8], + body_bytes: Bytes, ) -> Result { - // Only read the body when the rule maps it onto the message. - let json_body = match entry.body { + let mapping = match &entry.request_body { + RequestBody::Parsed(mapping) => mapping, + RequestBody::RawRoot | RequestBody::RawField(..) => &NO_PARSED_BODY, + }; + // Only parse the body when the rule maps it onto the message. + let json_body = match mapping { request::BodyMapping::None => serde_json::Value::Null, - _ => body::parse_body(body::content_type(headers), body_bytes) + _ => body::parse_body(body::content_type(headers), &body_bytes) .map_err(|e| format!("failed to parse request body: {e}"))?, }; @@ -557,105 +947,40 @@ fn decode_request( let query_pairs = request::parse_query(raw_query)?; let input_desc = entry.method.input(); - let request_json = request::build_request_json( - &input_desc, - &entry.body, - json_body, - path_params, - &query_pairs, - )?; - - DynamicMessage::deserialize(input_desc, request_json) - .map_err(|e| format!("failed to decode request: {e}")) -} - -/// The 400 answer to a request [`decode_request`] could not map, in the same -/// error body the upstream's own errors get on this route. -fn bad_request(entry: &RouteEntry, message: String) -> Response { - error::status_to_response_with_details( - &tonic::Status::invalid_argument(message), - entry.error_details.as_deref(), - ) -} - -/// Generic transcoding handler. -async fn transcode_handler( - State(proxy_state): State, - headers: HeaderMap, - Path(path_params): Path>, - RawQuery(raw_query): RawQuery, - body_bytes: axum::body::Bytes, - entry: std::sync::Arc, -) -> Response { - let channel = proxy_state.grpc_channel(); - - let request_msg = match decode_request( - &entry, - &headers, - &path_params, - raw_query.as_deref(), - &body_bytes, - ) { - Ok(msg) => msg, - Err(message) => return bad_request(&entry, message), - }; - - let grpc_metadata = - metadata::http_headers_to_grpc_metadata(&headers, proxy_state.forwarded_headers()); - let mut grpc_request = tonic::Request::new(request_msg); - *grpc_request.metadata_mut() = grpc_metadata; - metadata::apply_request_deadline(&mut grpc_request, &headers); - - let output_desc = entry.method.output(); - let grpc_codec = - codec::DynamicCodec::with_required_check(output_desc.clone(), entry.response_has_required); - let grpc_path = entry.grpc_path.clone(); - - let mut grpc_client = Grpc::new(channel); - if let Err(e) = grpc_client.ready().await { - let status = tonic::Status::unavailable(format!("gRPC upstream not ready: {e}")); - return error::status_to_response_with_details(&status, entry.error_details.as_deref()); - } - - match grpc_client.unary(grpc_request, grpc_path, grpc_codec).await { - Ok(response) => { - let response_msg = response.into_inner(); - let serialize_opts = response_serialize_options(); - match response_msg - .serialize_with_options(serde_json::value::Serializer, &serialize_opts) - { - Ok(json_value) => { - // `response_body` returns just that subfield as the HTTP body. - let out = match &entry.response_body { - Some(path) => request::extract_response_body(&json_value, path) - .unwrap_or_else(|| { - tracing::warn!( - response_body = %path, - "configured response_body path not found in response; \ - returning null" - ); - serde_json::Value::Null - }), - None => json_value, - }; - (StatusCode::OK, Json(out)).into_response() - } - Err(e) => { - tracing::error!("Failed to serialize gRPC response: {e}"); - error::status_to_response_with_details( - &tonic::Status::internal("failed to serialize response"), - entry.error_details.as_deref(), - ) - } - } + let request_json = + request::build_request_json(&input_desc, mapping, json_body, path_params, &query_pairs)?; + + let mut message = DynamicMessage::deserialize(input_desc, request_json) + .map_err(|e| format!("failed to decode request: {e}"))?; + match &entry.request_body { + RequestBody::Parsed(_) => {} + RequestBody::RawRoot => { + httpbody::fill(&mut message, request_content_type(headers)?, body_bytes); } - Err(status) => { - error::status_to_response_with_details(&status, entry.error_details.as_deref()) + RequestBody::RawField(field, http_body) => { + let mut inner = DynamicMessage::new(http_body.clone()); + httpbody::fill(&mut inner, request_content_type(headers)?, body_bytes); + message.set_field(field, prost_reflect::Value::Message(inner)); } } + Ok(message) +} + +/// The request's full `Content-Type` value (parameters included) for an +/// HttpBody, empty when absent. `HttpBody.content_type` is a proto string, so +/// a value that is not visible ASCII is rejected rather than altered. +fn request_content_type(headers: &HeaderMap) -> Result { + match headers.get(CONTENT_TYPE) { + None => Ok(String::new()), + Some(value) => value + .to_str() + .map(str::to_owned) + .map_err(|_| "request Content-Type is not a visible ASCII string".to_string()), + } } -/// Extract HTTP route entries from proto descriptors. +/// 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 { let http_ext = match pool.get_extension_by_name("google.api.http") { Some(ext) => ext, @@ -669,7 +994,7 @@ fn extract_routes(pool: &DescriptorPool) -> Vec { for service in pool.services() { for method in service.methods() { - if method.is_client_streaming() || method.is_server_streaming() { + if method.is_client_streaming() { continue; } @@ -682,74 +1007,32 @@ fn extract_routes(pool: &DescriptorPool) -> Vec { } }; - let response_has_required = codec::has_required_fields(&method.output()); - for binding in extract_http_bindings(&method, &http_ext) { - entries.push(RouteEntry { - http_path: binding.http_path, - http_method: binding.http_method, - grpc_path: grpc_path.clone(), - method: method.clone(), - body: binding.body, - response_body: binding.response_body, - // Decided per mounted path in `routes_with_options`. - error_details: None, - ndjson_envelope: false, - response_has_required, - }); - } - } - } - - entries -} - -/// Extract server-streaming HTTP route entries. -fn extract_streaming_routes(pool: &DescriptorPool) -> Vec { - let http_ext = match pool.get_extension_by_name("google.api.http") { - Some(ext) => ext, - None => return Vec::new(), - }; - - let mut entries = Vec::new(); - - for service in pool.services() { - for method in service.methods() { - if !method.is_server_streaming() || method.is_client_streaming() { - continue; - } - - let grpc_path = format!("/{}/{}", service.full_name(), method.name()); - let grpc_path: axum::http::uri::PathAndQuery = match grpc_path.parse() { - Ok(p) => p, - Err(e) => { - tracing::error!("skipping route with invalid gRPC path '{grpc_path}': {e}"); - continue; + let input = method.input(); + let output = method.output(); + let streaming = method.is_server_streaming(); + let response_has_required = codec::has_required_fields(&output); + for binding in rule::http_bindings(&method, &http_ext) { + if streaming { + tracing::info!( + "Registering streaming route: {} {} → {}", + binding.method.as_str(), + binding.path, + grpc_path + ); } - }; - - let response_has_required = codec::has_required_fields(&method.output()); - for binding in extract_http_bindings(&method, &http_ext) { - tracing::info!( - "Registering streaming route: {} {} → {}", - match binding.http_method { - HttpMethod::Get => "GET", - HttpMethod::Post => "POST", - _ => "OTHER", - }, - binding.http_path, - grpc_path - ); entries.push(RouteEntry { - http_path: binding.http_path, - http_method: binding.http_method, + http_path: binding.path, + http_method: binding.method, grpc_path: grpc_path.clone(), method: method.clone(), - body: binding.body, - response_body: binding.response_body, + streaming, + request_body: request_body(&input, binding.body), + response: response_shape(&output, binding.response_body), // Decided per mounted path in `routes_with_options`. error_details: None, ndjson_envelope: false, response_has_required, + denied_headers: Arc::default(), }); } } @@ -758,97 +1041,31 @@ fn extract_streaming_routes(pool: &DescriptorPool) -> Vec { entries } -/// A single HTTP binding parsed from a `google.api.http` rule. -struct HttpBinding { - http_method: HttpMethod, - http_path: String, - body: request::BodyMapping, - response_body: Option, -} - -/// Extract all HTTP bindings (the primary rule plus any `additional_bindings`) -/// from a method's `google.api.http` extension. -fn extract_http_bindings( - method: &MethodDescriptor, - http_ext: &prost_reflect::ExtensionDescriptor, -) -> Vec { - let options = method.options(); - if !options.has_extension(http_ext) { - return Vec::new(); - } - - let prost_reflect::Value::Message(rule_msg) = options.get_extension(http_ext).into_owned() - else { - return Vec::new(); - }; - - collect_bindings(&rule_msg) -} - -/// Collect the primary binding plus every `additional_bindings` entry from an -/// `HttpRule` message. -fn collect_bindings(rule_msg: &DynamicMessage) -> Vec { - let mut bindings = Vec::new(); - if let Some(binding) = parse_http_rule(rule_msg) { - bindings.push(binding); - } - - // additional_bindings is a repeated HttpRule; each carries its own - // method/path/body. The proto forbids nesting them further. - if let Some(field) = rule_msg.get_field_by_name("additional_bindings") { - if let prost_reflect::Value::List(list) = field.into_owned() { - for item in list { - if let prost_reflect::Value::Message(sub) = item { - if let Some(binding) = parse_http_rule(&sub) { - bindings.push(binding); - } - } - } - } +/// How a binding's `body` rule reaches `input`: raw into an HttpBody (the +/// input itself with `body: "*"`, or the HttpBody field `body` names), parsed +/// otherwise. +fn request_body(input: &MessageDescriptor, mapping: request::BodyMapping) -> RequestBody { + match &mapping { + request::BodyMapping::Root if httpbody::is_http_body(input) => RequestBody::RawRoot, + request::BodyMapping::Field(name) => match httpbody::http_body_field(input, name) { + Some((field, http_body)) => RequestBody::RawField(field, http_body), + None => RequestBody::Parsed(mapping), + }, + _ => RequestBody::Parsed(mapping), } - - bindings } -/// Parse a single `HttpRule` message into a binding (method+path required). -fn parse_http_rule(rule_msg: &DynamicMessage) -> Option { - let (http_method, http_path) = [ - ("get", HttpMethod::Get), - ("post", HttpMethod::Post), - ("put", HttpMethod::Put), - ("delete", HttpMethod::Delete), - ("patch", HttpMethod::Patch), - ] - .into_iter() - .find_map( - |(name, http_method)| match rule_msg.get_field_by_name(name)?.into_owned() { - prost_reflect::Value::String(path) if !path.is_empty() => Some((http_method, path)), - _ => None, +/// What a binding answers with: the raw body of an HttpBody (the output itself, +/// or the HttpBody `response_body` names), JSON otherwise. +fn response_shape(output: &MessageDescriptor, response_body: Option) -> ResponseShape { + match response_body { + None if httpbody::is_http_body(output) => ResponseShape::HttpBody(Vec::new()), + None => ResponseShape::Json(None), + Some(path) => match httpbody::http_body_path(output, &path) { + Some(fields) => ResponseShape::HttpBody(fields), + None => ResponseShape::Json(Some(path)), }, - )?; - - let body = rule_msg - .get_field_by_name("body") - .and_then(|v| match v.into_owned() { - prost_reflect::Value::String(s) => Some(request::BodyMapping::parse(&s)), - _ => None, - }) - .unwrap_or(request::BodyMapping::None); - - let response_body = - rule_msg - .get_field_by_name("response_body") - .and_then(|v| match v.into_owned() { - prost_reflect::Value::String(s) if !s.is_empty() => Some(s), - _ => None, - }); - - Some(HttpBinding { - http_method, - http_path, - body, - response_body, - }) + } } /// Convert a `google.api.http` path template to axum 0.8 path syntax. diff --git a/src/transcode/response.rs b/src/transcode/response.rs new file mode 100644 index 0000000..b2aca8e --- /dev/null +++ b/src/transcode/response.rs @@ -0,0 +1,191 @@ +//! Upstream response metadata → HTTP response headers and status. +//! +//! The upstream's response metadata is its HTTP response headers, as in Envoy's +//! transcoder, where gRPC and HTTP share one HTTP/2 stream. Only what belongs +//! to gRPC itself or to the upstream connection is held back, plus +//! `x-http-code`, which sets the status of a successful unary answer +//! (grpc-gateway's convention). + +use axum::body::Body; +use axum::http::header::{Entry, OccupiedEntry, CONTENT_TYPE}; +use axum::http::{HeaderMap, HeaderName, HeaderValue, StatusCode}; +use axum::response::Response; +use tonic::metadata::MetadataMap; + +/// Response metadata key whose value sets the HTTP status of a successful +/// unary call. +pub(crate) const HTTP_CODE_KEY: &str = "x-http-code"; + +/// Whether a response metadata key never becomes an HTTP response header: +/// gRPC's own keys (`grpc-*`, the binary `-bin` encoding, `content-type`, which +/// the transcoder sets), and the hop-by-hop and framing fields that describe +/// the upstream connection rather than the response (RFC 9110 §7.6.1, +/// RFC 9110 §8.6 for `content-length`). +fn is_withheld(name: &str) -> bool { + name.starts_with("grpc-") + || name.ends_with("-bin") + || matches!( + name, + "content-type" + | "connection" + | "keep-alive" + | "proxy-connection" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + | "content-length" + ) +} + +/// What `x-http-code` said across the metadata absorbed so far. +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +enum HttpCode { + #[default] + Absent, + Set(StatusCode), + /// Not an integer in 200-599, or given more than once. + Invalid, +} + +/// The `x-http-code` value is not a single integer in 200-599. +#[derive(Debug, PartialEq, Eq)] +pub(crate) struct InvalidHttpCode; + +/// HTTP response headers collected from the upstream's response metadata. +#[derive(Debug, Default)] +pub(crate) struct UpstreamHeaders { + headers: HeaderMap, + http_code: HttpCode, +} + +/// How [`UpstreamHeaders::absorb`] treats the values of the key it is on. +enum Current<'a> { + Forward(OccupiedEntry<'a, HeaderValue>), + HttpCode, + Drop, +} + +impl UpstreamHeaders { + /// Add the forwardable entries of `metadata`, in order, after the ones + /// already absorbed: a key present in both initial metadata and trailers + /// keeps both values. Keys in `deny` are dropped like withheld ones. + pub(crate) fn absorb(&mut self, metadata: MetadataMap, deny: &[HeaderName]) { + let headers = &mut self.headers; + let http_code = &mut self.http_code; + let mut current = Current::Drop; + // A header map yields a name only with the first value of each key. + for (name, value) in metadata.into_headers() { + if let Some(name) = name { + // Release the previous key's entry before borrowing the map again. + current = Current::Drop; + if name.as_str() == HTTP_CODE_KEY { + record_http_code(http_code, &value); + current = Current::HttpCode; + } else if !is_withheld(name.as_str()) && !deny.contains(&name) { + current = Current::Forward(match headers.entry(name) { + Entry::Occupied(mut entry) => { + entry.append(value); + entry + } + Entry::Vacant(entry) => entry.insert_entry(value), + }); + } + continue; + } + match &mut current { + Current::Forward(entry) => entry.append(value), + Current::HttpCode => record_http_code(http_code, &value), + Current::Drop => {} + } + } + } + + /// The status `x-http-code` sets, if any. + pub(crate) fn status(&self) -> Result, InvalidHttpCode> { + match self.http_code { + HttpCode::Absent => Ok(None), + HttpCode::Set(status) => Ok(Some(status)), + HttpCode::Invalid => Err(InvalidHttpCode), + } + } + + pub(crate) fn into_headers(self) -> HeaderMap { + self.headers + } +} + +/// Fold one `x-http-code` value into `slot`: the first valid one sets it, a +/// second one of any kind makes it invalid (two statuses cannot both apply). +fn record_http_code(slot: &mut HttpCode, value: &HeaderValue) { + *slot = match (*slot, parse_http_code(value)) { + (HttpCode::Absent, Some(status)) => HttpCode::Set(status), + _ => HttpCode::Invalid, + }; +} + +/// An `x-http-code` value: exactly three ASCII digits naming a status in +/// 200-599. `1xx` is excluded because an interim response cannot be the final +/// answer (RFC 9110 §15.2). +fn parse_http_code(value: &HeaderValue) -> Option { + let bytes = value.as_bytes(); + if bytes.len() != 3 || !bytes.iter().all(u8::is_ascii_digit) { + return None; + } + let code = bytes + .iter() + .fold(0u16, |code, digit| code * 10 + u16::from(digit - b'0')); + if !(200..=599).contains(&code) { + return None; + } + StatusCode::from_u16(code).ok() +} + +/// `response` with `upstream` as its headers, the ones it already carries +/// (`Content-Type`, an SSE `Cache-Control`) replacing same-named upstream +/// values: what the proxy writes describes the body it writes. +pub(crate) fn with_upstream_headers(response: Response, upstream: HeaderMap) -> Response { + if upstream.is_empty() { + return response; + } + let (mut parts, body) = response.into_parts(); + let own = std::mem::replace(&mut parts.headers, upstream); + let mut last: Option = None; + for (name, value) in own { + match name { + Some(name) => { + parts.headers.insert(&name, value); + last = Some(name); + } + None => { + if let Some(name) = &last { + parts.headers.append(name, value); + } + } + } + } + Response::from_parts(parts, body) +} + +/// A response with `status`, `headers` and `body`, typed by `content_type`. +/// `204` and `304` carry neither content nor a `Content-Type`, whatever the +/// upstream returned (RFC 9110 §15.3.5, §15.4.5). +pub(crate) fn build( + status: StatusCode, + mut headers: HeaderMap, + content_type: Option, + body: Body, +) -> Response { + let no_content = matches!(status, StatusCode::NO_CONTENT | StatusCode::NOT_MODIFIED); + let body = if no_content { Body::empty() } else { body }; + if let (Some(content_type), false) = (content_type, no_content) { + headers.insert(CONTENT_TYPE, content_type); + } + let mut response = Response::new(body); + *response.status_mut() = status; + *response.headers_mut() = headers; + response +} + +#[cfg(test)] +mod tests; diff --git a/src/transcode/response/tests.rs b/src/transcode/response/tests.rs new file mode 100644 index 0000000..3dcb6dd --- /dev/null +++ b/src/transcode/response/tests.rs @@ -0,0 +1,262 @@ +use super::*; +use tonic::metadata::{AsciiMetadataValue, BinaryMetadataValue}; + +fn metadata(entries: &[(&'static str, &'static str)]) -> MetadataMap { + let mut md = MetadataMap::new(); + for (key, value) in entries { + md.append(*key, AsciiMetadataValue::from_static(value)); + } + md +} + +fn values<'a>(headers: &'a HeaderMap, name: &str) -> Vec<&'a str> { + headers + .get_all(name) + .iter() + .map(|v| v.to_str().unwrap()) + .collect() +} + +#[test] +fn application_metadata_becomes_headers() { + let mut upstream = UpstreamHeaders::default(); + upstream.absorb( + metadata(&[ + ("cache-control", "no-store"), + ("www-authenticate", "Bearer error=\"invalid_token\""), + ("dpop-nonce", "n-1"), + ]), + &[], + ); + let headers = upstream.into_headers(); + assert_eq!(values(&headers, "cache-control"), ["no-store"]); + assert_eq!( + values(&headers, "www-authenticate"), + ["Bearer error=\"invalid_token\""] + ); + assert_eq!(values(&headers, "dpop-nonce"), ["n-1"]); +} + +#[test] +fn grpc_and_connection_keys_are_withheld() { + // gRPC's own keys, content-type (the transcoder sets it) and the + // hop-by-hop / framing fields describe the upstream stream, not the + // HTTP response. + let withheld = [ + ("grpc-status", "0"), + ("grpc-message", "ok"), + ("grpc-encoding", "gzip"), + ("grpc-accept-encoding", "gzip"), + ("content-type", "application/grpc"), + ("connection", "close"), + ("keep-alive", "timeout=5"), + ("proxy-connection", "keep-alive"), + ("te", "trailers"), + ("trailer", "grpc-status"), + ("transfer-encoding", "chunked"), + ("upgrade", "h2c"), + ("content-length", "12"), + ("x-http-code", "302"), + ]; + let mut md = metadata(&withheld); + md.append_bin("trace-bin", BinaryMetadataValue::from_bytes(b"\x00\x01")); + md.append("x-kept", AsciiMetadataValue::from_static("yes")); + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(md, &[]); + let headers = upstream.into_headers(); + assert_eq!( + headers.keys().map(|k| k.as_str()).collect::>(), + ["x-kept"] + ); +} + +#[test] +fn deny_list_drops_its_keys_only() { + let mut upstream = UpstreamHeaders::default(); + upstream.absorb( + metadata(&[ + ("x-debug-trace", "t"), + ("x-debug-trace", "u"), + ("location", "/next"), + ]), + &[HeaderName::from_static("x-debug-trace")], + ); + let headers = upstream.into_headers(); + assert!(headers.get("x-debug-trace").is_none()); + assert_eq!(values(&headers, "location"), ["/next"]); +} + +#[test] +fn repeated_values_and_trailers_keep_every_value_in_order() { + // Repeated metadata values are repeated header fields; a key sent in + // initial metadata and trailers keeps both values, initial first. + let mut upstream = UpstreamHeaders::default(); + upstream.absorb( + metadata(&[("set-cookie", "a=1"), ("set-cookie", "b=2"), ("x-one", "1")]), + &[], + ); + upstream.absorb(metadata(&[("set-cookie", "c=3"), ("x-two", "2")]), &[]); + let headers = upstream.into_headers(); + assert_eq!(values(&headers, "set-cookie"), ["a=1", "b=2", "c=3"]); + assert_eq!(values(&headers, "x-one"), ["1"]); + assert_eq!(values(&headers, "x-two"), ["2"]); +} + +#[test] +fn http_code_sets_the_status_and_is_not_forwarded() { + let mut upstream = UpstreamHeaders::default(); + upstream.absorb( + metadata(&[("x-http-code", "302"), ("location", "/cb")]), + &[], + ); + assert_eq!(upstream.status(), Ok(Some(StatusCode::FOUND))); + assert!(upstream.into_headers().get("x-http-code").is_none()); +} + +#[test] +fn absent_http_code_leaves_the_status_alone() { + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(metadata(&[("x-other", "1")]), &[]); + assert_eq!(upstream.status(), Ok(None)); +} + +#[test] +fn http_code_range_edges() { + for (value, expected) in [ + ("200", Some(StatusCode::OK)), + ("599", Some(StatusCode::from_u16(599).unwrap())), + ("400", Some(StatusCode::BAD_REQUEST)), + ("199", None), + ("600", None), + ("100", None), + ("000", None), + ] { + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(metadata(&[("x-http-code", value)]), &[]); + let status = upstream.status(); + match expected { + Some(code) => assert_eq!(status, Ok(Some(code)), "{value}"), + None => assert_eq!(status, Err(InvalidHttpCode), "{value}"), + } + } +} + +#[test] +fn http_code_that_is_not_three_digits_is_invalid() { + for value in ["+200", "20", "2000", " 200", "200 ", "2x0", "", "0x1f"] { + let mut md = MetadataMap::new(); + md.append("x-http-code", AsciiMetadataValue::try_from(value).unwrap()); + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(md, &[]); + assert_eq!(upstream.status(), Err(InvalidHttpCode), "{value:?}"); + } +} + +#[test] +fn http_code_given_twice_is_invalid_even_if_equal() { + // Twice in one map, and once in initial metadata plus once in trailers. + let mut upstream = UpstreamHeaders::default(); + upstream.absorb( + metadata(&[("x-http-code", "400"), ("x-http-code", "400")]), + &[], + ); + assert_eq!(upstream.status(), Err(InvalidHttpCode)); + + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(metadata(&[("x-http-code", "400")]), &[]); + upstream.absorb(metadata(&[("x-http-code", "401")]), &[]); + assert_eq!(upstream.status(), Err(InvalidHttpCode)); +} + +#[test] +fn http_code_cannot_be_denied_into_forwarding() { + // The deny-list only removes more; x-http-code is consumed either way. + let mut upstream = UpstreamHeaders::default(); + upstream.absorb( + metadata(&[("x-http-code", "201")]), + &[HeaderName::from_static("x-http-code")], + ); + assert_eq!(upstream.status(), Ok(Some(StatusCode::CREATED))); + assert!(upstream.into_headers().is_empty()); +} + +#[test] +fn own_headers_replace_same_named_upstream_values() { + let mut upstream = HeaderMap::new(); + upstream.append("cache-control", HeaderValue::from_static("max-age=60")); + upstream.append("cache-control", HeaderValue::from_static("public")); + upstream.append("x-upstream", HeaderValue::from_static("1")); + let mut response = Response::new(Body::empty()); + response + .headers_mut() + .insert("cache-control", HeaderValue::from_static("no-cache")); + response + .headers_mut() + .insert(CONTENT_TYPE, HeaderValue::from_static("text/event-stream")); + let response = with_upstream_headers(response, upstream); + let headers = response.headers(); + assert_eq!(values(headers, "cache-control"), ["no-cache"]); + assert_eq!(values(headers, "content-type"), ["text/event-stream"]); + assert_eq!(values(headers, "x-upstream"), ["1"]); +} + +#[test] +fn own_multi_valued_header_keeps_all_its_values() { + let mut response = Response::new(Body::empty()); + response + .headers_mut() + .append("vary", HeaderValue::from_static("accept")); + response + .headers_mut() + .append("vary", HeaderValue::from_static("origin")); + let mut upstream = HeaderMap::new(); + upstream.insert("vary", HeaderValue::from_static("cookie")); + let response = with_upstream_headers(response, upstream); + assert_eq!(values(response.headers(), "vary"), ["accept", "origin"]); +} + +async fn body_bytes(response: Response) -> Vec { + axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap() + .to_vec() +} + +#[tokio::test] +async fn build_sets_status_headers_and_content_type() { + let mut headers = HeaderMap::new(); + headers.insert("cache-control", HeaderValue::from_static("no-store")); + let response = build( + StatusCode::BAD_REQUEST, + headers, + Some(HeaderValue::from_static("application/json")), + Body::from("{}"), + ); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!(response.headers()["cache-control"], "no-store"); + assert_eq!(response.headers()[CONTENT_TYPE], "application/json"); + assert_eq!(body_bytes(response).await, b"{}"); +} + +#[tokio::test] +async fn build_without_content_type_sets_none() { + let response = build(StatusCode::FOUND, HeaderMap::new(), None, Body::empty()); + assert!(response.headers().get(CONTENT_TYPE).is_none()); + assert!(body_bytes(response).await.is_empty()); +} + +#[tokio::test] +async fn no_content_statuses_drop_body_and_content_type() { + // RFC 9110 §15.3.5 / §15.4.5: 204 and 304 cannot carry content. + for status in [StatusCode::NO_CONTENT, StatusCode::NOT_MODIFIED] { + let response = build( + status, + HeaderMap::new(), + Some(HeaderValue::from_static("application/json")), + Body::from("{}"), + ); + assert_eq!(response.status(), status); + assert!(response.headers().get(CONTENT_TYPE).is_none(), "{status}"); + assert!(body_bytes(response).await.is_empty(), "{status}"); + } +} diff --git a/src/transcode/rule.rs b/src/transcode/rule.rs new file mode 100644 index 0000000..a8652f7 --- /dev/null +++ b/src/transcode/rule.rs @@ -0,0 +1,140 @@ +//! `google.api.http` rule parsing: the HTTP method and path template of every +//! binding of an RPC, with its `body` and `response_body` settings. The +//! transcoded routes and the OpenAPI document both read bindings from here, so +//! the two cannot disagree on what an RPC is mounted at. + +use axum::http::Method; +use prost_reflect::{DynamicMessage, ExtensionDescriptor, MethodDescriptor, Value}; + +use super::request::BodyMapping; + +/// The HTTP method a binding answers. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum RouteMethod { + /// One method: a standard pattern, or a `custom` rule naming any token. + One(Method), + /// Every method: a `custom` rule with `kind: "*"`. + Any, +} + +impl RouteMethod { + /// The method token, or `*` for [`RouteMethod::Any`]. + pub(crate) fn as_str(&self) -> &str { + match self { + Self::One(method) => method.as_str(), + Self::Any => "*", + } + } +} + +/// One HTTP binding of an RPC. +#[derive(Debug, Clone)] +pub(crate) struct HttpBinding { + pub(crate) method: RouteMethod, + /// Path template, as written in the rule. + pub(crate) path: String, + pub(crate) body: BodyMapping, + /// `response_body`, when set. + pub(crate) response_body: Option, +} + +/// Every binding of `method`: its `google.api.http` rule plus the rule's +/// `additional_bindings`, in that order. Empty when the RPC has no rule. +pub(crate) fn http_bindings( + method: &MethodDescriptor, + http_ext: &ExtensionDescriptor, +) -> Vec { + let options = method.options(); + if !options.has_extension(http_ext) { + return Vec::new(); + } + match options.get_extension(http_ext).as_ref() { + Value::Message(rule) => collect_bindings(rule), + _ => Vec::new(), + } +} + +/// The binding of an `HttpRule` message plus every `additional_bindings` entry. +pub(crate) fn collect_bindings(rule: &DynamicMessage) -> Vec { + let mut bindings = Vec::new(); + bindings.extend(parse_http_rule(rule)); + + // additional_bindings is a repeated HttpRule; each carries its own + // pattern, body and response_body. The proto forbids nesting them further. + if let Some(field) = rule.get_field_by_name("additional_bindings") { + if let Value::List(list) = field.as_ref() { + for item in list { + if let Value::Message(sub) = item { + bindings.extend(parse_http_rule(sub)); + } + } + } + } + + bindings +} + +/// The binding one `HttpRule` describes, or `None` when it sets no pattern (or +/// an unusable `custom` one). +fn parse_http_rule(rule: &DynamicMessage) -> Option { + let (method, path) = standard_pattern(rule).or_else(|| custom_pattern(rule))?; + let body = string_field(rule, "body") + .map(|body| BodyMapping::parse(&body)) + .unwrap_or(BodyMapping::None); + Some(HttpBinding { + method, + path, + body, + response_body: string_field(rule, "response_body"), + }) +} + +/// The `get` / `put` / `post` / `delete` / `patch` member of the `pattern` +/// oneof, when one is set. +fn standard_pattern(rule: &DynamicMessage) -> Option<(RouteMethod, String)> { + [ + ("get", Method::GET), + ("put", Method::PUT), + ("post", Method::POST), + ("delete", Method::DELETE), + ("patch", Method::PATCH), + ] + .into_iter() + .find_map(|(name, method)| Some((RouteMethod::One(method), string_field(rule, name)?))) +} + +/// The `custom` member of the `pattern` oneof (`CustomHttpPattern {kind, +/// path}`). `google/api/http.proto` defines `kind: "*"` as "every method"; any +/// other kind is the method token itself, case-sensitive (RFC 9110 §9.1), so +/// `"head"` is an extension method, not `HEAD`. +fn custom_pattern(rule: &DynamicMessage) -> Option<(RouteMethod, String)> { + let custom = rule.get_field_by_name("custom")?; + let Value::Message(custom) = custom.as_ref() else { + return None; + }; + let kind = string_field(custom, "kind")?; + let path = string_field(custom, "path")?; + let method = if kind == "*" { + RouteMethod::Any + } else { + match Method::from_bytes(kind.as_bytes()) { + Ok(method) => RouteMethod::One(method), + Err(_) => { + tracing::warn!(%kind, %path, "custom HTTP rule kind is not a method token; skipping it"); + return None; + } + } + }; + Some((method, path)) +} + +/// A string field of `msg` that is present and non-empty. +fn string_field(msg: &DynamicMessage, name: &str) -> Option { + match msg.get_field_by_name(name)?.as_ref() { + Value::String(s) if !s.is_empty() => Some(s.clone()), + _ => None, + } +} + +#[cfg(test)] +mod tests; diff --git a/src/transcode/rule/tests.rs b/src/transcode/rule/tests.rs new file mode 100644 index 0000000..eb6203c --- /dev/null +++ b/src/transcode/rule/tests.rs @@ -0,0 +1,168 @@ +use super::*; +use prost_reflect::DescriptorPool; + +/// A standalone `HttpRule`-shaped descriptor (self-referential +/// `additional_bindings`, a `CustomHttpPattern` for `custom`) so the binding +/// parser can be tested without the google.api extension wiring. +fn http_rule_descriptor() -> prost_reflect::MessageDescriptor { + use prost_reflect::prost::Message; + use prost_reflect::prost_types::{ + field_descriptor_proto::{Label, Type}, + DescriptorProto, FieldDescriptorProto, FileDescriptorProto, FileDescriptorSet, + }; + + let str_field = |name: &str, num: i32| FieldDescriptorProto { + name: Some(name.to_string()), + number: Some(num), + label: Some(Label::Optional as i32), + r#type: Some(Type::String as i32), + ..Default::default() + }; + let message_field = + |name: &str, num: i32, label: Label, type_name: &str| FieldDescriptorProto { + name: Some(name.to_string()), + number: Some(num), + label: Some(label as i32), + r#type: Some(Type::Message as i32), + type_name: Some(type_name.to_string()), + ..Default::default() + }; + let custom = DescriptorProto { + name: Some("CustomHttpPattern".to_string()), + field: vec![str_field("kind", 1), str_field("path", 2)], + ..Default::default() + }; + let rule = DescriptorProto { + name: Some("HttpRule".to_string()), + field: vec![ + str_field("get", 2), + str_field("put", 3), + str_field("post", 4), + str_field("delete", 5), + str_field("patch", 6), + str_field("body", 7), + message_field("custom", 8, Label::Optional, ".gapi.CustomHttpPattern"), + str_field("response_body", 12), + message_field("additional_bindings", 11, Label::Repeated, ".gapi.HttpRule"), + ], + ..Default::default() + }; + let file = FileDescriptorProto { + name: Some("http.proto".to_string()), + package: Some("gapi".to_string()), + message_type: vec![rule, custom], + syntax: Some("proto3".to_string()), + ..Default::default() + }; + let fds = FileDescriptorSet { file: vec![file] }; + let pool = DescriptorPool::decode(fds.encode_to_vec().as_slice()).unwrap(); + pool.get_message_by_name("gapi.HttpRule").unwrap() +} + +/// An `HttpRule` with the `custom` pattern `{kind, path}`. +fn custom_rule(kind: &str, path: &str) -> DynamicMessage { + let rule_desc = http_rule_descriptor(); + let custom_desc = rule_desc + .parent_pool() + .get_message_by_name("gapi.CustomHttpPattern") + .unwrap(); + let mut custom = DynamicMessage::new(custom_desc); + custom.set_field_by_name("kind", Value::String(kind.into())); + custom.set_field_by_name("path", Value::String(path.into())); + let mut rule = DynamicMessage::new(rule_desc); + rule.set_field_by_name("custom", Value::Message(custom)); + rule +} + +#[test] +fn collect_bindings_reads_body_response_and_additional() { + let desc = http_rule_descriptor(); + + // additional_bindings entry: POST /v1/items with whole-body mapping. + let mut extra = DynamicMessage::new(desc.clone()); + extra.set_field_by_name("post", Value::String("/v1/items".into())); + extra.set_field_by_name("body", Value::String("*".into())); + + // primary rule: GET /v1/items/{id}, returns only the `result` subfield. + let mut rule = DynamicMessage::new(desc); + rule.set_field_by_name("get", Value::String("/v1/items/{id}".into())); + rule.set_field_by_name("response_body", Value::String("result".into())); + rule.set_field_by_name( + "additional_bindings", + Value::List(vec![Value::Message(extra)]), + ); + + let bindings = collect_bindings(&rule); + assert_eq!(bindings.len(), 2); + + // Primary: GET, no body, response_body = result. + assert_eq!(bindings[0].method, RouteMethod::One(Method::GET)); + assert_eq!(bindings[0].path, "/v1/items/{id}"); + assert_eq!(bindings[0].body, BodyMapping::None); + assert_eq!(bindings[0].response_body.as_deref(), Some("result")); + + // Additional: POST, whole-body mapping, no response_body. + assert_eq!(bindings[1].method, RouteMethod::One(Method::POST)); + assert_eq!(bindings[1].path, "/v1/items"); + assert_eq!(bindings[1].body, BodyMapping::Root); + assert_eq!(bindings[1].response_body, None); +} + +#[test] +fn custom_rule_binds_its_kind_as_the_method() { + let bindings = collect_bindings(&custom_rule("HEAD", "/v1/items/{id}")); + assert_eq!(bindings.len(), 1); + assert_eq!(bindings[0].method, RouteMethod::One(Method::HEAD)); + assert_eq!(bindings[0].path, "/v1/items/{id}"); + assert_eq!(bindings[0].method.as_str(), "HEAD"); +} + +#[test] +fn custom_rule_with_star_kind_binds_every_method() { + let bindings = collect_bindings(&custom_rule("*", "/v1/auth/verify")); + assert_eq!(bindings.len(), 1); + assert_eq!(bindings[0].method, RouteMethod::Any); + assert_eq!(bindings[0].method.as_str(), "*"); +} + +#[test] +fn custom_rule_kind_is_case_sensitive() { + // RFC 9110 §9.1: method tokens are case-sensitive, so a lowercase kind is + // an extension method of that exact spelling, never folded into HEAD. + let bindings = collect_bindings(&custom_rule("head", "/v1/items")); + assert_eq!(bindings.len(), 1); + assert_ne!(bindings[0].method, RouteMethod::One(Method::HEAD)); + assert_eq!(bindings[0].method.as_str(), "head"); +} + +#[test] +fn custom_rule_with_an_invalid_kind_is_skipped() { + // A space is not a token character; such a rule cannot be routed at all. + assert!(collect_bindings(&custom_rule("NOT A METHOD", "/v1/items")).is_empty()); +} + +#[test] +fn custom_rule_without_kind_or_path_is_skipped() { + assert!(collect_bindings(&custom_rule("", "/v1/items")).is_empty()); + assert!(collect_bindings(&custom_rule("HEAD", "")).is_empty()); +} + +#[test] +fn custom_rule_in_additional_bindings_is_collected() { + let mut rule = DynamicMessage::new(http_rule_descriptor()); + rule.set_field_by_name("get", Value::String("/v1/items".into())); + rule.set_field_by_name( + "additional_bindings", + Value::List(vec![Value::Message(custom_rule("OPTIONS", "/v1/items"))]), + ); + let bindings = collect_bindings(&rule); + let methods: Vec<&str> = bindings.iter().map(|b| b.method.as_str()).collect(); + assert_eq!(methods, ["GET", "OPTIONS"]); +} + +#[test] +fn rule_without_a_pattern_yields_no_binding() { + let mut rule = DynamicMessage::new(http_rule_descriptor()); + rule.set_field_by_name("body", Value::String("*".into())); + assert!(collect_bindings(&rule).is_empty()); +} diff --git a/src/transcode/tests.rs b/src/transcode/tests.rs index 4acafde..dcb8f02 100644 --- a/src/transcode/tests.rs +++ b/src/transcode/tests.rs @@ -1,91 +1,5 @@ use super::*; - -/// Build a standalone `HttpRule`-shaped descriptor (self-referential -/// `additional_bindings`) so the binding parser can be tested without the -/// google.api extension wiring. -fn http_rule_descriptor() -> prost_reflect::MessageDescriptor { - use prost_reflect::prost::Message; - use prost_reflect::prost_types::{ - field_descriptor_proto::{Label, Type}, - DescriptorProto, FieldDescriptorProto, FileDescriptorProto, FileDescriptorSet, - }; - - let str_field = |name: &str, num: i32| FieldDescriptorProto { - name: Some(name.to_string()), - number: Some(num), - label: Some(Label::Optional as i32), - r#type: Some(Type::String as i32), - ..Default::default() - }; - let rule = DescriptorProto { - name: Some("HttpRule".to_string()), - field: vec![ - str_field("get", 2), - str_field("put", 3), - str_field("post", 4), - str_field("delete", 5), - str_field("patch", 6), - str_field("body", 7), - str_field("response_body", 12), - FieldDescriptorProto { - name: Some("additional_bindings".to_string()), - number: Some(11), - label: Some(Label::Repeated as i32), - r#type: Some(Type::Message as i32), - type_name: Some(".gapi.HttpRule".to_string()), - ..Default::default() - }, - ], - ..Default::default() - }; - let file = FileDescriptorProto { - name: Some("http.proto".to_string()), - package: Some("gapi".to_string()), - message_type: vec![rule], - syntax: Some("proto3".to_string()), - ..Default::default() - }; - let fds = FileDescriptorSet { file: vec![file] }; - let pool = DescriptorPool::decode(fds.encode_to_vec().as_slice()).unwrap(); - pool.get_message_by_name("gapi.HttpRule").unwrap() -} - -#[test] -fn collect_bindings_reads_body_response_and_additional() { - let desc = http_rule_descriptor(); - - // additional_bindings entry: POST /v1/items with whole-body mapping. - let mut extra = DynamicMessage::new(desc.clone()); - extra.set_field_by_name("post", prost_reflect::Value::String("/v1/items".into())); - extra.set_field_by_name("body", prost_reflect::Value::String("*".into())); - - // primary rule: GET /v1/items/{id}, returns only the `result` subfield. - let mut rule = DynamicMessage::new(desc); - rule.set_field_by_name("get", prost_reflect::Value::String("/v1/items/{id}".into())); - rule.set_field_by_name( - "response_body", - prost_reflect::Value::String("result".into()), - ); - rule.set_field_by_name( - "additional_bindings", - prost_reflect::Value::List(vec![prost_reflect::Value::Message(extra)]), - ); - - let bindings = collect_bindings(&rule); - assert_eq!(bindings.len(), 2); - - // Primary: GET, no body, response_body = result. - assert!(matches!(bindings[0].http_method, HttpMethod::Get)); - assert_eq!(bindings[0].http_path, "/v1/items/{id}"); - assert_eq!(bindings[0].body, request::BodyMapping::None); - assert_eq!(bindings[0].response_body.as_deref(), Some("result")); - - // Additional: POST, whole-body mapping, no response_body. - assert!(matches!(bindings[1].http_method, HttpMethod::Post)); - assert_eq!(bindings[1].http_path, "/v1/items"); - assert_eq!(bindings[1].body, request::BodyMapping::Root); - assert_eq!(bindings[1].response_body, None); -} +use axum::routing::get; #[test] fn test_proto_path_to_axum() { diff --git a/tests/common/mod.rs b/tests/common/mod.rs index 5671ae8..f050d98 100644 --- a/tests/common/mod.rs +++ b/tests/common/mod.rs @@ -22,11 +22,28 @@ message HttpRule { string post = 4; string delete = 5; string patch = 6; + CustomHttpPattern custom = 8; } string body = 7; string response_body = 12; repeated HttpRule additional_bindings = 11; } +message CustomHttpPattern { + string kind = 1; + string path = 2; +} +"#; + +/// `google/api/httpbody.proto`. +const HTTPBODY_PROTO: &str = r#" +syntax = "proto3"; +package google.api; +import "google/protobuf/any.proto"; +message HttpBody { + string content_type = 1; + bytes data = 2; + repeated google.protobuf.Any extensions = 3; +} "#; const ANNOTATIONS_PROTO: &str = r#" @@ -39,8 +56,8 @@ extend google.protobuf.MethodOptions { } "#; -/// Serves the google.api sources and one test file from memory, and -/// descriptor.proto from protox's bundled Google files. +/// Serves the google.api sources and one test file from memory, and the +/// `google/protobuf` files from protox's bundled Google files. struct TestProtos { name: &'static str, source: &'static str, @@ -51,6 +68,7 @@ impl protox::file::FileResolver for TestProtos { let source = match name { "google/api/http.proto" => HTTP_PROTO, "google/api/annotations.proto" => ANNOTATIONS_PROTO, + "google/api/httpbody.proto" => HTTPBODY_PROTO, _ if name == self.name => self.source, _ => return protox::file::GoogleFileResolver::new().open_file(name), }; diff --git a/tests/hooks.rs b/tests/hooks.rs index 8bcb7a0..ebba571 100644 --- a/tests/hooks.rs +++ b/tests/hooks.rs @@ -155,6 +155,51 @@ async fn verify_endpoint_is_backed_by_the_decider() { assert_eq!(denied.status(), StatusCode::UNAUTHORIZED); } +#[tokio::test] +async fn options_verify_request_is_decided_not_answered_by_cors() { + // A forward-auth sub-request carries the original request's method. An + // OPTIONS without `Access-Control-Request-Method` is not a CORS preflight + // (Fetch standard, CORS-preflight request), so the CORS layer must not + // answer it with 200: that would let an unauthenticated OPTIONS through + // the gate. It reaches the decider, which denies it. + let app = server().router().unwrap(); + let resp = app + .clone() + .oneshot( + axum::http::Request::builder() + .method(Method::OPTIONS) + .uri("/auth/verify") + .header("origin", "https://app.example.com") + .header("x-forwarded-method", "OPTIONS") + .header("x-forwarded-uri", "/v1/things") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + // Still a CORS response: the browser may read it. + assert!(resp.headers().contains_key("access-control-allow-origin")); + + // A real preflight is still answered by the CORS layer. + let preflight = app + .oneshot( + axum::http::Request::builder() + .method(Method::OPTIONS) + .uri("/auth/verify") + .header("origin", "https://app.example.com") + .header("access-control-request-method", "POST") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(preflight.status(), StatusCode::OK); + assert!(preflight + .headers() + .contains_key("access-control-allow-methods")); +} + #[tokio::test] async fn verify_redirect_becomes_401_with_location() { let app = server().router().unwrap(); diff --git a/tests/upstream_controls.rs b/tests/upstream_controls.rs new file mode 100644 index 0000000..274280d --- /dev/null +++ b/tests/upstream_controls.rs @@ -0,0 +1,887 @@ +//! The upstream decides the HTTP answer: response metadata becomes response +//! headers, `x-http-code` sets the status of a successful unary call, +//! `google.api.HttpBody` carries a raw body both ways, and `custom` rules bind +//! any method. +//! +//! Runs the proxy (through its public `ProxyServer`) in front of a real tonic +//! gRPC server, so metadata, trailers and trailers-only errors go over the +//! actual HTTP/2 stream. + +mod common; + +use std::convert::Infallible; +use std::future::{ready, Ready}; +use std::pin::Pin; +use std::task::{Context, Poll}; + +use axum::body::Body; +use bytes::Bytes; +use futures::stream::BoxStream; +use http::{HeaderMap, HeaderName, Method, StatusCode}; +use http_body_util::BodyExt; +use prost::Message as _; +use prost_reflect::{DescriptorPool, DynamicMessage, Value as PbValue}; +use serde_json::{json, Value}; +use structured_proxy::transcode::codec::DynamicCodec; +use structured_proxy::transcode::error::ErrorDetailsPolicy; +use structured_proxy::ProxyServer; +use tonic::metadata::{AsciiMetadataValue, BinaryMetadataValue}; +use tower::ServiceExt; + +// --- descriptors ------------------------------------------------------------ + +const CONTROLS_PROTO: &str = r#" +syntax = "proto3"; +package test.v1; +import "google/api/annotations.proto"; +import "google/api/httpbody.proto"; + +message Req { + string name = 1; +} +message Reply { + string name = 1; +} +message TokenRequest { + string grant_type = 1; +} +message Upload { + string name = 1; + google.api.HttpBody file = 2; +} + +service Controls { + rpc Get(Req) returns (Reply) { + option (google.api.http) = { get: "/v1/things/{name}" }; + } + rpc Trailers(Req) returns (Reply) { + option (google.api.http) = { get: "/v1/trailers" }; + } + rpc Token(TokenRequest) returns (google.api.HttpBody) { + option (google.api.http) = { post: "/v1/token" body: "*" }; + } + rpc Authorize(Req) returns (google.api.HttpBody) { + option (google.api.http) = { get: "/v1/authorize" }; + } + rpc Jwks(Req) returns (google.api.HttpBody) { + option (google.api.http) = { get: "/v1/jwks" }; + } + rpc Echo(google.api.HttpBody) returns (google.api.HttpBody) { + option (google.api.http) = { post: "/v1/echo" body: "*" }; + } + rpc Put(Upload) returns (Reply) { + option (google.api.http) = { put: "/v1/uploads/{name}" body: "file" }; + } + rpc Probe(Req) returns (Reply) { + option (google.api.http) = { custom: { kind: "HEAD" path: "/v1/probe/{name}" } }; + } + rpc Verify(Req) returns (Reply) { + option (google.api.http) = { custom: { kind: "*" path: "/v1/verify" } }; + } + rpc Dav(Req) returns (Reply) { + option (google.api.http) = { + custom: { kind: "PROPFIND" path: "/v1/dav" } + additional_bindings { get: "/v1/dav" } + }; + } + rpc Watch(Req) returns (stream Reply) { + option (google.api.http) = { get: "/v1/things/{name}/watch" }; + } + rpc Download(Req) returns (stream google.api.HttpBody) { + option (google.api.http) = { get: "/v1/files/{name}" }; + } +} +"#; + +fn pool() -> DescriptorPool { + common::compile("test/v1/controls.proto", CONTROLS_PROTO) +} + +// --- upstream --------------------------------------------------------------- + +/// A message of `type_name` with string / bytes fields set. +fn message(pool: &DescriptorPool, type_name: &str, fields: &[(&str, PbValue)]) -> DynamicMessage { + let mut msg = DynamicMessage::new(pool.get_message_by_name(type_name).unwrap()); + for (name, value) in fields { + msg.set_field_by_name(name, value.clone()); + } + msg +} + +fn string(msg: &DynamicMessage, field: &str) -> String { + match msg.get_field_by_name(field).as_deref() { + Some(PbValue::String(s)) => s.clone(), + _ => String::new(), + } +} + +fn bytes_field(msg: &DynamicMessage, field: &str) -> Bytes { + match msg.get_field_by_name(field).as_deref() { + Some(PbValue::Bytes(b)) => b.clone(), + _ => Bytes::new(), + } +} + +fn reply(pool: &DescriptorPool, name: &str) -> DynamicMessage { + message( + pool, + "test.v1.Reply", + &[("name", PbValue::String(name.into()))], + ) +} + +fn http_body(pool: &DescriptorPool, content_type: &str, data: &'static [u8]) -> DynamicMessage { + message( + pool, + "google.api.HttpBody", + &[ + ("content_type", PbValue::String(content_type.into())), + ("data", PbValue::Bytes(Bytes::from_static(data))), + ], + ) +} + +fn ascii(value: &str) -> AsciiMetadataValue { + AsciiMetadataValue::try_from(value).unwrap() +} + +/// `UNAUTHENTICATED` whose trailers-only response carries a `WWW-Authenticate` +/// challenge (RFC 6750 §3). +fn invalid_token() -> tonic::Status { + let mut status = tonic::Status::unauthenticated("token expired"); + status.metadata_mut().insert( + "www-authenticate", + ascii("Bearer error=\"invalid_token\", error_description=\"expired\""), + ); + status +} + +/// Unary handlers, by RPC name. +#[derive(Clone)] +struct Unary { + pool: DescriptorPool, + rpc: String, +} + +impl tonic::server::UnaryService for Unary { + type Response = DynamicMessage; + type Future = Ready, tonic::Status>>; + + fn call(&mut self, request: tonic::Request) -> Self::Future { + let pool = &self.pool; + let req = request.into_inner(); + let result = match self.rpc.as_str() { + "Get" => get(pool, &string(&req, "name")), + "Token" => Ok(token(pool, &string(&req, "grant_type"))), + "Authorize" => { + let mut resp = tonic::Response::new(http_body(pool, "", b"")); + resp.metadata_mut().insert("x-http-code", ascii("302")); + resp.metadata_mut().insert( + "location", + ascii("https://rp.example/cb?code=abc&state=xyz"), + ); + Ok(resp) + } + "Jwks" => Ok(tonic::Response::new(http_body( + pool, + "application/jwk-set+json", + br#"{"keys":[]}"#, + ))), + // The raw body and its content type come back as they arrived. + "Echo" => Ok(tonic::Response::new(req)), + "Put" => { + let file = match req.get_field_by_name("file").as_deref() { + Some(PbValue::Message(file)) => file.clone(), + _ => panic!("upload without a file"), + }; + let summary = format!( + "{}|{}|{}", + string(&req, "name"), + string(&file, "content_type"), + String::from_utf8(bytes_field(&file, "data").to_vec()).unwrap() + ); + Ok(tonic::Response::new(reply(pool, &summary))) + } + "Probe" => { + let mut resp = tonic::Response::new(reply(pool, "probed")); + resp.metadata_mut() + .insert("x-probe", ascii(&string(&req, "name"))); + Ok(resp) + } + "Verify" => { + let mut resp = tonic::Response::new(reply(pool, "verified")); + resp.metadata_mut().insert("x-user-id", ascii("u-42")); + Ok(resp) + } + "Dav" => Ok(tonic::Response::new(reply(pool, "dav"))), + other => panic!("unexpected unary RPC {other}"), + }; + ready(result) + } +} + +/// `Get`: behaviour chosen by `name`. +fn get( + pool: &DescriptorPool, + name: &str, +) -> Result, tonic::Status> { + let mut resp = tonic::Response::new(reply(pool, name)); + let md = resp.metadata_mut(); + match name { + "plain" => {} + "meta" => { + md.insert("cache-control", ascii("no-store")); + md.append("set-cookie", ascii("a=1")); + md.append("set-cookie", ascii("b=2")); + md.insert("x-debug", ascii("internal")); + md.insert("grpc-extra", ascii("1")); + md.insert_bin("x-trace-bin", BinaryMetadataValue::from_bytes(b"\x00\x01")); + } + "created" => { + md.insert("x-http-code", ascii("201")); + md.insert("location", ascii("/v1/things/created")); + } + "no-content" => { + md.insert("x-http-code", ascii("204")); + } + "bad-code" => { + md.insert("x-http-code", ascii("abc")); + md.insert("x-leak", ascii("must not reach the client")); + } + "unauth" => return Err(invalid_token()), + "corrupt-unauth" => { + // Details that are not a google.rpc.Status: a broken error, whose + // metadata must not ride on the generic INTERNAL either. + let mut status = tonic::Status::with_details( + tonic::Code::Unauthenticated, + "nope", + Bytes::from_static(b"\xff\xff\xff"), + ); + status + .metadata_mut() + .insert("www-authenticate", ascii("Bearer")); + return Err(status); + } + other => panic!("unexpected Get name {other}"), + } + Ok(resp) +} + +/// RFC 6749 token endpoint: success, or the §5.2 error body with 400; both +/// with `Cache-Control: no-store` (§5.1). +fn token(pool: &DescriptorPool, grant_type: &str) -> tonic::Response { + let (body, code): (&'static [u8], Option<&str>) = match grant_type { + "authorization_code" => (br#"{"access_token":"at","token_type":"Bearer"}"#, None), + _ => ( + br#"{"error":"invalid_grant","error_description":"code expired"}"#, + Some("400"), + ), + }; + let mut resp = tonic::Response::new(http_body(pool, "application/json;charset=UTF-8", body)); + resp.metadata_mut() + .insert("cache-control", ascii("no-store")); + resp.metadata_mut().insert("pragma", ascii("no-cache")); + if let Some(code) = code { + resp.metadata_mut().insert("x-http-code", ascii(code)); + } + resp +} + +type ReplyStream = BoxStream<'static, Result>; + +/// Server-streaming handlers, by RPC name. +#[derive(Clone)] +struct Streaming { + pool: DescriptorPool, + rpc: String, +} + +impl tonic::server::ServerStreamingService for Streaming { + type Response = DynamicMessage; + type ResponseStream = ReplyStream; + type Future = Ready, tonic::Status>>; + + fn call(&mut self, request: tonic::Request) -> Self::Future { + let pool = &self.pool; + let name = string(request.get_ref(), "name"); + let items: Vec> = + match (self.rpc.as_str(), name.as_str()) { + ("Watch", "denied") => return ready(Err(invalid_token())), + ("Watch", _) => vec![Ok(reply(pool, "one")), Ok(reply(pool, "two"))], + ("Download", "broken") => vec![ + Ok(http_body(pool, "text/csv", b"a,b\n")), + Err(tonic::Status::internal("disk failed")), + ], + ("Download", "empty") => Vec::new(), + ("Download", _) => vec![ + Ok(http_body(pool, "text/csv", b"a,b\n")), + // Only the first message's content type counts. + Ok(http_body(pool, "text/plain", b"1,2\n")), + Ok(http_body(pool, "", b"")), + Ok(http_body(pool, "", b"3,4\n")), + ], + (rpc, _) => panic!("unexpected streaming RPC {rpc}"), + }; + let mut resp = tonic::Response::new(Box::pin(futures::stream::iter(items)) as ReplyStream); + resp.metadata_mut().insert("x-stream", ascii("1")); + resp.metadata_mut() + .insert("cache-control", ascii("max-age=60")); + ready(Ok(resp)) + } +} + +/// A successful `Trailers` answer written by hand, since tonic's server API +/// cannot set trailers on success: `x-both` in the headers and the trailers, +/// `x-trailer` only in the trailers, and gRPC's own trailer keys. +fn trailers_response(pool: &DescriptorPool) -> http::Response { + let payload = reply(pool, "trailed").encode_to_vec(); + let mut frame = Vec::with_capacity(5 + payload.len()); + frame.push(0); + frame.extend_from_slice(&u32::try_from(payload.len()).unwrap().to_be_bytes()); + frame.extend_from_slice(&payload); + let mut trailers = HeaderMap::new(); + trailers.insert("grpc-status", "0".parse().unwrap()); + trailers.insert("x-both", "trailer".parse().unwrap()); + trailers.insert("x-trailer", "t".parse().unwrap()); + let frames: Vec, Infallible>> = vec![ + Ok(http_body::Frame::data(Bytes::from(frame))), + Ok(http_body::Frame::trailers(trailers)), + ]; + let body = http_body_util::StreamBody::new(futures::stream::iter(frames)); + http::Response::builder() + .header("content-type", "application/grpc") + .header("x-both", "initial") + .body(tonic::body::Body::new(body)) + .unwrap() +} + +/// The `test.v1.Controls` gRPC service, dispatching by method path. +#[derive(Clone)] +struct Controls { + pool: DescriptorPool, +} + +impl tonic::server::NamedService for Controls { + const NAME: &'static str = "test.v1.Controls"; +} + +impl tower::Service> for Controls { + type Response = http::Response; + type Error = Infallible; + type Future = + Pin> + Send>>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, req: http::Request) -> Self::Future { + let pool = self.pool.clone(); + Box::pin(async move { + let rpc = req + .uri() + .path() + .strip_prefix("/test.v1.Controls/") + .unwrap() + .to_owned(); + if rpc == "Trailers" { + // Read the request to its end before answering. + let _request = req.into_body().collect().await.unwrap(); + return Ok(trailers_response(&pool)); + } + let method = pool + .get_service_by_name("test.v1.Controls") + .unwrap() + .methods() + .find(|m| m.name() == rpc) + .unwrap(); + let mut grpc = tonic::server::Grpc::new(DynamicCodec::new(method.input())); + let resp = if method.is_server_streaming() { + grpc.server_streaming(Streaming { pool, rpc }, req).await + } else { + grpc.unary(Unary { pool, rpc }, req).await + }; + Ok(resp) + }) + } +} + +// --- proxy harness ---------------------------------------------------------- + +/// The proxy router in front of a fresh upstream, built by `configure`. +async fn proxy_with(configure: impl FnOnce(ProxyServer) -> ProxyServer) -> axum::Router { + let pool = pool(); + let upstream = common::serve(Controls { pool: pool.clone() }).await; + let server = ProxyServer::from_yaml_str(&format!("upstream:\n default: \"{upstream}\"\n")) + .unwrap() + .with_descriptors(pool); + configure(server).router().unwrap() +} + +/// The proxy router in front of a fresh upstream, with default settings. +async fn proxy() -> axum::Router { + let pool = pool(); + let upstream = common::serve(Controls { pool: pool.clone() }).await; + common::proxy(&upstream, pool, ErrorDetailsPolicy::default()) +} + +/// Send `request`; returns the status, the headers and the raw body. +async fn send(app: &axum::Router, request: http::Request) -> (StatusCode, HeaderMap, Bytes) { + let resp = app.clone().oneshot(request).await.unwrap(); + let (parts, body) = resp.into_parts(); + let body = axum::body::to_bytes(body, usize::MAX).await.unwrap(); + (parts.status, parts.headers, body) +} + +async fn call(app: &axum::Router, method: Method, path: &str) -> (StatusCode, HeaderMap, Bytes) { + let request = http::Request::builder() + .method(method) + .uri(path) + .body(Body::empty()) + .unwrap(); + send(app, request).await +} + +fn values<'a>(headers: &'a HeaderMap, name: &str) -> Vec<&'a str> { + headers + .get_all(name) + .iter() + .map(|v| v.to_str().unwrap()) + .collect() +} + +fn json_body(body: &Bytes) -> Value { + serde_json::from_slice(body).unwrap() +} + +// --- response metadata → headers --------------------------------------------- + +#[tokio::test] +async fn unary_metadata_becomes_response_headers() { + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/things/meta").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(json_body(&body), json!({"name": "meta"})); + assert_eq!(values(&headers, "cache-control"), ["no-store"]); + assert_eq!(values(&headers, "set-cookie"), ["a=1", "b=2"]); + assert_eq!(values(&headers, "x-debug"), ["internal"]); + // gRPC's own keys and the binary encoding stay behind; the content type + // is the transcoder's. + assert!(headers.get("grpc-extra").is_none()); + assert!(headers.get("x-trace-bin").is_none()); + assert!(headers.get("grpc-status").is_none()); + assert!(headers.get("grpc-accept-encoding").is_none()); + assert_eq!(values(&headers, "content-type"), ["application/json"]); +} + +#[tokio::test] +async fn plain_answer_adds_no_upstream_headers() { + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/things/plain").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(json_body(&body), json!({"name": "plain"})); + assert!(headers.get("cache-control").is_none()); + assert!(headers.get("x-http-code").is_none()); +} + +#[tokio::test] +async fn trailers_of_a_successful_call_become_headers_after_initial_metadata() { + // A key in both the initial metadata and the trailers keeps both values, + // initial first; gRPC's trailer keys stay behind. + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/trailers").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(json_body(&body), json!({"name": "trailed"})); + assert_eq!(values(&headers, "x-both"), ["initial", "trailer"]); + assert_eq!(values(&headers, "x-trailer"), ["t"]); + assert!(headers.get("grpc-status").is_none()); +} + +#[tokio::test] +async fn deny_list_from_the_builder_drops_its_keys() { + let app = proxy_with(|server| { + server.with_denied_response_headers([HeaderName::from_static("x-debug")]) + }) + .await; + let (_, headers, _) = call(&app, Method::GET, "/v1/things/meta").await; + assert!(headers.get("x-debug").is_none()); + assert_eq!(values(&headers, "cache-control"), ["no-store"]); +} + +#[tokio::test] +async fn deny_list_from_yaml_drops_its_keys() { + let pool = pool(); + let upstream = common::serve(Controls { pool: pool.clone() }).await; + let app = ProxyServer::from_yaml_str(&format!( + "upstream:\n default: \"{upstream}\"\nresponse_headers:\n deny: [\"X-Debug\", \"set-cookie\"]\n" + )) + .unwrap() + .with_descriptors(pool) + .router() + .unwrap(); + let (_, headers, _) = call(&app, Method::GET, "/v1/things/meta").await; + assert!(headers.get("x-debug").is_none()); + assert!(headers.get("set-cookie").is_none()); + assert_eq!(values(&headers, "cache-control"), ["no-store"]); +} + +#[tokio::test] +async fn trailers_only_error_carries_its_metadata() { + // RFC 6750 §3: the 401 carries the upstream's `WWW-Authenticate`. + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/things/unauth").await; + assert_eq!(status, StatusCode::UNAUTHORIZED); + assert_eq!( + values(&headers, "www-authenticate"), + ["Bearer error=\"invalid_token\", error_description=\"expired\""] + ); + assert_eq!(json_body(&body)["error"], "UNAUTHENTICATED"); + assert_eq!(values(&headers, "content-type"), ["application/json"]); +} + +#[tokio::test] +async fn malformed_error_status_carries_no_upstream_metadata() { + // The answer is the generic INTERNAL, so nothing of the broken error, + // headers included, reaches the client. + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/things/corrupt-unauth").await; + assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); + assert_eq!(json_body(&body)["error"], "INTERNAL"); + assert!(headers.get("www-authenticate").is_none()); +} + +// --- x-http-code ------------------------------------------------------------ + +#[tokio::test] +async fn http_code_sets_the_status_of_a_successful_call() { + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/things/created").await; + assert_eq!(status, StatusCode::CREATED); + assert_eq!(values(&headers, "location"), ["/v1/things/created"]); + assert!(headers.get("x-http-code").is_none()); + assert_eq!(json_body(&body), json!({"name": "created"})); +} + +#[tokio::test] +async fn http_code_204_answers_without_content() { + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/things/no-content").await; + assert_eq!(status, StatusCode::NO_CONTENT); + assert!(body.is_empty()); + assert!(headers.get("content-type").is_none()); +} + +#[tokio::test] +async fn invalid_http_code_is_a_malformed_upstream_internal() { + // Never a partial response: no other upstream header rides along. + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/things/bad-code").await; + assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); + assert_eq!( + json_body(&body), + json!({ + "error": "INTERNAL", + "message": "upstream returned a malformed response", + "code": 13, + "details": [] + }) + ); + assert!(headers.get("x-leak").is_none()); + assert!(headers.get("x-http-code").is_none()); +} + +#[tokio::test] +async fn redirect_with_location_and_an_empty_body() { + // RFC 6749 §4.1.2: the authorization endpoint answers 302 + Location. + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/authorize").await; + assert_eq!(status, StatusCode::FOUND); + assert_eq!( + values(&headers, "location"), + ["https://rp.example/cb?code=abc&state=xyz"] + ); + assert!(body.is_empty()); + assert!(headers.get("content-type").is_none()); +} + +// --- google.api.HttpBody ---------------------------------------------------- + +async fn post_form( + app: &axum::Router, + path: &str, + form: &'static str, +) -> (StatusCode, HeaderMap, Bytes) { + let request = http::Request::post(path) + .header("content-type", "application/x-www-form-urlencoded") + .body(Body::from(form)) + .unwrap(); + send(app, request).await +} + +#[tokio::test] +async fn token_error_is_the_rfc_6749_body_with_400_and_no_store() { + // RFC 6749 §5.2: 400, the upstream's own JSON body, `Cache-Control: + // no-store` (§5.1). Not the transcoder's error body. + let app = proxy().await; + let (status, headers, body) = post_form(&app, "/v1/token", "grant_type=refresh_token").await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(values(&headers, "cache-control"), ["no-store"]); + assert_eq!(values(&headers, "pragma"), ["no-cache"]); + assert_eq!( + values(&headers, "content-type"), + ["application/json;charset=UTF-8"] + ); + assert_eq!( + json_body(&body), + json!({"error": "invalid_grant", "error_description": "code expired"}) + ); +} + +#[tokio::test] +async fn token_success_is_the_raw_json_with_no_store() { + let app = proxy().await; + let (status, headers, body) = + post_form(&app, "/v1/token", "grant_type=authorization_code").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(values(&headers, "cache-control"), ["no-store"]); + assert_eq!(&body[..], br#"{"access_token":"at","token_type":"Bearer"}"#); +} + +#[tokio::test] +async fn http_body_response_keeps_the_content_type_and_bytes() { + // RFC 7517 §8.5: a JWK Set is served as application/jwk-set+json. + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/jwks").await; + assert_eq!(status, StatusCode::OK); + assert_eq!( + values(&headers, "content-type"), + ["application/jwk-set+json"] + ); + assert_eq!(&body[..], br#"{"keys":[]}"#); +} + +#[tokio::test] +async fn http_body_request_receives_the_raw_body_and_content_type() { + // Bytes that are neither JSON nor UTF-8 arrive untouched, with the full + // Content-Type value (parameters included). + let app = proxy().await; + let raw: &'static [u8] = b"\x89PNG\r\n\x1a\n\x00\xff"; + let request = http::Request::post("/v1/echo") + .header("content-type", "image/png; name=x") + .body(Body::from(raw)) + .unwrap(); + let (status, headers, body) = send(&app, request).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(values(&headers, "content-type"), ["image/png; name=x"]); + assert_eq!(&body[..], raw); +} + +#[tokio::test] +async fn http_body_request_without_a_body_is_empty() { + let app = proxy().await; + let request = http::Request::post("/v1/echo").body(Body::empty()).unwrap(); + let (status, headers, body) = send(&app, request).await; + assert_eq!(status, StatusCode::OK); + assert!(body.is_empty()); + assert!(headers.get("content-type").is_none()); +} + +#[tokio::test] +async fn http_body_field_receives_the_raw_body_next_to_path_fields() { + let app = proxy().await; + let request = http::Request::put("/v1/uploads/report") + .header("content-type", "text/csv") + .body(Body::from("a,b")) + .unwrap(); + let (status, body) = common::send(&app, request).await; + assert_eq!(status, StatusCode::OK); + assert_eq!( + serde_json::from_str::(&body).unwrap(), + json!({"name": "report|text/csv|a,b"}) + ); +} + +#[tokio::test] +async fn http_body_request_with_a_non_ascii_content_type_is_rejected() { + // HttpBody.content_type is a proto string; bytes that are not visible + // ASCII are refused before the upstream is called. + let app = proxy().await; + let request = http::Request::post("/v1/echo") + .header( + "content-type", + http::HeaderValue::from_bytes(b"text/plain; x=\xe9").unwrap(), + ) + .body(Body::from("x")) + .unwrap(); + let (status, _, body) = send(&app, request).await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(json_body(&body)["error"], "INVALID_ARGUMENT"); +} + +// --- server streaming ------------------------------------------------------- + +#[tokio::test] +async fn streaming_initial_metadata_becomes_headers() { + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/things/x/watch").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(values(&headers, "x-stream"), ["1"]); + assert_eq!(values(&headers, "content-type"), ["application/x-ndjson"]); + let lines: Vec = std::str::from_utf8(&body) + .unwrap() + .lines() + .map(|line| serde_json::from_str(line).unwrap()) + .collect(); + assert_eq!(lines, [json!({"name": "one"}), json!({"name": "two"})]); +} + +#[tokio::test] +async fn sse_keeps_its_own_cache_control_over_the_upstream_one() { + // What the proxy writes describes the body it writes: SSE must not be + // cached, whatever the upstream asked for. + let app = proxy().await; + let request = http::Request::get("/v1/things/x/watch") + .header("accept", "text/event-stream") + .body(Body::empty()) + .unwrap(); + let (status, headers, _) = send(&app, request).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(values(&headers, "content-type"), ["text/event-stream"]); + assert_eq!(values(&headers, "cache-control"), ["no-cache"]); + assert_eq!(values(&headers, "x-stream"), ["1"]); +} + +#[tokio::test] +async fn refused_stream_carries_its_metadata() { + let app = proxy().await; + let (status, headers, _) = call(&app, Method::GET, "/v1/things/denied/watch").await; + assert_eq!(status, StatusCode::UNAUTHORIZED); + assert!(values(&headers, "www-authenticate")[0].starts_with("Bearer error=\"invalid_token\"")); +} + +#[tokio::test] +async fn streaming_http_body_is_chunked_raw_data() { + // Content type from the first message; every message's data in order. + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/files/report").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(values(&headers, "content-type"), ["text/csv"]); + assert_eq!(values(&headers, "x-stream"), ["1"]); + assert_eq!(&body[..], b"a,b\n1,2\n3,4\n"); +} + +#[tokio::test] +async fn streaming_http_body_ignores_sse_negotiation() { + let app = proxy().await; + let request = http::Request::get("/v1/files/report") + .header("accept", "text/event-stream") + .body(Body::empty()) + .unwrap(); + let (_, headers, body) = send(&app, request).await; + assert_eq!(values(&headers, "content-type"), ["text/csv"]); + assert_eq!(&body[..], b"a,b\n1,2\n3,4\n"); +} + +#[tokio::test] +async fn empty_streaming_http_body_is_an_empty_ok() { + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/files/empty").await; + assert_eq!(status, StatusCode::OK); + assert!(body.is_empty()); + assert!(headers.get("content-type").is_none()); +} + +#[tokio::test] +async fn streaming_http_body_failing_mid_stream_aborts_the_body() { + // A raw body has no in-band error frame: the transfer is cut short so the + // client cannot take the partial file for a complete one. + let app = proxy().await; + let request = http::Request::get("/v1/files/broken") + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(request).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + assert!(axum::body::to_bytes(resp.into_body(), usize::MAX) + .await + .is_err()); +} + +// --- custom rules ------------------------------------------------------------ + +#[tokio::test] +async fn custom_head_rule_routes_head_to_the_rpc() { + let app = proxy().await; + let (status, headers, body) = call(&app, Method::HEAD, "/v1/probe/disk").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(values(&headers, "x-probe"), ["disk"]); + assert!(body.is_empty()); + // Only HEAD is bound on that path. + let (status, _, _) = call(&app, Method::GET, "/v1/probe/disk").await; + assert_eq!(status, StatusCode::METHOD_NOT_ALLOWED); +} + +#[tokio::test] +async fn star_rule_routes_every_method_to_the_rpc() { + // A forward-auth sub-request arrives with the original request's method. + let app = proxy().await; + for method in [ + Method::GET, + Method::POST, + Method::PUT, + Method::DELETE, + Method::PATCH, + Method::OPTIONS, + Method::from_bytes(b"PROPFIND").unwrap(), + ] { + let (status, headers, body) = call(&app, method.clone(), "/v1/verify").await; + assert_eq!(status, StatusCode::OK, "{method}"); + assert_eq!(values(&headers, "x-user-id"), ["u-42"], "{method}"); + assert_eq!(json_body(&body), json!({"name": "verified"}), "{method}"); + } +} + +#[tokio::test] +async fn extension_method_rule_routes_next_to_a_standard_one() { + let app = proxy().await; + let propfind = Method::from_bytes(b"PROPFIND").unwrap(); + let (status, _, body) = call(&app, propfind, "/v1/dav").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(json_body(&body), json!({"name": "dav"})); + let (status, _, body) = call(&app, Method::GET, "/v1/dav").await; + assert_eq!(status, StatusCode::OK); + assert_eq!(json_body(&body), json!({"name": "dav"})); + // A method nobody binds is 405 with the full Allow list (RFC 9110 + // §15.5.6). + let (status, headers, _) = call(&app, Method::from_bytes(b"MKCOL").unwrap(), "/v1/dav").await; + assert_eq!(status, StatusCode::METHOD_NOT_ALLOWED); + assert_eq!(values(&headers, "allow"), ["PROPFIND, GET, HEAD"]); + let (status, _, _) = call(&app, Method::DELETE, "/v1/dav").await; + assert_eq!(status, StatusCode::METHOD_NOT_ALLOWED); +} + +#[tokio::test] +async fn star_rule_collides_with_another_method_on_its_path() { + // `*` claims every method, so another binding on the same path is a + // conflict the router refuses up front instead of an axum panic. + const CLASH_PROTO: &str = r#" +syntax = "proto3"; +package test.v1; +import "google/api/annotations.proto"; +message Req { string name = 1; } +service Clash { + rpc Any(Req) returns (Req) { + option (google.api.http) = { custom: { kind: "*" path: "/v1/x" } }; + } + rpc Get(Req) returns (Req) { + option (google.api.http) = { get: "/v1/x" }; + } +} +"#; + let pool = common::compile("test/v1/clash.proto", CLASH_PROTO); + let err = ProxyServer::from_yaml_str("upstream:\n default: \"http://127.0.0.1:1\"\n") + .unwrap() + .with_descriptors(pool) + .router() + .expect_err("a `*` rule next to another method must be rejected"); + assert!(err.to_string().contains("more than one endpoint"), "{err}"); +} From 10fe639263926d358c19ca205bd50278d05fd66e Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Sun, 27 Sep 2026 04:10:57 +0300 Subject: [PATCH 2/3] fix(transcode): tighten preflight, raw bodies and error headers MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - CORS: a preflight needs Origin as well as Access-Control-Request-Method; any other OPTIONS reaches its route - Raw HttpBody field: a query key naming the body field no longer binds into it (the body wins over the query, as for a parsed body) - Aliases run through the same path-template conversion as the route they alias ({path=**} becomes {*path}) - response_body segments resolve through proto field names, not the serialized JSON keys, so multi-word fields no longer yield null - 205 Reset Content carries no content or Content-Type (RFC 9110 §15.3.6), like 204 and 304 - A call that fails after the upstream sent response headers keeps their metadata on the error response, ahead of the failure's own - OpenAPI registers the schema of message-typed query parameters Each fix carries a regression test that failed before it. --- README.md | 13 +-- src/cors.rs | 20 ++--- src/cors/tests.rs | 10 +++ src/openapi.rs | 2 + src/openapi/tests.rs | 24 ++++++ src/transcode/mod.rs | 140 ++++++++++++++++++++++---------- src/transcode/response.rs | 9 +- src/transcode/response/tests.rs | 8 +- src/transcode/tests.rs | 92 +++++++++++++++++++++ tests/upstream_controls.rs | 78 +++++++++++++++++- 10 files changed, 333 insertions(+), 63 deletions(-) diff --git a/README.md b/README.md index 18279d8..4dbad42 100644 --- a/README.md +++ b/README.md @@ -440,8 +440,9 @@ grpc-gateway do, so the same service works behind any of them. 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, so a `401` -carries its `WWW-Authenticate`), and the initial metadata of a server-streaming +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: @@ -469,8 +470,9 @@ 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 `{"error": "INTERNAL", "code": 13, "message": "upstream returned a malformed response", "details": []}` -(500), with nothing else of the upstream's answer. `204` and `304` are sent -without a body or `Content-Type` (RFC 9110 §15.3.5, §15.4.5). Errors keep the +(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. @@ -511,7 +513,8 @@ 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 `Access-Control-Request-Method`) is +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. ## Library Usage diff --git a/src/cors.rs b/src/cors.rs index 13bf480..533bb0c 100644 --- a/src/cors.rs +++ b/src/cors.rs @@ -4,13 +4,14 @@ //! no route could serve `OPTIONS`: not a `custom` rule, not a `*` rule, not the //! forward-auth endpoint, whose sub-request carries the original request's //! method and would be allowed through by that `200`. The Fetch standard (§3.2.2, -//! CORS-preflight request) makes a preflight an `OPTIONS` request with -//! `Access-Control-Request-Method`; any other `OPTIONS` is an ordinary request. Such a request passes the CORS layer under a stand-in -//! method, so it gets the response headers of an ordinary CORS request, and has -//! its method restored before anything else sees it. +//! CORS request and CORS-preflight request) makes a preflight an `OPTIONS` +//! request carrying both `Origin` and `Access-Control-Request-Method`; any other +//! `OPTIONS` is an ordinary request. Such a request passes the CORS layer under a +//! stand-in method, so it gets the response headers of an ordinary CORS request, +//! and has its method restored before anything else sees it. use axum::extract::Request; -use axum::http::header::ACCESS_CONTROL_REQUEST_METHOD; +use axum::http::header::{ACCESS_CONTROL_REQUEST_METHOD, ORIGIN}; use axum::http::Method; use axum::middleware::{self, Next}; use axum::response::Response; @@ -42,11 +43,10 @@ where /// Outermost: hide an ordinary `OPTIONS` from the CORS layer. async fn disguise_options(mut request: Request, next: Next) -> Response { - if request.method() == Method::OPTIONS - && !request - .headers() - .contains_key(ACCESS_CONTROL_REQUEST_METHOD) - { + let headers = request.headers(); + let preflight = + headers.contains_key(ORIGIN) && headers.contains_key(ACCESS_CONTROL_REQUEST_METHOD); + if request.method() == Method::OPTIONS && !preflight { *request.method_mut() = stand_in(); request.extensions_mut().insert(OrdinaryOptions); } diff --git a/src/cors/tests.rs b/src/cors/tests.rs index 22698b9..9bff9ce 100644 --- a/src/cors/tests.rs +++ b/src/cors/tests.rs @@ -49,6 +49,16 @@ async fn options_without_origin_reaches_the_route() { assert_eq!(body, "OPTIONS"); } +#[tokio::test] +async fn options_with_a_request_method_but_no_origin_reaches_the_route() { + // A preflight is a CORS request, and a CORS request carries `Origin` + // (Fetch §3.2.2): without it this is an ordinary OPTIONS, which a + // forward-auth check must see rather than the CORS layer's 200. + let (status, _, body) = send(options(&[("access-control-request-method", "PUT")])).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(body, "OPTIONS"); +} + #[tokio::test] async fn preflight_is_answered_by_the_cors_layer() { let (status, headers, body) = send(options(&[ diff --git a/src/openapi.rs b/src/openapi.rs index 745615b..ce95bfa 100644 --- a/src/openapi.rs +++ b/src/openapi.rs @@ -270,6 +270,8 @@ fn build_operation( !path_params.contains(&f.name().to_string()) && body_field != Some(f.name()) }) .map(|field| { + // A message-typed parameter refers to its schema by `$ref`. + register_nested(&field, schemas); json!({ "name": field.name(), "in": "query", diff --git a/src/openapi/tests.rs b/src/openapi/tests.rs index 1ebc56e..8b4b6c5 100644 --- a/src/openapi/tests.rs +++ b/src/openapi/tests.rs @@ -135,6 +135,12 @@ message Item { message Note { string text = 1; } +message Filter { + string q = 1; +} +message Search { + Filter filter = 1; +} message Upload { string name = 1; google.api.HttpBody file = 2; @@ -161,6 +167,9 @@ service Api { rpc Put(Upload) returns (google.api.HttpBody) { option (google.api.http) = { put: "/v1/uploads/{name}" body: "file" }; } + rpc Find(Search) returns (Item) { + option (google.api.http) = { get: "/v1/find" }; + } rpc Raw(google.api.HttpBody) returns (google.api.HttpBody) { option (google.api.http) = { post: "/v1/raw" body: "*" }; } @@ -269,6 +278,21 @@ fn body_rule_decides_body_and_query_fields() { assert_eq!(touch["parameters"].as_array().unwrap().len(), 3); } +#[test] +fn message_typed_query_parameter_schema_is_registered() { + // `filter` is bound only through the query; its `$ref` must resolve to a + // schema in `components`. + let spec = spec(&[]); + let param = &spec["paths"]["/v1/find"]["get"]["parameters"][0]; + assert_eq!(param["in"], "query"); + assert_eq!(param["schema"]["$ref"], "#/components/schemas/Filter"); + assert!( + spec["components"]["schemas"].get("Filter").is_some(), + "{}", + spec["components"]["schemas"] + ); +} + #[test] fn http_body_request_and_response_are_raw_content() { let spec = spec(&[]); diff --git a/src/transcode/mod.rs b/src/transcode/mod.rs index c69267b..2e50be8 100644 --- a/src/transcode/mod.rs +++ b/src/transcode/mod.rs @@ -76,9 +76,15 @@ enum RequestBody { Parsed(request::BodyMapping), /// Raw bytes and `Content-Type` into the input message, a `google.api.HttpBody`. RawRoot, - /// Raw bytes and `Content-Type` into the HttpBody field (of that type) `body` - /// names; the other fields still come from path and query. - RawField(FieldDescriptor, MessageDescriptor), + /// Raw bytes and `Content-Type` into the HttpBody field `body` names; the + /// other fields still come from path and query. + RawField { + /// The rule's `body` mapping, which keeps the field out of query binding. + mapping: request::BodyMapping, + field: FieldDescriptor, + /// The field's type, a `google.api.HttpBody`. + http_body: MessageDescriptor, + }, } /// What the HTTP response body is made of. @@ -375,7 +381,7 @@ fn route_bindings(pool: &DescriptorPool, aliases: &[AliasConfig]) -> Vec fields.remove(segment), + let field = desc.as_ref().and_then(|d| d.get_field_by_name(segment)); + desc = match field.as_ref().map(FieldDescriptor::kind) { + Some(prost_reflect::Kind::Message(inner)) => Some(inner), + _ => None, + }; + value = match (value, field) { + (Some(serde_json::Value::Object(mut fields)), Some(field)) => { + fields.remove(field.json_name()) + } _ => None, }; } @@ -596,11 +612,18 @@ async fn call_unary( .server_streaming(request, entry.grpc_path.clone(), entry.codec()) .await?; let (initial, mut stream, _) = response.into_parts(); - let message = stream - .message() - .await? - .ok_or_else(|| tonic::Status::internal("Missing response message."))?; - let trailers = stream.trailers().await?; + let message = match stream.message().await { + Ok(Some(message)) => message, + Ok(None) => { + let status = tonic::Status::internal("Missing response message."); + return Err(with_initial(status, initial)); + } + Err(status) => return Err(with_initial(status, initial)), + }; + let trailers = match stream.trailers().await { + Ok(trailers) => trailers, + Err(status) => return Err(with_initial(status, initial)), + }; Ok(UnaryAnswer { initial, message, @@ -608,6 +631,26 @@ async fn call_unary( }) } +/// A failure that came after the upstream's response headers, with their +/// metadata put ahead of the failure's own: both belong to the error +/// response, as tonic's `Grpc::unary` keeps them. +fn with_initial(mut status: tonic::Status, initial: MetadataMap) -> tonic::Status { + let own = std::mem::take(status.metadata_mut()).into_headers(); + let mut merged = initial.into_headers(); + let mut last: Option = None; + // A header map yields a name only with the first value of each key. + for (name, value) in own { + if let Some(name) = name { + merged.append(&name, value); + last = Some(name); + } else if let Some(name) = &last { + merged.append(name, value); + } + } + *status.metadata_mut() = MetadataMap::from_headers(merged); + status +} + /// Serve a unary RPC. async fn unary_call( mut client: Grpc, @@ -685,11 +728,10 @@ fn invalid_content_type(entry: &RouteEntry) -> Response { } /// The HTTP response to a failed call, carrying the failure's metadata as -/// headers unless its details were malformed and the answer is the generic -/// `INTERNAL`. Only the failure's own metadata is used (a trailers-only -/// response, or the trailers ending the call): a failure the proxy's client -/// raises itself, such as an undecodable message, has none, so nothing of an -/// answer the proxy rejected reaches the client. +/// headers (a trailers-only response, or the response headers and trailers of +/// a call that failed after them, see [`with_initial`]) unless its details +/// were malformed and the answer is the generic `INTERNAL`, which then carries +/// nothing of the upstream's answer. fn upstream_error(mut status: tonic::Status, entry: &RouteEntry) -> Response { let (response, faithful) = error::render_response(&status, entry.error_details.as_deref()); let metadata = std::mem::take(status.metadata_mut()); @@ -726,13 +768,12 @@ async fn streaming_call( Err(status) => return upstream_error(status, &entry), }; let (initial, stream, _) = response.into_parts(); + if matches!(entry.response, ResponseShape::HttpBody(_)) { + return http_body_stream(stream, entry, initial).await; + } let mut upstream = UpstreamHeaders::default(); upstream.absorb(initial, &entry.denied_headers); let headers = upstream.into_headers(); - - if matches!(entry.response, ResponseShape::HttpBody(_)) { - return http_body_stream(stream, entry, headers).await; - } // Once the upstream accepted the call, the response starts at once rather // than waiting for the first item, so headers and SSE keep-alives are not // held back; an error that comes before the first message is a terminal @@ -752,19 +793,23 @@ async fn streaming_call( /// A server-streaming HttpBody response: `Content-Type` from the first /// message, so the headers wait for it, then every message's `data` as it -/// arrives. An error before the first message is an ordinary error response; -/// after it, the raw body has no in-band error frame, so the body is aborted +/// arrives. An error before the first message is an ordinary error response, +/// with the initial metadata the upstream sent before it; after the first +/// message, the raw body has no in-band error frame, so the body is aborted /// and the client sees a truncated transfer instead of a clean end. async fn http_body_stream( mut stream: tonic::Streaming, entry: Arc, - headers: HeaderMap, + initial: MetadataMap, ) -> Response { let first = match stream.message().await { Ok(Some(message)) => entry.raw_body(message).unwrap_or_default(), Ok(None) => httpbody::RawBody::default(), - Err(status) => return upstream_error(status, &entry), + Err(status) => return upstream_error(with_initial(status, initial), &entry), }; + let mut upstream = UpstreamHeaders::default(); + upstream.absorb(initial, &entry.denied_headers); + let headers = upstream.into_headers(); let content_type = match content_type_header(&first.content_type) { Ok(content_type) => content_type, Err(()) => return invalid_content_type(&entry), @@ -930,15 +975,19 @@ fn decode_request( raw_query: Option<&str>, body_bytes: Bytes, ) -> Result { - let mapping = match &entry.request_body { - RequestBody::Parsed(mapping) => mapping, - RequestBody::RawRoot | RequestBody::RawField(..) => &NO_PARSED_BODY, - }; - // Only parse the body when the rule maps it onto the message. - let json_body = match mapping { - request::BodyMapping::None => serde_json::Value::Null, - _ => body::parse_body(body::content_type(headers), &body_bytes) - .map_err(|e| format!("failed to parse request body: {e}"))?, + // A raw field keeps its `body` mapping with a null placeholder, so query + // binding leaves that field to the body as it does for a parsed one. + let (mapping, json_body) = match &entry.request_body { + RequestBody::Parsed(request::BodyMapping::None) => { + (&NO_PARSED_BODY, serde_json::Value::Null) + } + RequestBody::Parsed(mapping) => ( + mapping, + body::parse_body(body::content_type(headers), &body_bytes) + .map_err(|e| format!("failed to parse request body: {e}"))?, + ), + RequestBody::RawRoot => (&NO_PARSED_BODY, serde_json::Value::Null), + RequestBody::RawField { mapping, .. } => (mapping, serde_json::Value::Null), }; // Query string → field bindings (fields not bound by path or body). @@ -957,7 +1006,9 @@ fn decode_request( RequestBody::RawRoot => { httpbody::fill(&mut message, request_content_type(headers)?, body_bytes); } - RequestBody::RawField(field, http_body) => { + RequestBody::RawField { + field, http_body, .. + } => { let mut inner = DynamicMessage::new(http_body.clone()); httpbody::fill(&mut inner, request_content_type(headers)?, body_bytes); message.set_field(field, prost_reflect::Value::Message(inner)); @@ -1043,15 +1094,22 @@ fn extract_routes(pool: &DescriptorPool) -> Vec { /// How a binding's `body` rule reaches `input`: raw into an HttpBody (the /// input itself with `body: "*"`, or the HttpBody field `body` names), parsed -/// otherwise. +/// otherwise. `body` names a field of the input itself, never a dotted path: +/// `google/api/http.proto` requires "the referred field must not be a repeated +/// field and must be present at the top-level of request message type". fn request_body(input: &MessageDescriptor, mapping: request::BodyMapping) -> RequestBody { - match &mapping { - request::BodyMapping::Root if httpbody::is_http_body(input) => RequestBody::RawRoot, - request::BodyMapping::Field(name) => match httpbody::http_body_field(input, name) { - Some((field, http_body)) => RequestBody::RawField(field, http_body), - None => RequestBody::Parsed(mapping), + let raw_field = match &mapping { + request::BodyMapping::Root if httpbody::is_http_body(input) => return RequestBody::RawRoot, + request::BodyMapping::Field(name) => httpbody::http_body_field(input, name), + _ => None, + }; + match raw_field { + Some((field, http_body)) => RequestBody::RawField { + mapping, + field, + http_body, }, - _ => RequestBody::Parsed(mapping), + None => RequestBody::Parsed(mapping), } } diff --git a/src/transcode/response.rs b/src/transcode/response.rs index b2aca8e..4062967 100644 --- a/src/transcode/response.rs +++ b/src/transcode/response.rs @@ -168,15 +168,18 @@ pub(crate) fn with_upstream_headers(response: Response, upstream: HeaderMap) -> } /// A response with `status`, `headers` and `body`, typed by `content_type`. -/// `204` and `304` carry neither content nor a `Content-Type`, whatever the -/// upstream returned (RFC 9110 §15.3.5, §15.4.5). +/// `204`, `205` and `304` carry neither content nor a `Content-Type`, whatever +/// the upstream returned (RFC 9110 §15.3.5, §15.3.6, §15.4.5). pub(crate) fn build( status: StatusCode, mut headers: HeaderMap, content_type: Option, body: Body, ) -> Response { - let no_content = matches!(status, StatusCode::NO_CONTENT | StatusCode::NOT_MODIFIED); + let no_content = matches!( + status, + StatusCode::NO_CONTENT | StatusCode::RESET_CONTENT | StatusCode::NOT_MODIFIED + ); let body = if no_content { Body::empty() } else { body }; if let (Some(content_type), false) = (content_type, no_content) { headers.insert(CONTENT_TYPE, content_type); diff --git a/src/transcode/response/tests.rs b/src/transcode/response/tests.rs index 3dcb6dd..4598064 100644 --- a/src/transcode/response/tests.rs +++ b/src/transcode/response/tests.rs @@ -247,8 +247,12 @@ async fn build_without_content_type_sets_none() { #[tokio::test] async fn no_content_statuses_drop_body_and_content_type() { - // RFC 9110 §15.3.5 / §15.4.5: 204 and 304 cannot carry content. - for status in [StatusCode::NO_CONTENT, StatusCode::NOT_MODIFIED] { + // RFC 9110 §15.3.5 / §15.3.6 / §15.4.5: 204, 205 and 304 carry no content. + for status in [ + StatusCode::NO_CONTENT, + StatusCode::RESET_CONTENT, + StatusCode::NOT_MODIFIED, + ] { let response = build( status, HeaderMap::new(), diff --git a/src/transcode/tests.rs b/src/transcode/tests.rs index dcb8f02..b4f370c 100644 --- a/src/transcode/tests.rs +++ b/src/transcode/tests.rs @@ -1,6 +1,98 @@ use super::*; use axum::routing::get; +/// Serves a minimal `google/api` from memory plus one test file. +struct OneApi(&'static str); + +impl protox::file::FileResolver for OneApi { + fn open_file(&self, name: &str) -> Result { + let source = match name { + "google/api/annotations.proto" => { + r#"syntax = "proto3"; +package google.api; +import "google/protobuf/descriptor.proto"; +message HttpRule { + oneof pattern { string get = 2; string post = 4; } + string body = 7; + string response_body = 12; +} +extend google.protobuf.MethodOptions { HttpRule http = 72295728; } +"# + } + "api.proto" => self.0, + _ => return protox::file::GoogleFileResolver::new().open_file(name), + }; + protox::file::File::from_source(name, source) + } +} + +/// A descriptor pool compiled from one annotated `.proto` source. +fn api_pool(source: &'static str) -> DescriptorPool { + protox::Compiler::with_file_resolver(OneApi(source)) + .open_file("api.proto") + .unwrap() + .descriptor_pool() +} + +#[test] +fn response_body_names_proto_fields_not_json_keys() { + // `response_body` is a proto field path (`user_info`), while the + // serialized message uses JSON names (`userInfo`); multi-word fields must + // still resolve, at every level. + let pool = api_pool( + r#"syntax = "proto3"; +package t; +message Inner { string display_name = 1; } +message Resp { Inner user_info = 1; } +"#, + ); + let inner_desc = pool.get_message_by_name("t.Inner").unwrap(); + let mut inner = DynamicMessage::new(inner_desc); + inner.set_field_by_name("display_name", prost_reflect::Value::String("Ann".into())); + let mut resp = DynamicMessage::new(pool.get_message_by_name("t.Resp").unwrap()); + resp.set_field_by_name("user_info", prost_reflect::Value::Message(inner)); + + let json = |path| { + serde_json::from_slice::(&json_body(&resp, Some(path)).unwrap()).unwrap() + }; + assert_eq!(json("user_info"), serde_json::json!({"displayName": "Ann"})); + assert_eq!(json("user_info.display_name"), serde_json::json!("Ann")); + // A path that names no field is JSON null, as before. + assert_eq!(json("user_info.missing"), serde_json::Value::Null); + assert_eq!(json("userInfo"), serde_json::Value::Null); +} + +#[test] +fn alias_paths_are_converted_like_the_route_they_alias() { + // An alias keeps the route's field template, so it must go through the + // same template conversion: `{path=**}` is axum's `{*path}`, never a + // literal capture named `path=**` that leaves the field unbound. + let pool = api_pool( + r#"syntax = "proto3"; +package t; +import "google/api/annotations.proto"; +message Req { string path = 1; } +service S { + rpc Get(Req) returns (Req) { option (google.api.http) = { get: "/v1/files/{path=**}" }; } + rpc Watch(Req) returns (stream Req) { option (google.api.http) = { get: "/v1/logs/{path=*}" }; } +} +"#, + ); + let alias: AliasConfig = serde_yaml::from_str("from: /api/{path}\nto: /v1").unwrap(); + let paths = route_paths(&pool, &[alias]); + for expected in [ + "/v1/files/{*path}", + "/api/files/{*path}", + "/v1/logs/{path}", + "/api/logs/{path}", + ] { + assert!( + paths.contains(&("GET".to_owned(), expected.to_owned())), + "{expected} missing from {paths:?}" + ); + } +} + #[test] fn test_proto_path_to_axum() { // axum 0.8: proto `{param}` IS the native capture syntax, pass through verbatim. diff --git a/tests/upstream_controls.rs b/tests/upstream_controls.rs index 274280d..d1e0ab2 100644 --- a/tests/upstream_controls.rs +++ b/tests/upstream_controls.rs @@ -57,6 +57,9 @@ service Controls { rpc Trailers(Req) returns (Reply) { option (google.api.http) = { get: "/v1/trailers" }; } + rpc LateError(Req) returns (Reply) { + option (google.api.http) = { get: "/v1/late-error" }; + } rpc Token(TokenRequest) returns (google.api.HttpBody) { option (google.api.http) = { post: "/v1/token" body: "*" }; } @@ -313,6 +316,8 @@ impl tonic::server::ServerStreamingService for Streaming { Err(tonic::Status::internal("disk failed")), ], ("Download", "empty") => Vec::new(), + // Headers sent, then a failure before the first message. + ("Download", "late-denied") => vec![Err(invalid_token())], ("Download", _) => vec![ Ok(http_body(pool, "text/csv", b"a,b\n")), // Only the first message's content type counts. @@ -355,6 +360,28 @@ fn trailers_response(pool: &DescriptorPool) -> http::Response .unwrap() } +/// A `LateError` answer written by hand: response headers with metadata +/// (`x-initial`), then no message and an `UNAUTHENTICATED` in the trailers, +/// which carry the challenge. tonic's server API would send a trailers-only +/// response instead. +fn late_error_response() -> http::Response { + let mut trailers = HeaderMap::new(); + trailers.insert("grpc-status", "16".parse().unwrap()); + trailers.insert("grpc-message", "expired".parse().unwrap()); + trailers.insert( + "www-authenticate", + "Bearer error=\"invalid_token\"".parse().unwrap(), + ); + let frames: Vec, Infallible>> = + vec![Ok(http_body::Frame::trailers(trailers))]; + let body = http_body_util::StreamBody::new(futures::stream::iter(frames)); + http::Response::builder() + .header("content-type", "application/grpc") + .header("x-initial", "1") + .body(tonic::body::Body::new(body)) + .unwrap() +} + /// The `test.v1.Controls` gRPC service, dispatching by method path. #[derive(Clone)] struct Controls { @@ -384,10 +411,14 @@ impl tower::Service> for Controls { .strip_prefix("/test.v1.Controls/") .unwrap() .to_owned(); - if rpc == "Trailers" { + if rpc == "Trailers" || rpc == "LateError" { // Read the request to its end before answering. let _request = req.into_body().collect().await.unwrap(); - return Ok(trailers_response(&pool)); + return Ok(if rpc == "Trailers" { + trailers_response(&pool) + } else { + late_error_response() + }); } let method = pool .get_service_by_name("test.v1.Controls") @@ -539,6 +570,34 @@ async fn trailers_only_error_carries_its_metadata() { assert_eq!(values(&headers, "content-type"), ["application/json"]); } +#[tokio::test] +async fn unary_error_after_headers_carries_initial_metadata_and_trailers() { + // The upstream sent response headers, then failed in its trailers: the + // error keeps both, initial metadata first, as tonic's own unary call + // does. + let app = proxy().await; + let (status, headers, body) = call(&app, Method::GET, "/v1/late-error").await; + assert_eq!(status, StatusCode::UNAUTHORIZED); + assert_eq!(values(&headers, "x-initial"), ["1"]); + assert_eq!( + values(&headers, "www-authenticate"), + ["Bearer error=\"invalid_token\""] + ); + assert_eq!(json_body(&body)["message"], "expired"); +} + +#[tokio::test] +async fn http_body_stream_failing_before_its_first_message_keeps_initial_metadata() { + // Headers were sent before the failure, so they belong to the error + // response, next to the failure's own metadata. + let app = proxy().await; + let (status, headers, _) = call(&app, Method::GET, "/v1/files/late-denied").await; + assert_eq!(status, StatusCode::UNAUTHORIZED); + assert_eq!(values(&headers, "x-stream"), ["1"]); + assert_eq!(values(&headers, "cache-control"), ["max-age=60"]); + assert!(values(&headers, "www-authenticate")[0].starts_with("Bearer error=")); +} + #[tokio::test] async fn malformed_error_status_carries_no_upstream_metadata() { // The answer is the generic INTERNAL, so nothing of the broken error, @@ -701,6 +760,21 @@ async fn http_body_field_receives_the_raw_body_next_to_path_fields() { ); } +#[tokio::test] +async fn query_key_naming_the_raw_body_field_does_not_break_the_upload() { + // The body binds `file`, and the body wins over the query: a `file` + // query parameter (or one under it) is ignored rather than bound into + // the HttpBody field before the raw body replaces it. + let app = proxy().await; + let request = http::Request::put("/v1/uploads/report?file=x&file.content_type=y") + .header("content-type", "text/csv") + .body(Body::from("a,b")) + .unwrap(); + let (status, _, body) = send(&app, request).await; + assert_eq!(status, StatusCode::OK, "{body:?}"); + assert_eq!(json_body(&body), json!({"name": "report|text/csv|a,b"})); +} + #[tokio::test] async fn http_body_request_with_a_non_ascii_content_type_is_rejected() { // HttpBody.content_type is a proto string; bytes that are not visible From 167ef1de3f25c3899f7ee4e5058777aa01155b9a Mon Sep 17 00:00:00 2001 From: Dmitry Prudnikov Date: Sun, 27 Sep 2026 04:34:56 +0300 Subject: [PATCH 3/3] fix(transcode): method before body, raw root query, OpenAPI schemas - Extension-method fallback picks the binding by method before reading anything, so an unbound method gets its 405 without its body being buffered (it could get 413 instead) - An HttpBody input with body "*" binds nothing from the query: every field comes from the raw body, and a stray query key no longer fails the request - OpenAPI: a self-referencing message is registered once and referenced instead of recursing until the stack overflows (router build aborted) - OpenAPI: HEAD operations describe no response content - OpenAPI: a JSON response_body is described by the selected field's schema, not the whole response message Each fix carries a regression test that failed before it. --- README.md | 6 +++-- src/openapi.rs | 55 +++++++++++++++++++++++++++++++++++++- src/openapi/tests.rs | 54 +++++++++++++++++++++++++++++++++++++ src/transcode/mod.rs | 25 ++++++++++------- tests/upstream_controls.rs | 31 +++++++++++++++++++++ 5 files changed, 158 insertions(+), 13 deletions(-) diff --git a/README.md b/README.md index 4dbad42..29d4e73 100644 --- a/README.md +++ b/README.md @@ -481,8 +481,10 @@ answer with `x-http-code` and that body. Server-streaming calls ignore the key. 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; the other fields still come from the path and -query. A server-streaming `HttpBody` writes each message's `data` as it +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` diff --git a/src/openapi.rs b/src/openapi.rs index ce95bfa..4b8b951 100644 --- a/src/openapi.rs +++ b/src/openapi.rs @@ -76,6 +76,9 @@ pub fn generate(pool: &DescriptorPool, config: &OpenApiConfig, aliases: &[AliasC } for path in targets { let mut operation = operation.clone(); + if http_method == "head" { + strip_response_content(&mut operation); + } operation["operationId"] = json!(unique_id(&mut operation_ids, &id)); add_path_operation(&mut paths, &path, http_method, operation); } @@ -153,6 +156,18 @@ fn operation_methods(method: &RouteMethod) -> Vec<&'static str> { vec![operation] } +/// Drop the content of every response: a HEAD response carries none +/// (RFC 9110 §9.3.2). +fn strip_response_content(operation: &mut Value) { + if let Some(responses) = operation["responses"].as_object_mut() { + for response in responses.values_mut() { + if let Some(response) = response.as_object_mut() { + response.remove("content"); + } + } + } +} + /// `base`, or `base_2`, `base_3`, ... when taken: OpenAPI requires every /// `operationId` to be unique, and one RPC can be mounted several times /// (additional bindings, aliases, a `*` rule). @@ -308,6 +323,22 @@ fn build_operation( }, }, }) + } else if let Some(path) = &binding.response_body { + // The answer is only the field `response_body` names. + let schema = match response_field(&output, path) { + Some(field) => { + register_nested(&field, schemas); + field_to_schema(&field) + } + // The runtime answers JSON `null` for a path that names no field. + None => json!({ "nullable": true }), + }; + json!({ + "200": { + "description": "Success", + "content": { "application/json": { "schema": schema } }, + }, + }) } else if output.full_name() == "google.protobuf.Empty" { json!({ "200": { "description": "Success (empty response)" } }) } else { @@ -350,10 +381,32 @@ fn build_operation( op } -/// Register the schema of a message-typed field so its `$ref` resolves. +/// The field a (possibly dotted) `response_body` path names in `output`, +/// walking singular message fields. +fn response_field(output: &MessageDescriptor, path: &str) -> Option { + let mut desc = output.clone(); + let mut segments = path.split('.').peekable(); + while let Some(segment) = segments.next() { + let field = desc.get_field_by_name(segment)?; + if segments.peek().is_none() { + return Some(field); + } + match field.kind() { + Kind::Message(inner) if !field.is_list() && !field.is_map() => desc = inner, + _ => return None, + } + } + None +} + +/// Register the schema of a message-typed field so its `$ref` resolves. A +/// placeholder goes in before the fields are walked, so a message that +/// contains itself (directly or through others) is referenced rather than +/// expanded without end. fn register_nested(field: &FieldDescriptor, schemas: &mut Map) { if let Kind::Message(nested) = field.kind() { if !is_well_known(&nested) && !schemas.contains_key(nested.name()) { + schemas.insert(nested.name().to_string(), json!({ "type": "object" })); let nested_schema = message_to_schema(&nested, &[], schemas); schemas.insert(nested.name().to_string(), nested_schema); } diff --git a/src/openapi/tests.rs b/src/openapi/tests.rs index 8b4b6c5..57abf22 100644 --- a/src/openapi/tests.rs +++ b/src/openapi/tests.rs @@ -145,6 +145,10 @@ message Upload { string name = 1; google.api.HttpBody file = 2; } +message Tree { + string name = 1; + Tree child = 2; +} service Api { rpc Head(Item) returns (Item) { @@ -162,8 +166,15 @@ service Api { body: "*" additional_bindings { put: "/v1/items/{name}" body: "note" } additional_bindings { post: "/v1/items:touch" } + additional_bindings { get: "/v1/items/{name}/note" response_body: "note" } }; } + rpc Walk(Tree) returns (Tree) { + option (google.api.http) = { get: "/v1/tree" }; + } + rpc Plant(Tree) returns (Tree) { + option (google.api.http) = { post: "/v1/tree" body: "*" }; + } rpc Put(Upload) returns (google.api.HttpBody) { option (google.api.http) = { put: "/v1/uploads/{name}" body: "file" }; } @@ -293,6 +304,49 @@ fn message_typed_query_parameter_schema_is_registered() { ); } +#[test] +fn self_referencing_message_does_not_recurse_forever() { + // `Tree.child` is a `Tree`: in the query (GET) and in the body (POST) its + // schema is registered once and referenced, instead of being expanded + // until the stack overflows. + let spec = spec(&[]); + let tree = &spec["components"]["schemas"]["Tree"]; + assert_eq!( + tree["properties"]["child"]["$ref"], + "#/components/schemas/Tree" + ); + assert_eq!( + spec["paths"]["/v1/tree"]["get"]["parameters"][1]["schema"]["$ref"], + "#/components/schemas/Tree" + ); + assert!(spec["paths"]["/v1/tree"]["post"]["requestBody"].is_object()); +} + +#[test] +fn head_operations_describe_no_response_content() { + // A HEAD response has no content (RFC 9110 §9.3.2), whether the binding + // is a `HEAD` rule or the HEAD operation of a `*` rule. + let spec = spec(&[]); + for path in ["/v1/items/{name}", "/v1/verify"] { + let head = &spec["paths"][path]["head"]["responses"]["200"]; + assert!(head.is_object(), "{path}"); + assert!(head.get("content").is_none(), "{path}: {head}"); + } + // The other operations of the `*` rule keep their content. + assert!(spec["paths"]["/v1/verify"]["get"]["responses"]["200"]["content"].is_object()); +} + +#[test] +fn json_response_body_is_described_by_the_selected_field() { + // `response_body: "note"` answers with the `note` field only. + let spec = spec(&[]); + let content = &spec["paths"]["/v1/items/{name}/note"]["get"]["responses"]["200"]["content"]; + assert_eq!( + content["application/json"]["schema"]["$ref"], + "#/components/schemas/Note" + ); +} + #[test] fn http_body_request_and_response_are_raw_content() { let spec = spec(&[]); diff --git a/src/transcode/mod.rs b/src/transcode/mod.rs index 2e50be8..1b0613a 100644 --- a/src/transcode/mod.rs +++ b/src/transcode/mod.rs @@ -343,16 +343,16 @@ impl PathRoutes { let allow = HeaderValue::from_str(&allow.join(", ")) .expect("method tokens are valid header value characters"); let extension: Arc<[(Method, Arc)]> = extension.into(); + // The method decides before anything is extracted, so a request no + // binding answers gets its 405 without its body being read. router.fallback( - move |method: Method, - state: State, - headers: HeaderMap, - path_params: Path, - raw_query: RawQuery, - body: Bytes| async move { - match extension.iter().find(|(bound, _)| *bound == method) { + move |State(state): State, request: axum::extract::Request| async move { + match extension + .iter() + .find(|(bound, _)| bound == request.method()) + { Some((_, entry)) => { - handle(state, headers, path_params, raw_query, body, entry.clone()).await + axum::handler::Handler::call(endpoint!(entry.clone()), request, state).await } None => (StatusCode::METHOD_NOT_ALLOWED, [(ALLOW, allow)]).into_response(), } @@ -992,8 +992,13 @@ fn decode_request( // Query string → field bindings (fields not bound by path or body). // A malformed query is a client error: reject it rather than silently - // dropping every query-bound field. - let query_pairs = request::parse_query(raw_query)?; + // dropping every query-bound field. An HttpBody input with `body: "*"` + // takes every field from the raw body, so its query binds nothing + // (google/api/http.proto: with `*` there are no HTTP parameters). + let query_pairs = match entry.request_body { + RequestBody::RawRoot => Vec::new(), + _ => request::parse_query(raw_query)?, + }; let input_desc = entry.method.input(); let request_json = diff --git a/tests/upstream_controls.rs b/tests/upstream_controls.rs index d1e0ab2..558c789 100644 --- a/tests/upstream_controls.rs +++ b/tests/upstream_controls.rs @@ -735,6 +735,22 @@ async fn http_body_request_receives_the_raw_body_and_content_type() { assert_eq!(&body[..], raw); } +#[tokio::test] +async fn query_on_a_whole_message_http_body_route_is_ignored() { + // With `body: "*"` on an HttpBody input every field comes from the body + // (google/api/http.proto: no HTTP parameters with `*`), so a query that + // names an HttpBody field is ignored instead of failing the request. + let app = proxy().await; + let request = http::Request::post("/v1/echo?extensions=x&content_type=y&data=z") + .header("content-type", "text/plain") + .body(Body::from("raw")) + .unwrap(); + let (status, headers, body) = send(&app, request).await; + assert_eq!(status, StatusCode::OK, "{body:?}"); + assert_eq!(values(&headers, "content-type"), ["text/plain"]); + assert_eq!(&body[..], b"raw"); +} + #[tokio::test] async fn http_body_request_without_a_body_is_empty() { let app = proxy().await; @@ -933,6 +949,21 @@ async fn extension_method_rule_routes_next_to_a_standard_one() { assert_eq!(status, StatusCode::METHOD_NOT_ALLOWED); } +#[tokio::test] +async fn unbound_method_is_405_before_its_body_is_read() { + // The method decides first: a method nobody binds on the path is 405 + // even with a body over the extractor limit, which is never buffered. + let app = proxy().await; + let request = http::Request::builder() + .method(Method::DELETE) + .uri("/v1/dav") + .body(Body::from(vec![b'x'; 3 * 1024 * 1024])) + .unwrap(); + let (status, headers, _) = send(&app, request).await; + assert_eq!(status, StatusCode::METHOD_NOT_ALLOWED); + assert_eq!(values(&headers, "allow"), ["PROPFIND, GET, HEAD"]); +} + #[tokio::test] async fn star_rule_collides_with_another_method_on_its_path() { // `*` claims every method, so another binding on the same path is a