fix(proxy): strip caller credentials from websocket passthrough (#43855)

* fix(proxy): strip caller credentials from websocket passthrough

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): cover configured x-api-key in websocket passthrough credential test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: oliver <oliver@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-30 11:22:25 -07:00 • committed by GitHub
parent 79756cbb9b
commit 3930c5bab6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 63 additions and 2 deletions

View file

@ -2290,7 +2290,7 @@ def _upstream_close_to_relay(task_results: Iterable[object]) -> Close | None:
return upstream_close
_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("authorization", "x-api-key", "x-goog-user-project"))
_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("x-goog-user-project",))
def _with_trace_context(headers: Mapping[str, str], parent_span: object) -> dict[str, str]:

View file

@ -6068,7 +6068,68 @@ async def test_websocket_passthrough_propagates_active_trace_context(
propagated = get_current_span(TraceContextTextMapPropagator().extract(captured["headers"]))
assert propagated.get_span_context().trace_id == span.get_span_context().trace_id
assert propagated.get_span_context().span_id == span.get_span_context().span_id
assert captured["headers"].get("authorization") == ("Bearer client" if forward_headers else None)
assert "authorization" not in captured["headers"]
@pytest.mark.asyncio
async def test_websocket_passthrough_never_forwards_caller_credentials_upstream(monkeypatch):
from starlette.websockets import WebSocketState
captured: dict[str, dict[str, str]] = {}
upstream_ws = FakeUpstreamWebSocket("{}")
def fake_connect(target, additional_headers):
captured["headers"] = additional_headers
return FakeUpstreamConnect(upstream_ws)
websocket = MagicMock()
websocket.accept = AsyncMock()
websocket.send_text = AsyncMock()
websocket.send_bytes = AsyncMock()
websocket.receive = AsyncMock(return_value={"type": "websocket.disconnect"})
websocket.close = AsyncMock()
websocket.headers = {
"authorization": "Bearer sk-caller-virtual-key",
"api-key": "sk-caller-virtual-key",
"x-api-key": "sk-caller-virtual-key",
"x-goog-api-key": "sk-caller-virtual-key",
"x-goog-user-project": "caller-project",
}
websocket.client_state = WebSocketState.CONNECTED
websocket.application_state = WebSocketState.CONNECTED
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_success_hook = AsyncMock()
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_worker = MagicMock()
mock_worker.ensure_initialized_and_enqueue = MagicMock(side_effect=lambda async_coroutine: async_coroutine.close())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
monkeypatch.setattr(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect",
fake_connect,
)
monkeypatch.setattr(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER",
mock_worker,
)
await websocket_passthrough_request(
websocket=websocket,
target="wss://upstream.example.test/v1/realtime",
custom_headers={
"Authorization": "Bearer upstream-admin-secret",
"x-api-key": "upstream-admin-key",
},
user_api_key_dict=UserAPIKeyAuth(),
forward_headers=True,
endpoint="/realtime",
accept_websocket=True,
)
assert all("sk-caller-virtual-key" not in value for value in captured["headers"].values())
assert captured["headers"]["Authorization"] == "Bearer upstream-admin-secret"
assert captured["headers"]["x-api-key"] == "upstream-admin-key"
assert captured["headers"]["x-goog-user-project"] == "caller-project"
class ClosingUpstreamWebSocket: