mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): preserve OpenAI WS query params and provider auth
Forward realtime model query string, keep OPENAI_API_KEY (forward_headers=False), satisfy ruff strict gates, sync dashboard OpenAPI types, and cover the behavior in tests.
This commit is contained in:
parent
8d6b8d2ce9
commit
55b52970ea
3 changed files with 128 additions and 9 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
76
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
76
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue