fix(chatgpt): retain gateway headers across realtime connections

This commit is contained in:
jibanez-staticduo 2026-09-09 23:50:03 +02:00
parent bcac3b7234
commit 4b85d94de4
No known key found for this signature in database
6 changed files with 67 additions and 6 deletions

View file

@ -19,6 +19,7 @@ class CodexRealtimeCall(BaseModel):
model: str model: str
alias: str alias: str
api_base: str | None = None api_base: str | None = None
extra_headers: Mapping[str, str] | None = None
owner: str owner: str
expires_at: float expires_at: float
@ -26,6 +27,7 @@ class CodexRealtimeCall(BaseModel):
class ChatGPTCallRouting(BaseModel): class ChatGPTCallRouting(BaseModel):
model: str model: str
api_base: str | None = None api_base: str | None = None
extra_headers: Mapping[str, str] | None = None
class CodexSidebandRequest(TypedDict): class CodexSidebandRequest(TypedDict):
@ -33,6 +35,7 @@ class CodexSidebandRequest(TypedDict):
model: ReadOnly[str] model: ReadOnly[str]
chatgpt_realtime_call_id: ReadOnly[str] chatgpt_realtime_call_id: ReadOnly[str]
query_params: ReadOnly[RealtimeQueryParams] query_params: ReadOnly[RealtimeQueryParams]
extra_headers: ReadOnly[Mapping[str, str] | None]
def build_call_request( def build_call_request(
@ -67,6 +70,7 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire
owner=owner, owner=owner,
expires_at=expires_at, expires_at=expires_at,
api_base=routing.api_base, 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}", model=f"chatgpt/{call.model}",
chatgpt_realtime_call_id=call.call_id, chatgpt_realtime_call_id=call.call_id,
query_params=RealtimeQueryParams(model=call.model), query_params=RealtimeQueryParams(model=call.model),
extra_headers=call.extra_headers,
) )

View file

@ -12,11 +12,19 @@ from litellm.types.router import GenericLiteLLMParams
from litellm.utils import get_model_info from litellm.utils import get_model_info
from .authenticator import Authenticator from .authenticator import Authenticator
from .common_utils import without_oauth_identity_headers
from .responses.transformation import ChatGPTResponsesAPIConfig 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( 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 ) -> dict[str, str]: # mutable-ok: HTTP handler header contract
forwarded: Final = MappingProxyType( forwarded: Final = MappingProxyType(
{ {
@ -32,6 +40,7 @@ def realtime_headers(
litellm_params=params, litellm_params=params,
), ),
**forwarded, **forwarded,
**configured_realtime_headers(extra_headers),
} }
@ -48,9 +57,11 @@ class ChatGPTRealtime(OpenAIRealtime):
def get_api_base(api_base: str | None = None) -> str: 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") 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__() 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)) self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None))
def _get_additional_headers( def _get_additional_headers(

View file

@ -304,12 +304,13 @@ async def arealtime_calls(
api_version=litellm_params.api_version, api_version=litellm_params.api_version,
) )
if custom_llm_provider == "chatgpt": 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( response.extensions["chatgpt_realtime"] = MappingProxyType(
{ {
"model": model_name, "model": model_name,
"api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base),
"extra_headers": configured_realtime_headers(kwargs.get("extra_headers")),
} }
) )
return response return response
@ -460,7 +461,7 @@ async def _arealtime(
elif _custom_llm_provider == "chatgpt": elif _custom_llm_provider == "chatgpt":
from litellm.llms.chatgpt.realtime import ChatGPTRealtime 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, model=model,
websocket=websocket, websocket=websocket,
logging_obj=litellm_logging_obj, logging_obj=litellm_logging_obj,

View file

@ -14,10 +14,12 @@ def test_signaling_rejects_invalid_upstream_call_id(location):
def test_signaling_preserves_selected_model_for_sideband(): def test_signaling_preserves_selected_model_for_sideband():
response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, 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) call = parse_call_response(response, "voice", "owner", 1000)
request = build_sideband_request(call) request = build_sideband_request(call)
assert request["api_base"] == "https://voice.example/codex" assert request["api_base"] == "https://voice.example/codex"
assert request["model"] == "chatgpt/gpt-live-1-codex" assert request["model"] == "chatgpt/gpt-live-1-codex"
assert request["chatgpt_realtime_call_id"] == "rtc_provider" assert request["chatgpt_realtime_call_id"] == "rtc_provider"
assert request["query_params"] == {"model": "gpt-live-1-codex"} assert request["query_params"] == {"model": "gpt-live-1-codex"}
assert request["extra_headers"] == {"x-gateway-route": "voice"}

View file

@ -1,7 +1,11 @@
import asyncio
import json import json
from types import SimpleNamespace
from unittest.mock import AsyncMock
import httpx import httpx
import pytest import pytest
from websockets.asyncio.server import serve
import litellm import litellm
from litellm.llms.chatgpt.realtime import ChatGPTRealtime 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 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.asyncio
@pytest.mark.parametrize("api_base", [None, "https://voice.example/backend-api/codex"]) @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): 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, client=client,
) )
assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1") 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 requests[0].url.host == ("voice.example" if api_base else "chatgpt.com")
assert response.status_code == 201 assert response.status_code == 201
assert requests[0].url.path == "/backend-api/codex/realtime/calls" assert requests[0].url.path == "/backend-api/codex/realtime/calls"

View file

@ -52,11 +52,13 @@ def test_sideband_token_binds_owner_and_model(monkeypatch):
call_id="rtc_test", call_id="rtc_test",
model="gpt-live-1-codex", model="gpt-live-1-codex",
alias="gpt-live-1-codex", alias="gpt-live-1-codex",
extra_headers={"x-gateway-secret": "configured-secret"},
owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), owner=hashlib.sha256(b"Bearer test-owner").hexdigest(),
expires_at=time.time() + 300, expires_at=time.time() + 300,
) )
token = encode_call(call) token = encode_call(call)
assert "/" not in token assert "/" not in token
assert "configured-secret" not in token
assert decode_call(token, "Bearer test-owner") == call assert decode_call(token, "Bearer test-owner") == call
with pytest.raises(HTTPException) as error: with pytest.raises(HTTPException) as error:
decode_call(token, "Bearer different-owner") decode_call(token, "Bearer different-owner")