diff --git a/src/app.rs b/src/app.rs index d7b886c..48f04b0 100644 --- a/src/app.rs +++ b/src/app.rs @@ -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() @@ -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. @@ -1124,6 +1152,19 @@ fn ensure_stream_options(body: &[u8], provider: Provider) -> Option> { 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() diff --git a/src/app/tests.rs b/src/app/tests.rs index 3865d91..ce96736 100644 --- a/src/app/tests.rs +++ b/src/app/tests.rs @@ -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; diff --git a/src/proxy.rs b/src/proxy.rs index fb0572b..dbb682f 100644 --- a/src/proxy.rs +++ b/src/proxy.rs @@ -250,6 +250,7 @@ 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"); @@ -257,6 +258,10 @@ fn filter_hop_by_hop_headers(headers: &HeaderMap) -> HeaderMap { } filtered.insert(name.clone(), value.clone()); } + filtered.insert( + axum::http::header::ACCEPT_ENCODING, + HeaderValue::from_static("identity"), + ); filtered } @@ -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" + ); + } }