diff --git a/.github/codeql/codeql-config.yml b/.github/codeql/codeql-config.yml index 6e15c1069a3..480ba4b3d7e 100644 --- a/.github/codeql/codeql-config.yml +++ b/.github/codeql/codeql-config.yml @@ -12,6 +12,11 @@ queries: query-filters: - exclude: id: py/clear-text-logging-sensitive-data # CWE-312 + # CodeQL 2.27.0 exceeds its 2 GiB result-set limit while evaluating the + # repository-wide log-injection data-flow query. Keep the remaining Python + # security-and-quality queries enabled until the upstream query scales. + - exclude: + id: py/log-injection # CWE-117 - exclude: id: py/polynomial-redos # CWE-730 # Import resolution confuses stdlib types with management_endpoints/types.py. diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index fac0d766535..2eff6c951a5 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -103,6 +103,16 @@ jobs: outputs: decision: ${{ steps.changes.outputs.decision }} has-coverage: ${{ steps.tests.outputs.has-coverage }} + services: + redis: + image: ${{ inputs.artifact-name == 'proxy-auth' && 'redis:8.2.9-alpine@sha256:30abb90e62f14b737010746def3ba99cc79fe19dcdb3d37b41f21fc62e7da19d' || '' }} + ports: + - '127.0.0.1::6379' + options: >- + --health-cmd "redis-cli ping" + --health-interval 2s + --health-timeout 2s + --health-retries 15 steps: - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 @@ -183,6 +193,7 @@ jobs: TEST_TIMEOUT_SECONDS: ${{ inputs.test-timeout-seconds }} DIST: ${{ inputs.dist }} COVERAGE_CORE: sysmon + LITELLM_TEST_REDIS_PORT: ${{ job.services.redis.ports['6379'] }} run: | echo "has-coverage=false" >> "$GITHUB_OUTPUT" selection="${TEST_PATH}" diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index 9a85ced57f6..c3c0b026210 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -76,9 +76,14 @@ jobs: if: matrix.language == 'python' uses: advanced-security/filter-sarif@2da736ff05ef065cb2894ac6892e47b5eac2c3c0 # v1.1 with: + # These SHA-256 digests are opaque ownership/cache identifiers, not password hashes. patterns: | -litellm/llms/oci/common_utils.py:py/weak-sensitive-data-hashing -litellm/proxy/auth/password_policy.py:py/weak-sensitive-data-hashing + -litellm/proxy/_types.py:py/weak-sensitive-data-hashing + -litellm/proxy/realtime_endpoints/call_sessions.py:py/weak-sensitive-data-hashing + -litellm/proxy/realtime_endpoints/live.py:py/weak-sensitive-data-hashing + -litellm/proxy/utils.py:py/weak-sensitive-data-hashing input: sarif-results/python.sarif output: sarif-results/python.sarif diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index d3604a81e8f..e576dbb0ec8 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -123,6 +123,7 @@ jobs: tests/test_litellm/proxy/hooks tests/test_litellm/proxy/policy_engine tests/test_litellm/proxy/client + tests/local_testing/test_realtime_call_redis.py workers: 2 reruns: 2 timeout-minutes: 20 diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 51a4d8f716c..d8296774409 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -60,6 +60,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = ( # Tools / agents (registry & policy admin) "/v1/tool/", "/v1/agents", + "/v1/traces", # Guardrails admin "/v2/guardrails/", # MCP server admin + BYOK OAuth flow (UI-initiated) + dynamic per-server endpoints diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 6e91f5486d0..43f17cac646 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -114,6 +114,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = ( "/{provider}/", "/toolset/", # Realtime / streaming + "/v1/live", + "/live", "/v1/realtime", "/realtime", # Health & ops diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 2e7a6b178a8..7944201fd9f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -120,13 +120,13 @@ impl NativeTraceStorage { fn query<'py>( &self, py: Python<'py>, - query: &str, + sql: &str, #[pyo3(from_py_with = litellm_host_python::from_py_argument)] parameters: BTreeMap< String, Parameter, >, ) -> PyResult> { - let query = ReadQuery::parse(query).map_err(map_error)?; + let query = ReadQuery::parse(sql).map_err(map_error)?; let connection = self.reader.clone().ok_or_else(|| { PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL") })?; diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 238b7cc3fdd..189b8ae38d6 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -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, ValidationError from typing_extensions import ReadOnly, TypedDict import litellm @@ -2796,6 +2796,13 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor): result["response"].get("usage", {}) ) usage_objects.append(usage_object) + usage_objects.extend( + ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( # pyright: ignore[reportPrivateUsage] # reuse the existing Responses usage conversion for nested Live events + response.usage.model_dump() + ) + for response in _live_backend_responses(results) + if response.usage is not None + ) return usage_objects @staticmethod @@ -2978,7 +2985,15 @@ def handle_realtime_stream_cost_calculation( potential_model_names.append(litellm_model_name) input_cost_per_token, output_cost_per_token = _first_priced_realtime_token_costs( potential_model_names=potential_model_names, - combined_usage_object=combined_usage_object, + combined_usage_object=( + RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results( + [ # mutable-ok: collector requires a concrete event list + event for event in results if event.get("type") != "response.event" + ] + ) + if any(event.get("type") == "response.event" for event in results) + else combined_usage_object + ), custom_llm_provider=custom_llm_provider, data_residency=data_residency, ) @@ -2992,7 +3007,18 @@ 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, + ) + backend_cost: Final = sum( + _live_backend_response_cost(response, litellm_logging_obj) + for response in _live_backend_responses(results, litellm_logging_obj) + ) + total_cost: Final = ( + input_cost_per_token + output_cost_per_token + transcription_cost + live_audio_cost + backend_cost + ) _store_cost_breakdown_in_logging_obj( litellm_logging_obj=litellm_logging_obj, @@ -3000,13 +3026,122 @@ 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={ # mutable-ok: logging cost breakdown requires a concrete dict + name: cost + for name, cost in ( + ("transcription_cost", transcription_cost), + ("live_audio_cost", live_audio_cost), + ("live_backend_cost", backend_cost), + ) + if cost > 0 + } + or None, data_residency=data_residency, ) return total_cost +class _LiveBackendEvent(BaseModel): + type: str + response: object = None + + +class _LiveBackendEnvelope(BaseModel): + event: _LiveBackendEvent + + +def _live_backend_responses( + results: OpenAIRealtimeStreamList, logging_obj: LitellmLoggingObject | None = None +) -> tuple[ResponsesAPIResponse, ...]: + responses: Final = { # mutable-ok: deduplicate terminal backend responses by response id + response.id: response + for result in results + if result.get("type") == "response.event" + and (response := _live_backend_response(result, logging_obj)) is not None + } + return tuple(responses.values()) + + +def _mark_live_backend_accounting_incomplete(logging_obj: LitellmLoggingObject | None) -> None: + verbose_logger.warning("Live backend accounting incomplete: missing valid terminal usage or model pricing") + if logging_obj is not None: + logging_obj.model_call_details["realtime_backend_accounting_incomplete"] = True + + +def _live_backend_response( + result: Mapping[str, object], logging_obj: LitellmLoggingObject | None +) -> ResponsesAPIResponse | None: + try: + event: Final = _LiveBackendEnvelope.model_validate(result).event + except ValidationError: + return None + if event.type not in ("response.completed", "response.incomplete", "response.failed"): + return None + try: + response: Final = ResponsesAPIResponse.model_validate(event.response) + except ValidationError: + _mark_live_backend_accounting_incomplete(logging_obj) + return None + if response.usage is None: + _mark_live_backend_accounting_incomplete(logging_obj) + return None + return response + + +def _live_backend_response_cost(response: ResponsesAPIResponse, logging_obj: LitellmLoggingObject | None) -> float: + try: + return completion_cost( + completion_response=response, model=response.model, custom_llm_provider="openai", call_type="aresponses" + ) + except Exception: # noqa: BLE001 # preserve measured voice cost when backend pricing cannot be resolved + _mark_live_backend_accounting_incomplete(logging_obj) + return 0.0 + + +def handle_live_session_duration_cost( + results: OpenAIRealtimeStreamList, + custom_llm_provider: str, + litellm_model_name: str, +) -> float: + seconds: Final = max( + ( + duration + for event in results + if event.get("type") in ("session.closed", "session.usage.updated") + and (duration := _live_duration_seconds(event)) is not None + ), + default=0.0, + ) + initialization_seconds: Final = max( + ( + duration + for event in results + if event.get("type") == "litellm.live.initialization" + and (duration := _live_duration_seconds(event)) is not None + ), + default=0.0, + ) + try: + model_info: Final = litellm.get_model_info(model=litellm_model_name, custom_llm_provider=custom_llm_provider) + except Exception: # noqa: BLE001 # an unknown model simply has no per-second price to bill + return 0.0 + return max(seconds, initialization_seconds) * (model_info.get("input_cost_per_second") or 0.0) + + +def _live_duration_seconds(event: Mapping[str, object]) -> float | None: + from litellm.types.realtime import LiveSessionUsageEvent + + try: + usage: Final = LiveSessionUsageEvent.model_validate(event).usage + except ValidationError: + return None + raw_usage: Final = cast( # cast-ok: LiveSessionUsageEvent validated the usage mapping above + Mapping[str, object], event.get("usage") + ) + return usage.duration / (1 if "seconds" in raw_usage else 1000) + + def handle_realtime_transcription_cost_calculation( results: OpenAIRealtimeStreamList, custom_llm_provider: str, diff --git a/litellm/images/main.py b/litellm/images/main.py index 7dc68dafecc..64589300b13 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -9,6 +9,7 @@ if TYPE_CHECKING: from litellm.images.utils import ImageEditRequestUtils import httpx +from pydantic import TypeAdapter import litellm @@ -372,6 +373,7 @@ def image_generation( # Providers using llm_http_handler ######################################################### elif custom_llm_provider in ( + litellm.LlmProviders.CHATGPT, litellm.LlmProviders.RECRAFT, litellm.LlmProviders.AIML, litellm.LlmProviders.GEMINI, @@ -397,6 +399,9 @@ def image_generation( model=model, prompt=prompt, image_generation_provider_config=image_generation_config, + extra_headers=TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python( + extra_headers + ), image_generation_optional_request_params=optional_params, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params_dict, @@ -854,16 +859,28 @@ def image_edit( additional_drop_params=kwargs.get("additional_drop_params"), ) - if image_edit_provider_config.use_multipart_form_data() and ( - custom_llm_provider == "openai" - or custom_llm_provider == "azure" - or custom_llm_provider in litellm.openai_compatible_providers - ): + if ( + image_edit_provider_config.use_multipart_form_data() + and ( + custom_llm_provider == "openai" + or custom_llm_provider == "azure" + or custom_llm_provider in litellm.openai_compatible_providers + ) + ) or custom_llm_provider == litellm.LlmProviders.CHATGPT: image_edit_request_params.update( flatten_form_field_values( non_default_params, extra_body if isinstance(extra_body, dict) else None, ) + if image_edit_provider_config.use_multipart_form_data() + else { # mutable-ok: image provider update requires a concrete request-parameter dict + **non_default_params, + **( + extra_body + if isinstance(extra_body, dict) + else {} # mutable-ok: empty fallback is consumed immediately + ), + } ) # Pre Call logging @@ -962,9 +979,9 @@ def image_edit( @client async def aimage_edit( - image: FileTypes | list[FileTypes], - model: str, - prompt: str, + image: FileTypes | list[FileTypes] | None = None, + model: str = "", + prompt: str = "", mask: str | None = None, n: int | None = None, quality: str | ImageGenerationRequestQuality | None = None, @@ -1002,11 +1019,9 @@ async def aimage_edit( model=model, api_base=local_vars.get("base_url", None) ) - images: Final = image if isinstance(image, list) else [image] - func: Final = partial( image_edit, - image=images, + image=image, prompt=prompt, mask=mask, model=model, diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 071960ba65f..138b4297553 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -9,6 +9,7 @@ from types import MappingProxyType from typing import Final, Literal, TypedDict, cast from zoneinfo import ZoneInfo, ZoneInfoNotFoundError +from pydantic import TypeAdapter from typing_extensions import ReadOnly import litellm @@ -1859,7 +1860,34 @@ def calculate_image_response_cost_from_usage( custom_llm_provider=custom_llm_provider, model_info=model_info, ) - return prompt_cost + completion_cost + cached_details: Final = ( + input_tokens_details.get("cached_tokens_details") + if isinstance(input_tokens_details, dict) + else getattr(input_tokens_details, "cached_tokens_details", None) + ) + if cached_details is None: + return prompt_cost + completion_cost + catalog_model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider) + details_adapter: Final = TypeAdapter[object](object) + cached_token_details: Final = details_adapter.validate_python(cached_details) + input_token_details: Final = details_adapter.validate_python(input_tokens_details) + cached_text: Final = _get_token_detail_value(cached_token_details, "text_tokens") or 0 + cached_image: Final = _get_token_detail_value(cached_token_details, "image_tokens") or 0 + input_text_tokens: Final = _get_token_detail_value(input_token_details, "text_tokens") or 0 + input_image_tokens: Final = _get_token_detail_value(input_token_details, "image_tokens") or 0 + if not (0 <= cached_text <= input_text_tokens and 0 <= cached_image <= input_image_tokens): + raise ValueError("Image cached token counts exceed their input modality counts") + text_rate: Final = catalog_model_info.get("input_cost_per_token") or 0.0 + image_rate: Final = catalog_model_info.get("input_cost_per_image_token") + cache_text_rate: Final = catalog_model_info.get("cache_read_input_token_cost") + cache_image_rate: Final = catalog_model_info.get("cache_read_input_image_token_cost") + text_savings: Final = cached_text * (text_rate - cache_text_rate) if cache_text_rate is not None else 0.0 + image_savings: Final = ( + cached_image * ((image_rate if image_rate is not None else text_rate) - cache_image_rate) + if cache_image_rate is not None + else 0.0 + ) + return prompt_cost + completion_cost - text_savings - image_savings def calculate_image_response_web_search_cost( diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index d2fbb26bb02..64481d66059 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1,11 +1,14 @@ import asyncio import json import traceback -from collections.abc import Coroutine, Mapping, Sequence +from collections.abc import Awaitable, Callable, Coroutine, Mapping, Sequence +from contextvars import ContextVar from dataclasses import dataclass from enum import Enum, auto +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast +from pydantic import TypeAdapter from typing_extensions import ReadOnly import litellm @@ -14,9 +17,11 @@ from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESS from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend from litellm.types.llms.openai import ( + OpenAILiveResponseEvent, OpenAIRealtimeEvents, OpenAIRealtimeOutputItemDone, OpenAIRealtimeResponseDelta, + OpenAIRealtimeSessionClosed, OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents, ) @@ -25,6 +30,10 @@ from litellm.types.realtime import ALL_DELTA_TYPES from .litellm_logging import Logging as LiteLLMLogging from .realtime_errors import client_close_code, realtime_error_event, websocket_close_reason +realtime_attachment_cleanup: Final[ContextVar[Callable[[], Awaitable[None]] | None]] = ContextVar( + "realtime_attachment_cleanup", default=None +) + if TYPE_CHECKING: from websockets.asyncio.client import ClientConnection from websockets.exceptions import ConnectionClosed @@ -137,12 +146,25 @@ class RealTimeStreaming: force_transcription_model: str | None = None, event_normalizer: RealtimeEventNormalizer | None = None, logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER, + *, + account_usage: bool = True, + live_initialization_seconds: float = 0, ): 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] = [] + if account_usage and live_initialization_seconds > 0: + self.messages.append( + { # mutable-ok: initialization event is appended to the mutable event history + "type": "litellm.live.initialization", + "usage": { # mutable-ok: usage payload is consumed as part of the typed event + "seconds": live_initialization_seconds, + }, + } + ) self._backend_sent_frames: bool = False self.input_message: dict = {} self.input_messages: list[dict[str, str]] = [] @@ -254,6 +276,42 @@ 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") in ("session.closed", "session.usage.updated") and isinstance( + message_obj.get("usage"), dict + ): + self.messages.append(TypeAdapter(OpenAIRealtimeSessionClosed).validate_python(message_obj)) + return + if message_obj.get("type") == "response.event" and isinstance(message_obj.get("event"), dict): + nested: Final = message_obj["event"] + if nested.get("type") in ("response.completed", "response.incomplete", "response.failed"): + response: Final = nested.get("response") + response_mapping: Final = ( + TypeAdapter(Mapping[str, object]).validate_python(response) + if isinstance(response, Mapping) + else None + ) + # Retain billing evidence even when response content is excluded from logging. + filtered: Final[OpenAILiveResponseEvent] = { + "type": "response.event", + "event": { + "type": nested["type"], + "response": { + **MappingProxyType( + { + key: value + for key, value in response_mapping.items() + if key in ("id", "created_at", "model", "usage", "service_tier") + } + ), + "output": TypeAdapter(list[object]).validate_python(()), + } + if response_mapping is not None + else None, + }, + } + stored: Final = message_obj if self._should_store_message(message_obj) else filtered + self.messages.append(TypeAdapter(OpenAILiveResponseEvent).validate_python(stored)) + return if not self._should_store_message(message_obj): return try: @@ -408,8 +466,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 @@ -419,9 +479,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: @@ -1573,7 +1636,12 @@ class RealTimeStreaming: finally: forward_task.cancel() client_task.cancel() - await asyncio.gather(forward_task, client_task, return_exceptions=True) + try: + await asyncio.gather(forward_task, client_task, return_exceptions=True) + finally: + cleanup: Final = realtime_attachment_cleanup.get() + if not self._account_usage and cleanup is not None: + await cleanup() async def _close_client(self, close: BackendClose) -> None: redacted_message: Final = redact_internal_details_from_client_message(close.message) diff --git a/litellm/llms/base_llm/realtime/http_transformation.py b/litellm/llms/base_llm/realtime/http_transformation.py index 43a80edb493..80daebd88b2 100644 --- a/litellm/llms/base_llm/realtime/http_transformation.py +++ b/litellm/llms/base_llm/realtime/http_transformation.py @@ -7,6 +7,7 @@ These are HTTP (not WebSocket) endpoints used by the WebRTC flow: """ from abc import ABC, abstractmethod +from collections.abc import Mapping from typing import Final import httpx @@ -36,6 +37,14 @@ class BaseRealtimeHTTPConfig(ABC): explicit api_base → litellm.api_base → env var → hard-coded default """ + def resolve_api_base(self, api_base: str | None, dynamic_api_base: str | None) -> str: + return self.get_api_base(dynamic_api_base or api_base) + + def get_realtime_calls_extra_headers( + self, headers: dict[str, object] | None + ) -> dict[str, object] | None: # mutable-ok: shared HTTP handler accepts a mutable header dictionary + return headers + @abstractmethod def get_api_key( self, @@ -97,6 +106,11 @@ class BaseRealtimeHTTPConfig(ABC): "Authorization": f"Bearer {ephemeral_key}", } + def transform_realtime_calls_response( + self, response: httpx.Response, model: str, model_id: str | None, headers: Mapping[str, object] | None + ) -> httpx.Response: + return response + # ------------------------------------------------------------------ # # Error handling # # ------------------------------------------------------------------ # diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index 563826c2b93..4a5045e02e7 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -49,8 +49,9 @@ class Authenticator: self.auth_file = os.path.join(self.token_dir, os.getenv("CHATGPT_AUTH_FILE", "auth.json")) self._ensure_token_dir() - def get_api_base(self) -> str: - return os.getenv("CHATGPT_API_BASE") or os.getenv("OPENAI_CHATGPT_API_BASE") or CHATGPT_API_BASE + @staticmethod + def get_api_base(default_base: str = CHATGPT_API_BASE) -> str: + return os.getenv("CHATGPT_API_BASE") or os.getenv("OPENAI_CHATGPT_API_BASE") or default_base def get_access_token(self) -> str: auth_data: Final = self._read_auth_file() diff --git a/litellm/llms/chatgpt/chat/transformation.py b/litellm/llms/chatgpt/chat/transformation.py index 1b110704c8b..eceef637d90 100644 --- a/litellm/llms/chatgpt/chat/transformation.py +++ b/litellm/llms/chatgpt/chat/transformation.py @@ -23,8 +23,8 @@ class ChatGPTConfig(OpenAIConfig): super().__init__() self.authenticator = Authenticator() - def api_base_without_login(self) -> str: - return self.authenticator.get_api_base() + def api_base_without_login(self, api_base: str | None = None) -> str: + return api_base or self.authenticator.get_api_base() def _get_openai_compatible_provider_info( self, @@ -33,7 +33,7 @@ class ChatGPTConfig(OpenAIConfig): api_key: str | None, custom_llm_provider: str, ) -> tuple[str | None, str | None, str]: - dynamic_api_base: Final = self.api_base_without_login() + dynamic_api_base: Final = self.api_base_without_login(api_base) try: dynamic_api_key: Final = self.authenticator.get_access_token() except GetAccessTokenError as e: diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py new file mode 100644 index 00000000000..92595890c86 --- /dev/null +++ b/litellm/llms/chatgpt/codex.py @@ -0,0 +1,95 @@ +from collections.abc import Mapping +from typing import Final +from urllib.parse import urlsplit + +import httpx +from pydantic import BaseModel, Field, TypeAdapter +from typing_extensions import ReadOnly, TypedDict + +from litellm.types.realtime import RealtimeQueryParams, RealtimeSessionConfig + + +class CodexRealtimeOffer(BaseModel): + sdp: str = Field(min_length=1) + session: RealtimeSessionConfig + + +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 | tuple[str, ...]] | None = None + usage_supervised: bool = False + parallel_reserved: 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 | tuple[str, ...]] | None = None + + +class CodexSidebandRequest(TypedDict): + api_base: ReadOnly[str | None] + model: ReadOnly[str] + chatgpt_realtime_call_id: ReadOnly[str] + query_params: ReadOnly[RealtimeQueryParams] + extra_headers: ReadOnly[Mapping[str, str] | None] + extra_query: ReadOnly[Mapping[str, str | tuple[str, ...]] | None] + + +def build_call_request( + offer: CodexRealtimeOffer, query: Mapping[str, str], headers: Mapping[str, str] +) -> dict[str, object]: # mutable-ok: proxy processor enriches the request dictionary + return { # mutable-ok: proxy processor enriches the request dictionary + "model": offer.session.model, + "sdp_body": offer.sdp.encode(), + "session": offer.session.model_dump(exclude_none=True), + "openai_ephemeral_key": "", + "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 + key: value + for key, value in headers.items() + if key in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation") + }, + } + + +def parse_call_response(response: httpx.Response, alias: str, owner: str, expires_at: float) -> CodexRealtimeCall: + routing_data: Final = response.extensions.get("chatgpt_realtime") + if not routing_data: + raise ValueError("Direct call signaling requires a ChatGPT deployment") + routing: Final = ChatGPTCallRouting.model_validate(routing_data) + location: Final = TypeAdapter(str).validate_python(response.headers.get("location", "")) + call_id: Final[str] = urlsplit(location).path.rstrip("/").rsplit("/", 1)[-1] + 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, + ) + + +def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: + return CodexSidebandRequest( + api_base=call.api_base, + model=f"chatgpt/{call.model}", + chatgpt_realtime_call_id=call.call_id, + query_params=RealtimeQueryParams(model=call.model), + extra_headers=call.extra_headers, + extra_query=call.extra_query, + ) diff --git a/litellm/llms/chatgpt/common_utils.py b/litellm/llms/chatgpt/common_utils.py index fe33219f110..405631d8192 100644 --- a/litellm/llms/chatgpt/common_utils.py +++ b/litellm/llms/chatgpt/common_utils.py @@ -4,6 +4,8 @@ Constants and helpers for ChatGPT subscription OAuth. import os import platform +from collections.abc import Mapping +from types import MappingProxyType from typing import Any, Final from uuid import uuid4 @@ -105,6 +107,12 @@ You are producing plain text that will later be styled by the CLI. Follow these """ +def without_oauth_identity_headers(headers: Mapping[str, object]) -> Mapping[str, object]: + return MappingProxyType( + {key: value for key, value in headers.items() if key.lower() not in ("authorization", "chatgpt-account-id")} + ) + + class ChatGPTAuthError(BaseLLMException): def __init__( self, diff --git a/litellm/llms/chatgpt/images.py b/litellm/llms/chatgpt/images.py new file mode 100644 index 00000000000..554373aa64e --- /dev/null +++ b/litellm/llms/chatgpt/images.py @@ -0,0 +1,154 @@ +import base64 +import os +from collections.abc import Mapping, Sequence +from pathlib import Path +from types import MappingProxyType +from typing import Final + +from httpx._types import FileTypes as HTTPFileTypes +from httpx._types import RequestFiles +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter + +from litellm.images.utils import ImageEditRequestUtils +from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig +from litellm.llms.openai.image_generation.gpt_transformation import GPTImageGenerationConfig +from litellm.types.llms.openai import AllMessageValues, FileTypes +from litellm.types.router import GenericLiteLLMParams + +from .authenticator import Authenticator +from .common_utils import without_oauth_identity_headers +from .responses.transformation import ChatGPTResponsesAPIConfig + + +class ReferenceImage(BaseModel): + model_config = ConfigDict(extra="forbid") + image_url: str = Field(pattern=r"^(data:image/(png|jpeg|webp);base64,|https://)") + + +def encode_reference( + file: HTTPFileTypes | FileTypes, +) -> dict[str, str]: # mutable-ok: image handler requires dictionaries + content: Final = file[1] if isinstance(file, tuple) else file + raw: Final = ( + Path(os.fsdecode(content)).read_bytes() + if isinstance(content, os.PathLike) + else content.encode() + if isinstance(content, str) + else content + if isinstance(content, bytes) + else content.read() + ) + content_type: Final = ( + file[2] + if isinstance(file, tuple) and len(file) >= 3 and file[2] + else ImageEditRequestUtils.get_image_content_type(raw) + ) + if content_type not in ("image/png", "image/jpeg", "image/webp"): + raise ValueError("Reference images must be PNG, JPEG, or WEBP") + return { # mutable-ok: JSON request serialization + "image_url": f"data:{content_type};base64," + base64.b64encode(raw).decode("ascii") + } + + +def image_headers( + headers: Mapping[str, object], model: str, params: Mapping[str, object] +) -> dict[str, object]: # mutable-ok: image handler requires dictionaries + auth_headers: Final = ChatGPTResponsesAPIConfig().validate_environment( + headers={}, # mutable-ok: Responses adapter header contract + model=model, + litellm_params=GenericLiteLLMParams.model_validate(params), + ) + return { # mutable-ok: image handler requires dictionaries + **without_oauth_identity_headers(headers), + **auth_headers, + "accept": "application/json", + } + + +class ChatGPTImageGenerationConfig(GPTImageGenerationConfig): + def validate_environment( + self, + headers: Mapping[str, object], + model: str, + messages: Sequence[AllMessageValues], + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + api_key: str | None = None, + api_base: str | None = None, + ) -> dict[str, object]: # mutable-ok: image handler requires dictionaries + return image_headers(headers, model, litellm_params) + + def get_complete_url( + self, + api_base: str | None, + api_key: str | None, + model: str, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + stream: bool | None = None, + ) -> str: + return f"{(api_base or Authenticator.get_api_base()).rstrip('/')}/images/generations" + + def transform_image_generation_request( + self, + model: str, + prompt: str, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + headers: Mapping[str, object], + ) -> dict[str, object]: # mutable-ok: image handler requires dictionaries + return {"model": model, "prompt": prompt, **optional_params} # mutable-ok: JSON request serialization + + +class ChatGPTImageEditConfig(OpenAIImageEditConfig): + def validate_environment( + self, + headers: Mapping[str, object], + model: str, + api_key: str | None = None, + litellm_params: Mapping[str, object] | None = None, + api_base: str | None = None, + ) -> dict[str, object]: # mutable-ok: image handler requires dictionaries + return image_headers(headers, model, litellm_params or MappingProxyType({})) + + def get_complete_url(self, model: str, api_base: str | None, litellm_params: Mapping[str, object]) -> str: + return f"{(api_base or Authenticator.get_api_base()).rstrip('/')}/images/edits" + + def use_multipart_form_data(self) -> bool: + return False + + def transform_image_edit_request( + self, + model: str, + prompt: str | None, + image: FileTypes | Sequence[FileTypes] | None, + image_edit_optional_request_params: Mapping[str, object], + litellm_params: GenericLiteLLMParams, + headers: Mapping[str, object], + ) -> tuple[dict[str, object], RequestFiles]: # mutable-ok: image handler requires dictionaries + if image_edit_optional_request_params.get("mask") is not None: + raise ValueError("ChatGPT image editing does not support masks") + references: Final = getattr(litellm_params, "images", None) + if references is not None: + if image: + raise ValueError("Specify only one of image or images") + validated: Final = TypeAdapter(tuple[ReferenceImage, ...]).validate_python(references) + if not 1 <= len(validated) <= 5: + raise ValueError("images must contain between 1 and 5 reference images") + return { # mutable-ok: JSON request serialization + "prompt": prompt, + **image_edit_optional_request_params, + "model": model, # the authenticated alias wins over passthrough fields + "images": tuple(item.model_dump() for item in validated), + }, () + + inputs: Final = tuple(image) if isinstance(image, list) else (image,) if image is not None else () + encoded: Final = tuple(encode_reference(file) for file in inputs) + if not 1 <= len(encoded) <= 5: + raise ValueError("images must contain between 1 and 5 reference images") + return { # mutable-ok: JSON request serialization + "prompt": prompt, + **image_edit_optional_request_params, + "model": model, # the authenticated alias wins over passthrough fields + "images": encoded, + }, () diff --git a/litellm/llms/chatgpt/live.py b/litellm/llms/chatgpt/live.py new file mode 100644 index 00000000000..23ea0e77630 --- /dev/null +++ b/litellm/llms/chatgpt/live.py @@ -0,0 +1,196 @@ +import re +from collections.abc import Mapping +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, TypeAlias +from unicodedata import category +from urllib.parse import quote, unquote + +import httpx +from pydantic import JsonValue + +from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES +from litellm.llms.chatgpt.realtime import ( + ChatGPTRealtime, + configured_realtime_headers, + configured_realtime_query, + realtime_headers, +) +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # shared client has legacy untyped optional params + get_shared_realtime_ssl_context, +) +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders + +if TYPE_CHECKING: + from websockets.asyncio.client import ClientConnection + +LiveQuery: TypeAlias = Mapping[str, str | int | float | bool | None | tuple[str | int | float | bool | None, ...]] +LiveBody: TypeAlias = Mapping[str, JsonValue] +LiveOperation: TypeAlias = Literal["fork", "accept", "reject", "refer", "hangup", "content", "attach"] +_PATH: Final = re.compile(r"live/sessions(?:/([^/]+)/(fork|accept|reject|refer|hangup|content|attach))?\Z") +_ROUTING_QUERY: Final = frozenset(("model", "session_id", "call_id", "api_key", "api_base", "authorization")) + + +@dataclass(frozen=True, slots=True) +class LiveDeployment: + model: str + model_id: str | None = None + provider: Literal["chatgpt", "openai"] = "chatgpt" + api_base: str | None = None + api_key: str | None = field(default=None, repr=False) + extra_headers: Mapping[str, str] = field(default_factory=lambda: MappingProxyType({}), repr=False) + extra_query: LiveQuery = field(default_factory=lambda: MappingProxyType({}), repr=False) + + +def _validate_session_id(session_id: str) -> None: + candidate: str = session_id # rebind-ok: inspect every decoding layer without recursive stack exhaustion + while True: + if ( + not candidate + or candidate in (".", "..") + or any(char in ("/", "\\") or category(char) in ("Cc", "Cs") for char in candidate) + ): + raise ValueError("Invalid Live session ID") + decoded: str = unquote(candidate, errors="strict") + if decoded == candidate: + return + candidate = decoded + + +def _validate_path(path: str) -> None: + match: Final = _PATH.fullmatch(path) + if match is None: + raise ValueError("Invalid Live endpoint") + if match.group(1) is None: + return + session_id: Final = unquote(match.group(1), errors="strict") + if quote(session_id, safe="") != match.group(1): + raise ValueError("Noncanonical Live session path") + _validate_session_id(session_id) + + +def live_session_path(session_id: str, operation: LiveOperation) -> str: + _validate_session_id(session_id) + path: Final = f"live/sessions/{quote(session_id, safe='')}/{operation}" + _validate_path(path) + return path + + +class LiveTransport: + def __init__( + self, + deployment: LiveDeployment, + inbound_headers: Mapping[str, str], + *, + http_handler: AsyncHTTPHandler | None = None, + ) -> None: + self.deployment = deployment + self._http_handler = http_handler + params: Final = GenericLiteLLMParams.model_validate( + MappingProxyType( + { + "api_base": deployment.api_base, + "extra_query": deployment.extra_query, + } + ) + ) + self._query = configured_realtime_query(params) + self._headers = ( + realtime_headers(params, inbound_headers, deployment.extra_headers) + if deployment.provider == "chatgpt" + else MappingProxyType( + { + **MappingProxyType( + { + key.lower(): value + for key, value in inbound_headers.items() + if key.lower() in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation") + } + ), + **configured_realtime_headers(deployment.extra_headers), + "authorization": f"Bearer {deployment.api_key or ''}", + } + ) + ) + + def _url(self, path: str, query: LiveQuery | None, *, websocket: bool) -> str: + _validate_path(path) + base: Final = httpx.URL( + ChatGPTRealtime.get_api_base(self.deployment.api_base) + if self.deployment.provider == "chatgpt" + else self.deployment.api_base or "https://api.openai.com/v1" + ) + if base.scheme not in ("https", "http", "wss", "ws") or not base.host or base.userinfo or base.fragment: + raise ValueError("Invalid Live API base") + merged: Final = base.params.merge(query or MappingProxyType({})).merge(self._query) + safe_query: Final = tuple( + (key, value) for key, value in merged.multi_items() if key.lower() not in _ROUTING_QUERY + ) + return str( + base.copy_with( + scheme=("wss" if base.scheme in ("https", "wss") else "ws") + if websocket + else ("https" if base.scheme in ("https", "wss") else "http"), + path=f"{base.path.rstrip('/')}/{path}", + params=safe_query, + ) + ) + + async def request( + self, + method: str, + path: str, + body: LiveBody | None = None, + query: LiveQuery | None = None, + ) -> httpx.Response: + if (method, path.rsplit("/", 1)[-1]) not in ( + ("POST", "sessions"), + ("POST", "fork"), + ("POST", "accept"), + ("POST", "reject"), + ("POST", "refer"), + ("POST", "hangup"), + ("GET", "content"), + ): + raise ValueError("Invalid Live HTTP operation") + url: Final = self._url(path, query, websocket=False) + handler: Final = self._http_handler or get_async_httpx_client( + llm_provider=LlmProviders.CHATGPT if self.deployment.provider == "chatgpt" else LlmProviders.OPENAI, + params={"follow_redirects": False}, + ) + headers: Final = {**self._headers, "content-type": "application/json"} + if method == "GET": + return await handler.get(url, headers=headers, timeout=60, follow_redirects=False) + try: + return await handler.post( + url, + headers=headers, + json=dict(body) if body is not None else None, # mutable-ok: JSON encoder requires a concrete dict + timeout=60, + ) + except httpx.HTTPStatusError as error: + return error.response + + async def connect(self, path: str, query: LiveQuery | None = None) -> "ClientConnection": + import websockets + + class DirectConnect(websockets.connect): + def process_redirect(self, exc: Exception) -> Exception: + return exc + + if path != "live/sessions" and path.rsplit("/", 1)[-1] not in ("attach", "fork"): + raise ValueError("Invalid Live WebSocket operation") + url: Final = self._url(path, query, websocket=True) + ssl_context: Final = get_shared_realtime_ssl_context() if url.startswith("wss://") else None + return await DirectConnect( + url, + additional_headers=self._headers, + max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, + max_queue=16, + ssl=True if ssl_context is False else ssl_context, + open_timeout=20, + close_timeout=10, + ) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py new file mode 100644 index 00000000000..a5aad4ccb8d --- /dev/null +++ b/litellm/llms/chatgpt/realtime.py @@ -0,0 +1,290 @@ +from collections.abc import Mapping +from enum import Enum, auto +from types import MappingProxyType +from typing import TYPE_CHECKING, Final + +from httpx import URL, QueryParams, Response +from pydantic import TypeAdapter + +from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES +from litellm.exceptions import AuthenticationError +from litellm.llms.openai.realtime.handler import OpenAIRealtime +from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig +from litellm.types.realtime import RealtimeQueryParams +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import get_model_info + +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( + without_oauth_identity_headers(headers or MappingProxyType({})) + ) + return MappingProxyType({key.lower(): value for key, value in validated.items()}) + + +def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str | tuple[str, ...]]: + inbound: Final = TypeAdapter(Mapping[str, str]).validate_python( + getattr(params, "chatgpt_realtime_client_query", None) or MappingProxyType({}) + ) + configured: Final = TypeAdapter( + Mapping[str, str | int | float | bool | None | tuple[str | int | float | bool | None, ...]] + ).validate_python(getattr(params, "extra_query", None) or MappingProxyType({})) + merged: Final = QueryParams( + tuple((key, value) for key, value in inbound.items() if key in ("intent", "architecture")) + ).merge(configured) + return MappingProxyType( + {key: merged[key] if len(merged.get_list(key)) == 1 else tuple(merged.get_list(key)) for key in merged} + ) + + +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({}) + ) + configured: Final = TypeAdapter(Mapping[str, object]).validate_python( + getattr(params, "extra_headers", None) or MappingProxyType({}) + ) + return { # mutable-ok: HTTP handler header contract + **MappingProxyType( + { + key.lower(): value + for key, value in inbound.items() + if key.lower() in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation") + } + ), + **configured_realtime_headers(configured), + } + + +def realtime_headers( + params: GenericLiteLLMParams, headers: Mapping[str, str], extra_headers: Mapping[str, object] | None = None +) -> dict[str, str]: # mutable-ok: HTTP handler header contract + forwarded: Final = MappingProxyType( + { + key.lower(): value + for key, value in headers.items() + if key.lower() in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation") + } + ) + return { # mutable-ok: HTTP handler updates headers + **ChatGPTResponsesAPIConfig().validate_environment( + headers={}, # mutable-ok: Responses adapter header contract + model="", + litellm_params=params, + ), + **forwarded, + **configured_realtime_headers(extra_headers), + } + + +def realtime_endpoint(model: str) -> str: + try: + model_info: Final = get_model_info(model, custom_llm_provider="chatgpt") + except Exception: # noqa: BLE001 # get_model_info raises bare Exception for unmapped models + return "realtime" + return "live" if "/v1/live" in (model_info.get("supported_endpoints") or ()) else "realtime" + + +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: + from websockets.exceptions import ConnectionClosed + + if realtime_endpoint(model) == "live": + 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: + from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.types.utils import LlmProviders + + 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 QueryParams(self._extra_query).multi_items() + if key not in ("model", "call_id") + ), + ) + client: Final = get_async_httpx_client(llm_provider=LlmProviders.CHATGPT) + response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10) + response.raise_for_status() + + @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") + + def __init__( + self, + params: GenericLiteLLMParams, + headers: Mapping[str, str], + extra_headers: Mapping[str, object] | None = None, + ) -> None: + 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) + self._account_usage = accounts_for_call_usage(params) + + def _get_default_api_base(self) -> str: + return self.get_api_base() + + def _resolve_api_key(self, api_key: str | None) -> str: + return "chatgpt-oauth" + + def _accounts_for_call_usage(self) -> bool: + return self._account_usage + + def _get_additional_headers( + self, api_key: str, *, openai_beta_realtime: bool = False + ) -> dict[str, str]: # mutable-ok: HTTP handler header contract + return { # mutable-ok: HTTP handler updates headers + **(MappingProxyType({"OpenAI-Beta": "realtime=v1"}) if openai_beta_realtime else MappingProxyType({})), + **self._profile_headers, + } + + def _construct_url(self, api_base: str, query_params: RealtimeQueryParams) -> str: + 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 QueryParams(self._extra_query).multi_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=gateway_query + (() if endpoint == "live" else (("call_id", self._call_id),)), + ) + ) + return str( + base.copy_with( + scheme="wss" if base.scheme in ("https", "wss") else "ws", + path=f"{base.path.rstrip('/')}/{endpoint}", + params=QueryParams(TypeAdapter(Mapping[str, str | None]).validate_python(query_params)).merge( + tuple( + (key, value) + for key, value in QueryParams(self._extra_query).multi_items() + if key not in ("model", "call_id") + ) + ), + ) + ) + + +class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): + realtime_calls_json: Final = True + + def __init__(self, params: GenericLiteLLMParams, use_codex_backend: bool = True) -> None: + self._params = params + self._use_codex_backend = use_codex_backend + + def get_api_base( + self, + api_base: str | None, + **kwargs: object, # kwargs-ok: provider interface accepts optional credentials + ) -> str: + return api_base or (Authenticator.get_api_base() if self._use_codex_backend else ChatGPTRealtime.get_api_base()) + + def resolve_api_base(self, api_base: str | None, dynamic_api_base: str | None) -> str: + return self.get_api_base(api_base) + + def get_realtime_calls_extra_headers( + self, headers: dict[str, object] | None + ) -> dict[str, object]: # mutable-ok: shared HTTP handler accepts a mutable header dictionary + return {**realtime_call_headers(self._params)} # mutable-ok: shared HTTP header contract + + def get_api_key( + self, + api_key: str | None, + **kwargs: object, # kwargs-ok: provider interface accepts optional credentials + ) -> str: + return "chatgpt-oauth" + + def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: + query: Final = configured_realtime_query(self._params) + return str(URL(f"{self.get_api_base(api_base).rstrip('/')}/realtime/calls", params=query)) + + def transform_realtime_calls_response( + self, response: Response, model: str, model_id: str | None, headers: Mapping[str, object] | None + ) -> Response: + response.extensions["chatgpt_realtime"] = MappingProxyType( + { + "model": model, + "model_id": model_id, + "api_base": ChatGPTRealtime.get_api_base(self._params.api_base), + "extra_headers": configured_realtime_headers(headers), + "extra_query": configured_realtime_query(self._params), + } + ) + return response + + def get_realtime_calls_headers( + self, ephemeral_key: str + ) -> dict[str, str]: # mutable-ok: HTTP handler header contract + if ephemeral_key: + raise AuthenticationError( + message="ChatGPT realtime calls require an authenticated JSON or multipart offer", + llm_provider="chatgpt", + model="", + ) + return realtime_headers(self._params, MappingProxyType({})) + + def validate_environment( + self, + headers: Mapping[str, str], + model: str, + api_key: str | None = None, + ) -> dict[str, str]: # mutable-ok: HTTP handler header contract + return { # mutable-ok: HTTP handler updates headers + **realtime_headers(self._params, headers), + "Content-Type": "application/json", + } + + def get_complete_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: + return f"{self.get_api_base(api_base).rstrip('/')}/realtime/client_secrets" + + def get_transcription_session_url( + self, + api_base: str | None, + model: str, + api_version: str | None = None, + ) -> str: + return f"{self.get_api_base(api_base).rstrip('/')}/realtime/transcription_sessions" diff --git a/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index 9774b762396..3167e031cdb 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -106,6 +106,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): "reasoning", "previous_response_id", "truncation", + "text", } return {k: v for k, v in request.items() if k in allowed_keys} diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 8d65aa7b0ca..a4c9215c0ca 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6251,6 +6251,9 @@ class BaseLLMHTTPHandler: Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and header auth when available; falls back to the legacy OpenAI-style defaults. """ + from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( llm_provider=litellm.LlmProviders.OPENAI, @@ -6276,7 +6279,11 @@ class BaseLLMHTTPHandler: } if extra_headers: - headers.update(extra_headers) + headers.update( + without_oauth_identity_headers(extra_headers) + if isinstance(provider_config, ChatGPTRealtimeHTTPConfig) + else extra_headers + ) logging_obj.pre_call( input=request_data, @@ -6327,6 +6334,9 @@ class BaseLLMHTTPHandler: - sdp: the SDP offer (text) - session: JSON string with {"type": "realtime", "model": "...", ...} """ + from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( llm_provider=litellm.LlmProviders.OPENAI, @@ -6344,11 +6354,15 @@ class BaseLLMHTTPHandler: } if extra_headers: - headers.update(extra_headers) + headers.update( + without_oauth_identity_headers(extra_headers) + if isinstance(provider_config, ChatGPTRealtimeHTTPConfig) + else extra_headers + ) # Build multipart form data: sdp + session JSON session_data: Final = session_config or {} - if "type" not in session_data: + if "type" not in session_data and not getattr(provider_config, "realtime_calls_json", False): session_data["type"] = "realtime" if "model" not in session_data and model: session_data["model"] = model @@ -6371,6 +6385,13 @@ class BaseLLMHTTPHandler: ) try: + if getattr(provider_config, "realtime_calls_json", False): + return await async_httpx_client.post( + url=url, + headers=headers, + json={"sdp": sdp_text, "session": session_data}, # mutable-ok: JSON signaling payload + timeout=timeout, + ) return await async_httpx_client.post( url=url, headers=headers, @@ -6572,6 +6593,14 @@ class BaseLLMHTTPHandler: raise Exception(f"Unexpected error while closing WebSocket: {close_error}") return None + @staticmethod + def _image_extra_headers(custom_llm_provider: str, headers: Mapping[str, object]) -> Mapping[str, object]: + if custom_llm_provider == "chatgpt": + from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers + + return without_oauth_identity_headers(headers) + return headers + def image_edit_handler( self, model: str, @@ -6628,7 +6657,7 @@ class BaseLLMHTTPHandler: ) if extra_headers: - headers.update(extra_headers) + headers.update(self._image_extra_headers(custom_llm_provider, extra_headers)) api_base: Final = image_edit_provider_config.get_complete_url( model=model, @@ -6729,7 +6758,7 @@ class BaseLLMHTTPHandler: ) if extra_headers: - headers.update(extra_headers) + headers.update(self._image_extra_headers(custom_llm_provider, extra_headers)) api_base: Final = image_edit_provider_config.get_complete_url( model=model, @@ -6848,7 +6877,7 @@ class BaseLLMHTTPHandler: ) if extra_headers: - headers.update(extra_headers) + headers.update(self._image_extra_headers(custom_llm_provider, extra_headers)) api_base: Final = image_generation_provider_config.get_complete_url( model=model, @@ -6956,7 +6985,7 @@ class BaseLLMHTTPHandler: ) if extra_headers: - headers.update(extra_headers) + headers.update(self._image_extra_headers(custom_llm_provider, extra_headers)) api_base: Final = image_generation_provider_config.get_complete_url( model=model, diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index bdc3a6c7908..ff59211d0ac 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -39,6 +39,14 @@ class OpenAIRealtime(OpenAIChatCompletion): """ return "https://api.openai.com/" + def _resolve_api_key(self, api_key: str | None) -> str: + if api_key is None: + raise ValueError("api_key is required for OpenAI realtime calls") + return api_key + + def _accounts_for_call_usage(self) -> bool: + return True + def _get_additional_headers( self, api_key: str, @@ -118,6 +126,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 @@ -125,8 +134,7 @@ class OpenAIRealtime(OpenAIChatCompletion): if api_base is None: api_base = self._get_default_api_base() - if api_key is None: - raise ValueError("api_key is required for OpenAI realtime calls") + resolved_api_key: Final = self._resolve_api_key(api_key) # Use all query params if provided, else fallback to just model if query_params is None: @@ -144,12 +152,12 @@ class OpenAIRealtime(OpenAIChatCompletion): "If your client expects beta event names, add 'OpenAI-Beta: realtime=v1' " "to the WebSocket headers sent to the LiteLLM proxy." ) - headers: Final = self._get_additional_headers(api_key, openai_beta_realtime=openai_beta_realtime) + headers: Final = self._get_additional_headers(resolved_api_key, openai_beta_realtime=openai_beta_realtime) # Log a masked request preview consistent with other endpoints. logging_obj.pre_call( input=None, - api_key=api_key, + api_key=resolved_api_key, additional_args={ "api_base": url, "headers": headers, @@ -173,6 +181,7 @@ class OpenAIRealtime(OpenAIChatCompletion): model if (query_params or {}).get("intent") == "transcription" else None ), event_normalizer=self._make_event_normalizer(), + account_usage=account_usage and self._accounts_for_call_usage(), ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2a01c4fe862..98ff53ec904 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -31156,6 +31156,16 @@ "max_tokens": 8191, "mode": "embedding" }, + "chatgpt/gpt-live-1-codex": { + "litellm_provider": "chatgpt", + "mode": "realtime", + "supported_endpoints": [ + "/v1/realtime/calls", + "/v1/live" + ], + "supports_audio_input": true, + "supports_audio_output": true + }, "chatgpt/gpt-5.5": { "litellm_provider": "chatgpt", "source": "https://platform.openai.com/docs/models/gpt-5.5", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 0b687340ea5..482f541aba5 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -231,10 +231,15 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( "/watsonx/", ), ), + LazyFeature( + name="live", + module_path="litellm.proxy.realtime_endpoints.live", + path_prefixes=("/openai/v1/live/sessions", "/v1/live/sessions", "/live/sessions"), + ), LazyFeature( name="realtime", module_path="litellm.proxy.realtime_endpoints.endpoints", - path_prefixes=("/openai/v1/realtime", "/v1/realtime", "/realtime"), + path_prefixes=("/openai/v1/realtime", "/v1/realtime", "/realtime", "/openai/v1/live", "/v1/live", "/live"), ), LazyFeature( name="anthropic_passthrough", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index fa4b36a03aa..c72da5ab8bc 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -16782,6 +16782,844 @@ } } }, + "live": { + "components": { + "schemas": { + "HTTPValidationError": { + "properties": { + "detail": { + "items": { + "$ref": "#/components/schemas/ValidationError" + }, + "title": "Detail", + "type": "array" + } + }, + "title": "HTTPValidationError", + "type": "object" + }, + "ValidationError": { + "properties": { + "ctx": { + "title": "Context", + "type": "object" + }, + "input": { + "title": "Input" + }, + "loc": { + "items": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "integer" + } + ] + }, + "title": "Location", + "type": "array" + }, + "msg": { + "title": "Message", + "type": "string" + }, + "type": { + "title": "Error Type", + "type": "string" + } + }, + "required": [ + "loc", + "msg", + "type" + ], + "title": "ValidationError", + "type": "object" + } + } + }, + "paths": { + "/live/sessions": { + "post": { + "operationId": "create_live_session_live_sessions_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Create Live Session", + "tags": [ + "live" + ] + } + }, + "/live/sessions/{session_id}/accept": { + "post": { + "operationId": "control_live_session_live_sessions__session_id__accept_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/live/sessions/{session_id}/content": { + "get": { + "operationId": "control_live_session_live_sessions__session_id__content_get", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/live/sessions/{session_id}/fork": { + "post": { + "operationId": "fork_live_session_live_sessions__session_id__fork_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Fork Live Session", + "tags": [ + "live" + ] + } + }, + "/live/sessions/{session_id}/hangup": { + "post": { + "operationId": "control_live_session_live_sessions__session_id__hangup_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/live/sessions/{session_id}/refer": { + "post": { + "operationId": "control_live_session_live_sessions__session_id__refer_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/live/sessions/{session_id}/reject": { + "post": { + "operationId": "control_live_session_live_sessions__session_id__reject_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/openai/v1/live/sessions": { + "post": { + "operationId": "create_live_session_openai_v1_live_sessions_post_2", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Create Live Session", + "tags": [ + "live" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/accept": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__accept_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/content": { + "get": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__content_get_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/fork": { + "post": { + "operationId": "fork_live_session_openai_v1_live_sessions__session_id__fork_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Fork Live Session", + "tags": [ + "live" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/hangup": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__hangup_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/refer": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__refer_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/reject": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__reject_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/v1/live/sessions": { + "post": { + "operationId": "create_live_session_v1_live_sessions_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Create Live Session", + "tags": [ + "live" + ] + } + }, + "/v1/live/sessions/{session_id}/accept": { + "post": { + "operationId": "control_live_session_v1_live_sessions__session_id__accept_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/v1/live/sessions/{session_id}/content": { + "get": { + "operationId": "control_live_session_v1_live_sessions__session_id__content_get", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/v1/live/sessions/{session_id}/fork": { + "post": { + "operationId": "fork_live_session_v1_live_sessions__session_id__fork_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Fork Live Session", + "tags": [ + "live" + ] + } + }, + "/v1/live/sessions/{session_id}/hangup": { + "post": { + "operationId": "control_live_session_v1_live_sessions__session_id__hangup_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/v1/live/sessions/{session_id}/refer": { + "post": { + "operationId": "control_live_session_v1_live_sessions__session_id__refer_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + }, + "/v1/live/sessions/{session_id}/reject": { + "post": { + "operationId": "control_live_session_v1_live_sessions__session_id__reject_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "live" + ] + } + } + } + }, "llm_passthrough": { "components": { "schemas": { @@ -27695,6 +28533,284 @@ ] } }, + "/openai/v1/live": { + "post": { + "operationId": "proxy_live_calls_openai_v1_live_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Proxy Live Calls", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/live/sessions": { + "post": { + "operationId": "create_live_session_openai_v1_live_sessions_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Create Live Session", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/accept": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__accept_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/content": { + "get": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__content_get", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/fork": { + "post": { + "operationId": "fork_live_session_openai_v1_live_sessions__session_id__fork_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Fork Live Session", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/hangup": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__hangup_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/refer": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__refer_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/reject": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__reject_post", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "llm_passthrough" + ] + } + }, "/openai/v1/realtime/calls": { "post": { "operationId": "proxy_realtime_calls_openai_v1_realtime_calls_post", @@ -47472,6 +48588,19 @@ "realtime": { "components": { "schemas": { + "HTTPValidationError": { + "properties": { + "detail": { + "items": { + "$ref": "#/components/schemas/ValidationError" + }, + "title": "Detail", + "type": "array" + } + }, + "title": "HTTPValidationError", + "type": "object" + }, "RealtimeClientSecretResponse": { "description": "Response from POST /v1/realtime/client_secrets.\n\nBoth the top-level `value` and `session.client_secret.value`\nwill contain the encrypted token instead of the raw ephemeral key.\nThe `session` field is kept as a raw dict so unknown fields pass through.", "properties": { @@ -47528,10 +48657,606 @@ }, "title": "RealtimeTranscriptionSessionResponse", "type": "object" + }, + "ValidationError": { + "properties": { + "ctx": { + "title": "Context", + "type": "object" + }, + "input": { + "title": "Input" + }, + "loc": { + "items": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "integer" + } + ] + }, + "title": "Location", + "type": "array" + }, + "msg": { + "title": "Message", + "type": "string" + }, + "type": { + "title": "Error Type", + "type": "string" + } + }, + "required": [ + "loc", + "msg", + "type" + ], + "title": "ValidationError", + "type": "object" } } }, "paths": { + "/live": { + "post": { + "operationId": "proxy_live_calls_live_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Proxy Live Calls", + "tags": [ + "realtime" + ] + } + }, + "/live/sessions": { + "post": { + "operationId": "create_live_session_live_sessions_post_2", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Create Live Session", + "tags": [ + "realtime" + ] + } + }, + "/live/sessions/{session_id}/accept": { + "post": { + "operationId": "control_live_session_live_sessions__session_id__accept_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/live/sessions/{session_id}/content": { + "get": { + "operationId": "control_live_session_live_sessions__session_id__content_get_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/live/sessions/{session_id}/fork": { + "post": { + "operationId": "fork_live_session_live_sessions__session_id__fork_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Fork Live Session", + "tags": [ + "realtime" + ] + } + }, + "/live/sessions/{session_id}/hangup": { + "post": { + "operationId": "control_live_session_live_sessions__session_id__hangup_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/live/sessions/{session_id}/refer": { + "post": { + "operationId": "control_live_session_live_sessions__session_id__refer_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/live/sessions/{session_id}/reject": { + "post": { + "operationId": "control_live_session_live_sessions__session_id__reject_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/openai/v1/live": { + "post": { + "operationId": "proxy_live_calls_openai_v1_live_post_2", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Proxy Live Calls", + "tags": [ + "realtime" + ] + } + }, + "/openai/v1/live/sessions": { + "post": { + "operationId": "create_live_session_openai_v1_live_sessions_post_3", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Create Live Session", + "tags": [ + "realtime" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/accept": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__accept_post_3", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/content": { + "get": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__content_get_3", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/fork": { + "post": { + "operationId": "fork_live_session_openai_v1_live_sessions__session_id__fork_post_3", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Fork Live Session", + "tags": [ + "realtime" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/hangup": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__hangup_post_3", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/refer": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__refer_post_3", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/openai/v1/live/sessions/{session_id}/reject": { + "post": { + "operationId": "control_live_session_openai_v1_live_sessions__session_id__reject_post_3", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, "/openai/v1/realtime/calls": { "post": { "operationId": "proxy_realtime_calls_openai_v1_realtime_calls_post_2", @@ -47676,6 +49401,284 @@ ] } }, + "/v1/live": { + "post": { + "operationId": "proxy_live_calls_v1_live_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Proxy Live Calls", + "tags": [ + "realtime" + ] + } + }, + "/v1/live/sessions": { + "post": { + "operationId": "create_live_session_v1_live_sessions_post_2", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Create Live Session", + "tags": [ + "realtime" + ] + } + }, + "/v1/live/sessions/{session_id}/accept": { + "post": { + "operationId": "control_live_session_v1_live_sessions__session_id__accept_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/v1/live/sessions/{session_id}/content": { + "get": { + "operationId": "control_live_session_v1_live_sessions__session_id__content_get_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/v1/live/sessions/{session_id}/fork": { + "post": { + "operationId": "fork_live_session_v1_live_sessions__session_id__fork_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Fork Live Session", + "tags": [ + "realtime" + ] + } + }, + "/v1/live/sessions/{session_id}/hangup": { + "post": { + "operationId": "control_live_session_v1_live_sessions__session_id__hangup_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/v1/live/sessions/{session_id}/refer": { + "post": { + "operationId": "control_live_session_v1_live_sessions__session_id__refer_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, + "/v1/live/sessions/{session_id}/reject": { + "post": { + "operationId": "control_live_session_v1_live_sessions__session_id__reject_post_2", + "parameters": [ + { + "in": "path", + "name": "session_id", + "required": true, + "schema": { + "title": "Session Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Control Live Session", + "tags": [ + "realtime" + ] + } + }, "/v1/realtime/calls": { "post": { "operationId": "proxy_realtime_calls_v1_realtime_calls_post", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7b1ba2ec1ac..3923f838128 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -114,6 +114,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. @@ -271,8 +275,8 @@ class Litellm_EntityType(enum.Enum): def hash_token(token: str): import hashlib - # Hash the string using SHA-256 - hashed_token: Final = hashlib.sha256(token.encode()).hexdigest() + # This digest is an opaque lookup identifier, not a password hash. + hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token @@ -422,6 +426,36 @@ class LiteLLMRoutes(enum.Enum): "/realtime?{model}", "/v1/realtime?{model}", "/openai/v1/realtime?{model}", + "/live", + "/v1/live", + "/v1/live/{call_id}", + "/openai/v1/live", + "/live/{call_id}", + "/openai/v1/live/{call_id}", + "/live/sessions", + "/live/sessions/{session_id}/attach", + "/live/sessions/{session_id}/fork", + "/live/sessions/{session_id}/content", + "/live/sessions/{session_id}/accept", + "/live/sessions/{session_id}/reject", + "/live/sessions/{session_id}/refer", + "/live/sessions/{session_id}/hangup", + "/v1/live/sessions", + "/v1/live/sessions/{session_id}/attach", + "/v1/live/sessions/{session_id}/fork", + "/v1/live/sessions/{session_id}/content", + "/v1/live/sessions/{session_id}/accept", + "/v1/live/sessions/{session_id}/reject", + "/v1/live/sessions/{session_id}/refer", + "/v1/live/sessions/{session_id}/hangup", + "/openai/v1/live/sessions", + "/openai/v1/live/sessions/{session_id}/attach", + "/openai/v1/live/sessions/{session_id}/fork", + "/openai/v1/live/sessions/{session_id}/content", + "/openai/v1/live/sessions/{session_id}/accept", + "/openai/v1/live/sessions/{session_id}/reject", + "/openai/v1/live/sessions/{session_id}/refer", + "/openai/v1/live/sessions/{session_id}/hangup", # realtime (GA WebRTC HTTP routes) "/realtime/client_secrets", "/v1/realtime/client_secrets", diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3ec430332ee..f0e2bbcde83 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2398,11 +2398,18 @@ async def get_team_membership( user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None = None, proxy_logging_obj: ProxyLogging | None = None, + raise_on_error: bool = True, ) -> Optional["LiteLLM_TeamMembership"]: """ Returns team membership object if user is member of team. Do a isolated check for team membership vs. doing a combined key + team + user + team-membership check, as key might come in frequently for different users/teams. Larger call will slowdown query time. This way we get to cache the constant (key/team/user info) and only update based on the changing value (team membership). + + ``raise_on_error`` defaults to True because the callers that apply member-level limits -- the budget and + model-scope checks in ``common_checks``, the JWT team resolution, and the compact summary gate -- cannot + tell an absent row apart from a failed read, so swallowing an outage there hands the member whatever the + team allows. A caller that only attributes grants, and can proceed with the lists it already holds, + passes False and degrades to "no member-level scope". """ if user_id is None or team_id is None: return None @@ -2416,7 +2423,17 @@ async def get_team_membership( inflight: Final[object] = _team_membership_inflight.get(_key) if isinstance(inflight, asyncio.Task): - return _membership_from_shared_load(await asyncio.shield(inflight)) + try: + return _membership_from_shared_load(await asyncio.shield(inflight)) + except Exception: + verbose_proxy_logger.exception( + "Error getting team membership for user_id: %s, team_id: %s", + user_id, + team_id, + ) + if raise_on_error: + raise + return None if prisma_client is None: raise Exception("No db connected") @@ -2439,7 +2456,17 @@ async def get_team_membership( _team_membership_inflight.pop(_key, None) task.add_done_callback(_clear_inflight) - return _membership_from_shared_load(await asyncio.shield(task)) + try: + return _membership_from_shared_load(await asyncio.shield(task)) + except Exception: + verbose_proxy_logger.exception( + "Error getting team membership for user_id: %s, team_id: %s", + user_id, + team_id, + ) + if raise_on_error: + raise + return None def model_in_access_group(model: str, team_models: list[str] | None, llm_router: Router | None) -> bool: @@ -4702,6 +4729,8 @@ async def _team_member_granted_models( prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, + *, + strict_grant_lookup: bool = False, team_membership: LiteLLM_TeamMembership | None = None, team_membership_loaded: bool = False, ) -> Sequence[str]: @@ -4716,6 +4745,10 @@ async def _team_member_granted_models( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + # Spelled out because it is the one caller that wants the opposite of the default: outside + # strict mode this walk only attributes grants, so an unreadable member scope degrades to + # "no member-level scope" instead of failing the request. + raise_on_error=strict_grant_lookup, ) return () if team_membership is None else _member_allowed_models(team_membership) @@ -4726,6 +4759,8 @@ async def _org_granted_models( prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, + *, + strict_grant_lookup: bool = False, ) -> Sequence[str]: """The org allowlist reached through the key, or through its team when the key names no org.""" org_id: Final = valid_token.org_id or (team_object.organization_id if team_object is not None else None) @@ -4741,6 +4776,8 @@ async def _org_granted_models( ) except Exception as e: # noqa: BLE001 # fail-safe: attribution degrades to "no org grant", it must never break auth verbose_proxy_logger.debug("access group attribution: org lookup failed: %s", e) + if strict_grant_lookup: + raise return () return org_object.models if org_object is not None else () @@ -4752,6 +4789,8 @@ async def _granted_model_lists( prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, + *, + strict_grant_lookup: bool = False, team_membership: LiteLLM_TeamMembership | None = None, team_membership_loaded: bool = False, ) -> tuple[Sequence[str], ...]: @@ -4765,6 +4804,7 @@ async def _granted_model_lists( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + strict_grant_lookup=strict_grant_lookup, team_membership=team_membership, team_membership_loaded=team_membership_loaded, ), @@ -4775,6 +4815,7 @@ async def _granted_model_lists( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + strict_grant_lookup=strict_grant_lookup, ), ) @@ -4861,6 +4902,8 @@ async def collect_matched_model_access_groups( prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, + *, + strict_grant_lookup: bool = False, team_membership: LiteLLM_TeamMembership | None = None, team_membership_loaded: bool = False, ) -> tuple[str, ...]: @@ -4878,7 +4921,9 @@ async def collect_matched_model_access_groups( The whole walk is gated on the budget registry, because collecting every match costs a full scan of each allowlist where the plain access check stops at the first hit. An empty registry means no - group carries a budget, so there is nothing to attribute and no work worth doing. + group carries a budget, so there is nothing to attribute and no work worth doing. The strict + lookup mode is reserved for enforcement paths that must not treat an unavailable inherited grant + as absent; the default remains fail-safe attribution for ordinary request telemetry. """ if model is None or valid_token is None or llm_router is None or prisma_client is None: return () @@ -4908,6 +4953,7 @@ async def collect_matched_model_access_groups( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + strict_grant_lookup=strict_grant_lookup, team_membership=team_membership, team_membership_loaded=team_membership_loaded, ) @@ -6537,7 +6583,9 @@ def is_model_allowed_by_pattern(model: str, allowed_model_pattern: str) -> bool: bool: True if model matches the pattern, False otherwise """ if "*" in allowed_model_pattern: - pattern: Final = f"^{allowed_model_pattern.replace('*', '.*')}$" + # Treat the configured model pattern as a glob; only '*' is special. + escaped_pattern: Final = re.escape(allowed_model_pattern) + pattern: Final = "^" + escaped_pattern.replace("\\*", ".*") + "$" return bool(re.match(pattern, model)) return False diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c0123ae45a3..f026241d017 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1710,6 +1710,7 @@ _MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS: Final = ("/evals",) _MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS: Final = ( "/realtime/client_secrets", "/realtime/calls", + "/live/sessions", ) _MODEL_ROUTING_ID_FIELDS: Final = ( "file_id", @@ -1900,15 +1901,27 @@ def _extract_model_candidates_from_request( uses_completion_model_sources: Final = _route_matches_any_marker( route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS ) + uses_session_model: Final = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS + ) or route.rstrip("/") in ("/live", "/v1/live", "/openai/v1/live") + session: Final[object] = request_data.get("session") if uses_session_model else None + parsed_session: Final[object] = safe_json_loads(session) if isinstance(session, str) else session + session_model: Final[object] = parsed_session.get("model") if isinstance(parsed_session, dict) else None + if ( + uses_session_model + and not _route_matches_any_marker(route=route, markers=("/realtime/client_secrets",)) + and isinstance(session_model, str) + and session_model + ): + candidates.append(session_model) + return candidates body_model: Final = request_data.get("model") _append_model_candidates(candidates, body_model) if uses_body_target_model_sources or not body_model: _append_model_candidates(candidates, request_data.get("target_model_names")) - if _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS): - session: Final = request_data.get("session") - if isinstance(session, dict): - _append_model_candidates(candidates, session.get("model")) + if uses_session_model: + _append_model_candidates(candidates, TypeAdapter[object](object).validate_python(session_model)) if uses_completion_model_sources and isinstance(request_data.get("completion"), dict): _append_model_candidates(candidates, request_data["completion"].get("model")) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 4448d860217..ce39e3a6a20 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -2485,7 +2485,7 @@ class JWTAuthManager: return JWTIdentity(user_id=user_id if is_admin else canonical_id, user_object=user, agent_id=agent_id) @staticmethod - async def authorize_jwt( + async def authorize_jwt( # noqa: C901 # preserves the established JWT authorization flow split from auth_builder api_key: str, jwt_handler: JWTHandler, request_data: dict[str, object], diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 5194f62cf78..a3b77e6a32f 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -452,7 +452,7 @@ async def _check_key_model_budget_with_fallback( model=model_name, ) except litellm.BudgetExceededError as e: - if request_data.get("model") != model_name: + if request_data.get("model") != model_name or request.scope.get("litellm_pinned_realtime_model") == model_name: raise e fallback_model: Final = await model_max_budget_limiter.get_fallback_model_within_budget( user_api_key_dict=valid_token, @@ -641,6 +641,38 @@ def _apply_budget_limits_to_end_user_params( verbose_proxy_logger.debug("Applied budget limits to end user %s", end_user_id) +def get_websocket_api_key(websocket: WebSocket) -> str | None: + from litellm.proxy.proxy_server import general_settings + + custom_header: Final = general_settings.get("litellm_key_header_name") + if isinstance(custom_header, str): + if not websocket.headers.get(custom_header): + return None + request: Final = Request( + {"type": "http", "headers": websocket.scope.get("headers", [])} # mutable-ok: ASGI request scope + ) + return get_api_key_from_custom_header(request, custom_header) + custom_key: Final = websocket.headers.get("x-litellm-api-key") + if custom_key is not None: + return _get_bearer_token_or_received_api_key(custom_key) + authorization: Final = websocket.headers.get("authorization") + if authorization: + if not authorization.startswith("Bearer "): + raise HTTPException(status_code=403, detail="Invalid Authorization header format") + return authorization[len("Bearer ") :].strip() + api_key: Final = websocket.headers.get("api-key") + if api_key: + return api_key + return next( + ( + protocol.strip().removeprefix("openai-insecure-api-key.") + for protocol in websocket.headers.get("sec-websocket-protocol", "").split(",") + if protocol.strip().startswith("openai-insecure-api-key.") + ), + None, + ) + + async def user_api_key_auth_websocket(websocket: WebSocket) -> UserAPIKeyAuth: return await user_api_key_auth_websocket_for_model(websocket, model=websocket.query_params.get("model")) @@ -668,36 +700,29 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str request._url = websocket.url - async def return_body(): - return _realtime_request_body(model) - - request.body = return_body - - authorization: Final = websocket.headers.get("authorization") - # If no Authorization header, try the api-key header - if not authorization: - api_key = websocket.headers.get("api-key") - if not api_key: - # Try extracting from WebSocket subprotocol (browser clients) - for protocol in websocket.headers.get("sec-websocket-protocol", "").split(","): - protocol = protocol.strip() - if protocol.startswith("openai-insecure-api-key."): - api_key = protocol[len("openai-insecure-api-key.") :] - break - if not api_key: - await websocket.close(code=status.WS_1008_POLICY_VIOLATION) - raise HTTPException(status_code=403, detail="No API key provided") - else: - # Extract the API key from the Bearer token - if not authorization.startswith("Bearer "): - await websocket.close(code=status.WS_1008_POLICY_VIOLATION) - raise HTTPException(status_code=403, detail="Invalid Authorization header format") - - api_key = authorization[len("Bearer ") :].strip() + try: + api_key: Final = get_websocket_api_key(websocket) + except HTTPException: + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + raise + if not api_key: + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + raise HTTPException(status_code=403, detail="No API key provided") # Call user_api_key_auth with the extracted API key # Note: You'll need to modify this to work with WebSocket context if needed try: + from litellm.proxy.realtime_endpoints.call_sessions import decode_call + + call_token: Final = websocket.path_params.get("call_id") or websocket.query_params.get("call_id") + resolved_model: Final = decode_call(call_token, f"Bearer {api_key}").alias if call_token is not None else model + if call_token is not None: + request.scope["litellm_pinned_realtime_model"] = resolved_model + + async def return_body(): + return _realtime_request_body(resolved_model) + + request.body = return_body return await user_api_key_auth(request=request, api_key=f"Bearer {api_key}") except Exception as e: if is_invalid_virtual_key_error(e): diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e7a08711eb1..6184ac3250a 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2015,7 +2015,9 @@ class ProxyBaseLLMRequestProcessing: model: str | None = None, llm_router: Router | None = None, rate_limited_model: str | None = None, + *, skip_guardrails: bool = False, + internal_realtime_observer: bool = False, ) -> tuple[dict, LiteLLMLoggingObj]: start_time: Final = datetime.now() # start before calling guardrail hooks @@ -2209,6 +2211,11 @@ class ProxyBaseLLMRequestProcessing: data=self.data, call_type=route_type, skip_guardrails=skip_guardrails, + **( + MappingProxyType({"internal_realtime_observer": True}) + if internal_realtime_observer + else MappingProxyType({}) + ), ) await _enforce_guardrail_added_tag_budgets( data=self.data, diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index c99665986dd..a661a09f033 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -336,6 +336,17 @@ def model_access_group_cache_key(access_group_name: str) -> str: return f"model_access_group:{access_group_name}" +def live_model_access_group_limits_cache_key(access_group_name: str) -> str: + """Cache key the Live delegation gate stores one access group's full limit row under. + + The gate needs the rpm and tpm columns that ``model_access_group:{name}`` flattens away, so it + keeps its own entry next to the flattened one. Any eviction of the flattened entry must clear + this key too: the gate reads cache-first, and a raised or lowered group limit left cached here + keeps permitting or refusing managed delegation until the entry's TTL expires (LIT-3803). + """ + return f"live:model_access_group_limits:{access_group_name}" + + def model_access_group_registry_cache_key() -> str: """Cache key for the set of model access group names that have a budget row.""" return "model_access_group_registry" diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index e3485ebf25d..ddf571b6098 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -1,9 +1,10 @@ import asyncio import sys +from collections.abc import Mapping from datetime import datetime, timedelta from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter from typing_extensions import TypedDict import litellm @@ -12,7 +13,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, @@ -51,11 +52,81 @@ class CacheObject(TypedDict): request_count_end_user_id: dict | None +class _RealtimeAttachmentReservations(BaseModel): + cache_keys: tuple[str, ...] = () + global_acquired: bool = False + + def acquire(self, key: str) -> None: + self.cache_keys = tuple(dict.fromkeys((*self.cache_keys, key))) + + def acquire_global(self) -> None: + self.global_acquired = True + + def take(self) -> tuple[tuple[str, ...], bool]: + owned: Final = (self.cache_keys, self.global_acquired) + self.cache_keys = () + self.global_acquired = False + return owned + + +_RELEASE_REALTIME_COUNTER_LUA: Final = """ +local raw = redis.call('GET', KEYS[1]) +if not raw then return 0 end +local value = cjson.decode(raw) +value.current_requests = math.max(value.current_requests - 1, 0) +redis.call('SET', KEYS[1], cjson.encode(value), 'KEEPTTL') +return 1 +""" + + class _PROXY_MaxParallelRequestsHandler(CustomLogger): # Class variables or attributes def __init__(self, internal_usage_cache: InternalUsageCache): self.internal_usage_cache = internal_usage_cache + def begin_realtime_attachment(self, request_data: dict[str, object]) -> None: + request_data["_legacy_realtime_attachment_reservations"] = ( # rebind-ok: request-scoped cleanup receipt + _RealtimeAttachmentReservations() + ) + + async def async_release_realtime_attachment( + self, request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth + ) -> None: + receipt: Final = request_data.get("_legacy_realtime_attachment_reservations") + if not isinstance(receipt, _RealtimeAttachmentReservations): + return + keys, global_acquired = receipt.take() + if global_acquired: + await self.internal_usage_cache.async_increment_cache( + key="global_max_parallel_requests", + value=-1, + local_only=True, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + for key in keys: + await self._release_realtime_counter(key) + + async def _release_realtime_counter(self, key: str) -> None: + local: Final = self.internal_usage_cache.dual_cache.in_memory_cache + remote: Final = self.internal_usage_cache.dual_cache.redis_cache + raw: Final[object] = local.get_cache(key) + current: Final = TypeAdapter[Mapping[str, int] | None](Mapping[str, int] | None).validate_python(raw) + updated: Final = ( + { # mutable-ok: shared cache counter dict + **current, + "current_requests": max(current["current_requests"] - 1, 0), + } + if current is not None + else None + ) + if updated is not None: + local.set_cache(key, updated, ttl=60) + if remote is not None: + release: Final = remote.async_register_script(_RELEASE_REALTIME_COUNTER_LUA) + await release(keys=(key,), args=()) + if local.get_cache(key) is updated: + local.delete_cache(key) + def print_verbose(self, print_statement): try: verbose_proxy_logger.debug(print_statement) @@ -143,6 +214,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): litellm_parent_otel_span=user_api_key_dict.parent_otel_span, local_only=True, ) + receipt: Final = data.get("_legacy_realtime_attachment_reservations") + if isinstance(receipt, _RealtimeAttachmentReservations): + receipt.acquire(request_count_api_key) return new_val def time_to_next_minute(self) -> float: @@ -300,6 +374,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): local_only=True, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, ) + receipt: Final = data.get("_legacy_realtime_attachment_reservations") + if isinstance(receipt, _RealtimeAttachmentReservations): + receipt.acquire_global() _model = data.get("model", None) current_date: Final = datetime.now().strftime("%Y-%m-%d") @@ -481,6 +558,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): values_to_update_in_cache=values_to_update_in_cache, ) + if isinstance(data.get("_legacy_realtime_attachment_reservations"), _RealtimeAttachmentReservations): + await self.internal_usage_cache.async_batch_set_cache( + cache_list=values_to_update_in_cache, + ttl=60, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + return asyncio.create_task( self.internal_usage_cache.async_batch_set_cache( cache_list=values_to_update_in_cache, @@ -490,6 +574,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, ) @@ -522,7 +607,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 @@ -553,13 +638,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"], } @@ -594,13 +679,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"], } @@ -620,13 +705,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"], } @@ -646,13 +731,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"], } @@ -672,13 +757,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"], } @@ -695,6 +780,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) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 36896bd6d44..5ab3cadb66f 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -10,8 +10,8 @@ import itertools import logging import os import uuid -from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence, Set -from contextlib import asynccontextmanager +from collections.abc import AsyncGenerator, Awaitable, Callable, Generator, Mapping, Sequence, Set +from contextlib import asynccontextmanager, contextmanager from contextvars import ContextVar from dataclasses import dataclass, field from datetime import datetime, timezone @@ -68,6 +68,7 @@ from litellm.proxy.hooks.batch_enqueued_tokens import ( canonical_provider_batch_id, ) from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease, is_realtime_call_attachment from litellm.router_utils.add_retry_fallback_headers import ( ensure_response_additional_headers, response_has_hidden_params, @@ -396,6 +397,23 @@ end return results """ +PARALLEL_RENEW_SCRIPT: Final = """ +local clock = redis.call('TIME') +local now = tonumber(clock[1]) +local ttl = tonumber(ARGV[2]) +for i = 1, #KEYS do + local score = redis.call('ZSCORE', KEYS[i], ARGV[1]) + if not score or tonumber(score) <= now - ttl then + return {0} + end +end +for i = 1, #KEYS do + redis.call('ZADD', KEYS[i], 'XX', now, ARGV[1]) + redis.call('EXPIRE', KEYS[i], ttl) +end +return {1} +""" + TOKEN_INCREMENT_SCRIPT: Final = """ local results = {} @@ -523,9 +541,17 @@ class ParallelRequestGauge(TypedDict): descriptor_key: str +def _without_parallel_limit(descriptor: RateLimitDescriptor) -> RateLimitDescriptor: + rate_limit: Final[RateLimitDescriptorRateLimitObject] = { + **(descriptor.get("rate_limit") or MappingProxyType({})), + "max_parallel_requests": None, + } + return RateLimitDescriptor(key=descriptor["key"], value=descriptor["value"], rate_limit=rate_limit) + + class ParallelSlotAcquisition(TypedDict): slot_id: str - counter_keys: list[str] + counter_keys: Sequence[str] class RateLimitStatus(TypedDict): @@ -706,6 +732,15 @@ def get_request_stash() -> RequestRateLimiterStash | None: return _request_stash.get() +@contextmanager +def isolated_request_stash() -> Generator[None]: + token: Final = _request_stash.set(None) + try: + yield + finally: + _request_stash.reset(token) + + def get_or_create_request_stash() -> RequestRateLimiterStash: stash = _request_stash.get() if stash is None: @@ -756,6 +791,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): parallel_acquire_script: _AsyncLuaScript | None parallel_release_script: _AsyncLuaScript | None parallel_count_script: _AsyncLuaScript | None + parallel_renew_script: _AsyncLuaScript | None def __init__( self, @@ -797,6 +833,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.parallel_count_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( PARALLEL_COUNT_SCRIPT ) + self.parallel_renew_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( + PARALLEL_RENEW_SCRIPT + ) else: self.batch_rate_limiter_script = None self.batch_counter_read_script = None @@ -806,6 +845,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.parallel_acquire_script = None self.parallel_release_script = None self.parallel_count_script = None + self.parallel_renew_script = None self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) @@ -1168,7 +1208,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def in_memory_cache_sliding_window( self, - keys: list[str], + keys: Sequence[str], now_int: int, window_size: int, ) -> CacheCounterValues: @@ -1357,7 +1397,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return (None,) * len(calls) return tuple(batch.script(source, run, keys, args) for keys, args in calls) - def _group_keys_by_hash_tag(self, keys: list[str]) -> dict[str, list[str]]: + def _group_keys_by_hash_tag(self, keys: Sequence[str]) -> Mapping[str, Sequence[str]]: """ Group keys by their Redis hash tag to ensure cluster compatibility. @@ -1377,7 +1417,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): groups[slot_key].append(key) else: # For regular Redis, no grouping needed - process all keys together - groups[REDIS_NODE_HASHTAG_NAME] = keys + return MappingProxyType({REDIS_NODE_HASHTAG_NAME: keys}) return groups @@ -1753,12 +1793,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ gauge_keys: Final = [gauge["counter_key"] for gauge in gauges] + if self._is_redis_cluster() and self.parallel_acquire_script is not None: + return await self._check_cluster_parallel_gauges(gauges, slot_id, parent_otel_span, read_only) + if read_only: if self.parallel_count_script is not None: try: raw_counts: Final[list[CacheCounterValue]] = await self.parallel_count_script( keys=gauge_keys, - args=[PARALLEL_REQUEST_SLOT_TTL_SECONDS for _ in gauges], + args=tuple(PARALLEL_REQUEST_SLOT_TTL_SECONDS for _ in gauges), ) counts = [max(0, int(value)) for value in raw_counts] except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the local mirror unless fail-closed rejects @@ -1825,6 +1868,129 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async with self._check_and_increment_lock: return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span) + async def _check_cluster_parallel_gauges( + self, + gauges: Sequence[ParallelRequestGauge], + slot_id: str, + parent_otel_span: Span | None, + read_only: bool, + ) -> RateLimitResponse: + by_key: Final = MappingProxyType( + { + gauge["counter_key"]: min( + (candidate for candidate in gauges if candidate["counter_key"] == gauge["counter_key"]), + key=lambda candidate: candidate["limit"], + ) + for gauge in gauges + } + ) + groups: Final = self._group_keys_by_hash_tag(tuple(by_key)) + counts: Final[dict[str, int]] = {} # mutable-ok: gather independent Redis-slot results + attempted: Final[list[str]] = [] # mutable-ok: rollback includes requests whose responses were lost + try: + for keys in groups.values(): + if read_only: + if self.parallel_count_script is None: + raise RuntimeError("Redis cluster parallel count script is unavailable") + counts.update( + (key, max(0, int(count))) + for key, count in zip( + keys, + await self.parallel_count_script( + keys=keys, args=tuple(PARALLEL_REQUEST_SLOT_TTL_SECONDS for _ in keys) + ), + strict=True, + ) + ) + continue + if self.parallel_acquire_script is None: + raise RuntimeError("Redis cluster parallel acquire script is unavailable") + attempted.extend(keys) + acquire_args: list[object] = [] # mutable-ok: Redis EVAL args are flattened per slot below + for key in keys: + acquire_args.extend((by_key[key]["limit"], PARALLEL_REQUEST_SLOT_TTL_SECONDS, slot_id)) + (raw,) = (await self.parallel_acquire_script(keys=keys, args=tuple(acquire_args)),) + if int(raw[0]) == 1: + await self._rollback_cluster_parallel_slots(tuple(attempted), slot_id, parent_otel_span) + return RateLimitResponse( + overall_code="OVER_LIMIT", + statuses=[ # mutable-ok: response contract requires a list + self._gauge_status(by_key[keys[int(raw[1]) - 1]], int(raw[2]), "OVER_LIMIT") + ], + ) + counts.update((key, int(count)) for key, count in zip(keys, raw[1:], strict=True)) + for key in keys: + await self.internal_usage_cache.async_set_cache( + key=key, + value=counts[key], + ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + except BaseException: + if attempted: + await self._rollback_cluster_parallel_slots(tuple(attempted), slot_id, parent_otel_span) + raise + statuses: Final = tuple( + self._gauge_status( + gauge, + counts[gauge["counter_key"]], + "OVER_LIMIT" if read_only and counts[gauge["counter_key"]] >= gauge["limit"] else "OK", + ) + for gauge in gauges + ) + return RateLimitResponse( + overall_code="OVER_LIMIT" if any(item["code"] == "OVER_LIMIT" for item in statuses) else "OK", + statuses=list(statuses), # mutable-ok: RateLimitResponse contract requires a list + ) + + async def _rollback_cluster_parallel_slots( + self, counter_keys: tuple[str, ...], slot_id: str, parent_otel_span: Span | None + ) -> None: + rollback: Final = asyncio.create_task( + self._release_cluster_parallel_slots(counter_keys, slot_id, parent_otel_span) + ) + cancelled = False # rebind-ok: defer repeated caller cancellation until compensation finishes + while not rollback.done(): + try: + await asyncio.shield(rollback) + except asyncio.CancelledError: + cancelled = True + except Exception: # noqa: BLE001 # retrieve and report the completed task's exception below + break + try: + rollback.result() + except Exception: # noqa: BLE001 # preserve admission failure; unreachable Redis slots expire by TTL + verbose_proxy_logger.error("Could not roll back all Redis cluster parallel request slots") + if cancelled: + raise asyncio.CancelledError + + async def _release_cluster_parallel_slots( + self, counter_keys: tuple[str, ...], slot_id: str, parent_otel_span: Span | None + ) -> None: + first_error: Exception | None = None # rebind-ok: finish every shard before reporting the first failure + for keys in self._group_keys_by_hash_tag(counter_keys).values(): + try: + if self.parallel_release_script is None: + raise RuntimeError("Redis cluster parallel release script is unavailable") + for key, count in zip( + keys, + await self.parallel_release_script(keys=keys, args=tuple(slot_id for _ in keys)), + strict=True, + ): + await self.internal_usage_cache.async_set_cache( + key=key, + value=max(0, int(count)), + ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + except Exception as exc: # noqa: BLE001 # one unreachable shard must not strand the other shards + if first_error is None: + first_error = exc + if first_error is not None: + raise first_error + async def _read_local_gauge_counts( self, gauge_keys: list[str], @@ -1894,6 +2060,68 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): statuses.append(self._gauge_status(gauge, in_flight + 1, "OK")) return RateLimitResponse(overall_code="OK", statuses=statuses) + def transfer_realtime_call_slot(self, request_data: Mapping[str, object]) -> RealtimeCallLease | None: + call_id: Final = request_data.get("litellm_call_id") + if not isinstance(call_id, str): + return None + stash: Final = get_request_stash_for_call(call_id) + if stash is None or stash.parallel_slot is None: + return None + slot_id: Final = stash.parallel_slot["slot_id"] + counter_keys: Final = tuple(stash.parallel_slot["counter_keys"]) + stash.parallel_slot = None + + async def renew() -> bool: + return await self._renew_realtime_call_slot(slot_id, counter_keys) + + async def release() -> None: + await self._release_parallel_request_slots( + ParallelSlotAcquisition(slot_id=slot_id, counter_keys=counter_keys) + ) + + return RealtimeCallLease(renew=renew, release=release) + + async def _renew_realtime_call_slot(self, slot_id: str, counter_keys: tuple[str, ...]) -> bool: + if self.parallel_renew_script is not None: + try: + for keys in self._group_keys_by_hash_tag(counter_keys).values(): + if tuple( + await self.parallel_renew_script(keys=keys, args=(slot_id, PARALLEL_REQUEST_SLOT_TTL_SECONDS)) + ) != (1,): + return False + return True + except Exception: # noqa: BLE001 # Redis ownership cannot be established by a local count mirror + return False + async with self._check_and_increment_lock: + now: Final = self._get_current_time().timestamp() + cutoff: Final = now - PARALLEL_REQUEST_SLOT_TTL_SECONDS + values: Final[tuple[ParallelGaugeCacheValue | None, ...]] = tuple( + [ + await self.internal_usage_cache.async_get_cache( + key=counter_key, local_only=True, litellm_parent_otel_span=None + ) + for counter_key in counter_keys + ] + ) + if any( + not isinstance(value, dict) + or not isinstance(score := value.get(slot_id), (int, float)) + or score <= cutoff + for value in values + ): + return False + for counter_key, value in zip(counter_keys, values): + if not isinstance(value, dict): + return False + await self.internal_usage_cache.async_set_cache( + key=counter_key, + value={**value, slot_id: now}, # mutable-ok: slot registry readers require a concrete dict + ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, + local_only=True, + litellm_parent_otel_span=None, + ) + return True + async def _release_stashed_parallel_slot( self, stash: RequestRateLimiterStash | None, @@ -1931,11 +2159,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): slot_id: Final = acquisition["slot_id"] if not counter_keys or not slot_id: return + if self._is_redis_cluster() and self.parallel_release_script is not None: + await self._release_cluster_parallel_slots(tuple(counter_keys), slot_id, parent_otel_span) + return if self.parallel_release_script is not None: try: raw: Final[list[CacheCounterValue]] = await self.parallel_release_script( keys=counter_keys, - args=[slot_id for _ in counter_keys], + args=tuple(slot_id for _ in counter_keys), ) await self._mirror_released_parallel_slots(counter_keys, raw, parent_otel_span) return @@ -4020,6 +4251,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if d["key"] == "tag" and d["rate_limit"] is not None and d["rate_limit"].get("tokens_per_unit") is not None ) + effective_descriptors: Final = ( + tuple(_without_parallel_limit(descriptor) for descriptor in descriptors) + if call_type == "_arealtime" + and is_realtime_call_attachment(TypeAdapter[object](object).validate_python(data.get("websocket"))) + else descriptors + ) + # Only check rate limits if we have descriptors with actual limits if descriptors: # First pass: RPM and max_parallel_requests sliding-window check. @@ -4038,16 +4276,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # double-charge every request. parallel_counter_keys: Final = [ self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests") - for d in descriptors + for d in effective_descriptors if (d.get("rate_limit") or {}).get("max_parallel_requests") is not None ] parallel_slot_id: Final = uuid.uuid4().hex if parallel_counter_keys else None first_pass_descriptors: Final = ( - descriptors + effective_descriptors if self.tpm_reservation_enabled else tuple( - d for d in descriptors if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) + d + for d in effective_descriptors + if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) ) ) response: Final = await self.should_rate_limit( @@ -5242,6 +5482,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ await self._release_stashed_parallel_slot(get_request_stash(), None) + async def async_release_realtime_attachment( + self, request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth + ) -> None: + await self.async_post_call_failure_hook( + request_data={}, # mutable-ok: existing failure hook requires dict; attachment has no billable usage + original_exception=Exception("Realtime attachment completed"), + user_api_key_dict=user_api_key_dict, + ) + async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): """ Release completed-request slots and update rate limit headers in the response. diff --git a/litellm/proxy/hooks/realtime_call_lease.py b/litellm/proxy/hooks/realtime_call_lease.py new file mode 100644 index 00000000000..c33d81f92d3 --- /dev/null +++ b/litellm/proxy/hooks/realtime_call_lease.py @@ -0,0 +1,74 @@ +import asyncio +from collections.abc import Awaitable, Callable, Generator +from contextlib import contextmanager +from contextvars import ContextVar +from typing import Final + +_realtime_call_attachment: Final[ContextVar[object | None]] = ContextVar("realtime_call_attachment", default=None) + + +@contextmanager +def realtime_call_attachment(websocket: object) -> Generator[None]: + token: Final = _realtime_call_attachment.set(websocket) + try: + yield + finally: + _realtime_call_attachment.reset(token) + + +def is_realtime_call_attachment(websocket: object) -> bool: + bound: Final = _realtime_call_attachment.get() + return bound is not None and bound is websocket + + +class RealtimeCallLease: + def __init__( + self, + *, + renew: Callable[[], Awaitable[bool]], + release: Callable[[], Awaitable[None]], + interval: float = 300, + renewal_timeout: float = 10, + ) -> None: + self._renew = renew + self._release = release + self._interval = interval + self._renewal_timeout = renewal_timeout + self._failed = asyncio.Event() + self._heartbeat: asyncio.Task[None] | None = None + self._closing: asyncio.Task[None] | None = None + + def start(self) -> None: + if self._heartbeat is None and self._closing is None: + self._heartbeat = asyncio.create_task(self._run()) + + async def renew(self) -> bool: + if self._closing is not None or self._failed.is_set(): + return False + try: + renewed: Final = await asyncio.wait_for(self._renew(), timeout=self._renewal_timeout) + except Exception: # noqa: BLE001 # fail closed without exposing cache credentials + self._failed.set() + return False + if not renewed: + self._failed.set() + return renewed and not self._failed.is_set() and self._closing is None + + async def wait_failed(self) -> None: + await self._failed.wait() + + async def _run(self) -> None: + while await self.renew(): + await asyncio.sleep(self._interval) + self._failed.set() + + async def close(self) -> None: + if self._closing is None: + self._closing = asyncio.create_task(self._close()) + await asyncio.shield(self._closing) + + async def _close(self) -> None: + if self._heartbeat is not None: + self._heartbeat.cancel() + await asyncio.gather(self._heartbeat, return_exceptions=True) + await self._release() diff --git a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py index e960bdfe337..72e251e8a31 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -26,6 +26,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, + live_model_access_group_limits_cache_key, model_access_group_cache_key, model_access_group_registry_cache_key, ) @@ -200,7 +201,14 @@ async def _evict_model_access_group_cache_keys(access_group: str, auth_cache: Us ) await evict_and_broadcast( - cache_keys=(model_access_group_cache_key(access_group), model_access_group_registry_cache_key()), + cache_keys=( + model_access_group_cache_key(access_group), + # The Live delegation gate caches the same group's full limit row next to the flattened + # entry because it needs the rpm and tpm columns; leaving that entry behind keeps the + # old limit deciding managed delegation until its TTL expires. + live_model_access_group_limits_cache_key(access_group), + model_access_group_registry_cache_key(), + ), user_api_key_cache=auth_cache, ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7c1ab0711ea..5d079189707 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1613,6 +1613,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: except Exception as e: verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e) + from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS + + await CALL_SUPERVISORS.shutdown() + await _drain_spend_event_producer_on_shutdown() # Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect @@ -12899,6 +12903,28 @@ async def _reject_realtime_session( await _release_realtime_max_parallel_slot(user_api_key_dict) +_CODEX_LIVE_AUTH_DEPENDENCY: Final = Depends(user_api_key_auth_websocket) + +reserve_lazy_slot(app, "live") +reserve_lazy_slot(app, "realtime") + + +@app.websocket("/openai/v1/live/{call_id}") +@app.websocket("/live/{call_id}") +@app.websocket("/v1/live/{call_id}") +async def codex_live_sideband_endpoint( + websocket: WebSocket, + call_id: str, + user_api_key_dict: UserAPIKeyAuth = _CODEX_LIVE_AUTH_DEPENDENCY, +) -> None: + from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband + + await codex_realtime_sideband(websocket, call_id, user_api_key_dict) + + +@app.websocket("/v1/live") +@app.websocket("/live") +@app.websocket("/openai/v1/live") @app.websocket("/openai/v1/realtime") @app.websocket("/v1/realtime") @app.websocket("/realtime") @@ -12906,12 +12932,18 @@ async def realtime_websocket_endpoint( websocket: WebSocket, model: str | None = fastapi.Query(None, description="The model to use for the websocket connection."), intent: str | None = fastapi.Query(None, description="The intent of the websocket connection."), + call_id: str | None = None, guardrails: str | None = fastapi.Query( None, description="Comma-separated list of guardrail names to apply to this request.", ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), ): + if call_id is not None: + from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband + + await codex_realtime_sideband(websocket, call_id, user_api_key_dict) + return requested_protocols: Final = [ p.strip() for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",") if p.strip() ] @@ -12943,7 +12975,12 @@ async def realtime_websocket_endpoint( await websocket.accept(**accept_kwargs) # Only use explicit parameters, not all query params - query_params: Final = cast(RealtimeQueryParams, dict(_realtime_query_params_template(model, intent))) + query_params: Final = cast( + RealtimeQueryParams, + dict( # mutable-ok: FastAPI request query params must be materialized as a dict + _realtime_query_params_template(model, intent) + ((("call_id", call_id),) if call_id is not None else ()) + ), + ) data: dict[str, object] = { "model": route_model, diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py new file mode 100644 index 00000000000..997fd0e38a6 --- /dev/null +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -0,0 +1,562 @@ +import asyncio +import base64 +import hashlib +import json +import time +from collections.abc import Awaitable, Callable, Mapping +from contextlib import AsyncExitStack, nullcontext +from contextvars import Token +from types import MappingProxyType +from typing import Final, Literal + +import httpx +from fastapi import HTTPException, Request, Response, WebSocket +from pydantic import TypeAdapter +from starlette.formparsers import MultiPartException, MultiPartParser +from starlette.types import Message, Scope + +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, + RealTimeStreaming, + realtime_attachment_cleanup, +) +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.chatgpt.codex import ( + CodexRealtimeCall, + CodexRealtimeOffer, + build_call_request, + build_sideband_request, + parse_call_response, +) +from litellm.llms.chatgpt.realtime import ( + CallAccounting, + ChatGPTRealtime, + configured_realtime_headers, + realtime_endpoint, +) +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, + get_api_key_from_custom_header, + get_websocket_api_key, + user_api_key_auth, +) +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.common_utils.http_parsing_utils import ( + _normalize_media_type, # pyright: ignore[reportPrivateUsage] # reuse the shared HTTP media-type normalization contract +) +from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # existing built-in limiter has no public alias +) +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # existing built-in limiter has no public alias + isolated_request_stash, +) +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease, realtime_call_attachment +from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + release_or_invalidate_budget_reservation, +) +from litellm.types.realtime import RealtimeQueryParams +from litellm.types.router import GenericLiteLLMParams + + +async def supervise_codex_call( + request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth, lease: RealtimeCallLease | None = None +) -> None: + with isolated_request_stash(): + await _start_codex_supervisor(request, call, auth, lease) + + +async def _start_codex_supervisor( + request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth, lease: RealtimeCallLease | None +) -> None: + import litellm + from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, CallSupervisor + + async def receive() -> Message: + body: Final[RealtimeQueryParams] = {"model": call.alias} + message: Final[Message] = { + "type": "http.request", + "body": json.dumps(body).encode(), + "more_body": False, + } + return 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_scope: Final[Scope] = {**request.scope} + observer_request: Final = Request(observer_scope, receive=receive) + processed, logger = await process_codex_request( + observer_request, + { # mutable-ok: common request processing enriches metadata + **build_sideband_request(call), + "model": call.alias, + }, + auth, + call.alias, + "_arealtime", + internal_realtime_observer=True, + ) + pinned: Final = { # mutable-ok: logging and provider parameter contract + **processed, + **build_sideband_request(call), + "extra_headers": MappingProxyType( + { + **configured_realtime_headers( + TypeAdapter[Mapping[str, object] | None](Mapping[str, object] | None).validate_python( + processed.get("extra_headers") + ) + ), + **configured_realtime_headers(call.extra_headers), + } + ), + "litellm_metadata": { # mutable-ok: Logging.update_from_kwargs requires a dict to retain ownership metadata + **TypeAdapter(Mapping[str, object]).validate_python( + processed.get("litellm_metadata") or MappingProxyType({}) + ), + **( + MappingProxyType( + { + "model_info": { # mutable-ok: logging and cost callbacks require a concrete model-info dict + **litellm.get_model_info(model=call.model_id), + "id": call.model_id, + } + } + ) + if call.model_id is not None + else MappingProxyType({}) + ), + }, + } + logger.update_from_kwargs( + kwargs=pinned, + model=call.model, + user=None, + optional_params={}, # mutable-ok: logging contract + litellm_params={ # mutable-ok: Logging.update_from_kwargs pops metadata from its argument + **logger.litellm_params, + "litellm_metadata": pinned["litellm_metadata"], + "arealtime": True, + }, + 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) + + async def force_close_call() -> None: + await handler.hangup_call(api_base) + + frontend_scope: Final[Scope] = {**request.scope, "type": "websocket"} + frontend: Final = WebSocket(frontend_scope, receive=receive, send=send) + stream: Final = RealTimeStreaming(frontend, connection, logger, model=call.model, user_api_key_dict=auth) + supervisor: Final = CallSupervisor( + connection, + stream, + logger, + auth, + close_call, + force_close_call=force_close_call, + terminal_usage_required=realtime_endpoint(call.model) == "live", + lease=lease, + ) + 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: + encrypted: Final = encrypt_value_helper(call.model_dump_json()) + return "rtc_litellm_" + base64.urlsafe_b64encode(encrypted.encode()).decode().rstrip("=") + + +def decode_call(token: str, authorization: str) -> CodexRealtimeCall: + try: + if not token.startswith("rtc_litellm_"): + raise ValueError("Invalid call prefix") + encoded: Final = token.removeprefix("rtc_litellm_") + encrypted: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True) + plaintext: Final = decrypt_value_helper(encrypted.decode(), key="codex_realtime_call") + call: Final = CodexRealtimeCall.model_validate_json(plaintext or "") + except (ValueError, TypeError, UnicodeError) as exc: + raise HTTPException(403, "Invalid realtime call") from exc + if ( + call.expires_at < time.time() + or call.owner != hashlib.sha256(authorization.encode(), usedforsecurity=False).hexdigest() + ): + raise HTTPException(403, "Invalid or expired realtime call") + return call + + +MAX_REALTIME_OFFER_BYTES: Final = 8 * 1024 * 1024 + + +async def _cache_bounded_offer_body(request: Request) -> None: + content_length: int | None + try: + content_length = int(request.headers.get("content-length", "")) + except ValueError: + # A missing or non-numeric content length is checked while streaming below. + content_length = None + if content_length is not None and content_length > MAX_REALTIME_OFFER_BYTES: + raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit") + if hasattr(request, "_body"): + if len(request._body) > MAX_REALTIME_OFFER_BYTES: # pyright: ignore[reportPrivateUsage] # validate Starlette's cached body without consuming it again + raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit") + return + if request._form is not None and request._stream_consumed: # pyright: ignore[reportPrivateUsage] # a mixed-case empty form cache may leave the stream unread + return + body: Final = bytearray() + async for chunk in request.stream(): + if len(body) + len(chunk) > MAX_REALTIME_OFFER_BYTES: + raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit") + body.extend(chunk) + request._body = bytes(body) # pyright: ignore[reportPrivateUsage] # Starlette has no public setter for its shared body cache + + +async def read_codex_offer(request: Request) -> CodexRealtimeOffer: + await _cache_bounded_offer_body(request) + content_type: Final = request.headers.get("content-type", "") + if _normalize_media_type(content_type) == "multipart/form-data": + if content_type.split(";", 1)[0] != "multipart/form-data" and not await request.form(): + try: + request._form = await MultiPartParser(request.headers, request.stream()).parse() # pyright: ignore[reportPrivateUsage] # Starlette exposes no setter for its shared form cache; # rebind-ok: Request.close must own and close uploaded files + request.scope.pop("parsed_body", None) + except MultiPartException as exc: + raise HTTPException(400, "Invalid realtime multipart offer") from exc + form: Final = await request.form() + return CodexRealtimeOffer.model_validate( + MappingProxyType({"sdp": form.get("sdp"), "session": json.loads(str(form.get("session", "{}")))}) + ) + return CodexRealtimeOffer.model_validate(await request.json()) + + +async def process_codex_request( + request: Request, + data: dict[str, object], # mutable-ok: common request processor enriches this dictionary + 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 + + processor: Final = ProxyBaseLLMRequestProcessing(data=data) + processed, logging_obj = await processor.common_processing_pre_call_logic( + request=request, + general_settings=server.general_settings, + user_api_key_dict=auth, + version=server.version, + proxy_logging_obj=server.proxy_logging_obj, + proxy_config=server.proxy_config, + llm_router=server.llm_router, + user_model=TypeAdapter[str | None](str | None).validate_python(server.user_model), + user_temperature=TypeAdapter[float | None](float | None).validate_python(server.user_temperature), + user_request_timeout=server.user_request_timeout, + user_max_tokens=server.user_max_tokens, + user_api_base=TypeAdapter[str | None](str | None).validate_python(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 + + +async def create_codex_realtime_call(request: Request) -> Response: + try: + with isolated_request_stash(): + return await _create_codex_realtime_call(request) + finally: + await request.close() + + +async def _create_codex_realtime_call(request: Request) -> Response: + from litellm.proxy import proxy_server as server + + try: + offer: Final = await read_codex_offer(request) + except ValueError as exc: + raise HTTPException(400, "Invalid realtime offer: expected sdp and session") from exc + model: Final = offer.session.model + if not model: + raise HTTPException(400, "session.model is required") + auth: Final = await user_api_key_auth( + request=request, + api_key=request.headers.get("authorization", ""), + azure_api_key_header=request.headers.get("api-key", ""), + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + custom_litellm_key_header=request.headers.get("x-litellm-api-key"), + ) + selected_key, _ = get_api_key( + request=request, + api_key=request.headers.get("authorization", ""), + azure_api_key_header=request.headers.get("api-key", ""), + custom_litellm_key_header=request.headers.get("x-litellm-api-key"), + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + pass_through_endpoints=None, + route="/v1/realtime/calls", + ) + custom_header: Final = server.general_settings.get("litellm_key_header_name") + 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 + call_lease: RealtimeCallLease | None = None + lease_transferred = False # rebind-ok: failed startup leaves the signaling task responsible for its lease + preprocessing_started = False # rebind-ok: only refund reservations belonging to this signaling request + limiter: Final = server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") + try: + await can_key_call_resolved_model( + model=model, + llm_model_list=TypeAdapter[tuple[object, ...] | None](tuple[object, ...] | None).validate_python( + server.llm_model_list + ), + valid_token=auth, + llm_router=server.llm_router, + ) + live_signaling: Final = request.url.path.rstrip("/") in ("/live", "/v1/live", "/openai/v1/live") + query: Final = ( + MappingProxyType({"intent": "quicksilver", "architecture": "avas", **request.query_params}) + if live_signaling + else request.query_params + ) + data: Final = build_call_request(offer, query, request.headers) + signaling_auth: Final = auth.model_copy(update=MappingProxyType({"budget_reservation": None})) + if isinstance(limiter, _PROXY_MaxParallelRequestsHandler) and ( + auth.max_parallel_requests is not None + or server.general_settings.get("global_max_parallel_requests") is not None + ): + raise HTTPException(400, "Realtime calls with parallel limits require the V3 rate limiter") + preprocessing_started = True + processed, _ = await process_codex_request(request, data, signaling_auth, model, "arealtime_calls") + if isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + call_lease = limiter.transfer_realtime_call_slot(processed) + if call_lease is not None: + call_lease.start() + if not await call_lease.renew(): + raise HTTPException(503, "Realtime call quota reservation was lost") + with isolated_request_stash(): + result: Final = await server.route_request( + data=processed, + route_type="arealtime_calls", + llm_router=server.llm_router, + user_model=TypeAdapter[str | None](str | None).validate_python(server.user_model), + ) + try: + response: Final = await result + except BaseLLMException as exc: + raise HTTPException(exc.status_code, str(exc)) from exc + if not isinstance(response, httpx.Response): + raise HTTPException(502, "Invalid realtime signaling response") + if response.is_error: + return Response(response.content, status_code=response.status_code, media_type="application/json") + try: + call: Final = parse_call_response( + response, + alias=model, + owner=hashlib.sha256(f"Bearer {owner_key}".encode(), usedforsecurity=False).hexdigest(), + expires_at=time.time() + 3600, + ) + except ValueError as exc: + raise HTTPException(400, str(exc)) from exc + supervised_call: Final = call.model_copy( + update=MappingProxyType({"usage_supervised": True, "parallel_reserved": call_lease is not None}) + ) + token: Final = encode_call(supervised_call) + supervision_started = True + if call_lease is None: + await supervise_codex_call(request, supervised_call, auth) + else: + await supervise_codex_call(request, supervised_call, auth, call_lease) + lease_transferred = True + return Response( + response.content, + status_code=response.status_code, + media_type="application/sdp", + headers=MappingProxyType( + {"Location": f"/v1/live/{token}" if live_signaling else f"/v1/realtime/calls/{token}"} + ), + ) + finally: + try: + if call_lease is not None and not lease_transferred: + await call_lease.close() + finally: + try: + if preprocessing_started and isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + await asyncio.shield( + limiter.async_post_call_failure_hook( + request_data={}, # mutable-ok: existing failure-hook contract + original_exception=Exception("Realtime signaling completed without token usage"), + user_api_key_dict=auth, + ) + ) + finally: + 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: + import litellm + from litellm.proxy import proxy_server as server + + protocols: Final = tuple( + p.strip() for p in websocket.headers.get("sec-websocket-protocol", "").split(",") if p.strip() + ) + logging_obj: Logging | None = None # rebind-ok: cleanup needs the logger only after pre-call succeeds + attachment_limiter: _PROXY_MaxParallelRequestsHandler | _PROXY_MaxParallelRequestsHandler_v3 | None = None + cleanup_token: Token[Callable[[], Awaitable[None]] | None] | None = None + try: + try: + api_key: Final = get_websocket_api_key(websocket) + if not api_key: + raise HTTPException(403, "No API key provided") + call: Final = decode_call(token, f"Bearer {api_key}") + await can_key_call_resolved_model( + model=call.alias, + llm_model_list=TypeAdapter[tuple[object, ...] | None](tuple[object, ...] | None).validate_python( + server.llm_model_list + ), + valid_token=auth, + llm_router=server.llm_router, + ) + except (HTTPException, ProxyException): + await websocket.close(code=1008, reason="Invalid realtime call") + return + + async def receive() -> Message: + return { # mutable-ok: ASGI receive message + "type": "http.request", + "body": json.dumps({"model": call.alias}).encode(), # mutable-ok: JSON request serialization + "more_body": False, + } + + request: Final = Request( + { # mutable-ok: Starlette stores request state in the ASGI scope + **websocket.scope, + "type": "http", + "method": "POST", + "path": websocket.scope.get("path", "/v1/realtime"), + }, + receive=receive, + ) + data: Final = { # mutable-ok: common request processor enriches routing arguments + **build_sideband_request(call), + "model": call.alias, + "websocket": websocket, + "guardrails": [ # mutable-ok: guardrail processing expects a list + name.strip() for name in websocket.query_params.get("guardrails", "").split(",") if name.strip() + ], + } + limiter: Final = server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") + if call.usage_supervised and isinstance( + limiter, (_PROXY_MaxParallelRequestsHandler, _PROXY_MaxParallelRequestsHandler_v3) + ): + attachment_limiter = limiter + if isinstance(limiter, _PROXY_MaxParallelRequestsHandler): + limiter.begin_realtime_attachment(data) + try: + with realtime_call_attachment(websocket) if call.parallel_reserved else nullcontext(): + processed, logging_obj = await process_codex_request(request, data, auth, call.alias, "_arealtime") + except Exception: # noqa: BLE001 # custom hook exceptions must reject the connection + verbose_proxy_logger.exception("Realtime sideband pre-call rejected") + await websocket.close(code=1008, reason="Realtime pre-call rejected") + return + await websocket.accept( + subprotocol=next((p for p in protocols if not p.startswith("openai-insecure-api-key.")), None) + ) + if attachment_limiter is not None: + selected_limiter: Final = attachment_limiter + + async def release_attachment() -> None: + await selected_limiter.async_release_realtime_attachment(data, auth) + + cleanup_token = realtime_attachment_cleanup.set(release_attachment) + await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call + model=f"chatgpt/{call.model}", + websocket=websocket, + **MappingProxyType( + { + key: value + for key, value in { # mutable-ok: retain processed metadata with pinned routing + **processed, + **build_sideband_request(call), + "extra_headers": MappingProxyType( + { + **configured_realtime_headers( + TypeAdapter[Mapping[str, object] | None]( + Mapping[str, object] | None + ).validate_python(processed.get("extra_headers")) + ), + **configured_realtime_headers(call.extra_headers), + } + ), + "websocket": websocket, + "user_api_key_dict": auth, + "chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None, + }.items() + if key not in ("model", "websocket") + } + ), + ) + finally: + try: + if attachment_limiter is not None: + await attachment_limiter.async_release_realtime_attachment(data, auth) + finally: + if cleanup_token is not None: + realtime_attachment_cleanup.reset(cleanup_token) + if logging_obj is None or not logging_obj.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY): + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py new file mode 100644 index 00000000000..710220b17d2 --- /dev/null +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -0,0 +1,270 @@ +import asyncio +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import suppress +from typing import Final, Protocol + +from pydantic import BaseModel, ValidationError +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 +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease +from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + release_or_invalidate_budget_reservation, +) +from litellm.types.realtime import LiveSessionUsageEvent + + +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, + termination_timeout: float = 60, + logging_timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE, + terminal_usage_required: bool = True, + connected_ready: bool = False, + force_close_call: Callable[[], Awaitable[None]] | None = None, + lease: RealtimeCallLease | 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._lease = lease + 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 = connected_ready + if connected_ready: + self._ready.set() + self._terminal = False + self._terminal_usage_valid = False + self._close_confirmed = False + self._accounting_complete = 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 self._lease is not None and not await self._lease.renew(): + raise RuntimeError("Call observer lost its quota reservation during startup") + 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_after_failed_start() + raise + + async def _close_after_failed_start(self) -> None: + cleanup: Final = asyncio.create_task(self.close()) + while not cleanup.done(): + with suppress(asyncio.CancelledError): + await asyncio.shield(cleanup) + cleanup.result() + + 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 + try: + LiveSessionUsageEvent.model_validate_json(message) + except ValidationError: + self._terminal_usage_valid = False + else: + self._terminal_usage_valid = True + return + + def _usage_complete(self) -> bool: + if self._terminal_usage_required: + return self._terminal and self._terminal_usage_valid + return self._terminal or self._close_confirmed + + async def _run(self) -> None: + try: + await self._observe() + finally: + try: + if self._lease is not None: + await self._lease.close() + finally: + self._ready.set() + + async def _observe(self) -> None: + reader: Final = asyncio.create_task(self._read()) + stopped: Final = asyncio.create_task(self._stop.wait()) + lease_failed: Final = asyncio.create_task(self._lease.wait_failed()) if self._lease is not None else None + try: + await asyncio.wait( + (reader, stopped, lease_failed) if lease_failed is not None else (reader, stopped), + timeout=self._lifetime, + return_when=asyncio.FIRST_COMPLETED, + ) + finally: + try: + if not self._terminal: + deadline: Final = asyncio.get_running_loop().time() + self._termination_timeout + primary_deadline: Final = ( + deadline - self._termination_timeout / 2 + if self._terminal_usage_required and self._force_close_call is not None + else deadline + ) + try: + await asyncio.wait_for( + self._close_call(), timeout=max(0.0, primary_deadline - asyncio.get_running_loop().time()) + ) + 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, timeout=max(0.0, primary_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() + if lease_failed is not None: + lease_failed.cancel() + # Branching rather than a conditional star-unpacked tuple: the overload solver cannot + # bind one result type across a tuple whose length depends on the branch. + if lease_failed is not None: + await asyncio.gather(reader, stopped, lease_failed, return_exceptions=True) + else: + 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 asyncio.wait_for( + self._stream.log_messages(wait_for_dispatch=True), timeout=self._logging_timeout + ) + self._accounting_complete = not bool( + self._logging.model_call_details.get("realtime_backend_accounting_incomplete") + ) + except asyncio.TimeoutError: + verbose_proxy_logger.error("Realtime observer timed out dispatching usage accounting") + finally: + 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 + ) + 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], *, timeout: float | None = None) -> None: + try: + 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") + 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() diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index be2ac2ff33e..59c3f59e427 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -16,7 +16,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( + _normalize_media_type, # pyright: ignore[reportPrivateUsage] # reuse the shared HTTP media-type normalization contract + _read_request_body, +) from litellm.proxy.common_utils.openai_error_payload import ( error_status_code, openai_error_param, @@ -359,6 +362,15 @@ async def create_realtime_client_secret( return RealtimeClientSecretResponse(**upstream_json) +@router.post("/v1/live", tags=["realtime"]) # mutable-ok: FastAPI route metadata uses a mutable tag list +@router.post("/live", tags=["realtime"]) # mutable-ok: FastAPI route metadata uses a mutable tag list +@router.post("/openai/v1/live", tags=["realtime"]) # mutable-ok: FastAPI route metadata uses a mutable tag list +async def proxy_live_calls(request: Request) -> Response: + from litellm.proxy.realtime_endpoints.call_sessions import create_codex_realtime_call + + return await create_codex_realtime_call(request) + + @router.post( "/v1/realtime/calls", tags=["realtime"], @@ -375,6 +387,11 @@ async def proxy_realtime_calls( request: Request, fastapi_response: Response, ) -> Response: + if _normalize_media_type(request.headers.get("content-type", "")) in ("application/json", "multipart/form-data"): + from litellm.proxy.realtime_endpoints.call_sessions import create_codex_realtime_call + + return await create_codex_realtime_call(request) + from litellm.proxy.proxy_server import ( add_litellm_data_to_request, general_settings, @@ -410,6 +427,12 @@ async def proxy_realtime_calls( sdp_body: Final[bytes] = await request.body() decoded_payload: Final = _decode_realtime_token_payload(decrypted_token_value) + if decoded_payload is None and decrypted_token_value.lstrip().startswith(("{", "[")): + return Response( + content='{"error":"Invalid or expired token"}', + status_code=http_status.HTTP_401_UNAUTHORIZED, + media_type="application/json", + ) if decoded_payload is not None: # Check token expiry expires_at: Final = decoded_payload.get("expires_at") diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py new file mode 100644 index 00000000000..26724a5146d --- /dev/null +++ b/litellm/proxy/realtime_endpoints/live.py @@ -0,0 +1,1639 @@ +import asyncio +import base64 +import hashlib +import json +import time +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping +from contextlib import asynccontextmanager, nullcontext +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, TypeVar + +import httpx +from fastapi import APIRouter, HTTPException, Request, Response, WebSocket, WebSocketDisconnect +from pydantic import BaseModel, Field, JsonValue, TypeAdapter +from starlette.types import Message + +from litellm._logging import verbose_proxy_logger + +if TYPE_CHECKING: + from websockets.asyncio.client import ClientConnection + +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming +from litellm.llms.chatgpt.live import LiveDeployment, LiveOperation, LiveTransport, live_session_path +from litellm.models.budget import LiteLLM_BudgetTable +from litellm.models.team import LiteLLM_TeamTable +from litellm.proxy._types import ( + LiteLLM_ProjectTableCachedObj, + LiteLLM_TeamTableCachedObj, + LitellmUserRoles, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_checks import ( + _cache_team_object, # pyright: ignore[reportPrivateUsage] # same cache write the chat path performs + _get_team_object_from_cache, # pyright: ignore[reportPrivateUsage] # same cache read the chat path performs + can_key_call_resolved_model, # pyright: ignore[reportUnknownVariableType] # legacy authorization accepts untyped deployment lists + can_org_access_model, + can_user_call_model, + collect_matched_model_access_groups, + get_object_permission, + get_org_object, + get_project_object, + get_team_membership, + get_team_object, + get_user_object, +) +from litellm.proxy.auth.user_api_key_auth import get_websocket_api_key, user_api_key_auth +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.common_utils.user_api_key_cache import ( + get_management_object_ttl, + live_model_access_group_limits_cache_key, +) +from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # limiter class is the existing hook identity +) +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # slot transfer is available on this existing hook + isolated_request_stash, +) +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease, realtime_call_attachment +from litellm.proxy.realtime_endpoints.call_sessions import process_codex_request +from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, CallSupervisor +from litellm.proxy.spend_tracking.budget_reservation import ( + release_or_invalidate_budget_reservation, # pyright: ignore[reportUnknownVariableType] # budget helper accepts legacy reservation dicts +) +from litellm.repositories.budget_repository import BudgetRepository +from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE +from litellm.repositories.table_repositories import ModelAccessGroupBudgetRepository +from litellm.repositories.team_repository import TeamRepository + +_routes: Final = APIRouter() +_JSON: Final = TypeAdapter[JsonValue](JsonValue) +_EMPTY: Final[Mapping[str, JsonValue]] = MappingProxyType({}) +_CACHEABLE_MODEL = TypeVar("_CACHEABLE_MODEL", bound=BaseModel) +_MAPPING: Final = TypeAdapter(Mapping[str, object]) +_OBJECT: Final = TypeAdapter(Mapping[str, JsonValue]) +_DEPLOYMENT: Final = TypeAdapter(LiveDeployment) +_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...]) +_OBJECT_VALUE: Final = TypeAdapter(object) +_SEQUENCE: Final = TypeAdapter(tuple[object, ...]) +_PREFIX: Final = "live_litellm_" + + +def _json_value(value: object) -> JsonValue: + root: Final[list[JsonValue]] = [None] # mutable-ok: iterative conversion fills JSON output containers + pending: Final[ # mutable-ok: work stack carries mutable JSON output containers + list[tuple[object, dict[str, JsonValue] | list[JsonValue], str | int, int]] + ] = [(value, root, 0, 0)] # mutable-ok: traversal adds pending nodes + while pending: + source, parent, key, depth = pending.pop() + if depth > 256: + raise ValueError("Live JSON nesting exceeds the supported depth") + converted: JsonValue + if isinstance(source, Mapping): + entries: Mapping[str, object] = _MAPPING.validate_python(source) + converted = {name: None for name in entries} # mutable-ok: JSON wire objects require dicts + pending.extend((item, converted, name, depth + 1) for name, item in entries.items()) + elif isinstance(source, (tuple, list)): + items: tuple[object, ...] = TypeAdapter(tuple[object, ...]).validate_python(source) + array: list[JsonValue] = [None] * len(items) + pending.extend((item, array, index, depth + 1) for index, item in enumerate(items)) + converted = array + else: + converted = _JSON.validate_python(source) + match parent, key: + case dict(), str(): + parent[key] = converted + case list(), int(): + parent[key] = converted + case _: + raise TypeError("Invalid Live JSON conversion target") + return root[0] + + +def _object(value: object) -> Mapping[str, JsonValue]: + return _OBJECT.validate_python(_json_value(value)) + + +def _mutable( + value: Mapping[str, object], +) -> dict[str, object]: # mutable-ok: legacy proxy and ASGI contracts mutate inputs + return dict(value) # mutable-ok: make the mutable copy at the framework boundary + + +def _encode_json(value: object) -> str: + return json.dumps(_json_value(value)) + + +class _ConnectionState: + def __init__(self) -> None: + self.connection: ClientConnection | None = None + + +class LiveHandle(BaseModel): + session_id: str + alias: str + deployment: Mapping[str, JsonValue] + owner: str + expires_at: float + parallel_reserved: bool = False + initialization_seconds: float = 0 + policy: Mapping[str, JsonValue] = Field(default_factory=lambda: _EMPTY) + + +def encode_session(handle: LiveHandle) -> str: + encrypted: Final = encrypt_value_helper(handle.model_dump_json()) + return _PREFIX + base64.urlsafe_b64encode(encrypted.encode()).decode().rstrip("=") + + +def decode_session(token: str, owner: str) -> LiveHandle: + try: + if not token.startswith(_PREFIX): + raise ValueError("Invalid prefix") + encoded: Final = token[len(_PREFIX) :] + encrypted: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True) + handle: Final = LiveHandle.model_validate_json( + decrypt_value_helper(encrypted.decode(), key="live_session") or "" + ) + if handle.owner != owner or handle.expires_at <= time.time(): + raise ValueError("Invalid ownership or expiry") + live_session_path(handle.session_id, "attach") + return handle + except (ValueError, TypeError, UnicodeError) as exc: + raise HTTPException(403, "Invalid or expired Live session") from exc + + +def rewrite_session_ids(value: JsonValue | Mapping[str, JsonValue], raw_id: str, public_id: str) -> JsonValue: + if not isinstance(value, Mapping): + return _json_value(value) + session: Final = value.get("session") + return _json_value( + MappingProxyType( + { + **value, + **(MappingProxyType({"session_id": public_id}) if value.get("session_id") == raw_id else _EMPTY), + **( + MappingProxyType({"session": MappingProxyType({**session, "id": public_id})}) + if isinstance(session, Mapping) and session.get("id") == raw_id + else _EMPTY + ), + } + ) + ) + + +def _owner(auth: UserAPIKeyAuth) -> str: + if not auth.api_key: + raise HTTPException(403, "Live sessions require an authenticated API key") + return hashlib.sha256(auth.api_key.encode(), usedforsecurity=False).hexdigest() + + +async def _auth(request: Request) -> UserAPIKeyAuth: + return await user_api_key_auth( + request=request, + api_key=request.headers.get("authorization", ""), + azure_api_key_header=request.headers.get("api-key", ""), + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + custom_litellm_key_header=request.headers.get("x-litellm-api-key"), + ) + + +async def _body(request: Request) -> Mapping[str, JsonValue]: + from litellm.proxy.realtime_endpoints.call_sessions import MAX_REALTIME_OFFER_BYTES + + chunks: Final = bytearray() + async for chunk in request.stream(): + if len(chunks) + len(chunk) > MAX_REALTIME_OFFER_BYTES: + raise HTTPException(413, "Live request exceeds the 8 MiB limit") + chunks.extend(chunk) + try: + return _OBJECT.validate_json(bytes(chunks)) if chunks else _EMPTY + except ValueError as exc: + raise HTTPException(400, "Expected a JSON object") from exc + + +def _request(source: Request | WebSocket, body: Mapping[str, JsonValue], pinned_model: str | None = None) -> Request: + """Build the synthetic POST request that authenticates one Live operation. + + A copied scope can carry the body a previous authentication parsed, so the + cached ``parsed_body`` is dropped and the request parses ``body`` again. + ``litellm_pinned_realtime_model`` marks the requests whose model this endpoint + dispatches itself, so a key-level budget fallback cannot authorize a different + model than the one that is about to run. + """ + + async def receive() -> Message: + return _mutable( + MappingProxyType({"type": "http.request", "body": _encode_json(body).encode(), "more_body": False}) + ) + + scope: Final = _mutable(MappingProxyType({**source.scope, "type": "http", "method": "POST"})) + scope.pop("parsed_body", None) + if pinned_model is not None: + scope["litellm_pinned_realtime_model"] = pinned_model + return Request(scope, receive=receive) + + +def _session_model(body: Mapping[str, JsonValue], fallback: str | None = None) -> str: + session: Final = body.get("session", _EMPTY) + if not isinstance(session, Mapping): + raise HTTPException(400, "session must be a JSON object") + model: Final = session.get("model", fallback) + if not isinstance(model, str) or not model: + raise HTTPException(400, "session.model is required") + if fallback is not None and "model" in session: + raise HTTPException(400, "A session fork cannot change the authorized model") + return model + + +async def _live_organization_id(auth: UserAPIKeyAuth) -> str | None: + if auth.org_id is not None or auth.team_id is None: + return auth.org_id + from litellm.proxy import proxy_server as server + + try: + team_object: Final = await get_team_object( + team_id=auth.team_id, + prisma_client=server.prisma_client, + user_api_key_cache=server.user_api_key_cache, + parent_otel_span=auth.parent_otel_span, + proxy_logging_obj=server.proxy_logging_obj, + ) + except Exception as exc: + raise HTTPException(503, "Could not verify Live team organization model access") from exc + return team_object.organization_id + + +async def _authorize(model: str, auth: UserAPIKeyAuth) -> None: + from litellm.proxy import proxy_server as server + + await can_key_call_resolved_model( + model=model, + llm_model_list=list( # mutable-ok: legacy authorization requires a concrete list + TypeAdapter(tuple[Mapping[str, object], ...]).validate_python(server.llm_model_list or ()) # pyright: ignore[reportUnknownMemberType] # legacy registry is validated here + ), + valid_token=auth, + llm_router=server.llm_router, + ) + + if auth.user_id is not None and auth.team_id is None and auth.user_role != LitellmUserRoles.PROXY_ADMIN: + try: + user_object: Final = await get_user_object( + user_id=auth.user_id, + prisma_client=server.prisma_client, + user_api_key_cache=server.user_api_key_cache, + user_id_upsert=False, + parent_otel_span=auth.parent_otel_span, + proxy_logging_obj=server.proxy_logging_obj, + ) + except Exception as exc: + raise HTTPException(503, "Could not verify Live user model access") from exc + if user_object is None: + raise HTTPException(503, "Could not verify Live user model access") + await can_user_call_model(model=model, llm_router=server.llm_router, user_object=user_object) + + organization_id: Final = await _live_organization_id(auth) + if organization_id is not None: + try: + org_object: Final = await get_org_object( + org_id=organization_id, + prisma_client=server.prisma_client, + user_api_key_cache=server.user_api_key_cache, + parent_otel_span=auth.parent_otel_span, + proxy_logging_obj=server.proxy_logging_obj, + ) + except Exception as exc: + raise HTTPException(503, "Could not verify Live organization model access") from exc + if org_object is None: + raise HTTPException(503, "Could not verify Live organization model access") + can_org_access_model(model=model, org_object=org_object, llm_router=server.llm_router) + + +async def _deployment(model: str, processed: Mapping[str, object]) -> LiveDeployment: + from litellm.proxy import proxy_server as server + + if server.llm_router is None: + raise HTTPException(503, "Live requires a configured model deployment") + selected: Final = _MAPPING.validate_python( + await server.llm_router.async_get_available_deployment( # pyright: ignore[reportUnknownMemberType] # validate the legacy router result at this boundary + model=model, request_kwargs=_mutable(processed) + ) + ) + await server.llm_router.async_routing_strategy_pre_call_checks(_mutable(selected), None) # pyright: ignore[reportUnknownMemberType] # existing routing strategy has an untyped deployment contract + params: Final = _object(selected["litellm_params"]) + qualified: Final = str(params["model"]) + prefix, _, suffix = qualified.partition("/") + provider: Final = prefix if prefix in ("openai", "chatgpt") else "openai" + upstream: Final = suffix if prefix in ("openai", "chatgpt") else qualified + if provider == "chatgpt" and any( + params.get(key) is not None for key in ("chatgpt_auth_profile", "chatgpt_token_dir", "chatgpt_auth_file") + ): + raise HTTPException( + 400, "ChatGPT Live uses the proxy OAuth credentials; deployment auth overrides are unsupported" + ) + if prefix not in ("openai", "chatgpt"): + if "/" in qualified: + raise HTTPException(400, "Live requires an OpenAI or ChatGPT deployment") + return _DEPLOYMENT.validate_python( + _mutable( + MappingProxyType( + { + "model": upstream, + "provider": provider, + "model_id": str(_MAPPING.validate_python(selected["model_info"])["id"]), + "api_base": params.get("api_base"), + "api_key": params.get("api_key"), + "extra_headers": params.get("extra_headers") or _EMPTY, + "extra_query": params.get("extra_query") or _EMPTY, + } + ) + ) + ) + + +def _pinned(handle: LiveHandle) -> LiveDeployment: + return _DEPLOYMENT.validate_python(_mutable(handle.deployment)) + + +def _validate_pinned_deployment(handle: LiveHandle) -> LiveDeployment: + """Reject handles whose deployment was removed, blocked, or replaced.""" + from litellm.proxy import proxy_server as server + + deployment: Final = _pinned(handle) + router = server.llm_router + if router is None or deployment.model_id is None: + raise HTTPException(410, "Live session deployment is no longer available") + configured_raw = router.get_deployment(model_id=deployment.model_id) + if configured_raw is None: + raise HTTPException(410, "Live session deployment is no longer available") + configured_data: Final = ( + configured_raw.model_dump() + if hasattr(configured_raw, "model_dump") + else vars(configured_raw) + if not isinstance(configured_raw, Mapping) + else configured_raw + ) + configured: Final = _MAPPING.validate_python(configured_data) + model_info: Final = _MAPPING.validate_python(configured.get("model_info", _EMPTY)) + if model_info.get("blocked") is True: + raise HTTPException(410, "Live session deployment is no longer available") + params: Final = _object(configured["litellm_params"]) + qualified: Final = str(params.get("model", "")) + prefix, _, suffix = qualified.partition("/") + provider: Final = prefix if prefix in ("openai", "chatgpt") else "openai" + upstream: Final = suffix if prefix in ("openai", "chatgpt") else qualified + if any( + ( + deployment.model != upstream, + deployment.provider != provider, + str(model_info.get("id")) != deployment.model_id, + params.get("api_base") != deployment.api_base, + params.get("api_key") != deployment.api_key, + (params.get("extra_headers") or _EMPTY) != deployment.extra_headers, + (params.get("extra_query") or _EMPTY) != deployment.extra_query, + ) + ): + raise HTTPException(410, "Live session deployment is no longer available") + return deployment + + +def _new_handle( + session_id: str, + alias: str, + deployment: LiveDeployment, + auth: UserAPIKeyAuth, + lease: RealtimeCallLease | None, + initialization_seconds: float = 0, + policy: Mapping[str, JsonValue] | None = None, +) -> LiveHandle: + live_session_path(session_id, "attach") + routing: Final = _object( + MappingProxyType( + { + "model": deployment.model, + "provider": deployment.provider, + "model_id": deployment.model_id, + "api_base": deployment.api_base, + "api_key": deployment.api_key, + "extra_headers": deployment.extra_headers, + "extra_query": deployment.extra_query, + } + ) + ) + return LiveHandle( + session_id=session_id, + alias=alias, + deployment=routing, + owner=_owner(auth), + expires_at=time.time() + 30 * 86400, + parallel_reserved=lease is not None, + initialization_seconds=initialization_seconds, + policy=policy or _EMPTY, + ) + + +def _session_id(payload: Mapping[str, JsonValue]) -> str: + session: Final = payload.get("session") + if isinstance(session, Mapping) and isinstance(session.get("id"), str): + return TypeAdapter(str).validate_python(session["id"]) + raise HTTPException(502, "Upstream did not return a Live session ID") + + +class _BudgetOwnership: + def __init__(self, auth: UserAPIKeyAuth) -> None: + self.auth = auth + self.transferred = False + + def replace_auth(self, auth: UserAPIKeyAuth) -> None: + self.auth = auth + + +@asynccontextmanager +async def _budget_scope(auth: UserAPIKeyAuth) -> AsyncGenerator[_BudgetOwnership]: + ownership: Final = _BudgetOwnership(auth) + try: + yield ownership + finally: + if not ownership.transferred: + await release_or_invalidate_budget_reservation(budget_reservation=ownership.auth.budget_reservation) + + +async def _reauth(ownership: _BudgetOwnership, request: Request, body: Mapping[str, JsonValue], model: str) -> None: + await release_or_invalidate_budget_reservation(budget_reservation=ownership.auth.budget_reservation) + ownership.replace_auth(await _auth(_request(request, MappingProxyType({**body, "model": model}), model))) + + +def _policy_object(value: object) -> Mapping[str, JsonValue]: + try: + return _object(value) + except ValueError as exc: + raise HTTPException(400, "Live session delegation and responses must be JSON objects") from exc + + +def _policy_body(body: Mapping[str, JsonValue], source: LiveHandle | None) -> Mapping[str, JsonValue]: + current: Final = _policy_object(body.get("session", _EMPTY)) + inherited: Final = source.policy if source is not None else _EMPTY + parent_delegation: Final = _policy_object(inherited.get("delegation") or _EMPTY) + child_delegation: Final = _policy_object(current.get("delegation") or _EMPTY) + merged_delegation: Final = MappingProxyType( + { + **parent_delegation, + **child_delegation, + "responses": MappingProxyType( + { + **_policy_object(parent_delegation.get("responses") or _EMPTY), + **_policy_object(child_delegation.get("responses") or _EMPTY), + } + ), + } + ) + return _policy_object( + MappingProxyType( + { + **body, + "session": MappingProxyType( + { + **inherited, + **current, + **( + MappingProxyType({"delegation": merged_delegation}) + if parent_delegation or child_delegation + else _EMPTY + ), + } + ), + } + ) + ) + + +def _session_policy(body: Mapping[str, JsonValue], source: LiveHandle | None) -> Mapping[str, JsonValue]: + session: Final = _object(_policy_body(body, source)["session"]) + delegation: Final = _object(session.get("delegation") or _EMPTY) + responses: Final = _object(delegation.get("responses") or _EMPTY) + return _object( + MappingProxyType( + { + **( + MappingProxyType( + { + "delegation": MappingProxyType( + { + "type": delegation.get("type"), + "responses": MappingProxyType({"model": responses.get("model")}), + } + ) + } + ) + if delegation.get("type") == "responses" or responses.get("model") + else _EMPTY + ), + **(MappingProxyType({"client": session["client"]}) if "client" in session else _EMPTY), + } + ) + ) + + +def _managed_constraints(auth: UserAPIKeyAuth) -> bool: + from litellm.proxy import proxy_server as server + + if any( + value is not None + for value in ( + auth.rpm_limit, + auth.tpm_limit, + auth.team_rpm_limit, + auth.team_tpm_limit, + auth.user_rpm_limit, + auth.user_tpm_limit, + auth.organization_rpm_limit, + auth.organization_tpm_limit, + auth.team_member_rpm_limit, + auth.team_member_tpm_limit, + auth.end_user_rpm_limit, + auth.end_user_tpm_limit, + auth.max_parallel_requests, + auth.max_budget, + auth.team_max_budget, + auth.user_max_budget, + auth.end_user_max_budget, + auth.organization_max_budget, + ) + ): + return True + + # An admin-configured proxy-wide concurrency cap admits every key through the limiter, + # and delegated backend invocations never reach that admission, so it constrains the key + # the same way a key-level `max_parallel_requests` does. + if ( + _MAPPING.validate_python(getattr(server, "general_settings", None) or _EMPTY).get( + "global_max_parallel_requests" + ) + is not None + ): + return True + + direct_maps: Final[tuple[object, ...]] = ( + _OBJECT_VALUE.validate_python(getattr(auth, "model_max_budget", None)), + _OBJECT_VALUE.validate_python(getattr(auth, "user_model_max_budget", None)), + _OBJECT_VALUE.validate_python(getattr(auth, "end_user_model_max_budget", None)), + _OBJECT_VALUE.validate_python(getattr(auth, "rpm_limit_per_model", None)), + _OBJECT_VALUE.validate_python(getattr(auth, "tpm_limit_per_model", None)), + _OBJECT_VALUE.validate_python(getattr(auth, "budget_limits", None)), + ) + if any(_nonempty_limit_value(value) for value in direct_maps): + return True + + pending: Final[list[object]] = [ # mutable-ok: explicit metadata traversal stack + value + for value in ( + _OBJECT_VALUE.validate_python(getattr(auth, "team_metadata", None)), + _OBJECT_VALUE.validate_python(getattr(auth, "metadata", None)), + _OBJECT_VALUE.validate_python(getattr(auth, "organization_metadata", None)), + _OBJECT_VALUE.validate_python(getattr(auth, "project_metadata", None)), + ) + if isinstance(value, (Mapping, list, tuple)) + ] + visited: Final[set[int]] = set() + while pending: + current: object = pending.pop() + if id(current) in visited: + continue + visited.add(id(current)) + if len(visited) > 4096: + return True + if isinstance(current, Mapping): + entries: Mapping[str, object] = _MAPPING.validate_python(current) + for key, item in entries.items(): + if key in ( + "rpm_limit", + "tpm_limit", + "max_budget", + "model_rpm_limit", + "model_tpm_limit", + "model_itpm_limit", + "model_otpm_limit", + "model_max_budget", + "budget_limits", + ) and _nonempty_limit_value(item): + return True + if isinstance(item, (Mapping, list, tuple)): + nested: object = _OBJECT_VALUE.validate_python(item) + pending.append(nested) + elif isinstance(current, (list, tuple)): + sequence: object = _OBJECT_VALUE.validate_python(current) + pending.extend(_SEQUENCE.validate_python(sequence)) + return False + + +def _nonempty_limit_value(value: object) -> bool: + if value is None: + return False + if isinstance(value, Mapping): + mapping: Final[object] = _OBJECT_VALUE.validate_python(value) + return bool(_MAPPING.validate_python(mapping)) + if isinstance(value, (list, tuple)): + sequence: Final[object] = _OBJECT_VALUE.validate_python(value) + return bool(_SEQUENCE.validate_python(sequence)) + return True + + +def _restricted_model_list(models: object) -> bool: + values: Final = _MODEL_NAMES.validate_python(models or ()) + return bool(values) and "*" not in values and "all-proxy-models" not in values + + +def _restricted_models(auth: UserAPIKeyAuth) -> bool: + if _restricted_model_list(getattr(auth, "models", None)) or _restricted_model_list( + getattr(auth, "team_models", None) + ): + return True + if auth.user_id is not None and auth.user_role != LitellmUserRoles.PROXY_ADMIN: + return True + return bool( + auth.access_group_ids or auth.matched_model_access_groups or auth.project_id or auth.org_id or auth.team_id + ) + + +def _live_budget_configured(value: object, zero_is_limit: bool) -> bool: + if value is None: + return False + value_object: Final[object] = value + value_mapping: Final[Mapping[str, object] | None] = ( + _MAPPING.validate_python(value) if isinstance(value, Mapping) else None + ) + max_budget: Final[object] = _OBJECT_VALUE.validate_python( + value_mapping.get("max_budget") if value_mapping is not None else getattr(value_object, "max_budget", None) + ) + if max_budget is not None: + if zero_is_limit or (isinstance(max_budget, (int, float)) and max_budget > 0): + return True + if not isinstance(max_budget, (int, float)): + return True + fields: Final[tuple[object, ...]] = tuple( + _OBJECT_VALUE.validate_python( + value_mapping.get(field) if value_mapping is not None else getattr(value_object, field, None) + ) + for field in ("rpm_limit", "tpm_limit", "model_max_budget", "max_parallel_requests") + ) + return any(_nonempty_limit_value(field) for field in fields) + + +def _live_budget_scope_present(auth: UserAPIKeyAuth, model: str | None, llm_router: object | None) -> bool: + explicit_model_scope: Final = _restricted_model_list(getattr(auth, "models", None)) or _restricted_model_list( + getattr(auth, "team_models", None) + ) + model_group_lookup_scope: Final = bool( + model is not None and llm_router is not None and (explicit_model_scope or auth.org_id is not None) + ) + return bool( + auth.team_id + or auth.project_id + or auth.access_group_ids + or auth.matched_model_access_groups + or model_group_lookup_scope + ) + + +async def _live_team_membership(auth: UserAPIKeyAuth) -> object | None: + from litellm.proxy import proxy_server as server + + if auth.team_id is None or auth.user_id is None: + return None + return await get_team_membership( + user_id=auth.user_id, + team_id=auth.team_id, + prisma_client=server.prisma_client, + user_api_key_cache=server.user_api_key_cache, + proxy_logging_obj=server.proxy_logging_obj, + ) + + +async def _live_cached_object( + *, + key: str, + model_type: type[_CACHEABLE_MODEL], + load: Callable[[], Awaitable[_CACHEABLE_MODEL | None]], +) -> _CACHEABLE_MODEL | None: + """Read one management object through the proxy cache, storing the row when the read misses. + + ``auth_checks`` already caches these rows for the chat path, but the two getters this gate + would use are unusable there: ``get_team_object`` reports every failed database read as an + HTTP 404, and ``get_team_member_default_budget`` returns ``None`` when its read raises. Both + turn an outage into "no limit configured", and this gate answers that question by allowing + managed delegation, so an unreadable limit has to stay an error. The cache key and TTL stay + the shared ones, so the entry is still written, read, and invalidated like any other. + """ + from litellm.proxy import proxy_server as server + + cached: Final = await server.user_api_key_cache.async_get_cache(key=key, model_type=model_type) + if cached is not None: + return cached + loaded: Final = await load() + if loaded is not None: + await server.user_api_key_cache.async_set_cache( + key=key, + value=loaded, + model_type=model_type, + ttl=get_management_object_ttl(server.user_api_key_cache), + ) + return loaded + + +async def _live_team(auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None: + from litellm.proxy import proxy_server as server + + if auth.team_id is None: + return None + team_id: Final = auth.team_id + cached: Final = await _get_team_object_from_cache( + key=f"team_id:{team_id}", + user_api_key_cache=server.user_api_key_cache, + parent_otel_span=None, + ) + if cached is not None: + return cached + + row: Final = await TeamRepository(server.prisma_client).find_by_id(team_id, id_field="team_id") + if row is None: + return None + team: Final = LiteLLM_TeamTableCachedObj.model_validate(row.model_dump()) + if team.object_permission_id and not team.object_permission: + # The entry is written under the key the chat path reads, so it has to carry the same + # permission relation the chat path caches; a cache hit elsewhere must not see a team + # stripped of the permissions it was about to enforce. + try: + team.object_permission = await get_object_permission( + object_permission_id=team.object_permission_id, + prisma_client=server.prisma_client, + user_api_key_cache=server.user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=server.proxy_logging_obj, + ) + except Exception as exc: # noqa: BLE001 # same degradation as the chat path: cache the team without permissions and log it + verbose_proxy_logger.debug("Failed to load object_permission for Live team %s: %s", team_id, exc) + await _cache_team_object( + team_id=team_id, + team_table=team, + user_api_key_cache=server.user_api_key_cache, + proxy_logging_obj=server.proxy_logging_obj, + ) + return team + + +def _live_team_budget_configured(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | None) -> bool: + if team is None: + return False + team_budget_limits: Final = getattr(team, "budget_limits", None) + if _nonempty_limit_value(team_budget_limits): + return True + if any(getattr(team, field, None) is not None for field in ("rpm_limit", "tpm_limit", "max_budget")): + return True + if _nonempty_limit_value(getattr(team, "model_max_budget", None)): + return True + team_metadata_value: Final = getattr(team, "metadata", None) + team_metadata: Final = ( + team_metadata_value if team_metadata_value is not None else getattr(auth, "team_metadata", None) + ) + return _managed_constraints( + auth.model_copy(update=MappingProxyType({"team_metadata": team_metadata, "budget_limits": team_budget_limits})) + ) + + +async def _live_default_budget(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | None) -> LiteLLM_BudgetTable | None: + metadata_source: Final = ( + getattr(team, "metadata", None) if team is not None else getattr(auth, "team_metadata", None) + ) + default_id: Final = _MAPPING.validate_python(metadata_source or _EMPTY).get("team_member_budget_id") + if not isinstance(default_id, str) or auth.team_id is None or auth.user_id is None: + return None + from litellm.proxy import proxy_server as server + + async def load() -> LiteLLM_BudgetTable | None: + return await BudgetRepository(server.prisma_client).find_by_id(default_id, id_field="budget_id") + + return await _live_cached_object( + key=f"team_member_default_budget:{default_id}", + model_type=LiteLLM_BudgetTable, + load=load, + ) + + +async def _live_project(auth: UserAPIKeyAuth) -> LiteLLM_ProjectTableCachedObj | None: + from litellm.proxy import proxy_server as server + + if auth.project_id is None: + return None + return await get_project_object( + project_id=auth.project_id, + prisma_client=server.prisma_client, + user_api_key_cache=server.user_api_key_cache, + proxy_logging_obj=server.proxy_logging_obj, + ) + + +async def _live_project_budget_configured(auth: UserAPIKeyAuth, project: LiteLLM_ProjectTableCachedObj | None) -> bool: + if project is None: + return False + project_budget: Final = getattr(project, "litellm_budget_table", None) + if _live_budget_configured(project_budget, zero_is_limit=True): + return True + if _nonempty_limit_value(getattr(project, "model_rpm_limit", None)) or _nonempty_limit_value( + getattr(project, "model_tpm_limit", None) + ): + return True + project_metadata_value: Final = getattr(project, "metadata", None) + project_metadata: Final = ( + project_metadata_value if project_metadata_value is not None else getattr(auth, "project_metadata", None) + ) + return _managed_constraints(auth.model_copy(update=MappingProxyType({"project_metadata": project_metadata}))) + + +def _live_group_limits(row: object) -> LiteLLM_BudgetTable: + """The limit fields of the budget linked to one model access group row. + + The row arrives as a Prisma join, so the limits are read by name. A group with no linked + budget yields an empty budget table: it reads as no limit, which is what the gate needs, and + it stays cacheable so the group is not re-read on every request. + """ + budget: Final = getattr(row, "litellm_budget_table", None) + if budget is None: + return LiteLLM_BudgetTable() + return LiteLLM_BudgetTable.model_validate( + { # mutable-ok: field values are read from the joined row into a fresh validation mapping + field: getattr(budget, field, None) + for field in ("max_budget", "rpm_limit", "tpm_limit", "model_max_budget", "max_parallel_requests") + } + ) + + +async def _live_fetch_group_limits(groups: tuple[str, ...]) -> tuple[LiteLLM_BudgetTable, ...]: + """Fetch the linked budget of each group in chunked queries and cache one entry per group.""" + if not groups: + return () + from litellm.proxy import proxy_server as server + + # `find_many_in` cannot carry the budget join, so the group names are sliced here by hand. + table: Final = ModelAccessGroupBudgetRepository(server.prisma_client).table + unique_groups: Final = tuple(dict.fromkeys(groups)) + rows: Final = tuple( + [ + row + # comprehension-ok: one iteration per IN_LIST_CHUNK_SIZE slice of group names + for start in range(0, len(unique_groups), IN_LIST_CHUNK_SIZE) + for row in await table.find_many( + where={ # mutable-ok: Prisma serializes query filters from concrete dictionaries + "access_group_name": { # mutable-ok: Prisma serializes nested filters from concrete dictionaries + # bounded-ok: <= IN_LIST_CHUNK_SIZE (5,000) names, the loop slices groups by that size + "in": list(unique_groups[start : start + IN_LIST_CHUNK_SIZE]), + } + }, + include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries + ) + ] + ) + linked: Final = MappingProxyType({getattr(row, "access_group_name", None): _live_group_limits(row) for row in rows}) + limits: Final = tuple(linked.get(group) or LiteLLM_BudgetTable() for group in groups) + await asyncio.gather( + *( + server.user_api_key_cache.async_set_cache( + key=live_model_access_group_limits_cache_key(group), + value=limit, + model_type=LiteLLM_BudgetTable, + ttl=get_management_object_ttl(server.user_api_key_cache), + ) + for group, limit in zip(groups, limits) + ) + ) + return limits + + +async def _live_model_group_limits(groups: tuple[str, ...]) -> tuple[LiteLLM_BudgetTable, ...]: + """One cached budget entry per group, served from a single row batch on a cold miss.""" + from litellm.proxy import proxy_server as server + + cached: Final = await asyncio.gather( + *( + server.user_api_key_cache.async_get_cache( + key=live_model_access_group_limits_cache_key(group), + model_type=LiteLLM_BudgetTable, + ) + for group in groups + ) + ) + uncached: Final = tuple(group for group, entry in zip(groups, cached) if entry is None) + fetched: Final = MappingProxyType(dict(zip(uncached, await _live_fetch_group_limits(uncached)))) + return tuple(entry if entry is not None else fetched[group] for group, entry in zip(groups, cached)) + + +async def _live_model_group_budget_configured( + auth: UserAPIKeyAuth, + model: str | None, + team: LiteLLM_TeamTable | None, + project: LiteLLM_ProjectTableCachedObj | None, +) -> bool: + if model is None: + return False + from litellm.proxy import proxy_server as server + + matched_groups: Final = await collect_matched_model_access_groups( + model=model, + valid_token=auth, + team_object=team, + project_object=project, + llm_router=server.llm_router, + prisma_client=server.prisma_client, + user_api_key_cache=server.user_api_key_cache, + proxy_logging_obj=server.proxy_logging_obj, + strict_grant_lookup=True, + ) + if not matched_groups: + return False + # The shared group-budget helper flattens the row down to spend and max_budget, which would + # drop the rpm and tpm limits this gate exists to refuse, so the linked row is read in full. + limits: Final = await _live_model_group_limits(matched_groups) + return any(_live_budget_configured(limit, zero_is_limit=False) for limit in limits) + + +async def _managed_member_budget(auth: UserAPIKeyAuth, model: str | None = None) -> bool: + from litellm.proxy import proxy_server as server + + try: + if not _live_budget_scope_present(auth, model, server.llm_router): + return False + if server.prisma_client is None: + raise HTTPException(503, "Could not verify Live managed budgets") + membership: Final = await _live_team_membership(auth) + if _live_budget_configured(getattr(membership, "litellm_budget_table", None), zero_is_limit=True): + return True + team: Final = await _live_team(auth) + if _live_team_budget_configured(auth, team): + return True + default: Final = await _live_default_budget(auth, team) + if _live_budget_configured(default, zero_is_limit=False): + return True + project: Final = await _live_project(auth) + if await _live_project_budget_configured(auth, project): + return True + return await _live_model_group_budget_configured(auth, model, team, project) + except Exception as exc: + raise HTTPException(503, "Could not verify Live managed budgets") from exc + + +async def _authorize_delegation(body: Mapping[str, JsonValue], auth: UserAPIKeyAuth) -> None: + session: Final = body.get("session") + if not isinstance(session, Mapping): + return + delegation: Final = session.get("delegation") + if not isinstance(delegation, dict) or delegation.get("type") == "client": + return + responses: Final = delegation.get("responses") + if delegation.get("type") != "responses" and not isinstance(responses, dict): + return + model: Final = responses.get("model") if isinstance(responses, dict) else None + if _managed_constraints(auth) or await _managed_member_budget( + auth, + model if isinstance(model, str) else None, + ): + raise HTTPException( + 400, + "Managed Live delegation cannot enforce configured budgets or rate limits; use client delegation", + ) + if not isinstance(responses, dict): + if body.get("type") != "session.update" and _restricted_models(auth): + raise HTTPException(400, "Restricted keys require an explicit authorized delegation.responses.model") + return + if isinstance(model, str): + await _authorize(model, auth) + elif body.get("type") != "session.update" and _restricted_models(auth): + raise HTTPException(400, "Restricted keys require an explicit authorized delegation.responses.model") + + transport: Final = body.get("transport") + if isinstance(transport, dict) and transport.get("type") == "webrtc" and _restricted_models(auth): + client: Final = session.get("client") + channel: Final = client.get("data_channel") if isinstance(client, dict) else None + events: Final = channel.get("allowed_client_events") if isinstance(channel, dict) else None + if not isinstance(events, list) or any( + not isinstance(event, str) or event.strip() == "session.update" or "*" in event for event in events + ): + raise HTTPException( + 400, "Restricted keys must explicitly exclude session.update from WebRTC allowed_client_events" + ) + + +async def _authorize_fork_policy( + body: Mapping[str, JsonValue], source: LiveHandle | None, auth: UserAPIKeyAuth +) -> None: + if ( + source is not None + and source.policy.get("delegation") + and (_restricted_models(auth) or _managed_constraints(auth)) + ): + # Handles contain startup policy; later sideband or WebRTC updates can change the backend model. + session: Final = _policy_object(body.get("session", _EMPTY)) + delegation: Final = _policy_object(session.get("delegation") or _EMPTY) + responses: Final = _policy_object(delegation.get("responses") or _EMPTY) + if delegation.get("type") != "client" and not ( + delegation.get("type") == "responses" and isinstance(responses.get("model"), str) and responses.get("model") + ): + raise HTTPException( + 400, "Constrained-key forks require explicit client delegation or an authorized responses model" + ) + await _authorize_delegation(_policy_body(body, source), auth) + + +class _Prepared: + def __init__( + self, + processed: Mapping[str, object], + logger: Logging, + lease: RealtimeCallLease | None, + ownership: _BudgetOwnership | None = None, + ) -> None: + self.processed = processed + self.logger = logger + self.lease = lease + self.transferred = False + self.ownership = ownership + + def transfer(self) -> None: + self.transferred = True + if self.ownership is not None: + self.ownership.transferred = True + + +class _PrecallState: + def __init__(self) -> None: + self.prepared: _Prepared | None = None + self.lease: RealtimeCallLease | None = None + + +@asynccontextmanager +async def _precall( + request: Request, + auth: UserAPIKeyAuth, + model: str, + *, + attachment: object | None = None, + parallel_reserved: bool = False, + ownership: _BudgetOwnership | None = None, +) -> AsyncGenerator[_Prepared]: + from litellm.proxy import proxy_server as server + + limiter: Final = server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") + with isolated_request_stash(): + state: Final = _PrecallState() + signaling_auth: Final = auth.model_copy(update=MappingProxyType({"budget_reservation": None})) + try: + await _authorize(model, auth) + if isinstance(limiter, _PROXY_MaxParallelRequestsHandler) and ( + auth.max_parallel_requests is not None + or _MAPPING.validate_python(server.general_settings).get("global_max_parallel_requests") # pyright: ignore[reportUnknownMemberType] # legacy settings are validated here is not None + ): + raise HTTPException(400, "Live requires the V3 rate limiter") + payload: Final = _OBJECT.validate_json(await request.body()) + await _authorize_delegation(payload, auth) + data: Final = MappingProxyType( + {key: value for key, value in payload.items() if key in ("session", "transport")} + ) + with ( + realtime_call_attachment(attachment) if attachment is not None and parallel_reserved else nullcontext() + ): + processed, logger = await process_codex_request( + request, + _mutable( + MappingProxyType( + { + **data, + "model": model, + **(MappingProxyType({"websocket": attachment}) if attachment is not None else _EMPTY), + } + ) + ), + signaling_auth, + model, + "_arealtime" if attachment is not None else "arealtime_calls", + ) + if attachment is None and isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + state.lease = limiter.transfer_realtime_call_slot(processed) + if state.lease is not None: + state.lease.start() + if not await state.lease.renew(): + raise HTTPException(503, "Live quota reservation was lost") + await _authorize_delegation(_processed_body(payload, processed), auth) + state.prepared = _Prepared(processed, logger, state.lease, ownership) + yield state.prepared + finally: + try: + if state.prepared is None or not state.prepared.transferred: + if state.lease is not None: + await state.lease.close() + if ownership is None: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + finally: + if isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + await limiter.async_post_call_failure_hook( # pyright: ignore[reportUnknownMemberType] # hook accepts legacy mutable request data + request_data=_mutable(_EMPTY), + original_exception=Exception("Live signaling complete"), + user_api_key_dict=signaling_auth, + ) + + +async def _supervise( + request: Request, handle: LiveHandle, auth: UserAPIKeyAuth, logger: Logging, lease: RealtimeCallLease | None +) -> RealTimeStreaming: + with isolated_request_stash(): + return await _start_supervisor(request, handle, auth, lease) + + +async def _start_supervisor( + request: Request, handle: LiveHandle, auth: UserAPIKeyAuth, lease: RealtimeCallLease | None +) -> RealTimeStreaming: + deployment: Final = _pinned(handle) + transport: Final = LiveTransport(deployment, request.headers) + state: Final = _ConnectionState() + try: + processed, logger = await process_codex_request( + _request(request, MappingProxyType({"model": handle.alias})), + _mutable(MappingProxyType({"model": handle.alias})), + auth, + handle.alias, + "_arealtime", + internal_realtime_observer=True, + ) + import litellm + + metadata: Final = _mutable( + MappingProxyType( + { + **_MAPPING.validate_python(processed.get("litellm_metadata") or _EMPTY), + **( + MappingProxyType( + { + "model_info": _mutable( + MappingProxyType( + { + **litellm.get_model_info(model=deployment.model_id), + "id": deployment.model_id, + } + ) + ) + } + ) + if deployment.model_id is not None + else _EMPTY + ), + } + ) + ) + pinned: Final = _mutable(MappingProxyType({**processed, "litellm_metadata": metadata})) + logger.update_from_kwargs( # pyright: ignore[reportUnknownMemberType] # logging accepts legacy mutable provider parameters + kwargs=pinned, + model=deployment.model, + user=None, + optional_params=_mutable(_EMPTY), + litellm_params=_mutable( + MappingProxyType( + { + **_MAPPING.validate_python(logger.litellm_params), # pyright: ignore[reportUnknownMemberType] # legacy logging parameters are validated here + "litellm_metadata": metadata, + "arealtime": True, + } + ) + ), + custom_llm_provider=deployment.provider, + ) + state.connection = await transport.connect(live_session_path(handle.session_id, "attach")) + + async def receive() -> Message: + return _mutable(MappingProxyType({"type": "websocket.disconnect", "code": 1000})) + + async def send(message: Message) -> None: + return None + + frontend: Final = WebSocket( + _mutable(MappingProxyType({**request.scope, "type": "websocket"})), receive=receive, send=send + ) + stream: Final = RealTimeStreaming( + frontend, + state.connection, + logger, + model=_pinned(handle).model, + user_api_key_dict=auth, + live_initialization_seconds=handle.initialization_seconds, + ) + + async def hangup() -> None: + result: Final = await transport.request("POST", live_session_path(handle.session_id, "hangup")) + result.raise_for_status() + + supervisor: Final = CallSupervisor( + state.connection, + stream, + logger, + auth, + hangup, + force_close_call=hangup, + lease=lease, + terminal_usage_required=True, + connected_ready=True, + ) + await CALL_SUPERVISORS.start(supervisor) + return stream + except BaseException: + from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, # pyright: ignore[reportUnknownVariableType] # reservation helper has a legacy dict contract + ) + + try: + result: Final = await transport.request("POST", live_session_path(handle.session_id, "hangup")) + result.raise_for_status() + except Exception: # noqa: BLE001 # whatever hung up failed, the reservation must go back + await invalidate_budget_reservation_counters(budget_reservation=auth.budget_reservation) + finally: + if state.connection is not None: + await state.connection.close() + raise + + +def _processed_body(body: Mapping[str, JsonValue], processed: Mapping[str, object]) -> Mapping[str, JsonValue]: + return _object( + MappingProxyType( + { + **body, + **_object( + MappingProxyType( + {key: value for key, value in processed.items() if key in ("session", "transport")} + ) + ), + } + ) + ) + + +def _provider_body(body: Mapping[str, JsonValue], model: str) -> Mapping[str, JsonValue]: + session: Final = _object(body.get("session", _EMPTY)) + return _object(MappingProxyType({**body, "session": MappingProxyType({**session, "model": model})})) + + +def _response(response: httpx.Response, handle: LiveHandle | None = None) -> Response: + if handle is None: + return Response( + response.content, + status_code=response.status_code, + headers=MappingProxyType( + { + key: value + for key, value in response.headers.items() + if key.lower() in ("content-type", "content-disposition", "content-range", "accept-ranges") + } + ), + ) + return Response( + _encode_json( + rewrite_session_ids(_OBJECT.validate_json(response.content), handle.session_id, encode_session(handle)) + ), + status_code=response.status_code, + media_type="application/json", + ) + + +async def _create(request: Request, token: str | None = None) -> Response: + body: Final = await _body(request) + requested: Final = None if token else _session_model(body) + auth_body: Final = _EMPTY if token or requested is None else MappingProxyType({**body, "model": requested}) + auth: Final = await _auth(_request(request, auth_body, requested)) + async with _budget_scope(auth) as ownership: + source: Final = decode_session(token, _owner(auth)) if token else None + model: Final = _session_model(body, source.alias if source else None) + if source is not None: + await _reauth(ownership, request, body, model) + async with _precall( + _request(request, MappingProxyType({**body, "model": model})), ownership.auth, model, ownership=ownership + ) as prepared: + await _authorize_fork_policy(_processed_body(body, prepared.processed), source, ownership.auth) + deployment: Final = ( + _validate_pinned_deployment(source) if source else await _deployment(model, prepared.processed) + ) + transport: Final = LiveTransport(deployment, request.headers) + path: Final = live_session_path(source.session_id, "fork") if source else "live/sessions" + response: Final = await transport.request( + "POST", + path, + body=_processed_body(body, prepared.processed) + if source + else _provider_body(_processed_body(body, prepared.processed), deployment.model), + ) + if response.is_error: + return _response(response) + handle: Final = _new_handle( + _session_id(_OBJECT.validate_json(response.content)), + model, + deployment, + ownership.auth, + prepared.lease, + initialization_seconds=15, + policy=_session_policy(_processed_body(body, prepared.processed), source), + ) + await _supervise(request, handle, ownership.auth, prepared.logger, prepared.lease) + prepared.transfer() + return _response(response, handle) + + +@_routes.post("/sessions") +async def create_live_session(request: Request) -> Response: + return await _create(request) + + +@_routes.post("/sessions/{session_id}/fork") +async def fork_live_session(request: Request, session_id: str) -> Response: + return await _create(request, session_id) + + +@_routes.get("/sessions/{session_id}/content") +@_routes.post("/sessions/{session_id}/accept") +@_routes.post("/sessions/{session_id}/reject") +@_routes.post("/sessions/{session_id}/refer") +@_routes.post("/sessions/{session_id}/hangup") +async def control_live_session(request: Request, session_id: str) -> Response: + body: Final = await _body(request) + auth: Final = await _auth(_request(request, _EMPTY)) + async with _budget_scope(auth) as ownership: + operation: Final = TypeAdapter[LiveOperation](LiveOperation).validate_python( + request.url.path.rsplit("/", 1)[-1] + ) + if not session_id.startswith(_PREFIX): + return await _incoming_sip(request, session_id, operation, body, auth, ownership) + handle: Final = decode_session(session_id, _owner(auth)) + await _reauth(ownership, request, body, handle.alias) + async with _precall( + _request(request, MappingProxyType({"model": handle.alias})), + ownership.auth, + handle.alias, + attachment=request, + parallel_reserved=handle.parallel_reserved, + ownership=ownership, + ): + response: Final = await LiveTransport(_pinned(handle), request.headers).request( + request.method, live_session_path(handle.session_id, operation), body=body if body else None + ) + return _response(response) + + +async def _incoming_sip( + request: Request, + session_id: str, + operation: LiveOperation, + body: Mapping[str, JsonValue], + auth: UserAPIKeyAuth, + ownership: _BudgetOwnership, +) -> Response: + from litellm.proxy import proxy_server as server + + if auth.user_role != LitellmUserRoles.PROXY_ADMIN or operation not in ("accept", "reject"): + raise HTTPException(403, "Incoming SIP enrollment requires a proxy administrator") + alias: Final = request.headers.get("x-litellm-live-model") + if not alias: + raise HTTPException(400, "Incoming SIP requires x-litellm-live-model identifying one deployment") + configured: Final = tuple( + item + for item in TypeAdapter(tuple[Mapping[str, object], ...]).validate_python( + getattr(server, "llm_model_list", ()) or () + ) + if item.get("model_name") == alias + ) + if len(configured) != 1: + raise HTTPException(400, "Incoming SIP requires a model alias with exactly one deployment") + live_session_path(session_id, "attach") + if operation == "accept" and _session_model(body) != alias: + raise HTTPException(400, "session.model must match x-litellm-live-model") + await _reauth(ownership, request, body, alias) + async with _precall( + _request(request, MappingProxyType({**body, "model": alias})), ownership.auth, alias, ownership=ownership + ) as prepared: + deployment: Final = await _deployment(alias, prepared.processed) + response: Final = await LiveTransport(deployment, request.headers).request( + "POST", + live_session_path(session_id, operation), + body=_provider_body(_processed_body(body, prepared.processed), deployment.model) + if operation == "accept" + else body, + ) + if response.is_error or operation == "reject": + return _response(response) + handle: Final = _new_handle( + session_id, + alias, + deployment, + ownership.auth, + prepared.lease, + policy=_session_policy(_processed_body(body, prepared.processed), None), + ) + await _supervise(request, handle, ownership.auth, prepared.logger, prepared.lease) + prepared.transfer() + return Response( + response.content, + status_code=response.status_code, + headers=MappingProxyType({"x-litellm-live-session-id": encode_session(handle)}), + ) + + +class _PublicSocket: + def __init__( + self, + websocket: WebSocket, + handle: LiveHandle, + public_id: str, + auth: UserAPIKeyAuth, + observer: RealTimeStreaming | None = None, + ) -> None: + self.websocket = websocket + self.handle = handle + self.public_id = public_id + self.auth = auth + self.observer = observer + self.scope = websocket.scope + self.headers = websocket.headers + + async def send_text(self, data: str) -> None: + if self.observer is not None: + self.observer.store_message(data) # pyright: ignore[reportUnknownMemberType] # stream also accepts legacy dict events + await self.websocket.send_text( + _encode_json(rewrite_session_ids(_OBJECT.validate_json(data), self.handle.session_id, self.public_id)) + ) + + async def receive_text(self) -> str: + data: Final = await self.websocket.receive_text() + payload: Final = _OBJECT.validate_json(data) + if payload.get("type") == "session.start": + raise HTTPException(400, "Session has already started") + session: Final = payload.get("session") + if isinstance(session, Mapping) and "model" in session: + raise HTTPException(400, "Session model cannot change") + await _authorize_delegation(payload, self.auth) + return _encode_json(rewrite_session_ids(payload, self.public_id, self.handle.session_id)) + + async def close(self, code: int = 1000, reason: str | None = None) -> None: + await self.websocket.close(code=code, reason=reason) + + +class _StartupEvents: + def __init__(self) -> None: + self.messages: tuple[str, ...] = () + self.size = 0 + + def store(self, event: Mapping[str, JsonValue]) -> None: + message: Final = _encode_json(event) + if len(self.messages) >= 128 or self.size + len(message) > 8 * 1024 * 1024: + raise HTTPException(502, "Upstream exceeded the Live startup event limit") + self.messages = (*self.messages, message) + self.size += len(message) + + +async def _wait_started( + connection: "ClientConnection", websocket: WebSocket, startup: _StartupEvents | None = None +) -> Mapping[str, JsonValue]: + async def receive_started() -> Mapping[str, JsonValue]: + while True: + event: Mapping[str, JsonValue] = _OBJECT.validate_json(await connection.recv()) + if event.get("type") == "session.started": + return event + if startup is not None: + startup.store(event) + await websocket.send_json(event) + if event.get("type") in ("error", "session.closed"): + raise HTTPException(502, "Upstream did not start the session") + + return await asyncio.wait_for(receive_started(), 20) + + +@_routes.websocket("/sessions") +@_routes.websocket("/sessions/{session_id}/attach") +@_routes.websocket("/sessions/{session_id}/fork") +async def websocket_live_session(websocket: WebSocket, session_id: str | None = None) -> None: + state: Final = _ConnectionState() + try: + api_key: Final = get_websocket_api_key(websocket) + if not api_key: + raise HTTPException(403, "API key required") + auth_request: Final = _request(websocket, _EMPTY) + inbound_headers: Final = TypeAdapter(tuple[tuple[bytes, bytes], ...]).validate_python( + websocket.scope["headers"] + ) + auth_request.scope["headers"] = ( + *(item for item in inbound_headers if item[0].lower() != b"authorization"), + (b"authorization", f"Bearer {api_key}".encode()), + ) + auth: Final = await _auth(auth_request) + async with _budget_scope(auth) as ownership: + source: Final = decode_session(session_id, _owner(auth)) if session_id else None + attached: Final = source is not None and websocket.url.path.endswith("/attach") + await websocket.accept() + first: Final = ( + _EMPTY if attached else _OBJECT.validate_json(await asyncio.wait_for(websocket.receive_text(), 20)) + ) + if not attached and first.get("type") != "session.start": + raise HTTPException(400, "First message must be session.start") + model: Final = ( + source.alias + if attached and source is not None + else _session_model(first, source.alias if source else None) + ) + await _reauth(ownership, auth_request, first, model) + async with _precall( + _request(websocket, MappingProxyType({**first, "model": model})), + ownership.auth, + model, + attachment=websocket if attached else None, + parallel_reserved=source.parallel_reserved if attached and source is not None else False, + ownership=ownership, + ) as prepared: + if attached: + await _authorize_delegation( + _policy_body(_processed_body(first, prepared.processed), source), ownership.auth + ) + else: + await _authorize_fork_policy(_processed_body(first, prepared.processed), source, ownership.auth) + deployment: Final = ( + _validate_pinned_deployment(source) if source else await _deployment(model, prepared.processed) + ) + path: Final = ( + live_session_path(source.session_id, "attach" if attached else "fork") + if source + else "live/sessions" + ) + state.connection = await LiveTransport(deployment, websocket.headers).connect(path) + + async def start_session() -> tuple[ + LiveHandle, RealTimeStreaming | None, Mapping[str, JsonValue] | None + ]: + if attached and source is not None: + return source, None, None + if state.connection is None: + raise RuntimeError("Live connection was not established") + await state.connection.send( + _encode_json( + _processed_body(first, prepared.processed) + if source + else _provider_body(_processed_body(first, prepared.processed), deployment.model) + ) + ) + startup: Final = _StartupEvents() + initial: Final = await _wait_started(state.connection, websocket, startup) + handle: Final = _new_handle( + _session_id(initial), + model, + deployment, + ownership.auth, + prepared.lease, + policy=_session_policy(_processed_body(first, prepared.processed), source), + ) + observer: Final = await _supervise( + _request(websocket, MappingProxyType({"model": model})), + handle, + ownership.auth, + prepared.logger, + prepared.lease, + ) + for buffered in startup.messages: + observer.store_message(buffered) # pyright: ignore[reportUnknownMemberType] # stream also accepts legacy dict events + prepared.transfer() + return handle, observer, initial + + handle, observer, initial = await start_session() + public_id: Final = session_id if attached and session_id is not None else encode_session(handle) + if initial is not None: + await websocket.send_json(rewrite_session_ids(initial, handle.session_id, public_id)) + frontend: Final = _PublicSocket(websocket, handle, public_id, ownership.auth, observer) + stream: Final = RealTimeStreaming( + frontend, + state.connection, + prepared.logger, + model=deployment.model, + user_api_key_dict=auth, + request_data=_mutable(prepared.processed), + account_usage=False, + ) + await stream.bidirectional_forward() + except (HTTPException, ValueError, WebSocketDisconnect, asyncio.TimeoutError): + try: + await websocket.close(code=1008, reason="Live session rejected") + except RuntimeError: + # The peer may have closed the socket before the rejection response. + return + except Exception: # noqa: BLE001 # an unexpected failure still gets the client a 1011 close + try: + await websocket.close(code=1011, reason="Live upstream connection failed") + except RuntimeError: + # The peer may have closed the socket before the failure response. + return + finally: + if state.connection is not None: + await state.connection.close() + + +router: Final = APIRouter() +for _prefix in ("/v1/live", "/live", "/openai/v1/live"): + router.include_router(_routes, prefix=_prefix) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index c0e36e6e172..92b4182ae7a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2352,7 +2352,9 @@ class ProxyLogging: data: None, call_type: CallTypesLiteral, guardrails_only: bool = False, + *, skip_guardrails: bool = False, + internal_realtime_observer: bool = False, ) -> None: pass @@ -2363,7 +2365,9 @@ class ProxyLogging: data: dict, call_type: CallTypesLiteral, guardrails_only: bool = False, + *, skip_guardrails: bool = False, + internal_realtime_observer: bool = False, ) -> dict: pass @@ -2373,7 +2377,9 @@ class ProxyLogging: data: dict | None, call_type: CallTypesLiteral, guardrails_only: bool = False, + *, skip_guardrails: 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. @@ -2481,6 +2487,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: @@ -4611,7 +4621,7 @@ class PrismaClient: def hash_token(self, token: str): # Hash the string using SHA-256 - hashed_token: Final = hashlib.sha256(token.encode()).hexdigest() + hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token @@ -7115,7 +7125,7 @@ def hash_token(token: str): import hashlib # Hash the string using SHA-256 - hashed_token: Final = hashlib.sha256(token.encode()).hexdigest() + hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 28814741852..6dfc9196d01 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -6,6 +6,8 @@ from collections.abc import Mapping from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, cast +from pydantic import TypeAdapter + import litellm from litellm.constants import ( AZURE_OPENAI_AUDIO_PROVIDERS, @@ -89,6 +91,7 @@ def _get_realtime_http_provider_config( dynamic_api_base: str | None, dynamic_api_key: str | None, litellm_params: GenericLiteLLMParams, + is_call: bool = False, ) -> tuple["BaseRealtimeHTTPConfig | None", str, str]: """ Return (provider_config, resolved_api_base, resolved_api_key) for the @@ -106,13 +109,15 @@ def _get_realtime_http_provider_config( provider_config = ProviderConfigManager.get_provider_realtime_http_config( model="", provider=LlmProviders(custom_llm_provider), + params=litellm_params, + is_call=is_call, ) raw_api_base: Final = dynamic_api_base or litellm_params.api_base raw_api_key: Final = dynamic_api_key or litellm_params.api_key if provider_config is not None: - resolved_api_base = provider_config.get_api_base(api_base=raw_api_base) + resolved_api_base = provider_config.resolve_api_base(litellm_params.api_base, dynamic_api_base) resolved_api_key = provider_config.get_api_key(api_key=raw_api_key) else: # Fallback for providers without a dedicated HTTP config (treated as OpenAI-compatible). @@ -285,9 +290,18 @@ async def arealtime_calls( dynamic_api_base=dynamic_api_base, dynamic_api_key=dynamic_api_key, litellm_params=litellm_params, + is_call=True, ) if session is not None: session = _with_resolved_session_model(session, model_name) + supplied_headers: Final = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python( + kwargs.get("extra_headers") + ) + call_headers: Final = ( + provider_config.get_realtime_calls_extra_headers(supplied_headers) + if provider_config is not None + else supplied_headers + ) litellm_logging_obj.update_from_kwargs( kwargs=kwargs, model=model_name, @@ -295,7 +309,7 @@ async def arealtime_calls( litellm_params={"api_base": resolved_api_base}, custom_llm_provider=custom_llm_provider, ) - return await base_llm_http_handler.async_realtime_calls_handler( + response: Final = await base_llm_http_handler.async_realtime_calls_handler( api_base=resolved_api_base, openai_ephemeral_key=openai_ephemeral_key, sdp_body=sdp_body, @@ -304,10 +318,17 @@ async def arealtime_calls( provider_config=provider_config, model=model_name, session_config=session, - extra_headers=kwargs.get("extra_headers"), + extra_headers=call_headers, client=kwargs.get("client"), api_version=litellm_params.api_version, ) + return ( + provider_config.transform_realtime_calls_response( + response, model_name, litellm_logging_obj.get_router_model_id(), call_headers + ) + if provider_config is not None + else response + ) async def vertex_access_token_resolver( @@ -363,8 +384,10 @@ async def _arealtime( For PROXY use only. """ - headers = cast(dict | None, kwargs.get("headers")) - extra_headers: Final = cast(dict | None, kwargs.get("extra_headers")) + headers = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(kwargs.get("headers")) + extra_headers: Final = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python( + kwargs.get("extra_headers") + ) if headers is None: headers = {} if extra_headers is not None: @@ -404,7 +427,27 @@ async def _arealtime( model=model, provider=LlmProviders(_custom_llm_provider), ) - if provider_config is not None: + provider_handler: Final = ( + ProviderConfigManager.get_provider_realtime_handler( + LlmProviders(_custom_llm_provider), litellm_params, lambda: websocket.headers, headers + ) + if _custom_llm_provider in LlmProviders._member_map_.values() + else None + ) + if provider_handler is not None: + user_api_key_dict: Final = TypeAdapter[object](object).validate_python(kwargs.get("user_api_key_dict")) + await provider_handler.async_realtime( + model=model, + websocket=websocket, + logging_obj=litellm_logging_obj, + api_base=api_base or None, + api_key=api_key, + timeout=timeout, + query_params=query_params, + user_api_key_dict=user_api_key_dict, + litellm_metadata=_build_litellm_metadata(kwargs), + ) + elif provider_config is not None: await base_llm_http_handler.async_realtime( model=model, websocket=websocket, diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 206c0f78ed8..ff8bc198f27 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -31,7 +31,7 @@ class NativeTraceStorage: def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> Future[None]: ... def insert_rows(self, table: str, rows: Sequence[Mapping[str, JsonValue]]) -> Future[None]: ... def lens_query(self, name: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... - def query(self, query: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... + def query(self, sql: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... @final class NativeDiagnosticProcessor: diff --git a/litellm/types/images/main.py b/litellm/types/images/main.py index 5d80135a8a1..603b28e1081 100644 --- a/litellm/types/images/main.py +++ b/litellm/types/images/main.py @@ -16,7 +16,7 @@ class ImageEditOptionalRequestParams(TypedDict, total=False): input_fidelity: Literal["high", "low"] | None mask: str | None n: int | None - quality: Literal["high", "medium", "low", "standard", "auto"] | None + quality: Literal["high", "medium", "low", "standard", "auto", "xhigh", "max"] | None response_format: Literal["url", "b64_json"] | None size: str | None user: str | None diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 20214078852..afb7e3d8204 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -87,6 +87,11 @@ class ProviderConnection: litellm_credential_name: str | None = None configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None use_xai_oauth: bool | None = None + # ChatGPT OAuth deployment options; read from litellm_params by the ChatGPT + # adapters and listed as owned so they are never swept into extra_body. + chatgpt_auth_profile: str | None = None + chatgpt_token_dir: str | None = None + chatgpt_auth_file: str | None = None @dataclass(frozen=True, slots=True, kw_only=True) diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 99ab5920c4f..c7f9b22a16f 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -93,7 +93,7 @@ from litellm.types.responses.main import ( from .base import CachedTokensDetails -FileContent = IO[bytes] | bytes | PathLike +FileContent = IO[bytes] | bytes | PathLike[str] FileTypes = ( # file (or bytes) @@ -2081,6 +2081,16 @@ class OpenAIRealtimeStreamResponseBaseObject(TypedDict): type: str +class OpenAIRealtimeSessionClosed(TypedDict): + type: ReadOnly[Literal["session.closed", "session.usage.updated", "litellm.live.initialization"]] + usage: ReadOnly[Mapping[str, object]] + + +class OpenAILiveResponseEvent(TypedDict): + type: ReadOnly[Literal["response.event"]] + event: ReadOnly[Mapping[str, object]] + + class OpenAIRealtimeConversationObject(TypedDict, total=False): id: str object: Required[Literal["realtime.conversation"]] @@ -2341,6 +2351,8 @@ class OpenAIRealtimeEventTypes(Enum): OpenAIRealtimeEvents = ( OpenAIRealtimeStreamResponseBaseObject + | OpenAIRealtimeSessionClosed + | OpenAILiveResponseEvent | OpenAIRealtimeStreamSessionEvents | OpenAIRealtimeStreamResponseOutputItemAdded | OpenAIRealtimeResponseContentPartAdded @@ -2371,6 +2383,8 @@ class ImageGenerationRequestQuality(str, Enum): LOW = "low" MEDIUM = "medium" HIGH = "high" + XHIGH = "xhigh" + MAX = "max" AUTO = "auto" STANDARD = "standard" HD = "hd" diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index 855cccc8ddd..f48aeb5560e 100644 --- a/litellm/types/realtime.py +++ b/litellm/types/realtime.py @@ -1,6 +1,6 @@ from typing import Any, Literal -from pydantic import BaseModel +from pydantic import AliasChoices, BaseModel, Field from typing_extensions import ReadOnly, TypedDict from .llms.openai import ( @@ -12,6 +12,16 @@ from .llms.openai import ( ALL_DELTA_TYPES = Literal["text", "audio"] +class LiveSessionDurationUsage(BaseModel): + duration: float = Field( + strict=True, ge=0, allow_inf_nan=False, validation_alias=AliasChoices("seconds", "audio_duration_ms") + ) + + +class LiveSessionUsageEvent(BaseModel): + usage: LiveSessionDurationUsage + + class RealtimeResponseTransformInput(TypedDict): session_configuration_request: str | None current_output_item_id: ( @@ -49,6 +59,7 @@ class RealtimeModalityResponseTransformOutput(TypedDict): class RealtimeQueryParams(TypedDict, total=False): model: str intent: str | None + call_id: ReadOnly[str] # Add more fields as needed diff --git a/litellm/utils.py b/litellm/utils.py index b0a7e4f1a68..c144992e93d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -429,6 +429,7 @@ if TYPE_CHECKING: ) from litellm.llms.cohere.common_utils import CohereModelInfo from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + from litellm.llms.openai.realtime.handler import OpenAIRealtime from litellm.proxy._types import AllowedModelRegion from litellm.router_utils.get_retry_from_policy import ( get_num_retries_from_retry_policy, @@ -444,7 +445,7 @@ if TYPE_CHECKING: ChatCompletionToolCallFunctionChunk, ) from litellm.types.rerank import RerankResponse - from litellm.types.router import LiteLLM_Params + from litellm.types.router import GenericLiteLLMParams, LiteLLM_Params from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.base_llm.completion.transformation import BaseTextCompletionConfig @@ -6199,6 +6200,7 @@ def _get_model_info_helper( input_cost_per_second=_model_info.get("input_cost_per_second", None), input_cost_per_audio_token=_model_info.get("input_cost_per_audio_token", None), input_cost_per_image_token=_model_info.get("input_cost_per_image_token", None), + cache_read_input_image_token_cost=_model_info.get("cache_read_input_image_token_cost", None), input_cost_per_video_token=_model_info.get("input_cost_per_video_token", None), input_cost_per_audio_token_batches=_model_info.get("input_cost_per_audio_token_batches", None), input_cost_per_image_token_batches=_model_info.get("input_cost_per_image_token_batches", None), @@ -9527,6 +9529,10 @@ class ProviderConfigManager: model: str, provider: LlmProviders, ) -> BaseImageGenerationConfig | None: + if LlmProviders.CHATGPT == provider: + from litellm.llms.chatgpt.images import ChatGPTImageGenerationConfig + + return ChatGPTImageGenerationConfig() if LlmProviders.OPENAI == provider: from litellm.llms.openai.image_generation import ( get_openai_image_generation_config, @@ -9707,16 +9713,36 @@ class ProviderConfigManager: return MetaRealtimeConfig() return None + @staticmethod + def get_provider_realtime_handler( + provider: LlmProviders, + params: GenericLiteLLMParams, + get_headers: Callable[[], Mapping[str, str]], + extra_headers: Mapping[str, object] | None = None, + ) -> OpenAIRealtime | None: + if provider == LlmProviders.CHATGPT: + from litellm.llms.chatgpt.realtime import ChatGPTRealtime + + return ChatGPTRealtime(params, get_headers(), extra_headers) + return None + @staticmethod def get_provider_realtime_http_config( model: str, provider: LlmProviders, + params: GenericLiteLLMParams | None = None, + is_call: bool = False, ) -> BaseRealtimeHTTPConfig | None: """ Return the HTTP transformation config for realtime HTTP endpoints (POST /realtime/client_secrets and POST /realtime/calls). """ + if LlmProviders.CHATGPT == provider: + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + from litellm.types.router import GenericLiteLLMParams + + return ChatGPTRealtimeHTTPConfig(params or GenericLiteLLMParams(), use_codex_backend=is_call) if LlmProviders.OPENAI == provider: from litellm.llms.openai.realtime.http_transformation import ( OpenAIRealtimeHTTPConfig, @@ -9736,6 +9762,10 @@ class ProviderConfigManager: model: str, provider: LlmProviders, ) -> BaseImageEditConfig | None: + if LlmProviders.CHATGPT == provider: + from litellm.llms.chatgpt.images import ChatGPTImageEditConfig + + return ChatGPTImageEditConfig() if LlmProviders.OPENAI == provider: from litellm.llms.openai.image_edit import get_openai_image_edit_config diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 2a01c4fe862..98ff53ec904 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -31156,6 +31156,16 @@ "max_tokens": 8191, "mode": "embedding" }, + "chatgpt/gpt-live-1-codex": { + "litellm_provider": "chatgpt", + "mode": "realtime", + "supported_endpoints": [ + "/v1/realtime/calls", + "/v1/live" + ], + "supports_audio_input": true, + "supports_audio_output": true + }, "chatgpt/gpt-5.5": { "litellm_provider": "chatgpt", "source": "https://platform.openai.com/docs/models/gpt-5.5", diff --git a/tests/local_testing/test_realtime_call_redis.py b/tests/local_testing/test_realtime_call_redis.py new file mode 100644 index 00000000000..d55f3488409 --- /dev/null +++ b/tests/local_testing/test_realtime_call_redis.py @@ -0,0 +1,102 @@ +import asyncio +import os +from contextlib import AsyncExitStack +from datetime import datetime +from uuid import uuid4 + +import pytest +import pytest_asyncio + +import litellm +from litellm.caching.redis_cache import RedisCache +from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler +from litellm.proxy.hooks.parallel_request_limiter_v3 import PARALLEL_REQUEST_SLOT_TTL_SECONDS +from litellm.proxy.utils import InternalUsageCache + + +@pytest_asyncio.fixture(loop_scope="function") +async def isolated_test_redis(monkeypatch): + raw_port = os.environ.get("LITELLM_TEST_REDIS_PORT", "") + if not raw_port.isdecimal() or not 1 <= int(raw_port) <= 65535: + pytest.fail("Set LITELLM_TEST_REDIS_PORT to an isolated Redis server's loopback port") + for name in tuple(os.environ): + if name.startswith("REDIS_"): + monkeypatch.delenv(name) + namespace = f"litellm-lua-test-{uuid4().hex}" + cache = RedisCache( + host="127.0.0.1", + port=int(raw_port), + namespace=namespace, + client_name=namespace, + socket_timeout=2, + socket_connect_timeout=2, + ) + async with AsyncExitStack() as cleanup: + cleanup.callback(cache.redis_client.close) + cleanup.push_async_callback(cache.async_redis_conn_pool.disconnect) + client = cache.init_async_client() + cleanup.push_async_callback(cache.async_redis_conn_pool.disconnect) + cleanup.push_async_callback(client.aclose) + cleanup.callback(litellm.in_memory_llm_clients_cache.delete_cache, cache._get_async_client_cache_key()) + try: + await client.ping() + yield cache + finally: + async for key in client.scan_iter(match=f"{namespace}:*"): + await client.delete(key) + + +@pytest.mark.asyncio +async def test_concurrent_realtime_releases_update_redis_without_lost_decrement(isolated_test_redis): + remote = isolated_test_redis + first_cache, second_cache = DualCache(redis_cache=remote), DualCache(redis_cache=remote) + first, second = (_PROXY_MaxParallelRequestsHandler(InternalUsageCache(c)) for c in (first_cache, second_cache)) + auth = UserAPIKeyAuth(api_key="concurrent-key", max_parallel_requests=2) + first_data, second_data = {"model": "test"}, {"model": "test"} + first.begin_realtime_attachment(first_data) + second.begin_realtime_attachment(second_data) + await first.async_pre_call_hook(auth, first_cache, first_data, "_arealtime") + await second.async_pre_call_hook(auth, second_cache, second_data, "_arealtime") + key = f"concurrent-key::{datetime.now().strftime('%Y-%m-%d-%H-%M')}::request_count" + counter = {"current_requests": 2, "current_rpm": 2, "current_tpm": 17} + await first_cache.async_set_cache(key, counter) + await second_cache.async_set_cache(key, counter, local_only=True) + remote.redis_client.pexpire(remote.check_and_fix_namespace(key), 15000) + await asyncio.gather( + first.async_release_realtime_attachment(first_data, auth), + second.async_release_realtime_attachment(second_data, auth), + ) + expected = {"current_requests": 0, "current_rpm": 2, "current_tpm": 17} + assert await remote.async_get_cache(key) == expected + assert 0 < remote.redis_client.pttl(remote.check_and_fix_namespace(key)) <= 15000 + assert await first_cache.async_get_cache(key) == expected + assert await second_cache.async_get_cache(key) == expected + await first_cache.async_set_cache("missing", counter, local_only=True) + await first._release_realtime_counter("missing") + assert await remote.async_get_cache("missing") is None + assert await first_cache.async_get_cache("missing", local_only=True) is None + + +@pytest.mark.asyncio +async def test_realtime_lease_redis_renewal_is_atomic_and_does_not_resurrect(isolated_test_redis): + from litellm.proxy.hooks.parallel_request_limiter_v3 import PARALLEL_RENEW_SCRIPT + + client = isolated_test_redis.init_async_client() + first_key = isolated_test_redis.check_and_fix_namespace("first") + second_key = isolated_test_redis.check_and_fix_namespace("second") + now = (await client.time())[0] + await client.zadd(first_key, {"owner": now - 10, "other": now}) + await client.zadd(second_key, {"owner": now - PARALLEL_REQUEST_SLOT_TTL_SECONDS}) + renew = client.register_script(PARALLEL_RENEW_SCRIPT) + assert await renew(keys=[first_key, second_key], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] + assert await client.zscore(first_key, "owner") == now - 10 + await client.zadd(second_key, {"owner": now - 10}) + assert await renew(keys=[first_key, second_key], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [1] + assert await client.zscore(first_key, "owner") >= now + assert await client.ttl(first_key) > PARALLEL_REQUEST_SLOT_TTL_SECONDS - 10 + await client.zrem(second_key, "owner") + assert await renew(keys=[first_key, second_key], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] + assert await client.zscore(second_key, "owner") is None + assert await client.zscore(first_key, "other") == now diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py index eee985f0aca..f7c734bd6f4 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -365,6 +365,9 @@ async def test_unknown_invocation_target_leaves_billing_unset(monkeypatch: pytes ("/v1/realtime", "GET", True), ("/v1/realtime", "POST", False), ("/v1/realtime/client_secrets", "POST", False), + ("/live", "POST", False), + ("/v1/live", "POST", False), + ("/live/sessions/session/accept", "POST", False), ("/mcp/tools/call", "POST", True), ("/a2a/target/message/send", "POST", True), ("/v1/a2a/target/message/send", "POST", True), @@ -573,7 +576,7 @@ def test_registered_inference_routes_have_an_explicit_managed_access_decision(ro "/videos", "/batches", "/files", "/fine_tuning", "/assistants", "/threads", "/utils/", "/vector_stores", "/vector_store/", "/search", "/containers", "/skills", "/claude-code/", "/interactions", "/agents", "/responses/{", "/responses/input_tokens", - "/realtime/client_secrets", "/realtime/calls", "/realtime/transcription_sessions", + "/realtime/client_secrets", "/realtime/calls", "/realtime/transcription_sessions", "/live", )) or normalized in ("/models", "/cursor/models", "/cursor/v1/models") concrete: Final = route.split("?")[0].replace("{model}", "model").replace("{model_name:path}", "model") assert managed_agent_route_allowed(concrete, None) is not unsupported, route diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 83ac56c4c85..8a0b6cc0d77 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -828,7 +828,11 @@ def _azure_relay_router(): model_list=[ { "model_name": "gpt", - "litellm_params": {"model": "azure_ai/gpt-5.4-mini", "api_base": "https://a.services.ai.azure.com", "api_key": "k"}, + "litellm_params": { + "model": "azure_ai/gpt-5.4-mini", + "api_base": "https://a.services.ai.azure.com", + "api_key": "k", + }, }, { "model_name": "other-group", @@ -1279,7 +1283,8 @@ def test_get_model_from_request_handles_managed_id_decoder_failures(): "/openai/v1/realtime/calls", ], ) -def test_get_model_from_request_extracts_realtime_session_model(route): +@pytest.mark.parametrize("encoded", [False, True]) +def test_get_model_from_request_extracts_realtime_session_model(route, encoded): """The effective realtime model lives in ``session.model`` (not the top-level ``model``). It must be surfaced so can_key_call_model() can validate the model a restricted key is actually requesting. @@ -1289,13 +1294,49 @@ def test_get_model_from_request_extracts_realtime_session_model(route): """ assert ( get_model_from_request( - request_data={"session": {"type": "realtime", "model": "gpt-realtime"}}, + request_data={"session": '{"model":"gpt-realtime"}' if encoded else {"model": "gpt-realtime"}}, route=route, ) == "gpt-realtime" ) +@pytest.mark.parametrize("session", ['{"model":"actual-voice"}', {"model": "actual-voice"}]) +@pytest.mark.parametrize( + "route", + [ + "/v1/realtime/calls", + "/v1/live", + "/live", + "/openai/v1/live", + "/v1/live/sessions", + "/live/sessions", + "/openai/v1/live/sessions", + "/v1/live/sessions/incoming/accept", + ], +) +def test_realtime_calls_auth_uses_executed_session_model_despite_decoys(session, route): + assert ( + get_model_from_request( + request_data={"model": "body-decoy", "session": session}, + route=route, + request_query_params={"model": "query-decoy"}, + request_headers={"x-litellm-model": "header-decoy"}, + ) + == "actual-voice" + ) + + +@pytest.mark.parametrize("model", ["voice,alias", " voice "]) +def test_realtime_calls_auth_preserves_exact_session_model(model): + assert get_model_from_request(request_data={"session": {"model": model}}, route="/v1/realtime/calls") == model + + +@pytest.mark.parametrize("session", ["invalid", "null", "[]", "12", '"text"', "{}"]) +def test_realtime_model_extraction_ignores_invalid_serialized_session(session): + assert get_model_from_request(request_data={"session": session}, route="/v1/realtime/calls") is None + + def test_get_model_from_request_realtime_includes_top_level_and_session_model(): """When both top-level and session model are present, both are returned so neither path can smuggle a disallowed model past the model-access check.""" diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index d8ee58a52ea..ad48ee5c58b 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -17,6 +17,24 @@ from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin from litellm.proxy.auth.route_checks import RouteChecks +@pytest.mark.parametrize("route", ["/live", "/v1/live", "/v1/live/rtc_litellm_test"]) +def test_codex_live_routes_allow_inference_keys(route: str): + from litellm.proxy.auth.auth_checks import _allowed_routes_check + + assert RouteChecks.is_llm_api_route(route) + assert _allowed_routes_check(user_route=route, allowed_routes=["openai_routes"]) + token = UserAPIKeyAuth(allowed_routes=["llm_api_routes"]) + RouteChecks.is_virtual_key_allowed_to_call_route(route=route, valid_token=token) + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=None, + _user_role=LitellmUserRoles.INTERNAL_USER.value, + route=route, + request=Request({"type": "http", "path": route, "query_string": b"", "headers": []}), + valid_token=token, + request_data={}, + ) + + def test_non_admin_config_update_route_rejected(): """Test that non-admin users are rejected when trying to call /config/update""" diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index ef6832ef77b..81d9281c2e9 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -365,7 +365,7 @@ async def test_custom_auth_does_not_enforce_key_model_access_by_default(): async def test_post_custom_auth_expired_key_returns_unauthorized(): expired_token = UserAPIKeyAuth( token="test_token", - expires=datetime.now() - timedelta(minutes=1), + expires=datetime.now(timezone.utc) - timedelta(minutes=1), ) with pytest.raises(ProxyException) as exc_info: @@ -8456,6 +8456,52 @@ def test_user_api_key_auth_opens_a_datadog_span_for_accepted_and_rejected_keys(t assert [span for span in report["spans"] if span == auth_span] == [auth_span, auth_span] +@pytest.mark.asyncio +@pytest.mark.parametrize("attachment", ["path", "query"]) +@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom", "custom-mixed"]) +@pytest.mark.parametrize("query_model", [b"", b"model=unbudgeted"]) +async def test_sideband_auth_uses_encrypted_model_for_budget_checks(monkeypatch, attachment, credential, query_model): + import hashlib + import importlib + import time + from unittest.mock import AsyncMock + from fastapi import WebSocket + from litellm.llms.chatgpt.codex import CodexRealtimeCall + from litellm.proxy.realtime_endpoints.call_sessions import encode_call + + auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-sideband-budget-salt") + token = encode_call(CodexRealtimeCall( + call_id="rtc_test", model="gpt-live-1-codex", alias="budgeted-voice", + owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300, + )) + from litellm.proxy import proxy_server + monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential.startswith("custom") else {}) + seen = [] + + async def authenticate(request, api_key): + seen.append((await request.json(), api_key)) + return "authenticated-with-model" + + monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate) + websocket = WebSocket({ + "type": "websocket", "scheme": "ws", "server": ("localhost", 4000), + "path": "/v1/live/" + token if attachment == "path" else "/v1/realtime", + "path_params": {"call_id": token} if attachment == "path" else {}, + "query_string": query_model + (b"&call_id=" + token.encode() if attachment == "query" else b""), + "headers": { + "authorization": [(b"authorization", b"Bearer owner")], + "api-key": [(b"api-key", b"owner")], + "x-litellm-api-key": [(b"x-litellm-api-key", b"owner")], + "custom": [(b"x-proxy-key", b"Bearer owner")], + "custom-mixed": [(b"x-proxy-key", b"Bearer owner"), (b"authorization", b"Bearer other-owner")], + "subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")], + }[credential], + }, AsyncMock(), AsyncMock()) + assert await auth_module.user_api_key_auth_websocket(websocket) == "authenticated-with-model" + assert seen == [({"model": "budgeted-voice"}, "Bearer owner")] + + @pytest.mark.asyncio @pytest.mark.parametrize("is_proxy_admin", [False, True], ids=["standard-return", "proxy-admin-return"]) async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_admin): @@ -8572,6 +8618,87 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a assert token.jwt_claims == {"sub": "jwt-user"} +@pytest.mark.asyncio +@pytest.mark.parametrize("attachment", ["path", "query"]) +async def test_sideband_rejects_budget_fallback_before_rerouting(monkeypatch, attachment): + import hashlib + import importlib + import time + from types import SimpleNamespace + from unittest.mock import AsyncMock + from fastapi import HTTPException, WebSocket + from litellm.llms.chatgpt.codex import CodexRealtimeCall + from litellm.proxy.realtime_endpoints.call_sessions import encode_call + + auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-sideband-budget-salt") + token = encode_call(CodexRealtimeCall( + call_id="rtc_test", model="gpt-live-1-codex", alias="budgeted-voice", + owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300, + )) + limiter = SimpleNamespace( + is_key_within_model_budget=AsyncMock(side_effect=litellm.BudgetExceededError(current_cost=2, max_budget=1)), + get_fallback_model_within_budget=AsyncMock(return_value="cheap-voice"), + ) + auth = UserAPIKeyAuth(models=["budgeted-voice", "cheap-voice"]) + + async def authenticate(request, api_key): + data = await request.json() + await auth_module._check_key_model_budget_with_fallback(auth, limiter, data["model"], data, request) + return auth + + monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate) + monkeypatch.setattr(auth_module, "can_key_call_model", AsyncMock()) + send = AsyncMock() + websocket = WebSocket({ + "type": "websocket", "scheme": "ws", "server": ("localhost", 4000), + "path": "/v1/live/" + token if attachment == "path" else "/v1/realtime", + "path_params": {"call_id": token} if attachment == "path" else {}, + "query_string": b"call_id=" + token.encode() if attachment == "query" else b"", + "headers": [(b"authorization", b"Bearer owner")], + }, AsyncMock(), send) + with pytest.raises(HTTPException) as error: + await auth_module.user_api_key_auth_websocket(websocket) + assert error.value.status_code == 403 + limiter.get_fallback_model_within_budget.assert_not_awaited() + send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("custom_value", [None, b"Bearer different-owner"]) +async def test_sideband_custom_header_cannot_fall_back_to_other_credentials(monkeypatch, custom_value): + import hashlib + import importlib + import time + from unittest.mock import AsyncMock + from fastapi import HTTPException, WebSocket + from litellm.proxy import proxy_server + from litellm.llms.chatgpt.codex import CodexRealtimeCall + from litellm.proxy.realtime_endpoints.call_sessions import encode_call + + auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-custom-header-salt") + monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"}) + token = encode_call(CodexRealtimeCall( + call_id="rtc_test", model="gpt-live-1-codex", alias="voice", + owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300, + )) + authenticate = AsyncMock() + monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate) + send = AsyncMock() + websocket = WebSocket({ + "type": "websocket", "scheme": "ws", "server": ("localhost", 4000), + "path": "/v1/live/" + token, "path_params": {"call_id": token}, "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")] + + ([(b"x-proxy-key", custom_value)] if custom_value is not None else []), + }, AsyncMock(), send) + with pytest.raises(HTTPException) as error: + await auth_module.user_api_key_auth_websocket(websocket) + assert error.value.status_code == 403 + authenticate.assert_not_awaited() + send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""}) + + @pytest.mark.asyncio @pytest.mark.parametrize( "route", ["/v1/messages", "/messages", "/v1/chat/completions", "/chat/completions", "/v1/responses", "/responses"] @@ -8641,6 +8768,67 @@ async def test_claude_view_never_reinterprets_explicit_names(monkeypatch, layer) assert data["model"] == ("foo" if layer == "unclaimed" else encoded) +def _malformed_authorization_websocket(send): + from unittest.mock import AsyncMock + + from fastapi import WebSocket + + return WebSocket( + { + "type": "websocket", "scheme": "ws", "server": ("localhost", 4000), + "path": "/v1/realtime", "query_string": b"", + "headers": [(b"authorization", b"Token malformed")], + }, + AsyncMock(), + send, + ) + + +@pytest.mark.parametrize("authorization_value", ["Token malformed", "bearer lowercase"]) +def test_get_websocket_api_key_rejects_malformed_authorization(monkeypatch, authorization_value): + import importlib + from unittest.mock import AsyncMock + + from fastapi import HTTPException, WebSocket + + from litellm.proxy import proxy_server + + auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setattr(proxy_server, "general_settings", {}) + websocket = WebSocket( + { + "type": "websocket", "scheme": "ws", "server": ("localhost", 4000), + "path": "/v1/realtime", "query_string": b"", + "headers": [(b"authorization", authorization_value.encode())], + }, + AsyncMock(), + AsyncMock(), + ) + with pytest.raises(HTTPException) as error: + auth_module.get_websocket_api_key(websocket) + assert error.value.status_code == 403 + assert error.value.detail == "Invalid Authorization header format" + + +@pytest.mark.asyncio +async def test_websocket_auth_closes_policy_violation_on_malformed_authorization(monkeypatch): + import importlib + from unittest.mock import AsyncMock + + from fastapi import HTTPException + + from litellm.proxy import proxy_server + + auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setattr(proxy_server, "general_settings", {}) + send = AsyncMock() + websocket = _malformed_authorization_websocket(send) + with pytest.raises(HTTPException) as error: + await auth_module.user_api_key_auth_websocket(websocket) + assert error.value.status_code == 403 + assert error.value.detail == "Invalid Authorization header format" + send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""}) + ISSUER_ONE = "https://issuer-one.example.com" ISSUER_TWO = "https://issuer-two.example.com" @@ -9205,6 +9393,33 @@ async def test_router_settings_model_group_alias_authorizes_target_for_team(monk assert get_client_requested_model(request) == "AgentX-LLM" +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ["configured-voice", None]) +async def test_websocket_auth_explicit_model_overrides_query(monkeypatch, model): + import importlib + from fastapi import WebSocket + + from litellm.proxy import proxy_server + + auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setattr(proxy_server, "general_settings", {}) + seen = [] + + async def authenticate(request, api_key): + seen.append((await request.json(), api_key)) + return "authenticated" + + monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate) + websocket = WebSocket({ + "type": "websocket", "scheme": "ws", "server": ("localhost", 4000), + "path": "/v1/realtime", "path_params": {}, + "query_string": b"model=untrusted-query", + "headers": [(b"x-litellm-api-key", b"owner")], + }, AsyncMock(), AsyncMock()) + assert await auth_module.user_api_key_auth_websocket_for_model(websocket, model) == "authenticated" + assert seen == [({"model": model or ""}, "Bearer owner")] + + @pytest.mark.asyncio async def test_reserve_budget_after_common_checks_hands_the_reservation_to_the_request_state(): from fastapi import Request diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 3bf51f02d34..b5d83461e2e 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -1039,6 +1039,20 @@ def _raw_batches_request(body: Dict[str, Any]) -> MagicMock: request.headers = {"Content-Type": "application/json"} request.client = MagicMock() request.client.host = "127.0.0.1" + request.scope = { + "type": "http", + "asgi": {"version": "3.0", "spec_version": "2.3"}, + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": "/v1/batches", + "raw_path": b"/v1/batches", + "query_string": b"", + "root_path": "", + "headers": [(b"content-type", b"application/json"), (b"host", b"localhost")], + "client": ("127.0.0.1", 54321), + "server": ("localhost", 8000), + } request.body = AsyncMock(return_value=json.dumps(body).encode()) return request diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py index a83ddc69863..4a25d4dbe79 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -2,12 +2,16 @@ Unit Tests for the max parallel request limiter v1 for the proxy """ +import asyncio from datetime import datetime +from unittest.mock import AsyncMock, MagicMock, patch import pytest from litellm.caching.caching import DualCache +from litellm.caching.redis_cache import RedisCache from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.parallel_request_limiter import ( _PROXY_MaxParallelRequestsHandler, ) @@ -15,6 +19,98 @@ from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage +@pytest.mark.asyncio +async def test_realtime_release_preserves_newer_local_admission_while_redis_finishes(): + started, finish = asyncio.Event(), asyncio.Event() + + async def release(**kwargs): + started.set() + await finish.wait() + + remote = MagicMock(spec=RedisCache) + remote.async_register_script.return_value = AsyncMock(side_effect=release) + cache = DualCache(redis_cache=remote) + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache)) + await cache.async_set_cache("key", {"current_requests": 1, "current_rpm": 1, "current_tpm": 7}, local_only=True) + task = asyncio.create_task(handler._release_realtime_counter("key")) + await started.wait() + next_admission = {"current_requests": 1, "current_rpm": 2, "current_tpm": 7} + await cache.async_set_cache("key", next_admission, local_only=True) + finish.set() + await task + assert await cache.async_get_cache("key", local_only=True) == next_admission + + +@pytest.mark.asyncio +@pytest.mark.parametrize("reject_team", [False, True]) +async def test_realtime_attachment_releases_only_acquired_legacy_slots(reject_team): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache)) + auth = UserAPIKeyAuth( + api_key="attachment-key", + user_id="attachment-user", + team_id="attachment-team", + team_rpm_limit=0 if reject_team else 100, + max_parallel_requests=1, + end_user_id="attachment-end-user", + metadata={"model_rpm_limit": {"test-model": 100}}, + ) + data = {"model": "test-model", "metadata": {"global_max_parallel_requests": 10}} + minute = datetime.now().strftime("%Y-%m-%d-%H-%M") + team_key = f"attachment-team::{minute}::request_count" + await cache.async_set_cache(team_key, {"current_requests": 3, "current_tpm": 7, "current_rpm": 4}) + handler.begin_realtime_attachment(data) + if reject_team: + with pytest.raises(ProxyRateLimitError, match="Rate Limit Handler"): + await handler.async_pre_call_hook(auth, cache, data, "_arealtime") + else: + await handler.async_pre_call_hook(auth, cache, data, "_arealtime") + await handler.async_release_realtime_attachment(data, auth) + await handler.async_release_realtime_attachment(data, auth) + assert await cache.async_get_cache("global_max_parallel_requests") == 0 + assert await cache.async_get_cache(f"attachment-key::{minute}::request_count") == { + "current_requests": 0, + "current_tpm": 0, + "current_rpm": 1, + } + assert await cache.async_get_cache(f"attachment-user::{minute}::request_count") == { + "current_requests": 0, + "current_tpm": 0, + "current_rpm": 1, + } + assert await cache.async_get_cache(team_key) == { + "current_requests": 3, + "current_tpm": 7, + "current_rpm": 4 if reject_team else 5, + } + assert await cache.async_get_cache(f"attachment-key::test-model::{minute}::request_count") == { + "current_requests": 0, + "current_tpm": 0, + "current_rpm": 1, + } + end_user = await cache.async_get_cache(f"attachment-end-user::{minute}::request_count") + assert end_user == (None if reject_team else {"current_requests": 0, "current_tpm": 0, "current_rpm": 1}) + if not reject_team: + handler.begin_realtime_attachment(data) + await handler.async_pre_call_hook(auth, cache, data, "_arealtime") + await handler.async_release_realtime_attachment(data, auth) + + +@pytest.mark.asyncio +async def test_realtime_attachment_rejected_before_acquisition_preserves_other_slot(): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key="busy-key", max_parallel_requests=1) + minute = datetime.now().strftime("%Y-%m-%d-%H-%M") + key = f"busy-key::{minute}::request_count" + current = {"current_requests": 1, "current_tpm": 13, "current_rpm": 2} + await cache.async_set_cache(key, current) + data = {"model": "test-model"} + handler.begin_realtime_attachment(data) + with pytest.raises(ProxyRateLimitError, match="Rate Limit Handler"): + await handler.async_pre_call_hook(auth, cache, data, "_arealtime") + await handler.async_release_realtime_attachment(data, auth) + assert await cache.async_get_cache(key) == current @pytest.mark.asyncio async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token(): handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) @@ -62,9 +158,7 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ team_id = "litellm-team" end_user_id = "customer-1" - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) current_date = datetime.now().strftime("%Y-%m-%d") current_hour = datetime.now().strftime("%H") @@ -103,7 +197,52 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ key=f"{scope_id}::{precise_minute}::request_count", litellm_parent_otel_span=None, ) - assert current["current_tpm"] == 50, ( - f"expected 50 tokens counted for {scope_id}, " - f"got {current['current_tpm']}" - ) + assert current["current_tpm"] == 50, f"expected 50 tokens counted for {scope_id}, got {current['current_tpm']}" + + +@pytest.mark.asyncio +async def test_realtime_attachment_release_without_receipt_never_touches_counters(): + dual_cache = MagicMock() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache)) + auth = UserAPIKeyAuth(api_key="no-receipt") + await handler.async_release_realtime_attachment({}, auth) + await handler.async_release_realtime_attachment( + {"_legacy_realtime_attachment_reservations": {"cache_keys": [], "global_acquired": True}}, auth + ) + # A release without a matching begin (or with a foreign receipt shape) must not decrement anything. + assert dual_cache.mock_calls == [] + + +@pytest.mark.asyncio +async def test_failure_event_skips_realtime_observer_without_decrementing_slots(): + from datetime import datetime + + from litellm.proxy._types import InternalRequestOrigin + + def failure_kwargs() -> dict: + return { + "litellm_params": {"metadata": {"user_api_key": "observer-hash", "global_max_parallel_requests": 5}}, + "exception": RuntimeError("backend disconnected"), + } + + dual_cache = MagicMock() + dual_cache.async_get_cache = AsyncMock(return_value=None) + dual_cache.async_increment_cache = AsyncMock() + dual_cache.async_batch_set_cache = AsyncMock() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache)) + start = datetime.now() + end = datetime.now() + + kwargs = failure_kwargs() + kwargs["internal_request_origin"] = InternalRequestOrigin.REALTIME_OBSERVER + await handler.async_log_failure_event(kwargs, None, start, end) + # The observer-internal failure mirror must leave the client-facing slot untouched. + assert dual_cache.mock_calls == [] + + dual_cache.mock_calls.clear() + await handler.async_log_failure_event(failure_kwargs(), None, start, end) + assert dual_cache.async_increment_cache.await_count >= 1 + assert any( + call.kwargs.get("key") == "global_max_parallel_requests" and call.kwargs.get("value") == -1 + for call in dual_cache.async_increment_cache.await_args_list + ) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index b546b9eb965..2fa66221693 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -38,6 +38,8 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token +from litellm.proxy.hooks.parallel_request_limiter_v3 import isolated_request_stash +from litellm.proxy.hooks.realtime_call_lease import realtime_call_attachment from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import ( @@ -59,6 +61,94 @@ class TimeController: self._current += timedelta(seconds=seconds) +@pytest.mark.asyncio +async def test_realtime_lease_retains_quota_across_signaling_and_three_attachments(): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key="logical-owner", max_parallel_requests=1, rpm_limit=4, tpm_limit=100000) + data = {"model": "gpt-3.5-turbo", "litellm_call_id": "signaling"} + await handler.async_pre_call_hook(auth, cache, data, "arealtime_calls") + stash = get_request_stash() + assert stash.reserved_tokens > 0 + assert handler.transfer_realtime_call_slot({"litellm_call_id": "other-call"}) is None + assert stash.parallel_slot is not None + lease = handler.transfer_realtime_call_slot(data) + assert lease is not None + assert stash.parallel_slot is None + assert stash.reserved_tokens > 0 + assert not stash.reservation_released + assert handler.transfer_realtime_call_slot(data) is None + await handler.async_log_success_event( + kwargs={ + "litellm_call_id": "signaling", + "standard_logging_object": {"metadata": {"user_api_key_hash": auth.api_key}}, + }, + response_obj=ModelResponse(usage=Usage()), + start_time=datetime.now(), + end_time=datetime.now(), + ) + assert await cache.async_get_cache("{api_key:logical-owner}:tokens") == 0 + socket = object() + for attachment in range(3): + with isolated_request_stash(), realtime_call_attachment(socket): + attachment_data = { + "model": "gpt-3.5-turbo", + "litellm_call_id": f"attachment-{attachment}", + "websocket": socket, + } + await handler.async_pre_call_hook(auth, cache, attachment_data, "_arealtime") + assert get_request_stash().parallel_slot is None + await handler.async_release_realtime_attachment(attachment_data, auth) + with isolated_request_stash(), realtime_call_attachment(socket), pytest.raises(HTTPException) as error: + await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo", "websocket": socket}, "_arealtime") + assert error.value.status_code == 429 + assert "requests" in str(error.value.detail) + quota_only = auth.model_copy(update={"rpm_limit": None}) + with isolated_request_stash(), realtime_call_attachment(object()), pytest.raises(HTTPException) as error: + await handler.async_pre_call_hook( + quota_only, cache, {"model": "gpt-3.5-turbo", "websocket": socket}, "_arealtime" + ) + assert "max_parallel_requests" in str(error.value.detail) + with isolated_request_stash(), realtime_call_attachment(socket), pytest.raises(HTTPException) as error: + await handler.async_pre_call_hook( + quota_only, cache, {"model": "gpt-3.5-turbo", "websocket": socket}, "acompletion" + ) + assert "max_parallel_requests" in str(error.value.detail) + with isolated_request_stash(), pytest.raises(HTTPException) as error: + await handler.async_pre_call_hook(quota_only, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + assert "max_parallel_requests" in str(error.value.detail) + assert get_request_stash() is stash + await lease.close() + with isolated_request_stash(): + await handler.async_pre_call_hook(quota_only, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + + +@pytest.mark.asyncio +async def test_realtime_lease_renewal_preserves_quota_past_ttl_and_does_not_resurrect_expiry(): + cache = DualCache() + clock = TimeController() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache), time_provider=clock.now) + auth = UserAPIKeyAuth(api_key="long-call", max_parallel_requests=1) + data = {"model": "gpt-3.5-turbo", "litellm_call_id": "long-call"} + await handler.async_pre_call_hook(auth, cache, data, "arealtime_calls") + lease = handler.transfer_realtime_call_slot(data) + assert lease is not None + clock.advance(PARALLEL_REQUEST_SLOT_TTL_SECONDS - 1) + assert await lease.renew() + clock.advance(2) + with isolated_request_stash(), pytest.raises(HTTPException): + await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + + clock.advance(PARALLEL_REQUEST_SLOT_TTL_SECONDS) + assert not await lease.renew() + await asyncio.wait_for(lease.wait_failed(), 1) + with isolated_request_stash(): + await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + await lease.close() + with isolated_request_stash(), pytest.raises(HTTPException): + await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + + @pytest.fixture def time_controller(monkeypatch): controller = TimeController() @@ -81,14 +171,10 @@ def _isolated_request_stash(): (0.5, 50, 500), ], ) -def test_api_key_descriptor_applies_budget_throttle( - throttle_pct, expected_rpm, expected_tpm -): +def test_api_key_descriptor_applies_budget_throttle(throttle_pct, expected_rpm, expected_tpm): """The api_key rate-limit descriptor scales the key's configured TPM/RPM by the request-scoped budget_throttle_pct, leaving the configured limits intact.""" - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key=hash_token("sk-throttle"), rpm_limit=100, @@ -146,18 +232,12 @@ async def test_sliding_window_rate_limit_v3(monkeypatch, time_controller): window_starts[window_key] = now new_counter = 1 request_counts[counter_key] = new_counter - await local_cache.async_set_cache( - key=window_key, value=now, ttl=window_size - ) - await local_cache.async_set_cache( - key=counter_key, value=new_counter, ttl=window_size - ) + await local_cache.async_set_cache(key=window_key, value=now, ttl=window_size) + await local_cache.async_set_cache(key=counter_key, value=new_counter, ttl=window_size) else: new_counter = prev_counter + 1 request_counts[counter_key] = new_counter - await local_cache.async_set_cache( - key=counter_key, value=new_counter, ttl=window_size - ) + await local_cache.async_set_cache(key=counter_key, value=new_counter, ttl=window_size) results.append(now) results.append(new_counter) return results @@ -239,18 +319,12 @@ async def test_rate_limiter_script_return_values_v3(monkeypatch, time_controller window_starts[window_key] = now new_counter = 1 request_counts[counter_key] = new_counter - await local_cache.async_set_cache( - key=window_key, value=now, ttl=window_size - ) - await local_cache.async_set_cache( - key=counter_key, value=new_counter, ttl=window_size - ) + await local_cache.async_set_cache(key=window_key, value=now, ttl=window_size) + await local_cache.async_set_cache(key=counter_key, value=new_counter, ttl=window_size) else: new_counter = prev_counter + 1 request_counts[counter_key] = new_counter - await local_cache.async_set_cache( - key=counter_key, value=new_counter, ttl=window_size - ) + await local_cache.async_set_cache(key=counter_key, value=new_counter, ttl=window_size) results.append(now) results.append(new_counter) return results @@ -282,9 +356,7 @@ async def test_rate_limiter_script_return_values_v3(monkeypatch, time_controller new_window_value = await local_cache.async_get_cache(key=window_key) new_counter_value = await local_cache.async_get_cache(key=counter_key) - assert ( - new_window_value == window_value - ), "Window value should not change within window" + assert new_window_value == window_value, "Window value should not change within window" assert new_counter_value == 2, "Counter should be 2 after second request" # Wait for window to expire @@ -315,9 +387,7 @@ async def test_rate_limiter_script_return_values_v3(monkeypatch, time_controller ) @pytest.mark.flaky(reruns=3) @pytest.mark.asyncio -async def test_normal_router_call_tpm_v3( - monkeypatch, rate_limit_object, time_controller -): +async def test_normal_router_call_tpm_v3(monkeypatch, rate_limit_object, time_controller): """ Test normal router call with parallel request limiter v3 for TPM rate limiting """ @@ -393,18 +463,12 @@ async def test_normal_router_call_tpm_v3( window_starts[window_key] = now new_counter = 1 request_counts[counter_key] = new_counter - await local_cache.async_set_cache( - key=window_key, value=now, ttl=window_size - ) - await local_cache.async_set_cache( - key=counter_key, value=new_counter, ttl=window_size - ) + await local_cache.async_set_cache(key=window_key, value=now, ttl=window_size) + await local_cache.async_set_cache(key=counter_key, value=new_counter, ttl=window_size) else: new_counter = prev_counter + 1 request_counts[counter_key] = new_counter - await local_cache.async_set_cache( - key=counter_key, value=new_counter, ttl=window_size - ) + await local_cache.async_set_cache(key=counter_key, value=new_counter, ttl=window_size) results.append(now) results.append(new_counter) return results @@ -427,9 +491,7 @@ async def test_normal_router_call_tpm_v3( return None value = get_value_for_key(rate_limit_object, user_api_key_dict, "azure-model") - counter_key = parallel_request_handler.create_rate_limit_keys( - rate_limit_object, value, "tokens" - ) + counter_key = parallel_request_handler.create_rate_limit_keys(rate_limit_object, value, "tokens") # First request should succeed. Include messages + a tight max_tokens so # the atomic reserve_tpm_tokens path populates the :tokens counter with a @@ -442,12 +504,8 @@ async def test_normal_router_call_tpm_v3( "messages": [{"role": "user", "content": "hi"}], "max_tokens": 5, } - expected_reservation = parallel_request_handler._estimate_tokens_for_request( - data=pre_call_data - ) - assert ( - expected_reservation < 10 - ), "Test premise: reservation must fit under tpm_limit=10" + expected_reservation = parallel_request_handler._estimate_tokens_for_request(data=pre_call_data) + assert expected_reservation < 10, "Test premise: reservation must fit under tpm_limit=10" await parallel_request_handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -475,15 +533,11 @@ async def test_normal_router_call_tpm_v3( counter_value = await local_cache.async_get_cache(key=counter_key) print(f"local_cache: {local_cache.in_memory_cache.cache_dict}") - assert ( - counter_value is not None - ), f"Counter value should be stored in cache for {counter_key}" + assert counter_value is not None, f"Counter value should be stored in cache for {counter_key}" # Manually increment the token counter to simulate token usage from previous call # This simulates what would happen after a successful call - await local_cache.async_increment_cache( - key=counter_key, value=15, ttl=2 - ) # Use up most of our 10 token limit + await local_cache.async_increment_cache(key=counter_key, value=15, ttl=2) # Use up most of our 10 token limit # Make another request to test rate limiting - this should fail as we've consumed tokens with pytest.raises(HTTPException) as exc_info: @@ -531,17 +585,13 @@ async def test_token_rate_limit_type_respected_v3(monkeypatch, token_rate_limit_ _api_key = hash_token(_api_key) user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, tpm_limit=100) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock the get_rate_limit_type method directly since it imports general_settings internally def mock_get_rate_limit_type(): return token_rate_limit_type - monkeypatch.setattr( - parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type - ) + monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type) # Create a mock response with different token counts mock_usage = Usage(prompt_tokens=20, completion_tokens=30, total_tokens=50) @@ -590,9 +640,9 @@ async def test_token_rate_limit_type_respected_v3(monkeypatch, token_rate_limit_ ) # Verify that the correct token count was used based on the rate limit type - assert ( - len(captured_operations) == 1 - ), "Should have 1 operation: the TPM increment (parallel slots are released via the gauge, not the pipeline)" + assert len(captured_operations) == 1, ( + "Should have 1 operation: the TPM increment (parallel slots are released via the gauge, not the pipeline)" + ) tpm_operation = None for op in captured_operations: @@ -609,9 +659,9 @@ async def test_token_rate_limit_type_respected_v3(monkeypatch, token_rate_limit_ "total": mock_usage.total_tokens, # 50 } - assert ( - tpm_operation["increment_value"] == expected_tokens[token_rate_limit_type] - ), f"Expected {expected_tokens[token_rate_limit_type]} tokens for type '{token_rate_limit_type}', got {tpm_operation['increment_value']}" + assert tpm_operation["increment_value"] == expected_tokens[token_rate_limit_type], ( + f"Expected {expected_tokens[token_rate_limit_type]} tokens for type '{token_rate_limit_type}', got {tpm_operation['increment_value']}" + ) @pytest.mark.parametrize( @@ -628,9 +678,7 @@ async def test_token_rate_limit_type_respected_v3(monkeypatch, token_rate_limit_ ], ) @pytest.mark.asyncio -async def test_async_log_success_event_counts_non_chat_response_tokens( - monkeypatch, response_obj -): +async def test_async_log_success_event_counts_non_chat_response_tokens(monkeypatch, response_obj): """ Embedding and text completion responses must increment the TPM counter, not just chat completion ModelResponse objects. @@ -638,12 +686,8 @@ async def test_async_log_success_event_counts_non_chat_response_tokens( monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") _api_key = hash_token("sk-12345") - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) - monkeypatch.setattr( - parallel_request_handler, "get_rate_limit_type", lambda: "total" - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", lambda: "total") mock_kwargs = { "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, @@ -669,9 +713,7 @@ async def test_async_log_success_event_counts_non_chat_response_tokens( end_time=datetime.now(), ) - tpm_operation = next( - (op for op in captured_operations if op["key"].endswith(":tokens")), None - ) + tpm_operation = next((op for op in captured_operations if op["key"].endswith(":tokens")), None) assert tpm_operation is not None, "Should have a TPM increment operation" assert tpm_operation["increment_value"] == 50 @@ -687,9 +729,7 @@ async def test_async_log_failure_event_v3(): _api_key = "sk-12345" _api_key = hash_token(_api_key) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" await _seed_max_parallel_requests_slots(local_cache, counter_key, ["slot-a", "slot-b"]) @@ -739,29 +779,18 @@ async def test_failure_event_without_acquired_slot_does_not_release_v3(): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" - await _seed_max_parallel_requests_slots( - local_cache, counter_key, ["slot-a", "slot-b", "slot-c"] - ) + await _seed_max_parallel_requests_slots(local_cache, counter_key, ["slot-a", "slot-b", "slot-c"]) await handler.async_log_failure_event( - kwargs={ - "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}} - }, + kwargs={"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}}, response_obj=None, start_time=None, end_time=None, ) - assert ( - handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) - == 3 - ) + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 3 @pytest.mark.asyncio @@ -812,9 +841,7 @@ async def test_rejected_request_does_not_consume_parallel_slot_v3(): rejected requests that should have been admitted after a release. """ local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) _api_key = hash_token("sk-12345") user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=1) @@ -867,9 +894,7 @@ async def test_parallel_gauge_uses_atomic_redis_script_v3(): and an over-limit script result maps to a 429 without occupying a slot. """ local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) _api_key = hash_token("sk-12345") user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=5) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" @@ -894,9 +919,7 @@ async def test_parallel_gauge_uses_atomic_redis_script_v3(): stashed_slot_id = stashed_acquisition["slot_id"] assert isinstance(stashed_slot_id, str) and stashed_slot_id assert stashed_acquisition["counter_keys"] == [counter_key] - assert captured_calls == [ - ([counter_key], [5, PARALLEL_REQUEST_SLOT_TTL_SECONDS, stashed_slot_id]) - ] + assert captured_calls == [([counter_key], [5, PARALLEL_REQUEST_SLOT_TTL_SECONDS, stashed_slot_id])] assert ( await handler.internal_usage_cache.async_get_cache( key=counter_key, litellm_parent_otel_span=None, local_only=True @@ -943,9 +966,7 @@ async def test_should_rate_limit_only_called_when_limits_exist_v3(): _api_key = "sk-12345" _api_key = hash_token(_api_key) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock should_rate_limit to track if it's called should_rate_limit_called = False @@ -974,9 +995,7 @@ async def test_should_rate_limit_only_called_when_limits_exist_v3(): call_type="", ) - assert ( - not should_rate_limit_called - ), "should_rate_limit should not be called when no rate limits are configured" + assert not should_rate_limit_called, "should_rate_limit should not be called when no rate limits are configured" # Test 2: API key rate limits configured - should_rate_limit SHOULD be called should_rate_limit_called = False @@ -992,9 +1011,7 @@ async def test_should_rate_limit_only_called_when_limits_exist_v3(): call_type="", ) - assert ( - should_rate_limit_called - ), "should_rate_limit should be called when API key rate limits are configured" + assert should_rate_limit_called, "should_rate_limit should be called when API key rate limits are configured" # Test 3: User rate limits configured - should_rate_limit SHOULD be called should_rate_limit_called = False @@ -1011,9 +1028,7 @@ async def test_should_rate_limit_only_called_when_limits_exist_v3(): call_type="", ) - assert ( - should_rate_limit_called - ), "should_rate_limit should be called when user rate limits are configured" + assert should_rate_limit_called, "should_rate_limit should be called when user rate limits are configured" # Test 4: Team rate limits configured - should_rate_limit SHOULD be called should_rate_limit_called = False @@ -1030,9 +1045,7 @@ async def test_should_rate_limit_only_called_when_limits_exist_v3(): call_type="", ) - assert ( - should_rate_limit_called - ), "should_rate_limit should be called when team rate limits are configured" + assert should_rate_limit_called, "should_rate_limit should be called when team rate limits are configured" # Test 5: End user rate limits configured - should_rate_limit SHOULD be called should_rate_limit_called = False @@ -1049,9 +1062,7 @@ async def test_should_rate_limit_only_called_when_limits_exist_v3(): call_type="", ) - assert ( - should_rate_limit_called - ), "should_rate_limit should be called when end user rate limits are configured" + assert should_rate_limit_called, "should_rate_limit should be called when end user rate limits are configured" # Test 6: Max parallel requests configured - should_rate_limit SHOULD be called should_rate_limit_called = False @@ -1067,9 +1078,7 @@ async def test_should_rate_limit_only_called_when_limits_exist_v3(): call_type="", ) - assert ( - should_rate_limit_called - ), "should_rate_limit should be called when max parallel requests are configured" + assert should_rate_limit_called, "should_rate_limit should be called when max parallel requests are configured" @pytest.mark.asyncio @@ -1085,9 +1094,7 @@ async def test_model_specific_rate_limits_only_called_when_configured_v3(): _api_key = "sk-12345" _api_key = hash_token(_api_key) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock should_rate_limit to track if it's called should_rate_limit_called = False @@ -1103,9 +1110,7 @@ async def test_model_specific_rate_limits_only_called_when_configured_v3(): should_rate_limit_called = False user_api_key_dict_with_model_limits = UserAPIKeyAuth( api_key=_api_key, - metadata={ - "model_tpm_limit": {"gpt-4": 1000} - }, # Rate limit for gpt-4, not gpt-3.5-turbo + metadata={"model_tpm_limit": {"gpt-4": 1000}}, # Rate limit for gpt-4, not gpt-3.5-turbo ) await parallel_request_handler.async_pre_call_hook( @@ -1115,17 +1120,15 @@ async def test_model_specific_rate_limits_only_called_when_configured_v3(): call_type="", ) - assert ( - not should_rate_limit_called - ), "should_rate_limit should not be called when model-specific limits don't match requested model" + assert not should_rate_limit_called, ( + "should_rate_limit should not be called when model-specific limits don't match requested model" + ) # Test 2: Model-specific rate limits configured for requested model - SHOULD be called should_rate_limit_called = False user_api_key_dict_with_matching_model_limits = UserAPIKeyAuth( api_key=_api_key, - metadata={ - "model_tpm_limit": {"gpt-3.5-turbo": 1000} - }, # Rate limit for requested model + metadata={"model_tpm_limit": {"gpt-3.5-turbo": 1000}}, # Rate limit for requested model ) await parallel_request_handler.async_pre_call_hook( @@ -1135,9 +1138,9 @@ async def test_model_specific_rate_limits_only_called_when_configured_v3(): call_type="", ) - assert ( - should_rate_limit_called - ), "should_rate_limit should be called when model-specific limits match requested model" + assert should_rate_limit_called, ( + "should_rate_limit should be called when model-specific limits match requested model" + ) @pytest.mark.asyncio @@ -1164,9 +1167,7 @@ async def test_tpm_api_key_rate_limits_v3(): user_api_key_dict.metadata["model_rpm_limit"] = rpms local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock should_rate_limit to capture the descriptors captured_descriptors = None @@ -1224,15 +1225,11 @@ async def test_tpm_api_key_rate_limits_v3(): break assert model_per_key_descriptor is not None, "Api-Key descriptor should be present" - assert ( - model_per_key_descriptor["value"] == f"{_api_key_hash}:{model}" - ), "Api-Key value should combine api_key and model" - assert ( - model_per_key_descriptor["rate_limit"]["requests_per_unit"] == rpm_limit - ), "Api-Key RPM limit should be set" - assert ( - model_per_key_descriptor["rate_limit"]["tokens_per_unit"] == tpm_limit - ), "Api-Key TPM limit should be set" + assert model_per_key_descriptor["value"] == f"{_api_key_hash}:{model}", ( + "Api-Key value should combine api_key and model" + ) + assert model_per_key_descriptor["rate_limit"]["requests_per_unit"] == rpm_limit, "Api-Key RPM limit should be set" + assert model_per_key_descriptor["rate_limit"]["tokens_per_unit"] == tpm_limit, "Api-Key TPM limit should be set" @pytest.mark.asyncio @@ -1259,9 +1256,7 @@ async def test_rpm_api_key_rate_limits_v3(): user_api_key_dict.metadata["model_rpm_limit"] = rpms local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock should_rate_limit to capture the descriptors captured_descriptors = None @@ -1319,15 +1314,11 @@ async def test_rpm_api_key_rate_limits_v3(): break assert model_per_key_descriptor is not None, "Api-Key descriptor should be present" - assert ( - model_per_key_descriptor["value"] == f"{_api_key_hash}:{model}" - ), "Api-Key value should combine api_key and model" - assert ( - model_per_key_descriptor["rate_limit"]["requests_per_unit"] == rpm_limit - ), "Api-Key RPM limit should be set" - assert ( - model_per_key_descriptor["rate_limit"]["tokens_per_unit"] == tpm_limit - ), "Api-Key TPM limit should be set" + assert model_per_key_descriptor["value"] == f"{_api_key_hash}:{model}", ( + "Api-Key value should combine api_key and model" + ) + assert model_per_key_descriptor["rate_limit"]["requests_per_unit"] == rpm_limit, "Api-Key RPM limit should be set" + assert model_per_key_descriptor["rate_limit"]["tokens_per_unit"] == tpm_limit, "Api-Key TPM limit should be set" @pytest.mark.asyncio @@ -1349,9 +1340,7 @@ async def test_team_member_rate_limits_v3(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock should_rate_limit to capture the descriptors captured_descriptors = None @@ -1383,18 +1372,12 @@ async def test_team_member_rate_limits_v3(): team_member_descriptor = descriptor break - assert ( - team_member_descriptor is not None - ), "Team member descriptor should be present" - assert ( - team_member_descriptor["value"] == f"{_team_id}:{_user_id}" - ), "Team member value should combine team_id and user_id" - assert ( - team_member_descriptor["rate_limit"]["requests_per_unit"] == 10 - ), "Team member RPM limit should be set" - assert ( - team_member_descriptor["rate_limit"]["tokens_per_unit"] == 1000 - ), "Team member TPM limit should be set" + assert team_member_descriptor is not None, "Team member descriptor should be present" + assert team_member_descriptor["value"] == f"{_team_id}:{_user_id}", ( + "Team member value should combine team_id and user_id" + ) + assert team_member_descriptor["rate_limit"]["requests_per_unit"] == 10, "Team member RPM limit should be set" + assert team_member_descriptor["rate_limit"]["tokens_per_unit"] == 1000, "Team member TPM limit should be set" @pytest.mark.asyncio @@ -1417,9 +1400,7 @@ async def test_team_member_rate_limits_v3_raises_429_when_over_limit(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) captured_descriptors = None @@ -1495,9 +1476,7 @@ async def test_dynamic_rate_limiting_v3(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock should_rate_limit to track if limits are enforced captured_descriptors = [] @@ -1530,9 +1509,9 @@ async def test_dynamic_rate_limiting_v3(): break assert api_key_descriptor is not None, "API key descriptor should be present" - assert ( - api_key_descriptor["rate_limit"]["requests_per_unit"] is None - ), "RPM limit should be None when dynamic mode and no failures" + assert api_key_descriptor["rate_limit"]["requests_per_unit"] is None, ( + "RPM limit should be None when dynamic mode and no failures" + ) # Test 2: With failures - rate limits SHOULD be enforced (rpm_limit should be set) async def mock_check_with_failures(*args, **kwargs): @@ -1556,9 +1535,9 @@ async def test_dynamic_rate_limiting_v3(): break assert api_key_descriptor is not None, "API key descriptor should be present" - assert ( - api_key_descriptor["rate_limit"]["requests_per_unit"] == 2 - ), "RPM limit should be enforced when dynamic mode and failures detected" + assert api_key_descriptor["rate_limit"]["requests_per_unit"] == 2, ( + "RPM limit should be enforced when dynamic mode and failures detected" + ) @pytest.mark.flaky(retries=3, delay=2) @@ -1604,9 +1583,7 @@ async def test_async_increment_tokens_with_ttl_preservation(): ) local_cache = DualCache(redis_cache=redis_cache) - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Verify Redis connection is working try: @@ -1616,9 +1593,7 @@ async def test_async_increment_tokens_with_ttl_preservation(): # Verify the TTL preservation script is registered if parallel_request_handler.token_increment_script is None: - pytest.skip( - "Token increment script not available - Redis Lua scripting may not be supported" - ) + pytest.skip("Token increment script not available - Redis Lua scripting may not be supported") # Test keys - use hash tags to ensure they map to same Redis cluster slot # Use a unique suffix per test run to avoid stale state from prior runs @@ -1639,11 +1614,11 @@ async def test_async_increment_tokens_with_ttl_preservation(): # First increment: Create operations with mixed TTL scenarios pipeline_operations_first = [ + RedisPipelineIncrementOperation(key=test_key_with_ttl, increment_value=10.0, ttl=60), RedisPipelineIncrementOperation( - key=test_key_with_ttl, increment_value=10.0, ttl=60 - ), - RedisPipelineIncrementOperation( - key=test_key_without_ttl, increment_value=5.0, ttl=None # No TTL + key=test_key_without_ttl, + increment_value=5.0, + ttl=None, # No TTL ), ] @@ -1657,29 +1632,21 @@ async def test_async_increment_tokens_with_ttl_preservation(): # Verify keys exist and check initial TTL ttl_after_first = await redis_cache.async_get_ttl(test_key_with_ttl) - value_after_first_with_ttl = await redis_cache.async_get_cache( - test_key_with_ttl - ) - value_after_first_without_ttl = await redis_cache.async_get_cache( - test_key_without_ttl - ) + value_after_first_with_ttl = await redis_cache.async_get_cache(test_key_with_ttl) + value_after_first_without_ttl = await redis_cache.async_get_cache(test_key_without_ttl) - assert ( - value_after_first_with_ttl == 10.0 - ), f"First increment should set value to 10.0, got {value_after_first_with_ttl}" - assert ( - value_after_first_without_ttl == 5.0 - ), "First increment should set value to 5.0" - assert ( - ttl_after_first is not None and ttl_after_first > 0 - ), "Key with TTL should have positive TTL after first increment" + assert value_after_first_with_ttl == 10.0, ( + f"First increment should set value to 10.0, got {value_after_first_with_ttl}" + ) + assert value_after_first_without_ttl == 5.0, "First increment should set value to 5.0" + assert ttl_after_first is not None and ttl_after_first > 0, ( + "Key with TTL should have positive TTL after first increment" + ) assert ttl_after_first <= 60, "TTL should not exceed the set value" # Check TTL for key without TTL (should be None, meaning no expiry) ttl_no_ttl_key = await redis_cache.async_get_ttl(test_key_without_ttl) - assert ( - ttl_no_ttl_key is None - ), "Key without TTL should have no expiry (None from async_get_ttl)" + assert ttl_no_ttl_key is None, "Key without TTL should have no expiry (None from async_get_ttl)" # Wait a moment to ensure TTL decreases await asyncio.sleep(2) @@ -1687,10 +1654,14 @@ async def test_async_increment_tokens_with_ttl_preservation(): # Second increment: Same operations to test TTL preservation pipeline_operations_second = [ RedisPipelineIncrementOperation( - key=test_key_with_ttl, increment_value=15.0, ttl=60 # Same TTL value + key=test_key_with_ttl, + increment_value=15.0, + ttl=60, # Same TTL value ), RedisPipelineIncrementOperation( - key=test_key_without_ttl, increment_value=7.0, ttl=None # No TTL + key=test_key_without_ttl, + increment_value=7.0, + ttl=None, # No TTL ), ] @@ -1704,39 +1675,23 @@ async def test_async_increment_tokens_with_ttl_preservation(): # Verify TTL preservation and value updates ttl_after_second = await redis_cache.async_get_ttl(test_key_with_ttl) - value_after_second_with_ttl = await redis_cache.async_get_cache( - test_key_with_ttl - ) - value_after_second_without_ttl = await redis_cache.async_get_cache( - test_key_without_ttl - ) + value_after_second_with_ttl = await redis_cache.async_get_cache(test_key_with_ttl) + value_after_second_without_ttl = await redis_cache.async_get_cache(test_key_without_ttl) - assert ( - value_after_second_with_ttl == 25.0 - ), "Second increment should update value to 25.0" - assert ( - value_after_second_without_ttl == 12.0 - ), "Second increment should update value to 12.0" + assert value_after_second_with_ttl == 25.0, "Second increment should update value to 25.0" + assert value_after_second_without_ttl == 12.0, "Second increment should update value to 12.0" # Critical test: TTL should be preserved (not reset to 60) assert ttl_after_second is not None, "TTL should still exist" - assert ( - ttl_after_second < ttl_after_first - ), "TTL should have decreased (not been reset)" + assert ttl_after_second < ttl_after_first, "TTL should have decreased (not been reset)" assert ttl_after_second > 0, "TTL should still be positive" # TTL should not be close to the original 60 seconds (proving it wasn't reset) - assert ( - ttl_after_second < 59 - ), "TTL should be significantly less than original, proving preservation" + assert ttl_after_second < 59, "TTL should be significantly less than original, proving preservation" # Key without TTL should still have no expiry - ttl_no_ttl_key_after_second = await redis_cache.async_get_ttl( - test_key_without_ttl - ) - assert ( - ttl_no_ttl_key_after_second is None - ), "Key without TTL should still have no expiry" + ttl_no_ttl_key_after_second = await redis_cache.async_get_ttl(test_key_without_ttl) + assert ttl_no_ttl_key_after_second is None, "Key without TTL should still have no expiry" finally: # Clean up test keys @@ -1763,44 +1718,30 @@ async def test_async_increment_tokens_fallback_behavior(): from litellm.types.caching import RedisPipelineIncrementOperation local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock the token_increment_script to None to simulate unavailable script parallel_request_handler.token_increment_script = None # Mock the fallback method fallback_called = False - original_method = ( - parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline - ) + original_method = parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline async def mock_fallback(*args, **kwargs): nonlocal fallback_called fallback_called = True return await original_method(*args, **kwargs) - parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_fallback - ) + parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_fallback # Test operations - pipeline_operations = [ - RedisPipelineIncrementOperation( - key="test_fallback_key", increment_value=10.0, ttl=60 - ) - ] + pipeline_operations = [RedisPipelineIncrementOperation(key="test_fallback_key", increment_value=10.0, ttl=60)] # Execute increment - await parallel_request_handler.async_increment_tokens_with_ttl_preservation( - pipeline_operations=pipeline_operations - ) + await parallel_request_handler.async_increment_tokens_with_ttl_preservation(pipeline_operations=pipeline_operations) # Verify fallback was called - assert ( - fallback_called - ), "Fallback method should be called when Lua script is not available" + assert fallback_called, "Fallback method should be called when Lua script is not available" # Redis Cluster Compatibility Tests @@ -1811,9 +1752,7 @@ def test_group_keys_by_hash_tag_regular_redis(): For regular Redis, all keys should be grouped together under a single group. """ local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Test keys with different hash tags test_keys = [ @@ -1833,9 +1772,7 @@ def test_group_keys_by_hash_tag_regular_redis(): # Verify all keys are in single group for regular Redis assert len(groups) == 1, f"Expected 1 group for regular Redis, got {len(groups)}" assert "all_keys" in groups, "Expected 'all_keys' group for regular Redis" - assert set(groups["all_keys"]) == set( - test_keys - ), "All keys should be in single group" + assert set(groups["all_keys"]) == set(test_keys), "All keys should be in single group" def test_group_keys_by_hash_tag_redis_cluster(): @@ -1847,9 +1784,7 @@ def test_group_keys_by_hash_tag_redis_cluster(): from unittest.mock import patch local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock _is_redis_cluster to return True with patch.object(handler, "_is_redis_cluster", return_value=True): @@ -1869,17 +1804,13 @@ def test_group_keys_by_hash_tag_redis_cluster(): # All group keys should start with "slot_" for group_key in groups.keys(): - assert group_key.startswith( - "slot_" - ), f"Group key {group_key} should start with 'slot_'" + assert group_key.startswith("slot_"), f"Group key {group_key} should start with 'slot_'" # Verify all original keys are present across groups all_grouped_keys = [] for group_keys in groups.values(): all_grouped_keys.extend(group_keys) - assert set(all_grouped_keys) == set( - test_keys - ), "All keys should be present in groups" + assert set(all_grouped_keys) == set(test_keys), "All keys should be present in groups" def test_keyslot_for_redis_cluster(): @@ -1887,9 +1818,7 @@ def test_keyslot_for_redis_cluster(): Test the keyslot calculation for Redis cluster. """ local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Test basic key slot1 = handler.keyslot_for_redis_cluster("user:1000") @@ -1917,18 +1846,14 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): from unittest.mock import AsyncMock, patch local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock _is_redis_cluster to return True for this test with patch.object(handler, "_is_redis_cluster", return_value=True): # Mock script that simulates Redis cluster slot conflict mock_script = AsyncMock() mock_script.side_effect = [ - Exception( - "EVALSHA - all keys must map to the same key slot" - ), # First group fails + Exception("EVALSHA - all keys must map to the same key slot"), # First group fails [1234, 1, 1234, 2], # Second group succeeds ] handler.batch_rate_limiter_script = mock_script @@ -1945,9 +1870,7 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): ] # Execute the method - results = await handler._execute_redis_batch_rate_limiter_script( - keys_to_fetch=test_keys, now_int=1234 - ) + results = await handler._execute_redis_batch_rate_limiter_script(keys_to_fetch=test_keys, now_int=1234) # Verify results: 2 from fallback + 4 from successful script = 6 total assert len(results) == 6, f"Expected 6 results, got {len(results)}" @@ -1972,9 +1895,7 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): # Should have processed all keys (some might be duplicated due to fallback) unique_processed_keys = set(all_processed_keys) - assert ( - len(unique_processed_keys) >= 2 - ), "Should have processed at least some keys" + assert len(unique_processed_keys) >= 2, "Should have processed at least some keys" @pytest.mark.asyncio @@ -2000,9 +1921,7 @@ async def test_multiple_rate_limits_per_descriptor(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock should_rate_limit to return a response with multiple statuses where one hits the limit # This simulates the case where we have more statuses than descriptors due to multiple rate limit types @@ -2080,9 +1999,7 @@ async def test_missing_descriptor_fallback(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock should_rate_limit to return a status with descriptor_key that doesn't match descriptors async def mock_should_rate_limit(descriptors, **kwargs): @@ -2126,9 +2043,7 @@ async def test_get_rate_limit_type_default_is_total(monkeypatch): This verifies the change from 'output' to 'total' as the default value. """ local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock general_settings to return empty dict (no token_rate_limit_type set) import litellm.proxy.proxy_server as proxy_server @@ -2138,9 +2053,7 @@ async def test_get_rate_limit_type_default_is_total(monkeypatch): try: result = parallel_request_handler.get_rate_limit_type() - assert ( - result == "total" - ), f"Default rate limit type should be 'total', got '{result}'" + assert result == "total", f"Default rate limit type should be 'total', got '{result}'" finally: monkeypatch.setattr(proxy_server, "general_settings", original_settings) @@ -2151,23 +2064,17 @@ async def test_get_rate_limit_type_invalid_falls_back_to_total(monkeypatch): Test that get_rate_limit_type falls back to 'total' when an invalid value is specified. """ local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock general_settings to return an invalid token_rate_limit_type import litellm.proxy.proxy_server as proxy_server original_settings = getattr(proxy_server, "general_settings", {}) - monkeypatch.setattr( - proxy_server, "general_settings", {"token_rate_limit_type": "invalid_type"} - ) + monkeypatch.setattr(proxy_server, "general_settings", {"token_rate_limit_type": "invalid_type"}) try: result = parallel_request_handler.get_rate_limit_type() - assert ( - result == "total" - ), f"Invalid rate limit type should fall back to 'total', got '{result}'" + assert result == "total", f"Invalid rate limit type should fall back to 'total', got '{result}'" finally: monkeypatch.setattr(proxy_server, "general_settings", original_settings) @@ -2181,9 +2088,7 @@ async def test_get_rate_limit_type_invalid_falls_back_to_total(monkeypatch): ], ) @pytest.mark.asyncio -async def test_async_log_success_event_with_dict_usage( - monkeypatch, token_rate_limit_type, expected_field -): +async def test_async_log_success_event_with_dict_usage(monkeypatch, token_rate_limit_type, expected_field): """ Test that async_log_success_event correctly handles usage as a dict (Responses API format). @@ -2195,17 +2100,13 @@ async def test_async_log_success_event_with_dict_usage( _api_key = "sk-12345" _api_key = hash_token(_api_key) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock the get_rate_limit_type method def mock_get_rate_limit_type(): return token_rate_limit_type - monkeypatch.setattr( - parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type - ) + monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type) # Create a mock response object with usage as a dict (Responses API format) from litellm.types.utils import BaseLiteLLMOpenAIResponseObject @@ -2268,9 +2169,9 @@ async def test_async_log_success_event_with_dict_usage( "total": 60, # total_tokens } - assert ( - tpm_operation["increment_value"] == expected_tokens[token_rate_limit_type] - ), f"Expected {expected_tokens[token_rate_limit_type]} tokens for type '{token_rate_limit_type}', got {tpm_operation['increment_value']}" + assert tpm_operation["increment_value"] == expected_tokens[token_rate_limit_type], ( + f"Expected {expected_tokens[token_rate_limit_type]} tokens for type '{token_rate_limit_type}', got {tpm_operation['increment_value']}" + ) @pytest.mark.asyncio @@ -2285,17 +2186,13 @@ async def test_async_log_success_event_with_dict_usage_missing_fields(monkeypatc _api_key = "sk-12345" _api_key = hash_token(_api_key) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock the get_rate_limit_type method def mock_get_rate_limit_type(): return "output" - monkeypatch.setattr( - parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type - ) + monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type) # Create a mock response object with usage as a dict missing some fields mock_response = MagicMock() @@ -2306,9 +2203,7 @@ async def test_async_log_success_event_with_dict_usage_missing_fields(monkeypatc } from litellm.types.utils import BaseLiteLLMOpenAIResponseObject - mock_response.__class__ = type( - "MockResponse", (BaseLiteLLMOpenAIResponseObject,), {} - ) + mock_response.__class__ = type("MockResponse", (BaseLiteLLMOpenAIResponseObject,), {}) # Create mock kwargs for the success event mock_kwargs = { @@ -2364,9 +2259,7 @@ async def test_execute_token_increment_script_cluster_compatibility(): from litellm.types.caching import RedisPipelineIncrementOperation local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Mock _is_redis_cluster to return True for this test with patch.object(handler, "_is_redis_cluster", return_value=True): @@ -2404,18 +2297,14 @@ async def test_execute_token_increment_script_cluster_compatibility(): "{api_key:sk-123}:max_parallel_requests", "{user:user-456}:tokens", } - assert ( - set(all_processed_keys) == expected_keys - ), "All operation keys should be processed" + assert set(all_processed_keys) == expected_keys, "All operation keys should be processed" # Verify args structure is correct for each call for call_args in call_args_list: keys = call_args[1]["keys"] args = call_args[1]["args"] # Each key should have 2 args (increment_value, ttl) - assert ( - len(args) == len(keys) * 2 - ), f"Each key should have 2 args, got {len(args)} args for {len(keys)} keys" + assert len(args) == len(keys) * 2, f"Each key should have 2 args, got {len(args)} args for {len(keys)} keys" @pytest.mark.asyncio @@ -2438,9 +2327,7 @@ async def test_agent_level_rate_limit_descriptors(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) mock_agent = AgentResponse( agent_id=_agent_id, @@ -2505,9 +2392,7 @@ async def test_agent_session_rate_limit_descriptors(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) mock_agent = AgentResponse( agent_id=_agent_id, @@ -2574,9 +2459,7 @@ async def test_agent_session_rate_limit_skipped_without_session_id(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) mock_agent = AgentResponse( agent_id=_agent_id, @@ -2609,8 +2492,7 @@ async def test_agent_session_rate_limit_skipped_without_session_id(): # should_rate_limit should not have been called (no agent-level limits, only session limits # but no session_id) assert captured_descriptors is None, ( - "No descriptors should be created when agent has only session limits " - "but no session_id in request" + "No descriptors should be created when agent has only session limits but no session_id in request" ) @@ -2633,9 +2515,7 @@ async def test_agent_rate_limit_from_metadata_agent_id(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) mock_agent = AgentResponse( agent_id=_agent_id, @@ -2676,9 +2556,7 @@ async def test_agent_rate_limit_from_metadata_agent_id(): agent_descriptor = d break - assert ( - agent_descriptor is not None - ), "Agent descriptor should be created from metadata agent_id" + assert agent_descriptor is not None, "Agent descriptor should be created from metadata agent_id" assert agent_descriptor["value"] == _agent_id assert agent_descriptor["rate_limit"]["requests_per_unit"] == 25 @@ -2704,9 +2582,7 @@ async def test_agent_both_agent_and_session_rate_limits(): ) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) mock_agent = AgentResponse( agent_id=_agent_id, @@ -2773,16 +2649,12 @@ async def test_agent_rate_limit_tpm_increment_on_success(monkeypatch): _session_id = "sess_tpm_test" local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) def mock_get_rate_limit_type(): return "total" - monkeypatch.setattr( - parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type - ) + monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type) mock_usage = Usage(prompt_tokens=20, completion_tokens=30, total_tokens=50) mock_response = ModelResponse( @@ -2893,18 +2765,12 @@ async def test_agent_rate_limit_429_on_over_limit(monkeypatch, time_controller): window_starts[window_key] = now new_counter = 1 request_counts[counter_key] = new_counter - await local_cache.async_set_cache( - key=window_key, value=now, ttl=window_size - ) - await local_cache.async_set_cache( - key=counter_key, value=new_counter, ttl=window_size - ) + await local_cache.async_set_cache(key=window_key, value=now, ttl=window_size) + await local_cache.async_set_cache(key=counter_key, value=new_counter, ttl=window_size) else: new_counter = prev_counter + 1 request_counts[counter_key] = new_counter - await local_cache.async_set_cache( - key=counter_key, value=new_counter, ttl=window_size - ) + await local_cache.async_set_cache(key=counter_key, value=new_counter, ttl=window_size) results.append(now) results.append(new_counter) return results @@ -3042,9 +2908,7 @@ async def test_project_model_rate_limits_enforced_v3(): """ _api_key = hash_token("sk-project-test") local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) captured_descriptors = [] @@ -3072,13 +2936,9 @@ async def test_project_model_rate_limits_enforced_v3(): ) descriptor_keys = [d["key"] for d in captured_descriptors] - assert ( - "model_per_project" in descriptor_keys - ), f"Expected model_per_project descriptor, got: {descriptor_keys}" + assert "model_per_project" in descriptor_keys, f"Expected model_per_project descriptor, got: {descriptor_keys}" - model_per_project = next( - d for d in captured_descriptors if d["key"] == "model_per_project" - ) + model_per_project = next(d for d in captured_descriptors if d["key"] == "model_per_project") assert model_per_project["value"] == "proj-abc123:gpt-4" assert model_per_project["rate_limit"]["requests_per_unit"] == 5 assert model_per_project["rate_limit"]["tokens_per_unit"] == 1000 @@ -3089,9 +2949,7 @@ async def test_project_model_rate_limits_not_triggered_for_other_model_v3(): """Project model limits should not trigger for a model not in project_metadata.""" _api_key = hash_token("sk-project-test-2") local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) captured_descriptors = [] @@ -3118,9 +2976,9 @@ async def test_project_model_rate_limits_not_triggered_for_other_model_v3(): ) descriptor_keys = [d["key"] for d in captured_descriptors] - assert ( - "model_per_project" not in descriptor_keys - ), f"model_per_project should not be added for unrelated model, got: {descriptor_keys}" + assert "model_per_project" not in descriptor_keys, ( + f"model_per_project should not be added for unrelated model, got: {descriptor_keys}" + ) @pytest.mark.asyncio @@ -3131,9 +2989,7 @@ async def test_project_model_itpm_otpm_limits_enforced_v3(): """ _api_key = hash_token("sk-project-io-test") local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) captured_descriptors = [] @@ -3164,12 +3020,8 @@ async def test_project_model_itpm_otpm_limits_enforced_v3(): assert "model_per_project_otpm" in descriptor_keys assert "model_per_project" not in descriptor_keys - itpm_descriptor = next( - d for d in captured_descriptors if d["key"] == "model_per_project_itpm" - ) - otpm_descriptor = next( - d for d in captured_descriptors if d["key"] == "model_per_project_otpm" - ) + itpm_descriptor = next(d for d in captured_descriptors if d["key"] == "model_per_project_itpm") + otpm_descriptor = next(d for d in captured_descriptors if d["key"] == "model_per_project_otpm") assert itpm_descriptor["value"] == "proj-mantle:bedrock_mantle/claude-opus" assert itpm_descriptor["rate_limit"]["tokens_per_unit"] == 20000000 assert otpm_descriptor["value"] == "proj-mantle:bedrock_mantle/claude-opus" @@ -3181,9 +3033,7 @@ async def test_project_model_itpm_otpm_limits_not_triggered_for_other_model_v3() """Split project limits must not apply to an unrelated model.""" _api_key = hash_token("sk-project-io-test-2") local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) captured_descriptors = [] @@ -3218,9 +3068,7 @@ async def test_project_model_itpm_and_tpm_limits_coexist_v3(): """Combined project TPM and split ITPM/OTPM limits are enforced together.""" _api_key = hash_token("sk-project-io-test-3") local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) captured_descriptors = [] @@ -3262,9 +3110,7 @@ async def test_enforce_project_io_token_quota_for_frame_blocks_over_limit_otpm() limit and reject once a frame's estimated output tokens exceed it.""" _api_key = hash_token("sk-ws-frame-otpm") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth( api_key=_api_key, project_id="proj-mantle-ws", @@ -3296,9 +3142,7 @@ async def test_enforce_project_io_token_quota_for_frame_noop_without_project_lim the per-frame check (no descriptors to reserve against).""" _api_key = hash_token("sk-ws-frame-no-limits") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) await handler.enforce_project_io_token_quota_for_frame( @@ -3433,17 +3277,13 @@ async def test_chat_tpm_refund_and_slot_release_via_context_stash(monkeypatch): monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) _api_key = hash_token("sk-refund-lifecycle") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth( api_key=_api_key, tpm_limit=10_000, max_parallel_requests=2, ) - tokens_key = handler.create_rate_limit_keys( - key="api_key", value=_api_key, rate_limit_type="tokens" - ) + tokens_key = handler.create_rate_limit_keys(key="api_key", value=_api_key, rate_limit_type="tokens") parallel_key = f"{{api_key:{_api_key}}}:max_parallel_requests" await handler.async_pre_call_hook( @@ -3460,28 +3300,18 @@ async def test_chat_tpm_refund_and_slot_release_via_context_stash(monkeypatch): reserved = get_request_stash().reserved_tokens assert reserved > 0 assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == reserved - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=parallel_key) - ) == 1 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=parallel_key)) == 1 kwargs = {"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}} - await handler.async_log_failure_event( - kwargs=kwargs, response_obj=None, start_time=None, end_time=None - ) + await handler.async_log_failure_event(kwargs=kwargs, response_obj=None, start_time=None, end_time=None) assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 0 - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=parallel_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=parallel_key)) == 0 assert get_request_stash().reservation_released is True - await handler.async_log_failure_event( - kwargs=kwargs, response_obj=None, start_time=None, end_time=None - ) + await handler.async_log_failure_event(kwargs=kwargs, response_obj=None, start_time=None, end_time=None) assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 0 - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=parallel_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=parallel_key)) == 0 @pytest.mark.asyncio @@ -3525,9 +3355,7 @@ async def test_pre_call_hook_ignores_caller_supplied_stash_values(): async def spy_increment_pipeline(increment_list, **kwargs): refund_calls.append(increment_list) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - spy_increment_pipeline - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = spy_increment_pipeline await handler.async_post_call_failure_hook( request_data=data, @@ -3553,17 +3381,13 @@ async def test_log_events_from_nested_calls_leave_owner_stash_alone(monkeypatch) monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) _api_key = hash_token("sk-nested-guard") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth( api_key=_api_key, tpm_limit=10_000, max_parallel_requests=2, ) - tokens_key = handler.create_rate_limit_keys( - key="api_key", value=_api_key, rate_limit_type="tokens" - ) + tokens_key = handler.create_rate_limit_keys(key="api_key", value=_api_key, rate_limit_type="tokens") parallel_key = f"{{api_key:{_api_key}}}:max_parallel_requests" await handler.async_pre_call_hook( @@ -3588,33 +3412,23 @@ async def test_log_events_from_nested_calls_leave_owner_stash_alone(monkeypatch) "litellm_call_id": "nested-guardrail-call", "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, } - await handler.async_log_success_event( - kwargs=nested_kwargs, response_obj=None, start_time=None, end_time=None - ) - await handler.async_log_failure_event( - kwargs=nested_kwargs, response_obj=None, start_time=None, end_time=None - ) + await handler.async_log_success_event(kwargs=nested_kwargs, response_obj=None, start_time=None, end_time=None) + await handler.async_log_failure_event(kwargs=nested_kwargs, response_obj=None, start_time=None, end_time=None) assert stash.parallel_slot is not None assert stash.reservation_released is False - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=parallel_key) - ) == 1 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=parallel_key)) == 1 assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == reserved owner_kwargs = { "litellm_call_id": "owner-call-id", "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, } - await handler.async_log_failure_event( - kwargs=owner_kwargs, response_obj=None, start_time=None, end_time=None - ) + await handler.async_log_failure_event(kwargs=owner_kwargs, response_obj=None, start_time=None, end_time=None) assert stash.parallel_slot is None assert stash.reservation_released is True - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=parallel_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=parallel_key)) == 0 assert int(await local_cache.async_get_cache(key=tokens_key) or 0) == 0 @@ -3628,9 +3442,7 @@ async def test_stash_applies_when_owner_or_callback_call_id_missing(): do not thread ``litellm_call_id`` into their logging kwargs. """ local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) unclaimed = get_or_create_request_stash() unclaimed.reserved_tokens = 42 @@ -3642,9 +3454,7 @@ async def test_stash_applies_when_owner_or_callback_call_id_missing(): ) assert unclaimed.reservation_released is True - claimed = RequestRateLimiterStash( - owner_litellm_call_id="owner-1", reserved_tokens=42 - ) + claimed = RequestRateLimiterStash(owner_litellm_call_id="owner-1", reserved_tokens=42) _request_stash.set(claimed) await handler.async_log_failure_event( kwargs={"standard_logging_object": {}}, @@ -3827,9 +3637,7 @@ async def test_failure_event_settles_project_itpm_otpm_at_recovered_partial_usag def _make_mcp_handler(): local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) return handler, local_cache @@ -3856,9 +3664,7 @@ def test_mcp_per_key_descriptor_created_for_matching_server_v3(): metadata={"mcp_rpm_limit": {"github": 5}}, ) - descriptors = _build_mcp_descriptors( - handler, user_api_key_dict, {"mcp_server_name": "github"} - ) + descriptors = _build_mcp_descriptors(handler, user_api_key_dict, {"mcp_server_name": "github"}) descriptor = _find_descriptor(descriptors, "mcp_per_key") assert descriptor is not None @@ -3876,9 +3682,7 @@ def test_mcp_per_key_descriptor_skipped_for_non_matching_server_v3(): metadata={"mcp_rpm_limit": {"github": 5}}, ) - descriptors = _build_mcp_descriptors( - handler, user_api_key_dict, {"mcp_server_name": "slack"} - ) + descriptors = _build_mcp_descriptors(handler, user_api_key_dict, {"mcp_server_name": "slack"}) assert _find_descriptor(descriptors, "mcp_per_key") is None @@ -3935,9 +3739,7 @@ def test_mcp_per_team_descriptor_created_from_team_metadata_v3(): team_metadata={"mcp_rpm_limit": {"github": 3}}, ) - descriptors = _build_mcp_descriptors( - handler, user_api_key_dict, {"mcp_server_name": "github"} - ) + descriptors = _build_mcp_descriptors(handler, user_api_key_dict, {"mcp_server_name": "github"}) descriptor = _find_descriptor(descriptors, "mcp_per_team") assert descriptor is not None @@ -3956,9 +3758,7 @@ async def test_mcp_per_key_rpm_enforced_v3(monkeypatch): monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") api_key = hash_token("sk-mcp-enforce") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) window_starts: Dict[str, int] = {} request_counts: Dict[str, int] = {} @@ -4052,9 +3852,7 @@ def test_get_key_mcp_rpm_limit_precedence(): _TEST_SLOT_ID = "slot-disconnect-test" -async def _seed_max_parallel_requests_slots( - dual_cache: DualCache, counter_key: str, slot_ids: List[str] -) -> None: +async def _seed_max_parallel_requests_slots(dual_cache: DualCache, counter_key: str, slot_ids: List[str]) -> None: await dual_cache.async_set_cache( key=counter_key, value={slot_id: time.time() for slot_id in slot_ids}, @@ -4185,9 +3983,7 @@ async def _build_seeded_limiter(): """Build a v3 limiter whose api-key slot registry already holds the pre-call slot.""" api_key = hash_token("sk-disconnect") cache = DualCache() - limiter = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(cache) - ) + limiter = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) counter_key = f"{{api_key:{api_key}}}:max_parallel_requests" await _seed_max_parallel_requests_slots(cache, counter_key, [_TEST_SLOT_ID]) user_api_key_dict = UserAPIKeyAuth(api_key=api_key, max_parallel_requests=2) @@ -4222,16 +4018,12 @@ async def test_release_max_parallel_requests_on_disconnect_v3(): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=2) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" await _seed_max_parallel_requests_slots(local_cache, counter_key, [_TEST_SLOT_ID]) - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 1 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 1 get_or_create_request_stash().parallel_slot = ParallelSlotAcquisition( slot_id=_TEST_SLOT_ID, @@ -4240,9 +4032,7 @@ async def test_release_max_parallel_requests_on_disconnect_v3(): await handler.async_release_max_parallel_requests_on_disconnect(user_api_key_dict) assert get_request_stash().parallel_slot is None - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 0 @pytest.mark.asyncio @@ -4255,9 +4045,7 @@ async def test_release_on_disconnect_works_when_key_config_changed_v3(): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" await _seed_max_parallel_requests_slots(local_cache, counter_key, [_TEST_SLOT_ID]) @@ -4268,9 +4056,7 @@ async def test_release_on_disconnect_works_when_key_config_changed_v3(): await handler.async_release_max_parallel_requests_on_disconnect( UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None) ) - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 0 @pytest.mark.asyncio @@ -4286,9 +4072,7 @@ async def test_post_call_failure_hook_releases_parallel_slot_v3(): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=1) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" @@ -4299,18 +4083,14 @@ async def test_post_call_failure_hook_releases_parallel_slot_v3(): data=admitted_data, call_type="", ) - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 1 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 1 await handler.async_post_call_failure_hook( request_data=admitted_data, original_exception=Exception("guardrail rejected the request"), user_api_key_dict=user_api_key_dict, ) - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 0 await handler.async_log_failure_event( kwargs={ @@ -4320,9 +4100,7 @@ async def test_post_call_failure_hook_releases_parallel_slot_v3(): start_time=None, end_time=None, ) - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 0 await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -4341,9 +4119,7 @@ async def test_success_event_releases_parallel_slot_v3(monkeypatch): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) monkeypatch.setattr(handler, "get_rate_limit_type", lambda: "total") user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=1) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" @@ -4355,23 +4131,17 @@ async def test_success_event_releases_parallel_slot_v3(monkeypatch): data=admitted_data, call_type="", ) - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 1 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 1 await handler.async_log_success_event( kwargs={ "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, }, - response_obj=ModelResponse( - usage=Usage(prompt_tokens=5, completion_tokens=5, total_tokens=10) - ), + response_obj=ModelResponse(usage=Usage(prompt_tokens=5, completion_tokens=5, total_tokens=10)), start_time=datetime.now(), end_time=datetime.now(), ) - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 0 await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -4391,9 +4161,7 @@ async def test_read_only_gauge_check_counts_without_acquiring_v3(): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" descriptors = [ { @@ -4412,9 +4180,7 @@ async def test_read_only_gauge_check_counts_without_acquiring_v3(): handler.parallel_count_script = fake_count response = await handler.should_rate_limit(descriptors=descriptors, read_only=True) - assert captured_calls == [ - ([counter_key], [PARALLEL_REQUEST_SLOT_TTL_SECONDS]) - ] + assert captured_calls == [([counter_key], [PARALLEL_REQUEST_SLOT_TTL_SECONDS])] assert response["overall_code"] == "OK" assert response["statuses"] == [ { @@ -4431,9 +4197,7 @@ async def test_read_only_gauge_check_counts_without_acquiring_v3(): raise ConnectionError("redis unavailable") handler.parallel_count_script = failing_count - await _seed_max_parallel_requests_slots( - local_cache, counter_key, ["s1", "s2", "s3", "s4", "s5"] - ) + await _seed_max_parallel_requests_slots(local_cache, counter_key, ["s1", "s2", "s3", "s4", "s5"]) response = await handler.should_rate_limit(descriptors=descriptors, read_only=True) assert response["overall_code"] == "OVER_LIMIT" assert response["statuses"][0]["rate_limit_type"] == "max_parallel_requests" @@ -4448,9 +4212,7 @@ async def test_redis_release_script_updates_local_mirror_v3(): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" captured_calls = [] @@ -4488,12 +4250,8 @@ async def test_tpm_over_limit_rejection_releases_parallel_slot_v3(monkeypatch): monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) - user_api_key_dict = UserAPIKeyAuth( - api_key=_api_key, max_parallel_requests=5, tpm_limit=100 - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=5, tpm_limit=100) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" async def over_limit_reservation(descriptors, estimated_tokens, parent_otel_span=None): @@ -4520,9 +4278,7 @@ async def test_tpm_over_limit_rejection_releases_parallel_slot_v3(monkeypatch): call_type="", ) assert exc_info.value.status_code == 429 - assert handler._gauge_in_flight_from_cache_value( - await local_cache.async_get_cache(key=counter_key) - ) == 0 + assert handler._gauge_in_flight_from_cache_value(await local_cache.async_get_cache(key=counter_key)) == 0 @pytest.mark.asyncio @@ -4537,9 +4293,7 @@ async def test_in_memory_fallback_respects_mirrored_redis_count_v3(): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=5) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" @@ -4589,9 +4343,7 @@ async def test_release_max_parallel_requests_on_disconnect_noop_v3(): """ _api_key = hash_token("sk-12345") local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" await handler.async_release_max_parallel_requests_on_disconnect( @@ -4622,9 +4374,7 @@ async def test_async_streaming_data_generator_releases_counter_on_disconnect_v3( from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing limiter, cache, counter_key, user_api_key_dict = await _build_seeded_limiter() - assert limiter._gauge_in_flight_from_cache_value( - await cache.async_get_cache(key=counter_key) - ) == 1 + assert limiter._gauge_in_flight_from_cache_value(await cache.async_get_cache(key=counter_key)) == 1 proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter @@ -4655,9 +4405,7 @@ async def test_async_streaming_data_generator_releases_counter_on_disconnect_v3( await gen.aclose() await _drain_release_task() - assert limiter._gauge_in_flight_from_cache_value( - await cache.async_get_cache(key=counter_key) - ) == 0 + assert limiter._gauge_in_flight_from_cache_value(await cache.async_get_cache(key=counter_key)) == 0 @pytest.mark.parametrize("disconnect", ["cancel", "aclose"]) @@ -4704,14 +4452,10 @@ async def test_async_data_generator_releases_counter_on_disconnect_v3(disconnect else: await gen.aclose() await _drain_release_task() - assert limiter._gauge_in_flight_from_cache_value( - await cache.async_get_cache(key=counter_key) - ) == 0 + assert limiter._gauge_in_flight_from_cache_value(await cache.async_get_cache(key=counter_key)) == 0 finally: if saved_hook is not None: - proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = ( - saved_hook - ) + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = saved_hook else: proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None) @@ -4729,9 +4473,7 @@ async def test_async_data_generator_releases_counter_when_wrapped_v3(): import litellm.proxy.proxy_server as proxy_server class _PassthroughIteratorOverride(CustomLogger): - async def async_post_call_streaming_iterator_hook( - self, user_api_key_dict, response, request_data - ): + async def async_post_call_streaming_iterator_hook(self, user_api_key_dict, response, request_data): async for chunk in response: yield chunk @@ -4759,14 +4501,10 @@ async def test_async_data_generator_releases_counter_when_wrapped_v3(): await gen.__anext__() await gen.aclose() await _drain_release_task() - assert limiter._gauge_in_flight_from_cache_value( - await cache.async_get_cache(key=counter_key) - ) == 0 + assert limiter._gauge_in_flight_from_cache_value(await cache.async_get_cache(key=counter_key)) == 0 finally: if saved_hook is not None: - proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = ( - saved_hook - ) + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = saved_hook else: proxy_logging_obj.proxy_hook_mapping.pop("parallel_request_limiter", None) @@ -4774,18 +4512,14 @@ async def test_async_data_generator_releases_counter_when_wrapped_v3(): def test_tpm_reservation_enabled_by_default(monkeypatch): """Upfront TPM reservation is on unless explicitly disabled via env.""" monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) assert handler.tpm_reservation_enabled is True @pytest.mark.parametrize("value", ["false", "False", "FALSE"]) def test_tpm_reservation_disabled_via_env(monkeypatch, value): monkeypatch.setenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", value) - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) assert handler.tpm_reservation_enabled is False @@ -4797,9 +4531,7 @@ async def test_pre_call_hook_reserves_tpm_when_enabled(monkeypatch): only the reservation path owns it. """ monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth(api_key=hash_token("sk-tpm"), tpm_limit=10_000) @@ -4839,9 +4571,7 @@ async def test_pre_call_hook_skips_reservation_when_disabled(monkeypatch): pre-v1.82 post-call accounting behavior. """ monkeypatch.setenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", "false") - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth(api_key=hash_token("sk-tpm"), tpm_limit=10_000) @@ -4889,9 +4619,7 @@ async def test_per_tag_rate_limit_independent_counters_v3(monkeypatch): metadata={"tag_rpm_limit": {"cell-1": 2}}, ) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) async def call(tag: str) -> None: await handler.async_pre_call_hook( @@ -4925,9 +4653,7 @@ async def test_per_tag_descriptor_creation_v3(): api_key=_api_key, metadata={"tag_rpm_limit": {"cell-1": 5}}, ) - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) descriptors = handler._create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, @@ -4951,9 +4677,7 @@ async def test_per_tag_descriptor_absent_without_config_v3(): api_key=hash_token("sk-no-tag"), rpm_limit=10, ) - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) descriptors = handler._create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, @@ -4984,9 +4708,7 @@ async def test_per_tag_untagged_request_governed_by_key_limit_v3(monkeypatch): metadata={"tag_rpm_limit": {"cell-1": 2}}, ) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) async def call(metadata: dict) -> None: await handler.async_pre_call_hook( @@ -5183,9 +4905,7 @@ async def test_streaming_end_to_end_populates_slp_ratelimit_headers(monkeypatch) tpm_limit=10000, ) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) # Real pre-call: populates data and stashes the response into metadata # so the success callback can find it via litellm_params.metadata. @@ -5242,27 +4962,16 @@ async def test_streaming_end_to_end_populates_slp_ratelimit_headers(monkeypatch) call_type="acompletion", ) - additional_headers = ( - mock_kwargs["standard_logging_object"] - .get("hidden_params", {}) - .get("additional_headers", {}) - ) + additional_headers = mock_kwargs["standard_logging_object"].get("hidden_params", {}).get("additional_headers", {}) # api_key-scoped remaining/limit values are the baseline every request # emits and must always reach the SLP. - remaining_keys = [ - k for k in additional_headers if "-remaining-" in k - ] - assert ( - remaining_keys - ), f"streaming success must populate remaining values, got {additional_headers!r}" + remaining_keys = [k for k in additional_headers if "-remaining-" in k] + assert remaining_keys, f"streaming success must populate remaining values, got {additional_headers!r}" limit_keys = [k for k in additional_headers if "-limit-" in k] assert limit_keys, "streaming success must also populate limit values" - assert ( - additional_headers.get("x-ratelimit-api_key-remaining-requests") == 99 - ), ( - "api_key remaining requests should reflect the just-consumed slot;" - f" got {additional_headers!r}" + assert additional_headers.get("x-ratelimit-api_key-remaining-requests") == 99, ( + f"api_key remaining requests should reflect the just-consumed slot; got {additional_headers!r}" ) @@ -5282,9 +4991,7 @@ async def test_streaming_populates_model_per_key_ratelimit_headers(monkeypatch): }, ) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) async def _noop_increment(increment_list, **_): return True @@ -5334,9 +5041,7 @@ async def test_streaming_populates_model_per_key_ratelimit_headers(monkeypatch): hidden_params = mock_kwargs["standard_logging_object"].get("hidden_params") or {} additional_headers = hidden_params.get("additional_headers") or {} - assert ( - additional_headers.get("x-ratelimit-model_per_key-remaining-requests") == 99 - ), f"got {additional_headers!r}" + assert additional_headers.get("x-ratelimit-model_per_key-remaining-requests") == 99, f"got {additional_headers!r}" assert additional_headers.get("x-ratelimit-model_per_key-limit-requests") == 100 # response._hidden_params is also updated for late readers. @@ -5353,9 +5058,7 @@ async def test_async_log_success_event_no_mirror_when_no_snapshot(monkeypatch): """ monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") _api_key = hash_token("sk-stream-no-mirror") - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) async def _noop_increment(increment_list, **_): return True @@ -5397,9 +5100,7 @@ async def test_async_log_success_event_no_mirror_when_no_snapshot(monkeypatch): hidden_params = mock_kwargs["standard_logging_object"].get("hidden_params") or {} additional_headers = hidden_params.get("additional_headers") or {} ratelimit_keys = [k for k in additional_headers if k.startswith("x-ratelimit-")] - assert ( - not ratelimit_keys - ), f"no snapshot must produce no rate-limit headers, got {ratelimit_keys}" + assert not ratelimit_keys, f"no snapshot must produce no rate-limit headers, got {ratelimit_keys}" @pytest.mark.asyncio @@ -5419,9 +5120,7 @@ async def test_streaming_mirror_matches_non_streaming_header_shape(monkeypatch): }, ) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) async def _noop_increment(increment_list, **_): return True @@ -5456,15 +5155,11 @@ async def test_streaming_mirror_matches_non_streaming_header_shape(monkeypatch): user_api_key_dict=user_api_key_dict, response=non_stream_response, ) - non_stream_headers = non_stream_response._hidden_params.get( - "additional_headers", {} - ) + non_stream_headers = non_stream_response._hidden_params.get("additional_headers", {}) # Streaming path: async_logging_hook mirrors into standard_logging_object. stream_kwargs: Dict[str, Any] = { - "standard_logging_object": { - "metadata": {"user_api_key_hash": _api_key} - }, + "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, "litellm_params": {"metadata": data["metadata"]}, "model": "gpt-4o-mini", } @@ -5481,18 +5176,13 @@ async def test_streaming_mirror_matches_non_streaming_header_shape(monkeypatch): result=stream_response, call_type="acompletion", ) - stream_slp_headers = ( - stream_kwargs["standard_logging_object"] - .get("hidden_params", {}) - .get("additional_headers", {}) - ) + stream_slp_headers = stream_kwargs["standard_logging_object"].get("hidden_params", {}).get("additional_headers", {}) def _rl_only(headers: Dict[str, Any]) -> Dict[str, Any]: return {k: v for k, v in headers.items() if k.startswith("x-ratelimit-")} assert _rl_only(stream_slp_headers) == _rl_only(non_stream_headers), ( - f"streaming={_rl_only(stream_slp_headers)}" - f" non_streaming={_rl_only(non_stream_headers)}" + f"streaming={_rl_only(stream_slp_headers)} non_streaming={_rl_only(non_stream_headers)}" ) assert "x-ratelimit-model_per_key-remaining-requests" in stream_slp_headers @@ -5508,12 +5198,8 @@ async def test_async_log_success_event_counts_passthrough_reported_tokens(monkey monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") _api_key = hash_token("sk-passthrough") - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) - monkeypatch.setattr( - parallel_request_handler, "get_rate_limit_type", lambda: "total" - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", lambda: "total") captured_operations = [] @@ -5549,9 +5235,7 @@ async def test_async_log_success_event_counts_passthrough_reported_tokens(monkey @pytest.mark.parametrize("rate_limit_type", ["input", "output", "total"]) @pytest.mark.asyncio -async def test_aggregate_only_usage_charges_tpm_under_every_limit_type( - monkeypatch, rate_limit_type -): +async def test_aggregate_only_usage_charges_tpm_under_every_limit_type(monkeypatch, rate_limit_type): """ A pass-through target reports one total for the whole request and cannot split it into prompt/completion. Reading a split out of it yields 0, which @@ -5561,12 +5245,8 @@ async def test_aggregate_only_usage_charges_tpm_under_every_limit_type( monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") _api_key = hash_token("sk-aggregate-only") - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) - monkeypatch.setattr( - parallel_request_handler, "get_rate_limit_type", lambda: rate_limit_type - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", lambda: rate_limit_type) captured_operations = [] @@ -5601,9 +5281,7 @@ async def test_split_usage_still_respects_the_configured_limit_type(monkeypatch) monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") _api_key = hash_token("sk-split-usage") - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", lambda: "output") captured_operations = [] @@ -5645,9 +5323,7 @@ async def test_atomic_check_with_zero_increment_still_enforces_token_limit(): """ from litellm.proxy.hooks.parallel_request_limiter_v3 import RateLimitDescriptor - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) descriptor = RateLimitDescriptor( key="model_saturation_check", value="tpm-only-model", @@ -5662,15 +5338,9 @@ async def test_atomic_check_with_zero_increment_still_enforces_token_limit(): assert under_limit["overall_code"] == "OK" assert [s["rate_limit_type"] for s in under_limit["statuses"]] == ["tokens"] - counter_key = handler.create_rate_limit_keys( - "model_saturation_check", "tpm-only-model", "tokens" - ) + counter_key = handler.create_rate_limit_keys("model_saturation_check", "tpm-only-model", "tokens") await handler.async_increment_tokens_with_ttl_preservation( - pipeline_operations=[ - RedisPipelineIncrementOperation( - key=counter_key, increment_value=100, ttl=60 - ) - ], + pipeline_operations=[RedisPipelineIncrementOperation(key=counter_key, increment_value=100, ttl=60)], ) at_limit = await handler.atomic_check_and_increment_by_n( @@ -5715,9 +5385,7 @@ async def test_reserve_tpm_tokens_never_evaluates_the_requests_dimension(): its descriptors or an exhausted RPM budget would double-enforce here.""" from litellm.proxy.hooks.parallel_request_limiter_v3 import RateLimitDescriptor - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) descriptor = RateLimitDescriptor( key="api_key", value="reserve-test-key", @@ -5729,8 +5397,7 @@ async def test_reserve_tpm_tokens_never_evaluates_the_requests_dimension(): estimated_tokens=10, ) assert response["overall_code"] == "OK", ( - "an exhausted requests budget (limit 0) must not block the token " - f"reservation pass, got: {response}" + f"an exhausted requests budget (limit 0) must not block the token reservation pass, got: {response}" ) assert [s["rate_limit_type"] for s in response["statuses"]] == ["tokens"] @@ -5981,9 +5648,7 @@ async def test_estimated_output_tokens_resolution_precedence( """ monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth( api_key=hash_token(f"sk-estimate-{expected_output_estimate}-{tier}"), tpm_limit=1_000_000, @@ -6006,9 +5671,7 @@ async def test_request_max_tokens_outranks_configured_estimate(monkeypatch): """An explicit request-level max_tokens stays the top of the precedence order.""" monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth( api_key=hash_token("sk-estimate-explicit-max-tokens"), tpm_limit=1_000_000, @@ -6030,9 +5693,7 @@ async def test_configured_estimate_does_not_apply_to_embeddings(monkeypatch): """Embeddings generate no output, so a declared output estimate must not be reserved.""" monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth( api_key=hash_token("sk-estimate-embeddings"), tpm_limit=1_000_000, @@ -6059,9 +5720,7 @@ async def test_configured_estimate_applies_to_contentless_requests(monkeypatch): """ monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) configured = UserAPIKeyAuth( api_key=hash_token("sk-estimate-contentless-configured"), tpm_limit=1_000_000, @@ -6073,17 +5732,9 @@ async def test_configured_estimate_applies_to_contentless_requests(monkeypatch): ) assert ( - await _reserved_tokens_for( - handler, local_cache, configured, {"model": "gpt-4o-mini", "messages": []} - ) - == 2002 - ) - assert ( - await _reserved_tokens_for( - handler, local_cache, unconfigured, {"model": "gpt-4o-mini", "messages": []} - ) - == 1 + await _reserved_tokens_for(handler, local_cache, configured, {"model": "gpt-4o-mini", "messages": []}) == 2002 ) + assert await _reserved_tokens_for(handler, local_cache, unconfigured, {"model": "gpt-4o-mini", "messages": []}) == 1 @pytest.mark.asyncio @@ -6100,9 +5751,7 @@ async def test_declared_estimate_never_tightens_the_small_tpm_clamp(monkeypatch) """ monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) raised_data: Dict[str, Any] = {"model": "gpt-4o-mini", "messages": ONE_TOKEN_PROMPT} raised_reserved = await _reserved_tokens_for( handler, @@ -6154,9 +5803,7 @@ async def test_one_malformed_estimate_field_does_not_discard_the_other(monkeypat """ monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) broken_map = await _reserved_tokens_for( handler, @@ -6205,9 +5852,7 @@ async def test_declared_estimate_over_the_tpm_budget_is_honored_and_explained(mo """ monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth( api_key=hash_token(f"sk-estimate-over-budget-{declared}"), tpm_limit=5000, @@ -6246,9 +5891,7 @@ async def test_a_key_that_declared_nothing_is_never_blamed_for_a_declaration(mon """ monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): with pytest.raises(HTTPException) as exc_info: @@ -6276,9 +5919,7 @@ async def test_declared_estimate_inside_the_tpm_budget_is_not_explained(monkeypa """ monkeypatch.delenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", raising=False) local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): await handler.async_pre_call_hook( @@ -6311,9 +5952,7 @@ async def test_configured_estimate_blocks_the_overrun_the_static_floor_admits(mo async def admitted(metadata): local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth( api_key=hash_token(f"sk-overrun-{metadata}"), tpm_limit=8000, @@ -6395,15 +6034,10 @@ def test_conflicting_token_limits_reserve_the_larger_declared_budget(): A request declaring max_tokens=1 alongside max_completion_tokens=10000 previously reserved one output token while the provider stayed free to emit ten thousand. """ - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) bodies = _conflicting_budget_bodies() - reserved = { - label: handler._estimate_tokens_for_request(data=body) - for label, body in bodies.items() - } + reserved = {label: handler._estimate_tokens_for_request(data=body) for label, body in bodies.items()} assert reserved["both"] == reserved["only_large"] assert reserved["both"] > reserved["only_small"] @@ -6417,9 +6051,7 @@ def test_non_integer_output_budgets_still_reserve_their_declared_size(declared): output floor, so dropping it from the estimate under-reserves and reopens the same TPM bypass that reading both spellings was meant to close. """ - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) base = {"model": "gpt-5-chat", "messages": [{"role": "user", "content": "hi"}]} reserved = handler._estimate_tokens_for_request(data={**base, "max_tokens": declared}) @@ -6432,12 +6064,8 @@ def test_non_integer_output_budgets_still_reserve_their_declared_size(declared): async def test_conflicting_token_limits_cannot_bypass_tpm_reservation(): """The pre-call hook must refuse a request whose larger declared budget exceeds the TPM limit.""" local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) - user_api_key_dict = UserAPIKeyAuth( - api_key=hash_token("sk-conflicting-budgets"), tpm_limit=100, models=[] - ) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + user_api_key_dict = UserAPIKeyAuth(api_key=hash_token("sk-conflicting-budgets"), tpm_limit=100, models=[]) bodies = _conflicting_budget_bodies() await handler.async_pre_call_hook( @@ -6464,7 +6092,9 @@ async def test_conflicting_token_limits_cannot_bypass_tpm_reservation(): def _enqueued_test_handler() -> _PROXY_MaxParallelRequestsHandler: - return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) + return _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60)) + ) def _batch_response(batch_id: str, status: str): @@ -6670,6 +6300,201 @@ async def test_post_call_success_hook_leaves_raw_provider_dict_untouched(): assert response == {"id": "msg_123", "type": "message", "role": "assistant", "content": []} +class _ClusterParallelTransport: + def __init__(self, handler): + self.handler = handler + self.members = {} + self.now = 10000 + self.calls = [] + self.fail = None + self.lose_acquire_response = None + self.pause_acquire = None + self.entered = asyncio.Event() + + def script(self, operation): + async def run(*, keys, args): + self.calls.append((operation, tuple(keys))) + if len({self.handler.keyslot_for_redis_cluster(key) for key in keys}) > 1: + raise RuntimeError("CROSSSLOT Keys in request do not hash to the same slot") + if self.fail is not None and (operation, keys[0]) == self.fail: + raise RuntimeError("Shard unavailable") + if operation in ("acquire", "count"): + for key in keys: + self.members[key] = { + slot: score + for slot, score in self.members.get(key, {}).items() + if score > self.now - PARALLEL_REQUEST_SLOT_TTL_SECONDS + } + if operation == "acquire": + for index, key in enumerate(keys): + if len(self.members[key]) >= args[index * 3]: + return [1, index + 1, len(self.members[key])] + for index, key in enumerate(keys): + self.members[key][args[index * 3 + 2]] = self.now + if keys[0] == self.pause_acquire: + self.entered.set() + await asyncio.Event().wait() + if keys[0] == self.lose_acquire_response: + raise RuntimeError("Reply lost after Redis admitted slot") + return [0, *(len(self.members[key]) for key in keys)] + if operation == "count": + return [len(self.members.get(key, {})) for key in keys] + if operation == "renew": + if any(self.members.get(key, {}).get(args[0], 0) <= self.now - args[1] for key in keys): + return [0] + for key in keys: + self.members[key][args[0]] = self.now + return [1] + assert operation == "release" + for index, key in enumerate(keys): + self.members.get(key, {}).pop(args[index], None) + return [len(self.members.get(key, {})) for key in keys] + + return run + + +def _cluster_parallel_fixture(monkeypatch): + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(DualCache())) + monkeypatch.setattr(handler, "_is_redis_cluster", lambda: True) + transport = _ClusterParallelTransport(handler) + for operation in ("acquire", "count", "renew", "release"): + monkeypatch.setattr(handler, f"parallel_{operation}_script", transport.script(operation)) + gauges = [ + {"counter_key": "{api_key:owner}:max_parallel_requests", "limit": 1, "descriptor_key": "api_key"}, + {"counter_key": "{team:group}:max_parallel_requests", "limit": 2, "descriptor_key": "team"}, + {"counter_key": "{api_key:owner}:another-parallel-scope", "limit": 1, "descriptor_key": "extra"}, + ] + return handler, transport, gauges + + +@pytest.mark.asyncio +async def test_cluster_parallel_slots_admit_count_renew_release_across_hash_slots(monkeypatch): + handler, transport, gauges = _cluster_parallel_fixture(monkeypatch) + keys = tuple(gauge["counter_key"] for gauge in gauges) + transport.members[keys[1]] = {"unrelated": transport.now} + result = await handler._check_parallel_request_gauges(gauges, "owner") + assert result["overall_code"] == "OK" + assert len([call for call in transport.calls if call[0] == "acquire"]) == 2 + assert all("owner" in transport.members[key] for key in keys) + result = await handler._check_parallel_request_gauges(gauges, "reader", read_only=True) + assert result["overall_code"] == "OVER_LIMIT" + assert [status["descriptor_key"] for status in result["statuses"]] == ["api_key", "team", "extra"] + transport.now += PARALLEL_REQUEST_SLOT_TTL_SECONDS - 1 + assert await handler._renew_realtime_call_slot("owner", keys) + transport.now += 2 + result = await handler._check_parallel_request_gauges(gauges, "second") + assert result["overall_code"] == "OVER_LIMIT" + acquisition = ParallelSlotAcquisition(slot_id="owner", counter_keys=list(keys)) + await handler._release_parallel_request_slots(acquisition) + await handler._release_parallel_request_slots(acquisition) + assert all("owner" not in transport.members[key] for key in keys) + assert not await handler._renew_realtime_call_slot("owner", keys) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["limit", "unreachable", "lost_reply", "cancel"]) +async def test_cluster_parallel_acquire_rolls_back_attempted_shards_without_releasing_others(monkeypatch, failure): + handler, transport, gauges = _cluster_parallel_fixture(monkeypatch) + keys = tuple(gauge["counter_key"] for gauge in gauges) + transport.members[keys[1]] = {"unrelated": transport.now} + if failure == "limit": + gauges[1]["limit"] = 1 + elif failure == "unreachable": + transport.fail = ("acquire", keys[1]) + elif failure == "lost_reply": + transport.lose_acquire_response = keys[1] + else: + transport.pause_acquire = keys[1] + task = asyncio.create_task(handler._check_parallel_request_gauges(gauges, "owner")) + if failure == "cancel": + await asyncio.wait_for(transport.entered.wait(), timeout=1) + task.cancel() + if failure == "limit": + assert (await task)["overall_code"] == "OVER_LIMIT" + else: + with pytest.raises(asyncio.CancelledError if failure == "cancel" else RuntimeError): + await task + assert all("owner" not in transport.members.get(key, {}) for key in keys) + assert transport.members[keys[1]] == {"unrelated": transport.now} + released_keys = {key for operation, group in transport.calls if operation == "release" for key in group} + assert released_keys == set(keys) + + +@pytest.mark.asyncio +async def test_cluster_parallel_release_continues_after_shard_failure_and_renewal_fails_closed(monkeypatch): + handler, transport, gauges = _cluster_parallel_fixture(monkeypatch) + keys = tuple(gauge["counter_key"] for gauge in gauges) + assert (await handler._check_parallel_request_gauges(gauges, "owner"))["overall_code"] == "OK" + transport.fail = ("renew", keys[1]) + assert not await handler._renew_realtime_call_slot("owner", keys) + transport.fail = None + transport.members[keys[1]].pop("owner") + assert not await handler._renew_realtime_call_slot("owner", keys) + assert "owner" not in transport.members[keys[1]] + transport.fail = ("count", keys[1]) + with pytest.raises(RuntimeError, match="Shard unavailable"): + await handler._check_parallel_request_gauges(gauges, "reader", read_only=True) + transport.fail = ("release", keys[0]) + receipt = ParallelSlotAcquisition(slot_id="owner", counter_keys=list(keys)) + with pytest.raises(RuntimeError, match="Shard unavailable"): + await handler._release_parallel_request_slots(receipt) + assert "owner" not in transport.members[keys[1]] + transport.fail = None + await handler._release_parallel_request_slots(receipt) + assert all("owner" not in transport.members.get(key, {}) for key in keys) + + +@pytest.mark.asyncio +async def test_cluster_parallel_duplicate_scope_keeps_strictest_limit(monkeypatch): + handler, transport, gauges = _cluster_parallel_fixture(monkeypatch) + gauges.append({**gauges[0], "limit": 100}) + transport.members[gauges[0]["counter_key"]] = {"unrelated": transport.now} + assert (await handler._check_parallel_request_gauges(gauges, "owner"))["overall_code"] == "OVER_LIMIT" + assert transport.members[gauges[0]["counter_key"]] == {"unrelated": transport.now} + + +@pytest.mark.asyncio +async def test_cluster_rollback_waits_through_repeated_cancellation(monkeypatch): + handler, transport, gauges = _cluster_parallel_fixture(monkeypatch) + keys = tuple(gauge["counter_key"] for gauge in gauges) + transport.lose_acquire_response = keys[1] + release_entered, finish_release = asyncio.Event(), asyncio.Event() + release = handler.parallel_release_script + + async def blocked_release(*, keys, args): + release_entered.set() + await finish_release.wait() + return await release(keys=keys, args=args) + + monkeypatch.setattr(handler, "parallel_release_script", blocked_release) + task = asyncio.create_task(handler._check_parallel_request_gauges(gauges, "owner")) + await asyncio.wait_for(release_entered.wait(), timeout=1) + try: + for _ in range(3): + task.cancel() + await asyncio.sleep(0) + assert not task.done(), "admission returned while its Redis compensation was still running" + finally: + finish_release.set() + await asyncio.gather(task, return_exceptions=True) + await asyncio.sleep(0) + assert task.cancelled() + assert all("owner" not in transport.members.get(key, {}) for key in keys) + + +@pytest.mark.asyncio +async def test_standalone_realtime_renewal_keeps_single_atomic_batch(monkeypatch): + from unittest.mock import AsyncMock + + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(DualCache())) + monkeypatch.setattr(handler, "_is_redis_cluster", lambda: False) + renew = AsyncMock(return_value=[1]) + monkeypatch.setattr(handler, "parallel_renew_script", renew) + keys = ("{api_key:owner}:max_parallel_requests", "{team:group}:max_parallel_requests") + assert await handler._renew_realtime_call_slot("owner", keys) + renew.assert_awaited_once_with(keys=keys, args=("owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS)) + + class _OpenBreakerRedis: def async_register_script(self, script: str): async def refused(keys, args): @@ -6718,6 +6543,52 @@ async def test_an_open_circuit_breaker_reads_the_sliding_window_locally_without_ assert any("circuit breaker is open" in record.getMessage() for record in caplog.records) +@pytest.mark.asyncio +async def test_cluster_gauge_guards_fail_fast_when_scripts_are_unavailable(monkeypatch): + # _check_parallel_request_gauges only enters the cluster path with an acquire script in hand, so the + # guards below run only for direct cluster callers or script resets racing an in-flight batch. + handler, transport, gauges = _cluster_parallel_fixture(monkeypatch) + monkeypatch.setattr(handler, "parallel_count_script", None) + with pytest.raises(RuntimeError, match="Redis cluster parallel count script is unavailable"): + await handler._check_cluster_parallel_gauges(gauges, "reader", None, read_only=True) + monkeypatch.setattr(handler, "parallel_count_script", transport.script("count")) + monkeypatch.setattr(handler, "parallel_acquire_script", None) + with pytest.raises(RuntimeError, match="Redis cluster parallel acquire script is unavailable"): + await handler._check_cluster_parallel_gauges(gauges, "owner", None, read_only=False) + assert transport.calls == [] + + +@pytest.mark.asyncio +async def test_cluster_release_guards_when_release_script_unavailable(monkeypatch): + handler, transport, gauges = _cluster_parallel_fixture(monkeypatch) + keys = tuple(gauge["counter_key"] for gauge in gauges) + monkeypatch.setattr(handler, "parallel_release_script", None) + with pytest.raises(RuntimeError, match="Redis cluster parallel release script is unavailable"): + await handler._release_cluster_parallel_slots(keys, "owner", None) + # Every shard group is still attempted; the first shard's error is the one re-raised. + assert [operation for operation, _ in transport.calls] == [] + + +@pytest.mark.asyncio +async def test_cluster_rollback_swallows_release_failure_and_logs(monkeypatch, caplog): + handler, transport, gauges = _cluster_parallel_fixture(monkeypatch) + keys = tuple(gauge["counter_key"] for gauge in gauges) + assert (await handler._check_parallel_request_gauges(gauges, "owner"))["overall_code"] == "OK" + transport.fail = ("release", keys[0]) + await handler._rollback_cluster_parallel_slots(keys, "owner", None) + # The admission error must not be replaced by an unreachable compensation shard: the rollback + # exception is retrieved, reported once, and swallowed so the caller keeps its original failure. + assert "Could not roll back all Redis cluster parallel request slots" in caplog.text + released = {key for operation, group in transport.calls if operation == "release" for key in group} + assert released == set(keys) + transport.fail = None + assert "owner" not in transport.members[keys[1]] + + +# Note: _renew_realtime_call_slot's in-memory "return False" isinstance guard after the any() scan is +# unreachable for any real cache state (the scan rejects non-dict values first), so no test drives it. + + class _UnreachableRedis: def async_register_script(self, script: str): async def refused(keys, args): diff --git a/tests/test_litellm/proxy/hooks/test_realtime_call_lease.py b/tests/test_litellm/proxy/hooks/test_realtime_call_lease.py new file mode 100644 index 00000000000..cb8ee7fc6b9 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_realtime_call_lease.py @@ -0,0 +1,85 @@ +import asyncio +from unittest.mock import AsyncMock + +import pytest + +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease + + +@pytest.mark.asyncio +async def test_failed_renewal_signals_owner_and_close_releases_once(): + renew = AsyncMock(side_effect=[True, False, True]) + release = AsyncMock() + lease = RealtimeCallLease(renew=renew, release=release, interval=0.001) + lease.start() + await asyncio.wait_for(lease.wait_failed(), timeout=1) + assert renew.await_count == 2 + assert not await lease.renew() + assert renew.await_count == 2 + await asyncio.gather(lease.close(), lease.close()) + assert release.await_count == 1 + + +@pytest.mark.asyncio +async def test_renewal_exception_and_close_before_start(): + release = AsyncMock() + lease = RealtimeCallLease(renew=AsyncMock(side_effect=RuntimeError("backend")), release=release, interval=0.001) + lease.start() + await asyncio.wait_for(lease.wait_failed(), timeout=1) + await lease.close() + assert release.await_count == 1 + unused = RealtimeCallLease(renew=AsyncMock(), release=release) + await unused.close() + assert release.await_count == 2 + + +@pytest.mark.asyncio +async def test_renewal_timeout_signals_failure_without_start(): + lease = RealtimeCallLease(renew=asyncio.Event().wait, release=AsyncMock(), renewal_timeout=0.001) + assert not await lease.renew() + await asyncio.wait_for(lease.wait_failed(), timeout=1) + await lease.close() + + +@pytest.mark.asyncio +async def test_cancelled_close_still_releases_exactly_once(): + entered = asyncio.Event() + finish = asyncio.Event() + + async def release(): + entered.set() + await finish.wait() + + cleanup = AsyncMock(side_effect=release) + lease = RealtimeCallLease(renew=AsyncMock(return_value=True), release=cleanup) + lease.start() + closing = asyncio.create_task(lease.close()) + await asyncio.wait_for(entered.wait(), timeout=1) + closing.cancel() + with pytest.raises(asyncio.CancelledError): + await closing + finish.set() + await lease.close() + assert cleanup.await_count == 1 + + +@pytest.mark.asyncio +async def test_concurrent_renewal_cannot_restore_a_failed_lease(): + pending = asyncio.Event() + entered = asyncio.Event() + + async def delayed_success(): + entered.set() + await pending.wait() + return True + + renew = AsyncMock(side_effect=delayed_success) + lease = RealtimeCallLease(renew=renew, release=AsyncMock()) + first = asyncio.create_task(lease.renew()) + await asyncio.wait_for(entered.wait(), timeout=1) + renew.side_effect = None + renew.return_value = False + assert not await lease.renew() + pending.set() + assert not await first + await lease.close() diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py index 59c2921e0d0..cbd286ee4df 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py @@ -784,20 +784,22 @@ def _proxy_with_stubbed_reload(prisma): def _eviction_journal(access_group): - """Both auth cache keys, in the order a write path has to evict them.""" + """Every auth cache key that holds this group's limits, in the order a write path has to evict them.""" from litellm.proxy.common_utils.user_api_key_cache import ( + live_model_access_group_limits_cache_key, model_access_group_cache_key, model_access_group_registry_cache_key, ) return [ f"auth_cache.delete:{model_access_group_cache_key(access_group)}", + f"auth_cache.delete:{live_model_access_group_limits_cache_key(access_group)}", f"auth_cache.delete:{model_access_group_registry_cache_key()}", ] def _assert_evicted_after_write(journal, access_group, write_entry): - """Exactly the two keys, in order, after the DB write. Deliberately not a tail slice: what + """Exactly the cached keys, in order, after the DB write. Deliberately not a tail slice: what has to hold is that the eviction follows the write, not that nothing follows the eviction.""" evictions = [entry for entry in journal if entry.startswith("auth_cache.delete:")] assert evictions == _eviction_journal(access_group) @@ -1206,7 +1208,7 @@ async def test_list_access_groups_reports_a_budgetless_group_as_unbudgeted_rathe @pytest.mark.asyncio -async def test_put_access_group_budget_evicts_both_auth_cache_keys(): +async def test_put_access_group_budget_evicts_every_cached_limit_key(): """Auth reads the per-group row and the registry of budgeted groups cache-first with no freshness check, so a PUT that skips either eviction returns 200 and enforces nothing until the TTL expires. Both keys, after the write.""" @@ -1233,7 +1235,7 @@ async def test_put_access_group_budget_evicts_both_auth_cache_keys(): @pytest.mark.asyncio -async def test_delete_access_group_budget_evicts_both_auth_cache_keys(): +async def test_delete_access_group_budget_evicts_every_cached_limit_key(): """Clearing a budget has the same window as setting one: until both keys are dropped, auth keeps enforcing the budget that is already gone.""" from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( @@ -1252,7 +1254,7 @@ async def test_delete_access_group_budget_evicts_both_auth_cache_keys(): @pytest.mark.asyncio -async def test_deleting_the_access_group_evicts_both_auth_cache_keys(): +async def test_deleting_the_access_group_evicts_every_cached_limit_key(): """The group-delete cascade drops the budget row too, so it owes the same two evictions.""" from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( delete_access_group, diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py new file mode 100644 index 00000000000..ca993f778f4 --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -0,0 +1,1655 @@ +import hashlib +import time +from types import SimpleNamespace + +import pytest +from fastapi import HTTPException, WebSocket + +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.realtime_endpoints import call_sessions as codex +from litellm.llms.chatgpt.codex import CodexRealtimeCall +from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_call + + +@pytest.mark.asyncio +@pytest.mark.parametrize("multipart", [False, True]) +@pytest.mark.parametrize("content_length", [None, "1", "999999999"]) +async def test_oversized_offer_stops_before_auth_or_multipart_files(monkeypatch, multipart, content_length): + import json + from unittest.mock import AsyncMock, Mock + + from fastapi import Request + + monkeypatch.setattr(codex, "MAX_REALTIME_OFFER_BYTES", 1024) + if multipart: + body = ( + b'--Boundary\r\nContent-Disposition: form-data; name="extra"; filename="large.bin"\r\n\r\n' + + b"x" * 2048 + + b"\r\n--Boundary--\r\n" + ) + media_type = b"Multipart/Form-Data; boundary=Boundary" + else: + body = json.dumps({"sdp": "x" * 2048, "session": {"model": "voice"}}).encode() + media_type = b"application/json" + chunks = [body[offset : offset + 256] for offset in range(0, len(body), 256)] + received = [] + + async def receive(): + chunk = chunks.pop(0) + received.append(len(chunk)) + return {"type": "http.request", "body": chunk, "more_body": bool(chunks)} + + headers = [(b"content-type", media_type)] + if content_length is not None: + headers.append((b"content-length", content_length.encode())) + request = Request({"type": "http", "headers": headers}, receive) + if multipart: + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + + assert await _read_request_body(request) == {} + authenticate = AsyncMock() + create_file = Mock(side_effect=AssertionError("Oversized offers must not create temporary files")) + monkeypatch.setattr(codex, "user_api_key_auth", authenticate) + monkeypatch.setattr("starlette.formparsers.SpooledTemporaryFile", create_file) + with pytest.raises(HTTPException) as rejected: + await codex.create_codex_realtime_call(request) + assert rejected.value.status_code == 413 + assert sum(received) <= 1280 + assert chunks + authenticate.assert_not_awaited() + create_file.assert_not_called() + + +@pytest.mark.asyncio +async def test_offer_at_size_limit_keeps_body_available_for_custom_auth(monkeypatch): + import json + + from fastapi import Request + + monkeypatch.setattr(codex, "MAX_REALTIME_OFFER_BYTES", 1024) + empty = {"sdp": "", "session": {"model": "voice"}} + sdp = "x" * (1024 - len(json.dumps(empty).encode())) + body = json.dumps({"sdp": sdp, "session": {"model": "voice"}}).encode() + chunks = [body[:512], body[512:]] + + async def receive(): + return {"type": "http.request", "body": chunks.pop(0), "more_body": bool(chunks)} + + request = Request({"type": "http", "headers": [(b"content-type", b"application/json")]}, receive) + offer = await codex.read_codex_offer(request) + assert offer.sdp == sdp + assert offer.session.model == "voice" + assert await request.body() == body + assert not chunks + + +@pytest.mark.asyncio +async def test_empty_pre_read_multipart_offer_returns_invalid_offer(): + from fastapi import Request + + async def receive(): + return {"type": "http.request", "body": b"--Boundary--\r\n", "more_body": False} + + request = Request( + {"type": "http", "headers": [(b"content-type", b"multipart/form-data; boundary=Boundary")]}, receive + ) + assert not await request.form() + with pytest.raises(HTTPException) as rejected: + await codex.create_codex_realtime_call(request) + assert rejected.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_oversized_pre_read_offer_is_rejected_before_decoding(monkeypatch): + from fastapi import Request + + monkeypatch.setattr(codex, "MAX_REALTIME_OFFER_BYTES", 1024) + + async def receive(): + return {"type": "http.request", "body": b"x" * 2048, "more_body": False} + + request = Request({"type": "http", "headers": [(b"content-type", b"application/json")]}, receive) + await request.body() + with pytest.raises(HTTPException) as rejected: + await codex.read_codex_offer(request) + assert rejected.value.status_code == 413 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("malformed", [False, True]) +async def test_mixed_case_offer_preserves_boundary_metadata_and_closes_extra_files(monkeypatch, malformed): + from fastapi import Request + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + + boundary = "AbCdEf123" + fields = { + "sdp": "v=0", + "session": "invalid" if malformed else '{"model":"voice"}', + "metadata": '{"policy":"keep"}', + "extra_policy": "keep", + } + body = ( + "".join( + f'--{boundary}\r\nContent-Disposition: form-data; name="{name}"\r\n\r\n{value}\r\n' + for name, value in fields.items() + ) + + f'--{boundary}\r\nContent-Disposition: form-data; name="extra_file"; filename="test.txt"\r\n\r\nextra\r\n--{boundary}--\r\n' + ).encode() + + async def receive(): + return {"type": "http.request", "body": body, "more_body": False} + + request = Request( + {"type": "http", "headers": [(b"content-type", f'Multipart/Form-Data; boundary="{boundary}"'.encode())]}, + receive, + ) + assert not await request.form() + if malformed: + with pytest.raises(HTTPException) as error: + await codex.create_codex_realtime_call(request) + assert error.value.status_code == 400 + assert (await request.form())["extra_file"].file.closed + return + first = await codex.read_codex_offer(request) + second = await codex.read_codex_offer(request) + assert first == second + parsed = await _read_request_body(request) + assert parsed["metadata"] == {"policy": "keep"} + assert parsed["extra_policy"] == "keep" + assert not parsed["extra_file"].file.closed + + async def deny_auth(**kwargs): + assert kwargs["request"] is request + auth_form = await request.form() + assert auth_form["extra_policy"] == "keep" + assert await auth_form["extra_file"].read() == b"extra" + assert not auth_form["extra_file"].file.closed + raise HTTPException(403, "policy denied") + + monkeypatch.setattr(codex, "user_api_key_auth", deny_auth) + with pytest.raises(HTTPException, match="policy denied"): + await codex.create_codex_realtime_call(request) + assert parsed["extra_file"].file.closed + assert request.headers["content-type"] == f'Multipart/Form-Data; boundary="{boundary}"' + + +@pytest.mark.asyncio +@pytest.mark.parametrize("multipart", [False, True]) +@pytest.mark.parametrize("policy", ["budget", "personal_models"]) +@pytest.mark.parametrize("mixed_case", [False, True]) +@pytest.mark.parametrize("pre_read", [False, True]) +async def test_offer_auth_enforces_session_model_policy_before_upstream( + monkeypatch, multipart, policy, mixed_case, pre_read +): + import json + from unittest.mock import AsyncMock, MagicMock + + import httpx + from fastapi import Request, Response + + import litellm + from litellm.exceptions import BudgetExceededError + from litellm.proxy import proxy_server as server + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.auth.auth_checks import common_checks + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.realtime_endpoints.endpoints import proxy_realtime_calls + + session = {"model": "forbidden-voice"} + payload = ( + {"files": {"sdp": (None, "v=0"), "session": (None, json.dumps(session)), "model": (None, "body-decoy")}} + if multipart + else {"json": {"sdp": "v=0", "session": session, "model": "body-decoy"}} + ) + outbound = httpx.Request("POST", "http://localhost/v1/realtime/calls", **payload) + body = outbound.read() + content_type = outbound.headers["content-type"] + if mixed_case: + content_type = content_type.replace("multipart/form-data", "Multipart/Form-Data").replace( + "application/json", "Application/JSON" + ) + receives = [] + + async def receive(): + receives.append(True) + assert len(receives) == 1 + return {"type": "http.request", "body": body, "more_body": False} + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "query_string": b"model=query-decoy&policy=keep", + "client": ("127.0.0.7", 1234), + "headers": [ + (b"content-type", content_type.encode()), + (b"x-policy-key", b"Bearer test-key"), + (b"x-custom-policy", b"preserved"), + (b"x-litellm-model", b"header-decoy"), + ], + }, + receive, + ) + token = UserAPIKeyAuth(token="test-key", user_id="personal-user", model_max_budget={"forbidden-voice": 0}) + budget = AsyncMock(side_effect=BudgetExceededError(current_cost=1, max_budget=0)) + upstream = AsyncMock() + original_request = request + + async def custom_auth(request: Request, api_key: str): + assert request is original_request + assert request.headers["content-type"] == content_type + assert api_key == "test-key" + assert request.headers["x-custom-policy"] == "preserved" + assert request.query_params["policy"] == "keep" + assert request.client.host == "127.0.0.7" + if multipart: + assert (await request.form())["model"] == "body-decoy" + parsed = await _read_request_body(request) + assert parsed["model"] == "body-decoy" + assert isinstance(parsed["session"], str) is multipart + if policy == "personal_models": + await common_checks( + request_body=parsed, + team_object=None, + user_object=LiteLLM_UserTable( + user_id="personal-user", models=["allowed-voice", "body-decoy", "query-decoy", "header-decoy"] + ), + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route="/v1/realtime/calls", + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + skip_budget_checks=True, + ) + return token + + custom = AsyncMock(side_effect=custom_auth) + monkeypatch.setattr(server, "general_settings", {"litellm_key_header_name": "x-policy-key"}) + monkeypatch.setattr(server, "user_custom_auth", custom) + monkeypatch.setattr(server, "llm_router", None) + monkeypatch.setattr(server, "llm_model_list", []) + monkeypatch.setattr(server, "model_max_budget_limiter", SimpleNamespace(is_key_within_model_budget=budget)) + monkeypatch.setattr(server, "route_request", upstream) + monkeypatch.setattr(litellm, "enable_post_custom_auth_checks", True, raising=False) + if pre_read: + await _read_request_body(request) + with pytest.raises(ProxyException) as denied: + await proxy_realtime_calls(request, Response()) + if policy == "personal_models": + internal_message = getattr(denied.value, "internal_message", str(denied.value)) + assert "user not allowed to access model" in internal_message + assert "forbidden-voice" in internal_message + custom.assert_awaited_once() + upstream.assert_not_awaited() + if policy == "budget": + budget.assert_awaited_once() + assert budget.await_args.kwargs["model"] == "forbidden-voice" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("route_type", ["arealtime_calls", "_arealtime"]) +@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 + from litellm.proxy._types import UserAPIKeyAuth + 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, *, skip_guardrails=False, 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 + + router = Router(model_list=[{ + "model_name": "voice-policy", + "litellm_params": {"model": "openai/gpt-realtime-1.5", "api_key": "test", "guardrails": ["model-policy"]}, + }]) + monkeypatch.setattr(server, "llm_router", router) + 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, + internal_realtime_observer=observer, + ) + assert error.value.status_code == 403 + assert error.value.detail == "Model policy rejected request" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("logged_success", [False, True]) +@pytest.mark.parametrize("disconnect_error", [False, True]) +async def test_sideband_preserves_pending_cost_reconciliation(monkeypatch, logged_success, disconnect_error): + import litellm + from unittest.mock import AsyncMock + from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall(call_id="rtc_test", model="gpt-live-1-codex", alias="voice", + owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time()+300) + auth = UserAPIKeyAuth() + auth.budget_reservation = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []} + logger = SimpleNamespace(model_call_details={}) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(codex, "process_codex_request", AsyncMock(return_value=({}, logger))) + + async def forward(**kwargs): + if logged_success: + logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True + if disconnect_error: + raise RuntimeError("Backend disconnected") + + monkeypatch.setattr(litellm, "_arealtime", forward) + websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque", "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")]}, + AsyncMock(return_value={"type": "websocket.connect"}), AsyncMock()) + if disconnect_error: + with pytest.raises(RuntimeError, match="Backend disconnected"): + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + else: + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + assert auth.budget_reservation["finalized"] is not logged_success + + +@pytest.mark.asyncio +@pytest.mark.parametrize("ending", ["normal", "disconnect", "pre_call", "admission"]) +async def test_supervised_attachments_release_real_limiter_before_reconnect(monkeypatch, ending): + import asyncio + from unittest.mock import AsyncMock + + import litellm + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server as server + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + _request_stash, + get_request_stash, + ) + from litellm.proxy.utils import InternalUsageCache + + cache = DualCache() + limiter = _PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key="attachment-owner", max_parallel_requests=1, tpm_limit=10000) + token_key = limiter.create_rate_limit_keys(key="api_key", value=auth.api_key, rate_limit_type="tokens") + parallel_key = f"{{api_key:{auth.api_key}}}:max_parallel_requests" + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-live-1-codex", + alias="voice", + usage_supervised=True, + owner=hashlib.sha256(b"Bearer owner").hexdigest(), + expires_at=time.time() + 300, + ) + monkeypatch.setenv("LITELLM_SALT_KEY", "attachment-cleanup-test") + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda name: limiter)) + + async def process(request, data, selected_auth, model, call_type): + await limiter.async_pre_call_hook( + user_api_key_dict=selected_auth, + cache=cache, + data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}], "max_tokens": 50}, + call_type="completion", + ) + assert get_request_stash().reserved_tokens > 0 + if ending == "pre_call": + raise RuntimeError("Later policy rejected attachment") + return data, SimpleNamespace(model_call_details={}) + + async def forward(**kwargs): + if ending == "disconnect": + raise asyncio.CancelledError() + + monkeypatch.setattr(codex, "process_codex_request", process) + monkeypatch.setattr(litellm, "_arealtime", forward) + blocker_stash = None + blocker_reserved = 0 + if ending == "admission": + setup_token = _request_stash.set(None) + try: + await limiter.async_pre_call_hook( + user_api_key_dict=auth, + cache=cache, + data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}], "max_tokens": 50}, + call_type="completion", + ) + blocker_stash = get_request_stash() + blocker_reserved = blocker_stash.reserved_tokens + finally: + _request_stash.reset(setup_token) + for _ in range(3): + stash_token = _request_stash.set(None) + try: + websocket = WebSocket( + { + "type": "websocket", + "path": "/v1/live/opaque", + "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")], + }, + AsyncMock(return_value={"type": "websocket.connect"}), + AsyncMock(), + ) + if ending == "disconnect": + with pytest.raises(asyncio.CancelledError): + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + else: + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + assert limiter._gauge_in_flight_from_cache_value(await cache.async_get_cache(parallel_key)) == int( + ending == "admission" + ) + assert int(await cache.async_get_cache(token_key) or 0) == blocker_reserved + finally: + _request_stash.reset(stash_token) + if blocker_stash is not None: + cleanup_token = _request_stash.set(blocker_stash) + try: + await limiter.async_release_realtime_attachment({}, auth) + finally: + _request_stash.reset(cleanup_token) + + +def test_sideband_token_binds_owner_and_model(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-live-1-codex", + alias="gpt-live-1-codex", + extra_headers={"x-gateway-secret": "configured-secret"}, + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() + 300, + ) + token = encode_call(call) + assert "/" not in token + assert "configured-secret" not in token + assert decode_call(token, "Bearer test-owner") == call + with pytest.raises(HTTPException) as error: + decode_call(token, "Bearer different-owner") + assert error.value.status_code == 403 + with pytest.raises(HTTPException): + decode_call(token[:30] + "tampered" + token[30:], "Bearer test-owner") + + +def test_sideband_rejects_expired_token(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-realtime-1.5", + alias="gpt-realtime-1.5", + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() - 1, + ) + with pytest.raises(HTTPException): + decode_call(encode_call(call), "Bearer test-owner") + + +@pytest.mark.parametrize("token", ["", "rtc_other", "rtc_litellm_%%%%", "rtc_litellm_a"]) +def test_sideband_rejects_malformed_tokens(token): + with pytest.raises(HTTPException): + decode_call(token, "Bearer test-owner") + + +@pytest.mark.asyncio +async def test_sideband_rejects_revoked_model_access(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-live-1-codex", + alias="voice", + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() + 300, + ) + sent = [] + + async def receive(): + return {"type": "websocket.connect"} + + async def send(message): + sent.append(message) + + async def deny_model(**kwargs): + raise ProxyException("Model access revoked", "auth_error", "model", 403) + + monkeypatch.setattr(codex, "can_key_call_resolved_model", deny_model) + websocket = WebSocket( + {"type": "websocket", "headers": [(b"authorization", b"Bearer test-owner")]}, receive, send + ) + await codex.codex_realtime_sideband(websocket, encode_call(call), UserAPIKeyAuth()) + assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Invalid realtime call"}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("call_id", ["rtc_raw", "", "rtc_litellm_invalid"]) +async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id): + from unittest.mock import AsyncMock + from fastapi import WebSocket + from litellm.proxy import proxy_server as server + from litellm.proxy._types import UserAPIKeyAuth + + sent = [] + + async def receive(): + return {"type": "websocket.connect"} + + async def send(message): + sent.append(message) + + route = AsyncMock() + monkeypatch.setattr(server, "route_request", route) + websocket = WebSocket({"type": "websocket", "headers": [], "query_string": b""}, receive, send) + await server.realtime_websocket_endpoint( + websocket, model="gpt-realtime-1.5", call_id=call_id, + intent=None, guardrails=None, user_api_key_dict=UserAPIKeyAuth() + ) + assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Invalid realtime call"}] + route.assert_not_called() + + +@pytest.mark.asyncio +@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"]) +@pytest.mark.parametrize("signaling_path", ["/v1/realtime/calls", "/v1/live", "/live", "/openai/v1/live"]) +async def test_offer_exchange_wraps_call_and_filters_client_headers( + monkeypatch, multipart, credential, signaling_credential, signaling_path +): + import json + from unittest.mock import AsyncMock + + import httpx + from fastapi import Request, WebSocket + from litellm.proxy import common_request_processing, proxy_server + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.realtime_endpoints import call_sessions as codex + import litellm + + 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))}, + ) + else: + 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(): + return {"type": "http.request", "body": body, "more_body": False} + + signaling_headers = ( + [(b"authorization", b"Bearer other-owner"), (b"x-litellm-api-key", b"owner")] + 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": signaling_path, + "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(codex, "can_key_call_resolved_model", authorize) + + class Processor: + def __init__(self, data): + self.data = data + + async def common_processing_pre_call_logic(self, **kwargs): + assert isinstance(kwargs["user_api_key_dict"], UserAPIKeyAuth) + if kwargs["route_type"] == "_arealtime": + 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", + }, + "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) + + async def route(**kwargs): + data = kwargs["data"] + assert data["sdp_body"] == b"v=0\r\n" + assert data["session"] == session + assert data["chatgpt_realtime_client_headers"] == {"openai-alpha": "quicksilver=v2"} + assert "extra_headers" not in data + 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"}, + "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 + expected_prefix = "/v1/realtime/calls/" if signaling_path == "/v1/realtime/calls" else "/v1/live/" + assert response.headers["location"].startswith(expected_prefix) + assert response.body == b"v=0\r\nanswer" + token = response.headers["location"].rsplit("/", 1)[-1] + call = codex.decode_call(token, "Bearer owner") + 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() + + sent = [] + + async def send(message): + sent.append(message) + + async def receive_ws(): + return {"type": "websocket.connect"} + + credential_headers = { + "authorization": [(b"authorization", b"Bearer owner")], + "api-key": [(b"api-key", b"owner")], + "x-litellm-api-key": [(b"x-litellm-api-key", b"owner")], + "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, + ) + 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["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" + assert authorize.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body", [b"not json", b'{}', b'{"sdp":"v=0","session":{}}']) +async def test_invalid_offers_fail_before_authentication(monkeypatch, body): + from unittest.mock import AsyncMock + from fastapi import Request + from litellm.proxy.realtime_endpoints import call_sessions as codex + + async def receive(): + return {"type": "http.request", "body": body} + + request = Request({"type": "http", "headers": [(b"content-type", b"application/json")]}, receive) + authenticate = AsyncMock() + monkeypatch.setattr(codex, "user_api_key_auth", authenticate) + with pytest.raises(HTTPException) as error: + await codex.create_codex_realtime_call(request) + assert error.value.status_code == 400 + authenticate.assert_not_called() + + +@pytest.mark.asyncio +async def test_sideband_pre_call_block_prevents_upstream_connection(monkeypatch): + from unittest.mock import AsyncMock + from fastapi import WebSocket + import litellm + from litellm.proxy import common_request_processing + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.realtime_endpoints import call_sessions as codex + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall(call_id="rtc_test", model="gpt-live-1-codex", alias="voice", + owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time()+300) + token = encode_call(call) + sent = [] + + async def receive(): + return {"type": "websocket.connect"} + + async def send(message): + sent.append(message) + + class BlockingProcessor: + def __init__(self, data): + assert data["model"] == "voice" + + async def common_processing_pre_call_logic(self, **kwargs): + assert kwargs["route_type"] == "_arealtime" + raise HTTPException(403, "Policy blocked this call") + + forward = AsyncMock() + monkeypatch.setattr(litellm, "_arealtime", forward) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", BlockingProcessor) + websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque", "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")]}, receive, send) + 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]) +@pytest.mark.parametrize("invalidate_fails", [False, True]) +@pytest.mark.parametrize("close_fails", [False, True]) +async def test_supervisor_constructor_failure_closes_effective_connection( + monkeypatch, hangup_fails, invalidate_fails, close_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() + if close_fails: + connection.close = AsyncMock(side_effect=RuntimeError("socket cleanup secret")) + handlers = [] + invalidate = AsyncMock() + if invalidate_fails: + invalidate.side_effect = RuntimeError("counter cleanup secret") + 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 + if hangup_fails and invalidate_fails: + assert "Realtime startup cleanup could not invalidate budget counters" in caplog.text + elif invalidate_fails: + assert "Realtime startup cleanup could not invalidate budget counters" not in caplog.text + if close_fails: + assert "Realtime startup cleanup could not close observer socket" in caplog.text + assert "socket cleanup secret" not in caplog.text + assert "counter cleanup secret" not in caplog.text + + +@pytest.mark.asyncio +async def test_attachment_releases_quota_before_upstream_close_handshake(monkeypatch): + import asyncio + from unittest.mock import AsyncMock + + import websockets + + import litellm + from litellm.proxy import proxy_server as server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.hooks.parallel_request_limiter_v3 import _request_stash + from litellm.proxy.utils import ProxyLogging + + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", []) + proxy._add_proxy_hooks() + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr( + server, + "llm_router", + litellm.Router( + model_list=[ + {"model_name": "voice", "litellm_params": {"model": "openai/gpt-realtime-1.5", "api_key": "test"}} + ] + ), + ) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setenv("LITELLM_SALT_KEY", "attachment-close-order-test") + from litellm.llms.chatgpt.authenticator import Authenticator + + monkeypatch.setattr(Authenticator, "get_access_token", lambda self: "test-token") + monkeypatch.setattr(Authenticator, "get_account_id", lambda self: "test-account") + closing = asyncio.Event() + finish_close = asyncio.Event() + + class Backend: + async def recv(self, **kwargs): + await asyncio.Event().wait() + + async def send(self, value): + return None + + class Connection: + async def __aenter__(self): + return Backend() + + async def __aexit__(self, *args): + closing.set() + await finish_close.wait() + + monkeypatch.setattr(websockets, "connect", lambda *args, **kwargs: Connection()) + auth = UserAPIKeyAuth(api_key="close-order-owner", max_parallel_requests=1) + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-live-1-codex", + alias="voice", + usage_supervised=True, + owner=hashlib.sha256(b"Bearer owner").hexdigest(), + expires_at=time.time() + 300, + ) + ws = WebSocket( + { + "type": "websocket", + "path": "/v1/live/test", + "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")], + "scheme": "ws", + "server": ("localhost", 80), + }, + AsyncMock(side_effect=[{"type": "websocket.connect"}, {"type": "websocket.disconnect", "code": 1000}]), + AsyncMock(), + ) + token = _request_stash.set(None) + request = asyncio.create_task(codex.codex_realtime_sideband(ws, encode_call(call), auth)) + try: + await asyncio.wait_for(closing.wait(), timeout=5) + limiter = proxy.get_proxy_hook("parallel_request_limiter") + value = await proxy.internal_usage_cache.async_get_cache( + "{api_key:close-order-owner}:max_parallel_requests", litellm_parent_otel_span=None, local_only=True + ) + assert limiter._gauge_in_flight_from_cache_value(value) == 0 + finally: + finish_close.set() + await asyncio.wait_for(request, timeout=5) + _request_stash.reset(token) + value = await proxy.internal_usage_cache.async_get_cache( + "{api_key:close-order-owner}:max_parallel_requests", litellm_parent_otel_span=None, local_only=True + ) + assert limiter._gauge_in_flight_from_cache_value(value) == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", [None, "provider", "observer", "renewal", "legacy_key", "legacy_global"]) +async def test_signaling_keeps_or_releases_owned_call_lease(monkeypatch, failure): + import json + from unittest.mock import AsyncMock, MagicMock + + import httpx + from fastapi import Request + from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease + + from litellm.proxy import proxy_server as server + from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + + auth = UserAPIKeyAuth(max_parallel_requests=None if failure == "legacy_global" else 1) + lease = MagicMock(spec=RealtimeCallLease) + lease.renew = AsyncMock(return_value=failure != "renewal") + lease.close = AsyncMock() + legacy = failure in ("legacy_key", "legacy_global") + limiter = MagicMock(spec=_PROXY_MaxParallelRequestsHandler if legacy else _PROXY_MaxParallelRequestsHandler_v3) + if not legacy: + limiter.transfer_realtime_call_slot.return_value = lease + proxy = MagicMock() + proxy.get_proxy_hook.return_value = limiter + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr( + server, "general_settings", {"global_max_parallel_requests": 1} if failure == "legacy_global" else {} + ) + monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth)) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + process = AsyncMock(return_value=({}, None)) + monkeypatch.setattr(codex, "process_codex_request", process) + monkeypatch.setenv("LITELLM_SALT_KEY", "lease-transfer-test") + + async def route(**kwargs): + lease.start.assert_called_once() + if failure == "provider": + raise RuntimeError("Provider unavailable") + + async def respond(): + return httpx.Response( + 201, + text="v=0\r\n", + headers={"Location": "/v1/realtime/calls/rtc_lease"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}, + ) + + return respond() + + async def supervise(request, call, owner, selected_lease): + assert owner is auth + assert selected_lease is lease + assert call.parallel_reserved + if failure == "observer": + raise RuntimeError("Observer unavailable") + + monkeypatch.setattr(server, "route_request", route) + monkeypatch.setattr(codex, "supervise_codex_call", supervise) + 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")], + }, + AsyncMock( + return_value={ + "type": "http.request", + "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode(), + } + ), + ) + if failure is None: + response = await codex.create_codex_realtime_call(request) + token = response.headers["location"].rsplit("/", 1)[-1] + assert codex.decode_call(token, "Bearer owner").parallel_reserved + lease.close.assert_not_awaited() + else: + with pytest.raises((RuntimeError, HTTPException)) as raised: + await codex.create_codex_realtime_call(request) + if legacy: + assert raised.value.status_code == 400 + assert "V3 rate limiter" in raised.value.detail + process.assert_not_awaited() + lease.close.assert_not_awaited() + else: + if failure == "renewal": + assert raised.value.status_code == 503 + lease.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_signaling_rejection_after_admission_refunds_parallel_slot(monkeypatch): + import json + from unittest.mock import AsyncMock + + from fastapi import Request + + import litellm + from litellm.integrations.custom_logger import CustomLogger + from litellm.proxy import proxy_server as server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.utils import ProxyLogging + + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", []) + proxy._add_proxy_hooks() + limiter = proxy.get_proxy_hook("parallel_request_limiter") + key = "{api_key:rejected-signaling-owner}:max_parallel_requests" + + class Reject(CustomLogger): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + current = await proxy.internal_usage_cache.async_get_cache( + key, litellm_parent_otel_span=None, local_only=True + ) + assert limiter._gauge_in_flight_from_cache_value(current) == 1 + raise RuntimeError("Policy rejected after admission") + + litellm.callbacks.append(Reject()) + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr(server, "general_settings", {}) + monkeypatch.setattr( + server, + "llm_router", + litellm.Router( + model_list=[ + {"model_name": "voice", "litellm_params": {"model": "openai/gpt-realtime-1.5", "api_key": "test"}} + ] + ), + ) + auth = UserAPIKeyAuth(api_key="rejected-signaling-owner", max_parallel_requests=1) + monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth)) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + route = AsyncMock() + monkeypatch.setattr(server, "route_request", route) + for _ in range(2): + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "query_string": b"", + "headers": [(b"content-type", b"application/json")], + }, + AsyncMock( + return_value={ + "type": "http.request", + "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode(), + } + ), + ) + with pytest.raises(RuntimeError, match="Policy rejected after admission"): + await codex.create_codex_realtime_call(request) + current = await proxy.internal_usage_cache.async_get_cache(key, litellm_parent_otel_span=None, local_only=True) + assert limiter._gauge_in_flight_from_cache_value(current) == 0 + route.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("callback_order", ["before", "after", "cancel"]) +async def test_signaling_settles_tokens_once_with_isolated_sdk_callbacks(monkeypatch, callback_order): + import asyncio + import json + from datetime import datetime + from unittest.mock import AsyncMock + + import httpx + import litellm + from fastapi import Request + from litellm.proxy import proxy_server as server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.hooks.parallel_request_limiter_v3 import get_request_stash, isolated_request_stash + from litellm.proxy.utils import ProxyLogging + + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", []) + proxy._add_proxy_hooks() + limiter = proxy.get_proxy_hook("parallel_request_limiter") + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr(server, "general_settings", {}) + monkeypatch.setattr( + server, + "llm_router", + litellm.Router( + model_list=[ + {"model_name": "voice", "litellm_params": {"model": "openai/gpt-realtime-1.5", "api_key": "test"}} + ] + ), + ) + auth = UserAPIKeyAuth(api_key="signaling-settlement-owner", max_parallel_requests=1, tpm_limit=10000) + monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth)) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + supervisor = AsyncMock() + monkeypatch.setattr(codex, "supervise_codex_call", supervisor) + monkeypatch.setenv("LITELLM_SALT_KEY", "signaling-settlement-test") + ready, release_callback = asyncio.Event(), asyncio.Event() + callbacks = [] + + async def counter(kind): + return await proxy.internal_usage_cache.async_get_cache( + f"{{api_key:{auth.api_key}}}:{kind}", litellm_parent_otel_span=None, local_only=True + ) + + async def route(**kwargs): + assert get_request_stash() is None + assert await counter("tokens") > 0 + + async def callback(): + await release_callback.wait() + assert get_request_stash() is None + await limiter.async_log_success_event( + kwargs={ + "litellm_call_id": kwargs["data"]["litellm_call_id"], + "standard_logging_object": {"metadata": {"user_api_key_hash": auth.api_key}}, + }, + response_obj=litellm.ModelResponse(usage=litellm.Usage()), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + async def respond(): + assert get_request_stash() is None + ready.set() + if callback_order == "cancel": + await asyncio.Event().wait() + callbacks.append(asyncio.create_task(callback())) + if callback_order == "before": + release_callback.set() + await callbacks[0] + assert await counter("tokens") > 0 + return httpx.Response( + 201, + text="v=0\r\n", + headers={"Location": "/v1/realtime/calls/rtc_settlement"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}, + ) + + return respond() + + monkeypatch.setattr(server, "route_request", route) + 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")], + }, + AsyncMock( + return_value={ + "type": "http.request", + "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode(), + } + ), + ) + with isolated_request_stash(): + signaling = asyncio.create_task(codex.create_codex_realtime_call(request)) + await asyncio.wait_for(ready.wait(), timeout=2) + if callback_order == "cancel": + signaling.cancel() + with pytest.raises(asyncio.CancelledError): + await signaling + supervisor.assert_not_awaited() + else: + assert (await signaling).status_code == 201 + assert limiter._gauge_in_flight_from_cache_value(await counter("max_parallel_requests")) == 1 + await supervisor.call_args.args[3].close() + assert await counter("tokens") == 0 + release_callback.set() + await asyncio.gather(*callbacks) + assert await counter("tokens") == 0 + assert limiter._gauge_in_flight_from_cache_value(await counter("max_parallel_requests")) == 0 + + +@pytest.mark.asyncio +async def test_observer_startup_owns_call_lifecycle_with_synthetic_sockets(monkeypatch): + import asyncio + import json + from unittest.mock import AsyncMock, MagicMock + + from fastapi import Request + from starlette.websockets import WebSocketState + + from litellm.proxy.realtime_endpoints import call_supervision + + call = CodexRealtimeCall( + call_id="rtc_open", + model="gpt-live-1-codex", + alias="voice", + owner="owner", + api_base="https://voice.example/codex", + expires_at=time.time() + 60, + ) + auth = UserAPIKeyAuth() + logger = MagicMock() + logger.litellm_params = {} + observer: dict = {} + + async def process(request, data, _auth, _model, route_type, *, internal_realtime_observer=False): + observer["request"] = request + observer["data"] = data + observer["route_type"] = route_type + observer["internal_realtime_observer"] = internal_realtime_observer + return {"extra_headers": {}}, logger + + monkeypatch.setattr(codex, "process_codex_request", process) + + class Connection: + def __init__(self): + self.messages = asyncio.Queue() + self.close = AsyncMock() + + def __aiter__(self): + return self + + async def __anext__(self): + message = await self.messages.get() + if message is None: + raise StopAsyncIteration + return message + + connection = Connection() + await connection.messages.put(json.dumps({"type": "session.started"})) + stream_instance = MagicMock() + stream_instance.log_messages = AsyncMock() + terminations: list = [] + + class Handler: + def __init__(self, params, headers, extra_headers): + pass + + @staticmethod + def get_api_base(base): + return "https://gateway.test/v1" + + async def open_call_connection(self, model, base): + return connection + + async def close_call(self, opened, model, base): + terminations.append(("close", opened, model, base)) + + async def hangup_call(self, base): + terminations.append(("hangup", base)) + + monkeypatch.setattr(codex, "ChatGPTRealtime", Handler) + stream = MagicMock(return_value=stream_instance) + monkeypatch.setattr(codex, "RealTimeStreaming", stream) + started: list = [] + + original_start = call_supervision.CALL_SUPERVISORS.start + + async def capture_start(supervisor): + started.append(supervisor) + await original_start(supervisor) + + monkeypatch.setattr(call_supervision.CALL_SUPERVISORS, "start", capture_start) + request = Request({"type": "http", "headers": [], "method": "POST", "path": "/v1/realtime/calls"}) + + await codex.supervise_codex_call(request, call, auth) + + # The observer request serves the synthetic aliased-model body through the ASGI receive closure. + assert await observer["request"].json() == {"model": "voice"} + assert observer["data"]["model"] == "voice" + assert observer["route_type"] == "_arealtime" + assert observer["internal_realtime_observer"] is True + # The synthetic frontend completes the raw ASGI handshake through the swallow-and-return send closure. + frontend = stream.call_args.args[0] + await frontend.send({"type": "websocket.accept"}) + assert frontend.application_state is WebSocketState.CONNECTED + # The supervisor owns the opened connection: the startup socket stack released it without closing it. + assert len(started) == 1 + supervisor = started[0] + assert isinstance(supervisor, call_supervision.CallSupervisor) + assert supervisor._lease is None + assert supervisor._terminal_usage_required is (codex.realtime_endpoint(call.model) == "live") + connection.close.assert_not_awaited() + await supervisor._close_call() + assert terminations == [("close", connection, "gpt-live-1-codex", "https://gateway.test/v1")] + await supervisor._force_close_call() + assert terminations[-1] == ("hangup", "https://gateway.test/v1") + await connection.messages.put( + json.dumps({"type": "session.closed", "usage": {"total_tokens": 1}}) + ) + await supervisor.wait() + await call_supervision.CALL_SUPERVISORS.shutdown() + connection.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_mixed_case_multipart_without_boundary_returns_400(): + from fastapi import Request + from starlette.formparsers import MultiPartException + + async def receive(): + return {"type": "http.request", "body": b"sdp payload", "more_body": False} + + request = Request( + {"type": "http", "headers": [(b"content-type", b"Multipart/Form-Data; charset=utf-8")]}, + receive, + ) + assert not await request.form() + with pytest.raises(HTTPException) as error: + await codex.read_codex_offer(request) + assert error.value.status_code == 400 + assert error.value.detail == "Invalid realtime multipart offer" + assert isinstance(error.value.__cause__, MultiPartException) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("observer", [False, True]) +async def test_observer_processing_stamps_internal_request_origin(monkeypatch, observer): + from types import SimpleNamespace + + from fastapi import Request + + from litellm.proxy import common_request_processing + from litellm.proxy._types import InternalRequestOrigin + from litellm.proxy.realtime_endpoints.call_sessions import process_codex_request + + recorded: dict = {} + + class PassthroughProcessor: + def __init__(self, data): + self.data = data + + async def common_processing_pre_call_logic(self, **kwargs): + recorded["observer"] = kwargs.get("internal_realtime_observer", False) + return {**self.data, "extra_headers": {}}, SimpleNamespace(model_call_details={}) + + monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", PassthroughProcessor) + request = Request({"type": "http", "method": "POST", "path": "/v1/realtime", "headers": []}) + processed, logging_obj = await process_codex_request( + request, + {"model": "voice"}, + UserAPIKeyAuth(), + "voice", + "_arealtime", + internal_realtime_observer=observer, + ) + assert recorded["observer"] is observer + assert processed["model"] == "voice" + if observer: + assert logging_obj.model_call_details["internal_request_origin"] is InternalRequestOrigin.REALTIME_OBSERVER + else: + assert "internal_request_origin" not in logging_obj.model_call_details + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["upstream_exception", "invalid_response", "upstream_error_status", "unroutable"]) +async def test_signaling_response_paths_map_upstream_outcomes_to_http(monkeypatch, mode): + import json + from unittest.mock import AsyncMock + + import httpx + from fastapi import Request + + from litellm.llms.base_llm.chat.transformation import BaseLLMException + from litellm.proxy import common_request_processing, proxy_server + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + body = json.dumps({"sdp": "v=0\r\n", "session": {"model": "voice-alias"}}).encode() + + async def receive(): + return {"type": "http.request", "body": body, "more_body": False} + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "scheme": "http", + "server": ("localhost", 80), + "query_string": b"", + "headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer owner")], + }, + receive, + ) + monkeypatch.setattr(proxy_server, "master_key", "owner") + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + + class PassthroughProcessor: + def __init__(self, data): + self.data = data + + async def common_processing_pre_call_logic(self, **kwargs): + return self.data, None + + monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", PassthroughProcessor) + supervise = AsyncMock() + monkeypatch.setattr(codex, "supervise_codex_call", supervise) + + def route_returning(value): + async def route(**kwargs): + async def respond(): + return value + + return respond() + + return route + + def route_raising(exc): + async def route(**kwargs): + async def boom(): + raise exc + + return boom() + + return route + + if mode == "upstream_exception": + monkeypatch.setattr(proxy_server, "route_request", route_raising(BaseLLMException(429, "provider saturated"))) + with pytest.raises(HTTPException) as error: + await codex.create_codex_realtime_call(request) + assert error.value.status_code == 429 + assert "provider saturated" in str(error.value.detail) + elif mode == "invalid_response": + monkeypatch.setattr(proxy_server, "route_request", route_returning("not-an-http-response")) + with pytest.raises(HTTPException) as error: + await codex.create_codex_realtime_call(request) + assert error.value.status_code == 502 + assert error.value.detail == "Invalid realtime signaling response" + elif mode == "upstream_error_status": + monkeypatch.setattr( + proxy_server, "route_request", route_returning(httpx.Response(422, content=b'{"error":"invalid sdp"}')) + ) + response = await codex.create_codex_realtime_call(request) + assert response.status_code == 422 + assert response.body == b'{"error":"invalid sdp"}' + assert response.media_type == "application/json" + else: + monkeypatch.setattr(proxy_server, "route_request", route_returning(httpx.Response(201, content=b"v=0\r\n"))) + with pytest.raises(HTTPException) as error: + await codex.create_codex_realtime_call(request) + assert error.value.status_code == 400 + assert "ChatGPT deployment" in error.value.detail + supervise.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_sideband_begins_realtime_attachment_on_legacy_limiter(monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server as server + from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, + _RealtimeAttachmentReservations, + ) + from litellm.proxy.utils import InternalUsageCache + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-live-1-codex", + alias="voice", + usage_supervised=True, + owner=hashlib.sha256(b"Bearer owner").hexdigest(), + expires_at=time.time() + 300, + ) + auth = UserAPIKeyAuth() + logger = SimpleNamespace(model_call_details={}) + limiter = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(DualCache())) + proxy_logging = MagicMock() + proxy_logging.get_proxy_hook.return_value = limiter + monkeypatch.setattr(server, "proxy_logging_obj", proxy_logging) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + captured: dict = {} + + async def process(request, data, *_args, **_kwargs): + receipt = data.get("_legacy_realtime_attachment_reservations") + assert isinstance(receipt, _RealtimeAttachmentReservations) + assert receipt.cache_keys == () and receipt.global_acquired is False + captured["data"] = data + return {}, logger + + monkeypatch.setattr(codex, "process_codex_request", process) + monkeypatch.setattr(litellm, "_arealtime", AsyncMock()) + websocket = WebSocket( + { + "type": "websocket", + "path": "/v1/live/opaque", + "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")], + }, + AsyncMock(return_value={"type": "websocket.connect"}), + AsyncMock(), + ) + + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + + # The release consumed the receipt opened by begin_realtime_attachment before pre-call processing. + assert captured["data"]["_legacy_realtime_attachment_reservations"].take() == ((), False) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py new file mode 100644 index 00000000000..e7581faab0d --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -0,0 +1,858 @@ +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 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("lease_lost", [False, True]) +async def test_supervisor_holds_call_lease_until_terminal_accounting(lease_lost): + from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + lost = asyncio.Event() + lease = MagicMock(spec=RealtimeCallLease) + lease.wait_failed = lost.wait + + async def release(): + assert socket.closed + assert sink.logs == 1 + + lease.close = AsyncMock(side_effect=release) + + async def close(): + await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}}) + + terminate = AsyncMock(side_effect=close) + supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), terminate, lease=lease) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + lease.close.assert_not_awaited() + if lease_lost: + lost.set() + else: + await close() + await asyncio.wait_for(supervisor.wait(), 1) + assert terminate.await_count == int(lease_lost) + lease.close.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stalled_step", ["close", "drain"]) +async def test_live_initial_close_reserves_time_for_independent_hangup(stalled_step): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + close_cancelled = asyncio.Event() + + async def close(): + if stalled_step == "close": + try: + await asyncio.Event().wait() + finally: + close_cancelled.set() + + async def force_close(): + await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}}) + + force = AsyncMock(side_effect=force_close) + sink = Sink(logger) + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close, + force_close_call=force, + drain_timeout=1, + termination_timeout=0.08, + ) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await asyncio.wait_for(supervisor.close(), timeout=0.5) + force.assert_awaited_once() + assert close_cancelled.is_set() == (stalled_step == "close") + assert any(event["type"] == "session.closed" for event in sink.events) + assert not logger.model_call_details.get("realtime_usage_incomplete") + assert sink.logs == 1 + assert socket.closed + + +@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 + 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 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "duration,valid", [(0, True), (1000, True), (None, False), (-1, False), (True, False), ("1000", False)] +) +@pytest.mark.parametrize("duration_field", ["audio_duration_ms", "seconds"]) +async def test_live_terminal_requires_valid_duration_for_accounting(monkeypatch, duration, valid, duration_field): + from litellm.proxy.realtime_endpoints import call_supervision + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + invalidate = AsyncMock() + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) + close = AsyncMock() + 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 socket.messages.put( + {"type": "session.closed", **({"usage": {duration_field: duration}} if duration is not None else {})} + ) + await supervisor.wait() + close.assert_not_awaited() + force.assert_not_awaited() + assert socket.closed + assert bool(logger.model_call_details.get("realtime_usage_incomplete")) is not valid + assert invalidate.await_count == (0 if valid else 1) + + +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 +@pytest.mark.parametrize("cancel_count", [1, 2, 3]) +async def test_repeated_start_cancellation_keeps_lease_until_shutdown_finishes(cancel_count): + from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease + + reading = asyncio.Event() + close_entered = asyncio.Event() + allow_close = asyncio.Event() + released = asyncio.Event() + + class ObservedSocket(Socket): + async def __anext__(self): + reading.set() + return await super().__anext__() + + socket = ObservedSocket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + + async def close_call(): + close_entered.set() + await allow_close.wait() + await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}}) + + async def release(): + released.set() + + lease = RealtimeCallLease(renew=AsyncMock(return_value=True), release=release) + lease.start() + supervisor = CallSupervisor( + socket, sink, logger, UserAPIKeyAuth(), close_call, lease=lease, ready_timeout=10, termination_timeout=10 + ) + registry = CallSupervisors() + + async def signaling(): + transferred = False + try: + await registry.start(supervisor) + transferred = True + finally: + # The signaling endpoint retains lease ownership until registry startup succeeds. + if not transferred: + await lease.close() + + started = asyncio.create_task(signaling()) + shutdown = None + try: + await asyncio.wait_for(reading.wait(), timeout=1) + started.cancel() + await asyncio.wait_for(close_entered.wait(), timeout=1) + for _ in range(cancel_count - 1): + started.cancel() + done, _ = await asyncio.wait({started}, timeout=0.02) + assert not done + assert not released.is_set() + shutdown = asyncio.create_task(registry.shutdown()) + done, _ = await asyncio.wait({started, shutdown}, timeout=0.02) + assert not done + assert not released.is_set() + assert not socket.closed + assert sink.logs == 0 + allow_close.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(started, timeout=1) + await asyncio.wait_for(shutdown, timeout=1) + assert socket.closed + assert sink.logs == 1 + assert released.is_set() + assert sink.events[-1]["usage"]["audio_duration_ms"] == 1000 + finally: + allow_close.set() + await asyncio.wait_for(supervisor.wait(), timeout=1) + await asyncio.gather(started, return_exceptions=True) + if shutdown is not None: + await shutdown + await registry.shutdown() + await lease.close() + + +@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_invalidates_reservation_without_zeroing_spend(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() + 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_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() + 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 +async def test_shutdown_allows_hangup_longer_than_usage_drain_timeout(): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + hangup_started = asyncio.Event() + allow_hangup = asyncio.Event() + hangup_finished = asyncio.Event() + + async def hangup(): + hangup_started.set() + await allow_hangup.wait() + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + hangup_finished.set() + + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + hangup, + drain_timeout=0.01, + termination_timeout=1, + terminal_usage_required=False, + ) + registry = CallSupervisors() + await socket.messages.put({"type": "session.created"}) + await registry.start(supervisor) + shutdown = asyncio.create_task(registry.shutdown()) + try: + await asyncio.wait_for(hangup_started.wait(), timeout=1) + await asyncio.sleep(0.04) + assert not shutdown.done() + assert not socket.closed + assert not hangup_finished.is_set() + finally: + allow_hangup.set() + await asyncio.wait_for(shutdown, timeout=1) + assert hangup_finished.is_set() + assert socket.closed + assert sink.logs == 1 + assert sink.events[-1]["usage"]["total_tokens"] == 42 + assert not logger.model_call_details.get("realtime_usage_incomplete") + + +@pytest.mark.asyncio +async def test_termination_timeout_cancels_hangup_and_finishes_cleanup(): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + hangup_cancelled = asyncio.Event() + + async def hangup(): + try: + await asyncio.Event().wait() + finally: + hangup_cancelled.set() + + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + hangup, + drain_timeout=0.01, + termination_timeout=0.02, + terminal_usage_required=False, + ) + await socket.messages.put({"type": "session.created"}) + await supervisor.start() + await asyncio.wait_for(supervisor.close(), timeout=1) + assert hangup_cancelled.is_set() + assert socket.closed + assert sink.logs == 1 + assert logger.model_call_details["realtime_usage_incomplete"] is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("closure", ["eof", "normal_close", "error"]) +@pytest.mark.parametrize("hangup_succeeds", [True, False]) +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 + + +@pytest.mark.asyncio +async def test_live_connected_attach_is_ready_without_session_started_and_bills_once(): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + + async def close(): + await socket.messages.put({"type": "session.closed", "usage": {"seconds": 30}}) + + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close, + connected_ready=True, + ready_timeout=0.1, + ) + await supervisor.start() + await socket.messages.put({"type": "session.usage.updated", "usage": {"seconds": 15}}) + await supervisor.close() + await supervisor.close() + assert sink.logs == 1 + assert sink.events[-1] == {"type": "session.closed", "usage": {"seconds": 30}} + assert "realtime_usage_incomplete" not in logger.model_call_details + assert socket.closed + + +@pytest.mark.asyncio +async def test_live_missing_backend_accounting_invalidates_budget_after_dispatch(monkeypatch): + from litellm.proxy.realtime_endpoints import call_supervision + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + invalidate = AsyncMock() + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) + + class IncompleteSink(Sink): + async def log_messages(self, *, wait_for_dispatch=False): + await super().log_messages(wait_for_dispatch=wait_for_dispatch) + logger.model_call_details["realtime_backend_accounting_incomplete"] = True + + sink = IncompleteSink(logger) + supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock(), connected_ready=True) + await supervisor.start() + await socket.messages.put({"type": "session.closed", "usage": {"seconds": 30}}) + await supervisor.wait() + assert sink.logs == 1 + assert logger.model_call_details["realtime_accounting_incomplete"] is True + assert "realtime_usage_incomplete" not in logger.model_call_details + invalidate.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_repeated_start_is_rejected_while_observer_is_running(): + socket, sink, close_call, supervisor = fixture() + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + with pytest.raises(RuntimeError, match="Call observer already started"): + await supervisor.start() + # The rejected second start leaves the running observer untouched. + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 7}}) + await supervisor.wait() + await supervisor.close() + close_call.assert_not_awaited() + assert sink.logs == 1 + assert socket.closed + + +@pytest.mark.asyncio +async def test_start_rejects_quota_reservation_lost_during_startup(): + from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + + async def hangup(): + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + + close_call = AsyncMock(side_effect=hangup) + renew = AsyncMock(return_value=False) + release = AsyncMock() + lease = RealtimeCallLease(renew=renew, release=release, interval=3600) + lease.start() + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close_call, + ready_timeout=10, + lifetime=10, + drain_timeout=0.05, + lease=lease, + ) + await socket.messages.put({"type": "session.started"}) + with pytest.raises(RuntimeError, match="lost its quota reservation during startup"): + await supervisor.start() + assert socket.closed + close_call.assert_awaited_once() + assert renew.await_count >= 1 + release.assert_awaited_once() + assert sink.logs == 1 + + +@pytest.mark.asyncio +async def test_registry_watch_logs_observer_accounting_failure_without_payload(caplog, monkeypatch): + import logging + + from litellm.proxy.realtime_endpoints import call_supervision + + caplog.set_level(logging.ERROR, logger="LiteLLM Proxy") + + class FailingSink(Sink): + async def log_messages(self, *, wait_for_dispatch=False): + raise RuntimeError("observer accounting secret-token") + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = FailingSink(logger) + + async def hangup(): + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + + close_call = AsyncMock(side_effect=hangup) + invalidate = AsyncMock() + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close_call, + ready_timeout=5, + lifetime=5, + drain_timeout=0.05, + ) + registry = CallSupervisors() + await socket.messages.put({"type": "session.created"}) + await registry.start(supervisor) + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + with pytest.raises(RuntimeError, match="observer accounting secret-token"): + await supervisor.wait() + await registry.shutdown() + assert registry._calls == () + assert registry._tasks == () + assert socket.closed + close_call.assert_not_awaited() + invalidate.assert_awaited_once_with(budget_reservation=None) + assert logger.model_call_details["realtime_accounting_incomplete"] is True + proxy_logs = [record.getMessage() for record in caplog.records if record.name == "LiteLLM Proxy"] + assert any("Realtime observer accounting failed" in message for message in proxy_logs) + assert not any("secret-token" in message for message in proxy_logs) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py new file mode 100644 index 00000000000..034a90965d0 --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -0,0 +1,2392 @@ +import json +import time +from contextlib import asynccontextmanager +from types import MappingProxyType, SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, Mock + +import httpx +import pytest +from fastapi import FastAPI, HTTPException, Request +from fastapi.testclient import TestClient +from prisma.builder import QueryBuilder + +from litellm.llms.chatgpt.live import LiveDeployment +from litellm.models.budget import LiteLLM_BudgetTable +from litellm.models.organization import LiteLLM_OrganizationTable +from litellm.models.project import LiteLLM_ProjectTable +from litellm.models.team import LiteLLM_TeamTable +from litellm.models.team_membership import LiteLLM_TeamMembership +from litellm.models.user import LiteLLM_UserTable +from litellm.proxy._types import ModelAccessDeniedProxyException, UserAPIKeyAuth +from litellm.proxy.realtime_endpoints import live + + +@pytest.fixture(autouse=True) +def encryption_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-live-encryption-key") + + +def _auth_cache(initial=None): + values = dict(initial or {}) + + async def get(*, key, **kwargs): + return values.get(key) + + async def set(*, key, value, **kwargs): + values[key] = value + + # ``redis_cache`` is part of the UserApiKeyCache surface the auth helpers read: ``_cache_team_object`` + # compares it against the usage cache to decide whether one Redis holds both entries. The fake has no + # Redis at all, so it reports ``None``, the same way a standalone in-memory cache does. + return SimpleNamespace( + async_get_cache=AsyncMock(side_effect=get), + async_set_cache=AsyncMock(side_effect=set), + redis_cache=None, + ) + + +def handle(owner="owner", model_id=None): + deployment = {"model": "gpt-live", "provider": "openai"} + if model_id is not None: + deployment["model_id"] = model_id + return live._new_handle( + "sess_upstream", + "voice", + LiveDeployment(**deployment), + UserAPIKeyAuth(api_key=owner), + None, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("field", ["chatgpt_auth_profile", "chatgpt_token_dir", "chatgpt_auth_file"]) +async def test_live_rejects_unsupported_deployment_credentials(monkeypatch, field): + from litellm.proxy import proxy_server + + router = SimpleNamespace( + async_get_available_deployment=AsyncMock( + return_value={ + "litellm_params": {"model": "chatgpt/gpt-live-1", field: "other-account"}, + "model_info": {"id": "voice"}, + } + ), + async_routing_strategy_pre_call_checks=AsyncMock(), + ) + monkeypatch.setattr(proxy_server, "llm_router", router) + with pytest.raises(HTTPException) as rejected: + await live._deployment("voice", {}) + assert rejected.value.status_code == 400 + assert "deployment auth overrides are unsupported" in rejected.value.detail + + +def test_session_tokens_hide_credentials_and_enforce_owner_expiry_and_integrity(): + original = handle() + token = live.encode_session(original) + assert "deployment-a" not in token and "sess_upstream" not in token + assert live.decode_session(token, live._owner(UserAPIKeyAuth(api_key="owner"))) == original + for candidate, owner in ( + (token, "other-owner"), + ("sess_upstream", original.owner), + (token[:-8] + "aaaaaaaa", original.owner), + (live.encode_session(original.model_copy(update={"expires_at": time.time() - 1})), original.owner), + ): + with pytest.raises(HTTPException) as rejected: + live.decode_session(candidate, owner) + assert rejected.value.status_code == 403 + + +def test_handle_serializes_mappingproxy_without_losing_pinned_deployment(): + deployment = LiveDeployment( + model="gpt-live", + model_id="deployment-a", + api_base="https://upstream.test/v1", + extra_headers=MappingProxyType({"openai-beta": "test"}), + extra_query=MappingProxyType({"architecture": "test"}), + ) + original = live._new_handle("sess_upstream", "voice", deployment, UserAPIKeyAuth(api_key="owner"), None) + assert live._pinned(live.decode_session(live.encode_session(original), original.owner)) == deployment + + +def test_json_value_iteratively_serializes_nested_mappingproxy_tuple_and_shared_subtree(): + shared = MappingProxyType({"deep": (1, 2)}) + value = MappingProxyType({"left": shared, "right": (shared,)}) + + assert live._json_value(value) == {"left": {"deep": [1, 2]}, "right": [{"deep": [1, 2]}]} + + +@pytest.mark.parametrize("value", [{1: "invalid"}, {"invalid": object()}]) +def test_json_value_rejects_non_json_objects_and_keys(value): + with pytest.raises(ValueError, match="validation error"): + live._json_value(value) + + +def test_json_value_rejects_cycles_and_excessive_depth(): + cycle = {} + cycle["self"] = cycle + with pytest.raises(ValueError, match="depth"): + live._json_value(cycle) + + nested = json.loads('{"value":' * 257 + "null" + "}" * 257) + with pytest.raises(ValueError, match="depth"): + live._json_value(nested) + + +def test_only_protocol_session_ids_are_rewritten_and_application_values_survive(): + event = { + "type": "session.started", + "session": {"id": "raw", "instructions": "raw"}, + "session_id": "raw", + "delta": "raw", + "event": {"type": "response.output_text.delta", "delta": "raw", "session_id": "raw"}, + } + rewritten = live.rewrite_session_ids(event, "raw", "public") + assert rewritten["session"]["id"] == "public" + assert rewritten["session_id"] == "public" + assert rewritten["session"]["instructions"] == "raw" + assert rewritten["delta"] == "raw" + assert rewritten["event"] == event["event"] + assert event["session"]["id"] == "raw" + + +@pytest.fixture +def route_client(monkeypatch): + from litellm.proxy import proxy_server + + auth = UserAPIKeyAuth(api_key="owner") + deployment = LiveDeployment(model="gpt-live", provider="openai", api_key="upstream-key", model_id="deployment-a") + transport = SimpleNamespace( + request=AsyncMock( + return_value=httpx.Response( + 201, json={"session": {"id": "sess_upstream"}, "transport": {"type": "webrtc", "sdp": "answer"}} + ) + ) + ) + selected = AsyncMock(return_value=deployment) + supervised = AsyncMock() + authenticated_bodies = [] + authenticated_scopes = [] + + async def authenticate(request): + authenticated_bodies.append(await request.json()) + authenticated_scopes.append(request.scope) + return auth + + @asynccontextmanager + async def precall(request, auth, model, **kwargs): + yield live._Prepared({"model": model}, Mock(), None, kwargs.get("ownership")) + + monkeypatch.setattr(live, "_auth", authenticate) + monkeypatch.setattr(live, "_precall", precall) + monkeypatch.setattr(live, "_deployment", selected) + monkeypatch.setattr(live, "_supervise", supervised) + monkeypatch.setattr( + proxy_server, + "llm_router", + SimpleNamespace( + get_deployment=lambda model_id: ( + { + "model_name": "voice", + "litellm_params": {"model": "openai/gpt-live"}, + "model_info": {"id": model_id}, + } + if model_id == "deployment-a" + else None + ) + ), + ) + factory = Mock(return_value=transport) + monkeypatch.setattr(live, "LiveTransport", factory) + app = FastAPI() + app.include_router(live.router) + return SimpleNamespace( + client=TestClient(app), + transport=transport, + selected=selected, + supervised=supervised, + auth=auth, + factory=factory, + bodies=authenticated_bodies, + scopes=authenticated_scopes, + ) + + +@pytest.mark.parametrize("prefix", ["/v1/live", "/live", "/openai/v1/live"]) +def test_create_preserves_configuration_and_returns_owned_json_session(route_client, prefix): + body = { + "session": { + "model": "voice", + "instructions": "hello", + "input": [{"role": "user", "content": "hi"}], + "audio": {"output": {"voice": "marin"}}, + "delegation": {"type": "client"}, + "future_option": {"enabled": True}, + }, + "transport": {"type": "webrtc", "sdp": "offer"}, + "api_base": "https://untrusted.test", + } + result = route_client.client.post(prefix + "/sessions", json=body) + assert result.status_code == 201 + output = result.json() + assert output["transport"] == {"type": "webrtc", "sdp": "answer"} + original = live.decode_session(output["session"]["id"], live._owner(route_client.auth)) + assert original.session_id == "sess_upstream" + assert original.deployment["api_key"] == "upstream-key" + assert original.initialization_seconds == 15 + forwarded = route_client.transport.request.await_args.kwargs["body"] + assert forwarded["session"] == {**body["session"], "model": "gpt-live"} + assert route_client.bodies[0] == {**body, "model": "voice"} + route_client.supervised.assert_awaited_once() + assert route_client.factory.call_args.args[0].api_base is None + + +def test_synthetic_live_request_drops_the_cached_body_and_pins_the_model_on_request(): + source = Request({"type": "http", "method": "POST", "headers": [], "path": "/v1/live/sessions"}) + source.scope["parsed_body"] = (("model",), {"model": "stale-alias"}) + + pinned = live._request(source, MappingProxyType({"model": "voice"}), "voice") + assert "parsed_body" not in pinned.scope + assert pinned.scope["litellm_pinned_realtime_model"] == "voice" + + plain = live._request(source, MappingProxyType({"model": "voice"})) + assert "litellm_pinned_realtime_model" not in plain.scope + assert "parsed_body" not in plain.scope + assert source.scope["parsed_body"] == (("model",), {"model": "stale-alias"}) + + +@pytest.mark.asyncio +async def test_synthetic_live_request_sends_the_replaced_body(): + source = Request({"type": "http", "method": "POST", "headers": [], "path": "/v1/live/sessions"}) + source.scope["parsed_body"] = (("model",), {"model": "stale-alias"}) + + request = live._request(source, MappingProxyType({"model": "voice"})) + + assert await request.json() == {"model": "voice"} + + +def test_live_create_and_fork_pin_the_model_they_dispatch(route_client): + created = route_client.client.post( + "/v1/live/sessions", + json={"session": {"model": "voice"}, "transport": {"type": "webrtc", "sdp": "offer"}}, + ) + assert created.status_code == 201 + assert [scope.get("litellm_pinned_realtime_model") for scope in route_client.scopes] == ["voice"] + + token = live.encode_session(handle(model_id="deployment-a")) + forked = route_client.client.post( + f"/v1/live/sessions/{token}/fork", + json={"session": {}, "transport": {"type": "webrtc", "sdp": "offer"}}, + ) + assert forked.status_code == 201 + assert [scope.get("litellm_pinned_realtime_model") for scope in route_client.scopes[1:]] == [None, "voice"] + + +def test_fork_preserves_empty_overrides_and_pins_source_deployment(route_client): + source = handle(model_id="deployment-a") + token = live.encode_session(source) + body = {"session": {}, "transport": {"type": "webrtc", "sdp": "offer"}} + result = route_client.client.post(f"/v1/live/sessions/{token}/fork", json=body) + assert result.status_code == 201 + assert route_client.transport.request.await_args.args == ("POST", "live/sessions/sess_upstream/fork") + assert route_client.transport.request.await_args.kwargs["body"] == body + assert route_client.factory.call_args.args[0].model_id == "deployment-a" + route_client.selected.assert_not_awaited() + + +@pytest.mark.parametrize( + "configured", + [ + None, + { + "model_name": "voice", + "litellm_params": {"model": "openai/gpt-replaced"}, + "model_info": {"id": "deployment-a"}, + }, + { + "model_name": "voice", + "litellm_params": {"model": "openai/gpt-live"}, + "model_info": {"id": "deployment-a", "blocked": True}, + }, + ], + ids=["removed", "replaced", "blocked"], +) +def test_fork_rejects_removed_or_replaced_source_deployment(route_client, monkeypatch, configured): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server.llm_router, "get_deployment", lambda model_id: configured) + token = live.encode_session(handle(model_id="deployment-a")) + + result = route_client.client.post( + f"/v1/live/sessions/{token}/fork", + json={"session": {}, "transport": {"type": "webrtc", "sdp": "offer"}}, + ) + + assert result.status_code == 410 + assert "no longer available" in result.json()["detail"] + route_client.transport.request.assert_not_awaited() + + +def test_fork_cannot_change_model_even_to_same_alias(route_client): + token = live.encode_session(handle()) + response = route_client.client.post(f"/v1/live/sessions/{token}/fork", json={"session": {"model": "voice"}}) + assert response.status_code == 400 + route_client.transport.request.assert_not_awaited() + + +def test_cross_key_followup_and_raw_incoming_ids_never_contact_upstream(route_client): + token = live.encode_session(handle("different-key")) + for path in (f"{token}/hangup", "sess_other/accept", "sess_other/reject"): + response = route_client.client.post(f"/v1/live/sessions/{path}", json={}) + assert response.status_code == 403 + route_client.transport.request.assert_not_awaited() + + +def test_recording_preserves_binary_body_status_and_content_headers(route_client): + token = live.encode_session(handle()) + route_client.transport.request.return_value = httpx.Response( + 206, + content=b"\x00\xffrecording", + headers={ + "content-type": "video/mp4", + "content-disposition": "attachment; filename=recording.mp4", + "content-range": "bytes 0-10/20", + }, + ) + result = route_client.client.get(f"/v1/live/sessions/{token}/content") + assert result.status_code == 206 + assert result.content == b"\x00\xffrecording" + assert result.headers["content-type"] == "video/mp4" + assert result.headers["content-range"] == "bytes 0-10/20" + + +@pytest.mark.asyncio +async def test_pre_call_preserves_safe_body_options_and_refunds_failed_signaling(monkeypatch): + from litellm.proxy import proxy_server + + auth = UserAPIKeyAuth(api_key="owner") + authorize = AsyncMock() + process = AsyncMock(return_value=({"model": "voice"}, Mock())) + release = AsyncMock() + monkeypatch.setattr(live, "_authorize", authorize) + monkeypatch.setattr(live, "process_codex_request", process) + monkeypatch.setattr(live, "release_or_invalidate_budget_reservation", release) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda _: None)) + request = Request({"type": "http", "headers": []}) + synthetic = live._request( + request, + { + "session": {"model": "voice", "instructions": "safe"}, + "api_base": "https://untrusted.test", + "extra_headers": {"x-admin": "true"}, + }, + ) + with pytest.raises(RuntimeError, match="signaling failed"): + async with live._precall(synthetic, auth, "voice"): + raise RuntimeError("signaling failed") + authorize.assert_awaited_once_with("voice", auth) + data = process.await_args.args[1] + assert data["session"]["instructions"] == "safe" + assert "api_base" not in data and "extra_headers" not in data + release.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_transferred_supervisor_retains_budget_on_client_disconnect(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(live, "_authorize", AsyncMock()) + monkeypatch.setattr(live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, Mock()))) + release = AsyncMock() + monkeypatch.setattr(live, "release_or_invalidate_budget_reservation", release) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda _: None)) + request = live._request(Request({"type": "http", "headers": []}), {"model": "voice"}) + + async def disconnect_after_transfer(): + async with live._precall(request, UserAPIKeyAuth(api_key="owner"), "voice") as prepared: + prepared.transferred = True + raise RuntimeError("client disconnected") + + with pytest.raises(RuntimeError, match="client disconnected"): + await disconnect_after_transfer() + release.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_info_before_started_is_forwarded_without_rejecting_session(): + info = {"type": "info", "code": "data_channel_permissions", "message": "ready"} + started = {"type": "session.started", "session": {"id": "sess_upstream"}} + socket = SimpleNamespace(recv=AsyncMock(side_effect=[json.dumps(info), json.dumps(started)])) + client = SimpleNamespace(send_json=AsyncMock()) + assert await live._wait_started(socket, client) == started + client.send_json.assert_awaited_once_with(info) + + +@pytest.mark.asyncio +async def test_upstream_startup_error_is_forwarded_without_creating_session(): + error = {"type": "error", "error": {"code": "forbidden", "message": "Voice access denied"}} + socket = SimpleNamespace(recv=AsyncMock(return_value=json.dumps(error))) + client = SimpleNamespace(send_json=AsyncMock()) + with pytest.raises(HTTPException) as rejected: + await live._wait_started(socket, client) + assert rejected.value.status_code == 502 + client.send_json.assert_awaited_once_with(error) + + +@pytest.mark.asyncio +async def test_session_events_keep_stable_public_id_and_feed_shared_usage_sink(): + original = handle() + token = live.encode_session(original) + client = SimpleNamespace(send_text=AsyncMock(), scope={}, headers={}) + observer = SimpleNamespace(store_message=Mock()) + frontend = live._PublicSocket(client, original, token, UserAPIKeyAuth(api_key="owner"), observer) + event = { + "type": "session.updated", + "session": {"id": original.session_id}, + "event": {"type": "response.completed", "response": {"id": "resp_a"}}, + } + for _ in range(2): + await frontend.send_text(json.dumps(event)) + assert json.loads(client.send_text.await_args.args[0])["session"]["id"] == token + assert observer.store_message.call_count == 2 + + +@pytest.mark.asyncio +async def test_websocket_delegation_model_update_is_authorized_before_forwarding(monkeypatch): + authorize = AsyncMock(side_effect=HTTPException(403, "Model forbidden")) + monkeypatch.setattr(live, "_authorize", authorize) + message = {"type": "session.update", "session": {"delegation": {"responses": {"model": "unauthorized"}}}} + client = SimpleNamespace(receive_text=AsyncMock(return_value=json.dumps(message)), scope={}, headers={}) + auth = UserAPIKeyAuth(api_key="owner") + frontend = live._PublicSocket(client, handle(), "public", auth) + with pytest.raises(HTTPException) as rejected: + await frontend.receive_text() + assert rejected.value.status_code == 403 + authorize.assert_awaited_once_with("unauthorized", auth) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "limits", + [ + {"rpm_limit": 1}, + {"rpm_limit": 0}, + {"model_max_budget": {"backend": 1}}, + {"rpm_limit_per_model": {"backend": 0}}, + {"tpm_limit_per_model": {"backend": 0}}, + {"team_tpm_limit": 10}, + {"organization_rpm_limit": 0}, + {"organization_tpm_limit": 0}, + {"team_member_rpm_limit": 0}, + {"team_member_tpm_limit": 0}, + {"end_user_rpm_limit": 0}, + {"end_user_tpm_limit": 0}, + {"team_metadata": {"model_rpm_limit": {"backend": 1}}}, + {"metadata": {"scopes": [{"nested": {"model_tpm_limit": {"backend": 0}}}]}}, + ], +) +async def test_managed_delegation_fails_closed_for_unenforceable_constraints(limits): + body = {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}} + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation(body, UserAPIKeyAuth(api_key="owner", **limits)) + assert rejected.value.status_code == 400 + assert "client delegation" in rejected.value.detail + + +@pytest.mark.asyncio +@pytest.mark.parametrize("delegation", [{"type": "responses"}, {"type": "responses", "responses": {}}]) +async def test_restricted_session_update_can_retain_backend_delegation_model( + delegation, +): + body = {"type": "session.update", "session": {"delegation": delegation}} + result = await live._authorize_delegation(body, UserAPIKeyAuth(api_key="owner", models=["voice", "backend"])) + assert result is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("session", [{}, {"delegation": None}, {"delegation": {"type": "client"}}]) +async def test_constrained_webrtc_client_delegation_allows_frontend_updates( + session, +): + result = await live._authorize_delegation( + {"session": session, "transport": {"type": "webrtc"}}, + UserAPIKeyAuth(api_key="owner", models=["voice"], rpm_limit=10), + ) + assert result is None + + +@pytest.mark.parametrize( + "limits", + [ + {"max_budget": 1}, + {"team_max_budget": 1}, + {"user_max_budget": 1}, + {"end_user_max_budget": 1}, + {"organization_max_budget": 1}, + {"budget_limits": [{"budget_duration": "1d", "max_budget": 1}]}, + ], +) +def test_managed_constraints_detect_scalar_and_window_budgets(limits): + assert live._managed_constraints(UserAPIKeyAuth(api_key="owner", **limits)) is True + + +@pytest.mark.parametrize("global_limit", [None, 8]) +def test_managed_constraints_covers_configured_parallel_limits(monkeypatch, global_limit): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "general_settings", {"global_max_parallel_requests": global_limit}) + + key_limited = live._managed_constraints(UserAPIKeyAuth(api_key="owner", max_parallel_requests=2)) + globally_limited = live._managed_constraints(UserAPIKeyAuth(api_key="owner")) + assert key_limited is True + assert globally_limited is (global_limit is not None) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("limits", [{}, {"max_parallel_requests": 2}]) +async def test_managed_delegation_requires_client_delegation_under_a_concurrency_limit(monkeypatch, limits): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "general_settings", {}) + body = {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}} + auth = UserAPIKeyAuth(api_key="owner", **limits) + + if not limits: + assert await live._authorize_delegation(body, auth) is None + return + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation(body, auth) + assert rejected.value.status_code == 400 + assert "use client delegation" in str(rejected.value.detail) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "member_limit, default_limit, blocked", + [(0, None, True), (1, None, True), (None, 1, True), (None, 0, False), (None, None, False)], +) +async def test_managed_delegation_checks_authoritative_member_and_default_budget( + monkeypatch, member_limit, default_limit, blocked +): + auth = UserAPIKeyAuth( + api_key="owner", team_id="team", user_id="user", team_metadata={"team_member_budget_id": "budget"} + ) + from litellm.proxy import proxy_server + + membership = AsyncMock( + return_value=LiteLLM_TeamMembership( + user_id="user", team_id="team", litellm_budget_table=LiteLLM_BudgetTable(max_budget=member_limit) + ) + ) + db = SimpleNamespace( + litellm_teammembership=SimpleNamespace(find_unique=membership), + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=None)), + litellm_budgettable=SimpleNamespace( + find_unique=AsyncMock(return_value=LiteLLM_BudgetTable(max_budget=default_limit)) + ), + ) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache()) + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr(live, "_authorize", AsyncMock()) + body = {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}} + + if blocked: + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation(body, auth) + assert rejected.value.status_code == 400 + assert "client delegation" in rejected.value.detail + else: + await live._authorize_delegation(body, auth) + assert membership.await_args.kwargs["where"] == {"user_id_team_id": {"user_id": "user", "team_id": "team"}} + + +@pytest.mark.asyncio +async def test_managed_delegation_rejects_unverifiable_member_budget_but_allows_client(monkeypatch): + from litellm.proxy import proxy_server + + db = SimpleNamespace( + litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(side_effect=RuntimeError("Database unavailable"))) + ) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache()) + auth = UserAPIKeyAuth(api_key="owner", team_id="team", user_id="user") + body = {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}} + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation(body, auth) + assert rejected.value.status_code == 503 + await live._authorize_delegation({"session": {"delegation": {"type": "client"}}}, auth) + + +def test_managed_constraints_fails_closed_after_metadata_node_limit(): + auth = UserAPIKeyAuth(api_key="owner") + auth.metadata = {"items": [{} for _ in range(4097)]} + + assert live._managed_constraints(auth) is True + + +@pytest.mark.parametrize( + "limits", + [ + {"model_max_budget": {}}, + {"team_metadata": {"model_rpm_limit": {}}}, + {"metadata": {"nested": [{"model_max_budget": {}}]}}, + ], +) +def test_empty_model_limit_maps_do_not_mark_delegation_as_managed(monkeypatch, limits): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "general_settings", {}) + assert live._managed_constraints(UserAPIKeyAuth(api_key="owner", **limits)) is False + + +def test_managed_constraints_terminates_on_cyclic_metadata_without_a_limit(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "general_settings", {}) + metadata = {} + metadata["self"] = metadata + auth = UserAPIKeyAuth(api_key="owner") + auth.metadata = metadata + + assert live._managed_constraints(auth) is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "scope", + [ + {"access_group_ids": ["restricted-group"]}, + {"project_id": "restricted-project"}, + {"org_id": "restricted-org"}, + {"team_id": "restricted-team"}, + {"team_id": "restricted-team", "user_id": "restricted-member"}, + ], +) +async def test_restricted_scopes_cannot_delegate_without_an_explicit_backend_model(monkeypatch, scope): + monkeypatch.setattr(live, "_managed_member_budget", AsyncMock(return_value=False)) + body = {"session": {"delegation": {"type": "responses", "responses": {}}}} + + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation(body, UserAPIKeyAuth(api_key="owner", **scope)) + + assert rejected.value.status_code == 400 + assert "explicit authorized delegation.responses.model" in rejected.value.detail + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "scope", + [ + {"models": ["voice", "backend"]}, + {"access_group_ids": ["restricted-group"]}, + {"matched_model_access_groups": ["restricted-group"]}, + {"project_id": "restricted-project"}, + {"org_id": "restricted-org"}, + {"team_id": "restricted-team"}, + {"team_id": "restricted-team", "user_id": "restricted-member"}, + ], +) +async def test_restricted_webrtc_cannot_change_managed_model_outside_proxy(monkeypatch, scope): + monkeypatch.setattr(live, "_managed_member_budget", AsyncMock(return_value=False)) + monkeypatch.setattr(live, "_authorize", AsyncMock()) + body = { + "session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}, + "transport": {"type": "webrtc", "sdp": "offer"}, + } + auth = UserAPIKeyAuth(api_key="owner", **scope) + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation(body, auth) + assert rejected.value.status_code == 400 + body["session"]["client"] = {"data_channel": {"allowed_client_events": ["session.close"]}} + await live._authorize_delegation(body, auth) + + +@pytest.mark.asyncio +async def test_budget_scope_releases_on_ownership_decode_error(monkeypatch): + release = AsyncMock() + monkeypatch.setattr(live, "release_or_invalidate_budget_reservation", release) + auth = UserAPIKeyAuth(api_key="owner") + with pytest.raises(HTTPException): + async with live._budget_scope(auth): + live.decode_session("raw-session-id", live._owner(auth)) + release.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_model_authorization_rejection_still_releases_budget(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(live, "_authorize", AsyncMock(side_effect=HTTPException(403, "Model forbidden"))) + release = AsyncMock() + monkeypatch.setattr(live, "release_or_invalidate_budget_reservation", release) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda _: None)) + request = live._request(Request({"type": "http", "headers": []}), {"model": "forbidden"}) + with pytest.raises(HTTPException): + async with live._precall(request, UserAPIKeyAuth(api_key="owner"), "forbidden"): + pytest.fail("Upstream must not be reached") + release.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_precall_guardrail_mutations_are_used_without_forwarding_routing_options(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(live, "_authorize", AsyncMock()) + monkeypatch.setattr(live, "release_or_invalidate_budget_reservation", AsyncMock()) + process = AsyncMock( + return_value=( + { + "model": "voice", + "session": {"model": "voice", "instructions": "redacted"}, + "api_base": "https://internal.test", + }, + Mock(), + ) + ) + monkeypatch.setattr(live, "process_codex_request", process) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda _: None)) + body = {"type": "session.start", "session": {"model": "voice", "instructions": "sensitive"}} + request = live._request(Request({"type": "http", "headers": []}), body) + async with live._precall(request, UserAPIKeyAuth(api_key="owner"), "voice") as prepared: + forwarded = live._processed_body(body, prepared.processed) + assert forwarded["session"]["instructions"] == "redacted" + assert forwarded["type"] == "session.start" + assert "api_base" not in forwarded + assert process.await_args.args[1]["session"]["instructions"] == "sensitive" + + +@pytest.mark.asyncio +async def test_failed_live_signaling_releases_real_parallel_limiter_and_can_retry(monkeypatch): + import litellm + from litellm.proxy import proxy_server as server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.utils import ProxyLogging + + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", []) + proxy._add_proxy_hooks() + model_list = [{"model_name": "voice", "litellm_params": {"model": "openai/gpt-live-1", "api_key": "upstream-key"}}] + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr(server, "general_settings", {}) + monkeypatch.setattr(server, "llm_model_list", model_list) + monkeypatch.setattr(server, "llm_router", litellm.Router(model_list=model_list)) + auth = UserAPIKeyAuth(api_key="live-limiter-owner", max_parallel_requests=1) + authenticate = AsyncMock(return_value=auth) + monkeypatch.setattr(live, "user_api_key_auth", authenticate) + transport = SimpleNamespace( + request=AsyncMock(return_value=httpx.Response(403, json={"error": {"code": "forbidden"}})) + ) + monkeypatch.setattr(live, "LiveTransport", Mock(return_value=transport)) + supervisor = AsyncMock() + monkeypatch.setattr(live, "_supervise", supervisor) + for _ in range(2): + request = live._request( + Request( + { + "type": "http", + "method": "POST", + "path": "/v1/live/sessions", + "query_string": b"", + "headers": [(b"content-type", b"application/json")], + } + ), + {"session": {"model": "voice"}, "transport": {"type": "webrtc", "sdp": "offer"}}, + ) + result = await live.create_live_session(request) + assert result.status_code == 403 + current = await proxy.internal_usage_cache.async_get_cache( + "{api_key:live-limiter-owner}:max_parallel_requests", litellm_parent_otel_span=None, local_only=True + ) + limiter = proxy.get_proxy_hook("parallel_request_limiter") + assert limiter._gauge_in_flight_from_cache_value(current) == 0 + assert transport.request.await_count == 2 + supervisor.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_inherited_managed_fork_cannot_bypass_new_key_constraints(): + source = handle().model_copy( + update={"policy": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}} + ) + payload = live._policy_body({"session": {}, "transport": {"type": "webrtc", "sdp": "offer"}}, source) + assert payload["session"]["delegation"]["responses"]["model"] == "backend" + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation(payload, UserAPIKeyAuth(api_key="owner", rpm_limit_per_model={"backend": 1})) + assert rejected.value.status_code == 400 + + +@pytest.mark.parametrize("protocol", ["http", "websocket"]) +@pytest.mark.parametrize("startup_policy", [{"delegation": {"type": "responses", "responses": {"model": "allowed"}}}]) +@pytest.mark.parametrize("overrides", [{}, {"delegation": {"responses": {}}}]) +def test_restricted_fork_never_trusts_startup_delegation(route_client, protocol, startup_policy, overrides): + from starlette.websockets import WebSocketDisconnect + + # The source may now use a revoked backend, including after an unrestricted WebRTC update. + route_client.auth.models = ["voice", "allowed"] + token = live.encode_session(handle().model_copy(update={"policy": startup_policy})) + path = f"/v1/live/sessions/{token}/fork" + if protocol == "http": + response = route_client.client.post(path, json={"session": overrides}) + assert response.status_code == 400 + else: + with route_client.client.websocket_connect(path, headers={"Authorization": "Bearer owner"}) as ws: + ws.send_json({"type": "session.start", "session": overrides}) + with pytest.raises(WebSocketDisconnect) as rejected: + ws.receive_json() + assert rejected.value.code == 1008 + route_client.factory.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("startup_policy", [{}, {"delegation": {"type": "responses", "responses": {"model": "old"}}}]) +async def test_explicit_fork_backend_is_authorized_even_when_startup_policy_differs(monkeypatch, startup_policy): + authorize = AsyncMock() + monkeypatch.setattr(live, "_authorize", authorize) + source = handle().model_copy(update={"policy": startup_policy}) + auth = UserAPIKeyAuth(api_key="owner", models=["voice", "allowed"]) + body = {"session": {"delegation": {"type": "responses", "responses": {"model": "allowed"}}}} + await live._authorize_fork_policy(body, source, auth) + authorize.assert_awaited_once_with("allowed", auth) + authorize.side_effect = HTTPException(403, "Model revoked") + with pytest.raises(HTTPException) as rejected: + await live._authorize_fork_policy(body, source, auth) + assert rejected.value.status_code == 403 + + +def test_restricted_client_fork_can_inherit_delegation(route_client): + route_client.auth.models = ["voice"] + body = {"session": {}} + token = live.encode_session(handle(model_id="deployment-a")) + route_client.transport.request.return_value = httpx.Response( + 200, json={"session": {"id": "sess_fork"}, "transport": {"type": "webrtc", "sdp": "answer"}} + ) + response = route_client.client.post(f"/v1/live/sessions/{token}/fork", json=body) + assert response.status_code == 200 + assert response.json()["transport"]["sdp"] == "answer" + route_client.transport.request.assert_awaited_once_with("POST", "live/sessions/sess_upstream/fork", body=body) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("limits", [{"rpm_limit": 10}, {"tpm_limit": 100}, {"model_max_budget": {"backend": 1}}]) +async def test_client_fork_can_inherit_immutable_delegation_with_new_limits( + limits, +): + result = await live._authorize_fork_policy({"session": {}}, handle(), UserAPIKeyAuth(api_key="owner", **limits)) + assert result is None + + +@pytest.mark.asyncio +async def test_explicit_managed_fork_still_rejected_for_new_rate_limits(): + with pytest.raises(HTTPException) as rejected: + await live._authorize_fork_policy( + {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}}, + handle(), + UserAPIKeyAuth(api_key="owner", rpm_limit=10), + ) + assert rejected.value.status_code == 400 + assert "cannot enforce" in rejected.value.detail + + +@pytest.mark.asyncio +async def test_startup_usage_buffer_preserves_nested_events_until_supervisor_owns_accounting(): + usage = { + "type": "response.event", + "event": { + "type": "response.completed", + "response": { + "id": "resp_before_start", + "model": "backend", + "usage": {"input_tokens": 7, "output_tokens": 2}, + }, + }, + } + started = {"type": "session.started", "session": {"id": "sess_upstream"}} + backend = SimpleNamespace(recv=AsyncMock(side_effect=[json.dumps(usage), json.dumps(started)])) + client = SimpleNamespace(send_json=AsyncMock()) + startup = live._StartupEvents() + assert await live._wait_started(backend, client, startup) == started + assert [json.loads(message) for message in startup.messages] == [usage] + + +def test_primary_websocket_authenticates_model_and_keeps_public_event_shape(route_client): + from websockets.exceptions import ConnectionClosedOK + from websockets.frames import Close + + event = {"type": "future.live.event", "session_id": "sess_upstream", "payload": {"untouched": ["a", 2]}} + backend = SimpleNamespace( + send=AsyncMock(), + close=AsyncMock(), + recv=AsyncMock( + side_effect=[ + json.dumps({"type": "info", "code": "ready", "message": "Preparing"}), + json.dumps({"type": "session.started", "session": {"id": "sess_upstream"}}), + json.dumps(event), + ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True), + ] + ), + ) + route_client.transport.connect = AsyncMock(return_value=backend) + observer = SimpleNamespace(store_message=Mock()) + route_client.supervised.return_value = observer + with route_client.client.websocket_connect("/v1/live/sessions", headers={"Authorization": "Bearer owner"}) as ws: + ws.send_json({"type": "session.start", "session": {"model": "voice", "instructions": "hello", "unknown": True}}) + assert ws.receive_json()["type"] == "info" + started = ws.receive_json() + public_id = started["session"]["id"] + assert live.decode_session(public_id, live._owner(route_client.auth)).session_id == "sess_upstream" + forwarded = ws.receive_json() + assert forwarded == {**event, "session_id": public_id} + assert route_client.bodies[0] == {} + assert route_client.bodies[1]["model"] == "voice" + assert route_client.bodies[1]["session"]["instructions"] == "hello" + initial = json.loads(backend.send.await_args.args[0]) + assert initial == { + "type": "session.start", + "session": {"model": "gpt-live", "instructions": "hello", "unknown": True}, + } + assert any(json.loads(call.args[0]) == event for call in observer.store_message.call_args_list) + + +@pytest.mark.asyncio +async def test_successful_live_session_holds_real_parallel_slot_until_supervisor_releases(monkeypatch): + import litellm + from litellm.proxy import proxy_server as server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.utils import ProxyLogging + + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", []) + proxy._add_proxy_hooks() + models = [{"model_name": "voice", "litellm_params": {"model": "openai/gpt-live-1", "api_key": "upstream-key"}}] + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr(server, "general_settings", {}) + monkeypatch.setattr(server, "llm_model_list", models) + monkeypatch.setattr(server, "llm_router", litellm.Router(model_list=models)) + auth = UserAPIKeyAuth(api_key="live-held-slot", max_parallel_requests=1) + monkeypatch.setattr(live, "user_api_key_auth", AsyncMock(return_value=auth)) + transport = SimpleNamespace( + request=AsyncMock( + return_value=httpx.Response( + 201, json={"session": {"id": "sess_upstream"}, "transport": {"type": "webrtc", "sdp": "answer"}} + ) + ) + ) + monkeypatch.setattr(live, "LiveTransport", Mock(return_value=transport)) + leases = [] + + async def supervise(request, handle, auth, logger, lease): + leases.append(lease) + + monkeypatch.setattr(live, "_supervise", supervise) + + def request(): + return live._request( + Request( + { + "type": "http", + "method": "POST", + "path": "/v1/live/sessions", + "query_string": b"", + "headers": [(b"content-type", b"application/json")], + } + ), + {"session": {"model": "voice"}, "transport": {"type": "webrtc", "sdp": "offer"}}, + ) + + try: + result = await live.create_live_session(request()) + assert result.status_code == 201 + assert leases[0] is not None + with pytest.raises(HTTPException) as blocked: + await live.create_live_session(request()) + assert blocked.value.status_code == 429 + assert transport.request.await_count == 1 + finally: + for lease in leases: + if lease is not None: + await lease.close() + current = await proxy.internal_usage_cache.async_get_cache( + "{api_key:live-held-slot}:max_parallel_requests", litellm_parent_otel_span=None, local_only=True + ) + limiter = proxy.get_proxy_hook("parallel_request_limiter") + assert limiter._gauge_in_flight_from_cache_value(current) == 0 + + +@pytest.mark.asyncio +async def test_supervisor_logger_keeps_deployment_pricing(monkeypatch): + import litellm + + original = handle().model_copy( + update={"deployment": {**handle().deployment, "model_id": "deployment-priced"}, "initialization_seconds": 15} + ) + logger = Mock(litellm_params={}, model_call_details={}) + monkeypatch.setattr(live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, logger))) + monkeypatch.setattr(litellm, "get_model_info", Mock(return_value={"input_cost_per_second": 0.1})) + connection = SimpleNamespace(close=AsyncMock()) + transport = SimpleNamespace(connect=AsyncMock(return_value=connection), request=AsyncMock()) + monkeypatch.setattr(live, "LiveTransport", Mock(return_value=transport)) + started = AsyncMock() + monkeypatch.setattr(live, "CALL_SUPERVISORS", SimpleNamespace(start=started)) + request = Request({"type": "http", "method": "POST", "path": "/v1/live/sessions", "headers": []}) + stream = await live._start_supervisor(request, original, UserAPIKeyAuth(api_key="owner"), None) + assert stream.messages[0]["usage"]["seconds"] == 15 + started.assert_awaited_once() + metadata = logger.update_from_kwargs.call_args.kwargs["kwargs"]["litellm_metadata"] + assert metadata["model_info"]["id"] == "deployment-priced" + assert metadata["model_info"]["input_cost_per_second"] == 0.1 + connection.close.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_supervisor_startup_failure_closes_observer_and_invalidates_unconfirmed_hangup(monkeypatch): + from litellm.proxy.spend_tracking import budget_reservation + + logger = Mock(litellm_params={}, model_call_details={}) + monkeypatch.setattr(live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, logger))) + connection = SimpleNamespace(close=AsyncMock()) + transport = SimpleNamespace( + connect=AsyncMock(return_value=connection), + request=AsyncMock(side_effect=httpx.ConnectError("upstream unavailable")), + ) + monkeypatch.setattr(live, "LiveTransport", Mock(return_value=transport)) + monkeypatch.setattr( + live, "CALL_SUPERVISORS", SimpleNamespace(start=AsyncMock(side_effect=RuntimeError("cannot start"))) + ) + invalidate = AsyncMock() + monkeypatch.setattr(budget_reservation, "invalidate_budget_reservation_counters", invalidate) + request = Request({"type": "http", "method": "POST", "path": "/v1/live/sessions", "headers": []}) + with pytest.raises(RuntimeError, match="cannot start"): + await live._start_supervisor(request, handle(), UserAPIKeyAuth(api_key="owner"), None) + connection.close.assert_awaited_once() + invalidate.assert_awaited_once() + + +def test_admin_sip_accept_requires_exact_deployment_and_returns_owned_handle(route_client, monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy._types import LitellmUserRoles + + route_client.auth.user_role = LitellmUserRoles.PROXY_ADMIN + monkeypatch.setattr(proxy_server, "llm_model_list", [{"model_name": "voice"}]) + route_client.transport.request.return_value = httpx.Response(200, content=b"") + body = {"session": {"model": "voice", "type": "live", "instructions": "incoming"}} + missing = route_client.client.post("/v1/live/sessions/sess_incoming/accept", json=body) + assert missing.status_code == 400 + route_client.transport.request.assert_not_awaited() + accepted = route_client.client.post( + "/v1/live/sessions/sess_incoming/accept", json=body, headers={"x-litellm-live-model": "voice"} + ) + assert accepted.status_code == 200 and accepted.content == b"" + token = accepted.headers["x-litellm-live-session-id"] + owned = live.decode_session(token, live._owner(route_client.auth)) + assert owned.session_id == "sess_incoming" and owned.alias == "voice" + assert route_client.transport.request.await_args.kwargs["body"]["session"]["model"] == "gpt-live" + route_client.supervised.assert_awaited_once() + raw_hangup = route_client.client.post("/v1/live/sessions/sess_incoming/hangup") + assert raw_hangup.status_code == 403 + assert route_client.client.post(f"/v1/live/sessions/{token}/hangup").status_code == 200 + + +def test_admin_sip_cannot_enroll_through_alias_with_multiple_accounts(route_client, monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy._types import LitellmUserRoles + + route_client.auth.user_role = LitellmUserRoles.PROXY_ADMIN + monkeypatch.setattr(proxy_server, "llm_model_list", [{"model_name": "voice"}, {"model_name": "voice"}]) + response = route_client.client.post( + "/v1/live/sessions/sess_incoming/reject", json={"status_code": 603}, headers={"x-litellm-live-model": "voice"} + ) + assert response.status_code == 400 + route_client.transport.request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_primary_first_frame_timeout_releases_authenticated_reservation(monkeypatch): + import asyncio + + from fastapi import WebSocket + + messages = iter([{"type": "websocket.connect"}]) + + async def receive(): + try: + return next(messages) + except StopIteration: + raise asyncio.TimeoutError("first message timeout") + + sent = [] + + async def send(message): + sent.append(message) + + websocket = WebSocket( + { + "type": "websocket", + "path": "/v1/live/sessions", + "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")], + "scheme": "ws", + "server": ("localhost", 4000), + }, + receive, + send, + ) + monkeypatch.setattr(live, "_auth", AsyncMock(return_value=UserAPIKeyAuth(api_key="owner"))) + release = AsyncMock() + transport = Mock() + monkeypatch.setattr(live, "release_or_invalidate_budget_reservation", release) + monkeypatch.setattr(live, "LiveTransport", transport) + await live.websocket_live_session(websocket) + release.assert_awaited_once() + transport.assert_not_called() + assert sent[-1]["type"] == "websocket.close" and sent[-1]["code"] == 1008 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("key_limit,global_limit,allowed", [(None, None, True), (1, None, False), (None, 2, False)]) +async def test_legacy_limiter_only_blocks_sessions_requiring_parallel_leases( + monkeypatch, key_limit, global_limit, allowed +): + from litellm.proxy import proxy_server + from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler + + legacy = Mock(spec=_PROXY_MaxParallelRequestsHandler) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda _: legacy)) + monkeypatch.setattr(proxy_server, "general_settings", {"global_max_parallel_requests": global_limit}) + monkeypatch.setattr(live, "_authorize", AsyncMock()) + processor = AsyncMock(return_value=({"model": "voice"}, Mock())) + monkeypatch.setattr(live, "process_codex_request", processor) + monkeypatch.setattr(live, "release_or_invalidate_budget_reservation", AsyncMock()) + request = live._request(Request({"type": "http", "headers": []}), {"model": "voice"}) + auth = UserAPIKeyAuth(api_key="owner", max_parallel_requests=key_limit) + if allowed: + async with live._precall(request, auth, "voice") as prepared: + assert prepared.lease is None + processor.assert_awaited_once() + return + with pytest.raises(HTTPException) as rejected: + async with live._precall(request, auth, "voice"): + pytest.fail("Parallel-limited sessions require renewable leases") + assert rejected.value.status_code == 400 + processor.assert_not_awaited() + + +@pytest.mark.parametrize("delegation", ["invalid", {"type": "responses", "responses": "invalid"}]) +def test_malformed_delegation_returns_client_error_before_provider(route_client, delegation): + result = route_client.client.post( + "/v1/live/sessions", + json={"session": {"model": "voice", "delegation": delegation}, "transport": {"type": "webrtc", "sdp": "offer"}}, + ) + assert result.status_code == 400 + route_client.transport.request.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("parallel_reserved", [False, True]) +async def test_attachment_skips_parallel_admission_only_with_existing_session_lease(monkeypatch, parallel_reserved): + from litellm.proxy import proxy_server + from litellm.proxy.hooks.realtime_call_lease import is_realtime_call_attachment + + attachment = object() + observed = [] + + async def process(request, data, auth, model, route_type): + observed.append(is_realtime_call_attachment(attachment)) + return data, Mock() + + monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda _: None)) + monkeypatch.setattr(live, "_authorize", AsyncMock()) + monkeypatch.setattr(live, "process_codex_request", process) + monkeypatch.setattr(live, "release_or_invalidate_budget_reservation", AsyncMock()) + request = live._request(Request({"type": "http", "headers": []}), {"model": "voice"}) + async with live._precall( + request, UserAPIKeyAuth(api_key="owner"), "voice", attachment=attachment, parallel_reserved=parallel_reserved + ): + pass + assert observed == [parallel_reserved] + + +@pytest.mark.asyncio +async def test_live_observer_becomes_ready_without_session_started_event(monkeypatch): + import asyncio + + from litellm.proxy.realtime_endpoints.call_supervision import CallSupervisor + + class Observer: + def __init__(self): + self.queue = asyncio.Queue() + self.closed = False + + def __aiter__(self): + return self + + async def __anext__(self): + return await self.queue.get() + + async def close(self): + self.closed = True + + observer = Observer() + logger = Mock(litellm_params={}, model_call_details={}) + monkeypatch.setattr(live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, logger))) + sink = SimpleNamespace(store_message=Mock(), log_messages=AsyncMock()) + monkeypatch.setattr(live, "RealTimeStreaming", Mock(return_value=sink)) + + async def hangup(*args, **kwargs): + await observer.queue.put(json.dumps({"type": "session.closed", "usage": {"seconds": 0}})) + return httpx.Response(200, request=httpx.Request("POST", "https://upstream.test/hangup")) + + monkeypatch.setattr( + live, + "LiveTransport", + Mock(return_value=SimpleNamespace(connect=AsyncMock(return_value=observer), request=hangup)), + ) + supervisors = [] + + def build_supervisor(*args, **kwargs): + return CallSupervisor( + *args, **kwargs, ready_timeout=0.1, drain_timeout=0.1, termination_timeout=0.2, logging_timeout=0.2 + ) + + async def start(supervisor): + supervisors.append(supervisor) + await supervisor.start() + + monkeypatch.setattr(live, "CallSupervisor", build_supervisor) + monkeypatch.setattr(live, "CALL_SUPERVISORS", SimpleNamespace(start=start)) + request = Request({"type": "http", "method": "POST", "path": "/v1/live/sessions", "headers": []}) + try: + result = await asyncio.wait_for( + live._start_supervisor(request, handle(), UserAPIKeyAuth(api_key="owner"), None), 0.5 + ) + assert result is sink + sink.store_message.assert_not_called() + finally: + for supervisor in supervisors: + await supervisor.close() + assert observer.closed + + +@pytest.fixture +def isolated_live_model_auth(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(proxy_server, "llm_model_list", []) + monkeypatch.setattr(proxy_server, "llm_router", None) + + +@pytest.mark.asyncio +async def test_authorize_enforces_authoritative_personal_user_models(monkeypatch, isolated_live_model_auth): + user_loader = AsyncMock(return_value=LiteLLM_UserTable(user_id="user-only", models=["voice"])) + monkeypatch.setattr(live, "get_user_object", user_loader) + auth = UserAPIKeyAuth(api_key="owner", models=[], user_id="user-only") + + await live._authorize("voice", auth) + with pytest.raises(ModelAccessDeniedProxyException) as rejected: + await live._authorize("backend", auth) + assert "user can only access" in rejected.value.internal_message + + assert user_loader.await_count == 2 + + +@pytest.mark.asyncio +async def test_authorize_enforces_authoritative_organization_models(monkeypatch, isolated_live_model_auth): + org_loader = AsyncMock( + return_value=LiteLLM_OrganizationTable( + organization_id="org-only", + budget_id="budget", + created_by="admin", + updated_by="admin", + models=["voice"], + ) + ) + monkeypatch.setattr(live, "get_org_object", org_loader) + auth = UserAPIKeyAuth(api_key="owner", models=[], org_id="org-only") + + await live._authorize("voice", auth) + with pytest.raises(ModelAccessDeniedProxyException) as rejected: + await live._authorize("backend", auth) + assert "org can only access" in rejected.value.internal_message + + assert org_loader.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("scope", "auth_kwargs", "loader_name"), + [ + ("user", {"user_id": "user-only"}, "get_user_object"), + ("organization", {"org_id": "org-only"}, "get_org_object"), + ], +) +async def test_authorize_fails_closed_when_principal_grant_lookup_fails( + monkeypatch, isolated_live_model_auth, scope, auth_kwargs, loader_name +): + loader = AsyncMock(side_effect=RuntimeError(f"{scope} lookup unavailable")) + monkeypatch.setattr(live, loader_name, loader) + + with pytest.raises(HTTPException) as rejected: + await live._authorize("backend", UserAPIKeyAuth(api_key="owner", **auth_kwargs)) + + assert rejected.value.status_code == 503 + assert "verify Live" in str(rejected.value.detail) + + +def test_restricted_models_marks_user_scoped_identity_as_restricted(): + assert live._restricted_models(UserAPIKeyAuth(api_key="owner", user_id="user-only")) is True + + +@pytest.mark.asyncio +async def test_sparse_responses_update_without_model_remains_valid_for_user_scoped_identity(): + result = await live._authorize_delegation( + {"type": "session.update", "session": {"delegation": {"type": "responses", "responses": {}}}}, + UserAPIKeyAuth(api_key="owner", user_id="user-only"), + ) + assert result is None + + +def test_managed_constraints_uses_exact_metadata_keys(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "general_settings", {}) + for key in ("max_budget_alert_emails", "model_max_budget_usage"): + assert live._managed_constraints(UserAPIKeyAuth(api_key="owner", metadata={key: {"backend": 1}})) is False + + +@pytest.mark.asyncio +async def test_managed_budget_reads_authoritative_member_budget(monkeypatch): + from litellm.proxy import proxy_server + + membership = LiteLLM_TeamMembership( + user_id="member", + team_id="team", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1), + ) + + def find_unique(*, where, include): + QueryBuilder(method="find_unique", arguments={"where": where}).build_query() + return membership + + db = SimpleNamespace( + litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(side_effect=find_unique)), + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=None)), + ) + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + + assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member")) is True + assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member")) is True + db.litellm_teammembership.find_unique.assert_awaited_once() + assert cache.async_set_cache.await_args.kwargs["key"] == "team_membership:member:team" + + +@pytest.mark.asyncio +async def test_managed_budget_caches_missing_membership_sentinel(monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.user_api_key_cache import NO_TEAM_MEMBERSHIP_SENTINEL + + membership_lookup = AsyncMock(return_value=None) + db = SimpleNamespace( + litellm_teammembership=SimpleNamespace(find_unique=membership_lookup), + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=None)), + ) + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + auth = UserAPIKeyAuth(api_key="owner", team_id="missing-member-team", user_id="missing-member") + + assert await live._live_team_membership(auth) is None + assert cache.async_set_cache.await_args.kwargs["value"] == NO_TEAM_MEMBERSHIP_SENTINEL + assert await live._live_team_membership(auth) is None + membership_lookup.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_managed_budget_fails_closed_when_membership_repository_is_unreadable(monkeypatch): + from litellm.proxy import proxy_server + + failure = RuntimeError("database unavailable") + db = SimpleNamespace( + litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(side_effect=failure)), + ) + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + + with pytest.raises(HTTPException) as rejected: + await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member")) + + assert rejected.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_managed_budget_checks_project_team_and_model_group_tables(monkeypatch): + from litellm.proxy import proxy_server + + project = LiteLLM_ProjectTable( + project_id="project", + budget_id="project-budget", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1), + ) + team = LiteLLM_TeamTable(team_id="team", budget_limits=[{"budget_duration": "1d", "max_budget": 1}]) + group = SimpleNamespace( + access_group_name="group", + spend=0, + litellm_budget_table=SimpleNamespace(max_budget=1), + ) + + def find_project(*, where, include): + QueryBuilder(method="find_unique", arguments={"where": where}).build_query() + return project + + def find_groups(*, where, include): + QueryBuilder(method="find_many", arguments={"where": where}).build_query() + return [group] + + db = SimpleNamespace( + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)), + litellm_projecttable=SimpleNamespace(find_unique=AsyncMock(side_effect=find_project)), + litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(side_effect=find_groups)), + ) + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + monkeypatch.setattr(live, "collect_matched_model_access_groups", AsyncMock(return_value=("group",))) + + assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", project_id="project")) is True + assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team")) is True + assert ( + await live._managed_member_budget( + UserAPIKeyAuth(api_key="owner", matched_model_access_groups=["voice-group"]), model="backend" + ) + is True + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("backend_budget, blocked", [(None, False), (0, False), (1, True)]) +async def test_managed_budget_uses_the_delegated_model_group(monkeypatch, backend_budget, blocked): + from litellm.proxy import proxy_server + + rows = [ + SimpleNamespace(access_group_name="voice-group", spend=0, litellm_budget_table=SimpleNamespace(max_budget=1)), + SimpleNamespace( + access_group_name="backend-group", spend=0, litellm_budget_table=SimpleNamespace(max_budget=backend_budget) + ), + ] + + async def find_group_budgets(*, where, include): + return [row for row in rows if row.access_group_name in where["access_group_name"]["in"]] + + db = SimpleNamespace( + litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(side_effect=find_group_budgets)), + ) + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace()) + monkeypatch.setattr(live, "collect_matched_model_access_groups", AsyncMock(return_value=("backend-group",))) + + auth = UserAPIKeyAuth(api_key="owner", models=["voice", "backend-group"]) + assert await live._managed_member_budget(auth, model="backend") is blocked + assert await live._managed_member_budget(auth, model="backend") is blocked + db.litellm_modelaccessgroupbudgettable.find_many.assert_awaited_once() + assert cache.async_set_cache.await_args.kwargs["key"] == "live:model_access_group_limits:backend-group" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("limit_field", ["rpm_limit", "tpm_limit"]) +async def test_managed_budget_blocks_a_delegated_group_rate_limit_without_a_budget(monkeypatch, limit_field): + from litellm.proxy import proxy_server + + group_budget = SimpleNamespace(max_budget=None, rpm_limit=None, tpm_limit=None) + setattr(group_budget, limit_field, 100) + db = SimpleNamespace( + litellm_modelaccessgroupbudgettable=SimpleNamespace( + find_many=AsyncMock( + return_value=[SimpleNamespace(access_group_name="voice-group", litellm_budget_table=group_budget)] + ) + ), + ) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache()) + monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace()) + monkeypatch.setattr(live, "collect_matched_model_access_groups", AsyncMock(return_value=("voice-group",))) + + assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", models=["voice"]), model="backend") is True + + +@pytest.mark.asyncio +async def test_managed_budget_fails_closed_when_the_team_row_is_unreadable(monkeypatch): + from litellm.proxy import proxy_server + + db = SimpleNamespace( + litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(return_value=None)), + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(side_effect=RuntimeError("Database unavailable"))), + ) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache()) + monkeypatch.setattr(proxy_server, "llm_router", None) + + with pytest.raises(HTTPException) as rejected: + await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member")) + + assert rejected.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_live_team_caches_the_permission_relation_with_the_team(monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + team = LiteLLM_TeamTable(team_id="team", object_permission_id="perm-1") + db = SimpleNamespace( + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)), + litellm_objectpermissiontable=SimpleNamespace( + find_unique=AsyncMock(return_value=LiteLLM_ObjectPermissionTable(object_permission_id="perm-1")) + ), + ) + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + + loaded: Final = await live._live_team(UserAPIKeyAuth(api_key="owner", team_id="team")) + assert loaded is not None and loaded.object_permission is not None + + cached_entries: Final = [ + call.kwargs["value"] for call in cache.async_set_cache.await_args_list if call.kwargs["key"] == "team_id:team" + ] + assert len(cached_entries) == 1, "the team must be cached under the key the chat path reads" + assert cached_entries[0].object_permission is not None + + +@pytest.mark.asyncio +async def test_managed_budget_fails_closed_when_the_default_budget_is_unreadable(monkeypatch): + from litellm.proxy import proxy_server + + team = LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "budget-1"}) + db = SimpleNamespace( + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)), + litellm_budgettable=SimpleNamespace(find_unique=AsyncMock(side_effect=RuntimeError("Database unavailable"))), + ) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache()) + monkeypatch.setattr(proxy_server, "llm_router", None) + + with pytest.raises(HTTPException) as rejected: + await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member")) + + assert rejected.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_managed_budget_caches_the_team_and_its_default_budget(monkeypatch): + from litellm.proxy import proxy_server + + team = LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "budget-1"}) + budget = LiteLLM_BudgetTable(max_budget=5) + team_lookup = AsyncMock(return_value=team) + budget_lookup = AsyncMock(return_value=budget) + db = SimpleNamespace( + litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(return_value=None)), + litellm_teamtable=SimpleNamespace(find_unique=team_lookup), + litellm_budgettable=SimpleNamespace(find_unique=budget_lookup), + ) + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + monkeypatch.setattr(proxy_server, "llm_router", None) + auth = UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member") + + assert await live._managed_member_budget(auth) is True + assert await live._managed_member_budget(auth) is True + + team_lookup.assert_awaited_once() + budget_lookup.assert_awaited_once() + cached_keys: set[str] = {call.kwargs["key"] for call in cache.async_set_cache.await_args_list} + assert {"team_id:team", "team_member_default_budget:budget-1"} <= cached_keys + + +@pytest.mark.asyncio +async def test_responses_delegation_fails_closed_when_inherited_org_lookup_fails(monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy.auth import auth_checks + + team = LiteLLM_TeamTable(team_id="org-lookup-team", organization_id="org", models=["*"]) + group = SimpleNamespace( + access_group_name="backend-group", + spend=0, + litellm_budget_table=SimpleNamespace(max_budget=1), + ) + db = SimpleNamespace( + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)), + litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(return_value=[group])), + ) + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + monkeypatch.setattr( + proxy_server, + "llm_router", + SimpleNamespace(get_model_access_groups=lambda model_name, team_id=None: {"backend-group"}), + ) + monkeypatch.setattr( + auth_checks, + "get_org_object", + AsyncMock(side_effect=RuntimeError("organization lookup unavailable")), + ) + monkeypatch.setattr(live, "_authorize", AsyncMock()) + + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation( + {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}}, + UserAPIKeyAuth(api_key="owner", models=["*"], team_id="org-lookup-team"), + ) + + assert rejected.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_authorize_checks_organization_inherited_from_team(monkeypatch): + from litellm.proxy import proxy_server + + team = SimpleNamespace(organization_id="org") + org = SimpleNamespace(models=["backend"]) + team_loader = AsyncMock(return_value=team) + org_loader = AsyncMock(return_value=org) + org_check = Mock() + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr(proxy_server, "user_api_key_cache", object()) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", None) + monkeypatch.setattr(proxy_server, "llm_model_list", []) + monkeypatch.setattr(proxy_server, "llm_router", object()) + monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(live, "get_team_object", team_loader, raising=False) + monkeypatch.setattr(live, "get_org_object", org_loader) + monkeypatch.setattr(live, "can_org_access_model", org_check) + + await live._authorize("backend", UserAPIKeyAuth(api_key="owner", team_id="team")) + + team_loader.assert_awaited_once() + org_loader.assert_awaited_once() + org_check.assert_called_once_with(model="backend", org_object=org, llm_router=proxy_server.llm_router) + + +@pytest.mark.asyncio +async def test_managed_budget_fails_closed_when_member_scope_lookup_fails_after_snapshot(monkeypatch): + from litellm.proxy import proxy_server + + team = LiteLLM_TeamTable(team_id="team", models=["*"]) + membership = LiteLLM_TeamMembership( + user_id="member", + team_id="team", + litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["backend-group"]), + ) + group = SimpleNamespace( + access_group_name="backend-group", + litellm_budget_table=SimpleNamespace(max_budget=1), + ) + db = SimpleNamespace( + litellm_teammembership=SimpleNamespace( + find_unique=AsyncMock(side_effect=[membership, RuntimeError("membership lookup unavailable")]) + ), + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)), + litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(return_value=[group])), + ) + cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", None) + monkeypatch.setattr( + proxy_server, + "llm_router", + SimpleNamespace(get_model_access_groups=lambda model_name, team_id=None: {"backend-group"}), + ) + monkeypatch.setattr(live, "_authorize", AsyncMock()) + + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation( + {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}}, + UserAPIKeyAuth(api_key="owner", models=["*"], team_id="team", user_id="member"), + ) + + assert rejected.value.status_code == 503 + + +def test_json_conversion_rejects_unsupported_parent_container(monkeypatch): + class _ForeignMapping: + def validate_python(self, value): + return {1: "coerced"} + + monkeypatch.setattr(live, "_MAPPING", _ForeignMapping()) + with pytest.raises(TypeError, match="Invalid Live JSON conversion target"): + live._json_value(MappingProxyType({"nested": "value"})) + + +def test_rewrite_session_ids_serializes_non_object_events_without_touching_ids(): + assert live.rewrite_session_ids(["live", {"id": "raw"}], "raw", "public") == ["live", {"id": "raw"}] + assert live.rewrite_session_ids("public", "raw", "public") == "public" + + +def test_owner_requires_authenticated_api_key(): + with pytest.raises(HTTPException) as rejected: + live._owner(UserAPIKeyAuth()) + assert rejected.value.status_code == 403 + assert rejected.value.detail == "Live sessions require an authenticated API key" + + +def _streamed_request(chunks: list[bytes]) -> Request: + pending = list(chunks) + + async def receive(): + return {"type": "http.request", "body": pending.pop(0), "more_body": bool(pending)} + + return Request({"type": "http", "method": "POST", "headers": [], "query_string": b""}, receive=receive) + + +@pytest.mark.asyncio +async def test_body_rejects_streams_larger_than_the_offer_limit(): + with pytest.raises(HTTPException) as rejected: + await live._body(_streamed_request([b"a" * (8 * 1024 * 1024), b"b"])) + assert rejected.value.status_code == 413 + assert rejected.value.detail == "Live request exceeds the 8 MiB limit" + + +@pytest.mark.asyncio +async def test_body_rejects_json_that_is_not_an_object(): + with pytest.raises(HTTPException) as rejected: + await live._body(_streamed_request([b"[1,2]"])) + assert rejected.value.status_code == 400 + assert rejected.value.detail == "Expected a JSON object" + + +def test_session_model_requires_object_and_model(): + with pytest.raises(HTTPException) as not_object: + live._session_model({"session": "voice"}) + assert not_object.value.status_code == 400 and not_object.value.detail == "session must be a JSON object" + with pytest.raises(HTTPException) as no_model: + live._session_model({"session": {}}) + assert no_model.value.status_code == 400 and no_model.value.detail == "session.model is required" + + +@pytest.mark.asyncio +async def test_team_organization_lookup_maps_failures_to_service_unavailable(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr(proxy_server, "user_api_key_cache", object()) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", object()) + monkeypatch.setattr(live, "get_team_object", AsyncMock(side_effect=RuntimeError("database unavailable"))) + + with pytest.raises(HTTPException) as rejected: + await live._live_organization_id(UserAPIKeyAuth(api_key="owner", team_id="team")) + assert rejected.value.status_code == 503 + assert rejected.value.detail == "Could not verify Live team organization model access" + + +@pytest.mark.asyncio +async def test_direct_user_authorization_fails_closed_when_user_lookup_fails(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_model_list", []) + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr(proxy_server, "user_api_key_cache", object()) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", object()) + monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(live, "get_user_object", AsyncMock(side_effect=RuntimeError("database unavailable"))) + + with pytest.raises(HTTPException) as rejected: + await live._authorize("voice", UserAPIKeyAuth(api_key="owner", user_id="user")) + assert rejected.value.status_code == 503 + assert rejected.value.detail == "Could not verify Live user model access" + + +@pytest.mark.asyncio +async def test_direct_org_authorization_fails_closed_when_org_lookup_fails(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_model_list", []) + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr(proxy_server, "user_api_key_cache", object()) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", object()) + monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(live, "get_org_object", AsyncMock(side_effect=RuntimeError("database unavailable"))) + + with pytest.raises(HTTPException) as rejected: + await live._authorize("voice", UserAPIKeyAuth(api_key="owner", org_id="org-1")) + assert rejected.value.status_code == 503 + assert rejected.value.detail == "Could not verify Live organization model access" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "auth,lookup,expected", + [ + (UserAPIKeyAuth(api_key="owner", user_id="user"), "get_user_object", "Could not verify Live user model access"), + ( + UserAPIKeyAuth(api_key="owner", org_id="org-1"), + "get_org_object", + "Could not verify Live organization model access", + ), + ], + ids=["user", "organization"], +) +async def test_authorization_fails_closed_when_the_principal_row_is_missing(monkeypatch, auth, lookup, expected): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_model_list", []) + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr(proxy_server, "user_api_key_cache", object()) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", object()) + monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(live, "can_user_call_model", AsyncMock()) + monkeypatch.setattr(live, "can_org_access_model", Mock()) + monkeypatch.setattr(live, lookup, AsyncMock(return_value=None)) + + with pytest.raises(HTTPException) as rejected: + await live._authorize("voice", auth) + assert rejected.value.status_code == 503 + assert rejected.value.detail == expected + + +@pytest.mark.asyncio +async def test_deployment_requires_router_and_supported_provider(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_router", None) + with pytest.raises(HTTPException) as without_router: + await live._deployment("voice", {}) + assert without_router.value.status_code == 503 + assert without_router.value.detail == "Live requires a configured model deployment" + + monkeypatch.setattr( + proxy_server, + "llm_router", + SimpleNamespace( + async_get_available_deployment=AsyncMock( + return_value={ + "litellm_params": {"model": "bedrock/voice"}, + "model_info": {"id": "deployment-a"}, + } + ), + async_routing_strategy_pre_call_checks=AsyncMock(), + ), + ) + with pytest.raises(HTTPException) as wrong_provider: + await live._deployment("voice", {}) + assert wrong_provider.value.status_code == 400 + assert "OpenAI or ChatGPT" in wrong_provider.value.detail + + +@pytest.mark.parametrize("payload", [{}, {"session": {"id": 5}}, {"session": None}]) +def test_session_id_requires_upstream_string_id(payload): + with pytest.raises(HTTPException) as rejected: + live._session_id(payload) + assert rejected.value.status_code == 502 + assert rejected.value.detail == "Upstream did not return a Live session ID" + + +@pytest.mark.asyncio +async def test_live_team_membership_prefers_reservation_cache_and_sentinel(monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec + from litellm.proxy.common_utils.user_api_key_cache import NO_TEAM_MEMBERSHIP_SENTINEL + + membership = LiteLLM_TeamMembership(user_id="user", team_id="team") + auth = UserAPIKeyAuth(api_key="owner", user_id="user", team_id="team") + monkeypatch.setattr( + proxy_server, + "user_api_key_cache", + _auth_cache({"team_membership:user:team": CacheCodec.serialize(membership, model_type=LiteLLM_TeamMembership)}), + ) + restored = await live._live_team_membership(auth) + assert restored is not None and restored.user_id == "user" and restored.team_id == "team" + + monkeypatch.setattr( + proxy_server, + "user_api_key_cache", + _auth_cache({"team_membership:user:team": NO_TEAM_MEMBERSHIP_SENTINEL}), + ) + assert await live._live_team_membership(auth) is None + + +@pytest.mark.asyncio +async def test_live_team_uses_team_cache_before_database(monkeypatch): + from litellm.proxy import proxy_server + + team = SimpleNamespace(team_id="team", models=["*"]) + team_lookup = AsyncMock(side_effect=AssertionError("cache hit must not query the team table")) + monkeypatch.setattr( + proxy_server, + "prisma_client", + SimpleNamespace(db=SimpleNamespace(litellm_teamtable=SimpleNamespace(find_unique=team_lookup))), + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache({"team_id:team": team})) + assert await live._live_team(UserAPIKeyAuth(api_key="owner", team_id="team")) is team + team_lookup.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_live_team_caches_database_row_after_miss(monkeypatch): + from litellm.proxy import proxy_server + + team = LiteLLM_TeamTable(team_id="cache-miss-team") + team_lookup = AsyncMock(return_value=team) + cache = _auth_cache() + monkeypatch.setattr( + proxy_server, + "prisma_client", + SimpleNamespace(db=SimpleNamespace(litellm_teamtable=SimpleNamespace(find_unique=team_lookup))), + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + auth = UserAPIKeyAuth(api_key="owner", team_id="cache-miss-team") + + assert (await live._live_team(auth)).team_id == "cache-miss-team" + assert cache.async_set_cache.await_args.kwargs["key"] == "team_id:cache-miss-team" + assert (await live._live_team(auth)).team_id == "cache-miss-team" + team_lookup.assert_awaited_once() + + +@pytest.mark.parametrize( + "team", + [ + SimpleNamespace(budget_limits=3, rpm_limit=None, tpm_limit=None, max_budget=None, model_max_budget=None), + SimpleNamespace(budget_limits=None, rpm_limit=5, tpm_limit=None, max_budget=None, model_max_budget=None), + SimpleNamespace( + budget_limits=None, rpm_limit=None, tpm_limit=None, max_budget=None, model_max_budget={"voice": 1} + ), + ], + ids=["scalar-windows", "scalar-rpm", "model-max-budget"], +) +def test_team_budget_fields_short_circuit_before_metadata_scan(team): + assert live._live_team_budget_configured(UserAPIKeyAuth(api_key="owner"), team) is True + + +@pytest.mark.parametrize( + "value,zero_is_limit,expected", + [ + (None, False, False), + ({"max_budget": 0}, True, True), + ({"max_budget": 0}, False, False), + ({"max_budget": "unlimited"}, False, True), + ({"rpm_limit": 2}, False, True), + ], + ids=["missing", "zero-as-limit", "zero-unlimited", "non-numeric-limit", "other-limit"], +) +def test_live_budget_configured_separates_zero_from_non_numeric_limits(value, zero_is_limit, expected): + assert live._live_budget_configured(value, zero_is_limit=zero_is_limit) is expected + + +@pytest.mark.asyncio +async def test_live_default_budget_uses_cached_team_member_budget(monkeypatch): + from litellm.proxy import proxy_server + + budget = LiteLLM_BudgetTable(max_budget=1) + budget_lookup = AsyncMock(side_effect=AssertionError("cache hit must not query the budget table")) + monkeypatch.setattr( + proxy_server, + "prisma_client", + SimpleNamespace(db=SimpleNamespace(litellm_budgettable=SimpleNamespace(find_unique=budget_lookup))), + ) + monkeypatch.setattr( + proxy_server, "user_api_key_cache", _auth_cache({"team_member_default_budget:budget-1": budget}) + ) + team = SimpleNamespace(metadata={"team_member_budget_id": "budget-1"}) + auth = UserAPIKeyAuth(api_key="owner", user_id="user", team_id="team") + assert await live._live_default_budget(auth, team) is budget + budget_lookup.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_live_project_uses_cache_or_reports_missing_row(monkeypatch): + from litellm.proxy import proxy_server + + project = LiteLLM_ProjectTable(project_id="project-1") + project_lookup = AsyncMock(return_value=project) + monkeypatch.setattr( + proxy_server, + "prisma_client", + SimpleNamespace(db=SimpleNamespace(litellm_projecttable=SimpleNamespace(find_unique=project_lookup))), + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache({"project_id:project-1": project})) + auth = UserAPIKeyAuth(api_key="owner", project_id="project-1") + assert await live._live_project(auth) is project + project_lookup.assert_not_awaited() + + cache = _auth_cache() + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + assert (await live._live_project(auth)).project_id == "project-1" + assert cache.async_set_cache.await_args.kwargs["key"] == "project_id:project-1" + assert (await live._live_project(auth)).project_id == "project-1" + project_lookup.assert_awaited_once() + + project_lookup.reset_mock(return_value=True) + project_lookup.return_value = None + monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache()) + assert await live._live_project(auth) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "project, managed", + [ + ( + SimpleNamespace( + litellm_budget_table=None, + budget_id=None, + model_rpm_limit={"voice": 5}, + model_tpm_limit=None, + metadata=None, + ), + True, + ), + ( + SimpleNamespace( + litellm_budget_table=None, + budget_id=None, + model_rpm_limit=None, + model_tpm_limit=None, + metadata={"rpm_limit": 5}, + ), + True, + ), + ( + SimpleNamespace( + litellm_budget_table=None, budget_id=None, model_rpm_limit=None, model_tpm_limit=None, metadata=None + ), + False, + ), + ], + ids=["model-rate-limit", "metadata-limit", "nothing"], +) +async def test_project_budget_falls_back_to_rate_limits_and_metadata(monkeypatch, project, managed): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", object()) + assert await live._live_project_budget_configured(UserAPIKeyAuth(api_key="owner"), project) is managed + + +@pytest.mark.asyncio +async def test_model_group_budget_requires_model_name(): + assert await live._live_model_group_budget_configured(UserAPIKeyAuth(api_key="owner"), None, None, None) is False + + +@pytest.mark.asyncio +async def test_managed_member_budget_fails_closed_without_database(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(proxy_server, "llm_router", object()) + with pytest.raises(HTTPException) as rejected: + await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team"), "voice") + assert rejected.value.status_code == 503 + assert rejected.value.detail == "Could not verify Live managed budgets" + + +@pytest.mark.asyncio +async def test_backend_delegation_without_responses_contract_requires_named_model(monkeypatch): + monkeypatch.setattr(live, "_managed_member_budget", AsyncMock(return_value=False)) + + await live._authorize_delegation( + {"session": {"delegation": {"type": "backend"}}}, + UserAPIKeyAuth(api_key="owner"), + ) + + with pytest.raises(HTTPException) as rejected: + await live._authorize_delegation( + {"type": "session.start", "session": {"delegation": {"type": "responses"}}}, + UserAPIKeyAuth(api_key="owner", models=["voice"]), + ) + assert rejected.value.status_code == 400 + assert "explicit authorized delegation.responses.model" in rejected.value.detail + + +@pytest.mark.asyncio +async def test_precall_aborts_when_transferred_quota_lease_cannot_renew(monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + + lease = SimpleNamespace(start=Mock(), renew=AsyncMock(return_value=False), close=AsyncMock()) + limiter = Mock(spec=_PROXY_MaxParallelRequestsHandler_v3) + limiter.transfer_realtime_call_slot = Mock(return_value=lease) + limiter.async_post_call_failure_hook = AsyncMock() + monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda _: limiter)) + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(live, "_authorize", AsyncMock()) + monkeypatch.setattr(live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, Mock()))) + request = live._request(Request({"type": "http", "headers": []}), {"session": {"model": "voice"}}) + + with pytest.raises(HTTPException) as lost: + async with live._precall(request, UserAPIKeyAuth(api_key="owner"), "voice"): + pytest.fail("session must not start when the quota reservation is lost") + + assert lost.value.status_code == 503 and "quota reservation was lost" in lost.value.detail + lease.start.assert_called_once() + lease.close.assert_awaited_once() + limiter.async_post_call_failure_hook.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_supervise_starts_observer_under_isolated_request_stash(monkeypatch): + start = AsyncMock(return_value="stream") + monkeypatch.setattr(live, "_start_supervisor", start) + request = Request({"type": "http", "headers": []}) + auth = UserAPIKeyAuth(api_key="owner") + source = handle() + + assert await live._supervise(request, source, auth, None, None) == "stream" + assert ( + start.await_args.args[0] is request and start.await_args.args[1] is source and start.await_args.args[2] is auth + ) + + +@pytest.mark.asyncio +async def test_observer_frontend_swallows_traffic_and_hangup_checks_upstream_status(monkeypatch): + from starlette.websockets import WebSocketState + + connection = SimpleNamespace(close=AsyncMock()) + transport = SimpleNamespace( + connect=AsyncMock(return_value=connection), + request=AsyncMock( + return_value=httpx.Response( + 502, request=httpx.Request("POST", "http://upstream.test/live/sessions/sess_upstream/hangup") + ) + ), + ) + monkeypatch.setattr(live, "LiveTransport", Mock(return_value=transport)) + monkeypatch.setattr( + live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, Mock(litellm_params={}))) + ) + monkeypatch.setattr(live, "CALL_SUPERVISORS", SimpleNamespace(start=AsyncMock())) + supervisor_init = Mock() + monkeypatch.setattr(live, "CallSupervisor", supervisor_init) + stream_cls = Mock() + monkeypatch.setattr(live, "RealTimeStreaming", stream_cls) + request = Request({"type": "http", "headers": []}) + + await live._start_supervisor(request, handle(), UserAPIKeyAuth(api_key="owner"), None) + + frontend = stream_cls.call_args.args[0] + frontend.client_state = WebSocketState.CONNECTED + frontend.application_state = WebSocketState.CONNECTED + assert await frontend.receive() == {"type": "websocket.disconnect", "code": 1000} + assert await frontend.send({"type": "websocket.send", "text": "tick"}) is None + hangup = supervisor_init.call_args.args[4] + with pytest.raises(httpx.HTTPStatusError): + await hangup() + transport.request.assert_awaited_once_with("POST", "live/sessions/sess_upstream/hangup") + + +@pytest.mark.asyncio +async def test_observer_startup_failure_still_hangs_up_and_keeps_the_original_error(monkeypatch): + from litellm.proxy.spend_tracking import budget_reservation + + connection = SimpleNamespace(close=AsyncMock()) + transport = SimpleNamespace( + connect=AsyncMock(return_value=connection), + request=AsyncMock( + return_value=httpx.Response( + 200, request=httpx.Request("POST", "http://upstream.test/live/sessions/sess_upstream/hangup") + ) + ), + ) + monkeypatch.setattr(live, "LiveTransport", Mock(return_value=transport)) + monkeypatch.setattr( + live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, Mock(litellm_params={}))) + ) + monkeypatch.setattr(live, "CALL_SUPERVISORS", SimpleNamespace(start=AsyncMock())) + monkeypatch.setattr(live, "RealTimeStreaming", Mock()) + monkeypatch.setattr(live, "CallSupervisor", Mock(side_effect=RuntimeError("supervisor refused the call"))) + invalidate = AsyncMock() + monkeypatch.setattr(budget_reservation, "invalidate_budget_reservation_counters", invalidate) + request = Request({"type": "http", "headers": []}) + + with pytest.raises(RuntimeError, match="supervisor refused the call"): + await live._start_supervisor(request, handle(), UserAPIKeyAuth(api_key="owner"), None) + + transport.request.assert_awaited_once_with("POST", "live/sessions/sess_upstream/hangup") + invalidate.assert_not_awaited() + connection.close.assert_awaited_once() + + +def test_admin_sip_accept_rejects_model_mismatch_and_passes_upstream_errors_through(route_client, monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy._types import LitellmUserRoles + + route_client.auth.user_role = LitellmUserRoles.PROXY_ADMIN + monkeypatch.setattr(proxy_server, "llm_model_list", [{"model_name": "voice"}]) + mismatch = route_client.client.post( + "/v1/live/sessions/sess_incoming/accept", + json={"session": {"model": "other", "type": "live"}}, + headers={"x-litellm-live-model": "voice"}, + ) + assert mismatch.status_code == 400 and "must match" in mismatch.json()["detail"] + route_client.transport.request.assert_not_awaited() + + route_client.transport.request.return_value = httpx.Response(503, json={"error": "gateway down"}) + failed = route_client.client.post( + "/v1/live/sessions/sess_incoming/accept", + json={"session": {"model": "voice", "type": "live"}}, + headers={"x-litellm-live-model": "voice"}, + ) + assert failed.status_code == 503 and "x-litellm-live-session-id" not in failed.headers + route_client.supervised.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_public_sideband_rejects_restart_and_model_change_then_rewrites_ids(): + websocket = SimpleNamespace( + receive_text=AsyncMock(return_value=json.dumps({"type": "session.start"})), + send_text=AsyncMock(), + close=AsyncMock(), + scope={}, + headers=Mock(), + ) + public = live._PublicSocket(websocket, handle(), "public", UserAPIKeyAuth(api_key="owner")) + + with pytest.raises(HTTPException) as restart: + await public.receive_text() + assert restart.value.status_code == 400 and restart.value.detail == "Session has already started" + + websocket.receive_text.return_value = json.dumps({"type": "session.update", "session": {"model": "other"}}) + with pytest.raises(HTTPException) as model: + await public.receive_text() + assert model.value.status_code == 400 and model.value.detail == "Session model cannot change" + + websocket.receive_text.return_value = json.dumps({"type": "custom", "session_id": "public"}) + assert json.loads(await public.receive_text()) == {"type": "custom", "session_id": "sess_upstream"} + + +def test_startup_events_overflow_fails_closed(): + events = live._StartupEvents() + for _ in range(128): + events.store({"type": "info"}) + with pytest.raises(HTTPException) as overflowed: + events.store({"type": "info"}) + assert overflowed.value.status_code == 502 + + +def test_websocket_requires_api_key_then_session_start(route_client): + from starlette.websockets import WebSocketDisconnect + + with pytest.raises(WebSocketDisconnect) as anonymous: + with route_client.client.websocket_connect("/v1/live/sessions") as ws: + ws.receive_json() + assert anonymous.value.code == 1008 + + with route_client.client.websocket_connect("/v1/live/sessions", headers={"Authorization": "Bearer owner"}) as ws: + ws.send_json({"type": "ping"}) + with pytest.raises(WebSocketDisconnect) as wrong_first: + ws.receive_json() + assert wrong_first.value.code == 1008 + route_client.transport.request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_attached_socket_authorizes_policy_and_reuses_source_without_start(route_client, monkeypatch): + from starlette.websockets import WebSocket + + streams: list = [] + + class AttachedStream: + def __init__(self, *args, **kwargs): + self.args = args + self.bidirectional_forward = AsyncMock() + streams.append(self) + + backend = SimpleNamespace(send=AsyncMock(), close=AsyncMock(), recv=AsyncMock()) + route_client.transport.connect = AsyncMock(return_value=backend) + monkeypatch.setattr(live, "RealTimeStreaming", AttachedStream) + authorize = AsyncMock() + monkeypatch.setattr(live, "_authorize_delegation", authorize) + token = live.encode_session(handle(model_id="deployment-a")) + inbound = iter([{"type": "websocket.connect"}, {"type": "websocket.disconnect", "code": 1000}]) + sent: list = [] + + async def receive(): + return next(inbound) + + async def send(message): + sent.append(message) + + websocket = WebSocket( + { + "type": "websocket", + "path": f"/v1/live/sessions/{token}/attach", + "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")], + "scheme": "ws", + "server": ("testserver", 80), + "client": ("testclient", 50000), + "subprotocols": [], + }, + receive, + send, + ) + + await live.websocket_live_session(websocket, token) + + assert sent == [{"type": "websocket.accept", "subprotocol": None, "headers": []}] + authorize.assert_awaited_once() + assert authorize.await_args.args[1] is route_client.auth + route_client.transport.connect.assert_awaited_once_with("live/sessions/sess_upstream/attach") + backend.send.assert_not_awaited() + backend.close.assert_awaited_once() + frontend = streams[0].args[0] + assert frontend.public_id == token and frontend.handle.session_id == "sess_upstream" and frontend.observer is None + route_client.supervised.assert_not_awaited() + streams[0].bidirectional_forward.assert_awaited_once() + + +def test_websocket_connection_failure_closes_with_internal_error(route_client): + from starlette.websockets import WebSocketDisconnect + + route_client.transport.connect = AsyncMock(return_value=None) + with route_client.client.websocket_connect("/v1/live/sessions", headers={"Authorization": "Bearer owner"}) as ws: + ws.send_json({"type": "session.start", "session": {"model": "voice"}}) + with pytest.raises(WebSocketDisconnect) as internal: + ws.receive_json() + assert internal.value.code == 1011 + route_client.transport.request.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stage", ["rejected", "crashed"]) +async def test_websocket_close_after_asgi_completion_is_swallowed(monkeypatch, stage): + from fastapi import WebSocket + + class CompletedWebSocket(WebSocket): + async def close(self, code=1000, reason=None): + raise RuntimeError("ASGI send channel already completed") + + headers = [] if stage == "rejected" else [(b"authorization", b"Bearer owner")] + sent = [] + + async def receive(): + return {"type": "websocket.disconnect"} + + async def send(message): + sent.append(message) + + websocket = CompletedWebSocket( + { + "type": "websocket", + "path": "/v1/live/sessions", + "query_string": b"", + "headers": headers, + "scheme": "ws", + "server": ("localhost", 4000), + }, + receive, + send, + ) + if stage == "crashed": + monkeypatch.setattr(live, "_auth", AsyncMock(side_effect=ConnectionError("redis down"))) + + await live.websocket_live_session(websocket) + + assert sent == [] diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py index f5c97142dde..e835cea5f00 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py @@ -7,7 +7,7 @@ Tests for LiteLLM proxy realtime WebRTC HTTP endpoints: import json import time from collections.abc import Awaitable -from typing import Protocol +from typing import Final, Protocol from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -15,7 +15,6 @@ import pytest from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient - from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.encrypt_decrypt_utils import ( @@ -120,6 +119,114 @@ def proxy_app(monkeypatch): return proxy_server.app +@pytest.mark.parametrize("path", ["/v1/live", "/live", "/openai/v1/live"]) +@pytest.mark.parametrize("query", ["", "?intent=custom&architecture=custom"]) +def test_live_multipart_offer_runs_authenticated_call_pipeline( + proxy_app: FastAPI, monkeypatch: pytest.MonkeyPatch, path: str, query: str +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.realtime_endpoints import call_sessions + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-live-signaling-salt") + session: Final = {"model": "gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}} + authenticate: Final = AsyncMock(wraps=call_sessions.user_api_key_auth) + process: Final = AsyncMock(side_effect=lambda request, data, *args: (data, MagicMock())) + supervise: Final = AsyncMock() + upstream: Final = httpx.Response( + 201, + content=b"v=0\r\nanswer", + headers={"Location": "/v1/live/rtc_private"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}, + ) + + async def route(**kwargs: object) -> Awaitable[httpx.Response]: + assert kwargs["route_type"] == "arealtime_calls" + assert isinstance(kwargs["data"], dict) + assert kwargs["data"]["sdp_body"] == b"v=0\r\noffer" + assert kwargs["data"]["session"] == session + assert kwargs["data"]["chatgpt_realtime_client_query"] == ( + {"architecture": "custom", "intent": "custom"} + if query + else {"architecture": "avas", "intent": "quicksilver"} + ) + + async def respond() -> httpx.Response: + return upstream + + return respond() + + monkeypatch.setattr(call_sessions, "user_api_key_auth", authenticate) + monkeypatch.setattr(call_sessions, "process_codex_request", process) + monkeypatch.setattr(call_sessions, "supervise_codex_call", supervise) + monkeypatch.setattr(proxy_server, "route_request", route) + response: Final = TestClient(proxy_app).post( + f"{path}{query}", + headers={"Authorization": "Bearer sk-test-master-key"}, + files={ + "sdp": (None, "v=0\r\noffer", "application/sdp"), + "session": (None, json.dumps(session), "application/json"), + }, + ) + + assert response.status_code == 201 + assert response.content == b"v=0\r\nanswer" + assert response.headers["content-type"] == "application/sdp" + authenticate.assert_awaited_once() + process.assert_awaited_once() + assert process.await_args.args[3:] == ("gpt-live-1-codex", "arealtime_calls") + assert isinstance(process.await_args.args[2], UserAPIKeyAuth) + supervise.assert_awaited_once() + assert response.headers["location"].startswith("/v1/live/") + token: Final = response.headers["location"].rsplit("/", 1)[-1] + call: Final = call_sessions.decode_call(token, "Bearer sk-test-master-key") + assert call.call_id == "rtc_private" + assert call.alias == "gpt-live-1-codex" + assert call.usage_supervised + + +@pytest.mark.parametrize("path", ["/v1/live", "/live", "/openai/v1/live"]) +def test_live_multipart_offer_rejects_invalid_credentials_before_routing( + proxy_app: FastAPI, monkeypatch: pytest.MonkeyPatch, path: str +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.realtime_endpoints import call_sessions + + authenticate: Final = AsyncMock(wraps=call_sessions.user_api_key_auth) + route: Final = AsyncMock() + monkeypatch.setattr(call_sessions, "user_api_key_auth", authenticate) + monkeypatch.setattr(proxy_server, "route_request", route) + response: Final = TestClient(proxy_app).post( + path, + files={"sdp": (None, "v=0\r\n"), "session": (None, '{"model":"gpt-live-1-codex"}')}, + ) + + assert response.status_code == 401 + authenticate.assert_awaited_once() + route.assert_not_awaited() + + +def test_live_multipart_offer_rejects_model_outside_key_scope( + proxy_app: FastAPI, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.realtime_endpoints import call_sessions + + authenticate: Final = AsyncMock(return_value=UserAPIKeyAuth(models=["another-model"])) + route: Final = AsyncMock() + monkeypatch.setattr(call_sessions, "user_api_key_auth", authenticate) + monkeypatch.setattr(proxy_server, "route_request", route) + response: Final = TestClient(proxy_app).post( + "/v1/live", + headers={"Authorization": "Bearer restricted-key"}, + files={"sdp": (None, "v=0\r\n"), "session": (None, '{"model":"gpt-live-1-codex"}')}, + ) + + assert response.status_code == 403 + assert "gpt-live-1-codex" in response.text + authenticate.assert_awaited_once() + route.assert_not_awaited() + + @pytest.fixture def mock_route_request_client_secrets(): """Mock route_request to return a fake upstream client_secrets response.""" @@ -414,12 +521,109 @@ def test_realtime_calls_invalid_token_returns_401(proxy_app): assert "Invalid or expired token" in response.json().get("error", "") +@pytest.mark.parametrize("handle_kind", ["codex", "live"]) +@pytest.mark.parametrize("provider", ["chatgpt", "openai"]) +def test_legacy_sdp_rejects_handle_ciphertext_before_oauth_dispatch( + proxy_app, monkeypatch, tmp_path, handle_kind, provider +): + import base64 + import hashlib + + from litellm import Router + from litellm.llms.chatgpt.codex import CodexRealtimeCall + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy import proxy_server + from litellm.proxy.realtime_endpoints.call_sessions import encode_call + from litellm.proxy.realtime_endpoints.live import LiveHandle, encode_session + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-sdp-handle-salt") + monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path)) + monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json") + (tmp_path / "auth.json").write_text( + json.dumps( + {"access_token": "test-sdp-oauth", "account_id": "test-sdp-account", "expires_at": time.time() + 3600} + ) + ) + owner = hashlib.sha256(b"Bearer restricted-key", usedforsecurity=False).hexdigest() + handle = ( + encode_call( + CodexRealtimeCall( + call_id="rtc_allowed", + model="gpt-live-1-codex", + alias="allowed-voice", + owner=owner, + expires_at=time.time() + 3600, + ) + ) + if handle_kind == "codex" + else encode_session( + LiveHandle( + session_id="live_allowed", + alias="allowed-voice", + deployment={"model": "chatgpt/gpt-live-1-codex"}, + owner=owner, + expires_at=time.time() + 3600, + policy={}, + ) + ) + ) + encoded = handle.removeprefix("rtc_litellm_").removeprefix("live_litellm_") + ciphertext = base64.urlsafe_b64decode(encoded + "=" * (-len(encoded) % 4)).decode() + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(201, text="v=0\r\nanswer") + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + router = Router( + model_list=[ + { + "model_name": "forbidden-voice", + "litellm_params": { + "model": "chatgpt/gpt-live-1-codex" if provider == "chatgpt" else "openai/gpt-realtime" + }, + } + ] + ) + + async def add_data(data, **kwargs): + return data + + async def pre_call(user_api_key_dict, data, call_type): + return data + + async def route(data, route_type, **kwargs): + assert route_type == "arealtime_calls" + return router.arealtime_calls(**data, client=client) + + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "user_model", None) + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", add_data) + monkeypatch.setattr(proxy_server, "route_request", route) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock(pre_call_hook=AsyncMock(side_effect=pre_call), post_call_failure_hook=AsyncMock()), + ) + response = TestClient(proxy_app).post( + "/v1/realtime/calls?model=forbidden-voice", + headers={"Authorization": f"Bearer {ciphertext}", "Content-Type": "application/sdp"}, + content=b"v=0\r\noffer", + ) + assert response.status_code == 401, [(r.url.path, r.headers.get("authorization")) for r in requests] + assert not requests + + @pytest.mark.asyncio +@pytest.mark.parametrize("token_format", ["versioned", "legacy"]) async def test_realtime_calls_success_with_valid_encrypted_token( proxy_app, mock_route_request_realtime_calls, mock_add_litellm_data, mock_pre_call_hook, + token_format, ): """POST /v1/realtime/calls returns 201 with valid encrypted token from client_secrets.""" # Build a valid encrypted token (same format as client_secrets returns) @@ -431,7 +635,7 @@ async def test_realtime_calls_success_with_valid_encrypted_token( team_id=None, expires_at=future_expires_at, ) - encrypted_token = encrypt_value_helper(token_payload) + encrypted_token = encrypt_value_helper(token_payload if token_format == "versioned" else "fake_upstream_epk") client = TestClient(proxy_app) with ( diff --git a/tests/test_litellm/proxy/test_live_route_registration.py b/tests/test_litellm/proxy/test_live_route_registration.py new file mode 100644 index 00000000000..30d58d8a6e2 --- /dev/null +++ b/tests/test_litellm/proxy/test_live_route_registration.py @@ -0,0 +1,64 @@ +from unittest.mock import AsyncMock + +import httpx +import pytest +from fastapi import HTTPException +from fastapi.testclient import TestClient + + +@pytest.mark.asyncio +@pytest.mark.parametrize("prefix", ["/live", "/v1/live", "/openai/v1/live"]) +@pytest.mark.parametrize( + ("method", "suffix"), + [ + ("POST", ""), + ("POST", "/opaque/fork"), + ("POST", "/opaque/accept"), + ("POST", "/opaque/reject"), + ("POST", "/opaque/refer"), + ("POST", "/opaque/hangup"), + ("GET", "/opaque/content"), + ], +) +async def test_public_live_http_routes_reach_live_auth_before_generic_passthrough(monkeypatch, prefix, method, suffix): + from litellm.proxy import proxy_server + from litellm.proxy.realtime_endpoints import live + + authenticate = AsyncMock(side_effect=HTTPException(401, "Live authentication required")) + monkeypatch.setattr(live, "_auth", authenticate) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=proxy_server.app), base_url="http://proxy" + ) as client: + response = await client.request( + method, + prefix + "/sessions" + suffix, + json={"session": {"model": "voice"}, "transport": {"type": "webrtc", "sdp": "offer"}}, + ) + assert response.status_code == 401 + assert "Live authentication required" in response.text + authenticate.assert_awaited_once() + + +@pytest.mark.parametrize("prefix", ["/live", "/v1/live", "/openai/v1/live"]) +@pytest.mark.parametrize("suffix", ["", "/opaque/attach", "/opaque/fork"]) +def test_public_live_websockets_reach_live_auth_before_legacy_sideband(monkeypatch, prefix, suffix): + from starlette.websockets import WebSocketDisconnect + + from litellm.proxy import proxy_server + from litellm.proxy.realtime_endpoints import live + + authenticate = AsyncMock(side_effect=HTTPException(403, "Live authentication rejected")) + monkeypatch.setattr(live, "_auth", authenticate) + monkeypatch.setattr(proxy_server, "general_settings", {}) + # TestClient runs the proxy lifespan, and the boot check refuses a weak or unset master key + # before the app serves anything. Set a safe key here instead of relying on the ambient one, so + # the request really reaches the routes: with a key in place the legacy sideband dependency + # would reject the "Bearer test" header, so awaiting _auth still proves which auth ran first. + monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-live-route-registration-test-master-key") + # A previous proxy test may leave the module scheduler bound to a closed loop. + monkeypatch.setattr(proxy_server, "scheduler", None) + with TestClient(proxy_server.app) as client: + with pytest.raises(WebSocketDisconnect): + with client.websocket_connect(prefix + "/sessions" + suffix, headers={"authorization": "Bearer test"}): + pass + authenticate.assert_awaited_once() diff --git a/tests/test_litellm/proxy/test_native_compaction.py b/tests/test_litellm/proxy/test_native_compaction.py index 778c1715530..35662285f70 100644 --- a/tests/test_litellm/proxy/test_native_compaction.py +++ b/tests/test_litellm/proxy/test_native_compaction.py @@ -121,6 +121,11 @@ async def test_real_proxy_child_auth_privacy_and_body_policy( monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) monkeypatch.setattr(proxy_server, "llm_router", None) monkeypatch.setattr(proxy_server, "general_settings", {}) + # Pin the proxy-wide budget: authentication only reads the global spend when a + # proxy max budget is configured, and that read goes through the stub prisma + # client above. A budget left set by an earlier test on the same worker would + # turn this fixture's child requests into 401s. + monkeypatch.setattr(litellm, "max_budget", 0.0) monkeypatch.setattr(common_request_processing, "route_request", route) with inherit_message_logging_privacy(True): call: Final = with_proxy_compaction_executor( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 5dfd2f57ca6..7be2ad6e587 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3299,7 +3299,10 @@ async def test_add_proxy_budget_to_db_only_creates_user_no_keys(monkeypatch: pyt import litellm from litellm.proxy.proxy_server import ProxyStartupEvent - # Set up required litellm settings + # Set up required litellm settings. Through monkeypatch rather than plain + # assignment: `litellm.max_budget` is process-global, and any later test on + # this worker that authenticates reads the global proxy spend whenever a + # proxy budget is set, which needs a real prisma client. monkeypatch.setattr(litellm, "budget_duration", "30d") monkeypatch.setattr(litellm, "max_budget", 100.0) @@ -14810,6 +14813,31 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the ) +@pytest.mark.asyncio +async def test_login_throttle_config_settings_override_database(monkeypatch): + import litellm.proxy.proxy_server as ps + from litellm.proxy.proxy_server import ProxyConfig + + config = ProxyConfig() + configured = { + "max_failed_login_attempts_per_source": 5, + "failed_login_window_seconds": 60, + "failed_login_block_seconds": 120, + } + config.settings.load_yaml(configured) + monkeypatch.setattr(ps, "general_settings", config.settings) + await config._update_general_settings( + db_general_settings={ + "max_failed_login_attempts_per_source": 999, + "failed_login_window_seconds": 1, + "failed_login_block_seconds": 1, + } + ) + for key, value in configured.items(): + assert ps.general_settings[key] == value + assert config.settings.source(key) == "config" + + @pytest.mark.asyncio async def test_login_throttle_limits_from_the_config_file_outrank_the_database(monkeypatch): import litellm.proxy.proxy_server as ps diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 0a095183b6e..b83a6575106 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -1,6 +1,7 @@ import datetime as real_datetime import smtplib from typing import Final +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException @@ -11,15 +12,16 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.bug_report import ISSUE_URL_BASE from litellm.proxy._types import ProxyErrorTypes, UserAPIKeyAuth -from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy +from litellm.proxy.utils import ( + PrismaClient, + ProxyLogging, + get_custom_url, + handle_exception_on_proxy, + join_paths, +) from litellm.types.guardrails import GuardrailEventHooks -from unittest.mock import AsyncMock, 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/") @@ -2153,6 +2155,90 @@ async def test_proxy_only_error_5xx_keeps_traceback_and_runs_sync_callbacks(monk assert "test_proxy_utils" in captured["async_traceback"] +def test_create_model_info_response_resolves_alias_to_deployment_model(): + """A public model name that is not itself a cost-map key must not be resolved through + the fallback-generalization rules: `bedrock-claude-opus-5` matches the generic + claude-family baseline (200k/64k) by substring, while the deployment it fronts really + accepts 1M/128k. Regression for the /v1/models alias resolution introduced in v1.94.0.""" + from litellm import Router + + saved_model_cost = dict(litellm.model_cost) + try: + router = Router( + model_list=[ + { + "model_name": "bedrock-claude-opus-5", + "litellm_params": { + "custom_llm_provider": "bedrock", + "model": "bedrock/eu.anthropic.claude-opus-5", + }, + "model_info": {"base_model": "eu.anthropic.claude-opus-5"}, + } + ] + ) + + response = create_model_info_response(model_id="bedrock-claude-opus-5", provider="openai", llm_router=router) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(saved_model_cost) + + assert response["max_input_tokens"] == 1000000 + assert response["max_output_tokens"] == 128000 + + +def test_create_model_info_response_keeps_exact_alias_over_generalized_deployment_model(): + """Mirror of the alias bug: when the deployment points at a custom backend name that + only matches a generalization rule, the listed name's exact cost-map entry is the + better answer and must win.""" + from litellm import Router + + saved_model_cost = dict(litellm.model_cost) + try: + router = Router( + model_list=[ + { + "model_name": "claude-opus-5", + "litellm_params": { + "custom_llm_provider": "bedrock", + "model": "bedrock/my-claude-opus-5-provisioned", + }, + } + ] + ) + + response = create_model_info_response(model_id="claude-opus-5", provider="openai", llm_router=router) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(saved_model_cost) + + assert response["max_input_tokens"] == 1000000 + + +def test_create_model_info_response_falls_back_to_alias_for_opaque_deployment_name(): + """An Azure deployment named after the resource rather than the model has no cost-map + entry; the listed name still does, and must keep answering.""" + from litellm import Router + + saved_model_cost = dict(litellm.model_cost) + try: + router = Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "azure/my-gpt4o-deployment"}, + } + ] + ) + + response = create_model_info_response(model_id="gpt-4o", provider="openai", llm_router=router) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(saved_model_cost) + + assert response["max_input_tokens"] == 128000 + assert response["max_output_tokens"] == 16384 + + def test_create_model_info_response_resolves_mode_through_deployment_model(): """`mode` is derived from the same lookup, so an aliased embedding deployment currently reports no mode at all; it must report `embedding`.""" @@ -2169,9 +2255,7 @@ def test_create_model_info_response_resolves_mode_through_deployment_model(): ] ) - response = create_model_info_response( - model_id="my-embeddings", provider="openai", llm_router=router - ) + response = create_model_info_response(model_id="my-embeddings", provider="openai", llm_router=router) finally: litellm.model_cost.clear() litellm.model_cost.update(saved_model_cost) @@ -2289,7 +2373,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, ) @@ -2299,6 +2385,130 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp 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} + + class TestPrismaClientTokenAuthBehindThePool: """Behind the in-container pool the supervisor renews the writer's database token and hands the workers a loopback URL with a static password, so the diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index fc750d88e42..148b282c558 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -18,16 +18,27 @@ async def test_trace_reader_projects_connection_and_parameters(recording_server: recording_server.enqueue(ResponseSpec(body={"data": [{"trace_id": "trace-1"}]})) reader_url: Final = recording_server.base_url.replace("http://", "http://reader:p%40ss%2Fword%25@") storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, reader_url + "?database=wrong") - rows: Final = json.loads(await storage.query("trace_spans", {"trace_id": "trace-1"})) + response: Final = json.loads( + await storage.query("trace_spans", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""}) + ) request: Final = recording_server.requests[0] - parameters: Final = parse_qs(urlsplit(request.path).query) - assert rows == {"data": [{"trace_id": "trace-1"}]} + parameters: Final = parse_qs(urlsplit(request.path).query, keep_blank_values=True) + assert response == {"data": [{"trace_id": "trace-1"}]} assert b"o.TraceId = {trace_id:String}" in request.raw_body - assert parameters["database"] == ["trace_test"] - assert parameters["param_trace_id"] == ["trace-1"] - assert parameters["readonly"] == ["1"] - assert "user" not in parameters - assert "password" not in parameters + assert b"trace-1" not in request.raw_body + assert parameters == { + "database": ["trace_test"], + "param_trace_id": ["trace-1"], + "param_team_ids": ["[]"], + "param_api_key_hash": [""], + "param_trace_ref": [""], + "readonly": ["1"], + "default_format": ["JSON"], + "max_execution_time": ["10"], + "max_result_rows": ["1000"], + "result_overflow_mode": ["throw"], + "wait_end_of_query": ["1"], + } assert request.headers["authorization"] == "Basic " + base64.b64encode(b"reader:p@ss/word%").decode() @@ -36,7 +47,7 @@ async def test_trace_reader_rejects_success_status_with_embedded_error(recording recording_server.enqueue(ResponseSpec(body={"data": [], "exception": "query failed"})) storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, recording_server.base_url) with pytest.raises(RuntimeError, match="invalid or failed JSON"): - await storage.query("trace_spans", {}) + await storage.query("trace_spans", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""}) @pytest.mark.asyncio @@ -61,7 +72,9 @@ async def test_schema_binding_rejects_non_positive_retention() -> None: @pytest.mark.asyncio -async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement(recording_server: RecordingServer) -> None: +async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement( + recording_server: RecordingServer, +) -> None: recording_server.expected_requests = 2 recording_server.enqueue(ResponseSpec(body="")) recording_server.enqueue(ResponseSpec(status=403, body="denied")) @@ -73,9 +86,10 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement assert recording_server.requests[0].raw_body.startswith(b"CREATE DATABASE IF NOT EXISTS") assert recording_server.requests[1].raw_body.startswith(b"CREATE TABLE IF NOT EXISTS") assert "readonly" not in parse_qs(urlsplit(recording_server.requests[0].path).query) - assert recording_server.requests[0].headers["authorization"] == "Basic " + base64.b64encode( - b"writer:p@ss/word%" - ).decode() + assert ( + recording_server.requests[0].headers["authorization"] + == "Basic " + base64.b64encode(b"writer:p@ss/word%").decode() + ) @pytest.mark.asyncio @@ -87,11 +101,14 @@ async def test_insert_encodes_and_sends_rows(recording_server: RecordingServer) after: Final = time.time_ns() // 1_000_000 request: Final = recording_server.requests[0] row: Final = json.loads(gzip.decompress(request.raw_body)) + assert type(row["EngineReceivedMs"]) is int assert before <= row["EngineReceivedMs"] <= after assert row == { "Input": "hello", "Timestamp": "1970-01-01T00:00:01.23456789Z", "EngineReceivedMs": row["EngineReceivedMs"], } - assert parse_qs(urlsplit(request.path).query)["query"] == ["INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow"] + assert parse_qs(urlsplit(request.path).query)["query"] == [ + "INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow" + ] assert request.headers["content-encoding"] == "gzip" diff --git a/tests/unit/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py index 247d02298aa..b373fd95b28 100644 --- a/tests/unit/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -10,7 +10,7 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v import json import os import re -from typing import Any, Dict +from typing import Any, Dict, Final from unittest.mock import MagicMock, patch import httpx @@ -33,30 +33,10 @@ def _load_openapi_spec_dict() -> Dict[str, Any]: return response.json() except Exception as e: # pragma: no cover - defensive, env-dependent pytest.skip( - f"Skipping Google Interactions OpenAPI compliance tests - " - f"unable to load spec from {OPENAPI_SPEC_URL}: {e}" + f"Skipping Google Interactions OpenAPI compliance tests - unable to load spec from {OPENAPI_SPEC_URL}: {e}" ) -def _model_create_request_schema(spec_dict: Dict[str, Any]) -> Dict[str, Any]: - schemas = spec_dict["components"]["schemas"] - create_path = next(path for path in spec_dict["paths"] if path.endswith("/interactions")) - body_schema = spec_dict["paths"][create_path]["post"]["requestBody"]["content"]["application/json"]["schema"] - variants = [schemas[option["$ref"].split("/")[-1]] for option in body_schema.get("oneOf", []) if "$ref" in option] - return next(variant for variant in variants if "model" in variant.get("properties", {})) - - -def _interaction_resource_path(spec_dict: Dict[str, Any], method: str) -> str | None: - return next( - ( - path - for path, methods in spec_dict["paths"].items() - if re.search(r"/interactions/\{[^}]+\}$", path) and method in methods - ), - None, - ) - - def _declared_type_value(variant_schema: Dict[str, Any]) -> Any: """The single `type` value a union variant pins, whether spelled as a const or a 1-item enum.""" type_property = variant_schema.get("properties", {}).get("type", {}) @@ -64,6 +44,56 @@ def _declared_type_value(variant_schema: Dict[str, Any]) -> Any: return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None) +def _resolve_local_ref(spec_dict: dict[str, Any], schema: dict[str, Any]) -> dict[str, Any]: + """Resolve component references used by operations, schemas, and parameters.""" + if "$ref" not in schema: + return schema + reference: Final = schema["$ref"] + assert reference.startswith("#/components/"), f"Expected a local component reference: {reference}" + category, name = reference.removeprefix("#/components/").split("/") + return spec_dict["components"][category][name.replace("~1", "/").replace("~0", "~")] + + +def _interaction_operation( + spec_dict: dict[str, Any], method: str, *, individual: bool = False +) -> tuple[str, dict[str, Any]]: + """Match collection or item routes exactly, independent of placeholder names.""" + pattern: Final = r"(?:/[^/]+)*/interactions" + (r"/(\{[^/{}]+\})" if individual else "") + matches: Final = tuple( + (path, path_item, match) + for path, path_item in spec_dict["paths"].items() + if (match := re.fullmatch(pattern, path)) and method in path_item + ) + assert len(matches) == 1, f"Expected one {method.upper()} interactions endpoint, got {matches}" + path, path_item, match = matches[0] + operation: Final = path_item[method] + if individual: + parameter_name: Final = match.group(1)[1:-1] + parameters: Final = { + (parameter["name"], parameter["in"]): parameter + for raw_parameter in (*path_item.get("parameters", ()), *operation.get("parameters", ())) + for parameter in (_resolve_local_ref(spec_dict, raw_parameter),) + } + parameter: Final = parameters.get((parameter_name, "path")) + assert parameter is not None, f"{path} must declare its interaction ID path parameter" + assert parameter.get("required") is True, f"{path} must require its interaction ID" + parameter_schema: Final = _resolve_local_ref(spec_dict, parameter["schema"]) + assert parameter_schema.get("type") == "string", f"{path} must accept a string interaction ID" + return path, operation + + +def _model_request_schema(spec_dict: dict[str, Any]) -> dict[str, Any]: + """Find the model variant of the JSON body declared by the create operation.""" + _, operation = _interaction_operation(spec_dict, "post") + request_body: Final = _resolve_local_ref(spec_dict, operation["requestBody"]) + assert request_body.get("required") is True, "Creating an interaction must require a request body" + schema: Final = _resolve_local_ref(spec_dict, request_body["content"]["application/json"]["schema"]) + variants: Final = tuple(_resolve_local_ref(spec_dict, variant) for variant in schema.get("oneOf", (schema,))) + model_variants: Final = tuple(variant for variant in variants if "model" in variant.get("properties", {})) + assert len(model_variants) == 1, f"Expected one model request variant, got {model_variants}" + return model_variants[0] + + @pytest.fixture(scope="module") def spec_dict() -> Dict[str, Any]: """Load raw spec dict for manual validation.""" @@ -80,10 +110,14 @@ class TestRequestCompliance: """Tests that our request bodies match the OpenAPI spec.""" def test_create_model_interaction_request_schema(self, spec_dict): - schema = _model_create_request_schema(spec_dict) + """Verify the model request schema declared by POST /interactions.""" + schema = _model_request_schema(spec_dict) assert "model" in schema["required"] - assert "input" in schema["properties"] + for field in ("model", "input"): + assert field in schema["properties"] + assert schema["properties"][field].get("readOnly") is not True + assert _resolve_local_ref(spec_dict, schema["properties"][field]).get("readOnly") is not True # Check our supported optional fields exist in spec our_optional_fields = [ @@ -106,13 +140,8 @@ class TestRequestCompliance: def test_input_types_match_spec(self, spec_dict): """Verify input field supports string, Content, Content[], Turn[].""" - schema = _model_create_request_schema(spec_dict) - input_schema = schema["properties"]["input"] - - # The input property may be inline oneOf or a $ref to InteractionsInput - if "$ref" in input_schema: - ref_name = input_schema["$ref"].split("/")[-1] - input_schema = spec_dict["components"]["schemas"][ref_name] + schema = _model_request_schema(spec_dict) + input_schema = _resolve_local_ref(spec_dict, schema["properties"]["input"]) # Should be oneOf with multiple types assert "oneOf" in input_schema @@ -143,22 +172,18 @@ class TestRequestCompliance: discriminator = content_schema.get("discriminator") if discriminator is not None: - assert ( - discriminator.get("propertyName") == "type" - ), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'" + assert discriminator.get("propertyName") == "type", ( + f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'" + ) variant_names = [ - option["$ref"].split("/")[-1] - for option in content_schema.get("oneOf", []) - if "$ref" in option + option["$ref"].split("/")[-1] for option in content_schema.get("oneOf", []) if "$ref" in option ] assert variant_names, f"Content is not a union of named variants: {content_schema}" mapping = (discriminator or {}).get("mapping") or {} type_values = { - variant: mapping_value - for mapping_value, ref in mapping.items() - for variant in [ref.split("/")[-1]] + variant: mapping_value for mapping_value, ref in mapping.items() for variant in [ref.split("/")[-1]] } or { variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {})) for variant in variant_names @@ -209,7 +234,9 @@ class TestRequestCompliance: for option in spec_dict["components"]["schemas"]["Step"]["oneOf"] if "$ref" in option } - assert {"UserInputStep", "ModelOutputStep"} <= step_variants, f"Step union is missing role steps: {step_variants}" + assert {"UserInputStep", "ModelOutputStep"} <= step_variants, ( + f"Step union is missing role steps: {step_variants}" + ) for step_name, type_value in [("UserInputStep", "user_input"), ("ModelOutputStep", "model_output")]: step_schema = spec_dict["components"]["schemas"][step_name] @@ -279,9 +306,7 @@ class TestResponseCompliance: expected_fields = ["total_input_tokens", "total_output_tokens", "total_tokens"] for field in expected_fields: - assert ( - field in usage_schema["properties"] - ), f"Usage field '{field}' not in spec" + assert field in usage_schema["properties"], f"Usage field '{field}' not in spec" print(f"✓ Usage field '{field}' exists") @@ -300,9 +325,7 @@ class TestToolsCompliance: """Verify FunctionDeclaration schema for function tools.""" if "FunctionDeclaration" in spec_dict["components"]["schemas"]: func_schema = spec_dict["components"]["schemas"]["FunctionDeclaration"] - assert "name" in func_schema.get( - "properties", {} - ) or "name" in func_schema.get("required", []) + assert "name" in func_schema.get("properties", {}) or "name" in func_schema.get("required", []) print("✓ FunctionDeclaration schema found") else: print("⚠ FunctionDeclaration schema not found (may be nested)") @@ -313,33 +336,94 @@ class TestEndpointCompliance: def test_create_endpoint_exists(self, spec_dict): """Verify POST /interactions endpoint exists.""" - paths = spec_dict["paths"] - - # Find the create interactions endpoint - create_path = None - for path, methods in paths.items(): - if "interactions" in path and "post" in methods: - create_path = path - break - - assert create_path is not None, "POST /interactions endpoint not found" + create_path, _ = _interaction_operation(spec_dict, "post") print(f"✓ Create endpoint: POST {create_path}") def test_get_endpoint_exists(self, spec_dict): """Verify GET /interactions/{id} endpoint exists.""" - get_path = _interaction_resource_path(spec_dict, "get") - - assert get_path is not None, "GET /interactions/{id} endpoint not found" + get_path, _ = _interaction_operation(spec_dict, "get", individual=True) print(f"✓ Get endpoint: GET {get_path}") def test_delete_endpoint_exists(self, spec_dict): """Verify DELETE /interactions/{id} endpoint exists.""" - delete_path = _interaction_resource_path(spec_dict, "delete") - - assert delete_path is not None, "DELETE /interactions/{id} endpoint not found" + delete_path, _ = _interaction_operation(spec_dict, "delete", individual=True) print(f"✓ Delete endpoint: DELETE {delete_path}") +class TestOperationResolution: + """Keep structural resolution strict without depending on generated names.""" + + @pytest.mark.parametrize("as_union", [False, True]) + def test_model_schema_comes_from_create_operation(self, as_union): + model_schema: Final = {"properties": {"model": {"type": "string"}}, "required": ["model"]} + reference: Final = {"$ref": "#/components/schemas/RenamedModelRequest"} + body_schema: Final = ( + {"oneOf": [{"properties": {"agent": {"type": "string"}}}, reference]} if as_union else reference + ) + spec: Final = { + "paths": { + "/{version}/interactions": { + "post": { + "requestBody": {"required": True, "content": {"application/json": {"schema": body_schema}}} + } + } + }, + "components": { + "schemas": {"RenamedModelRequest": model_schema, "CreateModelInteractionParams": {"properties": {}}} + }, + } + assert _model_request_schema(spec) is model_schema + + @pytest.mark.parametrize("method,shared", [("get", False), ("delete", True)]) + def test_item_route_accepts_a_renamed_declared_identifier(self, method, shared): + parameter: Final = {"name": "renamedId", "in": "path", "required": True, "schema": {"type": "string"}} + parameters: Final = [{"$ref": "#/components/parameters/Identifier"}] + operation: Final = {"parameters": [] if shared else parameters} + path: Final = "/{version}/interactions/{renamedId}" + spec: Final = { + "paths": {path: {"parameters": parameters if shared else [], method: operation}}, + "components": {"parameters": {"Identifier": parameter}}, + } + assert _interaction_operation(spec, method, individual=True) == (path, operation) + + @pytest.mark.parametrize( + "path,parameter,error", + [ + ( + "/interactions/{id}/cancel", + {"required": True, "type": "string"}, + "Expected one GET interactions endpoint", + ), + ( + "/other_interactions/{id}", + {"required": True, "type": "string"}, + "Expected one GET interactions endpoint", + ), + ("/interactions/{id}", {"required": False, "type": "string"}, "must require its interaction ID"), + ("/interactions/{id}", {"required": True, "type": "integer"}, "must accept a string interaction ID"), + ], + ) + def test_item_route_rejects_incompatible_contracts(self, path, parameter, error): + spec: Final = { + "paths": { + path: { + "get": { + "parameters": [ + { + "name": "id", + "in": "path", + "required": parameter["required"], + "schema": {"type": parameter["type"]}, + } + ] + } + } + } + } + with pytest.raises(AssertionError, match=error): + _interaction_operation(spec, "get", individual=True) + + if __name__ == "__main__": # Quick manual test import httpx @@ -356,6 +440,4 @@ if __name__ == "__main__": if method in ["get", "post", "delete", "put", "patch"]: print(f" {method.upper()} {path}") - print( - f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}..." - ) + print(f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}...") diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 088247c2ea4..76132c7df33 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -3093,6 +3093,39 @@ def test_image_response_input_image_tokens_priced_at_image_rate(details_as_dict) assert round(cost, 12) == round(expected, 12) +@pytest.mark.parametrize("excess", ["text", "image"]) +def test_image_response_cached_modality_counts_cannot_exceed_inputs(excess): + """ + A cached_tokens_details entry larger than the matching input modality count + would turn cache reads into negative savings; the calculation must reject + the inconsistent usage instead of pricing it. + """ + from unittest.mock import patch + + from litellm.litellm_core_utils.llm_cost_calc.utils import ( + calculate_image_response_cost_from_usage, + ) + from litellm.types.utils import Usage + + cached: dict = ( + {"text_tokens": 11, "image_tokens": 0} if excess == "text" else {"text_tokens": 0, "image_tokens": 101} + ) + image_response = ImageResponse(data=[ImageObject(b64_json="x")]) + image_response.usage = Usage( + prompt_tokens=0, + completion_tokens=0, + total_tokens=212, + input_tokens=110, + input_tokens_details={"text_tokens": 10, "image_tokens": 100, "cached_tokens_details": cached}, + output_tokens=102, + output_tokens_details={"image_tokens": 102, "text_tokens": 0}, + ) + with pytest.raises(ValueError, match="Image cached token counts exceed their input modality counts"): + calculate_image_response_cost_from_usage( + model="gpt-image-2", + image_response=image_response, + custom_llm_provider="openai", + ) GEMINI_DAY0_LAUNCH_PRICING = [ ("gemini-3.6-flash", 7.5e-07, 3.75e-06, 7.5e-08), ("gemini/gemini-3.6-flash", 7.5e-07, 3.75e-06, 7.5e-08), diff --git a/tests/unit/litellm_core_utils/test_get_supported_openai_params.py b/tests/unit/litellm_core_utils/test_get_supported_openai_params.py index f9e285cf9fb..4cbbdf28efa 100644 --- a/tests/unit/litellm_core_utils/test_get_supported_openai_params.py +++ b/tests/unit/litellm_core_utils/test_get_supported_openai_params.py @@ -57,10 +57,11 @@ def test_base_model_is_additive_not_replacement(): assert real_only <= combined +@pytest.mark.usefixtures("local_model_cost_map") def test_base_model_adds_capabilities_the_real_model_lacks(): """Regression for #27717 (the behavior the union must preserve). - ``gemini-exp-9999`` isn't in the cost map so it advertises no reasoning support, + ``gemini-exp-9999`` isn't in the bundled cost map, so it advertises no reasoning support, but the registered ``gemini-3.1-pro-preview`` base_model does. The hint must add ``reasoning_effort``/``thinking`` without the call erroring.""" real_only = set(get_supported_openai_params(model="gemini-exp-9999", custom_llm_provider="gemini")) diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 7e6d4d24905..fbd3639eac0 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -3399,6 +3399,82 @@ async def test_refused_session_does_not_stamp_the_reservation_ownership_marker() 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() + + +@pytest.mark.asyncio +async def test_log_messages_flush_awaits_dispatch_instead_of_enqueueing(): + worker = MagicMock() + logger = MagicMock() + logger.model_call_details = {} + logger.dispatch_success_handlers = AsyncMock() + stream = RealTimeStreaming(MagicMock(), MagicMock(), logger, logging_worker=worker) + stream.store_message({"type": "session.created"}) + + await stream.log_messages(wait_for_dispatch=True) + + logger.dispatch_success_handlers.assert_awaited_once_with(stream.messages, prefer_async_handlers=True) + worker.ensure_initialized_and_enqueue.assert_not_called() + assert logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("account_usage", [False, True]) +async def test_attachment_cleanup_runs_in_owning_context_only(account_usage): + from litellm.litellm_core_utils.realtime_streaming import realtime_attachment_cleanup + + contexts = [] + + async def one(name): + task = asyncio.current_task() + callback = AsyncMock(side_effect=lambda: contexts.append((name, asyncio.current_task() is task))) + token = realtime_attachment_cleanup.set(callback) + try: + websocket = MagicMock() + websocket.receive_text = AsyncMock(side_effect=RuntimeError("disconnected")) + backend = MagicMock() + + async def recv(**kwargs): + await asyncio.Event().wait() + + backend.recv = recv + stream = RealTimeStreaming(websocket, backend, MagicMock(), account_usage=account_usage) + await stream.bidirectional_forward() + if account_usage: + callback.assert_not_awaited() + else: + callback.assert_awaited_once() + finally: + realtime_attachment_cleanup.reset(token) + + await asyncio.gather(one("first"), one("second")) + assert sorted(contexts) == ([] if account_usage else [("first", True), ("second", True)]) + assert realtime_attachment_cleanup.get() is None + + @pytest.mark.asyncio async def test_refused_session_stamps_the_failure_ownership_marker(): """LIT-6463: the enqueued failure callback releases the key's max_parallel_requests @@ -3550,3 +3626,81 @@ async def test_provider_bytes_are_sent_raw_after_pacing(): assert [call.args[0] for call in backend_ws.send.await_args_list] == [b"\x00\x01", '{"type":"endStream"}'] provider_config.pace_backend_send.assert_awaited_once_with(b"\x00\x01") + + +def test_public_live_accounting_survives_filtered_logging(monkeypatch): + monkeypatch.setattr(litellm, "logged_real_time_event_types", []) + stream = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + events = [ + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + { + "type": "response.event", + "event": { + "type": "response.completed", + "response": { + "id": "resp_one", + "model": "gpt-backend", + "output": [], + "usage": {"total_tokens": 12}, + }, + }, + }, + {"type": "session.closed", "usage": {"seconds": 30}}, + ] + for event in events: + stream.store_message({**event, "private_transcript": "do not retain"}) + stream.store_message( + {"type": "response.event", "event": {"type": "response.output_text.delta", "delta": "private"}} + ) + assert stream.messages == events + + +@pytest.mark.parametrize("terminal", ["response.completed", "response.incomplete", "response.failed"]) +@pytest.mark.parametrize("allowed", [[], ["response.event"], "*"]) +def test_live_terminal_logging_filters_content_and_preserves_accounting( + monkeypatch: pytest.MonkeyPatch, terminal: str, allowed: list[str] | str +) -> None: + from litellm.cost_calculator import _live_backend_responses + + monkeypatch.setattr(litellm, "logged_real_time_event_types", allowed) + stream = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + response = { + "id": "resp_private", + "created_at": 1, + "model": "gpt-backend", + "output": [{"type": "message", "content": [{"type": "output_text", "text": "private answer"}]}], + "instructions": "private instructions", + "metadata": {"private": "metadata"}, + "usage": {"input_tokens": 20, "output_tokens": 10, "total_tokens": 30}, + } + event = {"type": "response.event", "event": {"type": terminal, "response": response}} + stream.store_message(event) + + stored = stream.messages[0]["event"]["response"] + if allowed: + assert stored == response + else: + assert stored == { + key: value for key, value in response.items() if key not in ("output", "instructions", "metadata") + } | {"output": []} + measured = _live_backend_responses(stream.messages) + assert len(measured) == 1 + assert measured[0].id == "resp_private" + assert measured[0].model == "gpt-backend" + assert measured[0].usage.total_tokens == 30 + assert response["instructions"] == "private instructions" + assert response["output"][0]["content"][0]["text"] == "private answer" + + +@pytest.mark.parametrize("account_usage,expected", [(True, 1), (False, 0)]) +def test_live_initialization_is_retained_only_by_accounting_owner(account_usage, expected): + stream = RealTimeStreaming( + MagicMock(), + MagicMock(), + MagicMock(), + account_usage=account_usage, + live_initialization_seconds=15, + ) + assert len(stream.messages) == expected + if account_usage: + assert stream.messages == [{"type": "litellm.live.initialization", "usage": {"seconds": 15}}] diff --git a/tests/unit/llms/chatgpt/conftest.py b/tests/unit/llms/chatgpt/conftest.py new file mode 100644 index 00000000000..fb5b4cd5a1c --- /dev/null +++ b/tests/unit/llms/chatgpt/conftest.py @@ -0,0 +1,22 @@ +import json +import time + +import pytest + + +@pytest.fixture +def chatgpt_tokens(tmp_path, monkeypatch): + monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path)) + monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json") + for profile in ("default", "account2", "account3"): + name = "auth.json" if profile == "default" else profile + ".json" + (tmp_path / name).write_text( + json.dumps( + { + "access_token": "test-token-" + profile, + "account_id": "test-account-" + profile, + "expires_at": time.time() + 3600, + } + ) + ) + return str(tmp_path) diff --git a/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py b/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py index 0b04dd0ed78..14b91835a22 100644 --- a/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py +++ b/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py @@ -30,6 +30,29 @@ def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> Generator[None, Non class TestChatGPTResponsesAPITransformation: + def test_guardian_preserves_strict_output_schema(self): + text = { + "format": { + "type": "json_schema", + "name": "review", + "strict": True, + "schema": { + "type": "object", + "properties": {"allowed": {"type": "boolean"}}, + "required": ["allowed"], + "additionalProperties": False, + }, + } + } + request = ChatGPTResponsesAPIConfig().transform_responses_api_request( + model="codex-auto-review", + input="Review the command pwd", + response_api_optional_request_params={"text": text}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert request["text"] == text + @pytest.mark.parametrize( "model_name", [ diff --git a/tests/unit/llms/chatgpt/test_codex.py b/tests/unit/llms/chatgpt/test_codex.py new file mode 100644 index 00000000000..59107d9acff --- /dev/null +++ b/tests/unit/llms/chatgpt/test_codex.py @@ -0,0 +1,63 @@ +import hashlib +import time + +import httpx +import pytest + +from litellm.llms.chatgpt.codex import CodexRealtimeCall, build_sideband_request, parse_call_response + + +def test_encrypted_call_preserves_repeated_gateway_query(monkeypatch): + from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_call + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-repeated-query") + authorization = "Bearer test-owner" + call = CodexRealtimeCall( + call_id="rtc_repeated", + model="gpt-live-1-codex", + alias="voice", + owner=hashlib.sha256(authorization.encode()).hexdigest(), + expires_at=time.time() + 60, + extra_query={"tag": ["alpha +/&", "beta"], "gateway": "tenant"}, + ) + restored = decode_call(encode_call(call), authorization) + assert restored.extra_query == {"tag": ("alpha +/&", "beta"), "gateway": "tenant"} + assert build_sideband_request(restored)["extra_query"] == restored.extra_query + + +@pytest.mark.parametrize("location", ["", "/v1/realtime/calls/foreign-id"]) +def test_signaling_rejects_invalid_upstream_call_id(location): + response = httpx.Response(201, headers={"Location": location}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) + with pytest.raises(ValueError, match="String should match pattern"): + parse_call_response(response, "voice", "owner", 1000) + + +@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(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 + + +def test_signaling_requires_chatgpt_routing_extension(): + response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_unrouted"}) + with pytest.raises(ValueError, match="Direct call signaling requires a ChatGPT deployment"): + parse_call_response(response, "voice", "owner", 1000) diff --git a/tests/unit/llms/chatgpt/test_images.py b/tests/unit/llms/chatgpt/test_images.py new file mode 100644 index 00000000000..147c2d9702e --- /dev/null +++ b/tests/unit/llms/chatgpt/test_images.py @@ -0,0 +1,305 @@ +import base64 +import json +from typing import Final + +import httpx +import pytest + +import litellm +from litellm.llms.chatgpt.images import ChatGPTImageEditConfig, ChatGPTImageGenerationConfig +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.types.llms.openai import ImageGenerationRequestQuality +from litellm.types.router import GenericLiteLLMParams + + +@pytest.mark.parametrize( + "model,quality", + [ + ("gpt-image-2", ImageGenerationRequestQuality.AUTO), + ("gpt-image-2.5-flare", ImageGenerationRequestQuality.XHIGH), + ("gpt-image-2.5-flare", ImageGenerationRequestQuality.MAX), + ("gpt-image-2.5-sunburst", ImageGenerationRequestQuality.XHIGH), + ("gpt-image-2.5-sunburst", ImageGenerationRequestQuality.MAX), + ], +) +@pytest.mark.parametrize("editing", [False, True]) +def test_image_25_transmits_model_quality_and_transparency(model, quality, editing, chatgpt_tokens): + expected: Final = { + "model": model, + "prompt": "a red circle with transparent surroundings", + "quality": quality.value, + "background": "transparent", + "size": "2048x2048", + **({"images": [{"image_url": "data:image/png;base64,aGVsbG8="}]} if editing else {}), + } + + def respond(request): + assert str(request.url) == "https://chatgpt.com/backend-api/codex/images/" + ( + "edits" if editing else "generations" + ) + assert request.headers["content-type"] == "application/json" + assert json.loads(request.content) == expected + return httpx.Response( + 200, + json={"created": 1, "data": [{"b64_json": "aGVsbG8="}], "quality": quality.value}, + ) + + client: Final = HTTPHandler() + client.client = httpx.Client(transport=httpx.MockTransport(respond)) + operation: Final = litellm.image_edit if editing else litellm.image_generation + try: + response: Final = operation( + **{**expected, "model": "chatgpt/" + model, "quality": quality}, + client=client, + chatgpt_token_dir=chatgpt_tokens, + ) + assert response.data[0].b64_json == "aGVsbG8=" + assert response.quality == quality.value + finally: + client.client.close() + + +@pytest.mark.parametrize("model", ["gpt-image-2", "gpt-image-2.5-flare", "gpt-image-2.5-sunburst"]) +def test_json_edit_preserves_provider_params_and_extra_body_precedence(model, chatgpt_tokens): + references: Final = [{"image_url": "data:image/png;base64,aGVsbG8="}] + + def respond(request): + assert request.headers["content-type"] == "application/json" + assert json.loads(request.content) == { + "model": model, + "prompt": "red circle", + "images": references, + "seed": 7, + "provider_options": {"steps": 30, "enabled": True}, + "output_compression": 90, + } + return httpx.Response(200, json={"created": 1, "data": [{"b64_json": "aGVsbG8="}]}) + + with httpx.Client(transport=httpx.MockTransport(respond)) as http_client: + response: Final = litellm.image_edit( + model="chatgpt/" + model, + prompt="red circle", + images=references, + client=HTTPHandler(client=http_client), + chatgpt_token_dir=chatgpt_tokens, + seed=42, + output_compression=90, + extra_body={"seed": 7, "provider_options": {"steps": 30, "enabled": True}}, + ) + assert response.data[0].b64_json == "aGVsbG8=" + + +@pytest.mark.parametrize("api_base", [None, "https://image-gateway.test"]) +def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens, api_base): + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(200, json={"created": 1, "data": [{"b64_json": "aGVsbG8="}]}) + + client = HTTPHandler() + client.client = httpx.Client(transport=httpx.MockTransport(respond)) + result = litellm.image_generation( + model="chatgpt/gpt-image-2", + prompt="blue circle", + api_base=api_base, + client=client, + quality="auto", + size="auto", + background="auto", + extra_headers={"x-gateway-route": "images", "aUtHoRiZaTiOn": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, + ) + assert requests[0].headers["x-gateway-route"] == "images" + assert result.data[0].b64_json == "aGVsbG8=" + assert str(requests[0].url) == (api_base or "https://chatgpt.com/backend-api/codex") + "/images/generations" + assert requests[0].headers["authorization"] == "Bearer test-token-" + "default" + assert requests[0].headers["chatgpt-account-id"] == "test-account-" + "default" + assert b'"model":"gpt-image-2"' in requests[0].content + + +def test_codex_json_edit_survives_sdk_dispatch(chatgpt_tokens): + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(200, json={"created": 1, "data": [{"b64_json": "aGVsbG8="}]}) + + client = HTTPHandler() + client.client = httpx.Client(transport=httpx.MockTransport(respond)) + references = [{"image_url": "data:image/png;base64,aGVsbG8="}] + result = litellm.image_edit( + model="chatgpt/gpt-image-2", + prompt="red circle", + extra_headers={"x-gateway-route": "images", "authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, + images=references, + client=client, + quality="auto", + size="auto", + ) + assert result.data[0].b64_json == "aGVsbG8=" + assert str(requests[0].url) == "https://chatgpt.com/backend-api/codex/images/edits" + import json + + assert json.loads(requests[0].content)["images"] == references + + assert requests[0].headers["authorization"] == "Bearer test-token-default" + assert requests[0].headers["chatgpt-account-id"] == "test-account-default" + assert requests[0].headers["x-gateway-route"] == "images" + + +@pytest.mark.parametrize( + "references", [[], [{"image_url": "file:///etc/passwd"}], [{}], [{"image_url": "https://example.com/a.png"}] * 6] +) +def test_edit_rejects_invalid_references(references): + with pytest.raises(ValueError, match=r"images must contain|validation error"): + ChatGPTImageEditConfig().transform_image_edit_request( + "gpt-image-2", "edit", None, {}, GenericLiteLLMParams(images=references), {} + ) + + +def test_edit_converts_multipart_image_bytes(): + data, files = ChatGPTImageEditConfig().transform_image_edit_request( + "gpt-image-2", "edit", b"example", {}, GenericLiteLLMParams(), {} + ) + assert not files + assert base64.b64decode(data["images"][0]["image_url"].split(",", 1)[1]) == b"example" + + +def test_image_auth_does_not_accept_inbound_override(chatgpt_tokens): + headers = ChatGPTImageGenerationConfig().validate_environment( + {"authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, "gpt-image-2", [], {}, {"chatgpt_token_dir": chatgpt_tokens} + ) + assert httpx.Headers(headers)["authorization"] == "Bearer test-token-default" + assert httpx.Headers(headers)["chatgpt-account-id"] == "test-account-default" + + +@pytest.mark.asyncio +async def test_async_codex_edit_without_multipart_image(chatgpt_tokens): + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(200, json={"created": 1, "data": [{"b64_json": "aGVsbG8="}]}) + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + response = await litellm.aimage_edit( + model="chatgpt/gpt-image-2", + prompt="red circle", + extra_headers={"x-gateway-route": "images", "authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, + client=client, + images=[{"image_url": "data:image/png;base64,aGVsbG8="}], + chatgpt_auth_profile="account3", + ) + assert response.data[0].b64_json == "aGVsbG8=" + assert str(requests[0].url).endswith("/codex/images/edits") + assert requests[0].headers["content-type"] == "application/json" + await client.client.aclose() + + assert requests[0].headers["authorization"] == "Bearer test-token-default" + assert requests[0].headers["chatgpt-account-id"] == "test-account-default" + assert requests[0].headers["x-gateway-route"] == "images" + + +@pytest.mark.parametrize( + "references", + [None, [{"image_url": "data:image/png;base64,aGVsbG8="}]], + ids=["uploaded-image", "reference-images"], +) +def test_edit_keeps_the_authenticated_model_over_passthrough_fields(tmp_path, references): + image = None + if references is None: + image = tmp_path / "reference.png" + image.write_bytes(b"reference image bytes") + data, _ = ChatGPTImageEditConfig().transform_image_edit_request( + "gpt-image-2", + "edit", + image, + {"model": "gpt-image-2.5-flare", "size": "1024x1024"}, + GenericLiteLLMParams(images=references), + {}, + ) + assert data["model"] == "gpt-image-2" + assert data["size"] == "1024x1024" + + +@pytest.mark.parametrize("as_tuple", [False, True]) +def test_edit_accepts_filesystem_path(tmp_path, as_tuple): + image = tmp_path / "reference.png" + image.write_bytes(b"reference image bytes") + data, files = ChatGPTImageEditConfig().transform_image_edit_request( + "gpt-image-2", "edit", ("reference.png", image, "image/png") if as_tuple else image, + {}, GenericLiteLLMParams(), {} + ) + assert not files + assert data["images"] == ({"image_url": "data:image/png;base64," + base64.b64encode(image.read_bytes()).decode()},) + + +@pytest.mark.parametrize("env_name", ["CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"]) +@pytest.mark.parametrize("api_base", [None, "https://deployment.example/codex"]) +def test_image_routes_use_configured_gateway(monkeypatch, env_name, api_base, tmp_path): + token_path = tmp_path / "unavailable-token-directory" + token_path.write_text("not a directory") + monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(token_path)) + monkeypatch.delenv("CHATGPT_API_BASE", raising=False) + monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False) + monkeypatch.setenv(env_name, "https://gateway.example/codex/") + expected = api_base or "https://gateway.example/codex" + assert ChatGPTImageGenerationConfig().get_complete_url(api_base, None, "gpt-image-2", {}, {}) == ( + expected + "/images/generations" + ) + assert ChatGPTImageEditConfig().get_complete_url("gpt-image-2", api_base, {}) == expected + "/images/edits" + + +@pytest.mark.parametrize( + "reference", + [ + b"GIF89a" + b"\x00" * 32, + ("reference.gif", b"hello", "image/gif"), + ], + ids=["detected-gif", "declared-gif"], +) +def test_edit_rejects_non_bitmap_reference_content_type(reference): + with pytest.raises(ValueError, match="Reference images must be PNG, JPEG, or WEBP"): + ChatGPTImageEditConfig().transform_image_edit_request( + "gpt-image-2", "edit", reference, {}, GenericLiteLLMParams(), {} + ) + + +def test_edit_rejects_mask_before_any_provider_call(): + with pytest.raises(ValueError, match="ChatGPT image editing does not support masks"): + ChatGPTImageEditConfig().transform_image_edit_request( + "gpt-image-2", + "edit", + "data:image/png;base64,aGVsbG8=", + {"mask": "data:image/png;base64,aGVsbG8="}, + GenericLiteLLMParams(), + {}, + ) + + +def test_edit_rejects_image_and_images_together(): + with pytest.raises(ValueError, match="Specify only one of image or images"): + ChatGPTImageEditConfig().transform_image_edit_request( + "gpt-image-2", + "edit", + "data:image/png;base64,aGVsbG8=", + {}, + GenericLiteLLMParams(images=[{"image_url": "data:image/png;base64,aGVsbG8="}]), + {}, + ) + + +@pytest.mark.parametrize( + "images", + [ + [], + ["data:image/png;base64,aGVsbG8="] * 6, + ], + ids=["zero", "six"], +) +def test_edit_enforces_one_to_five_reference_images(images): + with pytest.raises(ValueError, match="images must contain between 1 and 5 reference images"): + ChatGPTImageEditConfig().transform_image_edit_request( + "gpt-image-2", "edit", images, {}, GenericLiteLLMParams(), {} + ) diff --git a/tests/unit/llms/chatgpt/test_live.py b/tests/unit/llms/chatgpt/test_live.py new file mode 100644 index 00000000000..998639349a6 --- /dev/null +++ b/tests/unit/llms/chatgpt/test_live.py @@ -0,0 +1,261 @@ +import json +from contextlib import asynccontextmanager +from unittest.mock import AsyncMock + +import httpx +import pytest + +from litellm.llms.chatgpt.live import LiveDeployment, LiveTransport, live_session_path +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + +@asynccontextmanager +async def live_handler(respond): + handler = AsyncHTTPHandler(transport=httpx.MockTransport(respond), follow_redirects=False) + try: + yield handler + finally: + await handler.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider", ["chatgpt", "openai"]) +@pytest.mark.parametrize("status", [201, 403, 429, 503]) +async def test_live_request_preserves_payload_status_and_selected_credentials(provider, status, chatgpt_tokens): + payload = { + "session": {"model": "deployment-model", "tools": [{"type": "function", "name": "lookup"}]}, + "transport": {"type": "webrtc", "sdp": "v=0\r\n"}, + "future_option": {"nested": [True, None, 3]}, + } + + def respond(request): + assert request.url.path == "/custom/v1/live/sessions" + assert request.headers["authorization"] == ( + "Bearer test-token-default" if provider == "chatgpt" else "Bearer deployment-key" + ) + assert request.headers.get("chatgpt-account-id") == ("test-account-default" if provider == "chatgpt" else None) + assert request.headers["x-gateway"] == "configured" + assert request.headers["openai-beta"] == "feature=v1" + assert "cookie" not in request.headers + assert json.loads(request.content) == payload + assert request.url.params.get_list("tag") == ["a +/&", "b"] + assert request.url.params["gateway"] == "trusted" + assert request.url.params["cursor"] == "opaque +/&" + assert not {"model", "call_id", "session_id", "api_key"}.intersection(request.url.params) + return httpx.Response(status, json={"result": "upstream"}, headers={"x-request-id": "provider-id"}) + + async with live_handler(respond) as handler: + transport = LiveTransport( + LiveDeployment( + model="deployment-model", + provider=provider, + api_key="deployment-key", + api_base="https://gateway.example/custom/v1/?gateway=base", + extra_headers={"x-gateway": "configured", "Authorization": "bad", "ChatGPT-Account-Id": "bad"}, + extra_query={"gateway": "trusted", "tag": ("a +/&", "b"), "model": "bad", "session_id": "bad"}, + ), + {"Authorization": "Bearer proxy-key", "Cookie": "private", "OpenAI-Beta": "feature=v1"}, + http_handler=handler, + ) + response = await transport.request( + "POST", + "live/sessions", + payload, + {"gateway": "untrusted", "cursor": "opaque +/&", "call_id": "bad", "api_key": "bad"}, + ) + assert response.status_code == status + assert response.json() == {"result": "upstream"} + assert response.headers["x-request-id"] == "provider-id" + assert not handler.client.is_closed + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["fork", "accept", "reject", "refer", "hangup", "content"]) +@pytest.mark.parametrize("status", [204, 404, 503]) +async def test_live_all_http_operations(operation, status): + def respond(request): + assert request.url.path == f"/v1/live/sessions/sess_new-ID/{operation}" + assert request.method == ("GET" if operation == "content" else "POST") + assert request.url.params["output_format"] == "json" + return httpx.Response(status) + + async with live_handler(respond) as handler: + transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler) + response = await transport.request( + "GET" if operation == "content" else "POST", + live_session_path("sess_new-ID", operation), + query={"output_format": "json"}, + ) + assert response.status_code == status + + +@pytest.mark.asyncio +async def test_live_request_uses_handler_methods(): + handler = AsyncMock(spec=AsyncHTTPHandler) + handler.get.return_value = httpx.Response(404) + handler.post.return_value = httpx.Response(503) + transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler) + + get_response = await transport.request("GET", live_session_path("sess_1", "content")) + post_response = await transport.request("POST", "live/sessions", {"transport": {"type": "webrtc"}}) + + assert get_response.status_code == 404 + assert post_response.status_code == 503 + handler.get.assert_awaited_once_with( + "https://api.openai.com/v1/live/sessions/sess_1/content", + headers={"authorization": "Bearer key", "content-type": "application/json"}, + timeout=60, + follow_redirects=False, + ) + handler.post.assert_awaited_once_with( + "https://api.openai.com/v1/live/sessions", + headers={"authorization": "Bearer key", "content-type": "application/json"}, + json={"transport": {"type": "webrtc"}}, + timeout=60, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["live/sessions", "live/sessions/sess_1/attach", "live/sessions/sess_1/fork"]) +async def test_live_websocket_paths_bounds_and_auth(path, chatgpt_tokens): + from websockets.asyncio.server import serve + + async def observe(connection): + assert connection.request.headers["Authorization"] == "Bearer test-token-default" + await connection.send(connection.request.path) + + async with serve(observe, "127.0.0.1", 0) as server: + port = server.sockets[0].getsockname()[1] + transport = LiveTransport( + LiveDeployment("model", api_base=f"http://127.0.0.1:{port}/v1", extra_query={"route": "a+&b"}), + {"Authorization": "Bearer proxy-key"}, + ) + connection = await transport.connect(path, {"checkpoint": "opaque+value"}) + try: + received = await connection.recv() + url = httpx.URL(f"http://127.0.0.1{received}") + assert url.path == f"/v1/{path}" + assert url.params["route"] == "a+&b" + assert url.params["checkpoint"] == "opaque+value" + finally: + await connection.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "path", + [ + "https://evil.example/live/sessions", + "//evil.example/live/sessions", + "live/sessions/../accept", + "live/sessions/sess%2Fbad/accept", + "live/sessions/sess%5Cbad/accept", + "live/sessions/sess%252Fbad/accept", + "live/sessions/%252e%252e/accept", + "live/sessions/%2E%2E/accept", + "live/sessions/sess%00bad/accept", + "live/sessions/sess%0Abad/accept", + "live/sessions/sess.foo?query/accept", + "live/sessions/sess_1/accept?url=https://evil.example", + "live/sessions/sess_1/accept#fragment", + "live/sessions/sess_1/accept\n", + ], +) +async def test_live_rejects_noncanonical_paths_before_network(path): + transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}) + with pytest.raises(ValueError, match=r"(?:Invalid|Noncanonical) Live"): + await transport.request("POST", path) + with pytest.raises(ValueError, match=r"(?:Invalid|Noncanonical) Live"): + await transport.connect(path) + + +@pytest.mark.parametrize( + "session_id", + ["../x", "sess/x", "sess%2Fx", "sess\\x", "sess%255cx", ".", "..", "%252e%252e", "", "sess\n", "sess\x00"], +) +def test_live_session_ids_cannot_inject_path_or_query(session_id): + with pytest.raises(ValueError, match="Invalid Live session ID"): + live_session_path(session_id, "content") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("session_id", ["sess.foo", "sess-\u00f1\u4e2d", "sess?x#y", "sess 50%", "x" * 1024]) +async def test_live_preserves_opaque_session_ids(session_id): + from urllib.parse import quote + + def respond(request): + assert request.url.raw_path == f"/v1/live/sessions/{quote(session_id, safe='')}/content".encode() + assert request.url.params == httpx.QueryParams() + return httpx.Response(200, json={"session_id": session_id}) + + async with live_handler(respond) as handler: + transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler) + response = await transport.request("GET", live_session_path(session_id, "content")) + assert response.json()["session_id"] == session_id + + +@pytest.mark.asyncio +async def test_live_does_not_redirect_credentials(): + def respond(request): + assert request.url.host == "api.openai.com" + return httpx.Response(307, headers={"location": "https://elsewhere.example/collect"}) + + async with live_handler(respond) as handler: + transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler) + response = await transport.request("POST", "live/sessions", {}) + assert response.status_code == 307 + + +@pytest.mark.asyncio +async def test_live_websocket_does_not_redirect_credentials(): + from websockets.asyncio.server import serve + from websockets.datastructures import Headers + from websockets.exceptions import InvalidStatus + from websockets.http11 import Response + + async def unused(connection): + pytest.fail("Redirected WebSocket must never open") + + def redirect(connection, request): + assert request.headers["authorization"] == "Bearer deployment-key" + return Response(307, "Temporary Redirect", Headers({"Location": "/elsewhere"})) + + async with serve(unused, "127.0.0.1", 0, process_request=redirect) as server: + port = server.sockets[0].getsockname()[1] + transport = LiveTransport( + LiveDeployment( + "model", provider="openai", api_key="deployment-key", api_base=f"http://127.0.0.1:{port}/v1" + ), + {}, + ) + with pytest.raises(InvalidStatus) as failure: + await transport.connect("live/sessions") + assert failure.value.response.status_code == 307 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "api_base", + [ + "ftp://gateway.example/v1", + "https://user:secret@gateway.example/v1", + "https://gateway.example/v1#fragment", + "not a url", + ], +) +async def test_live_rejects_invalid_api_base_before_network(api_base): + requests: list = [] + + def respond(request): + requests.append(request) + return httpx.Response(200, json={}) + + async with live_handler(respond) as handler: + transport = LiveTransport( + LiveDeployment("deployment-model", provider="openai", api_key="deployment-key", api_base=api_base), + {}, + http_handler=handler, + ) + with pytest.raises(ValueError, match="Invalid Live API base"): + await transport.request("POST", "live/sessions", {}) + assert requests == [] diff --git a/tests/unit/llms/chatgpt/test_realtime.py b/tests/unit/llms/chatgpt/test_realtime.py new file mode 100644 index 00000000000..327bd957c4e --- /dev/null +++ b/tests/unit/llms/chatgpt/test_realtime.py @@ -0,0 +1,519 @@ +import json +import sys +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +import litellm +from litellm.llms.chatgpt.realtime import ChatGPTRealtime +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.types.router import GenericLiteLLMParams + +pytestmark = pytest.mark.usefixtures("local_model_cost_map") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["closed", "network"]) +@pytest.mark.parametrize("hangup_status", [200, 503]) +async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, chatgpt_tokens, monkeypatch): + from websockets.exceptions import ConnectionClosedOK + from websockets.frames import Close + + from litellm.caching.llm_caching_handler import LLMClientCache + + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) + handler = ChatGPTRealtime( + GenericLiteLLMParams( + chatgpt_realtime_call_id="rtc_live_closed", + chatgpt_token_dir=chatgpt_tokens, + extra_query={"gateway": "tenant", "tag": ["alpha +/&", "beta"]}, + ), + {}, + {"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)) + try: + with patch("httpx.AsyncClient", return_value=client) as create_client: + for _ in range(2): + with pytest.raises(httpx.HTTPStatusError) if hangup_status == 503 else nullcontext(): + await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1") + assert not client.is_closed + create_client.assert_called_once() + finally: + await client.aclose() + assert len(requests) == 2 + assert requests[0].method == "POST" + assert requests[0].url.path == "/v1/realtime/calls/rtc_live_closed/hangup" + assert requests[0].url.params.get_list("tag") == ["alpha +/&", "beta"] + assert requests[0].url.params["gateway"] == "tenant" + assert requests[0].headers["x-gateway-token"] == "test-only" + assert requests[0].headers["Authorization"] == "Bearer test-token-default" + assert requests[0].extensions["timeout"]["read"] == 10 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"]) +@pytest.mark.parametrize("source", ["default", "explicit", "CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"]) +async def test_realtime_session_urls_honor_gateway(endpoint, source, chatgpt_tokens, monkeypatch): + monkeypatch.setenv("CHATGPT_TOKEN_DIR", chatgpt_tokens) + monkeypatch.delenv("CHATGPT_API_BASE", raising=False) + monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False) + gateway = "https://voice.example/custom/v1/" + if source in ("CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"): + monkeypatch.setenv(source, gateway) + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(200, json={"client_secret": {"value": "test-secret"}}) + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + kwargs = {"model": "chatgpt/gpt-realtime-1.5", "client": client} + if source == "explicit": + kwargs["api_base"] = gateway + try: + if endpoint == "client_secrets": + await litellm.acreate_realtime_client_secret(**kwargs) + else: + await litellm.acreate_realtime_transcription_session(**kwargs) + finally: + await client.client.aclose() + base = "https://api.openai.com/v1" if source == "default" else gateway.rstrip("/") + assert len(requests) == 1 + assert str(requests[0].url) == f"{base}/realtime/{endpoint}" + assert requests[0].headers["authorization"] == "Bearer test-token-default" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("inbound_headers", [{}, {"openai-alpha": "quicksilver=v2"}]) +@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 = [] + + def respond(request): + requests.append(request) + return httpx.Response(201, text="v=0\r\n", headers={"location": "/v1/realtime/calls/rtc_test"}) + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + router = litellm.Router( + model_list=[ + { + "model_name": "voice-gateway", + "litellm_params": { + "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", + "count": 7, + "fraction": 1.5, + "enabled": True, + "disabled": False, + "blank": None, + "tag": ["alpha +/&", "beta"], + "empty": [], + "model": "other-model", + "call_id": "rtc_wrong", + }, + }, + "model_info": {"id": "selected-gateway-deployment"}, + } + ], + num_retries=0, + ) + offer = CodexRealtimeOffer(sdp="v=0\r\n", session={"model": "voice-gateway"}) + try: + 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", + "count": "7", + "fraction": "1.5", + "enabled": "true", + "disabled": "false", + "blank": "", + "tag": "alpha +/&", + "model": "other-model", + "call_id": "rtc_wrong", + } + assert requests[0].url.params.get_list("tag") == ["alpha +/&", "beta"] + assert response.extensions["chatgpt_realtime"]["extra_query"] == { + **dict(requests[0].url.params), + "tag": ("alpha +/&", "beta"), + "empty": (), + } + 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.params.get_list("tag") == ["alpha +/&", "beta"] + assert sideband_url.path.endswith("/realtime" if endpoint == "realtime" else "/live/rtc_test") + finally: + await client.client.aclose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ["gpt-realtime-1.5", "gpt-live-1-codex"]) +@pytest.mark.parametrize("call_id", [None, "rtc_existing"]) +async def test_websocket_forwards_configured_headers_without_client_identity(model, call_id, chatgpt_tokens): + websocket = SimpleNamespace( + headers={"authorization": "Bearer client", "cookie": "private-cookie", "openai-alpha": "client-value"}, + scope={}, + receive_text=AsyncMock(side_effect=RuntimeError("client disconnected")), + send_text=AsyncMock(), + close=AsyncMock(), + ) + with patch("websockets.connect") as connect: + connect.return_value.__aenter__ = AsyncMock(side_effect=RuntimeError("stop before streaming")) + await litellm._arealtime( + model=f"chatgpt/{model}", + websocket=websocket, + api_base="https://voice.example/codex", + chatgpt_realtime_call_id=call_id, + query_params={"model": model, "intent": "client-intent"}, + extra_query={"intent": "configured-intent", "tag": ["alpha +/&", "beta"]}, + headers={"x-deployment-header": "configured"}, + extra_headers={ + "X-Gateway-Route": "voice", + "OpenAI-Alpha": "configured-value", + "aUtHoRiZaTiOn": "Bearer wrong", + "CHATGPT-ACCOUNT-ID": "wrong", + }, + ) + connect.assert_called_once() + headers = httpx.Headers(connect.call_args.kwargs["additional_headers"]) + upstream_url = httpx.URL(connect.call_args.args[0]) + assert upstream_url.params.get_list("intent") == ["configured-intent"] + assert upstream_url.params.get_list("tag") == ["alpha +/&", "beta"] + assert headers["x-deployment-header"] == "configured" + assert headers["x-gateway-route"] == "voice" + assert headers["openai-alpha"] == "configured-value" + assert headers["authorization"] == "Bearer test-token-default" + assert headers["chatgpt-account-id"] == "test-account-default" + assert "cookie" not in headers + + +@pytest.mark.asyncio +@pytest.mark.parametrize("api_base", [None, "https://voice.example/backend-api/codex"]) +async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, api_base): + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(201, text="v=0\r\n", headers={"location": "/v1/realtime/calls/rtc_test"}) + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + response = await litellm.arealtime_calls( + model="chatgpt/gpt-live-1-codex", + api_base=api_base, + openai_ephemeral_key="", + 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", + "aUtHoRiZaTiOn": "Bearer wrong", + "CHATGPT-ACCOUNT-ID": "wrong", + }, + 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 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" + assert requests[0].headers["chatgpt-account-id"] == "test-account-default" + assert requests[0].headers["openai-alpha"] == "quicksilver=v2" + assert requests[0].headers["x-gateway-route"] == "voice" + assert json.loads(requests[0].content) == { + "sdp": "v=0\r\n", + "session": {"model": "gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}}, + } + await client.client.aclose() + + +@pytest.mark.asyncio +async def test_chatgpt_call_rejects_ephemeral_key_before_oauth_dispatch(chatgpt_tokens): + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(201, text="v=0\r\n") + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + try: + with pytest.raises(litellm.AuthenticationError): + await litellm.arealtime_calls( + model="chatgpt/gpt-live-1-codex", + openai_ephemeral_key="legacy-ephemeral-key", + sdp_body=b"v=0\r\n", + client=client, + ) + assert not requests + finally: + await client.client.aclose() + + +@pytest.mark.asyncio +async def test_openai_call_preserves_explicit_identity_headers(): + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(201, text="v=0\r\n") + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + try: + response = await litellm.arealtime_calls( + model="openai/gpt-realtime-1.5", + openai_ephemeral_key="original-key", + sdp_body=b"v=0\r\n", + extra_headers={"Authorization": "Bearer explicit-key", "chatgpt-account-id": "custom-account"}, + client=client, + ) + assert response.status_code == 201 + assert requests[0].headers["authorization"] == "Bearer explicit-key" + assert requests[0].headers["chatgpt-account-id"] == "custom-account" + assert requests[0].headers["content-type"].startswith("multipart/form-data") + finally: + await client.client.aclose() + + +@pytest.mark.parametrize("model,endpoint", [("gpt-realtime-1.5", "realtime"), ("gpt-live-1-codex", "live")]) +def test_realtime_uses_platform_endpoint_with_oauth_headers(model, endpoint, chatgpt_tokens, local_model_cost_map): + handler = ChatGPTRealtime( + GenericLiteLLMParams(), + { + "authorization": "Bearer proxy-key", + "openai-alpha": "quicksilver=v2", + }, + ) + assert handler._construct_url("https://api.openai.com/v1", {"model": model}) == ( + f"wss://api.openai.com/v1/{endpoint}?model={model}" + ) + headers = handler._get_additional_headers("unused") + assert headers["Authorization"] == "Bearer test-token-default" + assert "authorization" not in headers + assert headers["openai-alpha"] == "quicksilver=v2" + + +@pytest.mark.parametrize("endpoint", ["live", "realtime"]) +def test_new_realtime_session_preserves_gateway_query(endpoint, chatgpt_tokens, local_model_cost_map): + model = "gpt-live-1-codex" if endpoint == "live" else "gpt-realtime-1.5" + handler = ChatGPTRealtime( + GenericLiteLLMParams( + chatgpt_token_dir=chatgpt_tokens, + chatgpt_realtime_client_query={"intent": "conversation", "architecture": "client-architecture"}, + extra_query={ + "gateway_token": "opaque +/& value", + "intent": "gateway-intent", + "architecture": "gateway-architecture", + "model": "other-model", + "call_id": "rtc_other", + }, + ), + {}, + ) + url = httpx.URL(handler._construct_url("https://gateway.example/v1", {"model": model, "intent": "query-intent"})) + assert url.path == f"/v1/{endpoint}" + assert dict(url.params) == { + "model": model, + "gateway_token": "opaque +/& value", + "intent": "gateway-intent", + "architecture": "gateway-architecture", + } + + +@pytest.mark.asyncio +async def test_openai_http_call_does_not_require_websockets(monkeypatch): + monkeypatch.delitem(sys.modules, "litellm.llms.chatgpt.realtime", raising=False) + for name in tuple(sys.modules): + if name == "websockets" or name.startswith("websockets."): + monkeypatch.delitem(sys.modules, name) + monkeypatch.setitem(sys.modules, "websockets", None) + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(201, text="v=0\r\n") + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + try: + response = await litellm.arealtime_calls( + model="openai/gpt-realtime-1.5", + openai_ephemeral_key="test-only", + sdp_body=b"v=0\r\n", + api_key="test-only", + client=client, + ) + assert response.status_code == 201 + assert len(requests) == 1 + assert requests[0].url.path == "/v1/realtime/calls" + finally: + await client.client.aclose() + + +@pytest.mark.parametrize("endpoint", ["live", "realtime"]) +@pytest.mark.parametrize("call_id", [None, "rtc_metadata"]) +def test_realtime_routes_new_models_using_registered_metadata(endpoint, call_id, chatgpt_tokens, local_model_cost_map): + model = "metadata-voice-model" + litellm.register_model({f"chatgpt/{model}": { + "litellm_provider": "chatgpt", "mode": "realtime", "supported_endpoints": [f"/v1/{endpoint}"] + }}) + handler = ChatGPTRealtime(GenericLiteLLMParams(chatgpt_realtime_call_id=call_id), {}) + expected = ( + f"wss://api.openai.com/v1/{endpoint}?model={model}" if call_id is None + else f"wss://api.openai.com/v1/live/{call_id}" if endpoint == "live" + else f"wss://api.openai.com/v1/realtime?call_id={call_id}" + ) + assert handler._construct_url("https://api.openai.com/v1", {"model": model}) == expected + + +def test_realtime_unknown_model_keeps_standard_endpoint(chatgpt_tokens, local_model_cost_map): + handler = ChatGPTRealtime(GenericLiteLLMParams(), {}) + assert handler._construct_url("https://api.openai.com/v1", {"model": "unknown-voice-model"}) == ( + "wss://api.openai.com/v1/realtime?model=unknown-voice-model" + ) + + +@pytest.mark.parametrize("env_name", ["CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"]) +@pytest.mark.parametrize("api_base", [None, "https://deployment.example/codex"]) +def test_realtime_routes_use_configured_gateway(monkeypatch, env_name, api_base, chatgpt_tokens): + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + + monkeypatch.delenv("CHATGPT_API_BASE", raising=False) + monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False) + monkeypatch.setenv(env_name, "https://gateway.example/codex/") + expected = api_base or "https://gateway.example/codex" + config = ChatGPTRealtimeHTTPConfig(GenericLiteLLMParams()) + assert config.get_realtime_calls_url(api_base, "gpt-live-1-codex") == expected + "/realtime/calls" + handler = ChatGPTRealtime(GenericLiteLLMParams(), {}) + 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" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ["gpt-live-1-codex", "gpt-realtime-1.5"]) +async def test_close_call_prefers_session_close_only_for_live_models(model, chatgpt_tokens): + handler = ChatGPTRealtime( + GenericLiteLLMParams(chatgpt_realtime_call_id="rtc_close", chatgpt_token_dir=chatgpt_tokens), + {}, + {}, + ) + connection = SimpleNamespace(send=AsyncMock()) + handler.hangup_call = AsyncMock() + await handler.close_call(connection, model, "https://gateway.example/v1") + if model == "gpt-live-1-codex": + connection.send.assert_awaited_once_with('{"type":"session.close"}') + handler.hangup_call.assert_not_awaited() + else: + connection.send.assert_not_awaited() + handler.hangup_call.assert_awaited_once_with("https://gateway.example/v1") diff --git a/tests/unit/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py index f3332cb513c..757c2ec18d2 100644 --- a/tests/unit/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py @@ -3835,6 +3835,125 @@ def test_image_edit_handler_keeps_the_sync_transform(): assert response.data[0].b64_json == "sync" +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"]) +@pytest.mark.parametrize("provider", ["chatgpt", "openai"]) +@pytest.mark.parametrize("authorization_header", ["Authorization", "aUtHoRiZaTiOn"]) +async def test_realtime_http_sessions_preserve_provider_identity( + endpoint, provider, authorization_header, tmp_path, monkeypatch +): + import time + + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig + from litellm.types.router import GenericLiteLLMParams + + monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path)) + monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json") + (tmp_path / "auth.json").write_text( + json.dumps({"access_token": "test-resolved", "account_id": "test-selected", "expires_at": time.time() + 3600}) + ) + config = ChatGPTRealtimeHTTPConfig(GenericLiteLLMParams()) if provider == "chatgpt" else OpenAIRealtimeHTTPConfig() + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(200, json={"id": "session-test"}) + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + try: + response = await BaseLLMHTTPHandler()._async_realtime_session_post( + endpoint=endpoint, + api_base="https://gateway.example/v1", + api_key="test-openai", + request_data={"session": {"model": "gpt-realtime-1.5"}}, + logging_obj=Mock(), + timeout=5, + provider_config=config, + model="gpt-realtime-1.5", + extra_headers={ + authorization_header: "Bearer test-override", + "CHATGPT-ACCOUNT-ID": "test-other-account", + "x-gateway-route": "required", + }, + client=client, + ) + assert response.status_code == 200 + assert not client.client.is_closed + finally: + await client.client.aclose() + assert len(requests) == 1 + assert requests[0].url.path == f"/v1/realtime/{endpoint}" + assert requests[0].headers["x-gateway-route"] == "required" + if provider == "chatgpt": + assert requests[0].headers.get_list("authorization") == ["Bearer test-resolved"] + assert requests[0].headers.get_list("chatgpt-account-id") == ["test-selected"] + else: + assert requests[0].headers.get_list("authorization")[-1] == "Bearer test-override" + assert requests[0].headers["chatgpt-account-id"] == "test-other-account" + + +class _ImageGenerationRecordingConfig(BaseImageGenerationConfig): + def get_supported_openai_params(self, model): + return ["size"] + + def map_openai_params(self, non_default_params, optional_params, model, drop_params): + optional_params.update(non_default_params) + return optional_params + + def validate_environment(self, headers, model, messages, optional_params, litellm_params, api_key=None, api_base=None): + return {"authorization": f"Bearer {api_key}"} + + def get_complete_url(self, api_base, api_key, model, optional_params, litellm_params, stream=None): + return "https://images.example/v1/generations" + + def transform_image_generation_request(self, model, prompt, optional_params, litellm_params, headers): + return {"model": model, "prompt": prompt} + + def transform_image_generation_response(self, model, raw_response, model_response, logging_obj, request_data, optional_params, litellm_params, encoding=None, api_key=None, json_mode=None): + return ImageResponse(data=[ImageObject(b64_json=raw_response.json()["created"])]) + + +def test_image_extra_headers_strips_oauth_identity_only_for_chatgpt(): + headers: Final = {"authorization": "Bearer oauth", "chatgpt-account-id": "acct-1", "x-router": "keep"} + assert BaseLLMHTTPHandler._image_extra_headers("openai", headers) is headers + stripped: Final = BaseLLMHTTPHandler._image_extra_headers("chatgpt", headers) + assert dict(stripped) == {"x-router": "keep"} + + +@pytest.mark.asyncio +async def test_async_image_generation_handler_merges_extra_headers_for_non_chatgpt(): + requests: Final = [] + + def respond(request): + requests.append(request) + return httpx.Response(200, json={"created": "ok"}) + + client: Final = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + try: + response: Final = await BaseLLMHTTPHandler().async_image_generation_handler( + model="image-model", + prompt="a red circle", + image_generation_provider_config=_ImageGenerationRecordingConfig(), + image_generation_optional_request_params={}, + custom_llm_provider="openai", + litellm_params={"api_key": "sk-image"}, + logging_obj=Mock(), + timeout=10, + extra_headers={"x-router-header": "routed"}, + api_key="sk-image", + client=client, + ) + finally: + await client.client.aclose() + assert requests[0].headers["x-router-header"] == "routed" + assert requests[0].headers["authorization"] == "Bearer sk-image" + assert requests[0].url == "https://images.example/v1/generations" + assert response.data[0].b64_json == "ok" + + class _ScriptedClientWebSocket(_FakeClientWebSocket): def __init__(self, messages: list[str], last_event_type: str) -> None: super().__init__() diff --git a/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py b/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py index f7a88b5ba63..774f7abd049 100644 --- a/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py +++ b/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py @@ -26,6 +26,15 @@ def test_openai_realtime_handler_url_construction(api_base): assert "model=gpt-4o-realtime-preview-2024-10-01" in url +def test_openai_realtime_handler_requires_api_key(): + from litellm.llms.openai.realtime.handler import OpenAIRealtime + + handler = OpenAIRealtime() + with pytest.raises(ValueError, match="api_key is required for OpenAI realtime calls"): + handler._resolve_api_key(None) + assert handler._resolve_api_key("sk-realtime-key") == "sk-realtime-key" + + def test_openai_realtime_handler_url_with_extra_params(): from litellm.llms.openai.realtime.handler import OpenAIRealtime from litellm.types.realtime import RealtimeQueryParams diff --git a/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py index ba2b26bf0a2..bcc024105a9 100644 --- a/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py @@ -308,6 +308,7 @@ class TestProcessResponse: ) +@pytest.mark.usefixtures("local_model_cost_map") class TestProcessEmbedContentResponseUsage: """Gemini Embedding 2 embedContent usageMetadata must drive spend. diff --git a/tests/unit/proxy/auth/test_auth_checks.py b/tests/unit/proxy/auth/test_auth_checks.py index 2538556d3b5..a5f00b01e0a 100644 --- a/tests/unit/proxy/auth/test_auth_checks.py +++ b/tests/unit/proxy/auth/test_auth_checks.py @@ -547,6 +547,8 @@ async def test_virtual_key_max_budget_check( False, ), # don't match on pattern ("openai/gpt-4o", ["openai/*"], True), # openai wildcard access + ("openai/gpt+4", ["openai/gpt+*"], True), # regex metacharacters stay literal + ("openai/gpttt4", ["openai/gpt+*"], False), # regex metacharacters do not overmatch ("gpt-4", ["gpt-3.5-turbo"], False), # model not in allowed list ("claude-3", [], True), # empty model list (allows all) ], diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index 9cdac341b1f..fc76cf84ef7 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -857,6 +857,7 @@ async def test_user_api_key_auth_websocket(): # Prepare a mock WebSocket object mock_websocket = MagicMock(spec=WebSocket) mock_websocket.query_params = {"model": "some_model"} + mock_websocket.path_params = {} mock_websocket.headers = {"authorization": "Bearer some_api_key"} # Mock the scope attribute that user_api_key_auth_websocket accesses mock_websocket.scope = {"headers": [(b"authorization", b"Bearer some_api_key")]} @@ -880,6 +881,7 @@ async def test_user_api_key_auth_websocket(): assert request_arg.headers["authorization"] == "Bearer some_api_key" assert mock_user_api_key_auth.call_args.kwargs["api_key"] == "Bearer some_api_key" + assert await request_arg.json() == {"model": "some_model"} @pytest.mark.asyncio @@ -893,6 +895,7 @@ async def test_user_api_key_auth_websocket_carries_asgi_path(): mock_websocket = MagicMock(spec=WebSocket) mock_websocket.query_params = {"model": "some_model"} + mock_websocket.path_params = {} mock_websocket.headers = {"authorization": "Bearer some_api_key"} mock_websocket.scope = { "type": "websocket", diff --git a/tests/unit/realtime_api/test_main.py b/tests/unit/realtime_api/test_main.py index 5d3276dfae1..59df424577c 100644 --- a/tests/unit/realtime_api/test_main.py +++ b/tests/unit/realtime_api/test_main.py @@ -1,4 +1,5 @@ import asyncio +import json import time from types import TracebackType from typing import Final @@ -26,6 +27,42 @@ class FakeLogging: pass +@pytest.mark.parametrize("provider", [litellm.LlmProviders.XAI, litellm.LlmProviders.OPENAI, litellm.LlmProviders.GEMINI]) +def test_realtime_handler_factory_does_not_read_headers_without_a_handler(provider): + from litellm.types.router import GenericLiteLLMParams + + read_headers = MagicMock(side_effect=AssertionError("Headers must not be read")) + assert realtime_main.ProviderConfigManager.get_provider_realtime_handler( + provider, GenericLiteLLMParams(), read_headers + ) is None + read_headers.assert_not_called() + + +def test_realtime_handler_factory_passes_actual_chatgpt_headers(tmp_path, monkeypatch): + from litellm.llms.chatgpt.realtime import ChatGPTRealtime + from litellm.types.router import GenericLiteLLMParams + + monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path)) + monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json") + (tmp_path / "auth.json").write_text( + json.dumps({"access_token": "factory-test-token", "account_id": "factory-account", "expires_at": time.time() + 3600}) + ) + params = GenericLiteLLMParams(litellm_session_id="factory-session") + headers = {"openai-alpha": "quicksilver=v2"} + extra_headers = {"x-gateway-route": "required"} + read_headers = MagicMock(return_value=headers) + result = realtime_main.ProviderConfigManager.get_provider_realtime_handler( + litellm.LlmProviders.CHATGPT, params, read_headers, extra_headers + ) + assert isinstance(result, ChatGPTRealtime) + read_headers.assert_called_once_with() + outgoing_headers = result._get_additional_headers("unused") + assert outgoing_headers["openai-alpha"] == headers["openai-alpha"] + assert outgoing_headers["x-gateway-route"] == extra_headers["x-gateway-route"] + assert outgoing_headers["session_id"] == "factory-session" + assert outgoing_headers["Authorization"] == "Bearer factory-test-token" + + def test_resolves_top_level_session_model(): resolved = _with_resolved_session_model({"model": "alias/gpt-realtime"}, "gpt-realtime") assert resolved == {"model": "gpt-realtime"} @@ -510,6 +547,32 @@ async def test_arealtime_azure_env_beta_protocol_wins_over_a_ga_client(monkeypat ) +@pytest.mark.parametrize("is_call", [False, True]) +@pytest.mark.parametrize("provider", ["chatgpt", "openai", "azure"]) +def test_realtime_http_provider_controls_dynamic_base_precedence(provider, is_call, monkeypatch): + from litellm.types.router import GenericLiteLLMParams + + monkeypatch.delenv("CHATGPT_API_BASE", raising=False) + monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False) + config, base, key = realtime_main._get_realtime_http_provider_config( + custom_llm_provider=provider, + dynamic_api_base="https://dynamic.example/v1", + dynamic_api_key="dynamic-key", + litellm_params=GenericLiteLLMParams(api_base="https://configured.example/v1"), + is_call=is_call, + ) + expected_base = "https://configured.example/v1" if provider == "chatgpt" else "https://dynamic.example/v1" + assert base == expected_base + assert key == ("chatgpt-oauth" if provider == "chatgpt" else "dynamic-key") + assert config is not None + if provider == "chatgpt": + assert config.get_realtime_calls_url(base, "gpt-realtime-1.5") == expected_base + "/realtime/calls" + else: + assert config.get_realtime_calls_extra_headers({"x-gateway-route": "required"}) == { + "x-gateway-route": "required" + } + + async def _vertex_provider_config_for(monkeypatch, model: str, vertex_location: str | None): from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import VertexChirpRealtimeConfig from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 36e188e82d6..ca0a0ceeb7d 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -1479,13 +1479,81 @@ def test_azure_ai_cache_cost_calculation(_local_model_cost_map): print(f"Output cost: {output_cost}, Expected: {expected_output_cost}") print(f"Total cost: {total_cost}") - assert abs(input_cost - expected_input_cost) < 1e-10, ( - f"Input cost mismatch: got {input_cost}, expected {expected_input_cost}" + assert ( + abs(input_cost - expected_input_cost) < 1e-10 + ), f"Input cost mismatch: got {input_cost}, expected {expected_input_cost}" + assert ( + abs(output_cost - expected_output_cost) < 1e-10 + ), f"Output cost mismatch: got {output_cost}, expected {expected_output_cost}" + + +AZURE_GPT_5_6_MAP_KEYS = ( + "azure/gpt-5.6", + "azure/gpt-5.6-sol", + "azure/gpt-5.6-terra", + "azure/gpt-5.6-luna", + "azure/us/gpt-5.6", + "azure/us/gpt-5.6-sol", + "azure/us/gpt-5.6-terra", + "azure/us/gpt-5.6-luna", + "azure/eu/gpt-5.6", + "azure/eu/gpt-5.6-sol", + "azure/eu/gpt-5.6-terra", + "azure/eu/gpt-5.6-luna", +) + + +def test_azure_gpt_5_6_cache_write_tokens_are_billed(_local_model_cost_map): + """ + Azure bills gpt-5.6 prompt cache writes at 1.25x the input rate on every + tier, but the azure entries carried no ``cache_creation_input_token_cost``, + so cache-write tokens were billed at the plain input rate instead. + """ + from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token + from litellm.types.utils import PromptTokensDetailsWrapper, Usage + + usage = Usage( + completion_tokens=100, + prompt_tokens=2000, + total_tokens=2100, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=0, text_tokens=687), + cache_creation_input_tokens=1313, ) - assert abs(output_cost - expected_output_cost) < 1e-10, ( - f"Output cost mismatch: got {output_cost}, expected {expected_output_cost}" + + input_cost, output_cost = generic_cost_per_token( + model="azure/gpt-5.6-luna", usage=usage, custom_llm_provider="azure" ) + assert input_cost == pytest.approx(687 * 2e-07 + 1313 * 2.5e-07) + assert output_cost == pytest.approx(100 * 1.2e-06) + + +@pytest.mark.parametrize("model", AZURE_GPT_5_6_MAP_KEYS) +def test_azure_gpt_5_6_rates_match_azure_price_page(_local_model_cost_map, model): + """ + Per the Azure OpenAI price page (rendered 2026-08-26): cache writes cost + 1.25x input on every gpt-5.6 tier, and Data Zone costs 1.1x Global for + standard and priority alike (us/eu priority rates previously sat at 1.25x). + """ + entry = litellm.model_cost[model] + input_keys = [key for key in entry if key.startswith("input_cost_per_token")] + assert input_keys + for key in input_keys: + suffix = key[len("input_cost_per_token") :] + assert entry["cache_creation_input_token_cost" + suffix] == pytest.approx(entry[key] * 1.25) + + zone = model.split("/")[1] + if zone in ("us", "eu"): + global_entry = litellm.model_cost["azure/" + model.split("/", 2)[2]] + prefixes = ("input_cost_per_token", "output_cost_per_token", "cache_read", "cache_creation") + token_cost_keys = [key for key in entry if key.startswith(prefixes)] + global_token_cost_keys = [key for key in global_entry if key.startswith(prefixes)] + assert len(token_cost_keys) >= 9 + assert set(token_cost_keys) <= set(global_token_cost_keys) + for key in token_cost_keys: + assert entry[key] == pytest.approx(global_entry[key] * 1.1), key + + def test_vertex_regional_deployment_costs_uplift_over_global(monkeypatch): """ @@ -3647,6 +3715,101 @@ def test_combine_usage_objects_sums_mirrored_cache_write_fields_once(): assert combined_pair.prompt_tokens_details.cache_creation_tokens == 100 +def test_completion_cost_prices_anthropic_shaped_cache_read_tokens(_local_model_cost_map): + """Regression: an Anthropic /v1/messages response reports cache reads as top-level + cache_read_input_tokens with input_tokens excluding them. Reading that usage as + Responses API usage dropped the cache tokens and billed the whole prompt at the + uncached input rate, overstating spend on cache hits.""" + + response = { + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "gpt-5.6-sol", + "stop_reason": "end_turn", + "content": [{"type": "text", "text": "1"}], + "usage": {"input_tokens": 3, "output_tokens": 5, "cache_read_input_tokens": 4014}, + } + + cost = litellm.completion_cost( + completion_response=response, + model="gpt-5.6-sol", + custom_llm_provider="openai", + ) + + assert cost == pytest.approx(3 * 4e-6 + 4014 * 4e-7 + 5 * 2e-5, rel=1e-9) + + +def _together_chat_response(model: str, prompt_tokens: int, completion_tokens: int, cached_tokens: int) -> ModelResponse: + return ModelResponse( + id="chatcmpl-together-cache", + choices=[{"finish_reason": "stop", "index": 0, "message": {"content": "acknowledged", "role": "assistant"}}], + created=1756164000, + model=model, + object="chat.completion", + usage=Usage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens), + ), + ) + + +def test_completion_cost_prices_together_cached_tokens_at_cache_read_rate(_local_model_cost_map): + """Regression: Together reports prompt_tokens_details.cached_tokens but no together_ai + registry entry carried cache_read_input_token_cost, so cache-hit tokens were priced at + 0.0 and spend on cache-heavy workloads was understated.""" + + cost = completion_cost( + completion_response=_together_chat_response( + model="deepseek-ai/DeepSeek-V4-Flash-0731", prompt_tokens=7864, completion_tokens=16, cached_tokens=7863 + ), + custom_llm_provider="together_ai", + ) + + assert cost == pytest.approx(1 * 1.4e-07 + 7863 * 3e-08 + 16 * 2.8e-07, rel=1e-9) + + +def test_completion_cost_together_mapped_model_skips_size_bucket(_local_model_cost_map): + """Regression: any together model whose name matches (\\d+b) was rewritten to a + together-ai-* size bucket before the registry lookup, so mapped models like + Muse-Glimmer-30B never used their per-model rates, cache fields included.""" + + cost = completion_cost( + completion_response=_together_chat_response( + model="meta-models/Muse-Glimmer-30B", prompt_tokens=63, completion_tokens=16, cached_tokens=0 + ), + custom_llm_provider="together_ai", + ) + + assert cost == pytest.approx(63 * 3.5e-07 + 16 * 1.5e-06, rel=1e-9) + + +def test_completion_cost_together_unmapped_model_still_uses_size_bucket(_local_model_cost_map): + cost = completion_cost( + completion_response=_together_chat_response( + model="qwen/Qwen2-72B-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0 + ), + custom_llm_provider="together_ai", + ) + + assert cost == pytest.approx((23 + 15) * 9e-07, rel=1e-9) + + +def test_completion_cost_together_metadata_only_model_still_uses_size_bucket(_local_model_cost_map): + assert "input_cost_per_token" not in litellm.model_cost["together_ai/togethercomputer/CodeLlama-34b-Instruct"] + + cost = completion_cost( + completion_response=_together_chat_response( + model="togethercomputer/CodeLlama-34b-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0 + ), + custom_llm_provider="together_ai", + ) + + assert cost == pytest.approx((23 + 15) * 8e-07, rel=1e-9) + + def test_select_model_name_strips_unregistered_alias_prefix(_local_model_cost_map): """A router-facing model_name alias containing "/" whose leading segment is NOT a registered provider must not be double-prefixed into a non-existent cost key. @@ -4233,6 +4396,106 @@ def test_collect_and_combine_realtime_usage_stores_partitioned_text_tokens() -> 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) + + +@pytest.mark.parametrize("public_live", [False, True]) +def test_live_terminal_duration_honors_deployment_override(monkeypatch, public_live): + 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(), [{"type": "session.closed", "usage": {"seconds": 4}}] if public_live else [_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) + + +@pytest.mark.parametrize("with_tokens", [False, True]) +@pytest.mark.parametrize("terminal_count", [1, 2]) +@pytest.mark.parametrize("duration_priced", [False, True]) +def test_live_terminal_with_response_done_preserves_configured_billing( + monkeypatch, with_tokens, terminal_count, duration_priced +): + monkeypatch.setitem( + litellm.model_cost, + "realtime-deployment-test", + { + "litellm_provider": "chatgpt", + "mode": "realtime", + **( + {"input_cost_per_second": 0.025} + if duration_priced + else {"input_cost_per_token": 0.001, "output_cost_per_token": 0.002} + ), + }, + ) + events = [ + { + "type": "response.done", + "response": { + "usage": ({"input_tokens": 10, "output_tokens": 5, "total_tokens": 15} if with_tokens else {}) + }, + }, + *(_live_terminal_event() for _ in range(terminal_count)), + ] + usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(events) + result = RealtimeAPITokenUsageProcessor.create_logging_realtime_object(usage, events) + assert completion_cost( + completion_response=result, + model="gpt-live-1-codex" if duration_priced else "gpt-realtime-1.5", + custom_llm_provider="chatgpt", + call_type="_arealtime", + custom_pricing=True, + router_model_id="realtime-deployment-test", + ) == pytest.approx(0.1 if duration_priced else (0.02 if with_tokens else 0)) def test_realtime_combine_sums_nested_cached_tokens_details(): results: OpenAIRealtimeStreamList = [ { @@ -4561,6 +4824,241 @@ def test_completion_cost_ocr_ignores_deployment_pricing_without_custom_pricing_f assert cost == 0.0 +@pytest.mark.parametrize("terminal", [False, True]) +def test_public_live_seconds_are_cumulative_and_backend_usage_is_separately_priced(monkeypatch, terminal): + monkeypatch.setitem( + litellm.model_cost, + "live-seconds-test", + { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.025, + "input_cost_per_token": 100, + "output_cost_per_token": 100, + }, + ) + monkeypatch.setitem( + litellm.model_cost, + "live-backend-test", + { + "litellm_provider": "openai", + "mode": "responses", + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + }, + ) + backend = { + "type": "response.event", + "event": { + "type": "response.completed", + "response": { + "id": "resp_backend", + "created_at": 1, + "model": "live-backend-test", + "output": [], + "usage": {"input_tokens": 20, "output_tokens": 10, "total_tokens": 30}, + }, + }, + } + events = [ + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + {"type": "session.usage.updated", "usage": {"seconds": 30}}, + backend, + backend, + {"type": "session.closed" if terminal else "session.usage.updated", "usage": {"seconds": 30}}, + ] + usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(events) + assert usage.total_tokens == 30 + assert handle_realtime_stream_cost_calculation(events, usage, "openai", "live-seconds-test") == pytest.approx(0.79) + + +@pytest.mark.parametrize("seconds", [-1, True, "30", float("inf"), float("nan"), None]) +def test_public_live_invalid_seconds_are_not_billed(monkeypatch, seconds): + monkeypatch.setitem( + litellm.model_cost, + "live-seconds-test", + { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.025, + }, + ) + assert ( + handle_realtime_stream_cost_calculation( + [{"type": "session.closed", "usage": {"seconds": seconds}}], Usage(), "openai", "live-seconds-test" + ) + == 0 + ) + + +@pytest.mark.parametrize("terminal", [False, True]) +def test_live_duration_does_not_regress_when_primary_and_observer_events_interleave(monkeypatch, terminal): + monkeypatch.setitem( + litellm.model_cost, + "live-interleaved-test", + {"litellm_provider": "openai", "mode": "realtime", "input_cost_per_second": 0.025}, + ) + events = [ + {"type": "session.closed" if terminal else "session.usage.updated", "usage": {"seconds": 30}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + ] + assert handle_realtime_stream_cost_calculation(events, Usage(), "openai", "live-interleaved-test") == pytest.approx( + 0.75 + ) + + +@pytest.mark.parametrize("seconds,expected", [(None, 15), (0, 15), (4, 15), (15, 15), (30, 30)]) +def test_live_webrtc_initialization_is_credited_against_duration(monkeypatch, seconds, expected): + monkeypatch.setitem( + litellm.model_cost, + "live-init-test", + { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.025, + }, + ) + events = [{"type": "litellm.live.initialization", "usage": {"seconds": 15}}] + if seconds is not None: + events.append({"type": "session.closed", "usage": {"seconds": seconds}}) + assert handle_realtime_stream_cost_calculation(events, Usage(), "openai", "live-init-test") == pytest.approx( + expected * 0.025 + ) + + +def test_live_invalid_terminal_retains_last_reported_partial_usage(monkeypatch): + monkeypatch.setitem( + litellm.model_cost, + "live-partial-test", + { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.025, + }, + ) + events = [ + {"type": "litellm.live.initialization", "usage": {"seconds": 15}}, + {"type": "session.usage.updated", "usage": {"seconds": 30}}, + {"type": "session.closed", "usage": {"seconds": "invalid"}}, + ] + assert handle_realtime_stream_cost_calculation(events, Usage(), "openai", "live-partial-test") == pytest.approx( + 0.75 + ) + + +@pytest.mark.parametrize( + "nested", + [ + {"type": "response.created", "response": {"model": "still-starting"}}, + {"type": "response.in_progress", "response": {}}, + {"type": "future.event", "response": ["unknown", "payload"]}, + {"type": "response.completed", "response": {"id": "broken", "usage": "invalid"}}, + ], +) +def test_live_partial_or_malformed_backend_events_preserve_duration(monkeypatch, nested): + from unittest.mock import MagicMock + + monkeypatch.setitem( + litellm.model_cost, + "live-resilient-test", + { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.025, + }, + ) + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + events = [ + {"type": "response.event", "event": nested}, + {"type": "session.closed", "usage": {"seconds": 30}}, + ] + usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(events) + assert usage.total_tokens == 0 + assert handle_realtime_stream_cost_calculation( + events, + usage, + "openai", + "live-resilient-test", + litellm_logging_obj=logger, + ) == pytest.approx(0.75) + assert bool(logger.model_call_details.get("realtime_backend_accounting_incomplete")) == ( + nested["type"] == "response.completed" + ) + + +@pytest.mark.parametrize("missing_usage_duplicate", [False, True]) +def test_live_missing_backend_price_preserves_duration_and_marks_accounting_incomplete( + monkeypatch, missing_usage_duplicate +): + from unittest.mock import MagicMock + + monkeypatch.setitem( + litellm.model_cost, + "live-resilient-test", + { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.025, + }, + ) + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + response = { + "id": "resp_unknown", + "created_at": 1, + "model": "unmapped-live-backend-price-test", + "output": [], + "usage": {"input_tokens": 20, "output_tokens": 10, "total_tokens": 30}, + } + events = [ + {"type": "response.event", "event": {"type": "response.completed", "response": response}}, + *( + [ + { + "type": "response.event", + "event": {"type": "response.completed", "response": {**response, "usage": None}}, + } + ] + if missing_usage_duplicate + else [] + ), + {"type": "session.closed", "usage": {"seconds": 30}}, + ] + usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(events) + assert usage.total_tokens == 30 + assert handle_realtime_stream_cost_calculation( + events, + usage, + "openai", + "live-resilient-test", + litellm_logging_obj=logger, + ) == pytest.approx(0.75) + assert logger.model_call_details["realtime_backend_accounting_incomplete"] is True + + +@pytest.mark.parametrize( + "envelope", + [ + {"type": "response.event"}, + {"type": "response.event", "event": {"response": {"id": "resp"}}}, + {"type": "response.event", "event": "not-an-object"}, + ], +) +def test_live_backend_malformed_envelope_is_dropped_without_accounting_flag(envelope): + """ + Malformed event envelopes (missing event, missing event.type, wrong shape) + are skipped silently: unlike a terminal response.completed that fails + response validation, they must not mark the call's accounting incomplete. + """ + from unittest.mock import MagicMock + + from litellm.cost_calculator import _live_backend_response + + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + assert _live_backend_response(envelope, logger) is None + assert "realtime_backend_accounting_incomplete" not in logger.model_call_details def test_completion_cost_prices_responses_websocket_turns_per_service_tier(): """Issue #41299: a session mixing default and priority turns must price each turn at its own returned service_tier, not the summed usage at a single tier.""" diff --git a/tests/unit/test_gpt_image_cost_calculator.py b/tests/unit/test_gpt_image_cost_calculator.py index d026285f9ae..5515e36d3e7 100644 --- a/tests/unit/test_gpt_image_cost_calculator.py +++ b/tests/unit/test_gpt_image_cost_calculator.py @@ -10,6 +10,8 @@ gpt-image-1 uses token-based pricing: - Image Output: $40.00/1M tokens """ +from typing import Final + import pytest import litellm @@ -37,6 +39,141 @@ def _use_local_model_cost_map(monkeypatch): class TestGPTImageCostCalculator: """Test the OpenAI gpt-image cost calculator""" + @pytest.mark.parametrize("family", ["flare", "sunburst"]) + @pytest.mark.parametrize("snapshot", ["", "-2026-09-08"]) + @pytest.mark.parametrize("call_type", ["image_generation", "image_edit"]) + @pytest.mark.parametrize("cached_text,cached_image", [(0, 0), (50, 500)]) + def test_image_25_official_prices(self, family, snapshot, call_type, cached_text, cached_image): + response: Final = ImageResponse( + created=1, + data=[], + usage={ + "input_tokens": 1100, + "output_tokens": 100, + "total_tokens": 1200, + "input_tokens_details": { + "text_tokens": 100, + "image_tokens": 1000, + "cached_tokens": cached_text + cached_image, + "cached_tokens_details": {"text_tokens": cached_text, "image_tokens": cached_image}, + }, + }, + ) + cost: Final = litellm.completion_cost( + model="gpt-image-2.5-" + family + snapshot, + completion_response=response, + call_type=call_type, + custom_llm_provider="openai", + ) + expected: Final = ( + (100 - cached_text) * 5e-6 + + cached_text * 1.25e-6 + + (1000 - cached_image) * 8e-6 + + cached_image * 2e-6 + + 100 * 30e-6 + ) + assert cost == pytest.approx(expected) + + def test_gpt_image_1_cost_with_text_only(self): + """Test cost calculation with only text input tokens""" + from litellm.llms.openai.image_generation.cost_calculator import cost_calculator + + usage = ImageUsage( + input_tokens=100, + output_tokens=5000, + total_tokens=5100, + input_tokens_details=ImageUsageInputTokensDetails( + text_tokens=100, + image_tokens=0, + ), + ) + + image_response = ImageResponse( + created=1234567890, + data=[ImageObject(url="http://example.com/image.jpg")], + ) + image_response.usage = usage + + cost = cost_calculator( + model="gpt-image-1", + image_response=image_response, + custom_llm_provider="openai", + ) + + # Expected cost: + # Text input: 100 * $5/1M = 0.0005 + # Image output: 5000 * $40/1M = 0.2 + # Total: 0.2005 + expected_cost = 0.0005 + 0.2 + assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}" + + def test_gpt_image_1_cost_with_image_input(self): + """Test cost calculation with both text and image input tokens (for edits)""" + from litellm.llms.openai.image_generation.cost_calculator import cost_calculator + + usage = ImageUsage( + input_tokens=600, + output_tokens=5000, + total_tokens=5600, + input_tokens_details=ImageUsageInputTokensDetails( + text_tokens=100, + image_tokens=500, + ), + ) + + image_response = ImageResponse( + created=1234567890, + data=[ImageObject(url="http://example.com/image.jpg")], + ) + image_response.usage = usage + + cost = cost_calculator( + model="gpt-image-1", + image_response=image_response, + custom_llm_provider="openai", + ) + + # Expected cost: + # Text input: 100 * $5/1M = 0.0005 + # Image input: 500 * $10/1M = 0.005 + # Image output: 5000 * $40/1M = 0.2 + # Total: 0.2055 + expected_cost = 0.0005 + 0.005 + 0.2 + assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}" + + def test_gpt_image_1_mini_cost(self): + """Test cost calculation for gpt-image-1-mini model""" + from litellm.llms.openai.image_generation.cost_calculator import cost_calculator + + usage = ImageUsage( + input_tokens=100, + output_tokens=5000, + total_tokens=5100, + input_tokens_details=ImageUsageInputTokensDetails( + text_tokens=100, + image_tokens=0, + ), + ) + + image_response = ImageResponse( + created=1234567890, + data=[ImageObject(url="http://example.com/image.jpg")], + ) + image_response.usage = usage + + cost = cost_calculator( + model="gpt-image-1-mini", + image_response=image_response, + custom_llm_provider="openai", + ) + + # Expected cost for gpt-image-1-mini: + # Text input: 100 * $2/1M = 0.0002 + # Image output: 5000 * $8/1M = 0.04 + # Total: 0.0402 + expected_cost = 0.0002 + 0.04 + assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}" + def test_gpt_image_1_cost_no_usage(self): """Test that cost returns 0 when no usage data is available""" from litellm.llms.openai.image_generation.cost_calculator import cost_calculator diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 4c36f99d2c8..7b1a7fbea9b 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -786,7 +786,6 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_read_input_token_cost_above_272k_tokens_flex": {"type": "number"}, "cache_read_input_token_cost_above_272k_tokens_ultrafast": {"type": "number"}, "cache_read_input_token_cost_above_512k_tokens": {"type": "number"}, - "input_cost_per_token_above_272k_tokens_ultrafast": {"type": "number"}, "cache_read_input_token_cost_batches": {"type": "number"}, "cache_read_input_token_cost_above_272k_tokens_batches": {"type": "number"}, "cache_creation_input_token_cost_above_1hr_above_200k_tokens": {"type": "number"}, @@ -809,6 +808,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "input_cost_per_token_above_200k_tokens_batches": {"type": "number"}, "input_cost_per_token_above_256k_tokens": {"type": "number"}, "input_cost_per_token_above_272k_tokens": {"type": "number"}, + "input_cost_per_token_above_272k_tokens_ultrafast": {"type": "number"}, "input_cost_per_token_above_512k_tokens": {"type": "number"}, "cache_read_input_token_cost_flex": {"type": "number"}, "cache_read_input_token_cost_priority": {"type": "number"}, @@ -1009,6 +1009,8 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "/v1/messages", "/v1/images/generations", "/v1/realtime", + "/v1/realtime/calls", + "/v1/live", "/v1/realtime/transcription_sessions", "/v1/images/variations", "/v1/images/edits", @@ -1289,7 +1291,11 @@ def test_openai_models_in_model_info(monkeypatch): model_map = litellm.model_cost violated_models = [] for model, info in model_map.items(): - if info.get("litellm_provider") == "openai" and info.get("supports_vision") is True: + if ( + info.get("litellm_provider") == "openai" + and info.get("supports_vision") is True + and info.get("mode") != "image_generation" + ): if info.get("supports_pdf_input") is not True: violated_models.append(model) assert len(violated_models) == 0, f"The following models should support pdf input: {violated_models}" diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index ab4f6c12431..e21c5d12d51 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -85,6 +85,9 @@ CONNECTION_NAMES: Final = ( "litellm_credential_name", "configurable_clientside_auth_params", "use_xai_oauth", + "chatgpt_auth_profile", + "chatgpt_auth_file", + "chatgpt_token_dir", "aws_batch_role_arn", "s3_bucket_name", "s3_region_name", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f709583c7dc..85c918312ba 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -8901,6 +8901,165 @@ export interface paths { patch?: never; trace?: never; }; + "/live": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: realtime_websocket_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_realtime_websocket_endpoint_get_5"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/live/": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: codex_live_sideband_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_codex_live_sideband_endpoint_get_2"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/live/sessions": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Create Live Session */ + post: operations["create_live_session_live_sessions_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/live/sessions/{session_id}/accept": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_live_sessions__session_id__accept_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/live/sessions/{session_id}/content": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Control Live Session */ + get: operations["control_live_session_live_sessions__session_id__content_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/live/sessions/{session_id}/fork": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Fork Live Session */ + post: operations["fork_live_session_live_sessions__session_id__fork_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/live/sessions/{session_id}/hangup": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_live_sessions__session_id__hangup_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/live/sessions/{session_id}/refer": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_live_sessions__session_id__refer_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/live/sessions/{session_id}/reject": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_live_sessions__session_id__reject_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/login": { parameters: { query?: never; @@ -10542,6 +10701,165 @@ export interface paths { patch?: never; trace?: never; }; + "/openai/v1/live": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: realtime_websocket_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_realtime_websocket_endpoint_get_4"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/openai/v1/live/": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: codex_live_sideband_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_codex_live_sideband_endpoint_get_3"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/openai/v1/live/sessions": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Create Live Session */ + post: operations["create_live_session_openai_v1_live_sessions_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/openai/v1/live/sessions/{session_id}/accept": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_openai_v1_live_sessions__session_id__accept_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/openai/v1/live/sessions/{session_id}/content": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Control Live Session */ + get: operations["control_live_session_openai_v1_live_sessions__session_id__content_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/openai/v1/live/sessions/{session_id}/fork": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Fork Live Session */ + post: operations["fork_live_session_openai_v1_live_sessions__session_id__fork_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/openai/v1/live/sessions/{session_id}/hangup": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_openai_v1_live_sessions__session_id__hangup_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/openai/v1/live/sessions/{session_id}/refer": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_openai_v1_live_sessions__session_id__refer_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/openai/v1/live/sessions/{session_id}/reject": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_openai_v1_live_sessions__session_id__reject_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/openai/v1/realtime": { parameters: { query?: never; @@ -19703,6 +20021,165 @@ export interface paths { patch?: never; trace?: never; }; + "/v1/live": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: realtime_websocket_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_realtime_websocket_endpoint_get_6"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/v1/live/": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: codex_live_sideband_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_codex_live_sideband_endpoint_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/v1/live/sessions": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Create Live Session */ + post: operations["create_live_session_v1_live_sessions_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/v1/live/sessions/{session_id}/accept": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_v1_live_sessions__session_id__accept_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/v1/live/sessions/{session_id}/content": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Control Live Session */ + get: operations["control_live_session_v1_live_sessions__session_id__content_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/v1/live/sessions/{session_id}/fork": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Fork Live Session */ + post: operations["fork_live_session_v1_live_sessions__session_id__fork_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/v1/live/sessions/{session_id}/hangup": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_v1_live_sessions__session_id__hangup_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/v1/live/sessions/{session_id}/refer": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_v1_live_sessions__session_id__refer_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/v1/live/sessions/{session_id}/reject": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Control Live Session */ + post: operations["control_live_session_v1_live_sessions__session_id__reject_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/mcp/access_groups": { parameters: { query?: never; @@ -60699,6 +61176,248 @@ export interface operations { }; }; }; + websocket_realtime_websocket_endpoint_get_5: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; + websocket_codex_live_sideband_endpoint_get_2: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; + create_live_session_live_sessions_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + control_live_session_live_sessions__session_id__accept_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_live_sessions__session_id__content_get: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + fork_live_session_live_sessions__session_id__fork_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_live_sessions__session_id__hangup_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_live_sessions__session_id__refer_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_live_sessions__session_id__reject_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; login_login_post: { parameters: { query?: never; @@ -63243,6 +63962,248 @@ export interface operations { }; }; }; + websocket_realtime_websocket_endpoint_get_4: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; + websocket_codex_live_sideband_endpoint_get_3: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; + create_live_session_openai_v1_live_sessions_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + control_live_session_openai_v1_live_sessions__session_id__accept_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_openai_v1_live_sessions__session_id__content_get: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + fork_live_session_openai_v1_live_sessions__session_id__fork_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_openai_v1_live_sessions__session_id__hangup_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_openai_v1_live_sessions__session_id__refer_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_openai_v1_live_sessions__session_id__reject_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; websocket_realtime_websocket_endpoint_get_3: { parameters: { query?: never; @@ -74488,6 +75449,248 @@ export interface operations { }; }; }; + websocket_realtime_websocket_endpoint_get_6: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; + websocket_codex_live_sideband_endpoint_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; + create_live_session_v1_live_sessions_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + control_live_session_v1_live_sessions__session_id__accept_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_v1_live_sessions__session_id__content_get: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + fork_live_session_v1_live_sessions__session_id__fork_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_v1_live_sessions__session_id__hangup_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_v1_live_sessions__session_id__refer_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + control_live_session_v1_live_sessions__session_id__reject_post: { + parameters: { + query?: never; + header?: never; + path: { + session_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; get_mcp_access_groups_v1_mcp_access_groups_get: { parameters: { query?: never;