diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index bc812547d6b..78e71963b7f 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -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") diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 10a888969de..357da68f5f1 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -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 diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index e5ec85b5c2a..bceab5db353 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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 diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 3dd74a64d8d..d0760617ba0 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -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" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 303db86b4a7..08a86d93644 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -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():