fix(proxy): negotiate client subprotocol on OpenAI websocket passthrough

This commit is contained in:
mateo-berri 2026-08-16 14:11:12 -07:00
parent a258b2b130
commit 862f33bbaa
2 changed files with 45 additions and 2 deletions

View file

@ -1993,6 +1993,13 @@ async def openai_websocket_proxy_route(
"Authorization": f"Bearer {openai_api_key}"
}
requested_subprotocols: Final = tuple(
protocol.strip()
for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",")
if protocol.strip()
)
await websocket.accept(subprotocol=requested_subprotocols[0] if requested_subprotocols else None)
await websocket_passthrough_request(
websocket=websocket,
target=wss_target,
@ -2000,7 +2007,7 @@ async def openai_websocket_proxy_route(
user_api_key_dict=user_api_key_dict,
forward_headers=False,
endpoint=websocket.url.path,
accept_websocket=True,
accept_websocket=False,
)

View file

@ -18,10 +18,12 @@ def test_openai_websocket_passthrough_routes_registered():
assert "/openai_passthrough/{endpoint:path}" in ws_paths
def _mock_websocket(path: str, query: str) -> MagicMock:
def _mock_websocket(path: str, query: str, headers: dict[str, str] | None = None) -> MagicMock:
websocket = MagicMock()
websocket.url.path = path
websocket.url.query = query
websocket.headers = headers or {}
websocket.accept = AsyncMock()
websocket.close = AsyncMock()
return websocket
@ -56,6 +58,39 @@ async def test_openai_websocket_forwards_query_and_keeps_provider_auth(prefix):
assert kwargs["custom_headers"] == {"Authorization": "Bearer sk-provider"}
assert kwargs["forward_headers"] is False
assert kwargs["endpoint"] == f"/{prefix}/v1/realtime"
assert kwargs["accept_websocket"] is False
websocket.accept.assert_awaited_once_with(subprotocol=None)
websocket.close.assert_not_awaited()
@pytest.mark.asyncio
async def test_openai_websocket_accepts_first_client_subprotocol():
websocket = _mock_websocket(
"/openai/v1/realtime",
"model=gpt-4o-realtime-preview",
headers={
"sec-websocket-protocol": "realtime, openai-insecure-api-key.sk-abc, openai-beta.realtime-v1"
},
)
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="sk-provider",
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.websocket_passthrough_request",
new_callable=AsyncMock,
) as mock_ws,
):
await openai_websocket_proxy_route(
websocket=websocket,
endpoint="v1/realtime",
user_api_key_dict=UserAPIKeyAuth(),
)
websocket.accept.assert_awaited_once_with(subprotocol="realtime")
assert mock_ws.await_args.kwargs["accept_websocket"] is False
websocket.close.assert_not_awaited()
@ -83,6 +118,7 @@ async def test_openai_websocket_rejects_model_restricted_keys(user_api_key_dict)
websocket.close.assert_awaited_once()
assert websocket.close.await_args.kwargs["code"] == 1008
websocket.accept.assert_not_awaited()
mock_ws.assert_not_awaited()