mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): finish call termination and type pinned sideband routing
This commit is contained in:
parent
7cd2c354d8
commit
bac2ad4abb
3 changed files with 98 additions and 18 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue