fix(chatgpt): finish call termination and type pinned sideband routing

This commit is contained in:
jibanez-staticduo 2026-09-10 13:17:05 +02:00
parent 7cd2c354d8
commit bac2ad4abb
No known key found for this signature in database
3 changed files with 98 additions and 18 deletions

View file

@ -2,12 +2,14 @@ import base64
import hashlib
import json
import time
from collections.abc import Mapping
from contextlib import AsyncExitStack
from types import MappingProxyType
from typing import Final, Literal
import httpx
from fastapi import HTTPException, Request, Response, WebSocket
from pydantic import TypeAdapter
from starlette.types import Message
from litellm._logging import verbose_proxy_logger
@ -44,10 +46,6 @@ from litellm.types.router import GenericLiteLLMParams
async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth) -> None:
from collections.abc import Mapping
from pydantic import TypeAdapter
import litellm
from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, CallSupervisor
@ -362,19 +360,27 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP
subprotocol=next((p for p in protocols if not p.startswith("openai-insecure-api-key.")), None)
)
await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call
**{ # mutable-ok: retain processed policy metadata while pinning the existing call's routing
**processed,
**build_sideband_request(call),
"extra_headers": MappingProxyType(
{
**configured_realtime_headers(processed.get("extra_headers")),
**configured_realtime_headers(call.extra_headers),
}
),
"websocket": websocket,
"user_api_key_dict": auth,
"chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None,
}
model=f"chatgpt/{call.model}",
websocket=websocket,
**{
key: value
for key, value in { # mutable-ok: retain processed metadata with pinned routing
**processed,
**build_sideband_request(call),
"extra_headers": MappingProxyType(
{
**configured_realtime_headers(
TypeAdapter(Mapping[str, object] | None).validate_python(processed.get("extra_headers"))
),
**configured_realtime_headers(call.extra_headers),
}
),
"websocket": websocket,
"user_api_key_dict": auth,
"chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None,
}.items()
if key not in ("model", "websocket")
},
)
finally:
if logging_obj is None or not logging_obj.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY):

View file

@ -44,6 +44,7 @@ class CallSupervisor:
ready_timeout: float = 20,
lifetime: float = 3600,
drain_timeout: float = 5,
termination_timeout: float = 60,
terminal_usage_required: bool = True,
) -> None:
self._upstream = upstream
@ -54,6 +55,7 @@ class CallSupervisor:
self._ready_timeout = ready_timeout
self._lifetime = lifetime
self._drain_timeout = drain_timeout
self._termination_timeout = termination_timeout
self._terminal_usage_required = terminal_usage_required
self._ready = asyncio.Event()
self._stop = asyncio.Event()
@ -112,7 +114,7 @@ class CallSupervisor:
try:
if not self._terminal:
try:
await asyncio.wait_for(self._close_call(), timeout=self._drain_timeout)
await asyncio.wait_for(self._close_call(), timeout=self._termination_timeout)
self._close_confirmed = True
except Exception: # noqa: BLE001 # provider exceptions can contain credentials
verbose_proxy_logger.error("Realtime observer could not terminate upstream call")

View file

@ -264,6 +264,78 @@ async def test_confirmed_hangup_without_terminal_usage_matches_protocol(terminal
close_call.assert_awaited_once()
@pytest.mark.asyncio
async def test_shutdown_allows_hangup_longer_than_usage_drain_timeout():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
hangup_started = asyncio.Event()
allow_hangup = asyncio.Event()
hangup_finished = asyncio.Event()
async def hangup():
hangup_started.set()
await allow_hangup.wait()
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
hangup_finished.set()
supervisor = CallSupervisor(
socket, sink, logger, UserAPIKeyAuth(), hangup, drain_timeout=0.01, termination_timeout=1
)
registry = CallSupervisors()
await socket.messages.put({"type": "session.created"})
await registry.start(supervisor)
shutdown = asyncio.create_task(registry.shutdown())
try:
await asyncio.wait_for(hangup_started.wait(), timeout=1)
await asyncio.sleep(0.04)
assert not shutdown.done()
assert not socket.closed
assert not hangup_finished.is_set()
finally:
allow_hangup.set()
await asyncio.wait_for(shutdown, timeout=1)
assert hangup_finished.is_set()
assert socket.closed
assert sink.logs == 1
assert sink.events[-1]["usage"]["total_tokens"] == 42
assert not logger.model_call_details.get("realtime_usage_incomplete")
@pytest.mark.asyncio
async def test_termination_timeout_cancels_hangup_and_finishes_cleanup():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
hangup_cancelled = asyncio.Event()
async def hangup():
try:
await asyncio.Event().wait()
finally:
hangup_cancelled.set()
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
hangup,
drain_timeout=0.01,
termination_timeout=0.02,
terminal_usage_required=False,
)
await socket.messages.put({"type": "session.created"})
await supervisor.start()
await asyncio.wait_for(supervisor.close(), timeout=1)
assert hangup_cancelled.is_set()
assert socket.closed
assert sink.logs == 1
assert logger.model_call_details["realtime_usage_incomplete"] is True
@pytest.mark.asyncio
@pytest.mark.parametrize("closure", ["eof", "normal_close", "error"])
@pytest.mark.parametrize("hangup_succeeds", [True, False])