Skip to content
Open
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
41 changes: 41 additions & 0 deletions src/app.rs
Original file line number Diff line number Diff line change
Expand Up @@ -566,6 +566,19 @@ async fn handle_request(
"Key resolved successfully"
);

if header_has_non_identity_encoding(&headers, "content-encoding") {
warn!(
method = %method,
path = %path,
virtual_key = %virtual_key,
"Rejected request with unscannable Content-Encoding"
);
return Err(error_response(
StatusCode::BAD_REQUEST,
"Content-Encoding not supported; send an identity-encoded body",
));
}

// 2. Read the body
let body_bytes: Bytes = body
.collect()
Expand Down Expand Up @@ -747,6 +760,21 @@ async fn handle_request(
response
};

if state.dlp_scanner.scan_responses()
&& header_has_non_identity_encoding(response.headers(), "content-encoding")
{
error!(
method = %method,
path = %path,
virtual_key = %virtual_key,
"Upstream response is compressed and cannot be scanned; failing closed"
);
return Err(error_response(
StatusCode::BAD_GATEWAY,
"Upstream response could not be scanned",
));
}

// 5. Response processing.
// Non-streaming responses are buffered so we can (a) record upstream
// token usage for stats and (b) optionally DLP-redact before sending.
Expand Down Expand Up @@ -1124,6 +1152,19 @@ fn ensure_stream_options(body: &[u8], provider: Provider) -> Option<Vec<u8>> {
serde_json::to_vec(&json).ok()
}

fn header_has_non_identity_encoding(headers: &HeaderMap, name: &str) -> bool {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.is_some_and(|value| {
value
.split(',')
.map(str::trim)
.filter(|token| !token.is_empty())
.any(|token| !token.eq_ignore_ascii_case("identity"))
})
}

fn error_response(status: StatusCode, message: &str) -> Response {
let body = serde_json::json!({ "error": message });
(status, axum::Json(body)).into_response()
Expand Down
96 changes: 96 additions & 0 deletions src/app/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2068,6 +2068,102 @@ fn test_ensure_stream_options_skips_anthropic() {
assert!(result.is_none());
}

#[tokio::test]
async fn test_request_rejects_content_encoding() {
let mock_server = MockServer::start().await;
let app = make_app(&mock_server.uri());

let body = r#"{"model":"gpt-4","messages":[]}"#;
let req = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("authorization", "Bearer vk-test-1")
.header("content-type", "application/json")
.header("content-encoding", "gzip")
.body(Body::from(body))
.unwrap();

let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}

#[tokio::test]
async fn test_request_allows_identity_content_encoding() {
let mock_server = MockServer::start().await;

Mock::given(method("POST"))
.and(path("/v1/chat/completions"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({"ok": true})))
.mount(&mock_server)
.await;

let app = make_app(&mock_server.uri());
let body = r#"{"model":"gpt-4","messages":[]}"#;
let req = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("authorization", "Bearer vk-test-1")
.header("content-type", "application/json")
.header("content-encoding", "identity")
.body(Body::from(body))
.unwrap();

let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}

#[tokio::test]
async fn test_forwards_identity_accept_encoding_upstream() {
let mock_server = MockServer::start().await;

Mock::given(method("GET"))
.and(path("/v1/models"))
.and(header("accept-encoding", "identity"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({"ok": true})))
.mount(&mock_server)
.await;

let app = make_app(&mock_server.uri());
let req = Request::builder()
.method("GET")
.uri("/v1/models")
.header("authorization", "Bearer vk-test-1")
.header("accept-encoding", "gzip")
.body(Body::empty())
.unwrap();

let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}

#[tokio::test]
async fn test_response_scan_fails_closed_on_compressed_body() {
let mock_server = MockServer::start().await;

Mock::given(method("POST"))
.and(path("/v1/chat/completions"))
.respond_with(
ResponseTemplate::new(200)
.insert_header("content-encoding", "gzip")
.set_body_bytes(vec![0x1f, 0x8b, 0x08, 0x00]),
)
.mount(&mock_server)
.await;

let app = make_app_with_redact(&mock_server.uri());
let body = r#"{"model":"gpt-4","messages":[{"role":"user","content":"hi"}]}"#;
let req = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("authorization", "Bearer vk-test-1")
.header("content-type", "application/json")
.body(Body::from(body))
.unwrap();

let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::BAD_GATEWAY);
}

#[test]
fn test_ensure_stream_options_preserves_existing() {
use crate::config::Provider;
Expand Down
27 changes: 27 additions & 0 deletions src/proxy.rs
Original file line number Diff line number Diff line change
Expand Up @@ -250,13 +250,18 @@ fn filter_hop_by_hop_headers(headers: &HeaderMap) -> HeaderMap {
|| name_str == "connection"
|| name_str == "content-length"
|| name_str == "transfer-encoding"
|| name_str == "accept-encoding"
|| name_str == "x-api-key"
{
trace!(header = %name_str, "Skipping hop-by-hop/auth header");
continue;
}
filtered.insert(name.clone(), value.clone());
}
filtered.insert(
axum::http::header::ACCEPT_ENCODING,
HeaderValue::from_static("identity"),
);
filtered
}

Expand Down Expand Up @@ -385,4 +390,26 @@ mod tests {
assert!(filtered.get("content-type").is_some());
assert!(filtered.get("x-custom").is_some());
}

#[test]
fn test_filter_forces_identity_accept_encoding() {
let mut headers = HeaderMap::new();
headers.insert("accept-encoding", "gzip, br".parse().unwrap());

let filtered = filter_hop_by_hop_headers(&headers);
assert_eq!(
filtered.get("accept-encoding").unwrap().to_str().unwrap(),
"identity"
);
}

#[test]
fn test_filter_sets_identity_when_absent() {
let headers = HeaderMap::new();
let filtered = filter_hop_by_hop_headers(&headers);
assert_eq!(
filtered.get("accept-encoding").unwrap().to_str().unwrap(),
"identity"
);
}
}