mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
fix(proxy): harden OpenAI websocket passthrough
- decode upstream first frame as utf-8 instead of ascii - reject model-restricted keys at connect to match HTTP model enforcement - log the actual request path for /openai_passthrough traffic
This commit is contained in:
parent
02d4f8d6a8
commit
a258b2b130
4 changed files with 179 additions and 15 deletions
|
|
@ -1934,6 +1934,20 @@ async def openai_proxy_route(
|
|||
)
|
||||
|
||||
|
||||
_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(
|
||||
|
|
@ -1942,6 +1956,13 @@ async def openai_websocket_proxy_route(
|
|||
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,
|
||||
|
|
@ -1978,7 +1999,7 @@ async def openai_websocket_proxy_route(
|
|||
custom_headers=custom_headers,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
forward_headers=False,
|
||||
endpoint=f"/openai/{endpoint}",
|
||||
endpoint=websocket.url.path,
|
||||
accept_websocket=True,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -2085,8 +2085,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 = json.loads(raw_response.decode("ascii"))
|
||||
raw_response = raw_response.encode("utf-8")
|
||||
setup_response: Final = 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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -4877,3 +4878,80 @@ async def test_unusable_upstream_cost_records_zero_not_the_flat_estimate():
|
|||
assert len(payloads) == 1
|
||||
assert payloads[0]["response_cost"] == 0.0
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
"""OpenAI passthrough must register WebSocket catch-all routes (#36088)."""
|
||||
"""OpenAI passthrough must register WebSocket catch-all routes (#36088)."""
|
||||
|
||||
from starlette.routing import WebSocketRoute
|
||||
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,
|
||||
|
|
@ -12,21 +13,23 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
|||
|
||||
|
||||
def test_openai_websocket_passthrough_routes_registered():
|
||||
ws_paths = {
|
||||
route.path
|
||||
for route in router.routes
|
||||
if isinstance(route, WebSocketRoute)
|
||||
}
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_websocket_forwards_query_and_keeps_provider_auth():
|
||||
def _mock_websocket(path: str, query: str) -> MagicMock:
|
||||
websocket = MagicMock()
|
||||
websocket.url.query = "model=gpt-4o-realtime-preview"
|
||||
websocket.url.path = path
|
||||
websocket.url.query = query
|
||||
websocket.close = AsyncMock()
|
||||
user = MagicMock()
|
||||
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(
|
||||
|
|
@ -45,10 +48,72 @@ async def test_openai_websocket_forwards_query_and_keeps_provider_auth():
|
|||
await openai_websocket_proxy_route(
|
||||
websocket=websocket,
|
||||
endpoint="v1/realtime",
|
||||
user_api_key_dict=user,
|
||||
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"
|
||||
websocket.close.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
|
||||
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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue