From 9f547c840bfbbc3d2e7eace0b756227ac13d3428 Mon Sep 17 00:00:00 2001 From: Graham King Date: Tue, 15 Sep 2026 14:39:21 -0400 Subject: [PATCH] feat(server): redact provider API keys from client-facing responses If the upstream provider (e.g. OpenAI, Anthropic, etc) accidentally includes our API key in a response, that would leak it to the user. That's unlikely, but it's also usually not a problem. The person using say Codex also owns the API key. In our case Switchyard owns the API key, so leaking it to the user is a big problem. QA noticed this. Now we redact the keys in transit. - Runner now collects provider_api_keys from DeploymentConfig at load time (and exposes with_provider_api_keys for programmatic hosts). - A new redact_response middleware strips configured secrets from response headers and buffered JSON bodies. It skips gzip encoded bodies. - SSE framing (sse.rs) redacts each event's data before it's written to the stream, so streamed responses don't need to be buffered. Assisted-by: Pi:GPT 6 Astra medium Signed-off-by: Graham King --- crates/switchyard-runner/src/config.rs | 20 +- crates/switchyard-runner/src/runner.rs | 14 ++ crates/switchyard-server/src/lib.rs | 23 ++- crates/switchyard-server/src/redaction.rs | 221 ++++++++++++++++++++++ crates/switchyard-server/src/response.rs | 4 +- crates/switchyard-server/src/sse.rs | 64 ++++--- 6 files changed, 313 insertions(+), 33 deletions(-) create mode 100644 crates/switchyard-server/src/redaction.rs diff --git a/crates/switchyard-runner/src/config.rs b/crates/switchyard-runner/src/config.rs index e2e1c7b14..60a8f3444 100644 --- a/crates/switchyard-runner/src/config.rs +++ b/crates/switchyard-runner/src/config.rs @@ -207,7 +207,8 @@ impl DeploymentConfig { } } - let clients = self.build_clients()?; + let mut provider_api_keys = Vec::new(); + let clients = self.build_clients(&mut provider_api_keys)?; let targets = self.build_targets(); let fallback_base_url = self.fallback_base_url()?; let mut routes = Vec::with_capacity(self.routes.len()); @@ -260,11 +261,16 @@ impl DeploymentConfig { ); routes.push((config.id.clone(), route)); } - let runner = Runner::new(routes).with_fallback_url(fallback_base_url); + let runner = Runner::new(routes) + .with_fallback_url(fallback_base_url) + .with_provider_api_keys(provider_api_keys); Ok(runner) } - fn build_clients(&self) -> RunnerResult>> { + fn build_clients( + &self, + provider_api_keys: &mut Vec, + ) -> RunnerResult>> { let mut models_by_client = self .llm_clients .keys() @@ -273,7 +279,13 @@ impl DeploymentConfig { for (name, client_config) in &self.llm_clients { validate_value("llm client name", name)?; - build_backend(name, client_config, &BTreeMap::new(), None)?; + let backend = build_backend(name, client_config, &BTreeMap::new(), None)?; + let (Backend::OpenAiChat(config) + | Backend::OpenAiResponses(config) + | Backend::Anthropic(config)) = backend; + if let Some(key) = config.api_key { + provider_api_keys.push(key); + } } for (target_name, target) in &self.targets { let client_config = self.llm_clients.get(&target.llm_client).ok_or_else(|| { diff --git a/crates/switchyard-runner/src/runner.rs b/crates/switchyard-runner/src/runner.rs index 36ecb6410..420df2ccc 100644 --- a/crates/switchyard-runner/src/runner.rs +++ b/crates/switchyard-runner/src/runner.rs @@ -17,6 +17,7 @@ use crate::{ModelCapabilities, Route, RunnerError}; pub struct Runner { routes: Vec<(ModelId, Route)>, fallback_base_url: Option, + provider_api_keys: Vec, } /// Borrowed model metadata returned while listing routes. @@ -61,9 +62,22 @@ impl Runner { Self { routes, fallback_base_url: None, + provider_api_keys: Vec::new(), } } + /// Registers deployment-owned API keys for server response redaction. + /// TOML loading registers these automatically; programmatic hosts must supply them. + pub fn with_provider_api_keys(mut self, keys: Vec) -> Self { + self.provider_api_keys = keys; + self + } + + /// Returns deployment-owned secrets for the server's response redactor. + pub fn provider_api_keys(&self) -> &[String] { + &self.provider_api_keys + } + pub(crate) fn with_fallback_url(mut self, fallback_base_url: Option) -> Self { self.fallback_base_url = fallback_base_url; self diff --git a/crates/switchyard-server/src/lib.rs b/crates/switchyard-server/src/lib.rs index d239d66ad..0aea334ad 100644 --- a/crates/switchyard-server/src/lib.rs +++ b/crates/switchyard-server/src/lib.rs @@ -6,6 +6,7 @@ pub mod config; mod metrics; mod observability; +mod redaction; mod response; mod routing_log; mod shutdown; @@ -157,6 +158,7 @@ struct DecisionLlmClientResponse { #[derive(Clone)] pub struct ServerState { runner: Arc, + redactor: Arc, fallback_http: reqwest::Client, metrics: prometheus::Registry, stats: StatsAccumulator, @@ -210,7 +212,9 @@ impl ServerState { metrics.clone(), runner.models().map(|model| model.algorithm), ); + let redactor = redaction::Redactor::new(runner.provider_api_keys()); Ok(Self { + redactor: Arc::new(redactor), runner: Arc::new(runner), fallback_http, metrics, @@ -518,6 +522,10 @@ fn primary_llm_routes() -> Router { fn finish_router(router: Router, state: ServerState) -> Router { router .layer(DefaultBodyLimit::max(DEFAULT_MAX_REQUEST_BODY_BYTES)) + .layer(axum::middleware::from_fn_with_state( + state.clone(), + redaction::redact_response, + )) // `layer` only wraps routes registered before it, so this stays last. .layer(axum::middleware::from_fn(stamp_request_start)) .with_state(state) @@ -1060,11 +1068,16 @@ async fn handle_llm_request( let upstream_headers = std::mem::take(&mut response.upstream_headers); let response_model = served_model.as_ref().map(ToString::to_string); - let mut response = - match into_http_response(response, wire_format, response_model, request_extensions) { - Ok(response) => response, - Err(error) => return server_error(error.to_string()), - }; + let mut response = match into_http_response( + response, + wire_format, + response_model, + request_extensions, + Arc::clone(&state.redactor), + ) { + Ok(response) => response, + Err(error) => return server_error(error.to_string()), + }; // Forward upstream headers before Switchyard writes its own so any header // this server emits always overrides an upstream echo of the same name. let response_headers = response.headers_mut(); diff --git a/crates/switchyard-server/src/redaction.rs b/crates/switchyard-server/src/redaction.rs new file mode 100644 index 000000000..cff0df75b --- /dev/null +++ b/crates/switchyard-server/src/redaction.rs @@ -0,0 +1,221 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Removes configured provider credentials from client-facing responses. + +use axum::body::{Body, to_bytes}; +use axum::extract::{Request, State}; +use axum::http::header::{CONTENT_ENCODING, CONTENT_LENGTH, CONTENT_TYPE}; +use axum::middleware::Next; +use axum::response::Response; + +use crate::{DEFAULT_MAX_REQUEST_BODY_BYTES, ServerState}; + +#[derive(Default)] +pub(crate) struct Redactor { + raw: Vec, + json: Vec, +} + +impl Redactor { + pub(crate) fn new(keys: &[String]) -> Self { + let mut raw: Vec = keys.iter().filter(|key| !key.is_empty()).cloned().collect(); + raw.sort_by(|a, b| b.len().cmp(&a.len()).then_with(|| a.cmp(b))); + raw.dedup(); + let mut json = Vec::with_capacity(raw.len()); + for key in &raw { + let Ok(encoded) = serde_json::to_string(key) else { + unreachable!("serde_json::to_string on a String cannot fail"); + }; + json.push(encoded[1..encoded.len() - 1].to_string()); + } + json.sort_by_key(|key| std::cmp::Reverse(key.len())); + Self { raw, json } + } + + pub(crate) fn json(&self, value: String) -> String { + replace(value, &self.json) + } + + pub(crate) fn text(&self, value: String) -> String { + replace(value, &self.raw) + } +} + +fn replace(mut value: String, secrets: &[String]) -> String { + for secret in secrets { + if value.contains(secret) { + value = value.replace(secret, "[REDACTED]"); + } + } + value +} + +pub(crate) async fn redact_response( + State(state): State, + request: Request, + next: Next, +) -> Response { + let mut response = next.run(request).await; + let redactor = &state.redactor; + if redactor.raw.is_empty() { + return response; + } + let is_json = response.headers().get(CONTENT_TYPE).is_some_and(|value| { + value.to_str().is_ok_and(|value| { + value.split(';').next().is_some_and(|mime| { + mime.trim() == "application/json" || mime.trim().ends_with("+json") + }) + }) + }); + let is_encoded = response + .headers() + .get_all(CONTENT_ENCODING) + .iter() + .any(|value| { + value.to_str().map_or(true, |value| { + value + .split(',') + .any(|encoding| !encoding.trim().eq_ignore_ascii_case("identity")) + }) + }); + for value in response.headers_mut().values_mut() { + if let Ok(text) = value.to_str() { + let sanitized = redactor.text(text.to_string()); + if sanitized != text { + // Replacement contains only visible ASCII and cannot invalidate a header. + if let Ok(header) = sanitized.parse() { + *value = header; + } + } + } + } + // SSE bodies are redacted per event before framing, without buffering the stream. + // The fallback proxy passes compressed bodies through without decoding them. + if is_json && !is_encoded { + let (mut parts, body) = response.into_parts(); + let body = match to_bytes(body, DEFAULT_MAX_REQUEST_BODY_BYTES).await { + Ok(bytes) => match String::from_utf8(bytes.to_vec()) { + Ok(json) => Body::from(redactor.json(json)), + Err(_) => return crate::server_error("Invalid upstream JSON"), + }, + Err(_) => return crate::server_error("Unable to read response body"), + }; + parts.headers.remove(CONTENT_LENGTH); + response = Response::from_parts(parts, body); + } + response +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use axum::Json; + use axum::response::IntoResponse; + use axum::routing::get; + use serde_json::json; + use switchyard_runner::Runner; + use switchyard_translation::{LlmStreamError, WireFormat}; + use tower::ServiceExt; + + use super::*; + + #[tokio::test] + async fn responses_redact_provider_keys_without_changing_errors() + -> Result<(), Box> { + let key = "synthetic-\"provider\\key"; + let error = json!({"error": { + "message": "upstream failed", + "type": "provider_error", + "debug": format!("Bearer {key}"), + key: [key, "ordinary diagnostics"] + }}); + let runner = Runner::new(Vec::new()).with_provider_api_keys(vec![key.to_string()]); + let state = ServerState::from_runner(runner)?; + let buffered = error.clone(); + let streamed = error.clone(); + let header = axum::http::HeaderValue::from_str(key)?; + // gzip-compressed {"ok":true}. + const GZIP_JSON: &[u8] = &[ + 31, 139, 8, 0, 0, 0, 0, 0, 2, 3, 171, 86, 202, 207, 86, 178, 42, 41, 42, 77, 173, 5, 0, + 144, 95, 212, 167, 11, 0, 0, 0, + ]; + let router = axum::Router::new() + .route( + "/compressed", + get(|| async { + ( + [ + (CONTENT_TYPE, "application/json"), + (CONTENT_ENCODING, "gzip"), + (CONTENT_LENGTH, "31"), + ], + GZIP_JSON, + ) + }), + ) + .route("/buffered", get(move || async move { Json(buffered) })) + .route( + "/stream", + get(move |State(state): State| async move { + let events = futures_util::stream::iter([ + Ok(json!({"choices": [], "model": "ordinary-model"})), + Err(LlmStreamError::Upstream(streamed)), + ]); + let mut response = crate::sse::frame_stream( + Box::pin(events), + WireFormat::OpenAiChat, + Arc::clone(&state.redactor), + ) + .into_response(); + response.headers_mut().insert("x-upstream-debug", header); + response + }), + ); + let app = crate::finish_router(router, state); + let compressed = app + .clone() + .oneshot(Request::builder().uri("/compressed").body(Body::empty())?) + .await?; + assert_eq!(compressed.status(), axum::http::StatusCode::OK); + assert_eq!(compressed.headers()[CONTENT_ENCODING], "gzip"); + assert_eq!(compressed.headers()[CONTENT_LENGTH], "31"); + assert_eq!( + to_bytes(compressed.into_body(), usize::MAX).await?.as_ref(), + GZIP_JSON + ); + for path in ["/buffered", "/stream"] { + let response = app + .clone() + .oneshot(Request::builder().uri(path).body(Body::empty())?) + .await?; + if path == "/stream" { + assert_eq!(response.headers()["x-upstream-debug"], "[REDACTED]"); + } + let body = + String::from_utf8(to_bytes(response.into_body(), usize::MAX).await?.to_vec())?; + let error: serde_json::Value = if path == "/stream" { + assert!(body.contains("ordinary-model")); + assert!(!body.contains("[DONE]")); + let data = body + .lines() + .filter_map(|line| line.strip_prefix("data: ")) + .next_back() + .ok_or("missing SSE error")?; + serde_json::from_str(data)? + } else { + serde_json::from_str(&body)? + }; + assert_eq!(error["error"]["message"], "upstream failed"); + assert_eq!(error["error"]["type"], "provider_error"); + assert_eq!(error["error"]["debug"], "Bearer [REDACTED]"); + assert_eq!( + error["error"]["[REDACTED]"], + json!(["[REDACTED]", "ordinary diagnostics"]) + ); + assert!(!body.contains("synthetic-")); + } + Ok(()) + } +} diff --git a/crates/switchyard-server/src/response.rs b/crates/switchyard-server/src/response.rs index 950ec6399..84dac5b41 100644 --- a/crates/switchyard-server/src/response.rs +++ b/crates/switchyard-server/src/response.rs @@ -4,6 +4,7 @@ //! Response encoding glue for libsy server endpoints. use std::error::Error; +use std::sync::Arc; use axum::Json; use axum::response::{IntoResponse, Response as HttpResponse}; @@ -24,6 +25,7 @@ pub(crate) fn into_http_response( target_format: WireFormat, served_model: Option, request_extensions: ProviderExtensions, + redactor: Arc, ) -> Result { match response.llm_response { LlmResponse::Agg(response) => { @@ -42,7 +44,7 @@ pub(crate) fn into_http_response( served_model, &request_extensions, )?; - Ok(frame_stream(events, target_format).into_response()) + Ok(frame_stream(events, target_format, redactor).into_response()) } } } diff --git a/crates/switchyard-server/src/sse.rs b/crates/switchyard-server/src/sse.rs index d94cb220e..e685ba5e5 100644 --- a/crates/switchyard-server/src/sse.rs +++ b/crates/switchyard-server/src/sse.rs @@ -4,12 +4,15 @@ //! SSE framing helpers for OpenAI, Anthropic, and Responses endpoints. use std::convert::Infallible; +use std::sync::Arc; use axum::response::sse::{Event, Sse}; use futures_util::Stream; use serde_json::{Value, json}; use switchyard_translation::{LlmStreamError, RawEventStream, WireFormat}; +use crate::redaction::Redactor; + /// Boxed stream type accepted by Axum's SSE response wrapper. pub(crate) type SseFrameStream = std::pin::Pin> + Send>>; @@ -18,32 +21,33 @@ pub(crate) type SseFrameStream = pub(crate) fn frame_stream( stream: RawEventStream, target_format: WireFormat, + redactor: Arc, ) -> Sse { let framed = async_stream::stream! { let mut stream = stream; let mut failed = false; while let Some(item) = futures_util::StreamExt::next(&mut stream).await { let event = match item { - Ok(value) => match frame_event(target_format, value) { + Ok(value) => match frame_event(target_format, value, &redactor) { Ok(event) => event, Err(error) => { failed = true; - error_event(target_format, error.to_string()) + error_event(target_format, error.to_string(), &redactor) } }, - // The upstream's own error event, already in the target format: forward it - // verbatim so its code and type survive, rather than synthesizing one. + // Preserve the upstream error's fields, apart from credential redaction, + // rather than replacing it with a synthesized error. Err(LlmStreamError::Upstream(value)) => { failed = true; - frame_event(target_format, value.clone()).unwrap_or_else(|error| { + frame_event(target_format, value.clone(), &redactor).unwrap_or_else(|error| { tracing::warn!(error = %error, "in-band error event could not be framed"); - error_event(target_format, value.to_string()) + error_event(target_format, value.to_string(), &redactor) }) } Err(LlmStreamError::Client(error)) => { tracing::warn!(error = %error, "stream iteration failed"); failed = true; - error_event(target_format, error.to_string()) + error_event(target_format, error.to_string(), &redactor) } }; yield Ok(event); @@ -62,41 +66,50 @@ pub(crate) fn frame_stream( Sse::new(Box::pin(framed) as SseFrameStream) } -fn frame_event(target_format: WireFormat, value: Value) -> Result { +fn frame_event( + target_format: WireFormat, + value: Value, + redactor: &Redactor, +) -> Result { + let data = redactor.json(serde_json::to_string(&value)?); match target_format { - WireFormat::OpenAiChat => Event::default().json_data(value), + WireFormat::OpenAiChat => Ok(Event::default().data(data)), WireFormat::AnthropicMessages | WireFormat::OpenAiResponses => { let event_type = value .get("type") .and_then(Value::as_str) .unwrap_or("message") .to_string(); - Event::default().event(event_type).json_data(value) + Ok(Event::default().event(event_type).data(data)) } } } -fn error_event(target_format: WireFormat, message: String) -> Event { +fn error_event(target_format: WireFormat, message: String, redactor: &Redactor) -> Event { match target_format { WireFormat::OpenAiChat => Event::default().data( - json!({ - "error": { - "message": message, - "type": "SwitchyardError", - } - }) - .to_string(), - ), - WireFormat::AnthropicMessages | WireFormat::OpenAiResponses => { - Event::default().event("error").data( + redactor.json( json!({ - "type": "error", "error": { "message": message, "type": "SwitchyardError", } }) .to_string(), + ), + ), + WireFormat::AnthropicMessages | WireFormat::OpenAiResponses => { + Event::default().event("error").data( + redactor.json( + json!({ + "type": "error", + "error": { + "message": message, + "type": "SwitchyardError", + } + }) + .to_string(), + ), ) } } @@ -117,7 +130,12 @@ mod tests { // Renders a framed body for one Chat stream. async fn chat_body(items: Vec>) -> TestResult { let stream: RawEventStream = Box::pin(stream::iter(items)); - let response = frame_stream(stream, WireFormat::OpenAiChat).into_response(); + let response = frame_stream( + stream, + WireFormat::OpenAiChat, + Arc::new(Redactor::default()), + ) + .into_response(); Ok(String::from_utf8( to_bytes(response.into_body(), usize::MAX).await?.to_vec(), )?)