mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): preserve observer quotas and bound call cleanup
This commit is contained in:
parent
bac2ad4abb
commit
c360d187f5
11 changed files with 464 additions and 46 deletions
|
|
@ -3,8 +3,9 @@ from enum import Enum, auto
|
|||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from httpx import URL
|
||||
from httpx import URL, QueryParams
|
||||
from pydantic import TypeAdapter
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.llms.openai.realtime.handler import OpenAIRealtime
|
||||
|
|
@ -40,11 +41,14 @@ 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(
|
||||
configured: Final = TypeAdapter(Mapping[str, str | int | float | bool | None]).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}
|
||||
{
|
||||
**{key: value for key, value in inbound.items() if key in ("intent", "architecture")},
|
||||
**QueryParams(configured),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -111,8 +115,12 @@ class ChatGPTRealtime(OpenAIRealtime):
|
|||
|
||||
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
|
||||
try:
|
||||
await connection.send('{"type":"session.close"}')
|
||||
return
|
||||
except (ConnectionClosed, OSError):
|
||||
await self.hangup_call(api_base)
|
||||
return
|
||||
await self.hangup_call(api_base)
|
||||
|
||||
async def hangup_call(self, api_base: str) -> None:
|
||||
|
|
|
|||
|
|
@ -97,6 +97,10 @@ class ReconcileOutcome(NamedTuple):
|
|||
live_after: frozenset[str] | None
|
||||
|
||||
|
||||
class InternalRequestOrigin(enum.Enum):
|
||||
REALTIME_OBSERVER = enum.auto()
|
||||
|
||||
|
||||
class SupportedDBObjectType(str, enum.Enum):
|
||||
"""
|
||||
Supported database object types for fine-grained DB storage control.
|
||||
|
|
|
|||
|
|
@ -1832,6 +1832,8 @@ class ProxyBaseLLMRequestProcessing:
|
|||
user_api_base: str | None = None,
|
||||
model: str | None = None,
|
||||
llm_router: Router | None = None,
|
||||
*,
|
||||
internal_realtime_observer: bool = False,
|
||||
) -> tuple[dict, LiteLLMLoggingObj]:
|
||||
start_time: Final = datetime.now() # start before calling guardrail hooks
|
||||
|
||||
|
|
@ -1996,6 +1998,11 @@ class ProxyBaseLLMRequestProcessing:
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
data=self.data,
|
||||
call_type=route_type,
|
||||
**(
|
||||
MappingProxyType({"internal_realtime_observer": True})
|
||||
if internal_realtime_observer
|
||||
else MappingProxyType({})
|
||||
),
|
||||
)
|
||||
if route_type == "aget_responses":
|
||||
attach_post_call_pipelines_to_retrieval(
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.exceptions import RateLimitType
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
|
||||
from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth
|
||||
from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, InternalRequestOrigin, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
get_key_model_rpm_limit,
|
||||
get_key_model_tpm_limit,
|
||||
|
|
@ -489,6 +489,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
)
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time):
|
||||
releases_slot: Final = kwargs.get("internal_request_origin") is not InternalRequestOrigin.REALTIME_OBSERVER
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
get_model_group_from_litellm_kwargs,
|
||||
)
|
||||
|
|
@ -521,7 +522,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
# Setup values
|
||||
# ------------
|
||||
|
||||
if global_max_parallel_requests is not None:
|
||||
if releases_slot and global_max_parallel_requests is not None:
|
||||
# get value from cache
|
||||
_key: Final = "global_max_parallel_requests"
|
||||
# decrement
|
||||
|
|
@ -552,13 +553,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
key=request_count_api_key,
|
||||
litellm_parent_otel_span=litellm_parent_otel_span,
|
||||
) or {
|
||||
"current_requests": 1,
|
||||
"current_requests": int(releases_slot),
|
||||
"current_tpm": 0,
|
||||
"current_rpm": 0,
|
||||
}
|
||||
|
||||
new_val = {
|
||||
"current_requests": max(current["current_requests"] - 1, 0),
|
||||
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
|
||||
"current_tpm": current["current_tpm"] + total_tokens,
|
||||
"current_rpm": current["current_rpm"],
|
||||
}
|
||||
|
|
@ -593,13 +594,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
key=request_count_api_key,
|
||||
litellm_parent_otel_span=litellm_parent_otel_span,
|
||||
) or {
|
||||
"current_requests": 1,
|
||||
"current_requests": int(releases_slot),
|
||||
"current_tpm": 0,
|
||||
"current_rpm": 0,
|
||||
}
|
||||
|
||||
new_val = {
|
||||
"current_requests": max(current["current_requests"] - 1, 0),
|
||||
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
|
||||
"current_tpm": current["current_tpm"] + total_tokens,
|
||||
"current_rpm": current["current_rpm"],
|
||||
}
|
||||
|
|
@ -619,13 +620,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
key=request_count_api_key,
|
||||
litellm_parent_otel_span=litellm_parent_otel_span,
|
||||
) or {
|
||||
"current_requests": 1,
|
||||
"current_tpm": total_tokens,
|
||||
"current_rpm": 1,
|
||||
"current_requests": int(releases_slot),
|
||||
"current_tpm": total_tokens if releases_slot else 0,
|
||||
"current_rpm": int(releases_slot),
|
||||
}
|
||||
|
||||
new_val = {
|
||||
"current_requests": max(current["current_requests"] - 1, 0),
|
||||
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
|
||||
"current_tpm": current["current_tpm"] + total_tokens,
|
||||
"current_rpm": current["current_rpm"],
|
||||
}
|
||||
|
|
@ -645,13 +646,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
key=request_count_api_key,
|
||||
litellm_parent_otel_span=litellm_parent_otel_span,
|
||||
) or {
|
||||
"current_requests": 1,
|
||||
"current_tpm": total_tokens,
|
||||
"current_rpm": 1,
|
||||
"current_requests": int(releases_slot),
|
||||
"current_tpm": total_tokens if releases_slot else 0,
|
||||
"current_rpm": int(releases_slot),
|
||||
}
|
||||
|
||||
new_val = {
|
||||
"current_requests": max(current["current_requests"] - 1, 0),
|
||||
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
|
||||
"current_tpm": current["current_tpm"] + total_tokens,
|
||||
"current_rpm": current["current_rpm"],
|
||||
}
|
||||
|
|
@ -671,13 +672,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
key=request_count_api_key,
|
||||
litellm_parent_otel_span=litellm_parent_otel_span,
|
||||
) or {
|
||||
"current_requests": 1,
|
||||
"current_tpm": total_tokens,
|
||||
"current_rpm": 1,
|
||||
"current_requests": int(releases_slot),
|
||||
"current_tpm": total_tokens if releases_slot else 0,
|
||||
"current_rpm": int(releases_slot),
|
||||
}
|
||||
|
||||
new_val = {
|
||||
"current_requests": max(current["current_requests"] - 1, 0),
|
||||
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
|
||||
"current_tpm": current["current_tpm"] + total_tokens,
|
||||
"current_rpm": current["current_rpm"],
|
||||
}
|
||||
|
|
@ -694,6 +695,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
self.print_verbose(e)
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if kwargs.get("internal_request_origin") is InternalRequestOrigin.REALTIME_OBSERVER:
|
||||
return
|
||||
try:
|
||||
self.print_verbose("Inside Max Parallel Request Failure Hook")
|
||||
litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs)
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ from litellm.llms.chatgpt.realtime import (
|
|||
configured_realtime_headers,
|
||||
realtime_endpoint,
|
||||
)
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy._types import InternalRequestOrigin, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
get_api_key,
|
||||
|
|
@ -73,6 +73,7 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth:
|
|||
auth,
|
||||
call.alias,
|
||||
"_arealtime",
|
||||
internal_realtime_observer=True,
|
||||
)
|
||||
pinned: Final = { # mutable-ok: logging and provider parameter contract
|
||||
**processed,
|
||||
|
|
@ -116,6 +117,9 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth:
|
|||
async def close_call() -> None:
|
||||
await handler.close_call(connection, call.model, api_base)
|
||||
|
||||
async def force_close_call() -> None:
|
||||
await handler.hangup_call(api_base)
|
||||
|
||||
frontend: Final = WebSocket(
|
||||
{**request.scope, "type": "websocket"}, receive=receive, send=send
|
||||
) # mutable-ok: ASGI scope
|
||||
|
|
@ -126,6 +130,7 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth:
|
|||
logger,
|
||||
auth,
|
||||
close_call,
|
||||
force_close_call=force_close_call,
|
||||
terminal_usage_required=realtime_endpoint(call.model) == "live",
|
||||
)
|
||||
supervision_owned = True
|
||||
|
|
@ -191,6 +196,8 @@ async def process_codex_request(
|
|||
auth: UserAPIKeyAuth,
|
||||
model: str,
|
||||
route_type: Literal["arealtime_calls", "_arealtime"],
|
||||
*,
|
||||
internal_realtime_observer: bool = False,
|
||||
) -> tuple[dict[str, object], Logging]: # mutable-ok: common request processor returns enriched routing arguments
|
||||
from litellm.proxy import proxy_server as server
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
|
@ -211,7 +218,14 @@ async def process_codex_request(
|
|||
user_api_base=server.user_api_base,
|
||||
model=model,
|
||||
route_type=route_type,
|
||||
**(
|
||||
MappingProxyType({"internal_realtime_observer": True})
|
||||
if internal_realtime_observer
|
||||
else MappingProxyType({})
|
||||
),
|
||||
)
|
||||
if internal_realtime_observer:
|
||||
logging_obj.model_call_details["internal_request_origin"] = InternalRequestOrigin.REALTIME_OBSERVER
|
||||
return processed, logging_obj
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from pydantic import BaseModel
|
|||
from websockets.exceptions import ConnectionClosedOK
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE
|
||||
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
|
||||
|
|
@ -45,23 +46,28 @@ class CallSupervisor:
|
|||
lifetime: float = 3600,
|
||||
drain_timeout: float = 5,
|
||||
termination_timeout: float = 60,
|
||||
logging_timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE,
|
||||
terminal_usage_required: bool = True,
|
||||
force_close_call: Callable[[], Awaitable[None]] | None = None,
|
||||
) -> None:
|
||||
self._upstream = upstream
|
||||
self._stream = stream
|
||||
self._logging = logging_obj
|
||||
self._auth = auth
|
||||
self._close_call = close_call
|
||||
self._force_close_call = force_close_call
|
||||
self._ready_timeout = ready_timeout
|
||||
self._lifetime = lifetime
|
||||
self._drain_timeout = drain_timeout
|
||||
self._termination_timeout = termination_timeout
|
||||
self._logging_timeout = logging_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._accounting_complete = False
|
||||
self._task: asyncio.Task[None] | None = None
|
||||
|
||||
async def start(self) -> None:
|
||||
|
|
@ -70,7 +76,7 @@ class CallSupervisor:
|
|||
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():
|
||||
if not self._started or self._terminal or self._task.done():
|
||||
raise RuntimeError("Call observer ended before session became available")
|
||||
except BaseException:
|
||||
await self.close()
|
||||
|
|
@ -113,12 +119,21 @@ class CallSupervisor:
|
|||
finally:
|
||||
try:
|
||||
if not self._terminal:
|
||||
deadline: Final = asyncio.get_running_loop().time() + self._termination_timeout
|
||||
try:
|
||||
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")
|
||||
await self._drain(reader)
|
||||
await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time()))
|
||||
if self._terminal_usage_required and not self._terminal and self._force_close_call is not None:
|
||||
remaining: Final = max(0.0, deadline - asyncio.get_running_loop().time())
|
||||
try:
|
||||
await asyncio.wait_for(self._force_close_call(), timeout=remaining)
|
||||
self._close_confirmed = True
|
||||
except Exception: # noqa: BLE001 # provider exceptions can contain credentials
|
||||
verbose_proxy_logger.error("Realtime observer independent hangup failed")
|
||||
await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time()))
|
||||
finally:
|
||||
stopped.cancel()
|
||||
reader.cancel()
|
||||
|
|
@ -132,9 +147,16 @@ class CallSupervisor:
|
|||
)
|
||||
try:
|
||||
try:
|
||||
await self._stream.log_messages(wait_for_dispatch=True)
|
||||
await asyncio.wait_for(
|
||||
self._stream.log_messages(wait_for_dispatch=True), timeout=self._logging_timeout
|
||||
)
|
||||
self._accounting_complete = True
|
||||
except asyncio.TimeoutError:
|
||||
verbose_proxy_logger.error("Realtime observer timed out dispatching usage accounting")
|
||||
finally:
|
||||
if self._started and not self._usage_complete():
|
||||
if not self._accounting_complete:
|
||||
self._logging.model_call_details["realtime_accounting_incomplete"] = True
|
||||
if self._started and (not self._usage_complete() or not self._accounting_complete):
|
||||
await invalidate_budget_reservation_counters(
|
||||
budget_reservation=self._auth.budget_reservation
|
||||
)
|
||||
|
|
@ -145,9 +167,12 @@ class CallSupervisor:
|
|||
finally:
|
||||
self._ready.set()
|
||||
|
||||
async def _drain(self, reader: asyncio.Task[None]) -> None:
|
||||
async def _drain(self, reader: asyncio.Task[None], *, timeout: float | None = None) -> None:
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.shield(reader), timeout=self._drain_timeout)
|
||||
await asyncio.wait_for(
|
||||
asyncio.shield(reader),
|
||||
timeout=self._drain_timeout if timeout is None else min(self._drain_timeout, timeout),
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
if not self._usage_complete():
|
||||
verbose_proxy_logger.error("Realtime observer timed out draining terminal usage")
|
||||
|
|
|
|||
|
|
@ -2046,6 +2046,8 @@ class ProxyLogging:
|
|||
data: None,
|
||||
call_type: CallTypesLiteral,
|
||||
guardrails_only: bool = False,
|
||||
*,
|
||||
internal_realtime_observer: bool = False,
|
||||
) -> None:
|
||||
pass
|
||||
|
||||
|
|
@ -2056,6 +2058,8 @@ class ProxyLogging:
|
|||
data: dict,
|
||||
call_type: CallTypesLiteral,
|
||||
guardrails_only: bool = False,
|
||||
*,
|
||||
internal_realtime_observer: bool = False,
|
||||
) -> dict:
|
||||
pass
|
||||
|
||||
|
|
@ -2065,6 +2069,8 @@ class ProxyLogging:
|
|||
data: dict | None,
|
||||
call_type: CallTypesLiteral,
|
||||
guardrails_only: bool = False,
|
||||
*,
|
||||
internal_realtime_observer: bool = False,
|
||||
) -> dict | None:
|
||||
"""
|
||||
Allows users to modify/reject the incoming request to the proxy, without having to deal with parsing Request body.
|
||||
|
|
@ -2163,6 +2169,10 @@ class ProxyLogging:
|
|||
|
||||
deferred_route_exc: SensitiveDataRouteException | None = None
|
||||
for _callback in caps.resolved_callbacks:
|
||||
if internal_realtime_observer and isinstance(
|
||||
_callback, (_PROXY_MaxParallelRequestsHandler, _PROXY_MaxParallelRequestsHandler_v3)
|
||||
):
|
||||
continue
|
||||
start_time = time.time()
|
||||
try:
|
||||
if isinstance(_callback, CustomGuardrail) and data is not None:
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
from contextlib import nullcontext
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
|
|
@ -11,6 +12,51 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
|||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure", ["closed", "network"])
|
||||
@pytest.mark.parametrize(
|
||||
"hangup_status, expectation", [(200, nullcontext()), (503, pytest.raises(httpx.HTTPStatusError))]
|
||||
)
|
||||
async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, expectation, chatgpt_tokens):
|
||||
from websockets.exceptions import ConnectionClosedOK
|
||||
from websockets.frames import Close
|
||||
|
||||
handler = ChatGPTRealtime(
|
||||
GenericLiteLLMParams(
|
||||
chatgpt_realtime_call_id="rtc_live_closed",
|
||||
chatgpt_token_dir=chatgpt_tokens,
|
||||
extra_query={"gateway": "tenant"},
|
||||
),
|
||||
{},
|
||||
{"x-gateway-token": "test-only"},
|
||||
)
|
||||
connection = SimpleNamespace(
|
||||
send=AsyncMock(
|
||||
side_effect=(
|
||||
ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True)
|
||||
if failure == "closed"
|
||||
else OSError("socket unavailable")
|
||||
)
|
||||
)
|
||||
)
|
||||
requests = []
|
||||
|
||||
def respond(request):
|
||||
requests.append(request)
|
||||
return httpx.Response(hangup_status)
|
||||
|
||||
client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
|
||||
with patch("httpx.AsyncClient", return_value=client):
|
||||
with expectation:
|
||||
await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1")
|
||||
assert len(requests) == 1
|
||||
assert requests[0].method == "POST"
|
||||
assert str(requests[0].url) == "https://gateway.example/v1/realtime/calls/rtc_live_closed/hangup?gateway=tenant"
|
||||
assert requests[0].headers["x-gateway-token"] == "test-only"
|
||||
assert requests[0].headers["Authorization"] == "Bearer test-token-default"
|
||||
assert client.is_closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"])
|
||||
@pytest.mark.parametrize("source", ["default", "explicit", "CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"])
|
||||
|
|
@ -47,8 +93,17 @@ async def test_realtime_session_urls_honor_gateway(endpoint, source, chatgpt_tok
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("inbound_headers", [{}, {"openai-alpha": "quicksilver=v2"}])
|
||||
async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, chatgpt_tokens, monkeypatch):
|
||||
from litellm.llms.chatgpt.codex import CodexRealtimeOffer, build_call_request
|
||||
@pytest.mark.parametrize("model, endpoint", [("gpt-live-1-codex", "live"), ("gpt-realtime-1.5", "realtime")])
|
||||
async def test_routed_call_preserves_deployment_gateway_headers(
|
||||
inbound_headers, model, endpoint, chatgpt_tokens, monkeypatch
|
||||
):
|
||||
from litellm.llms.chatgpt.codex import (
|
||||
CodexRealtimeCall,
|
||||
CodexRealtimeOffer,
|
||||
build_call_request,
|
||||
build_sideband_request,
|
||||
parse_call_response,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("CHATGPT_TOKEN_DIR", chatgpt_tokens)
|
||||
requests = []
|
||||
|
|
@ -64,11 +119,22 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers,
|
|||
{
|
||||
"model_name": "voice-gateway",
|
||||
"litellm_params": {
|
||||
"model": "chatgpt/gpt-live-1-codex",
|
||||
"model": f"chatgpt/{model}",
|
||||
"api_base": "https://voice.example/backend-api/codex",
|
||||
"extra_headers": {"x-gateway-route": "configured"},
|
||||
"extra_query": {"gateway_token": "configured", "intent": "pinned-intent"},
|
||||
"extra_query": {
|
||||
"gateway_token": "configured",
|
||||
"intent": "pinned-intent",
|
||||
"count": 7,
|
||||
"fraction": 1.5,
|
||||
"enabled": True,
|
||||
"disabled": False,
|
||||
"blank": None,
|
||||
"model": "other-model",
|
||||
"call_id": "rtc_wrong",
|
||||
},
|
||||
},
|
||||
"model_info": {"id": "selected-gateway-deployment"},
|
||||
}
|
||||
],
|
||||
num_retries=0,
|
||||
|
|
@ -84,11 +150,29 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers,
|
|||
"gateway_token": "configured",
|
||||
"intent": "pinned-intent",
|
||||
"architecture": "avas",
|
||||
"count": "7",
|
||||
"fraction": "1.5",
|
||||
"enabled": "true",
|
||||
"disabled": "false",
|
||||
"blank": "",
|
||||
"model": "other-model",
|
||||
"call_id": "rtc_wrong",
|
||||
}
|
||||
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
|
||||
call = parse_call_response(response, alias="voice-gateway", owner="test-owner", expires_at=1)
|
||||
restored = CodexRealtimeCall.model_validate_json(call.model_dump_json())
|
||||
assert restored.model_id == "selected-gateway-deployment"
|
||||
assert restored.model == model
|
||||
handler = ChatGPTRealtime(GenericLiteLLMParams.model_validate(build_sideband_request(restored)), {})
|
||||
sideband_url = httpx.URL(handler._construct_url(restored.api_base, {"model": restored.model}))
|
||||
assert {key: value for key, value in sideband_url.params.items() if key != "call_id"} == {
|
||||
key: value for key, value in requests[0].url.params.items() if key not in ("model", "call_id")
|
||||
}
|
||||
assert sideband_url.params.get("call_id") == ("rtc_test" if endpoint == "realtime" else None)
|
||||
assert sideband_url.path.endswith("/realtime" if endpoint == "realtime" else "/live/rtc_test")
|
||||
finally:
|
||||
await client.client.aclose()
|
||||
|
||||
|
|
|
|||
|
|
@ -13,7 +13,8 @@ from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_c
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route_type", ["arealtime_calls", "_arealtime"])
|
||||
async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type):
|
||||
@pytest.mark.parametrize("observer", [False, True])
|
||||
async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type, observer):
|
||||
from fastapi import Request
|
||||
from litellm import Router
|
||||
from litellm.proxy import proxy_server as server
|
||||
|
|
@ -21,7 +22,8 @@ async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type)
|
|||
from litellm.proxy.realtime_endpoints.call_sessions import process_codex_request
|
||||
|
||||
class PolicyHook:
|
||||
async def pre_call_hook(self, user_api_key_dict, data, call_type):
|
||||
async def pre_call_hook(self, user_api_key_dict, data, call_type, *, internal_realtime_observer=False):
|
||||
assert internal_realtime_observer is observer
|
||||
if "model-policy" in data.get("metadata", {}).get("guardrails", []):
|
||||
raise HTTPException(403, "Model policy rejected request")
|
||||
return data
|
||||
|
|
@ -34,7 +36,14 @@ async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type)
|
|||
monkeypatch.setattr(server, "proxy_logging_obj", PolicyHook())
|
||||
request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", "headers": [], "query_string": b"", "scheme": "http", "server": ("localhost", 80)})
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await process_codex_request(request, {"model": "voice-policy"}, UserAPIKeyAuth(), "voice-policy", route_type)
|
||||
await process_codex_request(
|
||||
request,
|
||||
{"model": "voice-policy"},
|
||||
UserAPIKeyAuth(),
|
||||
"voice-policy",
|
||||
route_type,
|
||||
internal_realtime_observer=observer,
|
||||
)
|
||||
assert error.value.status_code == 403
|
||||
assert error.value.detail == "Model policy rejected request"
|
||||
|
||||
|
|
|
|||
|
|
@ -30,6 +30,67 @@ class Socket:
|
|||
self.closed = True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("fallback", ["terminal", "no_terminal", "timeout"])
|
||||
async def test_live_unacknowledged_close_uses_bounded_independent_hangup(monkeypatch, fallback):
|
||||
from litellm.proxy.realtime_endpoints import call_supervision
|
||||
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
sink = Sink(logger)
|
||||
invalidate = AsyncMock()
|
||||
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
|
||||
|
||||
async def force_close():
|
||||
if fallback == "terminal":
|
||||
await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}})
|
||||
elif fallback == "timeout":
|
||||
await asyncio.Event().wait()
|
||||
|
||||
force = AsyncMock(side_effect=force_close)
|
||||
close = AsyncMock()
|
||||
supervisor = CallSupervisor(
|
||||
socket,
|
||||
sink,
|
||||
logger,
|
||||
UserAPIKeyAuth(),
|
||||
close,
|
||||
force_close_call=force,
|
||||
drain_timeout=0.01,
|
||||
termination_timeout=0.08,
|
||||
)
|
||||
await socket.messages.put({"type": "session.started"})
|
||||
await supervisor.start()
|
||||
await asyncio.wait_for(supervisor.close(), timeout=0.5)
|
||||
close.assert_awaited_once()
|
||||
force.assert_awaited_once()
|
||||
assert socket.closed
|
||||
if fallback == "terminal":
|
||||
invalidate.assert_not_awaited()
|
||||
assert not logger.model_call_details.get("realtime_usage_incomplete")
|
||||
else:
|
||||
invalidate.assert_awaited_once()
|
||||
assert logger.model_call_details["realtime_usage_incomplete"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_live_confirmed_terminal_does_not_force_hangup():
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
|
||||
async def close():
|
||||
await socket.messages.put({"type": "session.closed"})
|
||||
|
||||
force = AsyncMock()
|
||||
supervisor = CallSupervisor(socket, Sink(logger), logger, UserAPIKeyAuth(), close, force_close_call=force)
|
||||
await socket.messages.put({"type": "session.started"})
|
||||
await supervisor.start()
|
||||
await supervisor.close()
|
||||
force.assert_not_awaited()
|
||||
|
||||
|
||||
class Sink:
|
||||
def __init__(self, logger):
|
||||
self.logger = logger
|
||||
|
|
@ -178,7 +239,7 @@ async def test_observer_error_rejects_start(caplog):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_logging_releases_reservation(monkeypatch):
|
||||
async def test_failed_logging_invalidates_reservation_without_zeroing_spend(monkeypatch):
|
||||
from litellm.proxy.realtime_endpoints import call_supervision
|
||||
|
||||
socket = Socket()
|
||||
|
|
@ -187,18 +248,89 @@ async def test_failed_logging_releases_reservation(monkeypatch):
|
|||
sink = MagicMock()
|
||||
sink.log_messages = AsyncMock(side_effect=RuntimeError("logging unavailable"))
|
||||
release = AsyncMock()
|
||||
invalidate = AsyncMock()
|
||||
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
|
||||
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
|
||||
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)
|
||||
release.assert_not_awaited()
|
||||
invalidate.assert_awaited_once_with(budget_reservation=None)
|
||||
assert logger.model_call_details["realtime_accounting_incomplete"] is True
|
||||
assert socket.closed
|
||||
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_rejects_terminal_session_while_accounting_is_pending():
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
dispatch_started = asyncio.Event()
|
||||
allow_dispatch = asyncio.Event()
|
||||
|
||||
async def log_messages(*, wait_for_dispatch=False):
|
||||
dispatch_started.set()
|
||||
await allow_dispatch.wait()
|
||||
logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
|
||||
|
||||
sink = MagicMock()
|
||||
sink.log_messages = AsyncMock(side_effect=log_messages)
|
||||
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock())
|
||||
await socket.messages.put({"type": "session.created"})
|
||||
await socket.messages.put({"type": "session.closed"})
|
||||
startup = asyncio.create_task(supervisor.start())
|
||||
try:
|
||||
await asyncio.wait_for(dispatch_started.wait(), timeout=1)
|
||||
finally:
|
||||
allow_dispatch.set()
|
||||
with pytest.raises(RuntimeError, match="ended before"):
|
||||
await asyncio.wait_for(startup, timeout=1)
|
||||
assert socket.closed
|
||||
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_bounds_accounting_and_invalidates_partial_dispatch(monkeypatch):
|
||||
from litellm.proxy.realtime_endpoints import call_supervision
|
||||
|
||||
socket = Socket()
|
||||
logger = MagicMock(spec=Logging)
|
||||
logger.model_call_details = {}
|
||||
dispatch_cancelled = asyncio.Event()
|
||||
invalidate = AsyncMock()
|
||||
release = AsyncMock()
|
||||
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
|
||||
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
|
||||
|
||||
async def log_messages(*, wait_for_dispatch=False):
|
||||
try:
|
||||
await asyncio.Event().wait()
|
||||
finally:
|
||||
dispatch_cancelled.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)
|
||||
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), hangup, logging_timeout=0.01)
|
||||
registry = CallSupervisors()
|
||||
await socket.messages.put({"type": "session.created"})
|
||||
await registry.start(supervisor)
|
||||
await asyncio.wait_for(registry.shutdown(), timeout=1)
|
||||
assert dispatch_cancelled.is_set()
|
||||
assert socket.closed
|
||||
assert logger.model_call_details["realtime_accounting_incomplete"] is True
|
||||
assert not logger.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY)
|
||||
invalidate.assert_awaited_once_with(budget_reservation=None)
|
||||
release.assert_not_awaited()
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import datetime as real_datetime
|
||||
import smtplib
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -8,15 +9,10 @@ from litellm.caching.caching import DualCache
|
|||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import ProxyErrorTypes, UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.proxy.utils import ProxyLogging, get_custom_url, join_paths
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from litellm.proxy.utils import get_custom_url, join_paths
|
||||
|
||||
|
||||
def test_get_custom_url(monkeypatch):
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/litellm")
|
||||
custom_url = get_custom_url(request_base_url="http://0.0.0.0:4000", route="ui/")
|
||||
|
|
@ -2030,7 +2026,9 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp
|
|||
with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()):
|
||||
await proxy_logging_obj.post_call_failure_hook(
|
||||
request_data={"metadata": {}},
|
||||
original_exception=HTTPException(status_code=400, detail="Upstream passthrough request failed with status 400"),
|
||||
original_exception=HTTPException(
|
||||
status_code=400, detail="Upstream passthrough request failed with status 400"
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
traceback_str=upstream_traceback,
|
||||
)
|
||||
|
|
@ -2038,3 +2036,127 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp
|
|||
assert recorder.received_traceback is not None
|
||||
assert provider_key not in recorder.received_traceback
|
||||
assert "REDACTED" in recorder.received_traceback
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("limiter_version", [1, 3])
|
||||
@pytest.mark.parametrize("limit", ["rpm_limit", "max_parallel_requests"])
|
||||
async def test_internal_realtime_observer_preserves_quota_and_custom_hooks(monkeypatch, limiter_version, limit):
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_PROXY_MaxParallelRequestsHandler_v3,
|
||||
_request_stash,
|
||||
get_request_stash,
|
||||
)
|
||||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
|
||||
|
||||
observed = []
|
||||
|
||||
class Hook(CustomLogger):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
observed.append(call_type)
|
||||
return {**data, "extra_headers": {"x-hook": "required"}}
|
||||
|
||||
cache = DualCache()
|
||||
limiter_type = _PROXY_MaxParallelRequestsHandler if limiter_version == 1 else _PROXY_MaxParallelRequestsHandler_v3
|
||||
limiter = limiter_type(InternalUsageCache(dual_cache=cache))
|
||||
proxy = ProxyLogging(UserApiKeyCache())
|
||||
monkeypatch.setattr(litellm, "callbacks", [limiter, Hook()])
|
||||
token = _request_stash.set(None)
|
||||
try:
|
||||
auth = UserAPIKeyAuth(api_key="observer-quota-test", **{limit: 1})
|
||||
await proxy.pre_call_hook(
|
||||
auth, {"model": "voice", "litellm_call_id": "signaling", "metadata": {}}, "arealtime_calls"
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
initial_stash = get_request_stash()
|
||||
result = await proxy.pre_call_hook(
|
||||
auth,
|
||||
{"model": "voice", "litellm_call_id": "observer", "metadata": {}},
|
||||
"_arealtime",
|
||||
internal_realtime_observer=True,
|
||||
)
|
||||
assert result["extra_headers"] == {"x-hook": "required"}
|
||||
assert observed == ["arealtime_calls", "_arealtime"]
|
||||
if limiter_version == 3:
|
||||
assert get_request_stash() is initial_stash
|
||||
assert initial_stash.owner_litellm_call_id == "signaling"
|
||||
if limit == "max_parallel_requests":
|
||||
await limiter.async_log_success_event(
|
||||
{
|
||||
"litellm_call_id": "signaling",
|
||||
"litellm_params": {"metadata": {"user_api_key": auth.api_key, "user_api_key_model_max_budget": {}}},
|
||||
},
|
||||
litellm.ModelResponse(usage=litellm.Usage(total_tokens=0)),
|
||||
datetime.now(),
|
||||
datetime.now(),
|
||||
)
|
||||
if limiter_version == 3:
|
||||
assert initial_stash.parallel_slot is None
|
||||
await proxy.pre_call_hook(
|
||||
auth, {"model": "voice", "litellm_call_id": "next", "metadata": {}}, "arealtime_calls"
|
||||
)
|
||||
if limiter_version == 1 and limit == "max_parallel_requests":
|
||||
from litellm.proxy._types import InternalRequestOrigin
|
||||
|
||||
await asyncio.sleep(0)
|
||||
observer_kwargs = {
|
||||
"internal_request_origin": InternalRequestOrigin.REALTIME_OBSERVER,
|
||||
"litellm_call_id": "observer",
|
||||
"litellm_params": {"metadata": {"user_api_key": auth.api_key, "user_api_key_model_max_budget": {}}},
|
||||
}
|
||||
await limiter.async_log_success_event(
|
||||
observer_kwargs,
|
||||
litellm.ModelResponse(usage=litellm.Usage(total_tokens=17)),
|
||||
datetime.now(),
|
||||
datetime.now(),
|
||||
)
|
||||
current = await limiter.internal_usage_cache.async_get_cache(
|
||||
key=f"{auth.api_key}::{datetime.now():%Y-%m-%d-%H-%M}::request_count", litellm_parent_otel_span=None
|
||||
)
|
||||
assert current["current_requests"] == 1
|
||||
assert current["current_tpm"] == 17
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await proxy.pre_call_hook(
|
||||
auth,
|
||||
{"model": "voice", "litellm_call_id": "forged", "metadata": {}, "internal_realtime_observer": True},
|
||||
"_arealtime",
|
||||
)
|
||||
assert error.value.status_code == 429
|
||||
finally:
|
||||
_request_stash.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scope", ["key", "user", "team", "end_user"])
|
||||
async def test_internal_observer_missing_legacy_counter_only_adds_usage(scope):
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.proxy._types import InternalRequestOrigin
|
||||
from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
|
||||
from litellm.proxy.utils import InternalUsageCache
|
||||
|
||||
limiter = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache=DualCache()))
|
||||
metadata = {"user_api_key": "expired-key", "user_api_key_model_max_budget": {}}
|
||||
if scope in ("user", "team"):
|
||||
metadata[f"user_api_key_{scope}_id"] = "expired-scope"
|
||||
kwargs = {
|
||||
"internal_request_origin": InternalRequestOrigin.REALTIME_OBSERVER,
|
||||
"litellm_params": {"metadata": metadata},
|
||||
**({"user": "expired-scope"} if scope == "end_user" else {}),
|
||||
}
|
||||
await limiter.async_log_success_event(
|
||||
kwargs, litellm.ModelResponse(usage=litellm.Usage(total_tokens=23)), datetime.now(), datetime.now()
|
||||
)
|
||||
identity = "expired-key" if scope == "key" else "expired-scope"
|
||||
current = await limiter.internal_usage_cache.async_get_cache(
|
||||
key=f"{identity}::{datetime.now():%Y-%m-%d-%H-%M}::request_count", litellm_parent_otel_span=None
|
||||
)
|
||||
assert current == {"current_requests": 0, "current_tpm": 23, "current_rpm": 0}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue