mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
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:
parent
79756cbb9b
commit
3930c5bab6
2 changed files with 63 additions and 2 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue