mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): preserve deployment headers in routed signaling
This commit is contained in:
parent
ef8d102af5
commit
e5b2f10220
5 changed files with 83 additions and 21 deletions
|
|
@ -49,7 +49,7 @@ def build_call_request(
|
|||
"extra_query": { # mutable-ok: router request parameters
|
||||
key: value for key, value in query.items() if key in ("intent", "architecture")
|
||||
},
|
||||
"extra_headers": { # mutable-ok: router request headers
|
||||
"chatgpt_realtime_client_headers": { # mutable-ok: router request headers
|
||||
key: value
|
||||
for key, value in headers.items()
|
||||
if key in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation")
|
||||
|
|
|
|||
|
|
@ -23,6 +23,25 @@ def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping
|
|||
return MappingProxyType({key.lower(): value for key, value in validated.items()})
|
||||
|
||||
|
||||
def realtime_call_headers(params: GenericLiteLLMParams) -> dict[str, str]: # mutable-ok: HTTP handler header contract
|
||||
inbound: Final = TypeAdapter(Mapping[str, str]).validate_python(
|
||||
getattr(params, "chatgpt_realtime_client_headers", None) or MappingProxyType({})
|
||||
)
|
||||
configured: Final = TypeAdapter(Mapping[str, object]).validate_python(
|
||||
getattr(params, "extra_headers", None) or MappingProxyType({})
|
||||
)
|
||||
return { # mutable-ok: HTTP handler header contract
|
||||
**MappingProxyType(
|
||||
{
|
||||
key.lower(): value
|
||||
for key, value in inbound.items()
|
||||
if key.lower() in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation")
|
||||
}
|
||||
),
|
||||
**configured_realtime_headers(configured),
|
||||
}
|
||||
|
||||
|
||||
def realtime_headers(
|
||||
params: GenericLiteLLMParams, headers: Mapping[str, str], extra_headers: Mapping[str, object] | None = None
|
||||
) -> dict[str, str]: # mutable-ok: HTTP handler header contract
|
||||
|
|
|
|||
|
|
@ -261,6 +261,8 @@ async def arealtime_calls(
|
|||
timeout: float | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
from litellm.llms.chatgpt.realtime import realtime_call_headers
|
||||
|
||||
model_name = model or "gpt-4o-realtime-preview"
|
||||
litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj")
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
|
@ -283,6 +285,9 @@ async def arealtime_calls(
|
|||
)
|
||||
if session is not None:
|
||||
session = _with_resolved_session_model(session, model_name)
|
||||
call_headers: Final = (
|
||||
realtime_call_headers(litellm_params) if custom_llm_provider == "chatgpt" else kwargs.get("extra_headers")
|
||||
)
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model_name,
|
||||
|
|
@ -299,7 +304,7 @@ async def arealtime_calls(
|
|||
provider_config=provider_config,
|
||||
model=model_name,
|
||||
session_config=session,
|
||||
extra_headers=kwargs.get("extra_headers"),
|
||||
extra_headers=call_headers,
|
||||
client=kwargs.get("client"),
|
||||
api_version=litellm_params.api_version,
|
||||
)
|
||||
|
|
@ -310,7 +315,7 @@ async def arealtime_calls(
|
|||
{
|
||||
"model": model_name,
|
||||
"api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base),
|
||||
"extra_headers": configured_realtime_headers(kwargs.get("extra_headers")),
|
||||
"extra_headers": configured_realtime_headers(call_headers),
|
||||
}
|
||||
)
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -1,11 +1,9 @@
|
|||
import asyncio
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from websockets.asyncio.server import serve
|
||||
|
||||
import litellm
|
||||
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
|
||||
|
|
@ -13,16 +11,48 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
|||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("inbound_headers", [{}, {"openai-alpha": "quicksilver=v2"}])
|
||||
async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, chatgpt_tokens, monkeypatch):
|
||||
from litellm.llms.chatgpt.codex import CodexRealtimeOffer, build_call_request
|
||||
|
||||
monkeypatch.setenv("CHATGPT_TOKEN_DIR", chatgpt_tokens)
|
||||
requests = []
|
||||
|
||||
def respond(request):
|
||||
requests.append(request)
|
||||
return httpx.Response(201, text="v=0\r\n", headers={"location": "/v1/realtime/calls/rtc_test"})
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "voice-gateway",
|
||||
"litellm_params": {
|
||||
"model": "chatgpt/gpt-live-1-codex",
|
||||
"api_base": "https://voice.example/backend-api/codex",
|
||||
"extra_headers": {"x-gateway-route": "configured"},
|
||||
},
|
||||
}
|
||||
],
|
||||
num_retries=0,
|
||||
)
|
||||
offer = CodexRealtimeOffer(sdp="v=0\r\n", session={"model": "voice-gateway"})
|
||||
try:
|
||||
response = await router.arealtime_calls(**build_call_request(offer, {}, inbound_headers), client=client)
|
||||
assert requests[0].headers.get("x-gateway-route") == "configured"
|
||||
assert response.extensions["chatgpt_realtime"]["extra_headers"]["x-gateway-route"] == "configured"
|
||||
for name, value in inbound_headers.items():
|
||||
assert requests[0].headers[name] == value
|
||||
finally:
|
||||
await client.client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", ["gpt-realtime-1.5", "gpt-live-1-codex"])
|
||||
@pytest.mark.parametrize("call_id", [None, "rtc_existing"])
|
||||
async def test_websocket_forwards_configured_headers_without_client_identity(model, call_id, chatgpt_tokens):
|
||||
captured = asyncio.get_running_loop().create_future()
|
||||
|
||||
async def receive_connection(connection):
|
||||
captured.set_result(connection.request.headers)
|
||||
await connection.wait_closed()
|
||||
|
||||
websocket = SimpleNamespace(
|
||||
headers={"authorization": "Bearer client", "cookie": "private-cookie", "openai-alpha": "client-value"},
|
||||
scope={},
|
||||
|
|
@ -30,16 +60,23 @@ async def test_websocket_forwards_configured_headers_without_client_identity(mod
|
|||
send_text=AsyncMock(),
|
||||
close=AsyncMock(),
|
||||
)
|
||||
async with serve(receive_connection, "127.0.0.1", 0) as gateway:
|
||||
port = gateway.sockets[0].getsockname()[1]
|
||||
await asyncio.wait_for(litellm._arealtime(
|
||||
model=f"chatgpt/{model}", websocket=websocket, api_base=f"http://127.0.0.1:{port}",
|
||||
with patch("websockets.connect") as connect:
|
||||
connect.return_value.__aenter__ = AsyncMock(side_effect=RuntimeError("stop before streaming"))
|
||||
await litellm._arealtime(
|
||||
model=f"chatgpt/{model}",
|
||||
websocket=websocket,
|
||||
api_base="https://voice.example/codex",
|
||||
chatgpt_realtime_call_id=call_id,
|
||||
headers={"x-deployment-header": "configured"},
|
||||
extra_headers={"X-Gateway-Route": "voice", "OpenAI-Alpha": "configured-value",
|
||||
"aUtHoRiZaTiOn": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"},
|
||||
), timeout=10)
|
||||
headers = await asyncio.wait_for(captured, timeout=5)
|
||||
extra_headers={
|
||||
"X-Gateway-Route": "voice",
|
||||
"OpenAI-Alpha": "configured-value",
|
||||
"aUtHoRiZaTiOn": "Bearer wrong",
|
||||
"CHATGPT-ACCOUNT-ID": "wrong",
|
||||
},
|
||||
)
|
||||
connect.assert_called_once()
|
||||
headers = httpx.Headers(connect.call_args.kwargs["additional_headers"])
|
||||
assert headers["x-deployment-header"] == "configured"
|
||||
assert headers["x-gateway-route"] == "voice"
|
||||
assert headers["openai-alpha"] == "configured-value"
|
||||
|
|
|
|||
|
|
@ -199,7 +199,8 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
data = kwargs["data"]
|
||||
assert data["sdp_body"] == b"v=0\r\n"
|
||||
assert data["session"] == session
|
||||
assert data["extra_headers"] == {"openai-alpha": "quicksilver=v2"}
|
||||
assert data["chatgpt_realtime_client_headers"] == {"openai-alpha": "quicksilver=v2"}
|
||||
assert "extra_headers" not in data
|
||||
assert data["extra_query"] == {"intent": "quicksilver", "architecture": "avas"}
|
||||
|
||||
async def respond():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue