Merge pull request #36151 from LHMQ878/fix/36088-openai-ws-passthrough

fix(proxy): register WebSocket passthrough for OpenAI prefixes
This commit is contained in:
Mateo Wang 2026-08-17 10:00:49 -07:00 • committed by GitHub
commit cfb2eba7f9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 443 additions and 27 deletions

View file

@ -9,7 +9,7 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
import json
import os
import re
from typing import Any, Final, cast
from typing import Annotated, Any, Final, cast
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
@ -27,7 +27,7 @@ from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth, user_api_key_auth_websocket
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
@ -1972,6 +1972,104 @@ async def openai_proxy_route(
)
def _join_url_paths(base_url: httpx.URL, path: str, custom_llm_provider: litellm.LlmProviders) -> str:
"""
Properly joins a base URL with a path, preserving any existing path in the base URL.
"""
# Combine paths via the shared helper so any '..' in the path cannot
# climb above the configured base path.
joined_path_str = str(
base_url.copy_with(path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, path))
)
# Apply OpenAI-specific path handling for both branches
if custom_llm_provider == litellm.LlmProviders.OPENAI and "/v1/" not in joined_path_str:
# Insert v1 after api.openai.com for OpenAI requests
joined_path_str = joined_path_str.replace("api.openai.com/", "api.openai.com/v1/")
return joined_path_str
_OPENAI_WS_ALL_MODEL_ACCESS: Final = frozenset(
{
SpecialModelNames.all_proxy_models.value,
SpecialModelNames.all_team_models.value,
"*",
}
)
def _key_has_model_restrictions(user_api_key_dict: UserAPIKeyAuth) -> bool:
scoped_models: Final = (*user_api_key_dict.models, *user_api_key_dict.team_models)
return any(str(model) not in _OPENAI_WS_ALL_MODEL_ACCESS for model in scoped_models)
@router.websocket("/openai_passthrough/{endpoint:path}")
@router.websocket("/openai/{endpoint:path}")
async def openai_websocket_proxy_route(
websocket: WebSocket,
endpoint: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)],
) -> None:
"""WebSocket passthrough for OpenAI prefixes (realtime / responses.connect)."""
if _key_has_model_restrictions(user_api_key_dict):
await websocket.close(
code=1008,
reason="Keys with model restrictions cannot use OpenAI websocket passthrough",
)
return
base_target_url: Final = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/"
openai_api_key: Final = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.OPENAI.value,
region_name=None,
)
if openai_api_key is None:
await websocket.close(
code=1011,
reason="Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.",
)
return
raw_path: Final = httpx.URL(endpoint).path
encoded_endpoint: Final = raw_path if raw_path.startswith("/") else f"/{raw_path}"
base_url: Final = httpx.URL(base_target_url)
updated_url: Final = _join_url_paths(
base_url=base_url,
path=encoded_endpoint,
custom_llm_provider=litellm.LlmProviders.OPENAI,
)
wss_base: Final = (
"wss://" + updated_url[len("https://") :]
if updated_url.startswith("https://")
else "ws://" + updated_url[len("http://") :]
if updated_url.startswith("http://")
else updated_url
)
query_string: Final = websocket.url.query
wss_target: Final = f"{wss_base}{'&' if '?' in wss_base else '?'}{query_string}" if query_string else wss_base
custom_headers: Final = { # mutable-ok: websocket_passthrough_request requires a plain dict of upstream headers
"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,
custom_headers=custom_headers,
user_api_key_dict=user_api_key_dict,
forward_headers=False,
endpoint=websocket.url.path,
accept_websocket=False,
)
class BaseOpenAIPassThroughHandler:
@staticmethod
async def _base_openai_pass_through_handler(
@ -1991,7 +2089,7 @@ class BaseOpenAIPassThroughHandler:
# Construct the full target URL by properly joining the base URL and endpoint path
base_url: Final = httpx.URL(base_target_url)
updated_url: Final = BaseOpenAIPassThroughHandler._join_url_paths(
updated_url: Final = _join_url_paths(
base_url=base_url,
path=encoded_endpoint,
custom_llm_provider=custom_llm_provider,
@ -2050,24 +2148,6 @@ class BaseOpenAIPassThroughHandler:
request=request,
)
@staticmethod
def _join_url_paths(base_url: httpx.URL, path: str, custom_llm_provider: litellm.LlmProviders) -> str:
"""
Properly joins a base URL with a path, preserving any existing path in the base URL.
"""
# Combine paths via the shared helper so any '..' in the path cannot
# climb above the configured base path.
joined_path_str = str(
base_url.copy_with(path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, path))
)
# Apply OpenAI-specific path handling for both branches
if custom_llm_provider == litellm.LlmProviders.OPENAI and "/v1/" not in joined_path_str:
# Insert v1 after api.openai.com for OpenAI requests
joined_path_str = joined_path_str.replace("api.openai.com/", "api.openai.com/v1/")
return joined_path_str
@router.api_route(
"/cursor/{endpoint:path}",

View file

@ -2091,8 +2091,8 @@ async def websocket_passthrough_request(
raw_response = await upstream_ws.recv(decode=False)
# Ensure raw_response is bytes before decoding
if isinstance(raw_response, str):
raw_response = raw_response.encode("ascii")
setup_response: Final[Mapping[str, object]] = json.loads(raw_response.decode("ascii"))
raw_response = raw_response.encode("utf-8")
setup_response: Final[Mapping[str, object]] = json.loads(raw_response.decode("utf-8"))
verbose_proxy_logger.debug("Setup response: %s", setup_response)
# Extract model and provider from setup response for Vertex AI Live

View file

@ -19,6 +19,7 @@ import litellm
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
BaseOpenAIPassThroughHandler,
RouteChecks,
_join_url_paths,
azure_proxy_route,
bedrock_llm_proxy_route,
create_pass_through_route,
@ -74,7 +75,7 @@ class TestBaseOpenAIPassThroughHandler:
# Test joining base URL with no path and a path
base_url = httpx.URL("https://api.example.com")
path = "/v1/chat/completions"
result = BaseOpenAIPassThroughHandler._join_url_paths(
result = _join_url_paths(
base_url, path, litellm.LlmProviders.OPENAI.value
)
print(f"Base URL with no path: '{base_url}' + '{path}' → '{result}'")
@ -83,7 +84,7 @@ class TestBaseOpenAIPassThroughHandler:
# Test joining base URL with path and another path
base_url = httpx.URL("https://api.example.com/v1")
path = "/chat/completions"
result = BaseOpenAIPassThroughHandler._join_url_paths(
result = _join_url_paths(
base_url, path, litellm.LlmProviders.OPENAI.value
)
print(f"Base URL with path: '{base_url}' + '{path}' → '{result}'")
@ -92,7 +93,7 @@ class TestBaseOpenAIPassThroughHandler:
# Test with path not starting with slash
base_url = httpx.URL("https://api.example.com/v1")
path = "chat/completions"
result = BaseOpenAIPassThroughHandler._join_url_paths(
result = _join_url_paths(
base_url, path, litellm.LlmProviders.OPENAI.value
)
print(f"Path without leading slash: '{base_url}' + '{path}' → '{result}'")
@ -101,7 +102,7 @@ class TestBaseOpenAIPassThroughHandler:
# Test with base URL having trailing slash
base_url = httpx.URL("https://api.example.com/v1/")
path = "/chat/completions"
result = BaseOpenAIPassThroughHandler._join_url_paths(
result = _join_url_paths(
base_url, path, litellm.LlmProviders.OPENAI.value
)
print(f"Base URL with trailing slash: '{base_url}' + '{path}' → '{result}'")

View file

@ -30,6 +30,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
pass_through_request,
resolve_pass_through_request_timeout,
resolve_llm_passthrough_timeout,
websocket_passthrough_request,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
@ -4879,6 +4880,83 @@ async def test_unusable_upstream_cost_records_zero_not_the_flat_estimate():
assert payloads[0]["total_tokens"] == 1874
class FakeUpstreamWebSocket:
def __init__(self, first_frame: bytes):
self._first_frame = first_frame
self.close = AsyncMock()
async def recv(self, decode: bool = True):
return self._first_frame
def __aiter__(self):
return self
async def __anext__(self):
raise StopAsyncIteration
class FakeUpstreamConnect:
def __init__(self, upstream_ws: FakeUpstreamWebSocket):
self._upstream_ws = upstream_ws
async def __aenter__(self):
return self._upstream_ws
async def __aexit__(self, exc_type, exc, tb):
return False
@pytest.mark.asyncio
async def test_websocket_passthrough_forwards_non_ascii_first_frame():
from starlette.websockets import WebSocketState
first_frame = json.dumps(
{"type": "session.created", "session": {"instructions": "Hablas español, ¿sí?"}},
ensure_ascii=False,
).encode("utf-8")
upstream_ws = FakeUpstreamWebSocket(first_frame)
websocket = MagicMock()
websocket.accept = AsyncMock()
websocket.send_text = AsyncMock()
websocket.send_bytes = AsyncMock()
websocket.close = AsyncMock()
websocket.receive = AsyncMock(return_value={"type": "websocket.disconnect"})
websocket.headers = {}
websocket.client_state = WebSocketState.CONNECTED
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect",
return_value=FakeUpstreamConnect(upstream_ws),
),
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER"
) as mock_worker,
):
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.ensure_initialized_and_enqueue = MagicMock(
side_effect=lambda async_coroutine: async_coroutine.close()
)
await websocket_passthrough_request(
websocket=websocket,
target="wss://api.openai.com/v1/realtime?model=gpt-realtime",
custom_headers={"Authorization": "Bearer sk-test"},
user_api_key_dict=UserAPIKeyAuth(),
forward_headers=False,
endpoint="/openai/v1/realtime",
accept_websocket=True,
)
websocket.send_text.assert_awaited_once()
forwarded = json.loads(websocket.send_text.await_args.args[0])
assert forwarded["session"]["instructions"] == "Hablas español, ¿sí?"
assert all(call.kwargs.get("code") != 1011 for call in websocket.close.await_args_list)
def _passthrough_kwargs_for_reservation(
user_api_key_dict: UserAPIKeyAuth, parsed_body: Optional[dict] = None
) -> dict:

View file

@ -0,0 +1,181 @@
"""OpenAI passthrough must register WebSocket catch-all routes (#36088)."""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from starlette.routing import WebSocketRoute
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
openai_websocket_proxy_route,
router,
)
def test_openai_websocket_passthrough_routes_registered():
ws_paths = {route.path for route in router.routes if isinstance(route, WebSocketRoute)}
assert "/openai/{endpoint:path}" in ws_paths
assert "/openai_passthrough/{endpoint:path}" in ws_paths
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
@pytest.mark.asyncio
@pytest.mark.parametrize("prefix", ["openai", "openai_passthrough"])
async def test_openai_websocket_forwards_query_and_keeps_provider_auth(prefix):
websocket = _mock_websocket(f"/{prefix}/v1/realtime", "model=gpt-4o-realtime-preview")
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._join_url_paths",
return_value="https://api.openai.com/v1/realtime",
),
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(),
)
kwargs = mock_ws.await_args.kwargs
assert kwargs["target"] == "wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview"
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()
@pytest.mark.asyncio
async def test_openai_websocket_closes_cleanly_when_provider_credentials_missing():
websocket = _mock_websocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value=None,
),
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.close.assert_awaited_once()
assert websocket.close.await_args.kwargs["code"] == 1011
websocket.accept.assert_not_awaited()
mock_ws.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"user_api_key_dict",
[
UserAPIKeyAuth(models=["gpt-4o"]),
UserAPIKeyAuth(team_models=["gpt-4o-realtime-preview"]),
UserAPIKeyAuth(models=["all-team-models"], team_models=["gpt-4o"]),
],
)
async def test_openai_websocket_rejects_model_restricted_keys(user_api_key_dict):
websocket = _mock_websocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")
with 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=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()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"user_api_key_dict",
[
UserAPIKeyAuth(),
UserAPIKeyAuth(models=["all-proxy-models"]),
UserAPIKeyAuth(models=["*"]),
UserAPIKeyAuth(models=["all-team-models"], team_models=["all-proxy-models"]),
],
)
async def test_openai_websocket_allows_unrestricted_keys(user_api_key_dict):
websocket = _mock_websocket("/openai/v1/responses", "")
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/responses",
user_api_key_dict=user_api_key_dict,
)
mock_ws.assert_awaited_once()
websocket.close.assert_not_awaited()

View file

@ -8665,6 +8665,26 @@ export interface paths {
patch?: never;
trace?: never;
};
"/openai/": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/**
* WebSocket: openai_websocket_proxy_route
* @description WebSocket connection endpoint
*/
get: operations["websocket_openai_websocket_proxy_route_get"];
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/openai/deployments/{model}/chat/completions": {
parameters: {
query?: never;
@ -9165,6 +9185,26 @@ export interface paths {
patch: operations["openai_proxy_route_openai__endpoint__patch"];
trace?: never;
};
"/openai_passthrough/": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/**
* WebSocket: openai_websocket_proxy_route
* @description WebSocket connection endpoint
*/
get: operations["websocket_openai_websocket_proxy_route_get_2"];
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/openai_passthrough/{endpoint}": {
parameters: {
query?: never;
@ -47294,6 +47334,24 @@ export interface operations {
};
};
};
websocket_openai_websocket_proxy_route_get: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description WebSocket Protocol Switched */
101: {
headers: {
[name: string]: unknown;
};
content?: never;
};
};
};
chat_completion_openai_deployments__model__chat_completions_post: {
parameters: {
query?: never;
@ -48153,6 +48211,24 @@ export interface operations {
};
};
};
websocket_openai_websocket_proxy_route_get_2: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description WebSocket Protocol Switched */
101: {
headers: {
[name: string]: unknown;
};
content?: never;
};
};
};
openai_proxy_route_openai_passthrough__endpoint__get: {
parameters: {
query?: never;