Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 45 additions & 7 deletions crates/switchyard-server/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1342,7 +1342,12 @@ fn upstream_error(status: StatusCode, body: &str) -> Response {
.as_str()
.filter(|code| !code.is_empty())
.unwrap_or("upstream_error");
error_response(status, message, "upstream_error", code)
let mut response = error_response(status, message, "upstream_error", code);
// Provider messages and codes can quote request content; log only fixed metadata.
response
.extensions_mut()
.insert(RequestLogError(format!("upstream_error (HTTP {status})")));
response
}

// Keep error details until the endpoint chooses its response format.
Expand Down Expand Up @@ -1395,11 +1400,16 @@ impl ApiError {
}
}

fn render_error_response(response: Response, wire_format: WireFormat) -> Response {
let Some(error) = response.extensions().get::<ApiError>().cloned() else {
fn render_error_response(mut response: Response, wire_format: WireFormat) -> Response {
let Some(error) = response.extensions_mut().remove::<ApiError>() else {
return response;
};
error.into_response(wire_format)
let log_error = response.extensions_mut().remove::<RequestLogError>();
let mut rendered = error.into_response(wire_format);
if let Some(log_error) = log_error {
rendered.extensions_mut().insert(log_error);
}
rendered
}

fn anthropic_error_response(response: Response) -> Response {
Expand Down Expand Up @@ -1904,9 +1914,7 @@ mod tests {
let mut message = String::new();
event.record(
&mut |field: &tracing::field::Field, value: &dyn std::fmt::Debug| {
if field.name() == "message" {
message = format!("{value:?}");
}
message.push_str(&format!("{}={value:?} ", field.name()));
},
);
self.0.lock().push((*event.metadata().level(), message));
Expand Down Expand Up @@ -1990,6 +1998,36 @@ mod tests {
);
}

// The provider's message can quote request content, so the request log
// records only the error class while the client still sees the message.
#[test]
fn upstream_error_redacts_request_log_error() {
const LEAKED: &str = "SECRET-quoted-request-content";
let error = LlmClientError::UpstreamHttp {
status: StatusCode::BAD_GATEWAY,
body: format!(
r#"{{"error":{{"message":"validation failed: {LEAKED}","code":"invalid_request_{LEAKED}"}}}}"#
),
};
for wire_format in [
WireFormat::OpenAiChat,
WireFormat::OpenAiResponses,
WireFormat::AnthropicMessages,
] {
let response = render_error_response(client_error(&error), wire_format);
let events = captured_events(|| request_log_context().emit(&response));
assert_eq!(events.len(), 1);
assert!(!events[0].1.contains(LEAKED), "{}", events[0].1);
assert!(events[0].1.contains("upstream_error"), "{}", events[0].1);
let api_error = response
.extensions()
.get::<ApiError>()
.expect("ApiError extension");
assert!(api_error.message.contains(LEAKED), "{}", api_error.message);
assert_eq!(api_error.code, format!("invalid_request_{LEAKED}"));
}
}

// LiteLLM's cost header passes through to the client; auth headers do not.
#[test]
fn upstream_header_forwarding_covers_litellm_cost() {
Expand Down
72 changes: 71 additions & 1 deletion crates/switchyard-server/src/sse.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ use std::sync::Arc;
use axum::response::sse::{Event, Sse};
use futures_util::Stream;
use serde_json::{Value, json};
use switchyard_runner::stream_error_summary;
use switchyard_translation::{LlmStreamError, RawEventStream, WireFormat};

use crate::redaction::Redactor;
Expand Down Expand Up @@ -45,7 +46,14 @@ pub(crate) fn frame_stream(
})
}
Err(LlmStreamError::Client(error)) => {
tracing::warn!(error = %error, "stream iteration failed");
// The error text can quote request content, so the log
// records only the stable error class.
let summary = stream_error_summary(&error, None);
tracing::warn!(
error.kind = summary.kind.as_str(),
error.upstream_status = summary.upstream_status,
"stream iteration failed"
);
failed = true;
error_event(target_format, error.to_string(), &redactor)
}
Expand Down Expand Up @@ -127,6 +135,7 @@ mod tests {
use axum::{body::to_bytes, response::IntoResponse};
use futures_util::stream;
use switchyard_protocol::LlmClientError;
use tracing_subscriber::layer::SubscriberExt;

use super::*;

Expand Down Expand Up @@ -187,6 +196,67 @@ mod tests {
Ok(())
}

// Collects rendered warn events so the test can assert against the final
// log sink rather than a field mid-pipeline.
#[derive(Clone, Default)]
struct WarnCapture(Arc<std::sync::Mutex<Vec<String>>>);

impl tracing_subscriber::Layer<tracing_subscriber::Registry> for WarnCapture {
fn on_event(
&self,
event: &tracing::Event<'_>,
_ctx: tracing_subscriber::layer::Context<'_, tracing_subscriber::Registry>,
) {
if *event.metadata().level() != tracing::Level::WARN {
return;
}
let mut fields = String::new();
event.record(
&mut |field: &tracing::field::Field, value: &dyn std::fmt::Debug| {
fields.push_str(&format!("{}={value:?} ", field.name()));
},
);
self.0.lock().unwrap().push(fields);
}
}

// The client error text can quote request content, so the stream-failure
// warn log records only the stable error class.
#[test]
fn stream_client_error_warn_redacts_upstream_body() -> TestResult {
const LEAKED: &str = "SECRET-quoted-request-content";
let capture = WarnCapture::default();
let subscriber = tracing_subscriber::registry().with(capture.clone());
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()?;

let body = tracing::subscriber::with_default(subscriber, || {
runtime.block_on(chat_body(vec![Err(LlmStreamError::Client(
LlmClientError::UpstreamHttp {
status: axum::http::StatusCode::BAD_GATEWAY,
body: format!("upstream failed: {LEAKED}"),
},
))]))
})?;

// The client still sees the error text in-band.
assert!(body.contains(LEAKED), "{body}");

let events = capture.0.lock().unwrap().clone();
assert!(
events
.iter()
.any(|event| event.contains("stream iteration failed")),
"{events:?}"
);
assert!(
!events.iter().any(|event| event.contains(LEAKED)),
"{events:?}"
);
Ok(())
}

#[tokio::test]
async fn anthropic_stream_error_uses_api_error_type() -> TestResult {
let failure = LlmClientError::General("boom".to_string());
Expand Down
Loading