diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 49e4227e5d2..bc812547d6b 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -19,6 +19,7 @@ class CodexRealtimeCall(BaseModel): model: str alias: str api_base: str | None = None + extra_headers: Mapping[str, str] | None = None owner: str expires_at: float @@ -26,6 +27,7 @@ class CodexRealtimeCall(BaseModel): class ChatGPTCallRouting(BaseModel): model: str api_base: str | None = None + extra_headers: Mapping[str, str] | None = None class CodexSidebandRequest(TypedDict): @@ -33,6 +35,7 @@ class CodexSidebandRequest(TypedDict): model: ReadOnly[str] chatgpt_realtime_call_id: ReadOnly[str] query_params: ReadOnly[RealtimeQueryParams] + extra_headers: ReadOnly[Mapping[str, str] | None] def build_call_request( @@ -67,6 +70,7 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire owner=owner, expires_at=expires_at, api_base=routing.api_base, + extra_headers=routing.extra_headers, ) @@ -76,4 +80,5 @@ def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: model=f"chatgpt/{call.model}", chatgpt_realtime_call_id=call.call_id, query_params=RealtimeQueryParams(model=call.model), + extra_headers=call.extra_headers, ) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 26e2f0ffad9..97783c2bf38 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -12,11 +12,19 @@ from litellm.types.router import GenericLiteLLMParams from litellm.utils import get_model_info from .authenticator import Authenticator +from .common_utils import without_oauth_identity_headers from .responses.transformation import ChatGPTResponsesAPIConfig +def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping[str, str]: + validated: Final = TypeAdapter(Mapping[str, str]).validate_python( + without_oauth_identity_headers(headers or MappingProxyType({})) + ) + return MappingProxyType({key.lower(): value for key, value in validated.items()}) + + def realtime_headers( - params: GenericLiteLLMParams, headers: Mapping[str, str] + params: GenericLiteLLMParams, headers: Mapping[str, str], extra_headers: Mapping[str, object] | None = None ) -> dict[str, str]: # mutable-ok: HTTP handler header contract forwarded: Final = MappingProxyType( { @@ -32,6 +40,7 @@ def realtime_headers( litellm_params=params, ), **forwarded, + **configured_realtime_headers(extra_headers), } @@ -48,9 +57,11 @@ class ChatGPTRealtime(OpenAIRealtime): def get_api_base(api_base: str | None = None) -> str: return api_base or Authenticator.get_api_base(default_base="https://api.openai.com/v1") - def __init__(self, params: GenericLiteLLMParams, headers: Mapping[str, str]) -> None: + def __init__( + self, params: GenericLiteLLMParams, headers: Mapping[str, str], extra_headers: Mapping[str, object] | None = None + ) -> None: super().__init__() - self._profile_headers = realtime_headers(params, headers) + self._profile_headers = realtime_headers(params, headers, extra_headers) self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None)) def _get_additional_headers( diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 9d0d2075870..e5ec85b5c2a 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -304,12 +304,13 @@ async def arealtime_calls( api_version=litellm_params.api_version, ) if custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtime + from litellm.llms.chatgpt.realtime import ChatGPTRealtime, configured_realtime_headers response.extensions["chatgpt_realtime"] = MappingProxyType( { "model": model_name, "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), + "extra_headers": configured_realtime_headers(kwargs.get("extra_headers")), } ) return response @@ -460,7 +461,7 @@ async def _arealtime( elif _custom_llm_provider == "chatgpt": from litellm.llms.chatgpt.realtime import ChatGPTRealtime - await ChatGPTRealtime(litellm_params, websocket.headers).async_realtime( + await ChatGPTRealtime(litellm_params, websocket.headers, headers).async_realtime( model=model, websocket=websocket, logging_obj=litellm_logging_obj, diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 3ea5b1fffd2..06f04b9b81c 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -14,10 +14,12 @@ def test_signaling_rejects_invalid_upstream_call_id(location): def test_signaling_preserves_selected_model_for_sideband(): response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex"}}) + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", + "extra_headers": {"x-gateway-route": "voice"}}}) call = parse_call_response(response, "voice", "owner", 1000) request = build_sideband_request(call) assert request["api_base"] == "https://voice.example/codex" assert request["model"] == "chatgpt/gpt-live-1-codex" assert request["chatgpt_realtime_call_id"] == "rtc_provider" assert request["query_params"] == {"model": "gpt-live-1-codex"} + assert request["extra_headers"] == {"x-gateway-route": "voice"} diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index d41b64f3178..3dd74a64d8d 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -1,7 +1,11 @@ +import asyncio import json +from types import SimpleNamespace +from unittest.mock import AsyncMock import httpx import pytest +from websockets.asyncio.server import serve import litellm from litellm.llms.chatgpt.realtime import ChatGPTRealtime @@ -9,6 +13,41 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.router import GenericLiteLLMParams +@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={}, + receive_text=AsyncMock(side_effect=RuntimeError("client disconnected")), + 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}", + 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) + assert headers["x-deployment-header"] == "configured" + assert headers["x-gateway-route"] == "voice" + assert headers["openai-alpha"] == "configured-value" + assert headers["authorization"] == "Bearer test-token-default" + assert headers["chatgpt-account-id"] == "test-account-default" + assert "cookie" not in headers + + @pytest.mark.asyncio @pytest.mark.parametrize("api_base", [None, "https://voice.example/backend-api/codex"]) async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, api_base): @@ -36,6 +75,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap client=client, ) assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1") + assert response.extensions["chatgpt_realtime"]["extra_headers"] == {"openai-alpha": "quicksilver=v2", "x-gateway-route": "voice"} assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com") assert response.status_code == 201 assert requests[0].url.path == "/backend-api/codex/realtime/calls" 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 59c08c09b00..303db86b4a7 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -52,11 +52,13 @@ def test_sideband_token_binds_owner_and_model(monkeypatch): call_id="rtc_test", model="gpt-live-1-codex", alias="gpt-live-1-codex", + extra_headers={"x-gateway-secret": "configured-secret"}, owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), expires_at=time.time() + 300, ) token = encode_call(call) assert "/" not in token + assert "configured-secret" not in token assert decode_call(token, "Bearer test-owner") == call with pytest.raises(HTTPException) as error: decode_call(token, "Bearer different-owner")