diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index f9409ab366f..31a8836abf8 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -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 @@ -1939,8 +1939,8 @@ async def openai_proxy_route( async def openai_websocket_proxy_route( websocket: WebSocket, endpoint: str, - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), -): + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)], +) -> None: """WebSocket passthrough for OpenAI prefixes (realtime / responses.connect).""" base_target_url = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/" openai_api_key = passthrough_endpoint_router.get_credentials( @@ -1949,7 +1949,7 @@ async def openai_websocket_proxy_route( ) if openai_api_key is None: await websocket.close(code=1011) - raise Exception("Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.") + raise ValueError("Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.") encoded_endpoint = httpx.URL(endpoint).path if not encoded_endpoint.startswith("/"): @@ -1960,7 +1960,6 @@ async def openai_websocket_proxy_route( path=encoded_endpoint, custom_llm_provider=litellm.LlmProviders.OPENAI, ) - # HTTP(S) base -> WS(S) target for the upgrade. if updated_url.startswith("https://"): wss_target = "wss://" + updated_url[len("https://") :] elif updated_url.startswith("http://"): @@ -1968,12 +1967,17 @@ async def openai_websocket_proxy_route( else: wss_target = updated_url - return await websocket_passthrough_request( + query_string = websocket.url.query + if query_string: + separator = "&" if "?" in wss_target else "?" + wss_target = f"{wss_target}{separator}{query_string}" + + await websocket_passthrough_request( websocket=websocket, target=wss_target, custom_headers={"Authorization": f"Bearer {openai_api_key}"}, user_api_key_dict=user_api_key_dict, - forward_headers=True, + forward_headers=False, endpoint=f"/openai/{endpoint}", accept_websocket=True, ) diff --git a/tests/test_litellm/proxy/test_openai_ws_passthrough_routes.py b/tests/test_litellm/proxy/test_openai_ws_passthrough_routes.py index e0184c6c428..9101cd4b780 100644 --- a/tests/test_litellm/proxy/test_openai_ws_passthrough_routes.py +++ b/tests/test_litellm/proxy/test_openai_ws_passthrough_routes.py @@ -1,8 +1,14 @@ -"""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 -from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import router +import pytest + +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + openai_websocket_proxy_route, + router, +) def test_openai_websocket_passthrough_routes_registered(): @@ -13,3 +19,36 @@ def test_openai_websocket_passthrough_routes_registered(): } 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(): + websocket = MagicMock() + websocket.url.query = "model=gpt-4o-realtime-preview" + websocket.close = AsyncMock() + user = MagicMock() + + 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.BaseOpenAIPassThroughHandler._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=user, + ) + + 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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 752572f9863..d175f94c634 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -8483,6 +8483,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; @@ -8983,6 +9003,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; @@ -46427,6 +46467,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; @@ -47286,6 +47344,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;