mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): supervise call usage and preserve gateway query routing
This commit is contained in:
parent
a9cd1dc56d
commit
d6b8360ac3
16 changed files with 1195 additions and 48 deletions
|
|
@ -8,7 +8,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
||||
|
||||
from httpx import Response
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
|
||||
import litellm
|
||||
import litellm._logging
|
||||
|
|
@ -2563,7 +2563,12 @@ def handle_realtime_stream_cost_calculation(
|
|||
if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results)
|
||||
else 0.0
|
||||
)
|
||||
total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost
|
||||
live_audio_cost: Final = handle_live_session_duration_cost(
|
||||
results=results,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_model_name=litellm_model_name,
|
||||
)
|
||||
total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost + live_audio_cost
|
||||
|
||||
_store_cost_breakdown_in_logging_obj(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
|
|
@ -2571,13 +2576,47 @@ def handle_realtime_stream_cost_calculation(
|
|||
completion_tokens_cost_usd_dollar=output_cost_per_token,
|
||||
cost_for_built_in_tools_cost_usd_dollar=0.0,
|
||||
total_cost_usd_dollar=total_cost,
|
||||
additional_costs={"transcription_cost": transcription_cost} if transcription_cost > 0 else None,
|
||||
additional_costs={
|
||||
name: cost
|
||||
for name, cost in (("transcription_cost", transcription_cost), ("live_audio_cost", live_audio_cost))
|
||||
if cost > 0
|
||||
}
|
||||
or None,
|
||||
data_residency=data_residency,
|
||||
)
|
||||
|
||||
return total_cost
|
||||
|
||||
|
||||
class _LiveSessionDurationUsage(BaseModel):
|
||||
audio_duration_ms: float = Field(strict=True, ge=0, allow_inf_nan=False)
|
||||
|
||||
|
||||
class _LiveSessionClosedEvent(BaseModel):
|
||||
usage: _LiveSessionDurationUsage
|
||||
|
||||
|
||||
def handle_live_session_duration_cost(
|
||||
results: OpenAIRealtimeStreamList,
|
||||
custom_llm_provider: str,
|
||||
litellm_model_name: str,
|
||||
) -> float:
|
||||
if any(event.get("type") == "response.done" for event in results):
|
||||
return 0.0
|
||||
terminal: Final = next((event for event in reversed(results) if event.get("type") == "session.closed"), None)
|
||||
if terminal is None:
|
||||
return 0.0
|
||||
try:
|
||||
usage: Final = _LiveSessionClosedEvent.model_validate(terminal).usage
|
||||
except ValidationError:
|
||||
return 0.0
|
||||
try:
|
||||
model_info: Final = litellm.get_model_info(model=litellm_model_name, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
return 0.0
|
||||
return usage.audio_duration_ms / 1000 * (model_info.get("input_cost_per_second") or 0.0)
|
||||
|
||||
|
||||
def handle_realtime_transcription_cost_calculation(
|
||||
results: OpenAIRealtimeStreamList,
|
||||
custom_llm_provider: str,
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from dataclasses import dataclass
|
|||
from enum import Enum, auto
|
||||
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
|
|
@ -16,6 +17,7 @@ from litellm.types.llms.openai import (
|
|||
OpenAIRealtimeEvents,
|
||||
OpenAIRealtimeOutputItemDone,
|
||||
OpenAIRealtimeResponseDelta,
|
||||
OpenAIRealtimeSessionClosed,
|
||||
OpenAIRealtimeStreamResponseBaseObject,
|
||||
OpenAIRealtimeStreamSessionEvents,
|
||||
)
|
||||
|
|
@ -139,11 +141,14 @@ class RealTimeStreaming:
|
|||
force_transcription_model: str | None = None,
|
||||
event_normalizer: RealtimeEventNormalizer | None = None,
|
||||
logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER,
|
||||
*,
|
||||
account_usage: bool = True,
|
||||
):
|
||||
self.websocket: _ClientWebSocket = websocket
|
||||
self.backend_ws = backend_ws
|
||||
self.logging_obj = logging_obj
|
||||
self._logging_worker = logging_worker
|
||||
self._account_usage = account_usage
|
||||
self.messages: list[OpenAIRealtimeEvents] = []
|
||||
self._backend_sent_frames: bool = False
|
||||
self.input_message: dict = {}
|
||||
|
|
@ -256,6 +261,9 @@ class RealTimeStreaming:
|
|||
else:
|
||||
message_obj = cast(dict[str, Any], json.loads(cast(str, message)))
|
||||
self._collect_tool_calls_from_response_done(cast(dict, message_obj))
|
||||
if message_obj.get("type") == "session.closed" and isinstance(message_obj.get("usage"), dict):
|
||||
self.messages.append(TypeAdapter(OpenAIRealtimeSessionClosed).validate_python(message_obj))
|
||||
return
|
||||
if not self._should_store_message(message_obj):
|
||||
return
|
||||
try:
|
||||
|
|
@ -410,8 +418,10 @@ class RealTimeStreaming:
|
|||
if self.logging_obj:
|
||||
self.logging_obj.pre_call(input=message, api_key="")
|
||||
|
||||
async def log_messages(self):
|
||||
async def log_messages(self, *, wait_for_dispatch: bool = False):
|
||||
"""Log messages in list"""
|
||||
if not self._account_usage:
|
||||
return
|
||||
if self.logging_obj:
|
||||
if self.input_messages:
|
||||
self.logging_obj.model_call_details["messages"] = self.input_messages
|
||||
|
|
@ -421,9 +431,12 @@ class RealTimeStreaming:
|
|||
# Route through the bounded logging worker (per-coroutine timeout +
|
||||
# concurrency cap) instead of a bare create_task, so a slow callback
|
||||
# can't leave suspended tasks pinning each call's response in memory.
|
||||
self._logging_worker.ensure_initialized_and_enqueue(
|
||||
self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
|
||||
)
|
||||
if wait_for_dispatch:
|
||||
await self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
|
||||
else:
|
||||
self._logging_worker.ensure_initialized_and_enqueue(
|
||||
self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
|
||||
)
|
||||
self.logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
|
||||
|
||||
async def _send_to_backend(self, message: str) -> bool:
|
||||
|
|
|
|||
|
|
@ -17,17 +17,22 @@ class CodexRealtimeOffer(BaseModel):
|
|||
class CodexRealtimeCall(BaseModel):
|
||||
call_id: str = Field(pattern=r"^rtc_[A-Za-z0-9_-]+$")
|
||||
model: str
|
||||
model_id: str | None = None
|
||||
alias: str
|
||||
api_base: str | None = None
|
||||
extra_headers: Mapping[str, str] | None = None
|
||||
extra_query: Mapping[str, str] | None = None
|
||||
usage_supervised: bool = False
|
||||
owner: str
|
||||
expires_at: float
|
||||
|
||||
|
||||
class ChatGPTCallRouting(BaseModel):
|
||||
model: str
|
||||
model_id: str | None = None
|
||||
api_base: str | None = None
|
||||
extra_headers: Mapping[str, str] | None = None
|
||||
extra_query: Mapping[str, str] | None = None
|
||||
|
||||
|
||||
class CodexSidebandRequest(TypedDict):
|
||||
|
|
@ -36,6 +41,7 @@ class CodexSidebandRequest(TypedDict):
|
|||
chatgpt_realtime_call_id: ReadOnly[str]
|
||||
query_params: ReadOnly[RealtimeQueryParams]
|
||||
extra_headers: ReadOnly[Mapping[str, str] | None]
|
||||
extra_query: ReadOnly[Mapping[str, str] | None]
|
||||
|
||||
|
||||
def build_call_request(
|
||||
|
|
@ -46,7 +52,7 @@ def build_call_request(
|
|||
"sdp_body": offer.sdp.encode(),
|
||||
"session": offer.session.model_dump(exclude_none=True),
|
||||
"openai_ephemeral_key": "",
|
||||
"extra_query": { # mutable-ok: router request parameters
|
||||
"chatgpt_realtime_client_query": { # mutable-ok: router request parameters
|
||||
key: value for key, value in query.items() if key in ("intent", "architecture")
|
||||
},
|
||||
"chatgpt_realtime_client_headers": { # mutable-ok: router request headers
|
||||
|
|
@ -66,11 +72,13 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire
|
|||
return CodexRealtimeCall(
|
||||
call_id=call_id,
|
||||
model=routing.model,
|
||||
model_id=routing.model_id,
|
||||
alias=alias,
|
||||
owner=owner,
|
||||
expires_at=expires_at,
|
||||
api_base=routing.api_base,
|
||||
extra_headers=routing.extra_headers,
|
||||
extra_query=routing.extra_query,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -81,4 +89,5 @@ def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest:
|
|||
chatgpt_realtime_call_id=call.call_id,
|
||||
query_params=RealtimeQueryParams(model=call.model),
|
||||
extra_headers=call.extra_headers,
|
||||
extra_query=call.extra_query,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,10 +1,12 @@
|
|||
from collections.abc import Mapping
|
||||
from enum import Enum, auto
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from httpx import URL
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.llms.openai.realtime.handler import OpenAIRealtime
|
||||
from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig
|
||||
from litellm.types.realtime import RealtimeQueryParams
|
||||
|
|
@ -15,6 +17,17 @@ from .authenticator import Authenticator
|
|||
from .common_utils import without_oauth_identity_headers
|
||||
from .responses.transformation import ChatGPTResponsesAPIConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
|
||||
class CallAccounting(Enum):
|
||||
SUPERVISED = auto()
|
||||
|
||||
|
||||
def accounts_for_call_usage(params: GenericLiteLLMParams) -> bool:
|
||||
return getattr(params, "chatgpt_call_accounting", None) is not CallAccounting.SUPERVISED
|
||||
|
||||
|
||||
def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping[str, str]:
|
||||
validated: Final = TypeAdapter(Mapping[str, str]).validate_python(
|
||||
|
|
@ -23,6 +36,18 @@ def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping
|
|||
return MappingProxyType({key.lower(): value for key, value in validated.items()})
|
||||
|
||||
|
||||
def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str]:
|
||||
inbound: Final = TypeAdapter(Mapping[str, str]).validate_python(
|
||||
getattr(params, "chatgpt_realtime_client_query", None) or MappingProxyType({})
|
||||
)
|
||||
configured: Final = TypeAdapter(Mapping[str, str]).validate_python(
|
||||
getattr(params, "extra_query", None) or MappingProxyType({})
|
||||
)
|
||||
return MappingProxyType(
|
||||
{**{key: value for key, value in inbound.items() if key in ("intent", "architecture")}, **configured}
|
||||
)
|
||||
|
||||
|
||||
def realtime_call_headers(params: GenericLiteLLMParams) -> dict[str, str]: # mutable-ok: HTTP handler header contract
|
||||
inbound: Final = TypeAdapter(Mapping[str, str]).validate_python(
|
||||
getattr(params, "chatgpt_realtime_client_headers", None) or MappingProxyType({})
|
||||
|
|
@ -72,6 +97,40 @@ def realtime_endpoint(model: str) -> str:
|
|||
|
||||
|
||||
class ChatGPTRealtime(OpenAIRealtime):
|
||||
async def open_call_connection(self, model: str, api_base: str) -> "ClientConnection":
|
||||
import websockets
|
||||
|
||||
url: Final = self._construct_url(api_base, RealtimeQueryParams(model=model))
|
||||
return await websockets.connect(
|
||||
url,
|
||||
additional_headers=self._profile_headers,
|
||||
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
ssl=self._get_ssl_config(url),
|
||||
open_timeout=20,
|
||||
)
|
||||
|
||||
async def close_call(self, connection: "ClientConnection", model: str, api_base: str) -> None:
|
||||
if realtime_endpoint(model) == "live":
|
||||
await connection.send('{"type":"session.close"}')
|
||||
return
|
||||
await self.hangup_call(api_base)
|
||||
|
||||
async def hangup_call(self, api_base: str) -> None:
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
base: Final = URL(api_base)
|
||||
url: Final = base.copy_with(
|
||||
scheme="https" if base.scheme in ("https", "wss") else "http",
|
||||
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()
|
||||
|
||||
@staticmethod
|
||||
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")
|
||||
|
|
@ -85,6 +144,7 @@ class ChatGPTRealtime(OpenAIRealtime):
|
|||
super().__init__()
|
||||
self._profile_headers = realtime_headers(params, headers, extra_headers)
|
||||
self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None))
|
||||
self._extra_query = configured_realtime_query(params)
|
||||
|
||||
def _get_additional_headers(
|
||||
self, api_key: str, *, openai_beta_realtime: bool = False
|
||||
|
|
@ -98,13 +158,16 @@ class ChatGPTRealtime(OpenAIRealtime):
|
|||
base: Final = URL(api_base)
|
||||
endpoint: Final = realtime_endpoint(query_params.get("model", ""))
|
||||
if self._call_id:
|
||||
gateway_query: Final = tuple(
|
||||
(key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")
|
||||
)
|
||||
return str(
|
||||
base.copy_with(
|
||||
scheme="wss" if base.scheme in ("https", "wss") else "ws",
|
||||
path=f"{base.path.rstrip('/')}/{endpoint}/{self._call_id}"
|
||||
if endpoint == "live"
|
||||
else f"{base.path.rstrip('/')}/realtime",
|
||||
params=() if endpoint == "live" else (("call_id", self._call_id),),
|
||||
params=gateway_query + (() if endpoint == "live" else (("call_id", self._call_id),)),
|
||||
)
|
||||
)
|
||||
return str(
|
||||
|
|
@ -138,9 +201,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig):
|
|||
return "chatgpt-oauth"
|
||||
|
||||
def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
|
||||
query: Final = TypeAdapter(Mapping[str, str]).validate_python(
|
||||
getattr(self._params, "extra_query", None) or MappingProxyType({})
|
||||
)
|
||||
query: Final = configured_realtime_query(self._params)
|
||||
return str(URL(f"{self.get_api_base(api_base).rstrip('/')}/realtime/calls", params=query))
|
||||
|
||||
def get_realtime_calls_headers(
|
||||
|
|
|
|||
|
|
@ -117,6 +117,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
|
|||
query_params: RealtimeQueryParams | None = None,
|
||||
user_api_key_dict: object | None = None,
|
||||
litellm_metadata: dict | None = None,
|
||||
account_usage: bool = True,
|
||||
**kwargs: object,
|
||||
):
|
||||
import websockets
|
||||
|
|
@ -172,6 +173,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
|
|||
model if (query_params or {}).get("intent") == "transcription" else None
|
||||
),
|
||||
event_normalizer=self._make_event_normalizer(),
|
||||
account_usage=account_usage,
|
||||
)
|
||||
await realtime_streaming.bidirectional_forward()
|
||||
|
||||
|
|
|
|||
|
|
@ -1362,6 +1362,9 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e)
|
||||
|
||||
from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS
|
||||
|
||||
await CALL_SUPERVISORS.shutdown()
|
||||
await _flush_spend_logs_queue_on_shutdown()
|
||||
|
||||
await proxy_config.stop_config_sync_subscriber()
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import base64
|
|||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from contextlib import AsyncExitStack
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
|
|
@ -11,7 +12,7 @@ from starlette.types import Message
|
|||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY
|
||||
from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY, RealTimeStreaming
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.chatgpt.codex import (
|
||||
CodexRealtimeCall,
|
||||
|
|
@ -20,7 +21,12 @@ from litellm.llms.chatgpt.codex import (
|
|||
build_sideband_request,
|
||||
parse_call_response,
|
||||
)
|
||||
from litellm.llms.chatgpt.realtime import configured_realtime_headers
|
||||
from litellm.llms.chatgpt.realtime import (
|
||||
CallAccounting,
|
||||
ChatGPTRealtime,
|
||||
configured_realtime_headers,
|
||||
realtime_endpoint,
|
||||
)
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
|
|
@ -30,7 +36,126 @@ from litellm.proxy.auth.user_api_key_auth import (
|
|||
user_api_key_auth,
|
||||
)
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
|
||||
from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
invalidate_budget_reservation_counters,
|
||||
release_or_invalidate_budget_reservation,
|
||||
)
|
||||
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
|
||||
|
||||
async def receive() -> Message:
|
||||
return {
|
||||
"type": "http.request",
|
||||
"body": json.dumps({"model": call.alias}).encode(),
|
||||
"more_body": False,
|
||||
} # mutable-ok: ASGI message
|
||||
|
||||
async def send(_message: Message) -> None:
|
||||
return None
|
||||
|
||||
supervision_owned = False # rebind-ok: supervisor owns cleanup after construction
|
||||
effective_handler: ChatGPTRealtime | None = None # rebind-ok: reuse hook-enriched credentials for cleanup
|
||||
sockets: Final = AsyncExitStack()
|
||||
try:
|
||||
observer_request: Final = Request({**request.scope}, receive=receive) # mutable-ok: ASGI request scope
|
||||
processed, logger = await process_codex_request(
|
||||
observer_request,
|
||||
{
|
||||
**build_sideband_request(call),
|
||||
"model": call.alias,
|
||||
}, # mutable-ok: common request processing enriches metadata
|
||||
auth,
|
||||
call.alias,
|
||||
"_arealtime",
|
||||
)
|
||||
pinned: Final = { # mutable-ok: logging and provider parameter contract
|
||||
**processed,
|
||||
**build_sideband_request(call),
|
||||
"extra_headers": {
|
||||
**configured_realtime_headers(
|
||||
TypeAdapter(Mapping[str, object] | None).validate_python(processed.get("extra_headers"))
|
||||
),
|
||||
**configured_realtime_headers(call.extra_headers),
|
||||
},
|
||||
"litellm_metadata": {
|
||||
**TypeAdapter(Mapping[str, object]).validate_python(processed.get("litellm_metadata") or {}),
|
||||
**(
|
||||
{"model_info": {**litellm.get_model_info(model=call.model_id), "id": call.model_id}}
|
||||
if call.model_id is not None
|
||||
else {}
|
||||
),
|
||||
},
|
||||
}
|
||||
logger.update_from_kwargs(
|
||||
kwargs=pinned,
|
||||
model=call.model,
|
||||
user=None,
|
||||
optional_params={}, # mutable-ok: logging contract
|
||||
litellm_params={
|
||||
**logger.litellm_params,
|
||||
"litellm_metadata": pinned["litellm_metadata"],
|
||||
"arealtime": True,
|
||||
}, # mutable-ok: logging contract
|
||||
custom_llm_provider="chatgpt",
|
||||
)
|
||||
params: Final = GenericLiteLLMParams.model_validate(pinned)
|
||||
handler: Final = ChatGPTRealtime(
|
||||
params, request.headers, TypeAdapter(Mapping[str, object]).validate_python(pinned["extra_headers"])
|
||||
)
|
||||
effective_handler = handler
|
||||
api_base: Final = ChatGPTRealtime.get_api_base(call.api_base)
|
||||
connection: Final = await handler.open_call_connection(call.model, api_base)
|
||||
sockets.push_async_callback(connection.close)
|
||||
|
||||
async def close_call() -> None:
|
||||
await handler.close_call(connection, call.model, api_base)
|
||||
|
||||
frontend: Final = WebSocket(
|
||||
{**request.scope, "type": "websocket"}, receive=receive, send=send
|
||||
) # mutable-ok: ASGI scope
|
||||
stream: Final = RealTimeStreaming(frontend, connection, logger, model=call.model, user_api_key_dict=auth)
|
||||
supervisor: Final = CallSupervisor(
|
||||
connection,
|
||||
stream,
|
||||
logger,
|
||||
auth,
|
||||
close_call,
|
||||
terminal_usage_required=realtime_endpoint(call.model) == "live",
|
||||
)
|
||||
supervision_owned = True
|
||||
sockets.pop_all()
|
||||
await CALL_SUPERVISORS.start(supervisor)
|
||||
except BaseException:
|
||||
if not supervision_owned:
|
||||
try:
|
||||
fallback_handler: Final = effective_handler or ChatGPTRealtime(
|
||||
GenericLiteLLMParams.model_validate(build_sideband_request(call)),
|
||||
request.headers,
|
||||
call.extra_headers,
|
||||
)
|
||||
await fallback_handler.hangup_call(ChatGPTRealtime.get_api_base(call.api_base))
|
||||
except Exception: # noqa: BLE001 # preserve original failure without logging provider credentials
|
||||
verbose_proxy_logger.error("Realtime startup cleanup could not confirm upstream termination")
|
||||
try:
|
||||
await invalidate_budget_reservation_counters(budget_reservation=auth.budget_reservation)
|
||||
except Exception: # noqa: BLE001 # cleanup errors must not replace the original startup failure
|
||||
verbose_proxy_logger.error("Realtime startup cleanup could not invalidate budget counters")
|
||||
else:
|
||||
await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation)
|
||||
finally:
|
||||
try:
|
||||
await sockets.aclose()
|
||||
except Exception: # noqa: BLE001 # socket cleanup must preserve the original startup failure
|
||||
verbose_proxy_logger.error("Realtime startup cleanup could not close observer socket")
|
||||
raise
|
||||
|
||||
|
||||
def encode_call(call: CodexRealtimeCall) -> str:
|
||||
|
|
@ -126,6 +251,7 @@ async def create_codex_realtime_call(request: Request) -> Response:
|
|||
owner_key: Final = (
|
||||
get_api_key_from_custom_header(request, custom_header) if isinstance(custom_header, str) else selected_key
|
||||
)
|
||||
supervision_started = False # rebind-ok: transfer reservation ownership only after supervision is established
|
||||
try:
|
||||
await can_key_call_resolved_model(
|
||||
model=model,
|
||||
|
|
@ -134,7 +260,10 @@ async def create_codex_realtime_call(request: Request) -> Response:
|
|||
llm_router=server.llm_router,
|
||||
)
|
||||
data: Final = build_call_request(offer, request.query_params, request.headers)
|
||||
processed, _ = await process_codex_request(request, data, auth, model, "arealtime_calls")
|
||||
signaling_auth: Final = auth.model_copy(
|
||||
update={"budget_reservation": None}
|
||||
) # mutable-ok: Pydantic update contract
|
||||
processed, _ = await process_codex_request(request, data, signaling_auth, model, "arealtime_calls")
|
||||
result: Final = await server.route_request(
|
||||
data=processed,
|
||||
route_type="arealtime_calls",
|
||||
|
|
@ -158,7 +287,12 @@ async def create_codex_realtime_call(request: Request) -> Response:
|
|||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(400, str(exc)) from exc
|
||||
token: Final = encode_call(call)
|
||||
supervised_call: Final = call.model_copy(
|
||||
update={"usage_supervised": True}
|
||||
) # mutable-ok: Pydantic update contract
|
||||
token: Final = encode_call(supervised_call)
|
||||
supervision_started = True
|
||||
await supervise_codex_call(request, supervised_call, auth)
|
||||
return Response(
|
||||
response.content,
|
||||
status_code=response.status_code,
|
||||
|
|
@ -166,7 +300,8 @@ async def create_codex_realtime_call(request: Request) -> Response:
|
|||
headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}),
|
||||
)
|
||||
finally:
|
||||
await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation)
|
||||
if not supervision_started:
|
||||
await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation)
|
||||
|
||||
|
||||
async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAPIKeyAuth) -> None:
|
||||
|
|
@ -238,6 +373,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP
|
|||
),
|
||||
"websocket": websocket,
|
||||
"user_api_key_dict": auth,
|
||||
"chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None,
|
||||
}
|
||||
)
|
||||
finally:
|
||||
|
|
|
|||
187
litellm/proxy/realtime_endpoints/call_supervision.py
Normal file
187
litellm/proxy/realtime_endpoints/call_supervision.py
Normal file
|
|
@ -0,0 +1,187 @@
|
|||
import asyncio
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from contextlib import suppress
|
||||
from typing import Final, Protocol
|
||||
|
||||
from pydantic import BaseModel
|
||||
from websockets.exceptions import ConnectionClosedOK
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
invalidate_budget_reservation_counters,
|
||||
release_or_invalidate_budget_reservation,
|
||||
)
|
||||
|
||||
|
||||
class ObserverSocket(Protocol):
|
||||
def __aiter__(self) -> AsyncIterator[str | bytes]: ...
|
||||
|
||||
async def close(self) -> None: ...
|
||||
|
||||
|
||||
class UsageSink(Protocol):
|
||||
def store_message(self, message: str) -> None: ...
|
||||
|
||||
async def log_messages(self, *, wait_for_dispatch: bool = False) -> None: ...
|
||||
|
||||
|
||||
class _ObserverEvent(BaseModel):
|
||||
type: str
|
||||
|
||||
|
||||
class CallSupervisor:
|
||||
def __init__(
|
||||
self,
|
||||
upstream: ObserverSocket,
|
||||
stream: UsageSink,
|
||||
logging_obj: Logging,
|
||||
auth: UserAPIKeyAuth,
|
||||
close_call: Callable[[], Awaitable[None]],
|
||||
*,
|
||||
ready_timeout: float = 20,
|
||||
lifetime: float = 3600,
|
||||
drain_timeout: float = 5,
|
||||
terminal_usage_required: bool = True,
|
||||
) -> None:
|
||||
self._upstream = upstream
|
||||
self._stream = stream
|
||||
self._logging = logging_obj
|
||||
self._auth = auth
|
||||
self._close_call = close_call
|
||||
self._ready_timeout = ready_timeout
|
||||
self._lifetime = lifetime
|
||||
self._drain_timeout = drain_timeout
|
||||
self._terminal_usage_required = terminal_usage_required
|
||||
self._ready = asyncio.Event()
|
||||
self._stop = asyncio.Event()
|
||||
self._started = False
|
||||
self._terminal = False
|
||||
self._close_confirmed = False
|
||||
self._task: asyncio.Task[None] | None = None
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._task is not None:
|
||||
raise RuntimeError("Call observer already started")
|
||||
self._task = asyncio.create_task(self._run())
|
||||
try:
|
||||
await asyncio.wait_for(self._ready.wait(), timeout=self._ready_timeout)
|
||||
if not self._started or self._task.done():
|
||||
raise RuntimeError("Call observer ended before session became available")
|
||||
except BaseException:
|
||||
await self.close()
|
||||
raise
|
||||
|
||||
async def close(self) -> None:
|
||||
self._stop.set()
|
||||
await self.wait()
|
||||
|
||||
async def wait(self) -> None:
|
||||
if self._task is not None:
|
||||
await asyncio.shield(self._task)
|
||||
|
||||
async def _read(self) -> None:
|
||||
try:
|
||||
await self._read_events()
|
||||
except ConnectionClosedOK:
|
||||
return
|
||||
|
||||
async def _read_events(self) -> None:
|
||||
event: _ObserverEvent
|
||||
async for message in self._upstream:
|
||||
self._stream.store_message(message.decode("utf-8") if isinstance(message, bytes) else message)
|
||||
event = _ObserverEvent.model_validate_json(message)
|
||||
if event.type in ("session.started", "session.created"):
|
||||
self._started = True
|
||||
self._ready.set()
|
||||
if event.type == "session.closed":
|
||||
self._terminal = True
|
||||
return
|
||||
|
||||
def _usage_complete(self) -> bool:
|
||||
return self._terminal or (not self._terminal_usage_required and self._close_confirmed)
|
||||
|
||||
async def _run(self) -> None:
|
||||
reader: Final = asyncio.create_task(self._read())
|
||||
stopped: Final = asyncio.create_task(self._stop.wait())
|
||||
try:
|
||||
await asyncio.wait((reader, stopped), timeout=self._lifetime, return_when=asyncio.FIRST_COMPLETED)
|
||||
finally:
|
||||
try:
|
||||
if not self._terminal:
|
||||
try:
|
||||
await asyncio.wait_for(self._close_call(), timeout=self._drain_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")
|
||||
await self._drain(reader)
|
||||
finally:
|
||||
stopped.cancel()
|
||||
reader.cancel()
|
||||
await asyncio.gather(reader, stopped, return_exceptions=True)
|
||||
with suppress(Exception):
|
||||
await self._upstream.close()
|
||||
if not self._usage_complete():
|
||||
self._logging.model_call_details["realtime_usage_incomplete"] = True
|
||||
verbose_proxy_logger.error(
|
||||
"Realtime observer ended without terminal usage; recorded usage is partial"
|
||||
)
|
||||
try:
|
||||
try:
|
||||
await self._stream.log_messages(wait_for_dispatch=True)
|
||||
finally:
|
||||
if self._started and not self._usage_complete():
|
||||
await invalidate_budget_reservation_counters(
|
||||
budget_reservation=self._auth.budget_reservation
|
||||
)
|
||||
elif not self._logging.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY):
|
||||
await release_or_invalidate_budget_reservation(
|
||||
budget_reservation=self._auth.budget_reservation
|
||||
)
|
||||
finally:
|
||||
self._ready.set()
|
||||
|
||||
async def _drain(self, reader: asyncio.Task[None]) -> None:
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.shield(reader), timeout=self._drain_timeout)
|
||||
except asyncio.TimeoutError:
|
||||
if not self._usage_complete():
|
||||
verbose_proxy_logger.error("Realtime observer timed out draining terminal usage")
|
||||
except Exception: # noqa: BLE001 # cleanup must settle the socket even when reading or closing fails
|
||||
verbose_proxy_logger.error("Realtime observer could not drain terminal usage")
|
||||
return
|
||||
|
||||
|
||||
class CallSupervisors:
|
||||
def __init__(self) -> None:
|
||||
self._tasks: tuple[asyncio.Task[None], ...] = ()
|
||||
self._calls: tuple[CallSupervisor, ...] = ()
|
||||
|
||||
async def start(self, supervisor: CallSupervisor) -> None:
|
||||
self._calls = (*self._calls, supervisor)
|
||||
try:
|
||||
await supervisor.start()
|
||||
except BaseException:
|
||||
self._calls = tuple(call for call in self._calls if call is not supervisor)
|
||||
raise
|
||||
task: Final = asyncio.create_task(self._watch(supervisor))
|
||||
self._tasks = (*self._tasks, task)
|
||||
|
||||
async def _watch(self, supervisor: CallSupervisor) -> None:
|
||||
try:
|
||||
try:
|
||||
await supervisor.wait()
|
||||
except Exception: # noqa: BLE001 # task must be consumed without exposing provider exception payloads
|
||||
verbose_proxy_logger.error("Realtime observer accounting failed")
|
||||
finally:
|
||||
self._calls = tuple(call for call in self._calls if call is not supervisor)
|
||||
self._tasks = tuple(task for task in self._tasks if task is not asyncio.current_task())
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
await asyncio.gather(*(call.close() for call in self._calls), return_exceptions=True)
|
||||
await asyncio.gather(*self._tasks, return_exceptions=True)
|
||||
|
||||
|
||||
CALL_SUPERVISORS: Final = CallSupervisors()
|
||||
|
|
@ -309,13 +309,19 @@ async def arealtime_calls(
|
|||
api_version=litellm_params.api_version,
|
||||
)
|
||||
if custom_llm_provider == "chatgpt":
|
||||
from litellm.llms.chatgpt.realtime import ChatGPTRealtime, configured_realtime_headers
|
||||
from litellm.llms.chatgpt.realtime import (
|
||||
ChatGPTRealtime,
|
||||
configured_realtime_headers,
|
||||
configured_realtime_query,
|
||||
)
|
||||
|
||||
response.extensions["chatgpt_realtime"] = MappingProxyType(
|
||||
{
|
||||
"model": model_name,
|
||||
"model_id": litellm_logging_obj.get_router_model_id(),
|
||||
"api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base),
|
||||
"extra_headers": configured_realtime_headers(call_headers),
|
||||
"extra_query": configured_realtime_query(litellm_params),
|
||||
}
|
||||
)
|
||||
return response
|
||||
|
|
@ -464,7 +470,7 @@ async def _arealtime(
|
|||
litellm_metadata=_build_litellm_metadata(kwargs),
|
||||
)
|
||||
elif _custom_llm_provider == "chatgpt":
|
||||
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
|
||||
from litellm.llms.chatgpt.realtime import ChatGPTRealtime, accounts_for_call_usage
|
||||
|
||||
await ChatGPTRealtime(litellm_params, websocket.headers, headers).async_realtime(
|
||||
model=model,
|
||||
|
|
@ -476,6 +482,7 @@ async def _arealtime(
|
|||
query_params=query_params,
|
||||
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
||||
litellm_metadata=_build_litellm_metadata(kwargs),
|
||||
account_usage=accounts_for_call_usage(litellm_params),
|
||||
)
|
||||
elif _custom_llm_provider == "openai":
|
||||
api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or "https://api.openai.com/"
|
||||
|
|
|
|||
|
|
@ -2031,6 +2031,11 @@ class OpenAIRealtimeStreamResponseBaseObject(TypedDict):
|
|||
type: str
|
||||
|
||||
|
||||
class OpenAIRealtimeSessionClosed(TypedDict):
|
||||
type: ReadOnly[Literal["session.closed"]]
|
||||
usage: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class OpenAIRealtimeConversationObject(TypedDict, total=False):
|
||||
id: str
|
||||
object: Required[Literal["realtime.conversation"]]
|
||||
|
|
@ -2236,6 +2241,7 @@ class OpenAIRealtimeEventTypes(Enum):
|
|||
|
||||
OpenAIRealtimeEvents = (
|
||||
OpenAIRealtimeStreamResponseBaseObject
|
||||
| OpenAIRealtimeSessionClosed
|
||||
| OpenAIRealtimeStreamSessionEvents
|
||||
| OpenAIRealtimeStreamResponseOutputItemAdded
|
||||
| OpenAIRealtimeResponseContentPartAdded
|
||||
|
|
|
|||
|
|
@ -3412,3 +3412,29 @@ async def test_refused_session_does_not_stamp_the_reservation_ownership_marker()
|
|||
|
||||
assert session.logging.logged_failures == (upstream_close,)
|
||||
assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in session.logging.model_call_details
|
||||
|
||||
|
||||
def test_live_terminal_usage_survives_filtered_event_logging(monkeypatch):
|
||||
from litellm.cost_calculator import RealtimeAPITokenUsageProcessor
|
||||
|
||||
def terminal():
|
||||
return {"type": "session.closed", "usage": {"audio_duration_ms": 4000, "backend_model_usage": []}}
|
||||
|
||||
monkeypatch.setattr(litellm, "logged_real_time_event_types", [])
|
||||
stream = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
|
||||
event = {**terminal(), "private_transcript": "Do not retain this text"}
|
||||
stream.store_message(event)
|
||||
assert stream.messages == [terminal()]
|
||||
usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(stream.messages)
|
||||
assert usage.total_tokens == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_live_attachment_does_not_dispatch_duplicate_usage():
|
||||
worker = MagicMock()
|
||||
logger = MagicMock()
|
||||
stream = RealTimeStreaming(MagicMock(), MagicMock(), logger, logging_worker=worker, account_usage=False)
|
||||
stream.store_message({"type": "session.closed", "usage": {"audio_duration_ms": 4000}})
|
||||
await stream.log_messages()
|
||||
worker.ensure_initialized_and_enqueue.assert_not_called()
|
||||
logger.dispatch_success_handlers.assert_not_called()
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.chatgpt.codex import build_sideband_request, parse_call_response
|
||||
from litellm.llms.chatgpt.codex import CodexRealtimeCall, build_sideband_request, parse_call_response
|
||||
|
||||
|
||||
@pytest.mark.parametrize("location", ["", "/v1/realtime/calls/foreign-id"])
|
||||
|
|
@ -12,14 +12,25 @@ def test_signaling_rejects_invalid_upstream_call_id(location):
|
|||
parse_call_response(response, "voice", "owner", 1000)
|
||||
|
||||
|
||||
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",
|
||||
"extra_headers": {"x-gateway-route": "voice"}}})
|
||||
@pytest.mark.parametrize("extra_query", [None, {"gateway_token": "opaque +/& value"}])
|
||||
def test_signaling_preserves_selected_model_for_sideband(extra_query):
|
||||
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",
|
||||
"extra_headers": {"x-gateway-route": "voice"},
|
||||
**({"extra_query": extra_query} if extra_query is not None else {}),
|
||||
}
|
||||
},
|
||||
)
|
||||
call = parse_call_response(response, "voice", "owner", 1000)
|
||||
request = build_sideband_request(call)
|
||||
request = build_sideband_request(CodexRealtimeCall.model_validate_json(call.model_dump_json(exclude_none=True)))
|
||||
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"}
|
||||
assert request["extra_query"] == extra_query
|
||||
|
|
|
|||
|
|
@ -67,6 +67,7 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers,
|
|||
"model": "chatgpt/gpt-live-1-codex",
|
||||
"api_base": "https://voice.example/backend-api/codex",
|
||||
"extra_headers": {"x-gateway-route": "configured"},
|
||||
"extra_query": {"gateway_token": "configured", "intent": "pinned-intent"},
|
||||
},
|
||||
}
|
||||
],
|
||||
|
|
@ -74,8 +75,17 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers,
|
|||
)
|
||||
offer = CodexRealtimeOffer(sdp="v=0\r\n", session={"model": "voice-gateway"})
|
||||
try:
|
||||
response = await router.arealtime_calls(**build_call_request(offer, {}, inbound_headers), client=client)
|
||||
response = await router.arealtime_calls(
|
||||
**build_call_request(offer, {"intent": "quicksilver", "architecture": "avas"}, inbound_headers),
|
||||
client=client,
|
||||
)
|
||||
assert requests[0].headers.get("x-gateway-route") == "configured"
|
||||
assert dict(requests[0].url.params) == {
|
||||
"gateway_token": "configured",
|
||||
"intent": "pinned-intent",
|
||||
"architecture": "avas",
|
||||
}
|
||||
assert response.extensions["chatgpt_realtime"]["extra_query"] == dict(requests[0].url.params)
|
||||
assert response.extensions["chatgpt_realtime"]["extra_headers"]["x-gateway-route"] == "configured"
|
||||
for name, value in inbound_headers.items():
|
||||
assert requests[0].headers[name] == value
|
||||
|
|
@ -137,6 +147,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap
|
|||
sdp_body=b"v=0\r\n",
|
||||
session={"model": "chatgpt/gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}},
|
||||
extra_query={"intent": "quicksilver", "architecture": "avas"},
|
||||
chatgpt_realtime_client_query={"intent": "untrusted-override", "architecture": "avas", "untrusted": "bad"},
|
||||
extra_headers={
|
||||
"openai-alpha": "quicksilver=v2",
|
||||
"x-gateway-route": "voice",
|
||||
|
|
@ -146,9 +157,13 @@ 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 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 response.extensions["chatgpt_realtime"]["extra_query"] == {"intent": "quicksilver", "architecture": "avas"}
|
||||
assert requests[0].url.path == "/backend-api/codex/realtime/calls"
|
||||
assert requests[0].url.params["architecture"] == "avas"
|
||||
assert requests[0].headers["authorization"] == "Bearer test-token-" + "default"
|
||||
|
|
@ -244,3 +259,52 @@ def test_realtime_routes_use_configured_gateway(monkeypatch, env_name, api_base,
|
|||
assert handler._construct_url(handler.get_api_base(api_base), {"model": "gpt-realtime-1.5"}) == (
|
||||
expected.replace("https://", "wss://") + "/realtime?model=gpt-realtime-1.5"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model,endpoint", [("gpt-live-1-codex", "live"), ("gpt-realtime-1.5", "realtime")])
|
||||
def test_sideband_restores_gateway_query_without_overriding_call(model, endpoint, chatgpt_tokens):
|
||||
handler = ChatGPTRealtime(
|
||||
GenericLiteLLMParams(
|
||||
chatgpt_realtime_call_id="rtc_selected",
|
||||
extra_query={"gateway_token": "opaque +/& value", "model": "other", "call_id": "rtc_other"},
|
||||
),
|
||||
{},
|
||||
)
|
||||
url = httpx.URL(handler._construct_url("https://gateway.example/v1", {"model": model}))
|
||||
assert url.params["gateway_token"] == "opaque +/& value"
|
||||
assert "model" not in url.params
|
||||
if endpoint == "live":
|
||||
assert url.path == "/v1/live/rtc_selected"
|
||||
assert "call_id" not in url.params
|
||||
else:
|
||||
assert url.path == "/v1/realtime"
|
||||
assert url.params["call_id"] == "rtc_selected"
|
||||
|
||||
|
||||
def test_client_cannot_forge_supervised_call_accounting(chatgpt_tokens):
|
||||
from litellm.llms.chatgpt.realtime import CallAccounting, accounts_for_call_usage
|
||||
|
||||
assert accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting={"supervised": True}))
|
||||
assert accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting="supervised"))
|
||||
assert not accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting=CallAccounting.SUPERVISED))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", ["gpt-live-1-codex", "gpt-realtime-1.5"])
|
||||
async def test_supervisor_connection_preserves_call_routing(model, chatgpt_tokens):
|
||||
handler = ChatGPTRealtime(
|
||||
GenericLiteLLMParams(
|
||||
chatgpt_token_dir=chatgpt_tokens,
|
||||
chatgpt_realtime_call_id="rtc_owner",
|
||||
extra_query={"gateway_token": "a+b&c"},
|
||||
),
|
||||
{"openai-alpha": "quicksilver=v2"},
|
||||
{"x-gateway-token": "configured"},
|
||||
)
|
||||
connection = AsyncMock()
|
||||
with patch("websockets.connect", AsyncMock(return_value=connection)) as connect:
|
||||
assert await handler.open_call_connection(model, "https://gateway.example/v1") is connection
|
||||
url = httpx.URL(connect.call_args.args[0])
|
||||
assert url.params["gateway_token"] == "a+b&c"
|
||||
assert connect.call_args.kwargs["additional_headers"]["x-gateway-token"] == "configured"
|
||||
assert url.path.endswith("/rtc_owner") if model == "gpt-live-1-codex" else url.params["call_id"] == "rtc_owner"
|
||||
|
|
|
|||
|
|
@ -174,7 +174,9 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id
|
|||
@pytest.mark.parametrize("multipart", [False, True])
|
||||
@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom"])
|
||||
@pytest.mark.parametrize("signaling_credential", ["authorization", "api-key", "x-litellm-api-key", "mixed"])
|
||||
async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential, signaling_credential):
|
||||
async def test_offer_exchange_wraps_call_and_filters_client_headers(
|
||||
monkeypatch, multipart, credential, signaling_credential
|
||||
):
|
||||
import json
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
|
@ -188,11 +190,15 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime")
|
||||
session = {"model": "voice-alias", "audio": {"output": {"voice": "sol"}}}
|
||||
if multipart:
|
||||
body_request = httpx.Request("POST", "http://test/v1/realtime/calls", files={
|
||||
"sdp": (None, "v=0\r\n"), "session": (None, json.dumps(session))
|
||||
})
|
||||
body_request = httpx.Request(
|
||||
"POST",
|
||||
"http://test/v1/realtime/calls",
|
||||
files={"sdp": (None, "v=0\r\n"), "session": (None, json.dumps(session))},
|
||||
)
|
||||
else:
|
||||
body_request = httpx.Request("POST", "http://test/v1/realtime/calls", json={"sdp": "v=0\r\n", "session": session})
|
||||
body_request = httpx.Request(
|
||||
"POST", "http://test/v1/realtime/calls", json={"sdp": "v=0\r\n", "session": session}
|
||||
)
|
||||
body = body_request.read()
|
||||
|
||||
async def receive():
|
||||
|
|
@ -203,16 +209,30 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
if signaling_credential == "mixed"
|
||||
else [(signaling_credential.encode(), b"Bearer owner" if signaling_credential == "authorization" else b"owner")]
|
||||
)
|
||||
request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls",
|
||||
"scheme": "http", "server": ("localhost", 80),
|
||||
"query_string": b"intent=quicksilver&architecture=avas&untrusted=bad",
|
||||
"headers": [(b"content-type", body_request.headers["content-type"].encode()),
|
||||
*signaling_headers, *([(b"x-proxy-key", b"Bearer owner")] if credential == "custom" else []), (b"openai-alpha", b"quicksilver=v2"),
|
||||
(b"x-untrusted", b"bad")]}, receive)
|
||||
request = Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/realtime/calls",
|
||||
"scheme": "http",
|
||||
"server": ("localhost", 80),
|
||||
"query_string": b"intent=quicksilver&architecture=avas&untrusted=bad",
|
||||
"headers": [
|
||||
(b"content-type", body_request.headers["content-type"].encode()),
|
||||
*signaling_headers,
|
||||
*([(b"x-proxy-key", b"Bearer owner")] if credential == "custom" else []),
|
||||
(b"openai-alpha", b"quicksilver=v2"),
|
||||
(b"x-untrusted", b"bad"),
|
||||
],
|
||||
},
|
||||
receive,
|
||||
)
|
||||
auth = UserAPIKeyAuth()
|
||||
authorize = AsyncMock()
|
||||
monkeypatch.setattr(proxy_server, "master_key", "owner")
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential == "custom" else {})
|
||||
monkeypatch.setattr(
|
||||
proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential == "custom" else {}
|
||||
)
|
||||
monkeypatch.setattr(codex, "can_key_call_resolved_model", authorize)
|
||||
|
||||
class Processor:
|
||||
|
|
@ -225,7 +245,16 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
assert self.data["model"] == "voice-alias"
|
||||
assert self.data["guardrails"] == ["query-guardrail"]
|
||||
assert await kwargs["request"].json() == {"model": "voice-alias"}
|
||||
return {**self.data, "extra_headers": {"X-Hook-Required": "policy-value", "x-gateway-token": "untrusted-override", "Authorization": "Bearer untrusted"}, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None
|
||||
return {
|
||||
**self.data,
|
||||
"extra_headers": {
|
||||
"X-Hook-Required": "policy-value",
|
||||
"x-gateway-token": "untrusted-override",
|
||||
"Authorization": "Bearer untrusted",
|
||||
},
|
||||
"extra_query": {"gateway_token": "untrusted-override"},
|
||||
"metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"},
|
||||
}, None
|
||||
return self.data, None
|
||||
|
||||
monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor)
|
||||
|
|
@ -236,14 +265,28 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
assert data["session"] == session
|
||||
assert data["chatgpt_realtime_client_headers"] == {"openai-alpha": "quicksilver=v2"}
|
||||
assert "extra_headers" not in data
|
||||
assert data["extra_query"] == {"intent": "quicksilver", "architecture": "avas"}
|
||||
assert data["chatgpt_realtime_client_query"] == {"intent": "quicksilver", "architecture": "avas"}
|
||||
|
||||
async def respond():
|
||||
return httpx.Response(201, content=b"v=0\r\nanswer", headers={"Location": "/v1/realtime/calls/rtc_private"},
|
||||
extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", "extra_headers": {"X-Gateway-Token": "pinned-value"}}})
|
||||
return httpx.Response(
|
||||
201,
|
||||
content=b"v=0\r\nanswer",
|
||||
headers={"Location": "/v1/realtime/calls/rtc_private"},
|
||||
extensions={
|
||||
"chatgpt_realtime": {
|
||||
"model": "gpt-live-1-codex",
|
||||
"api_base": "https://voice.example/codex",
|
||||
"extra_headers": {"X-Gateway-Token": "pinned-value"},
|
||||
"extra_query": {"gateway_token": "pinned-query-value"},
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
return respond()
|
||||
|
||||
monkeypatch.setattr(proxy_server, "route_request", route)
|
||||
supervise = AsyncMock()
|
||||
monkeypatch.setattr(codex, "supervise_codex_call", supervise)
|
||||
response = await codex.create_codex_realtime_call(request)
|
||||
assert response.status_code == 201
|
||||
assert response.body == b"v=0\r\nanswer"
|
||||
|
|
@ -252,7 +295,11 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
assert call.call_id == "rtc_private"
|
||||
assert call.alias == "voice-alias"
|
||||
assert call.model == "gpt-live-1-codex"
|
||||
assert call.usage_supervised
|
||||
supervise.assert_awaited_once()
|
||||
assert "rtc_private" not in token
|
||||
assert "pinned-query-value" not in token
|
||||
assert call.extra_query == {"gateway_token": "pinned-query-value"}
|
||||
assert time.time() < call.expires_at < time.time() + 3601
|
||||
authorize.assert_awaited_once()
|
||||
|
||||
|
|
@ -271,16 +318,28 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
"custom": [(b"x-proxy-key", b"Bearer owner")],
|
||||
"subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")],
|
||||
}
|
||||
websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque",
|
||||
"query_string": b"guardrails=query-guardrail", "headers": credential_headers[credential]}, receive_ws, send)
|
||||
websocket = WebSocket(
|
||||
{
|
||||
"type": "websocket",
|
||||
"path": "/v1/live/opaque",
|
||||
"query_string": b"guardrails=query-guardrail",
|
||||
"headers": credential_headers[credential],
|
||||
},
|
||||
receive_ws,
|
||||
send,
|
||||
)
|
||||
forward = AsyncMock()
|
||||
monkeypatch.setattr(litellm, "_arealtime", forward)
|
||||
await codex.codex_realtime_sideband(websocket, token, auth)
|
||||
assert sent[0]["type"] == "websocket.accept"
|
||||
if credential == "subprotocol":
|
||||
assert sent[0]["subprotocol"] == "realtime"
|
||||
assert forward.await_args.kwargs["extra_headers"] == {"x-hook-required": "policy-value", "x-gateway-token": "pinned-value"}
|
||||
assert forward.await_args.kwargs["extra_headers"] == {
|
||||
"x-hook-required": "policy-value",
|
||||
"x-gateway-token": "pinned-value",
|
||||
}
|
||||
assert forward.await_args.kwargs["metadata"] == {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}
|
||||
assert forward.await_args.kwargs["extra_query"] == {"gateway_token": "pinned-query-value"}
|
||||
assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private"
|
||||
assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-codex"
|
||||
assert forward.await_args.kwargs["api_base"] == "https://voice.example/codex"
|
||||
|
|
@ -344,3 +403,158 @@ async def test_sideband_pre_call_block_prevents_upstream_connection(monkeypatch)
|
|||
await codex.codex_realtime_sideband(websocket, token, UserAPIKeyAuth())
|
||||
forward.assert_not_called()
|
||||
assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Realtime pre-call rejected"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("observer_fails", [False, True])
|
||||
async def test_signaling_transfers_reservation_only_to_ready_observer(monkeypatch, observer_fails):
|
||||
import json
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-reservation-transfer")
|
||||
reservation = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []}
|
||||
auth = UserAPIKeyAuth(budget_reservation=reservation)
|
||||
monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth))
|
||||
monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock())
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
process = AsyncMock(return_value=({}, None))
|
||||
monkeypatch.setattr(codex, "process_codex_request", process)
|
||||
|
||||
async def response():
|
||||
return httpx.Response(
|
||||
201,
|
||||
text="v=0\r\n",
|
||||
headers={"Location": "/v1/realtime/calls/rtc_ready"},
|
||||
extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}},
|
||||
)
|
||||
|
||||
async def route(**kwargs):
|
||||
return response()
|
||||
|
||||
monkeypatch.setattr(proxy_server, "route_request", route)
|
||||
|
||||
async def supervise(request, call, owner):
|
||||
assert owner is auth
|
||||
assert not owner.budget_reservation["finalized"]
|
||||
assert call.usage_supervised
|
||||
if observer_fails:
|
||||
await codex.release_or_invalidate_budget_reservation(budget_reservation=owner.budget_reservation)
|
||||
raise RuntimeError("Observer unavailable")
|
||||
|
||||
monkeypatch.setattr(codex, "supervise_codex_call", supervise)
|
||||
|
||||
async def receive():
|
||||
return {"type": "http.request", "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode()}
|
||||
|
||||
request = Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/realtime/calls",
|
||||
"query_string": b"",
|
||||
"headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer owner")],
|
||||
},
|
||||
receive,
|
||||
)
|
||||
if observer_fails:
|
||||
with pytest.raises(RuntimeError, match="Observer unavailable"):
|
||||
await codex.create_codex_realtime_call(request)
|
||||
else:
|
||||
assert (await codex.create_codex_realtime_call(request)).status_code == 201
|
||||
assert process.await_args.args[2].budget_reservation is None
|
||||
assert auth.budget_reservation["finalized"] is observer_fails
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_supervisor_policy_failure_hangs_up_before_releasing(monkeypatch):
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
call = CodexRealtimeCall(
|
||||
call_id="rtc_open", model="gpt-live-1-codex", alias="voice", owner="owner", expires_at=time.time() + 60
|
||||
)
|
||||
auth = UserAPIKeyAuth(budget_reservation={"reserved_cost": 0.5, "finalized": False, "entries": []})
|
||||
monkeypatch.setattr(codex, "process_codex_request", AsyncMock(side_effect=HTTPException(403, "Policy rejected")))
|
||||
closed = []
|
||||
|
||||
class Handler:
|
||||
def __init__(self, *args):
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def get_api_base(base):
|
||||
return "https://gateway.test/v1"
|
||||
|
||||
async def hangup_call(self, base):
|
||||
assert not auth.budget_reservation["finalized"]
|
||||
closed.append(base)
|
||||
|
||||
monkeypatch.setattr(codex, "ChatGPTRealtime", Handler)
|
||||
request = Request({"type": "http", "headers": [], "method": "POST", "path": "/v1/realtime/calls"})
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await codex.supervise_codex_call(request, call, auth)
|
||||
assert error.value.status_code == 403
|
||||
assert closed == ["https://gateway.test/v1"]
|
||||
assert auth.budget_reservation["finalized"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("hangup_fails", [False, True])
|
||||
async def test_supervisor_constructor_failure_closes_effective_connection(monkeypatch, hangup_fails, caplog):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
call = CodexRealtimeCall(
|
||||
call_id="rtc_open", model="gpt-live-1-codex", alias="voice", owner="owner", expires_at=time.time() + 60
|
||||
)
|
||||
auth = UserAPIKeyAuth(budget_reservation={"reserved_cost": 0.5, "finalized": False, "entries": []})
|
||||
logger = MagicMock()
|
||||
logger.litellm_params = {}
|
||||
connection = AsyncMock()
|
||||
handlers = []
|
||||
invalidate = AsyncMock()
|
||||
release = AsyncMock()
|
||||
monkeypatch.setattr(codex, "invalidate_budget_reservation_counters", invalidate, raising=False)
|
||||
monkeypatch.setattr(codex, "release_or_invalidate_budget_reservation", release)
|
||||
monkeypatch.setattr(
|
||||
codex, "process_codex_request", AsyncMock(return_value=({"extra_headers": {"x-hook": "effective"}}, logger))
|
||||
)
|
||||
|
||||
class Handler:
|
||||
def __init__(self, params, headers, extra_headers):
|
||||
self.headers = extra_headers
|
||||
handlers.append(self)
|
||||
|
||||
@staticmethod
|
||||
def get_api_base(base):
|
||||
return "https://gateway.test/v1"
|
||||
|
||||
async def open_call_connection(self, model, base):
|
||||
return connection
|
||||
|
||||
async def hangup_call(self, base):
|
||||
assert self.headers["x-hook"] == "effective"
|
||||
if hangup_fails:
|
||||
raise RuntimeError("private-cleanup-credential")
|
||||
|
||||
monkeypatch.setattr(codex, "ChatGPTRealtime", Handler)
|
||||
monkeypatch.setattr(codex, "RealTimeStreaming", MagicMock(side_effect=ValueError("original constructor failure")))
|
||||
request = Request({"type": "http", "headers": [], "method": "POST", "path": "/v1/realtime/calls"})
|
||||
with pytest.raises(ValueError, match="original constructor failure"):
|
||||
await codex.supervise_codex_call(request, call, auth)
|
||||
connection.close.assert_awaited_once()
|
||||
assert len(handlers) == 1
|
||||
if hangup_fails:
|
||||
invalidate.assert_awaited_once_with(budget_reservation=auth.budget_reservation)
|
||||
release.assert_not_awaited()
|
||||
else:
|
||||
release.assert_awaited_once_with(budget_reservation=auth.budget_reservation)
|
||||
invalidate.assert_not_awaited()
|
||||
assert "private-cleanup-credential" not in caplog.text
|
||||
|
|
|
|||
|
|
@ -0,0 +1,301 @@
|
|||
import asyncio
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.realtime_endpoints.call_supervision import CallSupervisor, CallSupervisors
|
||||
|
||||
|
||||
class Socket:
|
||||
def __init__(self):
|
||||
self.messages = asyncio.Queue()
|
||||
self.closed = False
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
message = await self.messages.get()
|
||||
if message is None:
|
||||
raise StopAsyncIteration
|
||||
if isinstance(message, Exception):
|
||||
raise message
|
||||
return json.dumps(message)
|
||||
|
||||
async def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
class Sink:
|
||||
def __init__(self, logger):
|
||||
self.logger = logger
|
||||
self.events = []
|
||||
self.logs = 0
|
||||
|
||||
def store_message(self, message):
|
||||
self.events.append(json.loads(message))
|
||||
|
||||
async def log_messages(self, *, wait_for_dispatch=False):
|
||||
assert wait_for_dispatch
|
||||
self.logs += 1
|
||||
self.logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
|
||||
|
||||
|
||||
def fixture(*, ready_timeout=1, lifetime=1):
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
sink = Sink(logger)
|
||||
|
||||
async def hangup():
|
||||
assert not socket.closed
|
||||
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
|
||||
|
||||
close_call = AsyncMock(side_effect=hangup)
|
||||
supervisor = CallSupervisor(
|
||||
socket,
|
||||
sink,
|
||||
logger,
|
||||
UserAPIKeyAuth(),
|
||||
close_call,
|
||||
ready_timeout=ready_timeout,
|
||||
lifetime=lifetime,
|
||||
drain_timeout=0.05,
|
||||
)
|
||||
return socket, sink, close_call, supervisor
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_observer_logs_webrtc_usage_without_client_sideband():
|
||||
socket, sink, close_call, supervisor = fixture()
|
||||
await socket.messages.put({"type": "session.started"})
|
||||
await supervisor.start()
|
||||
await socket.messages.put({"type": "response.done", "response": {"usage": {"total_tokens": 15}}})
|
||||
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 19}})
|
||||
await supervisor.wait()
|
||||
await supervisor.close()
|
||||
assert sink.logs == 1
|
||||
assert sink.events[-1]["usage"]["total_tokens"] == 19
|
||||
assert sink.events[1]["response"]["usage"]["total_tokens"] == 15
|
||||
assert socket.closed
|
||||
close_call.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_early_upstream_eof_rejects_start():
|
||||
socket, sink, close_call, supervisor = fixture()
|
||||
await socket.messages.put(None)
|
||||
with pytest.raises(RuntimeError, match="ended before"):
|
||||
await supervisor.start()
|
||||
assert socket.closed
|
||||
assert sink.logs == 1
|
||||
close_call.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancelled_start_hangs_up_and_drains_terminal_usage():
|
||||
socket, sink, close_call, supervisor = fixture()
|
||||
started = asyncio.create_task(supervisor.start())
|
||||
await asyncio.sleep(0)
|
||||
started.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await started
|
||||
close_call.assert_awaited_once()
|
||||
assert socket.closed
|
||||
assert sink.logs == 1
|
||||
assert sink.events[-1]["usage"]["total_tokens"] == 42
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_shutdown_drains_all_calls():
|
||||
registry = CallSupervisors()
|
||||
socket, sink, close_call, supervisor = fixture()
|
||||
await socket.messages.put({"type": "session.created"})
|
||||
await registry.start(supervisor)
|
||||
await registry.shutdown()
|
||||
await registry.shutdown()
|
||||
close_call.assert_awaited_once()
|
||||
assert socket.closed
|
||||
assert sink.logs == 1
|
||||
assert sink.events[-1]["usage"]["total_tokens"] == 42
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ready_timeout_hangs_up_before_returning_error():
|
||||
socket, sink, close_call, supervisor = fixture(ready_timeout=0.01)
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
await supervisor.start()
|
||||
close_call.assert_awaited_once()
|
||||
assert socket.closed
|
||||
assert sink.logs == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lifetime_limit_closes_call_and_collects_final_usage():
|
||||
socket, sink, close_call, supervisor = fixture(lifetime=0.01)
|
||||
await socket.messages.put({"type": "session.started"})
|
||||
await supervisor.start()
|
||||
await supervisor.wait()
|
||||
close_call.assert_awaited_once()
|
||||
assert socket.closed
|
||||
assert sink.events[-1]["usage"]["total_tokens"] == 42
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_socket_eof_after_ready_still_hangs_up_provider_call(monkeypatch):
|
||||
from litellm.proxy.realtime_endpoints import call_supervision
|
||||
|
||||
invalidate = AsyncMock()
|
||||
release = AsyncMock()
|
||||
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
|
||||
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
|
||||
socket, sink, close_call, supervisor = fixture()
|
||||
await socket.messages.put({"type": "session.started"})
|
||||
await supervisor.start()
|
||||
await socket.messages.put(None)
|
||||
await supervisor.wait()
|
||||
close_call.assert_awaited_once()
|
||||
assert socket.closed
|
||||
assert sink.logs == 1
|
||||
assert sink.logger.model_call_details["realtime_usage_incomplete"] is True
|
||||
invalidate.assert_awaited_once()
|
||||
release.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_observer_error_rejects_start(caplog):
|
||||
socket, sink, close_call, supervisor = fixture()
|
||||
await socket.messages.put(RuntimeError("private-provider-credential"))
|
||||
with pytest.raises(RuntimeError, match="ended before"):
|
||||
await supervisor.start()
|
||||
assert socket.closed
|
||||
close_call.assert_awaited_once()
|
||||
assert "private-provider-credential" not in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_logging_releases_reservation(monkeypatch):
|
||||
from litellm.proxy.realtime_endpoints import call_supervision
|
||||
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
sink = MagicMock()
|
||||
sink.log_messages = AsyncMock(side_effect=RuntimeError("logging unavailable"))
|
||||
release = AsyncMock()
|
||||
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
|
||||
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock())
|
||||
await socket.messages.put({"type": "session.started"})
|
||||
await supervisor.start()
|
||||
await socket.messages.put({"type": "session.closed"})
|
||||
with pytest.raises(RuntimeError, match="logging unavailable"):
|
||||
await supervisor.wait()
|
||||
release.assert_awaited_once_with(budget_reservation=None)
|
||||
assert socket.closed
|
||||
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_waits_for_usage_dispatch_completion():
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
dispatch_started = asyncio.Event()
|
||||
dispatch_complete = asyncio.Event()
|
||||
dispatch_finished = asyncio.Event()
|
||||
|
||||
async def log_messages(*, wait_for_dispatch=False):
|
||||
assert wait_for_dispatch
|
||||
dispatch_started.set()
|
||||
await dispatch_complete.wait()
|
||||
logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
|
||||
dispatch_finished.set()
|
||||
|
||||
async def hangup():
|
||||
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
|
||||
|
||||
sink = MagicMock()
|
||||
sink.log_messages = AsyncMock(side_effect=log_messages)
|
||||
registry = CallSupervisors()
|
||||
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), hangup)
|
||||
await socket.messages.put({"type": "session.created"})
|
||||
await registry.start(supervisor)
|
||||
shutdown = asyncio.create_task(registry.shutdown())
|
||||
try:
|
||||
await asyncio.wait_for(dispatch_started.wait(), timeout=1)
|
||||
assert not shutdown.done()
|
||||
assert not dispatch_finished.is_set()
|
||||
finally:
|
||||
dispatch_complete.set()
|
||||
await asyncio.wait_for(shutdown, timeout=1)
|
||||
assert dispatch_finished.is_set()
|
||||
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("terminal_usage_required", [True, False])
|
||||
async def test_confirmed_hangup_without_terminal_usage_matches_protocol(terminal_usage_required):
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
sink = Sink(logger)
|
||||
close_call = AsyncMock()
|
||||
supervisor = CallSupervisor(
|
||||
socket,
|
||||
sink,
|
||||
logger,
|
||||
UserAPIKeyAuth(),
|
||||
close_call,
|
||||
terminal_usage_required=terminal_usage_required,
|
||||
drain_timeout=0.01,
|
||||
)
|
||||
await socket.messages.put({"type": "session.created"})
|
||||
await supervisor.start()
|
||||
await socket.messages.put({"type": "response.done", "response": {"usage": {"total_tokens": 17}}})
|
||||
await supervisor.close()
|
||||
assert bool(logger.model_call_details.get("realtime_usage_incomplete")) == terminal_usage_required
|
||||
assert sink.events[-1]["response"]["usage"]["total_tokens"] == 17
|
||||
assert sink.logs == 1
|
||||
assert socket.closed
|
||||
close_call.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("closure", ["eof", "normal_close", "error"])
|
||||
@pytest.mark.parametrize("hangup_succeeds", [True, False])
|
||||
async def test_ga_observer_disconnect_requires_confirmed_hangup(closure, hangup_succeeds):
|
||||
from websockets.exceptions import ConnectionClosedOK
|
||||
from websockets.frames import Close
|
||||
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
sink = Sink(logger)
|
||||
close_call = AsyncMock(side_effect=None if hangup_succeeds else RuntimeError("unconfirmed hangup"))
|
||||
supervisor = CallSupervisor(
|
||||
socket,
|
||||
sink,
|
||||
logger,
|
||||
UserAPIKeyAuth(),
|
||||
close_call,
|
||||
terminal_usage_required=False,
|
||||
drain_timeout=0.01,
|
||||
)
|
||||
await socket.messages.put({"type": "session.created"})
|
||||
await supervisor.start()
|
||||
await socket.messages.put(
|
||||
None
|
||||
if closure == "eof"
|
||||
else ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True)
|
||||
if closure == "normal_close"
|
||||
else RuntimeError("observer failed")
|
||||
)
|
||||
await supervisor.wait()
|
||||
close_call.assert_awaited_once()
|
||||
assert bool(logger.model_call_details.get("realtime_usage_incomplete")) == (not hangup_succeeds)
|
||||
assert sink.logs == 1
|
||||
assert socket.closed
|
||||
|
|
@ -4751,3 +4751,71 @@ def test_collect_and_combine_realtime_usage_stores_partitioned_text_tokens() ->
|
|||
assert combined.completion_tokens_details.reasoning_tokens == 95
|
||||
assert combined.completion_tokens_details.text_tokens == 38
|
||||
assert combined.completion_tokens_details.audio_tokens == 0
|
||||
|
||||
|
||||
def _live_terminal_event(duration=4000):
|
||||
return {"type": "session.closed", "usage": {"audio_duration_ms": duration, "backend_model_usage": []}}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rate,expected", [(0.025, 0.1), (0, 0), (None, 0)])
|
||||
def test_live_terminal_duration_uses_configured_second_price(monkeypatch, rate, expected):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"live-priced-test",
|
||||
{"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": rate},
|
||||
)
|
||||
assert handle_realtime_stream_cost_calculation(
|
||||
[_live_terminal_event()], Usage(), "chatgpt", "live-priced-test"
|
||||
) == pytest.approx(expected)
|
||||
|
||||
|
||||
def test_live_terminal_duration_honors_deployment_override(monkeypatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"live-deployment-test",
|
||||
{"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025},
|
||||
)
|
||||
result = RealtimeAPITokenUsageProcessor.create_logging_realtime_object(Usage(), [_live_terminal_event()])
|
||||
assert completion_cost(
|
||||
completion_response=result,
|
||||
model="gpt-live-1",
|
||||
custom_llm_provider="chatgpt",
|
||||
call_type="_arealtime",
|
||||
custom_pricing=True,
|
||||
router_model_id="live-deployment-test",
|
||||
) == pytest.approx(0.1)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("duration", [-1, True, "4000", float("inf"), float("nan"), None])
|
||||
def test_live_terminal_invalid_duration_does_not_create_spend(monkeypatch, duration):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"live-priced-test",
|
||||
{"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025},
|
||||
)
|
||||
assert (
|
||||
handle_realtime_stream_cost_calculation(
|
||||
[_live_terminal_event(duration)], Usage(), "chatgpt", "live-priced-test"
|
||||
)
|
||||
== 0
|
||||
)
|
||||
|
||||
|
||||
def test_live_terminal_is_not_counted_twice(monkeypatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"live-priced-test",
|
||||
{"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025},
|
||||
)
|
||||
assert handle_realtime_stream_cost_calculation(
|
||||
[_live_terminal_event(), _live_terminal_event()], Usage(), "chatgpt", "live-priced-test"
|
||||
) == pytest.approx(0.1)
|
||||
assert (
|
||||
handle_realtime_stream_cost_calculation(
|
||||
[{"type": "response.done", "response": {"usage": {}}}, _live_terminal_event()],
|
||||
Usage(),
|
||||
"chatgpt",
|
||||
"live-priced-test",
|
||||
)
|
||||
== 0
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue