From bac2ad4abb7fab4926babae2641d9fb1cc5cfcaa Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 13:17:05 +0200 Subject: [PATCH] fix(chatgpt): finish call termination and type pinned sideband routing --- .../proxy/realtime_endpoints/call_sessions.py | 40 ++++++----- .../realtime_endpoints/call_supervision.py | 4 +- .../test_call_supervision.py | 72 +++++++++++++++++++ 3 files changed, 98 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index b0833b19414..b85df30bb31 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -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): diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index 98f018353ba..b574a229bbe 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -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") diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 642bfd4982b..264eccc787e 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -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])