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
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,
)

View file

@ -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(

View file

@ -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,

View file

@ -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"}

View file

@ -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"

View file

@ -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")