fix(chatgpt): protect realtime identity and reserve hangup time

This commit is contained in:
jibanez-staticduo 2026-09-10 17:08:07 +02:00
parent c360d187f5
commit 8ec259bdbf
No known key found for this signature in database
6 changed files with 201 additions and 21 deletions

View file

@ -5,7 +5,6 @@ from typing import TYPE_CHECKING, Final
from httpx import URL, QueryParams
from pydantic import TypeAdapter
from websockets.exceptions import ConnectionClosed
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.llms.openai.realtime.handler import OpenAIRealtime
@ -114,6 +113,8 @@ class ChatGPTRealtime(OpenAIRealtime):
)
async def close_call(self, connection: "ClientConnection", model: str, api_base: str) -> None:
from websockets.exceptions import ConnectionClosed
if realtime_endpoint(model) == "live":
try:
await connection.send('{"type":"session.close"}')
@ -124,7 +125,8 @@ class ChatGPTRealtime(OpenAIRealtime):
await self.hangup_call(api_base)
async def hangup_call(self, api_base: str) -> None:
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.utils import LlmProviders
base: Final = URL(api_base)
url: Final = base.copy_with(
@ -132,12 +134,9 @@ class ChatGPTRealtime(OpenAIRealtime):
path=f"{base.path.rstrip('/')}/realtime/calls/{self._call_id}/hangup",
params=tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")),
)
client: Final = AsyncHTTPHandler()
try:
response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10)
response.raise_for_status()
finally:
await client.close()
client: Final = get_async_httpx_client(llm_provider=LlmProviders.CHATGPT)
response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10)
response.raise_for_status()
@staticmethod
def get_api_base(api_base: str | None = None) -> str:
@ -182,7 +181,9 @@ class ChatGPTRealtime(OpenAIRealtime):
base.copy_with(
scheme="wss" if base.scheme in ("https", "wss") else "ws",
path=f"{base.path.rstrip('/')}/{endpoint}",
params=query_params,
params=QueryParams(TypeAdapter(Mapping[str, str | None]).validate_python(query_params)).merge(
tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id"))
),
)
)

View file

@ -6330,6 +6330,9 @@ class BaseLLMHTTPHandler:
Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and
header auth when available; falls back to the legacy OpenAI-style defaults.
"""
from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.OPENAI,
@ -6355,7 +6358,11 @@ class BaseLLMHTTPHandler:
}
if extra_headers:
headers.update(extra_headers)
headers.update(
without_oauth_identity_headers(extra_headers)
if isinstance(provider_config, ChatGPTRealtimeHTTPConfig)
else extra_headers
)
logging_obj.pre_call(
input=request_data,

View file

@ -120,12 +120,19 @@ class CallSupervisor:
try:
if not self._terminal:
deadline: Final = asyncio.get_running_loop().time() + self._termination_timeout
primary_deadline: Final = (
deadline - self._termination_timeout / 2
if self._terminal_usage_required and self._force_close_call is not None
else deadline
)
try:
await asyncio.wait_for(self._close_call(), timeout=self._termination_timeout)
await asyncio.wait_for(
self._close_call(), timeout=max(0.0, primary_deadline - asyncio.get_running_loop().time())
)
self._close_confirmed = True
except Exception: # noqa: BLE001 # provider exceptions can contain credentials
verbose_proxy_logger.error("Realtime observer could not terminate upstream call")
await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time()))
await self._drain(reader, timeout=max(0.0, primary_deadline - asyncio.get_running_loop().time()))
if self._terminal_usage_required and not self._terminal and self._force_close_call is not None:
remaining: Final = max(0.0, deadline - asyncio.get_running_loop().time())
try:

View file

@ -1,4 +1,5 @@
import json
import sys
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
@ -14,13 +15,14 @@ from litellm.types.router import GenericLiteLLMParams
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", ["closed", "network"])
@pytest.mark.parametrize(
"hangup_status, expectation", [(200, nullcontext()), (503, pytest.raises(httpx.HTTPStatusError))]
)
async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, expectation, chatgpt_tokens):
@pytest.mark.parametrize("hangup_status", [200, 503])
async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, chatgpt_tokens, monkeypatch):
from websockets.exceptions import ConnectionClosedOK
from websockets.frames import Close
from litellm.caching.llm_caching_handler import LLMClientCache
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache())
handler = ChatGPTRealtime(
GenericLiteLLMParams(
chatgpt_realtime_call_id="rtc_live_closed",
@ -46,15 +48,21 @@ async def test_live_closed_observer_uses_independent_hangup(failure, hangup_stat
return httpx.Response(hangup_status)
client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
with patch("httpx.AsyncClient", return_value=client):
with expectation:
await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1")
assert len(requests) == 1
try:
with patch("httpx.AsyncClient", return_value=client) as create_client:
for _ in range(2):
with pytest.raises(httpx.HTTPStatusError) if hangup_status == 503 else nullcontext():
await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1")
assert not client.is_closed
create_client.assert_called_once()
finally:
await client.aclose()
assert len(requests) == 2
assert requests[0].method == "POST"
assert str(requests[0].url) == "https://gateway.example/v1/realtime/calls/rtc_live_closed/hangup?gateway=tenant"
assert requests[0].headers["x-gateway-token"] == "test-only"
assert requests[0].headers["Authorization"] == "Bearer test-token-default"
assert client.is_closed
assert requests[0].extensions["timeout"]["read"] == 10
@pytest.mark.asyncio
@ -305,6 +313,63 @@ def test_realtime_uses_platform_endpoint_with_oauth_headers(model, endpoint, cha
assert headers["openai-alpha"] == "quicksilver=v2"
@pytest.mark.parametrize("endpoint", ["live", "realtime"])
def test_new_realtime_session_preserves_gateway_query(endpoint, chatgpt_tokens, local_model_cost_map):
model = "gpt-live-1-codex" if endpoint == "live" else "gpt-realtime-1.5"
handler = ChatGPTRealtime(
GenericLiteLLMParams(
chatgpt_token_dir=chatgpt_tokens,
chatgpt_realtime_client_query={"intent": "conversation", "architecture": "client-architecture"},
extra_query={
"gateway_token": "opaque +/& value",
"intent": "gateway-intent",
"architecture": "gateway-architecture",
"model": "other-model",
"call_id": "rtc_other",
},
),
{},
)
url = httpx.URL(handler._construct_url("https://gateway.example/v1", {"model": model, "intent": "query-intent"}))
assert url.path == f"/v1/{endpoint}"
assert dict(url.params) == {
"model": model,
"gateway_token": "opaque +/& value",
"intent": "gateway-intent",
"architecture": "gateway-architecture",
}
@pytest.mark.asyncio
async def test_openai_http_call_does_not_require_websockets(monkeypatch):
monkeypatch.delitem(sys.modules, "litellm.llms.chatgpt.realtime", raising=False)
for name in tuple(sys.modules):
if name == "websockets" or name.startswith("websockets."):
monkeypatch.delitem(sys.modules, name)
monkeypatch.setitem(sys.modules, "websockets", None)
requests = []
def respond(request):
requests.append(request)
return httpx.Response(201, text="v=0\r\n")
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
response = await litellm.arealtime_calls(
model="openai/gpt-realtime-1.5",
openai_ephemeral_key="test-only",
sdp_body=b"v=0\r\n",
api_key="test-only",
client=client,
)
assert response.status_code == 201
assert len(requests) == 1
assert requests[0].url.path == "/v1/realtime/calls"
finally:
await client.client.aclose()
@pytest.mark.parametrize("endpoint", ["live", "realtime"])
@pytest.mark.parametrize("call_id", [None, "rtc_metadata"])
def test_realtime_routes_new_models_using_registered_metadata(endpoint, call_id, chatgpt_tokens, local_model_cost_map):

View file

@ -3620,3 +3620,62 @@ def test_image_edit_handler_keeps_the_sync_transform():
assert config.transform_calls == ["sync"]
assert captured["body"] == {"transformed_by": "sync"}
assert response.data[0].b64_json == "sync"
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"])
@pytest.mark.parametrize("provider", ["chatgpt", "openai"])
@pytest.mark.parametrize("authorization_header", ["Authorization", "aUtHoRiZaTiOn"])
async def test_realtime_http_sessions_preserve_provider_identity(
endpoint, provider, authorization_header, tmp_path, monkeypatch
):
import time
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig
from litellm.types.router import GenericLiteLLMParams
monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path))
monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json")
(tmp_path / "auth.json").write_text(
json.dumps({"access_token": "test-resolved", "account_id": "test-selected", "expires_at": time.time() + 3600})
)
config = ChatGPTRealtimeHTTPConfig(GenericLiteLLMParams()) if provider == "chatgpt" else OpenAIRealtimeHTTPConfig()
requests = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"id": "session-test"})
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
response = await BaseLLMHTTPHandler()._async_realtime_session_post(
endpoint=endpoint,
api_base="https://gateway.example/v1",
api_key="test-openai",
request_data={"session": {"model": "gpt-realtime-1.5"}},
logging_obj=Mock(),
timeout=5,
provider_config=config,
model="gpt-realtime-1.5",
extra_headers={
authorization_header: "Bearer test-override",
"CHATGPT-ACCOUNT-ID": "test-other-account",
"x-gateway-route": "required",
},
client=client,
)
assert response.status_code == 200
assert not client.client.is_closed
finally:
await client.client.aclose()
assert len(requests) == 1
assert requests[0].url.path == f"/v1/realtime/{endpoint}"
assert requests[0].headers["x-gateway-route"] == "required"
if provider == "chatgpt":
assert requests[0].headers.get_list("authorization") == ["Bearer test-resolved"]
assert requests[0].headers.get_list("chatgpt-account-id") == ["test-selected"]
else:
assert requests[0].headers.get_list("authorization")[-1] == "Bearer test-override"
assert requests[0].headers["chatgpt-account-id"] == "test-other-account"

View file

@ -30,6 +30,47 @@ class Socket:
self.closed = True
@pytest.mark.asyncio
@pytest.mark.parametrize("stalled_step", ["close", "drain"])
async def test_live_initial_close_reserves_time_for_independent_hangup(stalled_step):
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
close_cancelled = asyncio.Event()
async def close():
if stalled_step == "close":
try:
await asyncio.Event().wait()
finally:
close_cancelled.set()
async def force_close():
await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}})
force = AsyncMock(side_effect=force_close)
sink = Sink(logger)
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close,
force_close_call=force,
drain_timeout=1,
termination_timeout=0.08,
)
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await asyncio.wait_for(supervisor.close(), timeout=0.5)
force.assert_awaited_once()
assert close_cancelled.is_set() == (stalled_step == "close")
assert any(event["type"] == "session.closed" for event in sink.events)
assert not logger.model_call_details.get("realtime_usage_incomplete")
assert sink.logs == 1
assert socket.closed
@pytest.mark.asyncio
@pytest.mark.parametrize("fallback", ["terminal", "no_terminal", "timeout"])
async def test_live_unacknowledged_close_uses_bounded_independent_hangup(monkeypatch, fallback):