From bce8aa1dcc4a224a6cd38f64d8c98700c437f696 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 07:20:37 +0200 Subject: [PATCH 01/90] feat(chatgpt): support Codex image and realtime routes --- docs/chatgpt-codex-routes.md | 52 +++++ litellm/images/main.py | 19 +- litellm/llms/chatgpt/images.py | 137 +++++++++++++ litellm/llms/chatgpt/realtime.py | 124 ++++++++++++ .../llms/chatgpt/responses/transformation.py | 1 + litellm/llms/custom_httpx/llm_http_handler.py | 9 +- litellm/proxy/_types.py | 3 + litellm/proxy/proxy_server.py | 24 ++- litellm/proxy/realtime_endpoints/codex.py | 188 ++++++++++++++++++ litellm/proxy/realtime_endpoints/endpoints.py | 5 + litellm/realtime_api/main.py | 33 ++- litellm/types/realtime.py | 1 + litellm/utils.py | 8 + tests/test_litellm/llms/chatgpt/conftest.py | 22 ++ .../test_chatgpt_responses_transformation.py | 23 +++ .../test_litellm/llms/chatgpt/test_images.py | 106 ++++++++++ .../llms/chatgpt/test_realtime.py | 57 ++++++ .../proxy/auth/test_route_checks.py | 18 ++ .../proxy/realtime_endpoints/test_codex.py | 76 +++++++ tests/test_litellm/realtime_api/test_main.py | 4 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 114 +++++++++++ 21 files changed, 1010 insertions(+), 14 deletions(-) create mode 100644 docs/chatgpt-codex-routes.md create mode 100644 litellm/llms/chatgpt/images.py create mode 100644 litellm/llms/chatgpt/realtime.py create mode 100644 litellm/proxy/realtime_endpoints/codex.py create mode 100644 tests/test_litellm/llms/chatgpt/conftest.py create mode 100644 tests/test_litellm/llms/chatgpt/test_images.py create mode 100644 tests/test_litellm/llms/chatgpt/test_realtime.py create mode 100644 tests/test_litellm/proxy/realtime_endpoints/test_codex.py diff --git a/docs/chatgpt-codex-routes.md b/docs/chatgpt-codex-routes.md new file mode 100644 index 00000000000..98cc1473385 --- /dev/null +++ b/docs/chatgpt-codex-routes.md @@ -0,0 +1,52 @@ +# ChatGPT OAuth routes used by Codex + +The ChatGPT provider supports image generation and editing, structured Responses output, Realtime WebSockets, and direct WebRTC offer exchange using the existing ChatGPT OAuth credentials configured with `CHATGPT_TOKEN_DIR` and `CHATGPT_AUTH_FILE` + +## Models and transports + +| Model | Proxy route | Upstream route | +| --- | --- | --- | +| `gpt-image-2` | `/v1/images/generations` | ChatGPT `/backend-api/codex/images/generations` | +| `gpt-image-2` | `/v1/images/edits` | ChatGPT `/backend-api/codex/images/edits` | +| `codex-auto-review` | `/v1/responses` | ChatGPT `/backend-api/codex/responses` | +| `gpt-realtime-1.5` | `/v1/realtime` WebSocket | OpenAI `/v1/realtime` with ChatGPT OAuth | +| `gpt-live-1-codex` | `/v1/realtime/calls`, then `/v1/live/{call_id}` | ChatGPT call signaling, OpenAI live sideband | +| `gpt-4o-mini-transcribe` | Nested in Realtime `session.audio.input.transcription.model` | Same Realtime session | + +Codex uses `gpt-image-2` for both image routes. There is no separate `gpt-image-2-edit` model in that contract. Memory extraction and consolidation use ordinary Responses models and need no additional transport + +Register deployments with `litellm_params.model: chatgpt/`. These auxiliary models do not need to appear in the conversation model selector. This contribution uses the existing single-account ChatGPT authentication contract + +## Images + +Both synchronous and asynchronous SDK calls support image generation and editing. Edits accept the usual `image` file input or Codex's JSON `images` array, containing one to five `{ "image_url": "data:image/png;base64,..." }` references. PNG, JPEG, WEBP, and HTTPS references are accepted. Masks and mixing `image` with `images` are rejected + +```python +await litellm.aimage_edit( + model="chatgpt/gpt-image-2", + prompt="Change the blue circle to red", + images=[{"image_url": image_data_url}], +) +``` + +## Guardian + +The Responses transformation preserves `text.format`, including strict JSON schemas. Custom Codex providers use `/responses` for `codex-auto-review`. The separate native `/guardian` route can be disabled by the upstream service; registering an alias does not enable that route + +## Realtime and GPT-Live + +Standard Realtime WebSockets use the OpenAI host with the selected ChatGPT access token and account header. Client protocol headers are forwarded without forwarding the client's proxy authorization header + +Codex WebRTC offers can use JSON or multipart `sdp` and `session` fields with an ordinary LiteLLM key. Signaling preserves the session shape and the `intent` and `architecture` query parameters. Existing raw SDP requests authenticated with an encrypted ephemeral key retain their existing route + +For GPT-Live, send `OpenAI-Alpha: quicksilver=v2`, `intent=quicksilver&architecture=avas`, and a Codex voice such as `sol`. Do not add the GA `session.type: realtime` field to a Frameless Bidi session. A generic `Voice session access denied` error can mean an unsupported voice; it is not sufficient evidence of missing account entitlement + +The returned `Location` contains an encrypted call identifier. The sideband connection must use the same LiteLLM bearer key and retain access to the requested alias. The identifier expires after one hour and remains valid across proxy workers sharing the same salt key. OAuth tokens are never returned to clients + +The direct `/v1/live?model=gpt-live-1-codex` WebSocket is forwarded, but upstream access can differ from WebRTC. A successful WebRTC call does not establish permission for direct live sessions. Upstream errors remain visible; the proxy does not replace the model, voice, or transport silently + +## Verification + +The integration was exercised with real image generation and JSON edits, strict Guardian JSON output, Realtime text and audio output, audio transcription, and a GPT-Live WebRTC session with a successful sideband context acknowledgment. Expired, malformed, tampered, and wrong-owner call identifiers have regression coverage + +References: [Codex source](https://github.com/openai/codex), [OpenClaw voice authentication](https://docs.openclaw.ai/providers/openai/voice-and-speech), and [Pi Codex signaling implementation](https://github.com/monotykamary/pi-better-openai/blob/main/src/live/transport.ts). These upstream capabilities may change independently of LiteLLM diff --git a/litellm/images/main.py b/litellm/images/main.py index 6a94e7c8df2..9a6f18f8ea1 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -102,7 +102,9 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse: ctx: Final = contextvars.copy_context() func_with_context: Final = partial(ctx.run, func) - _, custom_llm_provider, _, _ = get_llm_provider(model=model, api_base=kwargs.get("api_base", None)) + _, custom_llm_provider, _, _ = get_llm_provider( + model=model, api_base=kwargs.get("api_base", None), litellm_params=GenericLiteLLMParams(**kwargs) + ) # Await normally init_response: Final = await loop.run_in_executor(None, func_with_context) @@ -227,6 +229,7 @@ def image_generation( model=model, custom_llm_provider=custom_llm_provider, api_base=api_base, + litellm_params=GenericLiteLLMParams(**kwargs), ) else: model = "dall-e-2" @@ -377,6 +380,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, @@ -785,6 +789,7 @@ def image_edit( model, custom_llm_provider, _, _ = get_llm_provider( model=model or DEFAULT_IMAGE_ENDPOINT_MODEL, custom_llm_provider=custom_llm_provider, + litellm_params=litellm_params, ) # Check for custom provider @@ -965,9 +970,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,14 +1007,12 @@ async def aimage_edit( # get custom llm provider so we can use this for mapping exceptions if custom_llm_provider is None: _, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, api_base=local_vars.get("base_url", None) + model=model, api_base=local_vars.get("base_url", None), litellm_params=GenericLiteLLMParams(**kwargs) ) - 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/llms/chatgpt/images.py b/litellm/llms/chatgpt/images.py new file mode 100644 index 00000000000..1d296ff0fe6 --- /dev/null +++ b/litellm/llms/chatgpt/images.py @@ -0,0 +1,137 @@ +import base64 +from collections.abc import Mapping, Sequence +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.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 .common_utils import CHATGPT_API_BASE +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) -> dict[str, str]: # mutable-ok: image handler requires dictionaries + content: Final = file[1] if isinstance(file, tuple) else file + content_type: Final = file[2] if isinstance(file, tuple) and len(file) >= 3 else "image/png" + if content_type not in ("image/png", "image/jpeg", "image/webp"): + raise ValueError("Reference images must be PNG, JPEG, or WEBP") + raw: Final = ( + content.encode() if isinstance(content, str) else content if isinstance(content, bytes) else content.read() + ) + 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 {**headers, **auth_headers, "accept": "application/json"} # mutable-ok: JSON request serialization + + +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 CHATGPT_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 CHATGPT_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 | 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 + "model": model, + "prompt": prompt, + **image_edit_optional_request_params, + "images": tuple(item.model_dump() for item in validated), + }, () + + data, files = super().transform_image_edit_request( + model, + prompt, + image, + dict(image_edit_optional_request_params), # mutable-ok: parent edit adapter requires dictionaries + litellm_params, + dict(headers), # mutable-ok: parent edit adapter requires dictionaries + ) + parts: Final = files.items() if isinstance(files, Mapping) else files + encoded: Final = tuple(encode_reference(file) for field, file in parts if field == "image[]") + if not 1 <= len(encoded) <= 5: + raise ValueError("images must contain between 1 and 5 reference images") + return {**data, "images": encoded}, () # mutable-ok: JSON request serialization diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py new file mode 100644 index 00000000000..66a703eb616 --- /dev/null +++ b/litellm/llms/chatgpt/realtime.py @@ -0,0 +1,124 @@ +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final + +from httpx import URL +from pydantic import TypeAdapter + +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 .common_utils import CHATGPT_API_BASE +from .responses.transformation import ChatGPTResponsesAPIConfig + + +def realtime_headers( + params: GenericLiteLLMParams, headers: Mapping[str, str] +) -> 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, + } + + +class ChatGPTRealtime(OpenAIRealtime): + def __init__(self, params: GenericLiteLLMParams, headers: Mapping[str, str]) -> None: + super().__init__() + self._profile_headers = realtime_headers(params, headers) + self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None)) + + 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 = "live" if query_params.get("model") == "gpt-live-1-codex" else "realtime" + if self._call_id: + return str( + base.copy_with( + scheme="wss" if base.scheme in ("https", "wss") else "ws", + path=f"{base.path.rstrip('/')}/{endpoint}/{self._call_id}" + if endpoint == "live" + else f"{base.path.rstrip('/')}/realtime", + params=() if endpoint == "live" else (("call_id", self._call_id),), + ) + ) + return str( + base.copy_with( + scheme="wss" if base.scheme in ("https", "wss") else "ws", + path=f"{base.path.rstrip('/')}/{endpoint}", + params=query_params, + ) + ) + + +class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): + realtime_calls_json: Final = True + + def __init__(self, params: GenericLiteLLMParams) -> None: + self._params = params + + def get_api_base( + self, + api_base: str | None, + **kwargs: object, # kwargs-ok: provider interface accepts optional credentials + ) -> str: + return api_base or CHATGPT_API_BASE + + 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 = TypeAdapter(Mapping[str, str]).validate_python( + getattr(self._params, "extra_query", None) or MappingProxyType({}) + ) + return str(URL(f"{self.get_api_base(api_base).rstrip('/')}/realtime/calls", params=query)) + + def get_realtime_calls_headers( + self, ephemeral_key: str + ) -> dict[str, str]: # mutable-ok: HTTP handler header contract + 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 "https://api.openai.com/v1/realtime/client_secrets" + + def get_transcription_session_url( + self, + api_base: str | None, + model: str, + api_version: str | None = None, + ) -> str: + return "https://api.openai.com/v1/realtime/transcription_sessions" diff --git a/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index b96e06be3d8..1a5c6302a01 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -102,6 +102,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 7587b963a38..30f12e3a64d 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6411,7 +6411,7 @@ class BaseLLMHTTPHandler: # 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 @@ -6434,6 +6434,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, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d746cfd38d8..19fd19e5c8a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -399,6 +399,9 @@ class LiteLLMRoutes(enum.Enum): "/realtime?{model}", "/v1/realtime?{model}", "/openai/v1/realtime?{model}", + "/live", + "/v1/live", + "/v1/live/{call_id}", # realtime (GA WebRTC HTTP routes) "/realtime/client_secrets", "/v1/realtime/client_secrets", diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7219b373dc3..c8e3cffa521 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11592,6 +11592,19 @@ async def _reject_realtime_session( await _release_realtime_budget_reservation(user_api_key_dict) +@app.websocket("/v1/live/{call_id}") +async def codex_live_sideband_endpoint( + websocket: WebSocket, + call_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), +) -> None: + from litellm.proxy.realtime_endpoints.codex 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/realtime") @app.websocket("/v1/realtime") @app.websocket("/realtime") @@ -11599,12 +11612,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 and call_id.startswith("rtc_litellm_"): + from litellm.proxy.realtime_endpoints.codex 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() ] @@ -11635,7 +11654,10 @@ 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(_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/codex.py b/litellm/proxy/realtime_endpoints/codex.py new file mode 100644 index 00000000000..669b56d4371 --- /dev/null +++ b/litellm/proxy/realtime_endpoints/codex.py @@ -0,0 +1,188 @@ +import base64 +import hashlib +import json +import time +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import httpx +from fastapi import HTTPException, Request, Response, WebSocket +from pydantic import BaseModel, Field + +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation +from litellm.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 + alias: str + owner: str + expires_at: float + + +class ChatGPTCallRouting(BaseModel): + model: str + + +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()).hexdigest(): + raise HTTPException(403, "Invalid or expired realtime call") + return call + + +async def read_codex_offer(request: Request) -> CodexRealtimeOffer: + if request.headers.get("content-type", "").startswith("multipart/form-data"): + 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 create_codex_realtime_call(request: Request) -> Response: + from litellm.proxy import proxy_server as server + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + 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="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + custom_litellm_key_header=None, + ) + try: + await can_key_call_resolved_model( + model=model, + llm_model_list=server.llm_model_list, + valid_token=auth, + llm_router=server.llm_router, + ) + data: Final = { # mutable-ok: proxy processor enriches request data + "model": model, + "sdp_body": offer.sdp.encode(), + "session": offer.session.model_dump(exclude_none=True), + "openai_ephemeral_key": "", + "extra_query": { # mutable-ok: router request parameters + key: value for key, value in request.query_params.items() if key in ("intent", "architecture") + }, + "extra_headers": { # mutable-ok: router request headers + key: value + for key, value in request.headers.items() + if key in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation") + }, + } + processor: Final = ProxyBaseLLMRequestProcessing(data=data) + processed, _ = 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, + user_model=server.user_model, + user_temperature=server.user_temperature, + user_request_timeout=server.user_request_timeout, + user_max_tokens=server.user_max_tokens, + user_api_base=server.user_api_base, + model=model, + route_type="arealtime_calls", + ) + result: Final = await server.route_request( + data=processed, + route_type="arealtime_calls", + llm_router=server.llm_router, + user_model=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") + routing_data: Final = response.extensions.get("chatgpt_realtime") + if response.is_error: + return Response(response.content, status_code=response.status_code, media_type="application/json") + if not routing_data: + raise HTTPException(400, "Direct call signaling requires a ChatGPT deployment") + routing: Final = ChatGPTCallRouting.model_validate(routing_data) + call_id: Final = urlsplit(response.headers.get("location", "")).path.rstrip("/").rsplit("/", 1)[-1] + call: Final = CodexRealtimeCall( + call_id=call_id, + model=routing.model, + alias=model, + owner=hashlib.sha256(request.headers.get("authorization", "").encode()).hexdigest(), + expires_at=time.time() + 3600, + ) + token: Final = encode_call(call) + return Response( + response.content, + status_code=response.status_code, + media_type="application/sdp", + headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), + ) + finally: + 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 + + try: + try: + call: Final = decode_call(token, websocket.headers.get("authorization", "")) + await can_key_call_resolved_model( + model=call.alias, + llm_model_list=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 + await websocket.accept() + query: Final[RealtimeQueryParams] = {"model": call.model} + await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # internal proxy entrypoint for an already authorized call + model=f"chatgpt/{call.model}", + websocket=websocket, + chatgpt_realtime_call_id=call.call_id, + query_params=query, + user_api_key_dict=auth, + ) + finally: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index be2ac2ff33e..c7a20302a13 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -375,6 +375,11 @@ async def proxy_realtime_calls( request: Request, fastapi_response: Response, ) -> Response: + if request.headers.get("content-type", "").split(";", 1)[0] in ("application/json", "multipart/form-data"): + from litellm.proxy.realtime_endpoints.codex 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, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index b824a5928c6..32ae1b34485 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -89,7 +89,11 @@ def _get_realtime_http_provider_config( ) provider_config: BaseRealtimeHTTPConfig | None = None - if custom_llm_provider in LlmProviders._member_map_.values(): + if custom_llm_provider == "chatgpt": + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + + provider_config = ChatGPTRealtimeHTTPConfig(litellm_params) + elif custom_llm_provider in LlmProviders._member_map_.values(): provider_config = ProviderConfigManager.get_provider_realtime_http_config( model="", provider=LlmProviders(custom_llm_provider), @@ -137,6 +141,7 @@ async def acreate_realtime_client_secret( model=model_name, api_base=litellm_params.api_base, api_key=litellm_params.api_key, + litellm_params=litellm_params, ) ( provider_config, @@ -205,6 +210,7 @@ async def acreate_realtime_transcription_session( model=model_name, api_base=litellm_params.api_base, api_key=litellm_params.api_key, + litellm_params=litellm_params, ) ( provider_config, @@ -266,6 +272,7 @@ async def arealtime_calls( model=model_name, api_base=litellm_params.api_base, api_key=litellm_params.api_key, + litellm_params=litellm_params, ) provider_config, resolved_api_base, _ = _get_realtime_http_provider_config( custom_llm_provider=custom_llm_provider, @@ -282,7 +289,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, @@ -295,6 +302,13 @@ async def arealtime_calls( client=kwargs.get("client"), api_version=litellm_params.api_version, ) + if custom_llm_provider == "chatgpt": + response.extensions["chatgpt_realtime"] = MappingProxyType( + { + "model": model_name, + } + ) + return response async def vertex_access_token_resolver( @@ -366,6 +380,7 @@ async def _arealtime( model=model, api_base=api_base, api_key=api_key, + litellm_params=litellm_params, ) # If the client supplied `model` in the URL, ensure it uses the normalized @@ -439,6 +454,20 @@ async def _arealtime( user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata(kwargs), ) + elif _custom_llm_provider == "chatgpt": + from litellm.llms.chatgpt.realtime import ChatGPTRealtime + + await ChatGPTRealtime(litellm_params, websocket.headers).async_realtime( + model=model, + websocket=websocket, + logging_obj=litellm_logging_obj, + api_base=api_base or "https://api.openai.com/v1", + api_key="chatgpt-oauth", + timeout=timeout, + query_params=query_params, + user_api_key_dict=kwargs.get("user_api_key_dict"), + litellm_metadata=_build_litellm_metadata(kwargs), + ) elif _custom_llm_provider == "openai": api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or "https://api.openai.com/" # set API KEY diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index 17dc70126f3..ae3a3c6ef2f 100644 --- a/litellm/types/realtime.py +++ b/litellm/types/realtime.py @@ -49,6 +49,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 36b48d3b8d8..a3142cb54fe 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9040,6 +9040,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, @@ -9237,6 +9241,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/tests/test_litellm/llms/chatgpt/conftest.py b/tests/test_litellm/llms/chatgpt/conftest.py new file mode 100644 index 00000000000..fb5b4cd5a1c --- /dev/null +++ b/tests/test_litellm/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/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py b/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py index 8e0415d50de..aa69e45cada 100644 --- a/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py +++ b/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py @@ -19,6 +19,29 @@ from litellm.llms.chatgpt.responses.transformation import ChatGPTResponsesAPICon 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/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py new file mode 100644 index 00000000000..4ee9aae4717 --- /dev/null +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -0,0 +1,106 @@ +import base64 + +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.router import GenericLiteLLMParams + + +def test_generation_routes_with_chatgpt_oauth(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)) + result = litellm.image_generation( + model="chatgpt/gpt-image-2", + prompt="blue circle", + client=client, + quality="auto", + size="auto", + background="auto", + ) + assert result.data[0].b64_json == "aGVsbG8=" + assert str(requests[0].url) == "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", + 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 + + +@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"}, "gpt-image-2", [], {}, {"chatgpt_token_dir": chatgpt_tokens} + ) + assert headers["Authorization"] == "Bearer test-token-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", + 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() diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py new file mode 100644 index 00000000000..7ce8d5f298a --- /dev/null +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -0,0 +1,57 @@ +import json + +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 + + +@pytest.mark.asyncio +async def test_chatgpt_call_keeps_oauth_and_frameless_session(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)) + response = await litellm.arealtime_calls( + model="chatgpt/gpt-live-1-codex", + 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"}, + extra_headers={"openai-alpha": "quicksilver=v2"}, + client=client, + ) + assert response.status_code == 201 + 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 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.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): + 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" diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 0fd19f518d9..d8e01982717 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/realtime_endpoints/test_codex.py b/tests/test_litellm/proxy/realtime_endpoints/test_codex.py new file mode 100644 index 00000000000..fdba89daaf1 --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_codex.py @@ -0,0 +1,76 @@ +import hashlib +import time + +import pytest +from fastapi import HTTPException, WebSocket + +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.realtime_endpoints import codex +from litellm.proxy.realtime_endpoints.codex import CodexRealtimeCall, decode_call, encode_call + + +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", + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() + 300, + ) + token = encode_call(call) + assert "/" 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"}] diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index 761e87ac764..b45812f45d7 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -61,7 +61,7 @@ def _run_client_secret(session, model, monkeypatch): captured.update(kwargs) return object() - def mock_get_llm_provider(model, api_base, api_key): + def mock_get_llm_provider(model, api_base, api_key, litellm_params=None): return model, "openai", None, api_base monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) @@ -161,7 +161,7 @@ async def test_arealtime_vertex_branch_resolves_credentials_under_a_bound(monkey async def hanging_token_refresh(**kwargs): await asyncio.sleep(30) - def mock_get_llm_provider(model, api_base, api_key): + def mock_get_llm_provider(model, api_base, api_key, litellm_params=None): return model, "vertex_ai", None, api_base monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 83b0d58f2b2..a77332f1dcb 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -8276,6 +8276,26 @@ 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_4"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/login": { parameters: { query?: never; @@ -18396,6 +18416,46 @@ 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_5"]; + 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"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/mcp/access_groups": { parameters: { query?: never; @@ -50644,6 +50704,24 @@ 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; + }; + }; + }; login_login_post: { parameters: { query?: never; @@ -63000,6 +63078,42 @@ 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: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; get_mcp_access_groups_v1_mcp_access_groups_get: { parameters: { query?: never; From 58a56577844f629b9835d62edda89ede5234f32b Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 07:58:59 +0200 Subject: [PATCH 02/90] fix(chatgpt): validate sideband ownership and normalize image files --- docs/chatgpt-codex-routes.md | 52 ----- gateway/routes/allowlist.py | 2 + litellm/images/main.py | 8 +- .../chatgpt}/codex.py | 0 litellm/llms/chatgpt/images.py | 46 ++-- litellm/proxy/proxy_server.py | 6 +- litellm/proxy/realtime_endpoints/endpoints.py | 2 +- litellm/realtime_api/main.py | 4 - tests/test_litellm/llms/chatgpt/test_codex.py | 212 ++++++++++++++++++ .../test_litellm/llms/chatgpt/test_images.py | 12 + .../proxy/realtime_endpoints/test_codex.py | 76 ------- tests/test_litellm/realtime_api/test_main.py | 4 +- 12 files changed, 263 insertions(+), 161 deletions(-) delete mode 100644 docs/chatgpt-codex-routes.md rename litellm/{proxy/realtime_endpoints => llms/chatgpt}/codex.py (100%) create mode 100644 tests/test_litellm/llms/chatgpt/test_codex.py delete mode 100644 tests/test_litellm/proxy/realtime_endpoints/test_codex.py diff --git a/docs/chatgpt-codex-routes.md b/docs/chatgpt-codex-routes.md deleted file mode 100644 index 98cc1473385..00000000000 --- a/docs/chatgpt-codex-routes.md +++ /dev/null @@ -1,52 +0,0 @@ -# ChatGPT OAuth routes used by Codex - -The ChatGPT provider supports image generation and editing, structured Responses output, Realtime WebSockets, and direct WebRTC offer exchange using the existing ChatGPT OAuth credentials configured with `CHATGPT_TOKEN_DIR` and `CHATGPT_AUTH_FILE` - -## Models and transports - -| Model | Proxy route | Upstream route | -| --- | --- | --- | -| `gpt-image-2` | `/v1/images/generations` | ChatGPT `/backend-api/codex/images/generations` | -| `gpt-image-2` | `/v1/images/edits` | ChatGPT `/backend-api/codex/images/edits` | -| `codex-auto-review` | `/v1/responses` | ChatGPT `/backend-api/codex/responses` | -| `gpt-realtime-1.5` | `/v1/realtime` WebSocket | OpenAI `/v1/realtime` with ChatGPT OAuth | -| `gpt-live-1-codex` | `/v1/realtime/calls`, then `/v1/live/{call_id}` | ChatGPT call signaling, OpenAI live sideband | -| `gpt-4o-mini-transcribe` | Nested in Realtime `session.audio.input.transcription.model` | Same Realtime session | - -Codex uses `gpt-image-2` for both image routes. There is no separate `gpt-image-2-edit` model in that contract. Memory extraction and consolidation use ordinary Responses models and need no additional transport - -Register deployments with `litellm_params.model: chatgpt/`. These auxiliary models do not need to appear in the conversation model selector. This contribution uses the existing single-account ChatGPT authentication contract - -## Images - -Both synchronous and asynchronous SDK calls support image generation and editing. Edits accept the usual `image` file input or Codex's JSON `images` array, containing one to five `{ "image_url": "data:image/png;base64,..." }` references. PNG, JPEG, WEBP, and HTTPS references are accepted. Masks and mixing `image` with `images` are rejected - -```python -await litellm.aimage_edit( - model="chatgpt/gpt-image-2", - prompt="Change the blue circle to red", - images=[{"image_url": image_data_url}], -) -``` - -## Guardian - -The Responses transformation preserves `text.format`, including strict JSON schemas. Custom Codex providers use `/responses` for `codex-auto-review`. The separate native `/guardian` route can be disabled by the upstream service; registering an alias does not enable that route - -## Realtime and GPT-Live - -Standard Realtime WebSockets use the OpenAI host with the selected ChatGPT access token and account header. Client protocol headers are forwarded without forwarding the client's proxy authorization header - -Codex WebRTC offers can use JSON or multipart `sdp` and `session` fields with an ordinary LiteLLM key. Signaling preserves the session shape and the `intent` and `architecture` query parameters. Existing raw SDP requests authenticated with an encrypted ephemeral key retain their existing route - -For GPT-Live, send `OpenAI-Alpha: quicksilver=v2`, `intent=quicksilver&architecture=avas`, and a Codex voice such as `sol`. Do not add the GA `session.type: realtime` field to a Frameless Bidi session. A generic `Voice session access denied` error can mean an unsupported voice; it is not sufficient evidence of missing account entitlement - -The returned `Location` contains an encrypted call identifier. The sideband connection must use the same LiteLLM bearer key and retain access to the requested alias. The identifier expires after one hour and remains valid across proxy workers sharing the same salt key. OAuth tokens are never returned to clients - -The direct `/v1/live?model=gpt-live-1-codex` WebSocket is forwarded, but upstream access can differ from WebRTC. A successful WebRTC call does not establish permission for direct live sessions. Upstream errors remain visible; the proxy does not replace the model, voice, or transport silently - -## Verification - -The integration was exercised with real image generation and JSON edits, strict Guardian JSON output, Realtime text and audio output, audio transcription, and a GPT-Live WebRTC session with a successful sideband context acknowledgment. Expired, malformed, tampered, and wrong-owner call identifiers have regression coverage - -References: [Codex source](https://github.com/openai/codex), [OpenClaw voice authentication](https://docs.openclaw.ai/providers/openai/voice-and-speech), and [Pi Codex signaling implementation](https://github.com/monotykamary/pi-better-openai/blob/main/src/live/transport.ts). These upstream capabilities may change independently of LiteLLM diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 92b73867e67..2817792f27c 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -104,6 +104,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = ( "/{provider}/", "/toolset/", # Realtime / streaming + "/v1/live", + "/live", "/v1/realtime", "/realtime", # Health & ops diff --git a/litellm/images/main.py b/litellm/images/main.py index 9a6f18f8ea1..7290e537369 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -102,9 +102,7 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse: ctx: Final = contextvars.copy_context() func_with_context: Final = partial(ctx.run, func) - _, custom_llm_provider, _, _ = get_llm_provider( - model=model, api_base=kwargs.get("api_base", None), litellm_params=GenericLiteLLMParams(**kwargs) - ) + _, custom_llm_provider, _, _ = get_llm_provider(model=model, api_base=kwargs.get("api_base", None)) # Await normally init_response: Final = await loop.run_in_executor(None, func_with_context) @@ -229,7 +227,6 @@ def image_generation( model=model, custom_llm_provider=custom_llm_provider, api_base=api_base, - litellm_params=GenericLiteLLMParams(**kwargs), ) else: model = "dall-e-2" @@ -789,7 +786,6 @@ def image_edit( model, custom_llm_provider, _, _ = get_llm_provider( model=model or DEFAULT_IMAGE_ENDPOINT_MODEL, custom_llm_provider=custom_llm_provider, - litellm_params=litellm_params, ) # Check for custom provider @@ -1007,7 +1003,7 @@ async def aimage_edit( # get custom llm provider so we can use this for mapping exceptions if custom_llm_provider is None: _, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, api_base=local_vars.get("base_url", None), litellm_params=GenericLiteLLMParams(**kwargs) + model=model, api_base=local_vars.get("base_url", None) ) func: Final = partial( diff --git a/litellm/proxy/realtime_endpoints/codex.py b/litellm/llms/chatgpt/codex.py similarity index 100% rename from litellm/proxy/realtime_endpoints/codex.py rename to litellm/llms/chatgpt/codex.py diff --git a/litellm/llms/chatgpt/images.py b/litellm/llms/chatgpt/images.py index 1d296ff0fe6..ee6ba17e340 100644 --- a/litellm/llms/chatgpt/images.py +++ b/litellm/llms/chatgpt/images.py @@ -1,5 +1,7 @@ import base64 +import os from collections.abc import Mapping, Sequence +from pathlib import Path from types import MappingProxyType from typing import Final @@ -7,6 +9,7 @@ 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 @@ -21,14 +24,26 @@ class ReferenceImage(BaseModel): image_url: str = Field(pattern=r"^(data:image/(png|jpeg|webp);base64,|https://)") -def encode_reference(file: HTTPFileTypes) -> dict[str, str]: # mutable-ok: image handler requires dictionaries +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 - content_type: Final = file[2] if isinstance(file, tuple) and len(file) >= 3 else "image/png" + 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") - raw: Final = ( - content.encode() if isinstance(content, str) else content if isinstance(content, bytes) else content.read() - ) return { # mutable-ok: JSON request serialization "image_url": f"data:{content_type};base64," + base64.b64encode(raw).decode("ascii") } @@ -101,7 +116,7 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig): self, model: str, prompt: str | None, - image: FileTypes | None, + image: FileTypes | Sequence[FileTypes] | None, image_edit_optional_request_params: Mapping[str, object], litellm_params: GenericLiteLLMParams, headers: Mapping[str, object], @@ -122,16 +137,13 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig): "images": tuple(item.model_dump() for item in validated), }, () - data, files = super().transform_image_edit_request( - model, - prompt, - image, - dict(image_edit_optional_request_params), # mutable-ok: parent edit adapter requires dictionaries - litellm_params, - dict(headers), # mutable-ok: parent edit adapter requires dictionaries - ) - parts: Final = files.items() if isinstance(files, Mapping) else files - encoded: Final = tuple(encode_reference(file) for field, file in parts if field == "image[]") + 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 {**data, "images": encoded}, () # mutable-ok: JSON request serialization + return { # mutable-ok: JSON request serialization + "model": model, + "prompt": prompt, + **image_edit_optional_request_params, + "images": encoded, + }, () diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c8e3cffa521..9d55cf21d81 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11598,7 +11598,7 @@ async def codex_live_sideband_endpoint( call_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), ) -> None: - from litellm.proxy.realtime_endpoints.codex import codex_realtime_sideband + from litellm.llms.chatgpt.codex import codex_realtime_sideband await codex_realtime_sideband(websocket, call_id, user_api_key_dict) @@ -11619,8 +11619,8 @@ async def realtime_websocket_endpoint( ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), ): - if call_id is not None and call_id.startswith("rtc_litellm_"): - from litellm.proxy.realtime_endpoints.codex import codex_realtime_sideband + if call_id is not None: + from litellm.llms.chatgpt.codex import codex_realtime_sideband await codex_realtime_sideband(websocket, call_id, user_api_key_dict) return diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index c7a20302a13..76843d68353 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -376,7 +376,7 @@ async def proxy_realtime_calls( fastapi_response: Response, ) -> Response: if request.headers.get("content-type", "").split(";", 1)[0] in ("application/json", "multipart/form-data"): - from litellm.proxy.realtime_endpoints.codex import create_codex_realtime_call + from litellm.llms.chatgpt.codex import create_codex_realtime_call return await create_codex_realtime_call(request) diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 32ae1b34485..4a5859d718d 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -141,7 +141,6 @@ async def acreate_realtime_client_secret( model=model_name, api_base=litellm_params.api_base, api_key=litellm_params.api_key, - litellm_params=litellm_params, ) ( provider_config, @@ -210,7 +209,6 @@ async def acreate_realtime_transcription_session( model=model_name, api_base=litellm_params.api_base, api_key=litellm_params.api_key, - litellm_params=litellm_params, ) ( provider_config, @@ -272,7 +270,6 @@ async def arealtime_calls( model=model_name, api_base=litellm_params.api_base, api_key=litellm_params.api_key, - litellm_params=litellm_params, ) provider_config, resolved_api_base, _ = _get_realtime_http_provider_config( custom_llm_provider=custom_llm_provider, @@ -380,7 +377,6 @@ async def _arealtime( model=model, api_base=api_base, api_key=api_key, - litellm_params=litellm_params, ) # If the client supplied `model` in the URL, ensure it uses the normalized diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py new file mode 100644 index 00000000000..583b81cd507 --- /dev/null +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -0,0 +1,212 @@ +import hashlib +import time + +import pytest +from fastapi import HTTPException, WebSocket + +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.llms.chatgpt import codex +from litellm.llms.chatgpt.codex import CodexRealtimeCall, decode_call, encode_call + + +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", + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() + 300, + ) + token = encode_call(call) + assert "/" 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]) +async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart): + 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.llms.chatgpt import 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} + + request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", + "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", + "headers": [(b"content-type", body_request.headers["content-type"].encode()), + (b"authorization", b"Bearer owner"), (b"openai-alpha", b"quicksilver=v2"), + (b"x-untrusted", b"bad")]}, receive) + auth = UserAPIKeyAuth() + authenticate = AsyncMock(return_value=auth) + authorize = AsyncMock() + monkeypatch.setattr(codex, "user_api_key_auth", authenticate) + 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 kwargs["user_api_key_dict"] is auth + 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["extra_headers"] == {"openai-alpha": "quicksilver=v2"} + assert data["extra_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"}}) + return respond() + + monkeypatch.setattr(proxy_server, "route_request", route) + response = await codex.create_codex_realtime_call(request) + assert response.status_code == 201 + 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 "rtc_private" not in token + 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"} + + websocket = WebSocket({"type": "websocket", "headers": [(b"authorization", b"Bearer owner")]}, receive_ws, send) + forward = AsyncMock() + monkeypatch.setattr(litellm, "_arealtime", forward) + await codex.codex_realtime_sideband(websocket, token, auth) + assert sent[0]["type"] == "websocket.accept" + assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" + assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-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.llms.chatgpt import 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() diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 4ee9aae4717..bfdd1e53e4d 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -104,3 +104,15 @@ async def test_async_codex_edit_without_multipart_image(chatgpt_tokens): assert str(requests[0].url).endswith("/codex/images/edits") assert requests[0].headers["content-type"] == "application/json" await client.client.aclose() + + +@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()},) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_codex.py b/tests/test_litellm/proxy/realtime_endpoints/test_codex.py deleted file mode 100644 index fdba89daaf1..00000000000 --- a/tests/test_litellm/proxy/realtime_endpoints/test_codex.py +++ /dev/null @@ -1,76 +0,0 @@ -import hashlib -import time - -import pytest -from fastapi import HTTPException, WebSocket - -from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.proxy.realtime_endpoints import codex -from litellm.proxy.realtime_endpoints.codex import CodexRealtimeCall, decode_call, encode_call - - -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", - owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), - expires_at=time.time() + 300, - ) - token = encode_call(call) - assert "/" 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"}] diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index b45812f45d7..761e87ac764 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -61,7 +61,7 @@ def _run_client_secret(session, model, monkeypatch): captured.update(kwargs) return object() - def mock_get_llm_provider(model, api_base, api_key, litellm_params=None): + def mock_get_llm_provider(model, api_base, api_key): return model, "openai", None, api_base monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) @@ -161,7 +161,7 @@ async def test_arealtime_vertex_branch_resolves_credentials_under_a_bound(monkey async def hanging_token_refresh(**kwargs): await asyncio.sleep(30) - def mock_get_llm_provider(model, api_base, api_key, litellm_params=None): + def mock_get_llm_provider(model, api_base, api_key): return model, "vertex_ai", None, api_base monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) From 284cbc18ca0c69046b3bae6527fba00d3a743f97 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 08:43:31 +0200 Subject: [PATCH 03/90] refactor(chatgpt): keep HTTP handling inside the proxy --- litellm/llms/chatgpt/codex.py | 196 ++++------------ litellm/proxy/proxy_server.py | 4 +- .../proxy/realtime_endpoints/call_sessions.py | 156 ++++++++++++ litellm/proxy/realtime_endpoints/endpoints.py | 2 +- tests/test_litellm/llms/chatgpt/test_codex.py | 222 ++---------------- .../realtime_endpoints/test_call_sessions.py | 213 +++++++++++++++++ 6 files changed, 429 insertions(+), 364 deletions(-) create mode 100644 litellm/proxy/realtime_endpoints/call_sessions.py create mode 100644 tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 669b56d4371..ad801dcdcdb 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -1,21 +1,11 @@ -import base64 -import hashlib -import json -import time -from types import MappingProxyType +from collections.abc import Mapping from typing import Final from urllib.parse import urlsplit import httpx -from fastapi import HTTPException, Request, Response, WebSocket from pydantic import BaseModel, Field +from typing_extensions import ReadOnly, TypedDict -from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import can_key_call_resolved_model -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper -from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation from litellm.types.realtime import RealtimeQueryParams, RealtimeSessionConfig @@ -36,153 +26,49 @@ class ChatGPTCallRouting(BaseModel): model: str -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("=") +class CodexSidebandRequest(TypedDict): + model: ReadOnly[str] + chatgpt_realtime_call_id: ReadOnly[str] + query_params: ReadOnly[RealtimeQueryParams] -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()).hexdigest(): - raise HTTPException(403, "Invalid or expired realtime call") - return call +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": "", + "extra_query": { # mutable-ok: router request parameters + key: value for key, value in query.items() if key in ("intent", "architecture") + }, + "extra_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") + }, + } -async def read_codex_offer(request: Request) -> CodexRealtimeOffer: - if request.headers.get("content-type", "").startswith("multipart/form-data"): - 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 create_codex_realtime_call(request: Request) -> Response: - from litellm.proxy import proxy_server as server - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing - - 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="", - anthropic_api_key_header=None, - google_ai_studio_api_key_header=None, - azure_apim_header=None, - custom_litellm_key_header=None, +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) + call_id: Final = urlsplit(response.headers.get("location", "")).path.rstrip("/").rsplit("/", 1)[-1] + return CodexRealtimeCall( + call_id=call_id, + model=routing.model, + alias=alias, + owner=owner, + expires_at=expires_at, ) - try: - await can_key_call_resolved_model( - model=model, - llm_model_list=server.llm_model_list, - valid_token=auth, - llm_router=server.llm_router, - ) - data: Final = { # mutable-ok: proxy processor enriches request data - "model": model, - "sdp_body": offer.sdp.encode(), - "session": offer.session.model_dump(exclude_none=True), - "openai_ephemeral_key": "", - "extra_query": { # mutable-ok: router request parameters - key: value for key, value in request.query_params.items() if key in ("intent", "architecture") - }, - "extra_headers": { # mutable-ok: router request headers - key: value - for key, value in request.headers.items() - if key in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation") - }, - } - processor: Final = ProxyBaseLLMRequestProcessing(data=data) - processed, _ = 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, - user_model=server.user_model, - user_temperature=server.user_temperature, - user_request_timeout=server.user_request_timeout, - user_max_tokens=server.user_max_tokens, - user_api_base=server.user_api_base, - model=model, - route_type="arealtime_calls", - ) - result: Final = await server.route_request( - data=processed, - route_type="arealtime_calls", - llm_router=server.llm_router, - user_model=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") - routing_data: Final = response.extensions.get("chatgpt_realtime") - if response.is_error: - return Response(response.content, status_code=response.status_code, media_type="application/json") - if not routing_data: - raise HTTPException(400, "Direct call signaling requires a ChatGPT deployment") - routing: Final = ChatGPTCallRouting.model_validate(routing_data) - call_id: Final = urlsplit(response.headers.get("location", "")).path.rstrip("/").rsplit("/", 1)[-1] - call: Final = CodexRealtimeCall( - call_id=call_id, - model=routing.model, - alias=model, - owner=hashlib.sha256(request.headers.get("authorization", "").encode()).hexdigest(), - expires_at=time.time() + 3600, - ) - token: Final = encode_call(call) - return Response( - response.content, - status_code=response.status_code, - media_type="application/sdp", - headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), - ) - finally: - 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 - - try: - try: - call: Final = decode_call(token, websocket.headers.get("authorization", "")) - await can_key_call_resolved_model( - model=call.alias, - llm_model_list=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 - await websocket.accept() - query: Final[RealtimeQueryParams] = {"model": call.model} - await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # internal proxy entrypoint for an already authorized call - model=f"chatgpt/{call.model}", - websocket=websocket, - chatgpt_realtime_call_id=call.call_id, - query_params=query, - user_api_key_dict=auth, - ) - finally: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) +def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: + return { + "model": f"chatgpt/{call.model}", + "chatgpt_realtime_call_id": call.call_id, + "query_params": {"model": call.model}, + } diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9d55cf21d81..7fcf89904c2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11598,7 +11598,7 @@ async def codex_live_sideband_endpoint( call_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), ) -> None: - from litellm.llms.chatgpt.codex import codex_realtime_sideband + from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband await codex_realtime_sideband(websocket, call_id, user_api_key_dict) @@ -11620,7 +11620,7 @@ async def realtime_websocket_endpoint( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), ): if call_id is not None: - from litellm.llms.chatgpt.codex import codex_realtime_sideband + from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband await codex_realtime_sideband(websocket, call_id, user_api_key_dict) return diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py new file mode 100644 index 00000000000..202f919beb6 --- /dev/null +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -0,0 +1,156 @@ +import base64 +import hashlib +import json +import time +from types import MappingProxyType +from typing import Final + +import httpx +from fastapi import HTTPException, Request, Response, WebSocket + +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.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation + + +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()).hexdigest(): + raise HTTPException(403, "Invalid or expired realtime call") + return call + + +async def read_codex_offer(request: Request) -> CodexRealtimeOffer: + if request.headers.get("content-type", "").startswith("multipart/form-data"): + 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 create_codex_realtime_call(request: Request) -> Response: + from litellm.proxy import proxy_server as server + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + 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="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + custom_litellm_key_header=None, + ) + try: + await can_key_call_resolved_model( + model=model, + llm_model_list=server.llm_model_list, + valid_token=auth, + llm_router=server.llm_router, + ) + data: Final = build_call_request(offer, request.query_params, request.headers) + processor: Final = ProxyBaseLLMRequestProcessing(data=data) + processed, _ = 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, + user_model=server.user_model, + user_temperature=server.user_temperature, + user_request_timeout=server.user_request_timeout, + user_max_tokens=server.user_max_tokens, + user_api_base=server.user_api_base, + model=model, + route_type="arealtime_calls", + ) + result: Final = await server.route_request( + data=processed, + route_type="arealtime_calls", + llm_router=server.llm_router, + user_model=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(request.headers.get("authorization", "").encode()).hexdigest(), + expires_at=time.time() + 3600, + ) + except ValueError as exc: + raise HTTPException(400, str(exc)) from exc + token: Final = encode_call(call) + return Response( + response.content, + status_code=response.status_code, + media_type="application/sdp", + headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), + ) + finally: + 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 + + try: + try: + call: Final = decode_call(token, websocket.headers.get("authorization", "")) + await can_key_call_resolved_model( + model=call.alias, + llm_model_list=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 + await websocket.accept() + await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call + **build_sideband_request(call), + websocket=websocket, + user_api_key_dict=auth, + ) + finally: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index 76843d68353..32c6abdb7b6 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -376,7 +376,7 @@ async def proxy_realtime_calls( fastapi_response: Response, ) -> Response: if request.headers.get("content-type", "").split(";", 1)[0] in ("application/json", "multipart/form-data"): - from litellm.llms.chatgpt.codex import create_codex_realtime_call + from litellm.proxy.realtime_endpoints.call_sessions import create_codex_realtime_call return await create_codex_realtime_call(request) diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 583b81cd507..078c8d85ce3 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -1,212 +1,22 @@ -import hashlib -import time - +import httpx import pytest -from fastapi import HTTPException, WebSocket -from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.llms.chatgpt import codex -from litellm.llms.chatgpt.codex import CodexRealtimeCall, decode_call, encode_call +from litellm.llms.chatgpt.codex import build_sideband_request, parse_call_response -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", - owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), - expires_at=time.time() + 300, - ) - token = encode_call(call) - assert "/" 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") +@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): + parse_call_response(response, "voice", "owner", 1000) -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]) -async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart): - 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.llms.chatgpt import 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} - - request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", - "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", - "headers": [(b"content-type", body_request.headers["content-type"].encode()), - (b"authorization", b"Bearer owner"), (b"openai-alpha", b"quicksilver=v2"), - (b"x-untrusted", b"bad")]}, receive) - auth = UserAPIKeyAuth() - authenticate = AsyncMock(return_value=auth) - authorize = AsyncMock() - monkeypatch.setattr(codex, "user_api_key_auth", authenticate) - 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 kwargs["user_api_key_dict"] is auth - 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["extra_headers"] == {"openai-alpha": "quicksilver=v2"} - assert data["extra_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"}}) - return respond() - - monkeypatch.setattr(proxy_server, "route_request", route) - response = await codex.create_codex_realtime_call(request) - assert response.status_code == 201 - 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 "rtc_private" not in token - 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"} - - websocket = WebSocket({"type": "websocket", "headers": [(b"authorization", b"Bearer owner")]}, receive_ws, send) - forward = AsyncMock() - monkeypatch.setattr(litellm, "_arealtime", forward) - await codex.codex_realtime_sideband(websocket, token, auth) - assert sent[0]["type"] == "websocket.accept" - assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" - assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-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.llms.chatgpt import 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() +def test_signaling_preserves_selected_model_for_sideband(): + response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) + call = parse_call_response(response, "voice", "owner", 1000) + request = build_sideband_request(call) + 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"} 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..b7217964116 --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -0,0 +1,213 @@ +import hashlib +import time + +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 + + +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", + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() + 300, + ) + token = encode_call(call) + assert "/" 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]) +async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart): + 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} + + request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", + "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", + "headers": [(b"content-type", body_request.headers["content-type"].encode()), + (b"authorization", b"Bearer owner"), (b"openai-alpha", b"quicksilver=v2"), + (b"x-untrusted", b"bad")]}, receive) + auth = UserAPIKeyAuth() + authenticate = AsyncMock(return_value=auth) + authorize = AsyncMock() + monkeypatch.setattr(codex, "user_api_key_auth", authenticate) + 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 kwargs["user_api_key_dict"] is auth + 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["extra_headers"] == {"openai-alpha": "quicksilver=v2"} + assert data["extra_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"}}) + return respond() + + monkeypatch.setattr(proxy_server, "route_request", route) + response = await codex.create_codex_realtime_call(request) + assert response.status_code == 201 + 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 "rtc_private" not in token + 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"} + + websocket = WebSocket({"type": "websocket", "headers": [(b"authorization", b"Bearer owner")]}, receive_ws, send) + forward = AsyncMock() + monkeypatch.setattr(litellm, "_arealtime", forward) + await codex.codex_realtime_sideband(websocket, token, auth) + assert sent[0]["type"] == "websocket.accept" + assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" + assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-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() From daa00a3401b8923b7231689570a598ad85c2c6a7 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 08:50:24 +0200 Subject: [PATCH 04/90] refactor(chatgpt): construct typed sideband requests --- litellm/llms/chatgpt/codex.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index ad801dcdcdb..85388d8c386 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -67,8 +67,8 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: - return { - "model": f"chatgpt/{call.model}", - "chatgpt_realtime_call_id": call.call_id, - "query_params": {"model": call.model}, - } + return CodexSidebandRequest( + model=f"chatgpt/{call.model}", + chatgpt_realtime_call_id=call.call_id, + query_params=RealtimeQueryParams(model=call.model), + ) From 93285edba19531afb2e12d1c7f4947af44f951bf Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 10:56:04 +0200 Subject: [PATCH 05/90] test(chatgpt): assert invalid call identifier validation --- tests/test_litellm/llms/chatgpt/test_codex.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 078c8d85ce3..a893ddbfb12 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -8,7 +8,7 @@ from litellm.llms.chatgpt.codex import build_sideband_request, parse_call_respon 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): + with pytest.raises(ValueError, match="String should match pattern"): parse_call_response(response, "voice", "owner", 1000) From 3686f6a005cd1ca1915aa2542256ad8470eb35bf Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 11:17:09 +0200 Subject: [PATCH 06/90] fix(chatgpt): select live transport from model metadata --- litellm/llms/chatgpt/realtime.py | 11 +++++++- ...odel_prices_and_context_window_backup.json | 10 ++++++++ model_prices_and_context_window.json | 10 ++++++++ .../llms/chatgpt/test_realtime.py | 25 ++++++++++++++++++- 4 files changed, 54 insertions(+), 2 deletions(-) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 66a703eb616..93a79b7d159 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -9,6 +9,7 @@ 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 .common_utils import CHATGPT_API_BASE from .responses.transformation import ChatGPTResponsesAPIConfig @@ -34,6 +35,14 @@ def realtime_headers( } +def realtime_endpoint(model: str) -> str: + try: + model_info: Final = get_model_info(model, custom_llm_provider="chatgpt") + except Exception: + return "realtime" + return "live" if "/v1/live" in (model_info.get("supported_endpoints") or ()) else "realtime" + + class ChatGPTRealtime(OpenAIRealtime): def __init__(self, params: GenericLiteLLMParams, headers: Mapping[str, str]) -> None: super().__init__() @@ -50,7 +59,7 @@ class ChatGPTRealtime(OpenAIRealtime): def _construct_url(self, api_base: str, query_params: RealtimeQueryParams) -> str: base: Final = URL(api_base) - endpoint: Final = "live" if query_params.get("model") == "gpt-live-1-codex" else "realtime" + endpoint: Final = realtime_endpoint(query_params.get("model", "")) if self._call_id: return str( base.copy_with( diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 54ebdc85be9..ae4432e167d 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -27832,6 +27832,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.4": { "litellm_provider": "chatgpt", "max_input_tokens": 1050000, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 54ebdc85be9..ae4432e167d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -27832,6 +27832,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.4": { "litellm_provider": "chatgpt", "max_input_tokens": 1050000, diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 7ce8d5f298a..239a5a12280 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -40,7 +40,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens): @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): +def test_realtime_uses_platform_endpoint_with_oauth_headers(model, endpoint, chatgpt_tokens, local_model_cost_map): handler = ChatGPTRealtime( GenericLiteLLMParams(), { @@ -55,3 +55,26 @@ def test_realtime_uses_platform_endpoint_with_oauth_headers(model, endpoint, cha assert headers["Authorization"] == "Bearer test-token-default" assert "authorization" not in headers assert headers["openai-alpha"] == "quicksilver=v2" + + +@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" + ) From 8e5901e90aa8becc073df8ef8a5bc5c9b53dcdc1 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 11:21:29 +0200 Subject: [PATCH 07/90] chore(chatgpt): document unmapped model exception contract --- litellm/llms/chatgpt/realtime.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 93a79b7d159..831f67d82c5 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -38,7 +38,7 @@ def realtime_headers( def realtime_endpoint(model: str) -> str: try: model_info: Final = get_model_info(model, custom_llm_provider="chatgpt") - except Exception: + 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" From 36727584a9cbc2161269e70d05d3d4b0c975358f Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 11:45:39 +0200 Subject: [PATCH 08/90] fix(chatgpt): enforce sideband policies and websocket credentials --- .../proxy/realtime_endpoints/call_sessions.py | 104 ++++++++++++++---- .../realtime_endpoints/test_call_sessions.py | 59 +++++++++- 2 files changed, 138 insertions(+), 25 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 202f919beb6..2ea723335ee 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -3,11 +3,13 @@ import hashlib import json import time from types import MappingProxyType -from typing import Final +from typing import Final, Literal import httpx from fastapi import HTTPException, Request, Response, WebSocket +from starlette.types import Message +from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.chatgpt.codex import ( CodexRealtimeCall, @@ -52,10 +54,38 @@ async def read_codex_offer(request: Request) -> CodexRealtimeOffer: return CodexRealtimeOffer.model_validate(await request.json()) -async def create_codex_realtime_call(request: Request) -> Response: +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"], +) -> dict[str, object]: # 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, _ = 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, + user_model=server.user_model, + user_temperature=server.user_temperature, + user_request_timeout=server.user_request_timeout, + user_max_tokens=server.user_max_tokens, + user_api_base=server.user_api_base, + model=model, + route_type=route_type, + ) + return processed + + +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: @@ -80,22 +110,7 @@ async def create_codex_realtime_call(request: Request) -> Response: llm_router=server.llm_router, ) data: Final = build_call_request(offer, request.query_params, request.headers) - processor: Final = ProxyBaseLLMRequestProcessing(data=data) - processed, _ = 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, - user_model=server.user_model, - user_temperature=server.user_temperature, - user_request_timeout=server.user_request_timeout, - user_max_tokens=server.user_max_tokens, - user_api_base=server.user_api_base, - model=model, - route_type="arealtime_calls", - ) + processed: Final = await process_codex_request(request, data, auth, model, "arealtime_calls") result: Final = await server.route_request( data=processed, route_type="arealtime_calls", @@ -134,9 +149,16 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP 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() + ) + alternate_key: Final = websocket.headers.get("api-key") or next( + (p.removeprefix("openai-insecure-api-key.") for p in protocols if p.startswith("openai-insecure-api-key.")), "" + ) + authorization: Final = websocket.headers.get("authorization") or f"Bearer {alternate_key}" try: try: - call: Final = decode_call(token, websocket.headers.get("authorization", "")) + call: Final = decode_call(token, authorization) await can_key_call_resolved_model( model=call.alias, llm_model_list=server.llm_model_list, @@ -146,11 +168,47 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP except (HTTPException, ProxyException): await websocket.close(code=1008, reason="Invalid realtime call") return - await websocket.accept() - await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call + + 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), - websocket=websocket, - user_api_key_dict=auth, + "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() + ], + } + try: + processed: Final = await process_codex_request(request, data, auth, call.alias, "_arealtime") + except Exception: + 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) + ) + await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call + **{ # mutable-ok: retain processed policy metadata while pinning the existing call's routing + **processed, + **build_sideband_request(call), + "websocket": websocket, + "user_api_key_dict": auth, + } ) finally: await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index b7217964116..617cd644702 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -106,7 +106,8 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id @pytest.mark.asyncio @pytest.mark.parametrize("multipart", [False, True]) -async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart): +@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"]) +async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential): import json from unittest.mock import AsyncMock @@ -147,6 +148,11 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, async def common_processing_pre_call_logic(self, **kwargs): assert kwargs["user_api_key_dict"] is auth + 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, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None return self.data, None monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor) @@ -184,11 +190,20 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, async def receive_ws(): return {"type": "websocket.connect"} - websocket = WebSocket({"type": "websocket", "headers": [(b"authorization", b"Bearer owner")]}, receive_ws, send) + credential_headers = { + "authorization": [(b"authorization", b"Bearer owner")], + "api-key": [(b"api-key", b"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["metadata"] == {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"} assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-codex" assert authorize.await_count == 2 @@ -211,3 +226,43 @@ async def test_invalid_offers_fail_before_authentication(monkeypatch, body): 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"}] From d482638d5fa28fc31078cf8346cbd543ee7d82ea Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 12:15:25 +0200 Subject: [PATCH 09/90] chore(chatgpt): document custom hook rejection handling --- litellm/proxy/realtime_endpoints/call_sessions.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 2ea723335ee..a02c114bcab 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -195,7 +195,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP } try: processed: Final = await process_codex_request(request, data, auth, call.alias, "_arealtime") - except Exception: + except Exception: # noqa: BLE001 # user-defined pre-call hooks may raise any exception; always reject the connection verbose_proxy_logger.exception("Realtime sideband pre-call rejected") await websocket.close(code=1008, reason="Realtime pre-call rejected") return From d5f352823b9dcb34707722e66ad6e11ae5ef8178 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 12:16:23 +0200 Subject: [PATCH 10/90] style(chatgpt): keep hook rationale within line limit --- litellm/proxy/realtime_endpoints/call_sessions.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index a02c114bcab..f0105ea6cfd 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -195,7 +195,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP } try: processed: Final = await process_codex_request(request, data, auth, call.alias, "_arealtime") - except Exception: # noqa: BLE001 # user-defined pre-call hooks may raise any exception; always reject the connection + 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 From b0022e5d2e077791f1ff65cec29025ca85b4f4bb Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 12:38:42 +0200 Subject: [PATCH 11/90] test(chatgpt): allow Codex call endpoints in catalog schema --- tests/test_litellm/test_utils.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index fee5e3a2e4c..26f8e9778f5 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -1125,6 +1125,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", From 852dee23d4ae018f522f7a71d00b4597d8c6117c Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 12:52:03 +0200 Subject: [PATCH 12/90] fix(chatgpt): retain sideband routing and pending usage reservations --- litellm/llms/chatgpt/codex.py | 5 +++ .../proxy/realtime_endpoints/call_sessions.py | 16 +++++--- litellm/realtime_api/main.py | 7 +++- tests/test_litellm/llms/chatgpt/test_codex.py | 3 +- .../llms/chatgpt/test_realtime.py | 6 ++- .../realtime_endpoints/test_call_sessions.py | 39 ++++++++++++++++++- 6 files changed, 66 insertions(+), 10 deletions(-) diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 85388d8c386..49e4227e5d2 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -18,15 +18,18 @@ class CodexRealtimeCall(BaseModel): call_id: str = Field(pattern=r"^rtc_[A-Za-z0-9_-]+$") model: str alias: str + api_base: str | None = None owner: str expires_at: float class ChatGPTCallRouting(BaseModel): model: str + api_base: str | None = None class CodexSidebandRequest(TypedDict): + api_base: ReadOnly[str | None] model: ReadOnly[str] chatgpt_realtime_call_id: ReadOnly[str] query_params: ReadOnly[RealtimeQueryParams] @@ -63,11 +66,13 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire alias=alias, owner=owner, expires_at=expires_at, + api_base=routing.api_base, ) 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), diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index f0105ea6cfd..b1d1112af21 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -10,6 +10,8 @@ from fastapi import HTTPException, Request, Response, WebSocket from starlette.types import Message from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.chatgpt.codex import ( CodexRealtimeCall, @@ -60,12 +62,12 @@ async def process_codex_request( auth: UserAPIKeyAuth, model: str, route_type: Literal["arealtime_calls", "_arealtime"], -) -> dict[str, object]: # mutable-ok: common request processor returns enriched routing arguments +) -> 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, _ = await processor.common_processing_pre_call_logic( + processed, logging_obj = await processor.common_processing_pre_call_logic( request=request, general_settings=server.general_settings, user_api_key_dict=auth, @@ -80,7 +82,7 @@ async def process_codex_request( model=model, route_type=route_type, ) - return processed + return processed, logging_obj async def create_codex_realtime_call(request: Request) -> Response: @@ -110,7 +112,7 @@ async def create_codex_realtime_call(request: Request) -> Response: llm_router=server.llm_router, ) data: Final = build_call_request(offer, request.query_params, request.headers) - processed: Final = await process_codex_request(request, data, auth, model, "arealtime_calls") + processed, _ = await process_codex_request(request, data, auth, model, "arealtime_calls") result: Final = await server.route_request( data=processed, route_type="arealtime_calls", @@ -156,6 +158,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP (p.removeprefix("openai-insecure-api-key.") for p in protocols if p.startswith("openai-insecure-api-key.")), "" ) authorization: Final = websocket.headers.get("authorization") or f"Bearer {alternate_key}" + logging_obj: Logging | None = None # rebind-ok: cleanup needs the logger only after pre-call succeeds try: try: call: Final = decode_call(token, authorization) @@ -194,7 +197,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP ], } try: - processed: Final = await process_codex_request(request, data, auth, call.alias, "_arealtime") + 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") @@ -211,4 +214,5 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP } ) finally: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + 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/realtime_api/main.py b/litellm/realtime_api/main.py index 4a5859d718d..15339173a1d 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -99,7 +99,11 @@ def _get_realtime_http_provider_config( provider=LlmProviders(custom_llm_provider), ) - raw_api_base: Final = dynamic_api_base or litellm_params.api_base + raw_api_base: Final = ( + litellm_params.api_base or dynamic_api_base + if custom_llm_provider == "chatgpt" + else 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: @@ -303,6 +307,7 @@ async def arealtime_calls( response.extensions["chatgpt_realtime"] = MappingProxyType( { "model": model_name, + "api_base": litellm_params.api_base, } ) return response diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index a893ddbfb12..3ea5b1fffd2 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -14,9 +14,10 @@ def test_signaling_rejects_invalid_upstream_call_id(location): def test_signaling_preserves_selected_model_for_sideband(): response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex"}}) call = parse_call_response(response, "voice", "owner", 1000) request = build_sideband_request(call) + 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"} diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 239a5a12280..7d83d8c6856 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -10,7 +10,8 @@ from litellm.types.router import GenericLiteLLMParams @pytest.mark.asyncio -async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens): +@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): @@ -21,6 +22,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens): 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"}}}, @@ -28,6 +30,8 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens): extra_headers={"openai-alpha": "quicksilver=v2"}, client=client, ) + assert response.extensions["chatgpt_realtime"]["api_base"] == api_base + assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com") assert response.status_code == 201 assert requests[0].url.path == "/backend-api/codex/realtime/calls" assert requests[0].url.params["architecture"] == "avas" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 617cd644702..59c08c09b00 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -1,5 +1,6 @@ import hashlib import time +from types import SimpleNamespace import pytest from fastapi import HTTPException, WebSocket @@ -10,6 +11,41 @@ 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("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 + + def test_sideband_token_binds_owner_and_model(monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") call = CodexRealtimeCall( @@ -166,7 +202,7 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, 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"}}) + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex"}}) return respond() monkeypatch.setattr(proxy_server, "route_request", route) @@ -206,6 +242,7 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, assert forward.await_args.kwargs["metadata"] == {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"} 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 From 96594e7b0ef970255968703e306c313507e32b6c Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 13:30:28 +0200 Subject: [PATCH 13/90] fix(chatgpt): honor configured gateways across Codex transports --- litellm/llms/chatgpt/authenticator.py | 4 ++-- litellm/llms/chatgpt/images.py | 6 +++--- litellm/llms/chatgpt/realtime.py | 8 ++++++-- litellm/realtime_api/main.py | 6 ++++-- .../test_litellm/llms/chatgpt/test_images.py | 13 +++++++++++++ .../llms/chatgpt/test_realtime.py | 19 ++++++++++++++++++- 6 files changed, 46 insertions(+), 10 deletions(-) diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index 563826c2b93..a516fb2fe09 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -49,8 +49,8 @@ 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 + def get_api_base(self, 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/images.py b/litellm/llms/chatgpt/images.py index ee6ba17e340..97a0da5f518 100644 --- a/litellm/llms/chatgpt/images.py +++ b/litellm/llms/chatgpt/images.py @@ -15,7 +15,7 @@ from litellm.llms.openai.image_generation.gpt_transformation import GPTImageGene from litellm.types.llms.openai import AllMessageValues, FileTypes from litellm.types.router import GenericLiteLLMParams -from .common_utils import CHATGPT_API_BASE +from .authenticator import Authenticator from .responses.transformation import ChatGPTResponsesAPIConfig @@ -82,7 +82,7 @@ class ChatGPTImageGenerationConfig(GPTImageGenerationConfig): litellm_params: Mapping[str, object], stream: bool | None = None, ) -> str: - return f"{(api_base or CHATGPT_API_BASE).rstrip('/')}/images/generations" + return f"{(api_base or Authenticator().get_api_base()).rstrip('/')}/images/generations" def transform_image_generation_request( self, @@ -107,7 +107,7 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig): 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 CHATGPT_API_BASE).rstrip('/')}/images/edits" + return f"{(api_base or Authenticator().get_api_base()).rstrip('/')}/images/edits" def use_multipart_form_data(self) -> bool: return False diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 831f67d82c5..5be3e766b99 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -11,7 +11,7 @@ from litellm.types.realtime import RealtimeQueryParams from litellm.types.router import GenericLiteLLMParams from litellm.utils import get_model_info -from .common_utils import CHATGPT_API_BASE +from .authenticator import Authenticator from .responses.transformation import ChatGPTResponsesAPIConfig @@ -44,6 +44,10 @@ def realtime_endpoint(model: str) -> str: class ChatGPTRealtime(OpenAIRealtime): + @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]) -> None: super().__init__() self._profile_headers = realtime_headers(params, headers) @@ -90,7 +94,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): api_base: str | None, **kwargs: object, # kwargs-ok: provider interface accepts optional credentials ) -> str: - return api_base or CHATGPT_API_BASE + return api_base or Authenticator().get_api_base() def get_api_key( self, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 15339173a1d..9d0d2075870 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -304,10 +304,12 @@ async def arealtime_calls( api_version=litellm_params.api_version, ) if custom_llm_provider == "chatgpt": + from litellm.llms.chatgpt.realtime import ChatGPTRealtime + response.extensions["chatgpt_realtime"] = MappingProxyType( { "model": model_name, - "api_base": litellm_params.api_base, + "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), } ) return response @@ -462,7 +464,7 @@ async def _arealtime( model=model, websocket=websocket, logging_obj=litellm_logging_obj, - api_base=api_base or "https://api.openai.com/v1", + api_base=ChatGPTRealtime.get_api_base(api_base), api_key="chatgpt-oauth", timeout=timeout, query_params=query_params, diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index bfdd1e53e4d..911183ba50b 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -116,3 +116,16 @@ def test_edit_accepts_filesystem_path(tmp_path, as_tuple): ) 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): + 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" diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 7d83d8c6856..f66a3a92ac4 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -30,7 +30,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap extra_headers={"openai-alpha": "quicksilver=v2"}, client=client, ) - assert response.extensions["chatgpt_realtime"]["api_base"] == api_base + assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1") assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com") assert response.status_code == 201 assert requests[0].url.path == "/backend-api/codex/realtime/calls" @@ -82,3 +82,20 @@ def test_realtime_unknown_model_keeps_standard_endpoint(chatgpt_tokens, local_mo 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" + ) From 6e87b4b98514f428dd0f2492940371f921b73c70 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 13:32:33 +0200 Subject: [PATCH 14/90] fix(chatgpt): resolve gateway URLs without touching token storage --- litellm/llms/chatgpt/authenticator.py | 3 ++- litellm/llms/chatgpt/images.py | 4 ++-- litellm/llms/chatgpt/realtime.py | 4 ++-- tests/test_litellm/llms/chatgpt/test_images.py | 5 ++++- 4 files changed, 10 insertions(+), 6 deletions(-) diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index a516fb2fe09..4a5045e02e7 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -49,7 +49,8 @@ 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, default_base: str = CHATGPT_API_BASE) -> str: + @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: diff --git a/litellm/llms/chatgpt/images.py b/litellm/llms/chatgpt/images.py index 97a0da5f518..3467b68c31f 100644 --- a/litellm/llms/chatgpt/images.py +++ b/litellm/llms/chatgpt/images.py @@ -82,7 +82,7 @@ class ChatGPTImageGenerationConfig(GPTImageGenerationConfig): litellm_params: Mapping[str, object], stream: bool | None = None, ) -> str: - return f"{(api_base or Authenticator().get_api_base()).rstrip('/')}/images/generations" + return f"{(api_base or Authenticator.get_api_base()).rstrip('/')}/images/generations" def transform_image_generation_request( self, @@ -107,7 +107,7 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig): 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" + return f"{(api_base or Authenticator.get_api_base()).rstrip('/')}/images/edits" def use_multipart_form_data(self) -> bool: return False diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 5be3e766b99..26e2f0ffad9 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -46,7 +46,7 @@ def realtime_endpoint(model: str) -> str: class ChatGPTRealtime(OpenAIRealtime): @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") + return api_base or Authenticator.get_api_base(default_base="https://api.openai.com/v1") def __init__(self, params: GenericLiteLLMParams, headers: Mapping[str, str]) -> None: super().__init__() @@ -94,7 +94,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): api_base: str | None, **kwargs: object, # kwargs-ok: provider interface accepts optional credentials ) -> str: - return api_base or Authenticator().get_api_base() + return api_base or Authenticator.get_api_base() def get_api_key( self, diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 911183ba50b..6b69a83839d 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -120,7 +120,10 @@ def test_edit_accepts_filesystem_path(tmp_path, as_tuple): @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): +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/") From 57e831d44510c99c063580594056e5bf67696cce Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 19:44:19 +0200 Subject: [PATCH 15/90] fix(chatgpt): authenticate sideband models and forward image headers --- litellm/images/main.py | 1 + litellm/proxy/auth/user_api_key_auth.py | 22 ++++++---- .../test_litellm/llms/chatgpt/test_images.py | 2 + .../proxy/auth/test_user_api_key_auth.py | 41 +++++++++++++++++++ 4 files changed, 57 insertions(+), 9 deletions(-) diff --git a/litellm/images/main.py b/litellm/images/main.py index 7290e537369..11df9728ede 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -402,6 +402,7 @@ def image_generation( model=model, prompt=prompt, image_generation_provider_config=image_generation_config, + extra_headers=extra_headers, image_generation_optional_request_params=optional_params, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params_dict, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 6b000489d5a..1ab1811fef8 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -508,15 +508,6 @@ async def user_api_key_auth_websocket(websocket: WebSocket): request._url = websocket.url - query_params: Final = websocket.query_params - - model: Final = query_params.get("model") - - 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: @@ -542,6 +533,19 @@ async def user_api_key_auth_websocket(websocket: WebSocket): # 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") + model: Final = ( + decode_call(call_token, f"Bearer {api_key}").alias + if call_token is not None + else websocket.query_params.get("model") + ) + + async def return_body(): + return _realtime_request_body(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/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 6b69a83839d..0f4eafccfa4 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -25,7 +25,9 @@ def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens): quality="auto", size="auto", background="auto", + extra_headers={"x-gateway-route": "images"}, ) + assert requests[0].headers["x-gateway-route"] == "images" assert result.data[0].b64_json == "aGVsbG8=" assert str(requests[0].url) == "https://chatgpt.com/backend-api/codex/images/generations" assert requests[0].headers["authorization"] == "Bearer test-token-" + "default" 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 fd289e33ea6..d4dae5ee150 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 @@ -7133,3 +7133,44 @@ def test_user_api_key_auth_opens_a_datadog_span_for_accepted_and_rejected_keys(t assert report["outcomes"] == ["accepted", "rejected"] auth_span = "litellm.proxy.auth.user_api_key_auth.user_api_key_auth" 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"]) +@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, + )) + 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")], + "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")] From 511599d7b05d19d5148f884b82bca98c803a8ea5 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 20:45:55 +0200 Subject: [PATCH 16/90] fix(chatgpt): preserve explicit gateway during provider resolution --- litellm/llms/chatgpt/chat/transformation.py | 2 +- tests/test_litellm/llms/chatgpt/test_images.py | 6 ++++-- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/litellm/llms/chatgpt/chat/transformation.py b/litellm/llms/chatgpt/chat/transformation.py index e35408b0829..1926283e85e 100644 --- a/litellm/llms/chatgpt/chat/transformation.py +++ b/litellm/llms/chatgpt/chat/transformation.py @@ -30,7 +30,7 @@ class ChatGPTConfig(OpenAIConfig): api_key: str | None, custom_llm_provider: str, ) -> tuple[str | None, str | None, str]: - dynamic_api_base: Final = self.authenticator.get_api_base() + dynamic_api_base: Final = api_base or self.authenticator.get_api_base() try: dynamic_api_key: Final = self.authenticator.get_access_token() except GetAccessTokenError as e: diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 0f4eafccfa4..0f781475afe 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -9,7 +9,8 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.router import GenericLiteLLMParams -def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens): +@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): @@ -21,6 +22,7 @@ def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens): result = litellm.image_generation( model="chatgpt/gpt-image-2", prompt="blue circle", + api_base=api_base, client=client, quality="auto", size="auto", @@ -29,7 +31,7 @@ def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens): ) assert requests[0].headers["x-gateway-route"] == "images" assert result.data[0].b64_json == "aGVsbG8=" - assert str(requests[0].url) == "https://chatgpt.com/backend-api/codex/images/generations" + 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 From 09f52d9aa4e00f652c8cc8ff7d008880f3529edd Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 21:34:41 +0200 Subject: [PATCH 17/90] fix(ci): update websocket fixtures and vulnerable development parser --- tests/proxy_unit_tests/test_user_api_key_auth.py | 3 +++ ui/litellm-dashboard/package-lock.json | 6 +++--- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index 0cdf3500d50..7b8d20543a5 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -849,6 +849,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")]} @@ -872,6 +873,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 @@ -885,6 +887,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/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 44360e94392..1f459bc50ea 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -11402,9 +11402,9 @@ } }, "node_modules/smol-toml": { - "version": "1.6.1", - "resolved": "https://registry.npmjs.org/smol-toml/-/smol-toml-1.6.1.tgz", - "integrity": "sha512-dWUG8F5sIIARXih1DTaQAX4SsiTXhInKf1buxdY9DIg4ZYPZK5nGM1VRIYmEbDbsHt7USo99xSLFu5Q1IqTmsg==", + "version": "1.8.0", + "resolved": "https://registry.npmjs.org/smol-toml/-/smol-toml-1.8.0.tgz", + "integrity": "sha512-kCZr2V3ch9i00x8zXRhjUNVcjG9ijES5dDudkXvUVCT5QlJNQWElSJdZqyPemffHoLNUYwOcou0Fy+ojN0uHSQ==", "dev": true, "license": "BSD-3-Clause", "engines": { From fe9a4ffca6e4b7833b6b8bae8224e8d0c2560bec Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 21:48:44 +0200 Subject: [PATCH 18/90] fix(chatgpt): pin sideband budgets and preserve OAuth image identity --- litellm/llms/chatgpt/images.py | 12 ++++- litellm/llms/custom_httpx/llm_http_handler.py | 16 +++++-- litellm/proxy/auth/user_api_key_auth.py | 4 +- .../test_litellm/llms/chatgpt/test_images.py | 17 +++++-- .../proxy/auth/test_user_api_key_auth.py | 46 +++++++++++++++++++ 5 files changed, 86 insertions(+), 9 deletions(-) diff --git a/litellm/llms/chatgpt/images.py b/litellm/llms/chatgpt/images.py index 3467b68c31f..82845aa05f4 100644 --- a/litellm/llms/chatgpt/images.py +++ b/litellm/llms/chatgpt/images.py @@ -49,6 +49,12 @@ def encode_reference( } +def without_image_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")} + ) + + def image_headers( headers: Mapping[str, object], model: str, params: Mapping[str, object] ) -> dict[str, object]: # mutable-ok: image handler requires dictionaries @@ -57,7 +63,11 @@ def image_headers( model=model, litellm_params=GenericLiteLLMParams.model_validate(params), ) - return {**headers, **auth_headers, "accept": "application/json"} # mutable-ok: JSON request serialization + return { # mutable-ok: image handler requires dictionaries + **without_image_identity_headers(headers), + **auth_headers, + "accept": "application/json", + } class ChatGPTImageGenerationConfig(GPTImageGenerationConfig): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 30f12e3a64d..1c6cdefe907 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6638,6 +6638,14 @@ class BaseLLMHTTPHandler: else: raise Exception(f"Unexpected error while closing WebSocket: {close_error}") + @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.images import without_image_identity_headers + + return without_image_identity_headers(headers) + return headers + def image_edit_handler( self, model: str, @@ -6694,7 +6702,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, @@ -6793,7 +6801,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, @@ -6910,7 +6918,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, @@ -7017,7 +7025,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/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 2a1d0709212..f6f4d7bf1c4 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -302,7 +302,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, @@ -542,6 +542,8 @@ async def user_api_key_auth_websocket(websocket: WebSocket): if call_token is not None else websocket.query_params.get("model") ) + if call_token is not None: + request.scope["litellm_pinned_realtime_model"] = model async def return_body(): return _realtime_request_body(model) diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 0f781475afe..2c51b932e3e 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -27,7 +27,7 @@ def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens, api_base): quality="auto", size="auto", background="auto", - extra_headers={"x-gateway-route": "images"}, + 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=" @@ -50,6 +50,7 @@ def test_codex_json_edit_survives_sdk_dispatch(chatgpt_tokens): 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", @@ -61,6 +62,10 @@ def test_codex_json_edit_survives_sdk_dispatch(chatgpt_tokens): 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] @@ -82,9 +87,10 @@ def test_edit_converts_multipart_image_bytes(): def test_image_auth_does_not_accept_inbound_override(chatgpt_tokens): headers = ChatGPTImageGenerationConfig().validate_environment( - {"Authorization": "Bearer wrong"}, "gpt-image-2", [], {}, {"chatgpt_token_dir": chatgpt_tokens} + {"authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, "gpt-image-2", [], {}, {"chatgpt_token_dir": chatgpt_tokens} ) - assert headers["Authorization"] == "Bearer test-token-default" + assert httpx.Headers(headers)["authorization"] == "Bearer test-token-default" + assert httpx.Headers(headers)["chatgpt-account-id"] == "test-account-default" @pytest.mark.asyncio @@ -100,6 +106,7 @@ async def test_async_codex_edit_without_multipart_image(chatgpt_tokens): 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", @@ -109,6 +116,10 @@ async def test_async_codex_edit_without_multipart_image(chatgpt_tokens): 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("as_tuple", [False, True]) def test_edit_accepts_filesystem_path(tmp_path, as_tuple): 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 6df7f2118e0..886d3361c49 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 @@ -7290,3 +7290,49 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a assert token.team_member == Member(user_id="jwt-user", role="admin") assert token.team_member_spend == 1.5 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": ""}) From 023f3e93e748b068e064971ba6be7df6c6eef5cb Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 22:36:30 +0200 Subject: [PATCH 19/90] test(ui): follow configured model in reasoning preset assertion --- .../src/components/add_model/add_auto_router_tab.test.tsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx index 014854ac712..cd03cba3b31 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx @@ -1047,7 +1047,7 @@ describe("AddAutoRouterTab", () => { expect(vi.mocked(handleAddAutoRouterSubmit).mock.calls.at(-1)?.[0]).toMatchObject({ complexity_router_config: { tier_model_configs: { - REASONING: [{ model_name: "claude-opus-5", litellm_params: { reasoning_effort: "high" } }], + REASONING: [{ model_name: ANTHROPIC_TIERS.REASONING[0], litellm_params: { reasoning_effort: "high" } }], }, }, }); From 8339ebcc4de03eddfe0fbd573abd756e98705053 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 22:46:40 +0200 Subject: [PATCH 20/90] fix(chatgpt): preserve OAuth identity during realtime signaling --- litellm/llms/chatgpt/common_utils.py | 8 +++++ litellm/llms/chatgpt/images.py | 9 ++--- litellm/llms/custom_httpx/llm_http_handler.py | 13 +++++-- .../llms/chatgpt/test_realtime.py | 36 ++++++++++++++++++- 4 files changed, 55 insertions(+), 11 deletions(-) diff --git a/litellm/llms/chatgpt/common_utils.py b/litellm/llms/chatgpt/common_utils.py index 35e32e4172f..27fc28d2cd7 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 index 82845aa05f4..8366e5016aa 100644 --- a/litellm/llms/chatgpt/images.py +++ b/litellm/llms/chatgpt/images.py @@ -16,6 +16,7 @@ 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 @@ -49,12 +50,6 @@ def encode_reference( } -def without_image_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")} - ) - - def image_headers( headers: Mapping[str, object], model: str, params: Mapping[str, object] ) -> dict[str, object]: # mutable-ok: image handler requires dictionaries @@ -64,7 +59,7 @@ def image_headers( litellm_params=GenericLiteLLMParams.model_validate(params), ) return { # mutable-ok: image handler requires dictionaries - **without_image_identity_headers(headers), + **without_oauth_identity_headers(headers), **auth_headers, "accept": "application/json", } diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 1c6cdefe907..09f698891e5 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6390,6 +6390,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, @@ -6407,7 +6410,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 + ) # Build multipart form data: sdp + session JSON session_data: Final = session_config or {} @@ -6641,9 +6648,9 @@ class BaseLLMHTTPHandler: @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.images import without_image_identity_headers + from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers - return without_image_identity_headers(headers) + return without_oauth_identity_headers(headers) return headers def image_edit_handler( diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index f66a3a92ac4..d41b64f3178 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -27,7 +27,12 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap sdp_body=b"v=0\r\n", session={"model": "chatgpt/gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}}, extra_query={"intent": "quicksilver", "architecture": "avas"}, - extra_headers={"openai-alpha": "quicksilver=v2"}, + 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") @@ -36,6 +41,9 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap 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"}}}, @@ -43,6 +51,32 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap 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( From 4b85d94de4ba19cd5cf93a8387b3dbb195b89fad Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 23:50:03 +0200 Subject: [PATCH 21/90] fix(chatgpt): retain gateway headers across realtime connections --- litellm/llms/chatgpt/codex.py | 5 +++ litellm/llms/chatgpt/realtime.py | 17 ++++++-- litellm/realtime_api/main.py | 5 ++- tests/test_litellm/llms/chatgpt/test_codex.py | 4 +- .../llms/chatgpt/test_realtime.py | 40 +++++++++++++++++++ .../realtime_endpoints/test_call_sessions.py | 2 + 6 files changed, 67 insertions(+), 6 deletions(-) diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 49e4227e5d2..bc812547d6b 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -19,6 +19,7 @@ class CodexRealtimeCall(BaseModel): model: str alias: str api_base: str | None = None + extra_headers: Mapping[str, str] | None = None owner: str expires_at: float @@ -26,6 +27,7 @@ class CodexRealtimeCall(BaseModel): class ChatGPTCallRouting(BaseModel): model: str api_base: str | None = None + extra_headers: Mapping[str, str] | None = None class CodexSidebandRequest(TypedDict): @@ -33,6 +35,7 @@ class CodexSidebandRequest(TypedDict): model: ReadOnly[str] chatgpt_realtime_call_id: ReadOnly[str] query_params: ReadOnly[RealtimeQueryParams] + extra_headers: ReadOnly[Mapping[str, str] | None] def build_call_request( @@ -67,6 +70,7 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire owner=owner, expires_at=expires_at, api_base=routing.api_base, + extra_headers=routing.extra_headers, ) @@ -76,4 +80,5 @@ def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: model=f"chatgpt/{call.model}", chatgpt_realtime_call_id=call.call_id, query_params=RealtimeQueryParams(model=call.model), + extra_headers=call.extra_headers, ) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 26e2f0ffad9..97783c2bf38 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -12,11 +12,19 @@ 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 +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 realtime_headers( - params: GenericLiteLLMParams, headers: Mapping[str, str] + 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( { @@ -32,6 +40,7 @@ def realtime_headers( litellm_params=params, ), **forwarded, + **configured_realtime_headers(extra_headers), } @@ -48,9 +57,11 @@ class ChatGPTRealtime(OpenAIRealtime): 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]) -> None: + 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) + self._profile_headers = realtime_headers(params, headers, extra_headers) self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None)) def _get_additional_headers( diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 9d0d2075870..e5ec85b5c2a 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -304,12 +304,13 @@ async def arealtime_calls( api_version=litellm_params.api_version, ) if custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtime + from litellm.llms.chatgpt.realtime import ChatGPTRealtime, configured_realtime_headers response.extensions["chatgpt_realtime"] = MappingProxyType( { "model": model_name, "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), + "extra_headers": configured_realtime_headers(kwargs.get("extra_headers")), } ) return response @@ -460,7 +461,7 @@ async def _arealtime( elif _custom_llm_provider == "chatgpt": from litellm.llms.chatgpt.realtime import ChatGPTRealtime - await ChatGPTRealtime(litellm_params, websocket.headers).async_realtime( + await ChatGPTRealtime(litellm_params, websocket.headers, headers).async_realtime( model=model, websocket=websocket, logging_obj=litellm_logging_obj, diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 3ea5b1fffd2..06f04b9b81c 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -14,10 +14,12 @@ def test_signaling_rejects_invalid_upstream_call_id(location): def test_signaling_preserves_selected_model_for_sideband(): response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex"}}) + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", + "extra_headers": {"x-gateway-route": "voice"}}}) call = parse_call_response(response, "voice", "owner", 1000) request = build_sideband_request(call) 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"} diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index d41b64f3178..3dd74a64d8d 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -1,7 +1,11 @@ +import asyncio import json +from types import SimpleNamespace +from unittest.mock import AsyncMock import httpx import pytest +from websockets.asyncio.server import serve import litellm from litellm.llms.chatgpt.realtime import ChatGPTRealtime @@ -9,6 +13,41 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.router import GenericLiteLLMParams +@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): + captured = asyncio.get_running_loop().create_future() + + async def receive_connection(connection): + captured.set_result(connection.request.headers) + await connection.wait_closed() + + 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(), + ) + async with serve(receive_connection, "127.0.0.1", 0) as gateway: + port = gateway.sockets[0].getsockname()[1] + await asyncio.wait_for(litellm._arealtime( + model=f"chatgpt/{model}", websocket=websocket, api_base=f"http://127.0.0.1:{port}", + chatgpt_realtime_call_id=call_id, + headers={"x-deployment-header": "configured"}, + extra_headers={"X-Gateway-Route": "voice", "OpenAI-Alpha": "configured-value", + "aUtHoRiZaTiOn": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, + ), timeout=10) + headers = await asyncio.wait_for(captured, timeout=5) + 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): @@ -36,6 +75,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap client=client, ) assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1") + assert response.extensions["chatgpt_realtime"]["extra_headers"] == {"openai-alpha": "quicksilver=v2", "x-gateway-route": "voice"} assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com") assert response.status_code == 201 assert requests[0].url.path == "/backend-api/codex/realtime/calls" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 59c08c09b00..303db86b4a7 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -52,11 +52,13 @@ def test_sideband_token_binds_owner_and_model(monkeypatch): 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") From ef8d102af52a610704d67b24530fec8b5ca591ad Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 23:57:14 +0200 Subject: [PATCH 22/90] style(chatgpt): format realtime header configuration --- litellm/llms/chatgpt/realtime.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 97783c2bf38..10a888969de 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -58,7 +58,10 @@ class ChatGPTRealtime(OpenAIRealtime): 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 + 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) From e5b2f1022011c49b029829313b2324c51b3ed07c Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 03:12:31 +0200 Subject: [PATCH 23/90] fix(chatgpt): preserve deployment headers in routed signaling --- litellm/llms/chatgpt/codex.py | 2 +- litellm/llms/chatgpt/realtime.py | 19 +++++ litellm/realtime_api/main.py | 9 ++- .../llms/chatgpt/test_realtime.py | 71 ++++++++++++++----- .../realtime_endpoints/test_call_sessions.py | 3 +- 5 files changed, 83 insertions(+), 21 deletions(-) diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index bc812547d6b..78e71963b7f 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -49,7 +49,7 @@ def build_call_request( "extra_query": { # mutable-ok: router request parameters key: value for key, value in query.items() if key in ("intent", "architecture") }, - "extra_headers": { # mutable-ok: router request headers + "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") diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 10a888969de..357da68f5f1 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -23,6 +23,25 @@ def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping return MappingProxyType({key.lower(): value for key, value in validated.items()}) +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 diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index e5ec85b5c2a..bceab5db353 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -261,6 +261,8 @@ async def arealtime_calls( timeout: float | None = None, **kwargs, ): + from litellm.llms.chatgpt.realtime import realtime_call_headers + model_name = model or "gpt-4o-realtime-preview" litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj") litellm_params: Final = GenericLiteLLMParams(**kwargs) @@ -283,6 +285,9 @@ async def arealtime_calls( ) if session is not None: session = _with_resolved_session_model(session, model_name) + call_headers: Final = ( + realtime_call_headers(litellm_params) if custom_llm_provider == "chatgpt" else kwargs.get("extra_headers") + ) litellm_logging_obj.update_from_kwargs( kwargs=kwargs, model=model_name, @@ -299,7 +304,7 @@ 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, ) @@ -310,7 +315,7 @@ async def arealtime_calls( { "model": model_name, "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), - "extra_headers": configured_realtime_headers(kwargs.get("extra_headers")), + "extra_headers": configured_realtime_headers(call_headers), } ) return response diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 3dd74a64d8d..d0760617ba0 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -1,11 +1,9 @@ -import asyncio import json from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, patch import httpx import pytest -from websockets.asyncio.server import serve import litellm from litellm.llms.chatgpt.realtime import ChatGPTRealtime @@ -13,16 +11,48 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.router import GenericLiteLLMParams +@pytest.mark.asyncio +@pytest.mark.parametrize("inbound_headers", [{}, {"openai-alpha": "quicksilver=v2"}]) +async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, chatgpt_tokens, monkeypatch): + from litellm.llms.chatgpt.codex import CodexRealtimeOffer, build_call_request + + 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": "chatgpt/gpt-live-1-codex", + "api_base": "https://voice.example/backend-api/codex", + "extra_headers": {"x-gateway-route": "configured"}, + }, + } + ], + num_retries=0, + ) + offer = CodexRealtimeOffer(sdp="v=0\r\n", session={"model": "voice-gateway"}) + try: + response = await router.arealtime_calls(**build_call_request(offer, {}, inbound_headers), client=client) + assert requests[0].headers.get("x-gateway-route") == "configured" + assert response.extensions["chatgpt_realtime"]["extra_headers"]["x-gateway-route"] == "configured" + for name, value in inbound_headers.items(): + assert requests[0].headers[name] == value + 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): - captured = asyncio.get_running_loop().create_future() - - async def receive_connection(connection): - captured.set_result(connection.request.headers) - await connection.wait_closed() - websocket = SimpleNamespace( headers={"authorization": "Bearer client", "cookie": "private-cookie", "openai-alpha": "client-value"}, scope={}, @@ -30,16 +60,23 @@ async def test_websocket_forwards_configured_headers_without_client_identity(mod send_text=AsyncMock(), close=AsyncMock(), ) - async with serve(receive_connection, "127.0.0.1", 0) as gateway: - port = gateway.sockets[0].getsockname()[1] - await asyncio.wait_for(litellm._arealtime( - model=f"chatgpt/{model}", websocket=websocket, api_base=f"http://127.0.0.1:{port}", + 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, headers={"x-deployment-header": "configured"}, - extra_headers={"X-Gateway-Route": "voice", "OpenAI-Alpha": "configured-value", - "aUtHoRiZaTiOn": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, - ), timeout=10) - headers = await asyncio.wait_for(captured, timeout=5) + 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"]) assert headers["x-deployment-header"] == "configured" assert headers["x-gateway-route"] == "voice" assert headers["openai-alpha"] == "configured-value" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 303db86b4a7..08a86d93644 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -199,7 +199,8 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, data = kwargs["data"] assert data["sdp_body"] == b"v=0\r\n" assert data["session"] == session - assert data["extra_headers"] == {"openai-alpha": "quicksilver=v2"} + assert data["chatgpt_realtime_client_headers"] == {"openai-alpha": "quicksilver=v2"} + assert "extra_headers" not in data assert data["extra_query"] == {"intent": "quicksilver", "architecture": "avas"} async def respond(): From 297c9fdfe8178cb4a2c08f3911703e9eab4fd8d4 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 04:47:54 +0200 Subject: [PATCH 24/90] fix(chatgpt): apply model guardrails and realtime gateway URLs --- litellm/llms/chatgpt/realtime.py | 9 ++--- .../proxy/realtime_endpoints/call_sessions.py | 1 + litellm/realtime_api/main.py | 8 ++--- .../llms/chatgpt/test_realtime.py | 34 +++++++++++++++++++ .../realtime_endpoints/test_call_sessions.py | 28 +++++++++++++++ 5 files changed, 72 insertions(+), 8 deletions(-) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 357da68f5f1..c44a1e4f85e 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -119,15 +119,16 @@ class ChatGPTRealtime(OpenAIRealtime): class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): realtime_calls_json: Final = True - def __init__(self, params: GenericLiteLLMParams) -> None: + 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() + return api_base or (Authenticator.get_api_base() if self._use_codex_backend else ChatGPTRealtime.get_api_base()) def get_api_key( self, @@ -159,7 +160,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): } def get_complete_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: - return "https://api.openai.com/v1/realtime/client_secrets" + return f"{self.get_api_base(api_base).rstrip('/')}/realtime/client_secrets" def get_transcription_session_url( self, @@ -167,4 +168,4 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): model: str, api_version: str | None = None, ) -> str: - return "https://api.openai.com/v1/realtime/transcription_sessions" + return f"{self.get_api_base(api_base).rstrip('/')}/realtime/transcription_sessions" diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index b1d1112af21..38e1abba801 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -74,6 +74,7 @@ async def process_codex_request( version=server.version, proxy_logging_obj=server.proxy_logging_obj, proxy_config=server.proxy_config, + llm_router=server.llm_router, user_model=server.user_model, user_temperature=server.user_temperature, user_request_timeout=server.user_request_timeout, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index bceab5db353..849934b59cd 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -76,6 +76,7 @@ def _get_realtime_http_provider_config( dynamic_api_base: str | None, dynamic_api_key: str | None, litellm_params: GenericLiteLLMParams, + use_codex_backend: bool = False, ) -> tuple["BaseRealtimeHTTPConfig | None", str, str]: """ Return (provider_config, resolved_api_base, resolved_api_key) for the @@ -92,7 +93,7 @@ def _get_realtime_http_provider_config( if custom_llm_provider == "chatgpt": from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig - provider_config = ChatGPTRealtimeHTTPConfig(litellm_params) + provider_config = ChatGPTRealtimeHTTPConfig(litellm_params, use_codex_backend=use_codex_backend) elif custom_llm_provider in LlmProviders._member_map_.values(): provider_config = ProviderConfigManager.get_provider_realtime_http_config( model="", @@ -100,9 +101,7 @@ def _get_realtime_http_provider_config( ) raw_api_base: Final = ( - litellm_params.api_base or dynamic_api_base - if custom_llm_provider == "chatgpt" - else dynamic_api_base or litellm_params.api_base + litellm_params.api_base if custom_llm_provider == "chatgpt" else dynamic_api_base or litellm_params.api_base ) raw_api_key: Final = dynamic_api_key or litellm_params.api_key @@ -282,6 +281,7 @@ async def arealtime_calls( dynamic_api_base=dynamic_api_base, dynamic_api_key=dynamic_api_key, litellm_params=litellm_params, + use_codex_backend=True, ) if session is not None: session = _with_resolved_session_model(session, model_name) diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index d0760617ba0..b299e2a6842 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -11,6 +11,40 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.router import GenericLiteLLMParams +@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"}]) async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, chatgpt_tokens, monkeypatch): diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 08a86d93644..78c2dc628e0 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -11,6 +11,34 @@ 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("route_type", ["arealtime_calls", "_arealtime"]) +async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type): + 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): + 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) + 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]) From bba4fd50e8e13a40a93efcac1f9bdc9aa68f9529 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 05:29:08 +0200 Subject: [PATCH 25/90] fix(chatgpt): retain alternate signaling keys and hook headers --- .../proxy/realtime_endpoints/call_sessions.py | 30 ++++++++++++++++--- .../realtime_endpoints/test_call_sessions.py | 22 +++++++++----- 2 files changed, 41 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 38e1abba801..308cf389936 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -20,9 +20,10 @@ from litellm.llms.chatgpt.codex import ( build_sideband_request, parse_call_response, ) +from litellm.llms.chatgpt.realtime import configured_realtime_headers from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.auth.user_api_key_auth import get_api_key, get_api_key_from_custom_header, user_api_key_auth from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation @@ -99,11 +100,26 @@ async def create_codex_realtime_call(request: Request) -> Response: auth: Final = await user_api_key_auth( request=request, api_key=request.headers.get("authorization", ""), - azure_api_key_header="", + 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=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 ) try: await can_key_call_resolved_model( @@ -132,7 +148,7 @@ async def create_codex_realtime_call(request: Request) -> Response: call: Final = parse_call_response( response, alias=model, - owner=hashlib.sha256(request.headers.get("authorization", "").encode()).hexdigest(), + owner=hashlib.sha256(f"Bearer {owner_key}".encode()).hexdigest(), expires_at=time.time() + 3600, ) except ValueError as exc: @@ -210,6 +226,12 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP **{ # mutable-ok: retain processed policy metadata while pinning the existing call's routing **processed, **build_sideband_request(call), + "extra_headers": MappingProxyType( + { + **configured_realtime_headers(processed.get("extra_headers")), + **configured_realtime_headers(call.extra_headers), + } + ), "websocket": websocket, "user_api_key_dict": auth, } diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 78c2dc628e0..5d844b7be5d 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -173,7 +173,8 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id @pytest.mark.asyncio @pytest.mark.parametrize("multipart", [False, True]) @pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"]) -async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential): +@pytest.mark.parametrize("signaling_credential", ["authorization", "api-key", "x-litellm-api-key", "mixed"]) +async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential, signaling_credential): import json from unittest.mock import AsyncMock @@ -197,15 +198,21 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, 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": "/v1/realtime/calls", + "scheme": "http", "server": ("localhost", 80), "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", "headers": [(b"content-type", body_request.headers["content-type"].encode()), - (b"authorization", b"Bearer owner"), (b"openai-alpha", b"quicksilver=v2"), + *signaling_headers, (b"openai-alpha", b"quicksilver=v2"), (b"x-untrusted", b"bad")]}, receive) auth = UserAPIKeyAuth() - authenticate = AsyncMock(return_value=auth) authorize = AsyncMock() - monkeypatch.setattr(codex, "user_api_key_auth", authenticate) + monkeypatch.setattr(proxy_server, "master_key", "owner") + monkeypatch.setattr(proxy_server, "general_settings", {}) monkeypatch.setattr(codex, "can_key_call_resolved_model", authorize) class Processor: @@ -213,12 +220,12 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, self.data = data async def common_processing_pre_call_logic(self, **kwargs): - assert kwargs["user_api_key_dict"] is auth + 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, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None + return {**self.data, "extra_headers": {"X-Hook-Required": "policy-value", "x-gateway-token": "untrusted-override", "Authorization": "Bearer untrusted"}, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None return self.data, None monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor) @@ -233,7 +240,7 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, 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"}}) + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", "extra_headers": {"X-Gateway-Token": "pinned-value"}}}) return respond() monkeypatch.setattr(proxy_server, "route_request", route) @@ -270,6 +277,7 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, 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["chatgpt_realtime_call_id"] == "rtc_private" assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-codex" From a9cd1dc56d503ad61a738c8efb1bbb20b95da2e1 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 06:14:13 +0200 Subject: [PATCH 26/90] fix(chatgpt): share websocket credential selection for call ownership --- litellm/proxy/auth/user_api_key_auth.py | 61 ++++++++++++------- .../proxy/realtime_endpoints/call_sessions.py | 16 +++-- .../proxy/auth/test_user_api_key_auth.py | 42 ++++++++++++- .../realtime_endpoints/test_call_sessions.py | 8 ++- 4 files changed, 96 insertions(+), 31 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index f6f4d7bf1c4..fec5a1268d2 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -487,6 +487,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): # Accept the WebSocket connection @@ -509,27 +541,14 @@ async def user_api_key_auth_websocket(websocket: WebSocket): request._url = websocket.url - 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 diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 308cf389936..e02a66193b8 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -23,7 +23,12 @@ from litellm.llms.chatgpt.codex import ( from litellm.llms.chatgpt.realtime import configured_realtime_headers from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model -from litellm.proxy.auth.user_api_key_auth import get_api_key, get_api_key_from_custom_header, user_api_key_auth +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.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation @@ -171,14 +176,13 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP protocols: Final = tuple( p.strip() for p in websocket.headers.get("sec-websocket-protocol", "").split(",") if p.strip() ) - alternate_key: Final = websocket.headers.get("api-key") or next( - (p.removeprefix("openai-insecure-api-key.") for p in protocols if p.startswith("openai-insecure-api-key.")), "" - ) - authorization: Final = websocket.headers.get("authorization") or f"Bearer {alternate_key}" logging_obj: Logging | None = None # rebind-ok: cleanup needs the logger only after pre-call succeeds try: try: - call: Final = decode_call(token, authorization) + 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=server.llm_model_list, 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 886d3361c49..f5d802d843c 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 @@ -7137,7 +7137,7 @@ def test_user_api_key_auth_opens_a_datadog_span_for_accepted_and_rejected_keys(t @pytest.mark.asyncio @pytest.mark.parametrize("attachment", ["path", "query"]) -@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"]) +@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 @@ -7154,6 +7154,8 @@ async def test_sideband_auth_uses_encrypted_model_for_budget_checks(monkeypatch, 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): @@ -7169,6 +7171,9 @@ async def test_sideband_auth_uses_encrypted_model_for_budget_checks(monkeypatch, "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()) @@ -7336,3 +7341,38 @@ async def test_sideband_rejects_budget_fallback_before_rerouting(monkeypatch, at 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": ""}) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 5d844b7be5d..32b13ad678a 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -172,7 +172,7 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id @pytest.mark.asyncio @pytest.mark.parametrize("multipart", [False, True]) -@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"]) +@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom"]) @pytest.mark.parametrize("signaling_credential", ["authorization", "api-key", "x-litellm-api-key", "mixed"]) async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential, signaling_credential): import json @@ -207,12 +207,12 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, "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"openai-alpha", b"quicksilver=v2"), + *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", {}) + 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: @@ -267,6 +267,8 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, 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", From d6b8360ac3f3ca393a895e0f21e2a8f5bf2d17b8 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 12:52:04 +0200 Subject: [PATCH 27/90] fix(chatgpt): supervise call usage and preserve gateway query routing --- litellm/cost_calculator.py | 45 ++- .../litellm_core_utils/realtime_streaming.py | 21 +- litellm/llms/chatgpt/codex.py | 11 +- litellm/llms/chatgpt/realtime.py | 71 ++++- litellm/llms/openai/realtime/handler.py | 2 + litellm/proxy/proxy_server.py | 3 + .../proxy/realtime_endpoints/call_sessions.py | 148 ++++++++- .../realtime_endpoints/call_supervision.py | 187 +++++++++++ litellm/realtime_api/main.py | 11 +- litellm/types/llms/openai.py | 6 + .../test_realtime_streaming.py | 26 ++ tests/test_litellm/llms/chatgpt/test_codex.py | 23 +- .../llms/chatgpt/test_realtime.py | 68 +++- .../realtime_endpoints/test_call_sessions.py | 252 +++++++++++++-- .../test_call_supervision.py | 301 ++++++++++++++++++ tests/test_litellm/test_cost_calculator.py | 68 ++++ 16 files changed, 1195 insertions(+), 48 deletions(-) create mode 100644 litellm/proxy/realtime_endpoints/call_supervision.py create mode 100644 tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 814eaaf76f7..24e048eafd1 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, Field, ValidationError import litellm import litellm._logging @@ -2563,7 +2563,12 @@ def handle_realtime_stream_cost_calculation( if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results) else 0.0 ) - total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost + live_audio_cost: Final = handle_live_session_duration_cost( + results=results, + custom_llm_provider=custom_llm_provider, + litellm_model_name=litellm_model_name, + ) + total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost + live_audio_cost _store_cost_breakdown_in_logging_obj( litellm_logging_obj=litellm_logging_obj, @@ -2571,13 +2576,47 @@ def handle_realtime_stream_cost_calculation( completion_tokens_cost_usd_dollar=output_cost_per_token, cost_for_built_in_tools_cost_usd_dollar=0.0, total_cost_usd_dollar=total_cost, - additional_costs={"transcription_cost": transcription_cost} if transcription_cost > 0 else None, + additional_costs={ + name: cost + for name, cost in (("transcription_cost", transcription_cost), ("live_audio_cost", live_audio_cost)) + if cost > 0 + } + or None, data_residency=data_residency, ) return total_cost +class _LiveSessionDurationUsage(BaseModel): + audio_duration_ms: float = Field(strict=True, ge=0, allow_inf_nan=False) + + +class _LiveSessionClosedEvent(BaseModel): + usage: _LiveSessionDurationUsage + + +def handle_live_session_duration_cost( + results: OpenAIRealtimeStreamList, + custom_llm_provider: str, + litellm_model_name: str, +) -> float: + if any(event.get("type") == "response.done" for event in results): + return 0.0 + terminal: Final = next((event for event in reversed(results) if event.get("type") == "session.closed"), None) + if terminal is None: + return 0.0 + try: + usage: Final = _LiveSessionClosedEvent.model_validate(terminal).usage + except ValidationError: + return 0.0 + try: + model_info: Final = litellm.get_model_info(model=litellm_model_name, custom_llm_provider=custom_llm_provider) + except Exception: + return 0.0 + return usage.audio_duration_ms / 1000 * (model_info.get("input_cost_per_second") or 0.0) + + def handle_realtime_transcription_cost_calculation( results: OpenAIRealtimeStreamList, custom_llm_provider: str, diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 75046f2cf87..c5e6b0fafe3 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -6,6 +6,7 @@ from dataclasses import dataclass from enum import Enum, auto from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast +from pydantic import TypeAdapter from typing_extensions import ReadOnly import litellm @@ -16,6 +17,7 @@ from litellm.types.llms.openai import ( OpenAIRealtimeEvents, OpenAIRealtimeOutputItemDone, OpenAIRealtimeResponseDelta, + OpenAIRealtimeSessionClosed, OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents, ) @@ -139,11 +141,14 @@ class RealTimeStreaming: force_transcription_model: str | None = None, event_normalizer: RealtimeEventNormalizer | None = None, logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER, + *, + account_usage: bool = True, ): self.websocket: _ClientWebSocket = websocket self.backend_ws = backend_ws self.logging_obj = logging_obj self._logging_worker = logging_worker + self._account_usage = account_usage self.messages: list[OpenAIRealtimeEvents] = [] self._backend_sent_frames: bool = False self.input_message: dict = {} @@ -256,6 +261,9 @@ class RealTimeStreaming: else: message_obj = cast(dict[str, Any], json.loads(cast(str, message))) self._collect_tool_calls_from_response_done(cast(dict, message_obj)) + if message_obj.get("type") == "session.closed" and isinstance(message_obj.get("usage"), dict): + self.messages.append(TypeAdapter(OpenAIRealtimeSessionClosed).validate_python(message_obj)) + return if not self._should_store_message(message_obj): return try: @@ -410,8 +418,10 @@ class RealTimeStreaming: if self.logging_obj: self.logging_obj.pre_call(input=message, api_key="") - async def log_messages(self): + async def log_messages(self, *, wait_for_dispatch: bool = False): """Log messages in list""" + if not self._account_usage: + return if self.logging_obj: if self.input_messages: self.logging_obj.model_call_details["messages"] = self.input_messages @@ -421,9 +431,12 @@ class RealTimeStreaming: # Route through the bounded logging worker (per-coroutine timeout + # concurrency cap) instead of a bare create_task, so a slow callback # can't leave suspended tasks pinning each call's response in memory. - self._logging_worker.ensure_initialized_and_enqueue( - self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True) - ) + if wait_for_dispatch: + await self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True) + else: + self._logging_worker.ensure_initialized_and_enqueue( + self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True) + ) self.logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True async def _send_to_backend(self, message: str) -> bool: diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 78e71963b7f..f71a4d5bbad 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -17,17 +17,22 @@ class CodexRealtimeOffer(BaseModel): class CodexRealtimeCall(BaseModel): call_id: str = Field(pattern=r"^rtc_[A-Za-z0-9_-]+$") model: str + model_id: str | None = None alias: str api_base: str | None = None extra_headers: Mapping[str, str] | None = None + extra_query: Mapping[str, str] | None = None + usage_supervised: bool = False owner: str expires_at: float class ChatGPTCallRouting(BaseModel): model: str + model_id: str | None = None api_base: str | None = None extra_headers: Mapping[str, str] | None = None + extra_query: Mapping[str, str] | None = None class CodexSidebandRequest(TypedDict): @@ -36,6 +41,7 @@ class CodexSidebandRequest(TypedDict): chatgpt_realtime_call_id: ReadOnly[str] query_params: ReadOnly[RealtimeQueryParams] extra_headers: ReadOnly[Mapping[str, str] | None] + extra_query: ReadOnly[Mapping[str, str] | None] def build_call_request( @@ -46,7 +52,7 @@ def build_call_request( "sdp_body": offer.sdp.encode(), "session": offer.session.model_dump(exclude_none=True), "openai_ephemeral_key": "", - "extra_query": { # mutable-ok: router request parameters + "chatgpt_realtime_client_query": { # mutable-ok: router request parameters key: value for key, value in query.items() if key in ("intent", "architecture") }, "chatgpt_realtime_client_headers": { # mutable-ok: router request headers @@ -66,11 +72,13 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire return CodexRealtimeCall( call_id=call_id, model=routing.model, + model_id=routing.model_id, alias=alias, owner=owner, expires_at=expires_at, api_base=routing.api_base, extra_headers=routing.extra_headers, + extra_query=routing.extra_query, ) @@ -81,4 +89,5 @@ def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: chatgpt_realtime_call_id=call.call_id, query_params=RealtimeQueryParams(model=call.model), extra_headers=call.extra_headers, + extra_query=call.extra_query, ) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index c44a1e4f85e..f79713fe03f 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -1,10 +1,12 @@ from collections.abc import Mapping +from enum import Enum, auto from types import MappingProxyType -from typing import Final +from typing import TYPE_CHECKING, Final from httpx import URL from pydantic import TypeAdapter +from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.llms.openai.realtime.handler import OpenAIRealtime from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig from litellm.types.realtime import RealtimeQueryParams @@ -15,6 +17,17 @@ from .authenticator import Authenticator from .common_utils import without_oauth_identity_headers from .responses.transformation import ChatGPTResponsesAPIConfig +if TYPE_CHECKING: + from websockets.asyncio.client import ClientConnection + + +class CallAccounting(Enum): + SUPERVISED = auto() + + +def accounts_for_call_usage(params: GenericLiteLLMParams) -> bool: + return getattr(params, "chatgpt_call_accounting", None) is not CallAccounting.SUPERVISED + def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping[str, str]: validated: Final = TypeAdapter(Mapping[str, str]).validate_python( @@ -23,6 +36,18 @@ def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping return MappingProxyType({key.lower(): value for key, value in validated.items()}) +def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str]: + inbound: Final = TypeAdapter(Mapping[str, str]).validate_python( + getattr(params, "chatgpt_realtime_client_query", None) or MappingProxyType({}) + ) + configured: Final = TypeAdapter(Mapping[str, str]).validate_python( + getattr(params, "extra_query", None) or MappingProxyType({}) + ) + return MappingProxyType( + {**{key: value for key, value in inbound.items() if key in ("intent", "architecture")}, **configured} + ) + + def realtime_call_headers(params: GenericLiteLLMParams) -> dict[str, str]: # mutable-ok: HTTP handler header contract inbound: Final = TypeAdapter(Mapping[str, str]).validate_python( getattr(params, "chatgpt_realtime_client_headers", None) or MappingProxyType({}) @@ -72,6 +97,40 @@ def realtime_endpoint(model: str) -> str: class ChatGPTRealtime(OpenAIRealtime): + async def open_call_connection(self, model: str, api_base: str) -> "ClientConnection": + import websockets + + url: Final = self._construct_url(api_base, RealtimeQueryParams(model=model)) + return await websockets.connect( + url, + additional_headers=self._profile_headers, + max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, + ssl=self._get_ssl_config(url), + open_timeout=20, + ) + + async def close_call(self, connection: "ClientConnection", model: str, api_base: str) -> None: + if realtime_endpoint(model) == "live": + await connection.send('{"type":"session.close"}') + return + await self.hangup_call(api_base) + + async def hangup_call(self, api_base: str) -> None: + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + base: Final = URL(api_base) + url: Final = base.copy_with( + scheme="https" if base.scheme in ("https", "wss") else "http", + path=f"{base.path.rstrip('/')}/realtime/calls/{self._call_id}/hangup", + params=tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")), + ) + client: Final = AsyncHTTPHandler() + try: + response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10) + response.raise_for_status() + finally: + await client.close() + @staticmethod def get_api_base(api_base: str | None = None) -> str: return api_base or Authenticator.get_api_base(default_base="https://api.openai.com/v1") @@ -85,6 +144,7 @@ class ChatGPTRealtime(OpenAIRealtime): super().__init__() self._profile_headers = realtime_headers(params, headers, extra_headers) self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None)) + self._extra_query = configured_realtime_query(params) def _get_additional_headers( self, api_key: str, *, openai_beta_realtime: bool = False @@ -98,13 +158,16 @@ class ChatGPTRealtime(OpenAIRealtime): base: Final = URL(api_base) endpoint: Final = realtime_endpoint(query_params.get("model", "")) if self._call_id: + gateway_query: Final = tuple( + (key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id") + ) return str( base.copy_with( scheme="wss" if base.scheme in ("https", "wss") else "ws", path=f"{base.path.rstrip('/')}/{endpoint}/{self._call_id}" if endpoint == "live" else f"{base.path.rstrip('/')}/realtime", - params=() if endpoint == "live" else (("call_id", self._call_id),), + params=gateway_query + (() if endpoint == "live" else (("call_id", self._call_id),)), ) ) return str( @@ -138,9 +201,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): return "chatgpt-oauth" def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: - query: Final = TypeAdapter(Mapping[str, str]).validate_python( - getattr(self._params, "extra_query", None) or MappingProxyType({}) - ) + query: Final = configured_realtime_query(self._params) return str(URL(f"{self.get_api_base(api_base).rstrip('/')}/realtime/calls", params=query)) def get_realtime_calls_headers( diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index e3ecbac1a53..ca141c0b958 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -117,6 +117,7 @@ class OpenAIRealtime(OpenAIChatCompletion): query_params: RealtimeQueryParams | None = None, user_api_key_dict: object | None = None, litellm_metadata: dict | None = None, + account_usage: bool = True, **kwargs: object, ): import websockets @@ -172,6 +173,7 @@ class OpenAIRealtime(OpenAIChatCompletion): model if (query_params or {}).get("intent") == "transcription" else None ), event_normalizer=self._make_event_normalizer(), + account_usage=account_usage, ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8f631ccde52..a9ed397845b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1362,6 +1362,9 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: except Exception as e: verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e) + from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS + + await CALL_SUPERVISORS.shutdown() await _flush_spend_logs_queue_on_shutdown() await proxy_config.stop_config_sync_subscriber() diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index e02a66193b8..b0833b19414 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -2,6 +2,7 @@ import base64 import hashlib import json import time +from contextlib import AsyncExitStack from types import MappingProxyType from typing import Final, Literal @@ -11,7 +12,7 @@ from starlette.types import Message from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY +from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY, RealTimeStreaming from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.chatgpt.codex import ( CodexRealtimeCall, @@ -20,7 +21,12 @@ from litellm.llms.chatgpt.codex import ( build_sideband_request, parse_call_response, ) -from litellm.llms.chatgpt.realtime import configured_realtime_headers +from litellm.llms.chatgpt.realtime import ( + CallAccounting, + ChatGPTRealtime, + configured_realtime_headers, + realtime_endpoint, +) from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model from litellm.proxy.auth.user_api_key_auth import ( @@ -30,7 +36,126 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth, ) from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper -from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation +from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + release_or_invalidate_budget_reservation, +) +from litellm.types.router import GenericLiteLLMParams + + +async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth) -> None: + from collections.abc import Mapping + + from pydantic import TypeAdapter + + import litellm + from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, CallSupervisor + + async def receive() -> Message: + return { + "type": "http.request", + "body": json.dumps({"model": call.alias}).encode(), + "more_body": False, + } # mutable-ok: ASGI message + + async def send(_message: Message) -> None: + return None + + supervision_owned = False # rebind-ok: supervisor owns cleanup after construction + effective_handler: ChatGPTRealtime | None = None # rebind-ok: reuse hook-enriched credentials for cleanup + sockets: Final = AsyncExitStack() + try: + observer_request: Final = Request({**request.scope}, receive=receive) # mutable-ok: ASGI request scope + processed, logger = await process_codex_request( + observer_request, + { + **build_sideband_request(call), + "model": call.alias, + }, # mutable-ok: common request processing enriches metadata + auth, + call.alias, + "_arealtime", + ) + pinned: Final = { # mutable-ok: logging and provider parameter contract + **processed, + **build_sideband_request(call), + "extra_headers": { + **configured_realtime_headers( + TypeAdapter(Mapping[str, object] | None).validate_python(processed.get("extra_headers")) + ), + **configured_realtime_headers(call.extra_headers), + }, + "litellm_metadata": { + **TypeAdapter(Mapping[str, object]).validate_python(processed.get("litellm_metadata") or {}), + **( + {"model_info": {**litellm.get_model_info(model=call.model_id), "id": call.model_id}} + if call.model_id is not None + else {} + ), + }, + } + logger.update_from_kwargs( + kwargs=pinned, + model=call.model, + user=None, + optional_params={}, # mutable-ok: logging contract + litellm_params={ + **logger.litellm_params, + "litellm_metadata": pinned["litellm_metadata"], + "arealtime": True, + }, # mutable-ok: logging contract + custom_llm_provider="chatgpt", + ) + params: Final = GenericLiteLLMParams.model_validate(pinned) + handler: Final = ChatGPTRealtime( + params, request.headers, TypeAdapter(Mapping[str, object]).validate_python(pinned["extra_headers"]) + ) + effective_handler = handler + api_base: Final = ChatGPTRealtime.get_api_base(call.api_base) + connection: Final = await handler.open_call_connection(call.model, api_base) + sockets.push_async_callback(connection.close) + + async def close_call() -> None: + await handler.close_call(connection, call.model, api_base) + + frontend: Final = WebSocket( + {**request.scope, "type": "websocket"}, receive=receive, send=send + ) # mutable-ok: ASGI scope + stream: Final = RealTimeStreaming(frontend, connection, logger, model=call.model, user_api_key_dict=auth) + supervisor: Final = CallSupervisor( + connection, + stream, + logger, + auth, + close_call, + terminal_usage_required=realtime_endpoint(call.model) == "live", + ) + supervision_owned = True + sockets.pop_all() + await CALL_SUPERVISORS.start(supervisor) + except BaseException: + if not supervision_owned: + try: + fallback_handler: Final = effective_handler or ChatGPTRealtime( + GenericLiteLLMParams.model_validate(build_sideband_request(call)), + request.headers, + call.extra_headers, + ) + await fallback_handler.hangup_call(ChatGPTRealtime.get_api_base(call.api_base)) + except Exception: # noqa: BLE001 # preserve original failure without logging provider credentials + verbose_proxy_logger.error("Realtime startup cleanup could not confirm upstream termination") + try: + await invalidate_budget_reservation_counters(budget_reservation=auth.budget_reservation) + except Exception: # noqa: BLE001 # cleanup errors must not replace the original startup failure + verbose_proxy_logger.error("Realtime startup cleanup could not invalidate budget counters") + else: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + finally: + try: + await sockets.aclose() + except Exception: # noqa: BLE001 # socket cleanup must preserve the original startup failure + verbose_proxy_logger.error("Realtime startup cleanup could not close observer socket") + raise def encode_call(call: CodexRealtimeCall) -> str: @@ -126,6 +251,7 @@ async def create_codex_realtime_call(request: Request) -> Response: owner_key: Final = ( get_api_key_from_custom_header(request, custom_header) if isinstance(custom_header, str) else selected_key ) + supervision_started = False # rebind-ok: transfer reservation ownership only after supervision is established try: await can_key_call_resolved_model( model=model, @@ -134,7 +260,10 @@ async def create_codex_realtime_call(request: Request) -> Response: llm_router=server.llm_router, ) data: Final = build_call_request(offer, request.query_params, request.headers) - processed, _ = await process_codex_request(request, data, auth, model, "arealtime_calls") + signaling_auth: Final = auth.model_copy( + update={"budget_reservation": None} + ) # mutable-ok: Pydantic update contract + processed, _ = await process_codex_request(request, data, signaling_auth, model, "arealtime_calls") result: Final = await server.route_request( data=processed, route_type="arealtime_calls", @@ -158,7 +287,12 @@ async def create_codex_realtime_call(request: Request) -> Response: ) except ValueError as exc: raise HTTPException(400, str(exc)) from exc - token: Final = encode_call(call) + supervised_call: Final = call.model_copy( + update={"usage_supervised": True} + ) # mutable-ok: Pydantic update contract + token: Final = encode_call(supervised_call) + supervision_started = True + await supervise_codex_call(request, supervised_call, auth) return Response( response.content, status_code=response.status_code, @@ -166,7 +300,8 @@ async def create_codex_realtime_call(request: Request) -> Response: headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), ) finally: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + if not supervision_started: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAPIKeyAuth) -> None: @@ -238,6 +373,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP ), "websocket": websocket, "user_api_key_dict": auth, + "chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None, } ) finally: diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py new file mode 100644 index 00000000000..98f018353ba --- /dev/null +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -0,0 +1,187 @@ +import asyncio +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import suppress +from typing import Final, Protocol + +from pydantic import BaseModel +from websockets.exceptions import ConnectionClosedOK + +from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + release_or_invalidate_budget_reservation, +) + + +class ObserverSocket(Protocol): + def __aiter__(self) -> AsyncIterator[str | bytes]: ... + + async def close(self) -> None: ... + + +class UsageSink(Protocol): + def store_message(self, message: str) -> None: ... + + async def log_messages(self, *, wait_for_dispatch: bool = False) -> None: ... + + +class _ObserverEvent(BaseModel): + type: str + + +class CallSupervisor: + def __init__( + self, + upstream: ObserverSocket, + stream: UsageSink, + logging_obj: Logging, + auth: UserAPIKeyAuth, + close_call: Callable[[], Awaitable[None]], + *, + ready_timeout: float = 20, + lifetime: float = 3600, + drain_timeout: float = 5, + terminal_usage_required: bool = True, + ) -> None: + self._upstream = upstream + self._stream = stream + self._logging = logging_obj + self._auth = auth + self._close_call = close_call + self._ready_timeout = ready_timeout + self._lifetime = lifetime + self._drain_timeout = drain_timeout + self._terminal_usage_required = terminal_usage_required + self._ready = asyncio.Event() + self._stop = asyncio.Event() + self._started = False + self._terminal = False + self._close_confirmed = False + self._task: asyncio.Task[None] | None = None + + async def start(self) -> None: + if self._task is not None: + raise RuntimeError("Call observer already started") + self._task = asyncio.create_task(self._run()) + try: + await asyncio.wait_for(self._ready.wait(), timeout=self._ready_timeout) + if not self._started or self._task.done(): + raise RuntimeError("Call observer ended before session became available") + except BaseException: + await self.close() + raise + + async def close(self) -> None: + self._stop.set() + await self.wait() + + async def wait(self) -> None: + if self._task is not None: + await asyncio.shield(self._task) + + async def _read(self) -> None: + try: + await self._read_events() + except ConnectionClosedOK: + return + + async def _read_events(self) -> None: + event: _ObserverEvent + async for message in self._upstream: + self._stream.store_message(message.decode("utf-8") if isinstance(message, bytes) else message) + event = _ObserverEvent.model_validate_json(message) + if event.type in ("session.started", "session.created"): + self._started = True + self._ready.set() + if event.type == "session.closed": + self._terminal = True + return + + def _usage_complete(self) -> bool: + return self._terminal or (not self._terminal_usage_required and self._close_confirmed) + + async def _run(self) -> None: + reader: Final = asyncio.create_task(self._read()) + stopped: Final = asyncio.create_task(self._stop.wait()) + try: + await asyncio.wait((reader, stopped), timeout=self._lifetime, return_when=asyncio.FIRST_COMPLETED) + finally: + try: + if not self._terminal: + try: + await asyncio.wait_for(self._close_call(), timeout=self._drain_timeout) + self._close_confirmed = True + except Exception: # noqa: BLE001 # provider exceptions can contain credentials + verbose_proxy_logger.error("Realtime observer could not terminate upstream call") + await self._drain(reader) + finally: + stopped.cancel() + reader.cancel() + await asyncio.gather(reader, stopped, return_exceptions=True) + with suppress(Exception): + await self._upstream.close() + if not self._usage_complete(): + self._logging.model_call_details["realtime_usage_incomplete"] = True + verbose_proxy_logger.error( + "Realtime observer ended without terminal usage; recorded usage is partial" + ) + try: + try: + await self._stream.log_messages(wait_for_dispatch=True) + finally: + if self._started and not self._usage_complete(): + await invalidate_budget_reservation_counters( + budget_reservation=self._auth.budget_reservation + ) + elif not self._logging.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY): + await release_or_invalidate_budget_reservation( + budget_reservation=self._auth.budget_reservation + ) + finally: + self._ready.set() + + async def _drain(self, reader: asyncio.Task[None]) -> None: + try: + await asyncio.wait_for(asyncio.shield(reader), timeout=self._drain_timeout) + except asyncio.TimeoutError: + if not self._usage_complete(): + verbose_proxy_logger.error("Realtime observer timed out draining terminal usage") + except Exception: # noqa: BLE001 # cleanup must settle the socket even when reading or closing fails + verbose_proxy_logger.error("Realtime observer could not drain terminal usage") + return + + +class CallSupervisors: + def __init__(self) -> None: + self._tasks: tuple[asyncio.Task[None], ...] = () + self._calls: tuple[CallSupervisor, ...] = () + + async def start(self, supervisor: CallSupervisor) -> None: + self._calls = (*self._calls, supervisor) + try: + await supervisor.start() + except BaseException: + self._calls = tuple(call for call in self._calls if call is not supervisor) + raise + task: Final = asyncio.create_task(self._watch(supervisor)) + self._tasks = (*self._tasks, task) + + async def _watch(self, supervisor: CallSupervisor) -> None: + try: + try: + await supervisor.wait() + except Exception: # noqa: BLE001 # task must be consumed without exposing provider exception payloads + verbose_proxy_logger.error("Realtime observer accounting failed") + finally: + self._calls = tuple(call for call in self._calls if call is not supervisor) + self._tasks = tuple(task for task in self._tasks if task is not asyncio.current_task()) + + async def shutdown(self) -> None: + await asyncio.gather(*(call.close() for call in self._calls), return_exceptions=True) + await asyncio.gather(*self._tasks, return_exceptions=True) + + +CALL_SUPERVISORS: Final = CallSupervisors() diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 849934b59cd..23963475ed0 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -309,13 +309,19 @@ async def arealtime_calls( api_version=litellm_params.api_version, ) if custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtime, configured_realtime_headers + from litellm.llms.chatgpt.realtime import ( + ChatGPTRealtime, + configured_realtime_headers, + configured_realtime_query, + ) response.extensions["chatgpt_realtime"] = MappingProxyType( { "model": model_name, + "model_id": litellm_logging_obj.get_router_model_id(), "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), "extra_headers": configured_realtime_headers(call_headers), + "extra_query": configured_realtime_query(litellm_params), } ) return response @@ -464,7 +470,7 @@ async def _arealtime( litellm_metadata=_build_litellm_metadata(kwargs), ) elif _custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtime + from litellm.llms.chatgpt.realtime import ChatGPTRealtime, accounts_for_call_usage await ChatGPTRealtime(litellm_params, websocket.headers, headers).async_realtime( model=model, @@ -476,6 +482,7 @@ async def _arealtime( query_params=query_params, user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata(kwargs), + account_usage=accounts_for_call_usage(litellm_params), ) elif _custom_llm_provider == "openai": api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or "https://api.openai.com/" diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index b6da9490e01..1c87e355382 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2031,6 +2031,11 @@ class OpenAIRealtimeStreamResponseBaseObject(TypedDict): type: str +class OpenAIRealtimeSessionClosed(TypedDict): + type: ReadOnly[Literal["session.closed"]] + usage: ReadOnly[Mapping[str, object]] + + class OpenAIRealtimeConversationObject(TypedDict, total=False): id: str object: Required[Literal["realtime.conversation"]] @@ -2236,6 +2241,7 @@ class OpenAIRealtimeEventTypes(Enum): OpenAIRealtimeEvents = ( OpenAIRealtimeStreamResponseBaseObject + | OpenAIRealtimeSessionClosed | OpenAIRealtimeStreamSessionEvents | OpenAIRealtimeStreamResponseOutputItemAdded | OpenAIRealtimeResponseContentPartAdded diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 9c0f6f59463..cba46374837 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3412,3 +3412,29 @@ async def test_refused_session_does_not_stamp_the_reservation_ownership_marker() assert session.logging.logged_failures == (upstream_close,) assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in session.logging.model_call_details + + +def test_live_terminal_usage_survives_filtered_event_logging(monkeypatch): + from litellm.cost_calculator import RealtimeAPITokenUsageProcessor + + def terminal(): + return {"type": "session.closed", "usage": {"audio_duration_ms": 4000, "backend_model_usage": []}} + + monkeypatch.setattr(litellm, "logged_real_time_event_types", []) + stream = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + event = {**terminal(), "private_transcript": "Do not retain this text"} + stream.store_message(event) + assert stream.messages == [terminal()] + usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(stream.messages) + assert usage.total_tokens == 0 + + +@pytest.mark.asyncio +async def test_live_attachment_does_not_dispatch_duplicate_usage(): + worker = MagicMock() + logger = MagicMock() + stream = RealTimeStreaming(MagicMock(), MagicMock(), logger, logging_worker=worker, account_usage=False) + stream.store_message({"type": "session.closed", "usage": {"audio_duration_ms": 4000}}) + await stream.log_messages() + worker.ensure_initialized_and_enqueue.assert_not_called() + logger.dispatch_success_handlers.assert_not_called() diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 06f04b9b81c..402afaaa58c 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -1,7 +1,7 @@ import httpx import pytest -from litellm.llms.chatgpt.codex import build_sideband_request, parse_call_response +from litellm.llms.chatgpt.codex import CodexRealtimeCall, build_sideband_request, parse_call_response @pytest.mark.parametrize("location", ["", "/v1/realtime/calls/foreign-id"]) @@ -12,14 +12,25 @@ def test_signaling_rejects_invalid_upstream_call_id(location): parse_call_response(response, "voice", "owner", 1000) -def test_signaling_preserves_selected_model_for_sideband(): - response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", - "extra_headers": {"x-gateway-route": "voice"}}}) +@pytest.mark.parametrize("extra_query", [None, {"gateway_token": "opaque +/& value"}]) +def test_signaling_preserves_selected_model_for_sideband(extra_query): + response = httpx.Response( + 201, + headers={"Location": "/v1/realtime/calls/rtc_provider"}, + extensions={ + "chatgpt_realtime": { + "model": "gpt-live-1-codex", + "api_base": "https://voice.example/codex", + "extra_headers": {"x-gateway-route": "voice"}, + **({"extra_query": extra_query} if extra_query is not None else {}), + } + }, + ) call = parse_call_response(response, "voice", "owner", 1000) - request = build_sideband_request(call) + request = build_sideband_request(CodexRealtimeCall.model_validate_json(call.model_dump_json(exclude_none=True))) assert request["api_base"] == "https://voice.example/codex" assert request["model"] == "chatgpt/gpt-live-1-codex" assert request["chatgpt_realtime_call_id"] == "rtc_provider" assert request["query_params"] == {"model": "gpt-live-1-codex"} assert request["extra_headers"] == {"x-gateway-route": "voice"} + assert request["extra_query"] == extra_query diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index b299e2a6842..8b29d440b4c 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -67,6 +67,7 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, "model": "chatgpt/gpt-live-1-codex", "api_base": "https://voice.example/backend-api/codex", "extra_headers": {"x-gateway-route": "configured"}, + "extra_query": {"gateway_token": "configured", "intent": "pinned-intent"}, }, } ], @@ -74,8 +75,17 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, ) offer = CodexRealtimeOffer(sdp="v=0\r\n", session={"model": "voice-gateway"}) try: - response = await router.arealtime_calls(**build_call_request(offer, {}, inbound_headers), client=client) + response = await router.arealtime_calls( + **build_call_request(offer, {"intent": "quicksilver", "architecture": "avas"}, inbound_headers), + client=client, + ) assert requests[0].headers.get("x-gateway-route") == "configured" + assert dict(requests[0].url.params) == { + "gateway_token": "configured", + "intent": "pinned-intent", + "architecture": "avas", + } + assert response.extensions["chatgpt_realtime"]["extra_query"] == dict(requests[0].url.params) assert response.extensions["chatgpt_realtime"]["extra_headers"]["x-gateway-route"] == "configured" for name, value in inbound_headers.items(): assert requests[0].headers[name] == value @@ -137,6 +147,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap sdp_body=b"v=0\r\n", session={"model": "chatgpt/gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}}, extra_query={"intent": "quicksilver", "architecture": "avas"}, + chatgpt_realtime_client_query={"intent": "untrusted-override", "architecture": "avas", "untrusted": "bad"}, extra_headers={ "openai-alpha": "quicksilver=v2", "x-gateway-route": "voice", @@ -146,9 +157,13 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap client=client, ) assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1") - assert response.extensions["chatgpt_realtime"]["extra_headers"] == {"openai-alpha": "quicksilver=v2", "x-gateway-route": "voice"} + assert response.extensions["chatgpt_realtime"]["extra_headers"] == { + "openai-alpha": "quicksilver=v2", + "x-gateway-route": "voice", + } assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com") assert response.status_code == 201 + assert response.extensions["chatgpt_realtime"]["extra_query"] == {"intent": "quicksilver", "architecture": "avas"} assert requests[0].url.path == "/backend-api/codex/realtime/calls" assert requests[0].url.params["architecture"] == "avas" assert requests[0].headers["authorization"] == "Bearer test-token-" + "default" @@ -244,3 +259,52 @@ def test_realtime_routes_use_configured_gateway(monkeypatch, env_name, api_base, assert handler._construct_url(handler.get_api_base(api_base), {"model": "gpt-realtime-1.5"}) == ( expected.replace("https://", "wss://") + "/realtime?model=gpt-realtime-1.5" ) + + +@pytest.mark.parametrize("model,endpoint", [("gpt-live-1-codex", "live"), ("gpt-realtime-1.5", "realtime")]) +def test_sideband_restores_gateway_query_without_overriding_call(model, endpoint, chatgpt_tokens): + handler = ChatGPTRealtime( + GenericLiteLLMParams( + chatgpt_realtime_call_id="rtc_selected", + extra_query={"gateway_token": "opaque +/& value", "model": "other", "call_id": "rtc_other"}, + ), + {}, + ) + url = httpx.URL(handler._construct_url("https://gateway.example/v1", {"model": model})) + assert url.params["gateway_token"] == "opaque +/& value" + assert "model" not in url.params + if endpoint == "live": + assert url.path == "/v1/live/rtc_selected" + assert "call_id" not in url.params + else: + assert url.path == "/v1/realtime" + assert url.params["call_id"] == "rtc_selected" + + +def test_client_cannot_forge_supervised_call_accounting(chatgpt_tokens): + from litellm.llms.chatgpt.realtime import CallAccounting, accounts_for_call_usage + + assert accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting={"supervised": True})) + assert accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting="supervised")) + assert not accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting=CallAccounting.SUPERVISED)) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ["gpt-live-1-codex", "gpt-realtime-1.5"]) +async def test_supervisor_connection_preserves_call_routing(model, chatgpt_tokens): + handler = ChatGPTRealtime( + GenericLiteLLMParams( + chatgpt_token_dir=chatgpt_tokens, + chatgpt_realtime_call_id="rtc_owner", + extra_query={"gateway_token": "a+b&c"}, + ), + {"openai-alpha": "quicksilver=v2"}, + {"x-gateway-token": "configured"}, + ) + connection = AsyncMock() + with patch("websockets.connect", AsyncMock(return_value=connection)) as connect: + assert await handler.open_call_connection(model, "https://gateway.example/v1") is connection + url = httpx.URL(connect.call_args.args[0]) + assert url.params["gateway_token"] == "a+b&c" + assert connect.call_args.kwargs["additional_headers"]["x-gateway-token"] == "configured" + assert url.path.endswith("/rtc_owner") if model == "gpt-live-1-codex" else url.params["call_id"] == "rtc_owner" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 32b13ad678a..df1c53ec675 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -174,7 +174,9 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id @pytest.mark.parametrize("multipart", [False, True]) @pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom"]) @pytest.mark.parametrize("signaling_credential", ["authorization", "api-key", "x-litellm-api-key", "mixed"]) -async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential, signaling_credential): +async def test_offer_exchange_wraps_call_and_filters_client_headers( + monkeypatch, multipart, credential, signaling_credential +): import json from unittest.mock import AsyncMock @@ -188,11 +190,15 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") session = {"model": "voice-alias", "audio": {"output": {"voice": "sol"}}} if multipart: - body_request = httpx.Request("POST", "http://test/v1/realtime/calls", files={ - "sdp": (None, "v=0\r\n"), "session": (None, json.dumps(session)) - }) + body_request = httpx.Request( + "POST", + "http://test/v1/realtime/calls", + files={"sdp": (None, "v=0\r\n"), "session": (None, json.dumps(session))}, + ) else: - body_request = httpx.Request("POST", "http://test/v1/realtime/calls", json={"sdp": "v=0\r\n", "session": session}) + body_request = httpx.Request( + "POST", "http://test/v1/realtime/calls", json={"sdp": "v=0\r\n", "session": session} + ) body = body_request.read() async def receive(): @@ -203,16 +209,30 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, if signaling_credential == "mixed" else [(signaling_credential.encode(), b"Bearer owner" if signaling_credential == "authorization" else b"owner")] ) - request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", - "scheme": "http", "server": ("localhost", 80), - "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", - "headers": [(b"content-type", body_request.headers["content-type"].encode()), - *signaling_headers, *([(b"x-proxy-key", b"Bearer owner")] if credential == "custom" else []), (b"openai-alpha", b"quicksilver=v2"), - (b"x-untrusted", b"bad")]}, receive) + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "scheme": "http", + "server": ("localhost", 80), + "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", + "headers": [ + (b"content-type", body_request.headers["content-type"].encode()), + *signaling_headers, + *([(b"x-proxy-key", b"Bearer owner")] if credential == "custom" else []), + (b"openai-alpha", b"quicksilver=v2"), + (b"x-untrusted", b"bad"), + ], + }, + receive, + ) auth = UserAPIKeyAuth() authorize = AsyncMock() monkeypatch.setattr(proxy_server, "master_key", "owner") - monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential == "custom" else {}) + monkeypatch.setattr( + proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential == "custom" else {} + ) monkeypatch.setattr(codex, "can_key_call_resolved_model", authorize) class Processor: @@ -225,7 +245,16 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, assert self.data["model"] == "voice-alias" assert self.data["guardrails"] == ["query-guardrail"] assert await kwargs["request"].json() == {"model": "voice-alias"} - return {**self.data, "extra_headers": {"X-Hook-Required": "policy-value", "x-gateway-token": "untrusted-override", "Authorization": "Bearer untrusted"}, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None + return { + **self.data, + "extra_headers": { + "X-Hook-Required": "policy-value", + "x-gateway-token": "untrusted-override", + "Authorization": "Bearer untrusted", + }, + "extra_query": {"gateway_token": "untrusted-override"}, + "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}, + }, None return self.data, None monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor) @@ -236,14 +265,28 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, assert data["session"] == session assert data["chatgpt_realtime_client_headers"] == {"openai-alpha": "quicksilver=v2"} assert "extra_headers" not in data - assert data["extra_query"] == {"intent": "quicksilver", "architecture": "avas"} + assert data["chatgpt_realtime_client_query"] == {"intent": "quicksilver", "architecture": "avas"} async def respond(): - return httpx.Response(201, content=b"v=0\r\nanswer", headers={"Location": "/v1/realtime/calls/rtc_private"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", "extra_headers": {"X-Gateway-Token": "pinned-value"}}}) + return httpx.Response( + 201, + content=b"v=0\r\nanswer", + headers={"Location": "/v1/realtime/calls/rtc_private"}, + extensions={ + "chatgpt_realtime": { + "model": "gpt-live-1-codex", + "api_base": "https://voice.example/codex", + "extra_headers": {"X-Gateway-Token": "pinned-value"}, + "extra_query": {"gateway_token": "pinned-query-value"}, + } + }, + ) + return respond() monkeypatch.setattr(proxy_server, "route_request", route) + supervise = AsyncMock() + monkeypatch.setattr(codex, "supervise_codex_call", supervise) response = await codex.create_codex_realtime_call(request) assert response.status_code == 201 assert response.body == b"v=0\r\nanswer" @@ -252,7 +295,11 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, assert call.call_id == "rtc_private" assert call.alias == "voice-alias" assert call.model == "gpt-live-1-codex" + assert call.usage_supervised + supervise.assert_awaited_once() assert "rtc_private" not in token + assert "pinned-query-value" not in token + assert call.extra_query == {"gateway_token": "pinned-query-value"} assert time.time() < call.expires_at < time.time() + 3601 authorize.assert_awaited_once() @@ -271,16 +318,28 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, "custom": [(b"x-proxy-key", b"Bearer owner")], "subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")], } - websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque", - "query_string": b"guardrails=query-guardrail", "headers": credential_headers[credential]}, receive_ws, send) + websocket = WebSocket( + { + "type": "websocket", + "path": "/v1/live/opaque", + "query_string": b"guardrails=query-guardrail", + "headers": credential_headers[credential], + }, + receive_ws, + send, + ) forward = AsyncMock() monkeypatch.setattr(litellm, "_arealtime", forward) await codex.codex_realtime_sideband(websocket, token, auth) assert sent[0]["type"] == "websocket.accept" if credential == "subprotocol": assert sent[0]["subprotocol"] == "realtime" - assert forward.await_args.kwargs["extra_headers"] == {"x-hook-required": "policy-value", "x-gateway-token": "pinned-value"} + assert forward.await_args.kwargs["extra_headers"] == { + "x-hook-required": "policy-value", + "x-gateway-token": "pinned-value", + } assert forward.await_args.kwargs["metadata"] == {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"} + assert forward.await_args.kwargs["extra_query"] == {"gateway_token": "pinned-query-value"} assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-codex" assert forward.await_args.kwargs["api_base"] == "https://voice.example/codex" @@ -344,3 +403,158 @@ async def test_sideband_pre_call_block_prevents_upstream_connection(monkeypatch) await codex.codex_realtime_sideband(websocket, token, UserAPIKeyAuth()) forward.assert_not_called() assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Realtime pre-call rejected"}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("observer_fails", [False, True]) +async def test_signaling_transfers_reservation_only_to_ready_observer(monkeypatch, observer_fails): + import json + from unittest.mock import AsyncMock + + import httpx + from fastapi import Request + + from litellm.proxy import proxy_server + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-reservation-transfer") + reservation = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []} + auth = UserAPIKeyAuth(budget_reservation=reservation) + monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth)) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(proxy_server, "general_settings", {}) + process = AsyncMock(return_value=({}, None)) + monkeypatch.setattr(codex, "process_codex_request", process) + + async def response(): + return httpx.Response( + 201, + text="v=0\r\n", + headers={"Location": "/v1/realtime/calls/rtc_ready"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}, + ) + + async def route(**kwargs): + return response() + + monkeypatch.setattr(proxy_server, "route_request", route) + + async def supervise(request, call, owner): + assert owner is auth + assert not owner.budget_reservation["finalized"] + assert call.usage_supervised + if observer_fails: + await codex.release_or_invalidate_budget_reservation(budget_reservation=owner.budget_reservation) + raise RuntimeError("Observer unavailable") + + monkeypatch.setattr(codex, "supervise_codex_call", supervise) + + async def receive(): + return {"type": "http.request", "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode()} + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "query_string": b"", + "headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer owner")], + }, + receive, + ) + if observer_fails: + with pytest.raises(RuntimeError, match="Observer unavailable"): + await codex.create_codex_realtime_call(request) + else: + assert (await codex.create_codex_realtime_call(request)).status_code == 201 + assert process.await_args.args[2].budget_reservation is None + assert auth.budget_reservation["finalized"] is observer_fails + + +@pytest.mark.asyncio +async def test_supervisor_policy_failure_hangs_up_before_releasing(monkeypatch): + from unittest.mock import AsyncMock + + from fastapi import Request + + call = CodexRealtimeCall( + call_id="rtc_open", model="gpt-live-1-codex", alias="voice", owner="owner", expires_at=time.time() + 60 + ) + auth = UserAPIKeyAuth(budget_reservation={"reserved_cost": 0.5, "finalized": False, "entries": []}) + monkeypatch.setattr(codex, "process_codex_request", AsyncMock(side_effect=HTTPException(403, "Policy rejected"))) + closed = [] + + class Handler: + def __init__(self, *args): + pass + + @staticmethod + def get_api_base(base): + return "https://gateway.test/v1" + + async def hangup_call(self, base): + assert not auth.budget_reservation["finalized"] + closed.append(base) + + monkeypatch.setattr(codex, "ChatGPTRealtime", Handler) + request = Request({"type": "http", "headers": [], "method": "POST", "path": "/v1/realtime/calls"}) + with pytest.raises(HTTPException) as error: + await codex.supervise_codex_call(request, call, auth) + assert error.value.status_code == 403 + assert closed == ["https://gateway.test/v1"] + assert auth.budget_reservation["finalized"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("hangup_fails", [False, True]) +async def test_supervisor_constructor_failure_closes_effective_connection(monkeypatch, hangup_fails, caplog): + from unittest.mock import AsyncMock, MagicMock + + from fastapi import Request + + call = CodexRealtimeCall( + call_id="rtc_open", model="gpt-live-1-codex", alias="voice", owner="owner", expires_at=time.time() + 60 + ) + auth = UserAPIKeyAuth(budget_reservation={"reserved_cost": 0.5, "finalized": False, "entries": []}) + logger = MagicMock() + logger.litellm_params = {} + connection = AsyncMock() + handlers = [] + invalidate = AsyncMock() + release = AsyncMock() + monkeypatch.setattr(codex, "invalidate_budget_reservation_counters", invalidate, raising=False) + monkeypatch.setattr(codex, "release_or_invalidate_budget_reservation", release) + monkeypatch.setattr( + codex, "process_codex_request", AsyncMock(return_value=({"extra_headers": {"x-hook": "effective"}}, logger)) + ) + + class Handler: + def __init__(self, params, headers, extra_headers): + self.headers = extra_headers + handlers.append(self) + + @staticmethod + def get_api_base(base): + return "https://gateway.test/v1" + + async def open_call_connection(self, model, base): + return connection + + async def hangup_call(self, base): + assert self.headers["x-hook"] == "effective" + if hangup_fails: + raise RuntimeError("private-cleanup-credential") + + monkeypatch.setattr(codex, "ChatGPTRealtime", Handler) + monkeypatch.setattr(codex, "RealTimeStreaming", MagicMock(side_effect=ValueError("original constructor failure"))) + request = Request({"type": "http", "headers": [], "method": "POST", "path": "/v1/realtime/calls"}) + with pytest.raises(ValueError, match="original constructor failure"): + await codex.supervise_codex_call(request, call, auth) + connection.close.assert_awaited_once() + assert len(handlers) == 1 + if hangup_fails: + invalidate.assert_awaited_once_with(budget_reservation=auth.budget_reservation) + release.assert_not_awaited() + else: + release.assert_awaited_once_with(budget_reservation=auth.budget_reservation) + invalidate.assert_not_awaited() + assert "private-cleanup-credential" not in caplog.text 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..642bfd4982b --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -0,0 +1,301 @@ +import asyncio +import json +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.realtime_endpoints.call_supervision import CallSupervisor, CallSupervisors + + +class Socket: + def __init__(self): + self.messages = asyncio.Queue() + self.closed = False + + def __aiter__(self): + return self + + async def __anext__(self): + message = await self.messages.get() + if message is None: + raise StopAsyncIteration + if isinstance(message, Exception): + raise message + return json.dumps(message) + + async def close(self): + self.closed = True + + +class Sink: + def __init__(self, logger): + self.logger = logger + self.events = [] + self.logs = 0 + + def store_message(self, message): + self.events.append(json.loads(message)) + + async def log_messages(self, *, wait_for_dispatch=False): + assert wait_for_dispatch + self.logs += 1 + self.logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True + + +def fixture(*, ready_timeout=1, lifetime=1): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + + async def hangup(): + assert not socket.closed + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + + close_call = AsyncMock(side_effect=hangup) + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close_call, + ready_timeout=ready_timeout, + lifetime=lifetime, + drain_timeout=0.05, + ) + return socket, sink, close_call, supervisor + + +@pytest.mark.asyncio +async def test_observer_logs_webrtc_usage_without_client_sideband(): + socket, sink, close_call, supervisor = fixture() + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await socket.messages.put({"type": "response.done", "response": {"usage": {"total_tokens": 15}}}) + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 19}}) + await supervisor.wait() + await supervisor.close() + assert sink.logs == 1 + assert sink.events[-1]["usage"]["total_tokens"] == 19 + assert sink.events[1]["response"]["usage"]["total_tokens"] == 15 + assert socket.closed + close_call.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_early_upstream_eof_rejects_start(): + socket, sink, close_call, supervisor = fixture() + await socket.messages.put(None) + with pytest.raises(RuntimeError, match="ended before"): + await supervisor.start() + assert socket.closed + assert sink.logs == 1 + close_call.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_cancelled_start_hangs_up_and_drains_terminal_usage(): + socket, sink, close_call, supervisor = fixture() + started = asyncio.create_task(supervisor.start()) + await asyncio.sleep(0) + started.cancel() + with pytest.raises(asyncio.CancelledError): + await started + close_call.assert_awaited_once() + assert socket.closed + assert sink.logs == 1 + assert sink.events[-1]["usage"]["total_tokens"] == 42 + + +@pytest.mark.asyncio +async def test_worker_shutdown_drains_all_calls(): + registry = CallSupervisors() + socket, sink, close_call, supervisor = fixture() + await socket.messages.put({"type": "session.created"}) + await registry.start(supervisor) + await registry.shutdown() + await registry.shutdown() + close_call.assert_awaited_once() + assert socket.closed + assert sink.logs == 1 + assert sink.events[-1]["usage"]["total_tokens"] == 42 + + +@pytest.mark.asyncio +async def test_ready_timeout_hangs_up_before_returning_error(): + socket, sink, close_call, supervisor = fixture(ready_timeout=0.01) + with pytest.raises(asyncio.TimeoutError): + await supervisor.start() + close_call.assert_awaited_once() + assert socket.closed + assert sink.logs == 1 + + +@pytest.mark.asyncio +async def test_lifetime_limit_closes_call_and_collects_final_usage(): + socket, sink, close_call, supervisor = fixture(lifetime=0.01) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await supervisor.wait() + close_call.assert_awaited_once() + assert socket.closed + assert sink.events[-1]["usage"]["total_tokens"] == 42 + + +@pytest.mark.asyncio +async def test_socket_eof_after_ready_still_hangs_up_provider_call(monkeypatch): + from litellm.proxy.realtime_endpoints import call_supervision + + invalidate = AsyncMock() + release = AsyncMock() + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) + monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release) + socket, sink, close_call, supervisor = fixture() + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await socket.messages.put(None) + await supervisor.wait() + close_call.assert_awaited_once() + assert socket.closed + assert sink.logs == 1 + assert sink.logger.model_call_details["realtime_usage_incomplete"] is True + invalidate.assert_awaited_once() + release.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_observer_error_rejects_start(caplog): + socket, sink, close_call, supervisor = fixture() + await socket.messages.put(RuntimeError("private-provider-credential")) + with pytest.raises(RuntimeError, match="ended before"): + await supervisor.start() + assert socket.closed + close_call.assert_awaited_once() + assert "private-provider-credential" not in caplog.text + + +@pytest.mark.asyncio +async def test_failed_logging_releases_reservation(monkeypatch): + from litellm.proxy.realtime_endpoints import call_supervision + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = MagicMock() + sink.log_messages = AsyncMock(side_effect=RuntimeError("logging unavailable")) + release = AsyncMock() + monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release) + supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock()) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await socket.messages.put({"type": "session.closed"}) + with pytest.raises(RuntimeError, match="logging unavailable"): + await supervisor.wait() + release.assert_awaited_once_with(budget_reservation=None) + assert socket.closed + sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True) + + +@pytest.mark.asyncio +async def test_shutdown_waits_for_usage_dispatch_completion(): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + dispatch_started = asyncio.Event() + dispatch_complete = asyncio.Event() + dispatch_finished = asyncio.Event() + + async def log_messages(*, wait_for_dispatch=False): + assert wait_for_dispatch + dispatch_started.set() + await dispatch_complete.wait() + logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True + dispatch_finished.set() + + async def hangup(): + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + + sink = MagicMock() + sink.log_messages = AsyncMock(side_effect=log_messages) + registry = CallSupervisors() + supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), hangup) + await socket.messages.put({"type": "session.created"}) + await registry.start(supervisor) + shutdown = asyncio.create_task(registry.shutdown()) + try: + await asyncio.wait_for(dispatch_started.wait(), timeout=1) + assert not shutdown.done() + assert not dispatch_finished.is_set() + finally: + dispatch_complete.set() + await asyncio.wait_for(shutdown, timeout=1) + assert dispatch_finished.is_set() + sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("terminal_usage_required", [True, False]) +async def test_confirmed_hangup_without_terminal_usage_matches_protocol(terminal_usage_required): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + close_call = AsyncMock() + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close_call, + terminal_usage_required=terminal_usage_required, + drain_timeout=0.01, + ) + await socket.messages.put({"type": "session.created"}) + await supervisor.start() + await socket.messages.put({"type": "response.done", "response": {"usage": {"total_tokens": 17}}}) + await supervisor.close() + assert bool(logger.model_call_details.get("realtime_usage_incomplete")) == terminal_usage_required + assert sink.events[-1]["response"]["usage"]["total_tokens"] == 17 + assert sink.logs == 1 + assert socket.closed + close_call.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("closure", ["eof", "normal_close", "error"]) +@pytest.mark.parametrize("hangup_succeeds", [True, False]) +async def test_ga_observer_disconnect_requires_confirmed_hangup(closure, hangup_succeeds): + from websockets.exceptions import ConnectionClosedOK + from websockets.frames import Close + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + close_call = AsyncMock(side_effect=None if hangup_succeeds else RuntimeError("unconfirmed hangup")) + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close_call, + terminal_usage_required=False, + drain_timeout=0.01, + ) + await socket.messages.put({"type": "session.created"}) + await supervisor.start() + await socket.messages.put( + None + if closure == "eof" + else ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True) + if closure == "normal_close" + else RuntimeError("observer failed") + ) + await supervisor.wait() + close_call.assert_awaited_once() + assert bool(logger.model_call_details.get("realtime_usage_incomplete")) == (not hangup_succeeds) + assert sink.logs == 1 + assert socket.closed diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 8f8a7640c08..692939d3ed0 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -4751,3 +4751,71 @@ def test_collect_and_combine_realtime_usage_stores_partitioned_text_tokens() -> assert combined.completion_tokens_details.reasoning_tokens == 95 assert combined.completion_tokens_details.text_tokens == 38 assert combined.completion_tokens_details.audio_tokens == 0 + + +def _live_terminal_event(duration=4000): + return {"type": "session.closed", "usage": {"audio_duration_ms": duration, "backend_model_usage": []}} + + +@pytest.mark.parametrize("rate,expected", [(0.025, 0.1), (0, 0), (None, 0)]) +def test_live_terminal_duration_uses_configured_second_price(monkeypatch, rate, expected): + monkeypatch.setitem( + litellm.model_cost, + "live-priced-test", + {"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": rate}, + ) + assert handle_realtime_stream_cost_calculation( + [_live_terminal_event()], Usage(), "chatgpt", "live-priced-test" + ) == pytest.approx(expected) + + +def test_live_terminal_duration_honors_deployment_override(monkeypatch): + monkeypatch.setitem( + litellm.model_cost, + "live-deployment-test", + {"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025}, + ) + result = RealtimeAPITokenUsageProcessor.create_logging_realtime_object(Usage(), [_live_terminal_event()]) + assert completion_cost( + completion_response=result, + model="gpt-live-1", + custom_llm_provider="chatgpt", + call_type="_arealtime", + custom_pricing=True, + router_model_id="live-deployment-test", + ) == pytest.approx(0.1) + + +@pytest.mark.parametrize("duration", [-1, True, "4000", float("inf"), float("nan"), None]) +def test_live_terminal_invalid_duration_does_not_create_spend(monkeypatch, duration): + monkeypatch.setitem( + litellm.model_cost, + "live-priced-test", + {"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025}, + ) + assert ( + handle_realtime_stream_cost_calculation( + [_live_terminal_event(duration)], Usage(), "chatgpt", "live-priced-test" + ) + == 0 + ) + + +def test_live_terminal_is_not_counted_twice(monkeypatch): + monkeypatch.setitem( + litellm.model_cost, + "live-priced-test", + {"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025}, + ) + assert handle_realtime_stream_cost_calculation( + [_live_terminal_event(), _live_terminal_event()], Usage(), "chatgpt", "live-priced-test" + ) == pytest.approx(0.1) + assert ( + handle_realtime_stream_cost_calculation( + [{"type": "response.done", "response": {"usage": {}}}, _live_terminal_event()], + Usage(), + "chatgpt", + "live-priced-test", + ) + == 0 + ) From bac2ad4abb7fab4926babae2641d9fb1cc5cfcaa Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 13:17:05 +0200 Subject: [PATCH 28/90] fix(chatgpt): finish call termination and type pinned sideband routing --- .../proxy/realtime_endpoints/call_sessions.py | 40 ++++++----- .../realtime_endpoints/call_supervision.py | 4 +- .../test_call_supervision.py | 72 +++++++++++++++++++ 3 files changed, 98 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index b0833b19414..b85df30bb31 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -2,12 +2,14 @@ import base64 import hashlib import json import time +from collections.abc import Mapping from contextlib import AsyncExitStack from types import MappingProxyType from typing import Final, Literal import httpx from fastapi import HTTPException, Request, Response, WebSocket +from pydantic import TypeAdapter from starlette.types import Message from litellm._logging import verbose_proxy_logger @@ -44,10 +46,6 @@ from litellm.types.router import GenericLiteLLMParams async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth) -> None: - from collections.abc import Mapping - - from pydantic import TypeAdapter - import litellm from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, CallSupervisor @@ -362,19 +360,27 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP subprotocol=next((p for p in protocols if not p.startswith("openai-insecure-api-key.")), None) ) await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call - **{ # mutable-ok: retain processed policy metadata while pinning the existing call's routing - **processed, - **build_sideband_request(call), - "extra_headers": MappingProxyType( - { - **configured_realtime_headers(processed.get("extra_headers")), - **configured_realtime_headers(call.extra_headers), - } - ), - "websocket": websocket, - "user_api_key_dict": auth, - "chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None, - } + model=f"chatgpt/{call.model}", + websocket=websocket, + **{ + key: value + for key, value in { # mutable-ok: retain processed metadata with pinned routing + **processed, + **build_sideband_request(call), + "extra_headers": MappingProxyType( + { + **configured_realtime_headers( + TypeAdapter(Mapping[str, object] | None).validate_python(processed.get("extra_headers")) + ), + **configured_realtime_headers(call.extra_headers), + } + ), + "websocket": websocket, + "user_api_key_dict": auth, + "chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None, + }.items() + if key not in ("model", "websocket") + }, ) finally: if logging_obj is None or not logging_obj.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY): diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index 98f018353ba..b574a229bbe 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -44,6 +44,7 @@ class CallSupervisor: ready_timeout: float = 20, lifetime: float = 3600, drain_timeout: float = 5, + termination_timeout: float = 60, terminal_usage_required: bool = True, ) -> None: self._upstream = upstream @@ -54,6 +55,7 @@ class CallSupervisor: self._ready_timeout = ready_timeout self._lifetime = lifetime self._drain_timeout = drain_timeout + self._termination_timeout = termination_timeout self._terminal_usage_required = terminal_usage_required self._ready = asyncio.Event() self._stop = asyncio.Event() @@ -112,7 +114,7 @@ class CallSupervisor: try: if not self._terminal: try: - await asyncio.wait_for(self._close_call(), timeout=self._drain_timeout) + await asyncio.wait_for(self._close_call(), timeout=self._termination_timeout) self._close_confirmed = True except Exception: # noqa: BLE001 # provider exceptions can contain credentials verbose_proxy_logger.error("Realtime observer could not terminate upstream call") diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 642bfd4982b..264eccc787e 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -264,6 +264,78 @@ async def test_confirmed_hangup_without_terminal_usage_matches_protocol(terminal close_call.assert_awaited_once() +@pytest.mark.asyncio +async def test_shutdown_allows_hangup_longer_than_usage_drain_timeout(): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + hangup_started = asyncio.Event() + allow_hangup = asyncio.Event() + hangup_finished = asyncio.Event() + + async def hangup(): + hangup_started.set() + await allow_hangup.wait() + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + hangup_finished.set() + + supervisor = CallSupervisor( + socket, sink, logger, UserAPIKeyAuth(), hangup, drain_timeout=0.01, termination_timeout=1 + ) + registry = CallSupervisors() + await socket.messages.put({"type": "session.created"}) + await registry.start(supervisor) + shutdown = asyncio.create_task(registry.shutdown()) + try: + await asyncio.wait_for(hangup_started.wait(), timeout=1) + await asyncio.sleep(0.04) + assert not shutdown.done() + assert not socket.closed + assert not hangup_finished.is_set() + finally: + allow_hangup.set() + await asyncio.wait_for(shutdown, timeout=1) + assert hangup_finished.is_set() + assert socket.closed + assert sink.logs == 1 + assert sink.events[-1]["usage"]["total_tokens"] == 42 + assert not logger.model_call_details.get("realtime_usage_incomplete") + + +@pytest.mark.asyncio +async def test_termination_timeout_cancels_hangup_and_finishes_cleanup(): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + hangup_cancelled = asyncio.Event() + + async def hangup(): + try: + await asyncio.Event().wait() + finally: + hangup_cancelled.set() + + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + hangup, + drain_timeout=0.01, + termination_timeout=0.02, + terminal_usage_required=False, + ) + await socket.messages.put({"type": "session.created"}) + await supervisor.start() + await asyncio.wait_for(supervisor.close(), timeout=1) + assert hangup_cancelled.is_set() + assert socket.closed + assert sink.logs == 1 + assert logger.model_call_details["realtime_usage_incomplete"] is True + + @pytest.mark.asyncio @pytest.mark.parametrize("closure", ["eof", "normal_close", "error"]) @pytest.mark.parametrize("hangup_succeeds", [True, False]) From c360d187f5a78d43a3b8a9bd5771ddd4abe26d17 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 15:54:39 +0200 Subject: [PATCH 29/90] fix(chatgpt): preserve observer quotas and bound call cleanup --- litellm/llms/chatgpt/realtime.py | 18 ++- litellm/proxy/_types.py | 4 + litellm/proxy/common_request_processing.py | 7 + .../proxy/hooks/parallel_request_limiter.py | 39 ++--- .../proxy/realtime_endpoints/call_sessions.py | 16 ++- .../realtime_endpoints/call_supervision.py | 37 ++++- litellm/proxy/utils.py | 10 ++ .../llms/chatgpt/test_realtime.py | 92 +++++++++++- .../realtime_endpoints/test_call_sessions.py | 15 +- .../test_call_supervision.py | 136 +++++++++++++++++- tests/test_litellm/proxy/test_proxy_utils.py | 136 +++++++++++++++++- 11 files changed, 464 insertions(+), 46 deletions(-) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index f79713fe03f..27ce6cb3966 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -3,8 +3,9 @@ from enum import Enum, auto from types import MappingProxyType from typing import TYPE_CHECKING, Final -from httpx import URL +from httpx import URL, QueryParams from pydantic import TypeAdapter +from websockets.exceptions import ConnectionClosed from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.llms.openai.realtime.handler import OpenAIRealtime @@ -40,11 +41,14 @@ def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str] inbound: Final = TypeAdapter(Mapping[str, str]).validate_python( getattr(params, "chatgpt_realtime_client_query", None) or MappingProxyType({}) ) - configured: Final = TypeAdapter(Mapping[str, str]).validate_python( + configured: Final = TypeAdapter(Mapping[str, str | int | float | bool | None]).validate_python( getattr(params, "extra_query", None) or MappingProxyType({}) ) return MappingProxyType( - {**{key: value for key, value in inbound.items() if key in ("intent", "architecture")}, **configured} + { + **{key: value for key, value in inbound.items() if key in ("intent", "architecture")}, + **QueryParams(configured), + } ) @@ -111,8 +115,12 @@ class ChatGPTRealtime(OpenAIRealtime): async def close_call(self, connection: "ClientConnection", model: str, api_base: str) -> None: if realtime_endpoint(model) == "live": - await connection.send('{"type":"session.close"}') - return + try: + await connection.send('{"type":"session.close"}') + return + except (ConnectionClosed, OSError): + await self.hangup_call(api_base) + return await self.hangup_call(api_base) async def hangup_call(self, api_base: str) -> None: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 5a3633afa3e..78bfa2bf668 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -97,6 +97,10 @@ class ReconcileOutcome(NamedTuple): live_after: frozenset[str] | None +class InternalRequestOrigin(enum.Enum): + REALTIME_OBSERVER = enum.auto() + + class SupportedDBObjectType(str, enum.Enum): """ Supported database object types for fine-grained DB storage control. diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e3a2b892721..476f80b4158 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1832,6 +1832,8 @@ class ProxyBaseLLMRequestProcessing: user_api_base: str | None = None, model: str | None = None, llm_router: Router | None = None, + *, + internal_realtime_observer: bool = False, ) -> tuple[dict, LiteLLMLoggingObj]: start_time: Final = datetime.now() # start before calling guardrail hooks @@ -1996,6 +1998,11 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type, + **( + MappingProxyType({"internal_realtime_observer": True}) + if internal_realtime_observer + else MappingProxyType({}) + ), ) if route_type == "aget_responses": attach_post_call_pipelines_to_retrieval( diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index b313cb64c3f..2ec576b97b3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -12,7 +12,7 @@ from litellm._logging import verbose_proxy_logger from litellm.exceptions import RateLimitType from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs -from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth +from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, InternalRequestOrigin, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( get_key_model_rpm_limit, get_key_model_tpm_limit, @@ -489,6 +489,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time): + releases_slot: Final = kwargs.get("internal_request_origin") is not InternalRequestOrigin.REALTIME_OBSERVER from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, ) @@ -521,7 +522,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): # Setup values # ------------ - if global_max_parallel_requests is not None: + if releases_slot and global_max_parallel_requests is not None: # get value from cache _key: Final = "global_max_parallel_requests" # decrement @@ -552,13 +553,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): key=request_count_api_key, litellm_parent_otel_span=litellm_parent_otel_span, ) or { - "current_requests": 1, + "current_requests": int(releases_slot), "current_tpm": 0, "current_rpm": 0, } new_val = { - "current_requests": max(current["current_requests"] - 1, 0), + "current_requests": max(current["current_requests"] - int(releases_slot), 0), "current_tpm": current["current_tpm"] + total_tokens, "current_rpm": current["current_rpm"], } @@ -593,13 +594,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): key=request_count_api_key, litellm_parent_otel_span=litellm_parent_otel_span, ) or { - "current_requests": 1, + "current_requests": int(releases_slot), "current_tpm": 0, "current_rpm": 0, } new_val = { - "current_requests": max(current["current_requests"] - 1, 0), + "current_requests": max(current["current_requests"] - int(releases_slot), 0), "current_tpm": current["current_tpm"] + total_tokens, "current_rpm": current["current_rpm"], } @@ -619,13 +620,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): key=request_count_api_key, litellm_parent_otel_span=litellm_parent_otel_span, ) or { - "current_requests": 1, - "current_tpm": total_tokens, - "current_rpm": 1, + "current_requests": int(releases_slot), + "current_tpm": total_tokens if releases_slot else 0, + "current_rpm": int(releases_slot), } new_val = { - "current_requests": max(current["current_requests"] - 1, 0), + "current_requests": max(current["current_requests"] - int(releases_slot), 0), "current_tpm": current["current_tpm"] + total_tokens, "current_rpm": current["current_rpm"], } @@ -645,13 +646,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): key=request_count_api_key, litellm_parent_otel_span=litellm_parent_otel_span, ) or { - "current_requests": 1, - "current_tpm": total_tokens, - "current_rpm": 1, + "current_requests": int(releases_slot), + "current_tpm": total_tokens if releases_slot else 0, + "current_rpm": int(releases_slot), } new_val = { - "current_requests": max(current["current_requests"] - 1, 0), + "current_requests": max(current["current_requests"] - int(releases_slot), 0), "current_tpm": current["current_tpm"] + total_tokens, "current_rpm": current["current_rpm"], } @@ -671,13 +672,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): key=request_count_api_key, litellm_parent_otel_span=litellm_parent_otel_span, ) or { - "current_requests": 1, - "current_tpm": total_tokens, - "current_rpm": 1, + "current_requests": int(releases_slot), + "current_tpm": total_tokens if releases_slot else 0, + "current_rpm": int(releases_slot), } new_val = { - "current_requests": max(current["current_requests"] - 1, 0), + "current_requests": max(current["current_requests"] - int(releases_slot), 0), "current_tpm": current["current_tpm"] + total_tokens, "current_rpm": current["current_rpm"], } @@ -694,6 +695,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): self.print_verbose(e) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + if kwargs.get("internal_request_origin") is InternalRequestOrigin.REALTIME_OBSERVER: + return try: self.print_verbose("Inside Max Parallel Request Failure Hook") litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index b85df30bb31..8af489e8eb1 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -29,7 +29,7 @@ from litellm.llms.chatgpt.realtime import ( configured_realtime_headers, realtime_endpoint, ) -from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy._types import InternalRequestOrigin, ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model from litellm.proxy.auth.user_api_key_auth import ( get_api_key, @@ -73,6 +73,7 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: auth, call.alias, "_arealtime", + internal_realtime_observer=True, ) pinned: Final = { # mutable-ok: logging and provider parameter contract **processed, @@ -116,6 +117,9 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: async def close_call() -> None: await handler.close_call(connection, call.model, api_base) + async def force_close_call() -> None: + await handler.hangup_call(api_base) + frontend: Final = WebSocket( {**request.scope, "type": "websocket"}, receive=receive, send=send ) # mutable-ok: ASGI scope @@ -126,6 +130,7 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: logger, auth, close_call, + force_close_call=force_close_call, terminal_usage_required=realtime_endpoint(call.model) == "live", ) supervision_owned = True @@ -191,6 +196,8 @@ async def process_codex_request( auth: UserAPIKeyAuth, model: str, route_type: Literal["arealtime_calls", "_arealtime"], + *, + internal_realtime_observer: bool = False, ) -> tuple[dict[str, object], Logging]: # mutable-ok: common request processor returns enriched routing arguments from litellm.proxy import proxy_server as server from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing @@ -211,7 +218,14 @@ async def process_codex_request( user_api_base=server.user_api_base, model=model, route_type=route_type, + **( + MappingProxyType({"internal_realtime_observer": True}) + if internal_realtime_observer + else MappingProxyType({}) + ), ) + if internal_realtime_observer: + logging_obj.model_call_details["internal_request_origin"] = InternalRequestOrigin.REALTIME_OBSERVER return processed, logging_obj diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index b574a229bbe..e750ce61d6a 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -7,6 +7,7 @@ from pydantic import BaseModel from websockets.exceptions import ConnectionClosedOK from litellm._logging import verbose_proxy_logger +from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY from litellm.proxy._types import UserAPIKeyAuth @@ -45,23 +46,28 @@ class CallSupervisor: lifetime: float = 3600, drain_timeout: float = 5, termination_timeout: float = 60, + logging_timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE, terminal_usage_required: bool = True, + force_close_call: Callable[[], Awaitable[None]] | None = None, ) -> None: self._upstream = upstream self._stream = stream self._logging = logging_obj self._auth = auth self._close_call = close_call + self._force_close_call = force_close_call self._ready_timeout = ready_timeout self._lifetime = lifetime self._drain_timeout = drain_timeout self._termination_timeout = termination_timeout + self._logging_timeout = logging_timeout self._terminal_usage_required = terminal_usage_required self._ready = asyncio.Event() self._stop = asyncio.Event() self._started = False self._terminal = False self._close_confirmed = False + self._accounting_complete = False self._task: asyncio.Task[None] | None = None async def start(self) -> None: @@ -70,7 +76,7 @@ class CallSupervisor: self._task = asyncio.create_task(self._run()) try: await asyncio.wait_for(self._ready.wait(), timeout=self._ready_timeout) - if not self._started or self._task.done(): + if not self._started or self._terminal or self._task.done(): raise RuntimeError("Call observer ended before session became available") except BaseException: await self.close() @@ -113,12 +119,21 @@ class CallSupervisor: finally: try: if not self._terminal: + deadline: Final = asyncio.get_running_loop().time() + self._termination_timeout try: await asyncio.wait_for(self._close_call(), timeout=self._termination_timeout) self._close_confirmed = True except Exception: # noqa: BLE001 # provider exceptions can contain credentials verbose_proxy_logger.error("Realtime observer could not terminate upstream call") - await self._drain(reader) + await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time())) + if self._terminal_usage_required and not self._terminal and self._force_close_call is not None: + remaining: Final = max(0.0, deadline - asyncio.get_running_loop().time()) + try: + await asyncio.wait_for(self._force_close_call(), timeout=remaining) + self._close_confirmed = True + except Exception: # noqa: BLE001 # provider exceptions can contain credentials + verbose_proxy_logger.error("Realtime observer independent hangup failed") + await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time())) finally: stopped.cancel() reader.cancel() @@ -132,9 +147,16 @@ class CallSupervisor: ) try: try: - await self._stream.log_messages(wait_for_dispatch=True) + await asyncio.wait_for( + self._stream.log_messages(wait_for_dispatch=True), timeout=self._logging_timeout + ) + self._accounting_complete = True + except asyncio.TimeoutError: + verbose_proxy_logger.error("Realtime observer timed out dispatching usage accounting") finally: - if self._started and not self._usage_complete(): + if not self._accounting_complete: + self._logging.model_call_details["realtime_accounting_incomplete"] = True + if self._started and (not self._usage_complete() or not self._accounting_complete): await invalidate_budget_reservation_counters( budget_reservation=self._auth.budget_reservation ) @@ -145,9 +167,12 @@ class CallSupervisor: finally: self._ready.set() - async def _drain(self, reader: asyncio.Task[None]) -> None: + async def _drain(self, reader: asyncio.Task[None], *, timeout: float | None = None) -> None: try: - await asyncio.wait_for(asyncio.shield(reader), timeout=self._drain_timeout) + await asyncio.wait_for( + asyncio.shield(reader), + timeout=self._drain_timeout if timeout is None else min(self._drain_timeout, timeout), + ) except asyncio.TimeoutError: if not self._usage_complete(): verbose_proxy_logger.error("Realtime observer timed out draining terminal usage") diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 00ccad33b6d..b43835fd7d9 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2046,6 +2046,8 @@ class ProxyLogging: data: None, call_type: CallTypesLiteral, guardrails_only: bool = False, + *, + internal_realtime_observer: bool = False, ) -> None: pass @@ -2056,6 +2058,8 @@ class ProxyLogging: data: dict, call_type: CallTypesLiteral, guardrails_only: bool = False, + *, + internal_realtime_observer: bool = False, ) -> dict: pass @@ -2065,6 +2069,8 @@ class ProxyLogging: data: dict | None, call_type: CallTypesLiteral, guardrails_only: bool = False, + *, + internal_realtime_observer: bool = False, ) -> dict | None: """ Allows users to modify/reject the incoming request to the proxy, without having to deal with parsing Request body. @@ -2163,6 +2169,10 @@ class ProxyLogging: deferred_route_exc: SensitiveDataRouteException | None = None for _callback in caps.resolved_callbacks: + if internal_realtime_observer and isinstance( + _callback, (_PROXY_MaxParallelRequestsHandler, _PROXY_MaxParallelRequestsHandler_v3) + ): + continue start_time = time.time() try: if isinstance(_callback, CustomGuardrail) and data is not None: diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 8b29d440b4c..6db166e956b 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -1,4 +1,5 @@ import json +from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import AsyncMock, patch @@ -11,6 +12,51 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.router import GenericLiteLLMParams +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["closed", "network"]) +@pytest.mark.parametrize( + "hangup_status, expectation", [(200, nullcontext()), (503, pytest.raises(httpx.HTTPStatusError))] +) +async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, expectation, chatgpt_tokens): + from websockets.exceptions import ConnectionClosedOK + from websockets.frames import Close + + handler = ChatGPTRealtime( + GenericLiteLLMParams( + chatgpt_realtime_call_id="rtc_live_closed", + chatgpt_token_dir=chatgpt_tokens, + extra_query={"gateway": "tenant"}, + ), + {}, + {"x-gateway-token": "test-only"}, + ) + connection = SimpleNamespace( + send=AsyncMock( + side_effect=( + ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True) + if failure == "closed" + else OSError("socket unavailable") + ) + ) + ) + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(hangup_status) + + client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + with patch("httpx.AsyncClient", return_value=client): + with expectation: + await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1") + assert len(requests) == 1 + assert requests[0].method == "POST" + assert str(requests[0].url) == "https://gateway.example/v1/realtime/calls/rtc_live_closed/hangup?gateway=tenant" + assert requests[0].headers["x-gateway-token"] == "test-only" + assert requests[0].headers["Authorization"] == "Bearer test-token-default" + assert client.is_closed + + @pytest.mark.asyncio @pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"]) @pytest.mark.parametrize("source", ["default", "explicit", "CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"]) @@ -47,8 +93,17 @@ async def test_realtime_session_urls_honor_gateway(endpoint, source, chatgpt_tok @pytest.mark.asyncio @pytest.mark.parametrize("inbound_headers", [{}, {"openai-alpha": "quicksilver=v2"}]) -async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, chatgpt_tokens, monkeypatch): - from litellm.llms.chatgpt.codex import CodexRealtimeOffer, build_call_request +@pytest.mark.parametrize("model, endpoint", [("gpt-live-1-codex", "live"), ("gpt-realtime-1.5", "realtime")]) +async def test_routed_call_preserves_deployment_gateway_headers( + inbound_headers, model, endpoint, chatgpt_tokens, monkeypatch +): + from litellm.llms.chatgpt.codex import ( + CodexRealtimeCall, + CodexRealtimeOffer, + build_call_request, + build_sideband_request, + parse_call_response, + ) monkeypatch.setenv("CHATGPT_TOKEN_DIR", chatgpt_tokens) requests = [] @@ -64,11 +119,22 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, { "model_name": "voice-gateway", "litellm_params": { - "model": "chatgpt/gpt-live-1-codex", + "model": f"chatgpt/{model}", "api_base": "https://voice.example/backend-api/codex", "extra_headers": {"x-gateway-route": "configured"}, - "extra_query": {"gateway_token": "configured", "intent": "pinned-intent"}, + "extra_query": { + "gateway_token": "configured", + "intent": "pinned-intent", + "count": 7, + "fraction": 1.5, + "enabled": True, + "disabled": False, + "blank": None, + "model": "other-model", + "call_id": "rtc_wrong", + }, }, + "model_info": {"id": "selected-gateway-deployment"}, } ], num_retries=0, @@ -84,11 +150,29 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, "gateway_token": "configured", "intent": "pinned-intent", "architecture": "avas", + "count": "7", + "fraction": "1.5", + "enabled": "true", + "disabled": "false", + "blank": "", + "model": "other-model", + "call_id": "rtc_wrong", } assert response.extensions["chatgpt_realtime"]["extra_query"] == dict(requests[0].url.params) assert response.extensions["chatgpt_realtime"]["extra_headers"]["x-gateway-route"] == "configured" for name, value in inbound_headers.items(): assert requests[0].headers[name] == value + call = parse_call_response(response, alias="voice-gateway", owner="test-owner", expires_at=1) + restored = CodexRealtimeCall.model_validate_json(call.model_dump_json()) + assert restored.model_id == "selected-gateway-deployment" + assert restored.model == model + handler = ChatGPTRealtime(GenericLiteLLMParams.model_validate(build_sideband_request(restored)), {}) + sideband_url = httpx.URL(handler._construct_url(restored.api_base, {"model": restored.model})) + assert {key: value for key, value in sideband_url.params.items() if key != "call_id"} == { + key: value for key, value in requests[0].url.params.items() if key not in ("model", "call_id") + } + assert sideband_url.params.get("call_id") == ("rtc_test" if endpoint == "realtime" else None) + assert sideband_url.path.endswith("/realtime" if endpoint == "realtime" else "/live/rtc_test") finally: await client.client.aclose() diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index df1c53ec675..ef85e324680 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -13,7 +13,8 @@ from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_c @pytest.mark.asyncio @pytest.mark.parametrize("route_type", ["arealtime_calls", "_arealtime"]) -async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type): +@pytest.mark.parametrize("observer", [False, True]) +async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type, observer): from fastapi import Request from litellm import Router from litellm.proxy import proxy_server as server @@ -21,7 +22,8 @@ async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type) from litellm.proxy.realtime_endpoints.call_sessions import process_codex_request class PolicyHook: - async def pre_call_hook(self, user_api_key_dict, data, call_type): + async def pre_call_hook(self, user_api_key_dict, data, call_type, *, internal_realtime_observer=False): + assert internal_realtime_observer is observer if "model-policy" in data.get("metadata", {}).get("guardrails", []): raise HTTPException(403, "Model policy rejected request") return data @@ -34,7 +36,14 @@ async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type) monkeypatch.setattr(server, "proxy_logging_obj", PolicyHook()) request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", "headers": [], "query_string": b"", "scheme": "http", "server": ("localhost", 80)}) with pytest.raises(HTTPException) as error: - await process_codex_request(request, {"model": "voice-policy"}, UserAPIKeyAuth(), "voice-policy", route_type) + await process_codex_request( + request, + {"model": "voice-policy"}, + UserAPIKeyAuth(), + "voice-policy", + route_type, + internal_realtime_observer=observer, + ) assert error.value.status_code == 403 assert error.value.detail == "Model policy rejected request" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 264eccc787e..7d0d614f31d 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -30,6 +30,67 @@ class Socket: self.closed = True +@pytest.mark.asyncio +@pytest.mark.parametrize("fallback", ["terminal", "no_terminal", "timeout"]) +async def test_live_unacknowledged_close_uses_bounded_independent_hangup(monkeypatch, fallback): + from litellm.proxy.realtime_endpoints import call_supervision + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + invalidate = AsyncMock() + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) + + async def force_close(): + if fallback == "terminal": + await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}}) + elif fallback == "timeout": + await asyncio.Event().wait() + + force = AsyncMock(side_effect=force_close) + close = AsyncMock() + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close, + force_close_call=force, + drain_timeout=0.01, + termination_timeout=0.08, + ) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await asyncio.wait_for(supervisor.close(), timeout=0.5) + close.assert_awaited_once() + force.assert_awaited_once() + assert socket.closed + if fallback == "terminal": + invalidate.assert_not_awaited() + assert not logger.model_call_details.get("realtime_usage_incomplete") + else: + invalidate.assert_awaited_once() + assert logger.model_call_details["realtime_usage_incomplete"] is True + + +@pytest.mark.asyncio +async def test_live_confirmed_terminal_does_not_force_hangup(): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + + async def close(): + await socket.messages.put({"type": "session.closed"}) + + force = AsyncMock() + supervisor = CallSupervisor(socket, Sink(logger), logger, UserAPIKeyAuth(), close, force_close_call=force) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await supervisor.close() + force.assert_not_awaited() + + class Sink: def __init__(self, logger): self.logger = logger @@ -178,7 +239,7 @@ async def test_observer_error_rejects_start(caplog): @pytest.mark.asyncio -async def test_failed_logging_releases_reservation(monkeypatch): +async def test_failed_logging_invalidates_reservation_without_zeroing_spend(monkeypatch): from litellm.proxy.realtime_endpoints import call_supervision socket = Socket() @@ -187,18 +248,89 @@ async def test_failed_logging_releases_reservation(monkeypatch): sink = MagicMock() sink.log_messages = AsyncMock(side_effect=RuntimeError("logging unavailable")) release = AsyncMock() + invalidate = AsyncMock() monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release) + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock()) await socket.messages.put({"type": "session.started"}) await supervisor.start() await socket.messages.put({"type": "session.closed"}) with pytest.raises(RuntimeError, match="logging unavailable"): await supervisor.wait() - release.assert_awaited_once_with(budget_reservation=None) + release.assert_not_awaited() + invalidate.assert_awaited_once_with(budget_reservation=None) + assert logger.model_call_details["realtime_accounting_incomplete"] is True assert socket.closed sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True) +@pytest.mark.asyncio +async def test_start_rejects_terminal_session_while_accounting_is_pending(): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + dispatch_started = asyncio.Event() + allow_dispatch = asyncio.Event() + + async def log_messages(*, wait_for_dispatch=False): + dispatch_started.set() + await allow_dispatch.wait() + logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True + + sink = MagicMock() + sink.log_messages = AsyncMock(side_effect=log_messages) + supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock()) + await socket.messages.put({"type": "session.created"}) + await socket.messages.put({"type": "session.closed"}) + startup = asyncio.create_task(supervisor.start()) + try: + await asyncio.wait_for(dispatch_started.wait(), timeout=1) + finally: + allow_dispatch.set() + with pytest.raises(RuntimeError, match="ended before"): + await asyncio.wait_for(startup, timeout=1) + assert socket.closed + sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True) + + +@pytest.mark.asyncio +async def test_shutdown_bounds_accounting_and_invalidates_partial_dispatch(monkeypatch): + from litellm.proxy.realtime_endpoints import call_supervision + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + dispatch_cancelled = asyncio.Event() + invalidate = AsyncMock() + release = AsyncMock() + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) + monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release) + + async def log_messages(*, wait_for_dispatch=False): + try: + await asyncio.Event().wait() + finally: + dispatch_cancelled.set() + + async def hangup(): + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + + sink = MagicMock() + sink.log_messages = AsyncMock(side_effect=log_messages) + supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), hangup, logging_timeout=0.01) + registry = CallSupervisors() + await socket.messages.put({"type": "session.created"}) + await registry.start(supervisor) + await asyncio.wait_for(registry.shutdown(), timeout=1) + assert dispatch_cancelled.is_set() + assert socket.closed + assert logger.model_call_details["realtime_accounting_incomplete"] is True + assert not logger.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY) + invalidate.assert_awaited_once_with(budget_reservation=None) + release.assert_not_awaited() + sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True) + + @pytest.mark.asyncio async def test_shutdown_waits_for_usage_dispatch_completion(): socket = Socket() diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index b78ec7dcff6..2f7cc16a70d 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -1,5 +1,6 @@ import datetime as real_datetime import smtplib +from unittest.mock import MagicMock, patch import pytest from fastapi import HTTPException @@ -8,15 +9,10 @@ from litellm.caching.caching import DualCache from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import ProxyErrorTypes, UserAPIKeyAuth -from litellm.proxy.utils import ProxyLogging +from litellm.proxy.utils import ProxyLogging, get_custom_url, join_paths from litellm.types.guardrails import GuardrailEventHooks -from unittest.mock import MagicMock, patch - -from litellm.proxy.utils import get_custom_url, join_paths - - def test_get_custom_url(monkeypatch): monkeypatch.setenv("SERVER_ROOT_PATH", "/litellm") custom_url = get_custom_url(request_base_url="http://0.0.0.0:4000", route="ui/") @@ -2030,7 +2026,9 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()): await proxy_logging_obj.post_call_failure_hook( request_data={"metadata": {}}, - original_exception=HTTPException(status_code=400, detail="Upstream passthrough request failed with status 400"), + original_exception=HTTPException( + status_code=400, detail="Upstream passthrough request failed with status 400" + ), user_api_key_dict=UserAPIKeyAuth(), traceback_str=upstream_traceback, ) @@ -2038,3 +2036,127 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp assert recorder.received_traceback is not None assert provider_key not in recorder.received_traceback assert "REDACTED" in recorder.received_traceback + + +@pytest.mark.asyncio +@pytest.mark.parametrize("limiter_version", [1, 3]) +@pytest.mark.parametrize("limit", ["rpm_limit", "max_parallel_requests"]) +async def test_internal_realtime_observer_preserves_quota_and_custom_hooks(monkeypatch, limiter_version, limit): + import asyncio + from datetime import datetime + + from litellm.caching.caching import DualCache + from litellm.integrations.custom_logger import CustomLogger + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + _request_stash, + get_request_stash, + ) + from litellm.proxy.utils import InternalUsageCache, ProxyLogging + + observed = [] + + class Hook(CustomLogger): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + observed.append(call_type) + return {**data, "extra_headers": {"x-hook": "required"}} + + cache = DualCache() + limiter_type = _PROXY_MaxParallelRequestsHandler if limiter_version == 1 else _PROXY_MaxParallelRequestsHandler_v3 + limiter = limiter_type(InternalUsageCache(dual_cache=cache)) + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", [limiter, Hook()]) + token = _request_stash.set(None) + try: + auth = UserAPIKeyAuth(api_key="observer-quota-test", **{limit: 1}) + await proxy.pre_call_hook( + auth, {"model": "voice", "litellm_call_id": "signaling", "metadata": {}}, "arealtime_calls" + ) + await asyncio.sleep(0) + initial_stash = get_request_stash() + result = await proxy.pre_call_hook( + auth, + {"model": "voice", "litellm_call_id": "observer", "metadata": {}}, + "_arealtime", + internal_realtime_observer=True, + ) + assert result["extra_headers"] == {"x-hook": "required"} + assert observed == ["arealtime_calls", "_arealtime"] + if limiter_version == 3: + assert get_request_stash() is initial_stash + assert initial_stash.owner_litellm_call_id == "signaling" + if limit == "max_parallel_requests": + await limiter.async_log_success_event( + { + "litellm_call_id": "signaling", + "litellm_params": {"metadata": {"user_api_key": auth.api_key, "user_api_key_model_max_budget": {}}}, + }, + litellm.ModelResponse(usage=litellm.Usage(total_tokens=0)), + datetime.now(), + datetime.now(), + ) + if limiter_version == 3: + assert initial_stash.parallel_slot is None + await proxy.pre_call_hook( + auth, {"model": "voice", "litellm_call_id": "next", "metadata": {}}, "arealtime_calls" + ) + if limiter_version == 1 and limit == "max_parallel_requests": + from litellm.proxy._types import InternalRequestOrigin + + await asyncio.sleep(0) + observer_kwargs = { + "internal_request_origin": InternalRequestOrigin.REALTIME_OBSERVER, + "litellm_call_id": "observer", + "litellm_params": {"metadata": {"user_api_key": auth.api_key, "user_api_key_model_max_budget": {}}}, + } + await limiter.async_log_success_event( + observer_kwargs, + litellm.ModelResponse(usage=litellm.Usage(total_tokens=17)), + datetime.now(), + datetime.now(), + ) + current = await limiter.internal_usage_cache.async_get_cache( + key=f"{auth.api_key}::{datetime.now():%Y-%m-%d-%H-%M}::request_count", litellm_parent_otel_span=None + ) + assert current["current_requests"] == 1 + assert current["current_tpm"] == 17 + with pytest.raises(HTTPException) as error: + await proxy.pre_call_hook( + auth, + {"model": "voice", "litellm_call_id": "forged", "metadata": {}, "internal_realtime_observer": True}, + "_arealtime", + ) + assert error.value.status_code == 429 + finally: + _request_stash.reset(token) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("scope", ["key", "user", "team", "end_user"]) +async def test_internal_observer_missing_legacy_counter_only_adds_usage(scope): + from datetime import datetime + + from litellm.proxy._types import InternalRequestOrigin + from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler + from litellm.proxy.utils import InternalUsageCache + + limiter = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache=DualCache())) + metadata = {"user_api_key": "expired-key", "user_api_key_model_max_budget": {}} + if scope in ("user", "team"): + metadata[f"user_api_key_{scope}_id"] = "expired-scope" + kwargs = { + "internal_request_origin": InternalRequestOrigin.REALTIME_OBSERVER, + "litellm_params": {"metadata": metadata}, + **({"user": "expired-scope"} if scope == "end_user" else {}), + } + await limiter.async_log_success_event( + kwargs, litellm.ModelResponse(usage=litellm.Usage(total_tokens=23)), datetime.now(), datetime.now() + ) + identity = "expired-key" if scope == "key" else "expired-scope" + current = await limiter.internal_usage_cache.async_get_cache( + key=f"{identity}::{datetime.now():%Y-%m-%d-%H-%M}::request_count", litellm_parent_otel_span=None + ) + assert current == {"current_requests": 0, "current_tpm": 23, "current_rpm": 0} From 8ec259bdbfb82631d864d474bb3b604dc400be4e Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 17:08:07 +0200 Subject: [PATCH 30/90] fix(chatgpt): protect realtime identity and reserve hangup time --- litellm/llms/chatgpt/realtime.py | 19 +++-- litellm/llms/custom_httpx/llm_http_handler.py | 9 +- .../realtime_endpoints/call_supervision.py | 11 ++- .../llms/chatgpt/test_realtime.py | 83 +++++++++++++++++-- .../custom_httpx/test_llm_http_handler.py | 59 +++++++++++++ .../test_call_supervision.py | 41 +++++++++ 6 files changed, 201 insertions(+), 21 deletions(-) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 27ce6cb3966..559a6ef18f6 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -5,7 +5,6 @@ from typing import TYPE_CHECKING, Final from httpx import URL, QueryParams from pydantic import TypeAdapter -from websockets.exceptions import ConnectionClosed from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.llms.openai.realtime.handler import OpenAIRealtime @@ -114,6 +113,8 @@ class ChatGPTRealtime(OpenAIRealtime): ) 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"}') @@ -124,7 +125,8 @@ class ChatGPTRealtime(OpenAIRealtime): await self.hangup_call(api_base) async def hangup_call(self, api_base: str) -> None: - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + 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( @@ -132,12 +134,9 @@ class ChatGPTRealtime(OpenAIRealtime): path=f"{base.path.rstrip('/')}/realtime/calls/{self._call_id}/hangup", params=tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")), ) - client: Final = AsyncHTTPHandler() - try: - response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10) - response.raise_for_status() - finally: - await client.close() + 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: @@ -182,7 +181,9 @@ class ChatGPTRealtime(OpenAIRealtime): base.copy_with( scheme="wss" if base.scheme in ("https", "wss") else "ws", path=f"{base.path.rstrip('/')}/{endpoint}", - params=query_params, + params=QueryParams(TypeAdapter(Mapping[str, str | None]).validate_python(query_params)).merge( + tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")) + ), ) ) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 974f52b4268..44099efedeb 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6330,6 +6330,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, @@ -6355,7 +6358,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, diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index e750ce61d6a..ae04582a36b 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -120,12 +120,19 @@ class CallSupervisor: 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=self._termination_timeout) + 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, deadline - asyncio.get_running_loop().time())) + 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: diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 6db166e956b..b02b3a551f2 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -1,4 +1,5 @@ import json +import sys from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import AsyncMock, patch @@ -14,13 +15,14 @@ from litellm.types.router import GenericLiteLLMParams @pytest.mark.asyncio @pytest.mark.parametrize("failure", ["closed", "network"]) -@pytest.mark.parametrize( - "hangup_status, expectation", [(200, nullcontext()), (503, pytest.raises(httpx.HTTPStatusError))] -) -async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, expectation, chatgpt_tokens): +@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", @@ -46,15 +48,21 @@ async def test_live_closed_observer_uses_independent_hangup(failure, hangup_stat return httpx.Response(hangup_status) client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) - with patch("httpx.AsyncClient", return_value=client): - with expectation: - await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1") - assert len(requests) == 1 + 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 str(requests[0].url) == "https://gateway.example/v1/realtime/calls/rtc_live_closed/hangup?gateway=tenant" assert requests[0].headers["x-gateway-token"] == "test-only" assert requests[0].headers["Authorization"] == "Bearer test-token-default" - assert client.is_closed + assert requests[0].extensions["timeout"]["read"] == 10 @pytest.mark.asyncio @@ -305,6 +313,63 @@ def test_realtime_uses_platform_endpoint_with_oauth_headers(model, endpoint, cha 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): diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index cea2d439198..29a58b788f0 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -3620,3 +3620,62 @@ def test_image_edit_handler_keeps_the_sync_transform(): assert config.transform_calls == ["sync"] assert captured["body"] == {"transformed_by": "sync"} 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" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 7d0d614f31d..c14597a012a 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -30,6 +30,47 @@ class Socket: self.closed = True +@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 09e6b3a89d7469710ae13106d39e42c9a681dac4 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 19:11:03 +0200 Subject: [PATCH 31/90] fix(chatgpt): release sideband quotas and preserve repeated queries --- litellm/llms/chatgpt/codex.py | 6 +- litellm/llms/chatgpt/realtime.py | 32 +++--- .../proxy/hooks/parallel_request_limiter.py | 69 ++++++++++++- .../hooks/parallel_request_limiter_v3.py | 9 ++ .../proxy/realtime_endpoints/call_sessions.py | 22 ++++- .../realtime_endpoints/call_supervision.py | 21 +++- tests/test_litellm/llms/chatgpt/test_codex.py | 21 ++++ .../llms/chatgpt/test_realtime.py | 22 ++++- .../hooks/test_parallel_request_limiter.py | 83 ++++++++++++++-- .../realtime_endpoints/test_call_sessions.py | 97 +++++++++++++++++++ .../test_call_supervision.py | 37 ++++++- 11 files changed, 389 insertions(+), 30 deletions(-) diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index f71a4d5bbad..1a66f30008c 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -21,7 +21,7 @@ class CodexRealtimeCall(BaseModel): alias: str api_base: str | None = None extra_headers: Mapping[str, str] | None = None - extra_query: Mapping[str, str] | None = None + extra_query: Mapping[str, str | tuple[str, ...]] | None = None usage_supervised: bool = False owner: str expires_at: float @@ -32,7 +32,7 @@ class ChatGPTCallRouting(BaseModel): model_id: str | None = None api_base: str | None = None extra_headers: Mapping[str, str] | None = None - extra_query: Mapping[str, str] | None = None + extra_query: Mapping[str, str | tuple[str, ...]] | None = None class CodexSidebandRequest(TypedDict): @@ -41,7 +41,7 @@ class CodexSidebandRequest(TypedDict): chatgpt_realtime_call_id: ReadOnly[str] query_params: ReadOnly[RealtimeQueryParams] extra_headers: ReadOnly[Mapping[str, str] | None] - extra_query: ReadOnly[Mapping[str, str] | None] + extra_query: ReadOnly[Mapping[str, str | tuple[str, ...]] | None] def build_call_request( diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 559a6ef18f6..cc5773db66c 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -36,18 +36,18 @@ def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping return MappingProxyType({key.lower(): value for key, value in validated.items()}) -def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str]: +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]).validate_python( - getattr(params, "extra_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: value for key, value in inbound.items() if key in ("intent", "architecture")}, - **QueryParams(configured), - } + {key: merged[key] if len(merged.get_list(key)) == 1 else tuple(merged.get_list(key)) for key in merged} ) @@ -132,7 +132,11 @@ class ChatGPTRealtime(OpenAIRealtime): url: Final = base.copy_with( scheme="https" if base.scheme in ("https", "wss") else "http", path=f"{base.path.rstrip('/')}/realtime/calls/{self._call_id}/hangup", - params=tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")), + 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) @@ -166,7 +170,9 @@ class ChatGPTRealtime(OpenAIRealtime): endpoint: Final = realtime_endpoint(query_params.get("model", "")) if self._call_id: gateway_query: Final = tuple( - (key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id") + (key, value) + for key, value in QueryParams(self._extra_query).multi_items() + if key not in ("model", "call_id") ) return str( base.copy_with( @@ -182,7 +188,11 @@ class ChatGPTRealtime(OpenAIRealtime): 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 self._extra_query.items() if key not in ("model", "call_id")) + tuple( + (key, value) + for key, value in QueryParams(self._extra_query).multi_items() + if key not in ("model", "call_id") + ) ), ) ) diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 2ec576b97b3..819f0b8324a 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 @@ -50,11 +51,64 @@ 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 + + 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"] = _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, user_api_key_dict) + + async def _release_realtime_counter(self, key: str, user_api_key_dict: UserAPIKeyAuth) -> None: + raw: Final[object] = await self.internal_usage_cache.async_get_cache( + key=key, + local_only=True, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + if raw is None: + return + current: Final = TypeAdapter(Mapping[str, int]).validate_python(raw) + await self.internal_usage_cache.async_set_cache( + key=key, + value={**current, "current_requests": max(current["current_requests"] - 1, 0)}, + ttl=60, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + def print_verbose(self, print_statement): try: verbose_proxy_logger.debug(print_statement) @@ -142,6 +196,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: @@ -299,6 +356,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") @@ -480,6 +540,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, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index c6c3dde4b6e..905a2f02fb5 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4677,6 +4677,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) stash.parallel_slot = 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): """ Post-call hook to update rate limit headers in the response. diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 8af489e8eb1..1cec36876b2 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -38,6 +38,12 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth, ) from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.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 +) from litellm.proxy.spend_tracking.budget_reservation import ( invalidate_budget_reservation_counters, release_or_invalidate_budget_reservation, @@ -324,6 +330,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP 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 try: try: api_key: Final = get_websocket_api_key(websocket) @@ -364,6 +371,13 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP 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: processed, logging_obj = await process_codex_request(request, data, auth, call.alias, "_arealtime") except Exception: # noqa: BLE001 # custom hook exceptions must reject the connection @@ -397,5 +411,9 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP }, ) finally: - 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) + try: + if attachment_limiter is not None: + await attachment_limiter.async_release_realtime_attachment(data, auth) + finally: + 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 index ae04582a36b..ddf1e8f3f49 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -3,7 +3,7 @@ from collections.abc import AsyncIterator, Awaitable, Callable from contextlib import suppress from typing import Final, Protocol -from pydantic import BaseModel +from pydantic import BaseModel, Field, ValidationError from websockets.exceptions import ConnectionClosedOK from litellm._logging import verbose_proxy_logger @@ -33,6 +33,14 @@ class _ObserverEvent(BaseModel): type: str +class _LiveDurationUsage(BaseModel): + audio_duration_ms: float = Field(strict=True, ge=0, allow_inf_nan=False) + + +class _LiveTerminalEvent(BaseModel): + usage: _LiveDurationUsage + + class CallSupervisor: def __init__( self, @@ -66,6 +74,7 @@ class CallSupervisor: self._stop = asyncio.Event() self._started = False self._terminal = False + self._terminal_usage_valid = False self._close_confirmed = False self._accounting_complete = False self._task: asyncio.Task[None] | None = None @@ -106,10 +115,18 @@ class CallSupervisor: self._ready.set() if event.type == "session.closed": self._terminal = True + try: + _LiveTerminalEvent.model_validate_json(message) + except ValidationError: + self._terminal_usage_valid = False + else: + self._terminal_usage_valid = True return def _usage_complete(self) -> bool: - return self._terminal or (not self._terminal_usage_required and self._close_confirmed) + 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: reader: Final = asyncio.create_task(self._read()) diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 402afaaa58c..35606931e24 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -1,9 +1,30 @@ +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}, diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index b02b3a551f2..f3258fd20b4 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -27,7 +27,7 @@ async def test_live_closed_observer_uses_independent_hangup(failure, hangup_stat GenericLiteLLMParams( chatgpt_realtime_call_id="rtc_live_closed", chatgpt_token_dir=chatgpt_tokens, - extra_query={"gateway": "tenant"}, + extra_query={"gateway": "tenant", "tag": ["alpha +/&", "beta"]}, ), {}, {"x-gateway-token": "test-only"}, @@ -59,7 +59,9 @@ async def test_live_closed_observer_uses_independent_hangup(failure, hangup_stat await client.aclose() assert len(requests) == 2 assert requests[0].method == "POST" - assert str(requests[0].url) == "https://gateway.example/v1/realtime/calls/rtc_live_closed/hangup?gateway=tenant" + assert requests[0].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 @@ -138,6 +140,8 @@ async def test_routed_call_preserves_deployment_gateway_headers( "enabled": True, "disabled": False, "blank": None, + "tag": ["alpha +/&", "beta"], + "empty": [], "model": "other-model", "call_id": "rtc_wrong", }, @@ -163,10 +167,16 @@ async def test_routed_call_preserves_deployment_gateway_headers( "enabled": "true", "disabled": "false", "blank": "", + "tag": "alpha +/&", "model": "other-model", "call_id": "rtc_wrong", } - assert response.extensions["chatgpt_realtime"]["extra_query"] == dict(requests[0].url.params) + 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 @@ -180,6 +190,7 @@ async def test_routed_call_preserves_deployment_gateway_headers( 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() @@ -203,6 +214,8 @@ async def test_websocket_forwards_configured_headers_without_client_identity(mod 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", @@ -213,6 +226,9 @@ async def test_websocket_forwards_configured_headers_without_client_identity(mod ) 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" 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 0e2683dcbfd..0bf488016c2 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -7,6 +7,8 @@ from datetime import datetime import pytest from litellm.caching.caching import DualCache +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, ) @@ -14,6 +16,78 @@ from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage +@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.parametrize( "response_obj", [ @@ -39,9 +113,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") @@ -80,7 +152,4 @@ 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']}" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index ef85e324680..50e3aebfbe2 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -83,6 +83,103 @@ async def test_sideband_preserves_pending_cost_reconciliation(monkeypatch, logge 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( diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index c14597a012a..93b97a28fe6 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -147,6 +147,34 @@ class Sink: 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)] +) +async def test_live_terminal_requires_valid_duration_for_accounting(monkeypatch, duration, valid): + 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": {"audio_duration_ms": 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) @@ -454,7 +482,14 @@ async def test_shutdown_allows_hangup_longer_than_usage_drain_timeout(): hangup_finished.set() supervisor = CallSupervisor( - socket, sink, logger, UserAPIKeyAuth(), hangup, drain_timeout=0.01, termination_timeout=1 + 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"}) From 2642d9830234792ba80fbf372b971ef2f0a700d0 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 22:08:10 +0200 Subject: [PATCH 32/90] fix(chatgpt): release attachment quota before upstream websocket close --- .../litellm_core_utils/realtime_streaming.py | 14 ++- .../proxy/realtime_endpoints/call_sessions.py | 19 +++- .../test_realtime_streaming.py | 34 +++++++ .../realtime_endpoints/test_call_sessions.py | 91 +++++++++++++++++++ 4 files changed, 154 insertions(+), 4 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index c5e6b0fafe3..18f7eaa0e32 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1,7 +1,8 @@ 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 typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast @@ -26,6 +27,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 @@ -1581,7 +1586,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/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 1cec36876b2..6ed2e83ab1c 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -2,8 +2,9 @@ import base64 import hashlib import json import time -from collections.abc import Mapping +from collections.abc import Awaitable, Callable, Mapping from contextlib import AsyncExitStack +from contextvars import Token from types import MappingProxyType from typing import Final, Literal @@ -14,7 +15,11 @@ from starlette.types import Message from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY, RealTimeStreaming +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, @@ -331,6 +336,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP ) 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) @@ -387,6 +393,13 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP 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, @@ -415,5 +428,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP 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/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index cba46374837..5b263f9d557 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3438,3 +3438,37 @@ async def test_live_attachment_does_not_dispatch_duplicate_usage(): await stream.log_messages() worker.ensure_initialized_and_enqueue.assert_not_called() logger.dispatch_success_handlers.assert_not_called() + + +@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 diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 50e3aebfbe2..2591d45f9a5 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -664,3 +664,94 @@ async def test_supervisor_constructor_failure_closes_effective_connection(monkey release.assert_awaited_once_with(budget_reservation=auth.budget_reservation) invalidate.assert_not_awaited() assert "private-cleanup-credential" 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 From 34dcfbf833373c7fa47ae56600041d6708f68d5b Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 11 Sep 2026 02:25:21 +0200 Subject: [PATCH 33/90] fix(chatgpt): retain realtime call quotas and isolate provider routing --- .../base_llm/realtime/http_transformation.py | 14 + litellm/llms/chatgpt/codex.py | 1 + litellm/llms/chatgpt/realtime.py | 34 ++- litellm/llms/openai/realtime/handler.py | 17 +- .../proxy/hooks/parallel_request_limiter.py | 40 ++- .../hooks/parallel_request_limiter_v3.py | 118 +++++++- litellm/proxy/hooks/realtime_call_lease.py | 74 +++++ .../proxy/realtime_endpoints/call_sessions.py | 97 +++++-- .../realtime_endpoints/call_supervision.py | 29 +- litellm/realtime_api/main.py | 82 +++--- litellm/utils.py | 23 +- .../hooks/test_parallel_request_limiter.py | 105 +++++++ .../hooks/test_parallel_request_limiter_v3.py | 129 +++++++++ .../proxy/hooks/test_realtime_call_lease.py | 85 ++++++ .../realtime_endpoints/test_call_sessions.py | 274 ++++++++++++++++++ .../test_call_supervision.py | 36 +++ tests/test_litellm/realtime_api/test_main.py | 26 ++ 17 files changed, 1090 insertions(+), 94 deletions(-) create mode 100644 litellm/proxy/hooks/realtime_call_lease.py create mode 100644 tests/test_litellm/proxy/hooks/test_realtime_call_lease.py 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/codex.py b/litellm/llms/chatgpt/codex.py index 1a66f30008c..ea736c97a0d 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -23,6 +23,7 @@ class CodexRealtimeCall(BaseModel): 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 diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index cc5773db66c..f6aa04fc29d 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -3,7 +3,7 @@ from enum import Enum, auto from types import MappingProxyType from typing import TYPE_CHECKING, Final -from httpx import URL, QueryParams +from httpx import URL, QueryParams, Response from pydantic import TypeAdapter from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES @@ -156,6 +156,16 @@ class ChatGPTRealtime(OpenAIRealtime): 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 @@ -212,6 +222,14 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): ) -> 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, @@ -223,6 +241,20 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): 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 diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index ca141c0b958..a896594db41 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -38,6 +38,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, @@ -125,8 +133,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 +151,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,7 +180,7 @@ class OpenAIRealtime(OpenAIChatCompletion): model if (query_params or {}).get("intent") == "transcription" else None ), event_normalizer=self._make_event_normalizer(), - account_usage=account_usage, + account_usage=account_usage and self._accounts_for_call_usage(), ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 819f0b8324a..5bf9ab1b8f3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -68,6 +68,16 @@ class _RealtimeAttachmentReservations(BaseModel): 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): @@ -91,23 +101,23 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): litellm_parent_otel_span=user_api_key_dict.parent_otel_span, ) for key in keys: - await self._release_realtime_counter(key, user_api_key_dict) + await self._release_realtime_counter(key) - async def _release_realtime_counter(self, key: str, user_api_key_dict: UserAPIKeyAuth) -> None: - raw: Final[object] = await self.internal_usage_cache.async_get_cache( - key=key, - local_only=True, - litellm_parent_otel_span=user_api_key_dict.parent_otel_span, - ) - if raw is None: - return - current: Final = TypeAdapter(Mapping[str, int]).validate_python(raw) - await self.internal_usage_cache.async_set_cache( - key=key, - value={**current, "current_requests": max(current["current_requests"] - 1, 0)}, - ttl=60, - litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + 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).validate_python(raw) + updated: Final = ( + {**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: diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 905a2f02fb5..67412ccfb96 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -8,10 +8,12 @@ import asyncio import binascii import os import uuid -from collections.abc import Awaitable, Callable, Mapping, Sequence, Set +from collections.abc import Awaitable, Callable, Generator, Mapping, Sequence, Set +from contextlib import contextmanager from contextvars import ContextVar from dataclasses import dataclass, field from datetime import datetime +from types import MappingProxyType from typing import ( TYPE_CHECKING, Any, @@ -53,6 +55,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, @@ -304,6 +307,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 = {} @@ -416,6 +436,14 @@ 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] @@ -546,6 +574,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: @@ -595,6 +632,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, @@ -627,6 +665,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.token_increment_script = None @@ -635,6 +676,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)) @@ -1598,6 +1640,66 @@ 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=list(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: + result: Final = await self.parallel_renew_script( + keys=counter_keys, args=(slot_id, PARALLEL_REQUEST_SLOT_TTL_SECONDS) + ) + return tuple(result) == (1,) + 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}, + ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, + local_only=True, + litellm_parent_otel_span=None, + ) + return True + async def _release_parallel_request_slots( self, acquisition: ParallelSlotAcquisition, @@ -3477,6 +3579,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Org Level Rate Limits descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model)) + effective_descriptors: Final = ( + tuple(_without_parallel_limit(descriptor) for descriptor in descriptors) + if call_type == "_arealtime" and is_realtime_call_attachment(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. @@ -3495,16 +3603,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( 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/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 6ed2e83ab1c..8ea329945c5 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -1,9 +1,10 @@ +import asyncio import base64 import hashlib import json import time from collections.abc import Awaitable, Callable, Mapping -from contextlib import AsyncExitStack +from contextlib import AsyncExitStack, nullcontext from contextvars import Token from types import MappingProxyType from typing import Final, Literal @@ -48,24 +49,38 @@ from litellm.proxy.hooks.parallel_request_limiter import ( ) 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) -> None: +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: - return { + body: Final[RealtimeQueryParams] = {"model": call.alias} + message: Final[Message] = { "type": "http.request", - "body": json.dumps({"model": call.alias}).encode(), + "body": json.dumps(body).encode(), "more_body": False, - } # mutable-ok: ASGI message + } + return message async def send(_message: Message) -> None: return None @@ -143,6 +158,7 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, 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() @@ -241,6 +257,11 @@ async def process_codex_request( async def create_codex_realtime_call(request: Request) -> Response: + with isolated_request_stash(): + return await _create_codex_realtime_call(request) + + +async def _create_codex_realtime_call(request: Request) -> Response: from litellm.proxy import proxy_server as server try: @@ -275,6 +296,10 @@ async def create_codex_realtime_call(request: Request) -> Response: 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, @@ -286,17 +311,30 @@ async def create_codex_realtime_call(request: Request) -> Response: signaling_auth: Final = auth.model_copy( update={"budget_reservation": None} ) # mutable-ok: Pydantic update contract + 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") - result: Final = await server.route_request( - data=processed, - route_type="arealtime_calls", - llm_router=server.llm_router, - user_model=server.user_model, - ) - try: - response: Final = await result - except BaseLLMException as exc: - raise HTTPException(exc.status_code, str(exc)) from exc + 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=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: @@ -311,11 +349,15 @@ async def create_codex_realtime_call(request: Request) -> Response: except ValueError as exc: raise HTTPException(400, str(exc)) from exc supervised_call: Final = call.model_copy( - update={"usage_supervised": True} + update={"usage_supervised": True, "parallel_reserved": call_lease is not None} ) # mutable-ok: Pydantic update contract token: Final = encode_call(supervised_call) supervision_started = True - await supervise_codex_call(request, supervised_call, auth) + 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, @@ -323,8 +365,22 @@ async def create_codex_realtime_call(request: Request) -> Response: headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), ) finally: - if not supervision_started: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + 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: @@ -385,7 +441,8 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP if isinstance(limiter, _PROXY_MaxParallelRequestsHandler): limiter.begin_realtime_attachment(data) try: - processed, logging_obj = await process_codex_request(request, data, auth, call.alias, "_arealtime") + 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") diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index ddf1e8f3f49..8947edb8758 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -11,6 +11,7 @@ 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, @@ -57,6 +58,7 @@ class CallSupervisor: logging_timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE, terminal_usage_required: bool = True, force_close_call: Callable[[], Awaitable[None]] | None = None, + lease: RealtimeCallLease | None = None, ) -> None: self._upstream = upstream self._stream = stream @@ -64,6 +66,7 @@ class CallSupervisor: 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 @@ -85,6 +88,8 @@ class CallSupervisor: 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: @@ -129,10 +134,25 @@ class CallSupervisor: 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), timeout=self._lifetime, return_when=asyncio.FIRST_COMPLETED) + 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: @@ -161,7 +181,12 @@ class CallSupervisor: finally: stopped.cancel() reader.cancel() - await asyncio.gather(reader, stopped, return_exceptions=True) + if lease_failed is not None: + lease_failed.cancel() + await asyncio.gather( + *((reader, stopped, lease_failed) if lease_failed is not None else (reader, stopped)), + return_exceptions=True, + ) with suppress(Exception): await self._upstream.close() if not self._usage_complete(): diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 23963475ed0..0b8a32dcfc1 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -76,7 +76,7 @@ def _get_realtime_http_provider_config( dynamic_api_base: str | None, dynamic_api_key: str | None, litellm_params: GenericLiteLLMParams, - use_codex_backend: bool = False, + is_call: bool = False, ) -> tuple["BaseRealtimeHTTPConfig | None", str, str]: """ Return (provider_config, resolved_api_base, resolved_api_key) for the @@ -90,23 +90,19 @@ def _get_realtime_http_provider_config( ) provider_config: BaseRealtimeHTTPConfig | None = None - if custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig - - provider_config = ChatGPTRealtimeHTTPConfig(litellm_params, use_codex_backend=use_codex_backend) - elif custom_llm_provider in LlmProviders._member_map_.values(): + if custom_llm_provider in LlmProviders._member_map_.values(): 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 = ( - litellm_params.api_base if custom_llm_provider == "chatgpt" else dynamic_api_base or litellm_params.api_base - ) + 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). @@ -260,8 +256,6 @@ async def arealtime_calls( timeout: float | None = None, **kwargs, ): - from litellm.llms.chatgpt.realtime import realtime_call_headers - model_name = model or "gpt-4o-realtime-preview" litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj") litellm_params: Final = GenericLiteLLMParams(**kwargs) @@ -281,12 +275,14 @@ async def arealtime_calls( dynamic_api_base=dynamic_api_base, dynamic_api_key=dynamic_api_key, litellm_params=litellm_params, - use_codex_backend=True, + is_call=True, ) if session is not None: session = _with_resolved_session_model(session, model_name) call_headers: Final = ( - realtime_call_headers(litellm_params) if custom_llm_provider == "chatgpt" else kwargs.get("extra_headers") + provider_config.get_realtime_calls_extra_headers(kwargs.get("extra_headers")) + if provider_config is not None + else kwargs.get("extra_headers") ) litellm_logging_obj.update_from_kwargs( kwargs=kwargs, @@ -308,23 +304,13 @@ async def arealtime_calls( client=kwargs.get("client"), api_version=litellm_params.api_version, ) - if custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ( - ChatGPTRealtime, - configured_realtime_headers, - configured_realtime_query, + return ( + provider_config.transform_realtime_calls_response( + response, model_name, litellm_logging_obj.get_router_model_id(), call_headers ) - - response.extensions["chatgpt_realtime"] = MappingProxyType( - { - "model": model_name, - "model_id": litellm_logging_obj.get_router_model_id(), - "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), - "extra_headers": configured_realtime_headers(call_headers), - "extra_query": configured_realtime_query(litellm_params), - } - ) - return response + if provider_config is not None + else response + ) async def vertex_access_token_resolver( @@ -421,7 +407,26 @@ 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, websocket.headers, headers + ) + if _custom_llm_provider in LlmProviders._member_map_.values() + else None + ) + if provider_handler is not None: + 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=kwargs.get("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, @@ -469,21 +474,6 @@ async def _arealtime( user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata(kwargs), ) - elif _custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtime, accounts_for_call_usage - - await ChatGPTRealtime(litellm_params, websocket.headers, headers).async_realtime( - model=model, - websocket=websocket, - logging_obj=litellm_logging_obj, - api_base=ChatGPTRealtime.get_api_base(api_base), - api_key="chatgpt-oauth", - timeout=timeout, - query_params=query_params, - user_api_key_dict=kwargs.get("user_api_key_dict"), - litellm_metadata=_build_litellm_metadata(kwargs), - account_usage=accounts_for_call_usage(litellm_params), - ) elif _custom_llm_provider == "openai": api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or "https://api.openai.com/" # set API KEY diff --git a/litellm/utils.py b/litellm/utils.py index 5aab0210d4a..b4f1ce8afc2 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -419,6 +419,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.mistral.ocr.transformation import MistralOCRConfig + 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, @@ -434,7 +435,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 @@ -9211,16 +9212,36 @@ class ProviderConfigManager: return GeminiRealtimeConfig() return None + @staticmethod + def get_provider_realtime_handler( + provider: LlmProviders, + params: GenericLiteLLMParams, + headers: 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, 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, 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 0bf488016c2..3ee364a57a4 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -2,11 +2,17 @@ Unit Tests for the max parallel request limiter v1 for the proxy """ +import asyncio +import shutil +import socket +import subprocess from datetime import datetime +from unittest.mock import AsyncMock, MagicMock 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 ( @@ -16,6 +22,105 @@ from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage +@pytest.fixture +def isolated_legacy_redis(tmp_path): + executable = shutil.which("redis-server") + if executable is None: + pytest.skip("redis-server is required to exercise atomic Lua updates") + with socket.socket() as listener: + listener.bind(("127.0.0.1", 0)) + port = listener.getsockname()[1] + process = subprocess.Popen( + [ + executable, + "--bind", + "127.0.0.1", + "--port", + str(port), + "--save", + "", + "--appendonly", + "no", + "--dir", + str(tmp_path), + ], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + try: + import redis + + client = redis.Redis(host="127.0.0.1", port=port) + for _ in range(100): + try: + client.ping() + break + except redis.ConnectionError: + import time + + time.sleep(0.01) + else: + pytest.fail("isolated Redis did not start") + yield port + client.close() + finally: + process.terminate() + process.wait(timeout=5) + + +@pytest.mark.asyncio +async def test_concurrent_realtime_releases_update_redis_without_lost_decrement(isolated_legacy_redis): + remote = RedisCache(host="127.0.0.1", port=isolated_legacy_redis, namespace="legacy-test") + 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_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): 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 6d382370f5f..181b3d8901e 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 @@ -31,6 +31,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.utils import ( EmbeddingResponse, @@ -51,6 +53,133 @@ 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.mark.asyncio +async def test_realtime_lease_redis_renewal_is_atomic_and_does_not_resurrect(): + import shutil + import subprocess + import tempfile + from redis.asyncio import Redis + from redis.exceptions import ConnectionError as RedisConnectionError + from litellm.proxy.hooks.parallel_request_limiter_v3 import PARALLEL_RENEW_SCRIPT + + executable = shutil.which("redis-server") + if executable is None: + pytest.skip("redis-server is required for the Lua regression") + with tempfile.TemporaryDirectory(prefix="rtc-") as temporary: + socket = f"{temporary}/redis.sock" + process = subprocess.Popen( + [executable, "--port", "0", "--unixsocket", socket, "--save", "", "--appendonly", "no"], + stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) + client = Redis(unix_socket_path=socket) + try: + for attempt in range(100): + try: + await client.ping() + break + except RedisConnectionError: + await asyncio.sleep(0.01) + else: + pytest.fail("isolated Redis did not start") + now = (await client.time())[0] + await client.zadd("first", {"owner": now - 10, "other": now}) + await client.zadd("second", {"owner": now - PARALLEL_REQUEST_SLOT_TTL_SECONDS}) + renew = client.register_script(PARALLEL_RENEW_SCRIPT) + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] + assert await client.zscore("first", "owner") == now - 10 + await client.zadd("second", {"owner": now - 10}) + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [1] + assert await client.zscore("first", "owner") >= now + assert await client.ttl("first") > PARALLEL_REQUEST_SLOT_TTL_SECONDS - 10 + await client.zrem("second", "owner") + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] + assert await client.zscore("second", "owner") is None + assert await client.zscore("first", "other") == now + finally: + await client.aclose() + process.terminate() + process.wait(timeout=5) + + @pytest.fixture def time_controller(monkeypatch): controller = TimeController() 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/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 2591d45f9a5..75e1f64a65b 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -180,6 +180,7 @@ async def test_supervised_attachments_release_real_limiter_before_reconnect(monk 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( @@ -755,3 +756,276 @@ async def test_attachment_releases_quota_before_upstream_close_handshake(monkeyp "{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 diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 93b97a28fe6..3377f63433f 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -30,6 +30,42 @@ class Socket: 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): diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index 761e87ac764..3ab843dcea7 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -331,3 +331,29 @@ async def test_arealtime_azure_ai_on_a_foundry_host_connects_to_the_azure_openai "wss://my-project.services.ai.azure.com/openai/realtime" "?api-version=2024-10-01-preview&deployment=gpt-realtime-mini" ) + + +@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" + } From 464c1eb2bc30f4b18b62658cc33eefd728552e14 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 11 Sep 2026 10:20:26 +0200 Subject: [PATCH 34/90] fix(realtime): enforce multipart authorization and cluster call leases --- litellm/proxy/auth/auth_utils.py | 17 +- .../hooks/parallel_request_limiter_v3.py | 157 ++++++++++- .../realtime_endpoints/call_supervision.py | 9 +- pyproject.toml | 2 +- .../proxy/auth/test_auth_utils.py | 23 +- .../hooks/test_parallel_request_limiter.py | 65 ++--- .../hooks/test_parallel_request_limiter_v3.py | 252 +++++++++++++++--- .../realtime_endpoints/test_call_sessions.py | 96 +++++++ .../test_call_supervision.py | 80 ++++++ uv.lock | 72 ++++- 10 files changed, 662 insertions(+), 111 deletions(-) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index aa77835e858..ae19b570a66 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1865,15 +1865,26 @@ def _extract_model_candidates_from_request( uses_completion_model_sources: Final = _route_matches_any_marker( route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS ) + session: Final[object] = ( + request_data.get("session") + if _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS) + 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 ( + _route_matches_any_marker(route=route, markers=("/realtime/calls",)) + and isinstance(session_model, str) + and session_model + ): + return [session_model] 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")) + _append_model_candidates(candidates, 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/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 67412ccfb96..f65be9fcf21 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -446,7 +446,7 @@ def _without_parallel_limit(descriptor: RateLimitDescriptor) -> RateLimitDescrip class ParallelSlotAcquisition(TypedDict): slot_id: str - counter_keys: list[str] + counter_keys: Sequence[str] class RateLimitStatus(TypedDict): @@ -1038,7 +1038,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: @@ -1192,7 +1192,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): crc: Final = binascii.crc_hqx(key.encode("utf-8"), 0) return crc % REDIS_CLUSTER_SLOTS - 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. @@ -1212,7 +1212,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 @@ -1508,12 +1508,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, never a 500 @@ -1571,6 +1574,133 @@ 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) + (raw,) = ( + await self.parallel_acquire_script( + keys=keys, + args=tuple( + arg + for key in keys + for arg in (by_key[key]["limit"], PARALLEL_REQUEST_SLOT_TTL_SECONDS, slot_id) + ), + ), + ) + 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=[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), + ) + + 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], @@ -1656,7 +1786,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def release() -> None: await self._release_parallel_request_slots( - ParallelSlotAcquisition(slot_id=slot_id, counter_keys=list(counter_keys)) + ParallelSlotAcquisition(slot_id=slot_id, counter_keys=counter_keys) ) return RealtimeCallLease(renew=renew, release=release) @@ -1664,10 +1794,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _renew_realtime_call_slot(self, slot_id: str, counter_keys: tuple[str, ...]) -> bool: if self.parallel_renew_script is not None: try: - result: Final = await self.parallel_renew_script( - keys=counter_keys, args=(slot_id, PARALLEL_REQUEST_SLOT_TTL_SECONDS) - ) - return tuple(result) == (1,) + 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: @@ -1717,11 +1849,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), ) for counter_key, remaining in zip(counter_keys, raw): await self.internal_usage_cache.async_set_cache( diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index 8947edb8758..f511d943c3a 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -93,9 +93,16 @@ class CallSupervisor: 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() + 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() diff --git a/pyproject.toml b/pyproject.toml index 04f2f3fd1dd..8b058272b33 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -207,7 +207,7 @@ dev = [ "opentelemetry-instrumentation-fastapi==0.49b0", "langfuse==2.59.7", "fastapi-offline==1.7.6", - "fakeredis==2.34.1", + "fakeredis[lua]==2.34.1", "pytest-rerunfailures==15.1", "pytest-cov==5.0.0", "parameterized==0.9.0", diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 3c725148d3f..994226b0f1a 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -947,7 +947,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. @@ -957,13 +958,31 @@ 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"}]) +def test_realtime_calls_auth_uses_executed_session_model_despite_decoys(session): + assert ( + get_model_from_request( + request_data={"model": "body-decoy", "session": session}, + route="/v1/realtime/calls", + request_query_params={"model": "query-decoy"}, + request_headers={"x-litellm-model": "header-decoy"}, + ) + == "actual-voice" + ) + + +@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/hooks/test_parallel_request_limiter.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py index 3ee364a57a4..360da3fefd6 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -3,13 +3,12 @@ Unit Tests for the max parallel request limiter v1 for the proxy """ import asyncio -import shutil -import socket -import subprocess from datetime import datetime -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest +import pytest_asyncio +from fakeredis import FakeAsyncRedis, FakeRedis, FakeServer from litellm.caching.caching import DualCache from litellm.caching.redis_cache import RedisCache @@ -22,55 +21,23 @@ from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage -@pytest.fixture -def isolated_legacy_redis(tmp_path): - executable = shutil.which("redis-server") - if executable is None: - pytest.skip("redis-server is required to exercise atomic Lua updates") - with socket.socket() as listener: - listener.bind(("127.0.0.1", 0)) - port = listener.getsockname()[1] - process = subprocess.Popen( - [ - executable, - "--bind", - "127.0.0.1", - "--port", - str(port), - "--save", - "", - "--appendonly", - "no", - "--dir", - str(tmp_path), - ], - stdout=subprocess.DEVNULL, - stderr=subprocess.DEVNULL, - ) - try: - import redis - - client = redis.Redis(host="127.0.0.1", port=port) - for _ in range(100): - try: - client.ping() - break - except redis.ConnectionError: - import time - - time.sleep(0.01) - else: - pytest.fail("isolated Redis did not start") - yield port - client.close() - finally: - process.terminate() - process.wait(timeout=5) +@pytest_asyncio.fixture(loop_scope="function") +async def isolated_legacy_redis(): + server = FakeServer() + client = FakeRedis(server=server) + async with FakeAsyncRedis(server=server) as async_client: + with ( + patch("redis.Redis", autospec=True, return_value=client), + patch("redis.asyncio.BlockingConnectionPool", autospec=True, return_value=async_client.connection_pool), + patch("redis.asyncio.Redis", autospec=True, return_value=async_client), + ): + yield RedisCache(host="fake-legacy-redis", namespace="legacy-test") + client.close() @pytest.mark.asyncio async def test_concurrent_realtime_releases_update_redis_without_lost_decrement(isolated_legacy_redis): - remote = RedisCache(host="127.0.0.1", port=isolated_legacy_redis, namespace="legacy-test") + remote = isolated_legacy_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) 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 181b3d8901e..952da528eee 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 @@ -134,50 +134,24 @@ async def test_realtime_lease_renewal_preserves_quota_past_ttl_and_does_not_resu @pytest.mark.asyncio async def test_realtime_lease_redis_renewal_is_atomic_and_does_not_resurrect(): - import shutil - import subprocess - import tempfile - from redis.asyncio import Redis - from redis.exceptions import ConnectionError as RedisConnectionError + from fakeredis import FakeAsyncRedis from litellm.proxy.hooks.parallel_request_limiter_v3 import PARALLEL_RENEW_SCRIPT - executable = shutil.which("redis-server") - if executable is None: - pytest.skip("redis-server is required for the Lua regression") - with tempfile.TemporaryDirectory(prefix="rtc-") as temporary: - socket = f"{temporary}/redis.sock" - process = subprocess.Popen( - [executable, "--port", "0", "--unixsocket", socket, "--save", "", "--appendonly", "no"], - stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, - ) - client = Redis(unix_socket_path=socket) - try: - for attempt in range(100): - try: - await client.ping() - break - except RedisConnectionError: - await asyncio.sleep(0.01) - else: - pytest.fail("isolated Redis did not start") - now = (await client.time())[0] - await client.zadd("first", {"owner": now - 10, "other": now}) - await client.zadd("second", {"owner": now - PARALLEL_REQUEST_SLOT_TTL_SECONDS}) - renew = client.register_script(PARALLEL_RENEW_SCRIPT) - assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] - assert await client.zscore("first", "owner") == now - 10 - await client.zadd("second", {"owner": now - 10}) - assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [1] - assert await client.zscore("first", "owner") >= now - assert await client.ttl("first") > PARALLEL_REQUEST_SLOT_TTL_SECONDS - 10 - await client.zrem("second", "owner") - assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] - assert await client.zscore("second", "owner") is None - assert await client.zscore("first", "other") == now - finally: - await client.aclose() - process.terminate() - process.wait(timeout=5) + async with FakeAsyncRedis() as client: + now = (await client.time())[0] + await client.zadd("first", {"owner": now - 10, "other": now}) + await client.zadd("second", {"owner": now - PARALLEL_REQUEST_SLOT_TTL_SECONDS}) + renew = client.register_script(PARALLEL_RENEW_SCRIPT) + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] + assert await client.zscore("first", "owner") == now - 10 + await client.zadd("second", {"owner": now - 10}) + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [1] + assert await client.zscore("first", "owner") >= now + assert await client.ttl("first") > PARALLEL_REQUEST_SLOT_TTL_SECONDS - 10 + await client.zrem("second", "owner") + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] + assert await client.zscore("second", "owner") is None + assert await client.zscore("first", "other") == now @pytest.fixture @@ -6365,3 +6339,197 @@ 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)) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 75e1f64a65b..d65556ebd39 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -11,6 +11,102 @@ 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("policy", ["budget", "personal_models"]) +async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypatch, multipart, policy): + import json + from unittest.mock import AsyncMock, MagicMock + + import httpx + from fastapi import Request + + 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 + + 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() + + async def receive(): + 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": [ + *((key.lower(), value) for key, value in outbound.headers.raw), + (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() + + async def custom_auth(request: Request, api_key: str): + 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" + 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) + with pytest.raises(ProxyException) as denied: + await codex.create_codex_realtime_call(request) + if policy == "personal_models": + assert "user not allowed to access model" in str(denied.value) + assert "forbidden-voice" in str(denied.value) + 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]) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 3377f63433f..6603b4a0cdb 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -276,6 +276,86 @@ async def test_cancelled_start_hangs_up_and_drains_terminal_usage(): 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() diff --git a/uv.lock b/uv.lock index 0fe787645a2..9a6ca1fa56f 100644 --- a/uv.lock +++ b/uv.lock @@ -1875,6 +1875,11 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/49/b5/82f89307d0d769cd9bf46a54fb9136be08e4e57c5570ae421db4c9a2ba62/fakeredis-2.34.1-py3-none-any.whl", hash = "sha256:0107ec99d48913e7eec2a5e3e2403d1bd5f8aa6489d1a634571b975289c48f12", size = 122160, upload-time = "2026-02-25T13:17:49.701Z" }, ] +[package.optional-dependencies] +lua = [ + { name = "lupa" }, +] + [[package]] name = "fastapi" version = "0.136.3" @@ -4520,7 +4525,7 @@ dev = [ { name = "basedpyright" }, { name = "botocore-stubs" }, { name = "diff-cover" }, - { name = "fakeredis" }, + { name = "fakeredis", extra = ["lua"] }, { name = "fastapi-offline" }, { name = "hypothesis" }, { name = "keyring" }, @@ -4708,7 +4713,7 @@ dev = [ { name = "basedpyright", specifier = "==1.39.7" }, { name = "botocore-stubs", specifier = "==1.43.14" }, { name = "diff-cover", specifier = "==9.7.2" }, - { name = "fakeredis", specifier = "==2.34.1" }, + { name = "fakeredis", extras = ["lua"], specifier = "==2.34.1" }, { name = "fastapi-offline", specifier = "==1.7.6" }, { name = "hypothesis", specifier = "==6.165.10" }, { name = "keyring", specifier = "==25.7.0" }, @@ -5003,6 +5008,69 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/db/a4/441aee36c6f6b249823d20fd91f9be9ab89d7c5a8ae542a4a4ca6d342d56/lxml-6.1.1-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:ed21202aec73cda4d55d1ce57b389aadb90ffb044e6cd1080b8347efe1b1ec84", size = 3508989, upload-time = "2026-05-18T19:18:38.158Z" }, ] +[[package]] +name = "lupa" +version = "2.8" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/c3/a6/0f869fbb07c393f15473b1eefefb7b5bec162fb7481803d040ed4dc46002/lupa-2.8.tar.gz", hash = "sha256:d8022641b9ec8ecf2c5ecbe9f47e5a70e0b87c4b5ae921b92cb02a638e0acd08", size = 6156370, upload-time = "2026-04-15T20:08:30.534Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/09/21/9be4516ddd22f8eadba336d9ba065d17d79108465ae1b7f71424ab99b9d0/lupa-2.8-cp310-abi3-win32.whl", hash = "sha256:c2a5fd15dc62374e1661a55f01744c9ec1c56f291ba4a0749d3af2174556e78f", size = 1594887, upload-time = "2026-04-15T20:05:23.377Z" }, + { url = "https://files.pythonhosted.org/packages/2d/99/1557c9685d7034d9ce8dd2b54c40a26d6deb7c67c1fdb5c801abd1a02c3f/lupa-2.8-cp310-abi3-win_arm64.whl", hash = "sha256:9e304fb1c50cf23fd8882afbe1aa87525ef8a72667bcab3b37b2bbb2bc542269", size = 1371742, upload-time = "2026-04-15T20:05:27.417Z" }, + { url = "https://files.pythonhosted.org/packages/1c/34/05ce4745b191633f90ff1ab50f1a19a37da282bb0a41fb500d9157fc9b8f/lupa-2.8-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:97bd01e90b8031e56a5fd5bb70605aea09f1dba675c1140308a52780f93d06f1", size = 1202714, upload-time = "2026-04-15T20:05:31.088Z" }, + { url = "https://files.pythonhosted.org/packages/7d/d2/f70fdbeec2d4c69ee6a469e6cddde9635fff4af4e13fb652e6a1229eef51/lupa-2.8-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:0b5ebe1a13c45767919c86750b84fe2da9f6288b6f3cea4ce7660bb2abc9d921", size = 1857453, upload-time = "2026-04-15T20:05:34.611Z" }, + { url = "https://files.pythonhosted.org/packages/97/dc/6fcda0e36e75eb6cb98dc9190fa4737d727eeae29e58f892980b2c96b656/lupa-2.8-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:097e7d0f1719a88020b67c82e05d53d7973c166952393afcecfd8434c7e19a15", size = 2408890, upload-time = "2026-04-15T20:05:37.994Z" }, + { url = "https://files.pythonhosted.org/packages/58/29/7ea176eac3c1dac83d059762daa875ad1390decc0bf2c3b4c7bbfc1f1665/lupa-2.8-cp310-cp310-win_amd64.whl", hash = "sha256:7bb223ee8f72d0dc076b0d65296ee72f1c69450f9d2fed5315f7707d98c4a03d", size = 1910396, upload-time = "2026-04-15T20:05:41.163Z" }, + { url = "https://files.pythonhosted.org/packages/b7/0a/5a740717f27aa77481e6a61b97cf79d1e0c1ede729b1268caacded915326/lupa-2.8-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:b12e43c1fb787189dfc28cd604aef0baa2cb95e27da19498d520361d0ace070a", size = 1202376, upload-time = "2026-04-15T20:05:44.049Z" }, + { url = "https://files.pythonhosted.org/packages/1b/75/6b64d0098c64275a801896cb7a6a30e7e653d25fa102c64e747292afcdbb/lupa-2.8-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:f6f603391dffb256e36a79fd2044084d5f4b8a0a4c0e5ad291cd3ab3aaf1fd0a", size = 1839271, upload-time = "2026-04-15T20:05:47.399Z" }, + { url = "https://files.pythonhosted.org/packages/7b/2f/0d4f00563046ff616ef6a421f8b776a5ffb327f7b32ed69e856d52b917a8/lupa-2.8-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:9f6f41c91366e7d0d474f87d81c1274af861f40812bf729c9f97ab4c8f3c7ac8", size = 2376251, upload-time = "2026-04-15T20:05:49.891Z" }, + { url = "https://files.pythonhosted.org/packages/4c/8e/caa83237f427d9e85b7f02c816e7270c9c9571dec1673e06b0180402f70e/lupa-2.8-cp311-cp311-win_amd64.whl", hash = "sha256:f5a6af145b0ea818f01d27bfe2583a4b538570bef61d22c8773e0eccf011234c", size = 1923488, upload-time = "2026-04-15T20:05:52.954Z" }, + { url = "https://files.pythonhosted.org/packages/ad/0b/368f2f0bc750b25c69d4563e44f677925ab5dd3d2887f9b0c15465d21a2a/lupa-2.8-cp312-abi3-macosx_10_13_x86_64.whl", hash = "sha256:f4342f4de76ae7ce2ab0672d36003bdb7e1a33252f293b569298ddd792e70e33", size = 1194056, upload-time = "2026-04-15T20:05:55.794Z" }, + { url = "https://files.pythonhosted.org/packages/5b/0f/c89eb8dd36fdea4e50ae3f7f5275bea3b0cc5d4057b8ee7b3bbc78010422/lupa-2.8-cp312-abi3-manylinux2010_i686.manylinux_2_12_i686.manylinux_2_28_i686.whl", hash = "sha256:4203fa1659315e939a5304e75001b8cc14234fb3cbb3ed86c049b0cc5d90fcee", size = 1434278, upload-time = "2026-04-15T20:05:57.94Z" }, + { url = "https://files.pythonhosted.org/packages/47/30/c3b4d2cd8733621b404b8a4214e5f852955c4ba632546dc84123bea9ee89/lupa-2.8-cp312-abi3-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:81f2d843ce668b653146c007467570210ae44be51dac6926666c51d49536f307", size = 1150068, upload-time = "2026-04-15T20:06:01.04Z" }, + { url = "https://files.pythonhosted.org/packages/8d/d2/bac12c398519efafc6af84be1974edd0d7a4895fb4735b5c8d615d298595/lupa-2.8-cp312-abi3-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:d3d0cde2c77588d1c60875a4f34f059513476c6e1775351897195b51e0f3df08", size = 1409532, upload-time = "2026-04-15T20:06:03.592Z" }, + { url = "https://files.pythonhosted.org/packages/9c/6a/18b52e11962014026e07813530b0b108ee8bc0a2a13ef0eaea5d41dce023/lupa-2.8-cp312-abi3-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:9e0d11b8f3a8dac6413f704fef7161d048bb10c58bdac6cbffa5e60efa56e9a3", size = 1242687, upload-time = "2026-04-15T20:06:06.863Z" }, + { url = "https://files.pythonhosted.org/packages/b3/8e/7fd4eb049875f61429b96780d2eae4700f0e78fe0a52db8edb231b1cd09f/lupa-2.8-cp312-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:54cff414f21f8cd8c6be4aae52541f3b9cd39602b59e3a3db9b5c9f9f674ff18", size = 1856038, upload-time = "2026-04-15T20:06:09.358Z" }, + { url = "https://files.pythonhosted.org/packages/e9/f9/37ad9d2773d30f2931890d310a4bdce28d45484206e6f48bc18b0325eabd/lupa-2.8-cp312-abi3-musllinux_1_2_armv7l.whl", hash = "sha256:24b4d8af5558e549b70daf1547f5c1c1d664ecea9fc790f83efe5d75e9a93797", size = 1128982, upload-time = "2026-04-15T20:06:12.312Z" }, + { url = "https://files.pythonhosted.org/packages/57/31/c0fd7984c24844ea79caa45c0235f61a06b38fd69a839f6c62770f8d684a/lupa-2.8-cp312-abi3-musllinux_1_2_i686.whl", hash = "sha256:ce86dff1ee7f7cf45f5622065ae991949dd7bb1703581cbc58a630137bb7ccf9", size = 1457594, upload-time = "2026-04-15T20:06:15.881Z" }, + { url = "https://files.pythonhosted.org/packages/11/f5/a28e411be30ec1bf0db1eb0c087eebc73be9e7a1adcfe6ac209861ccc446/lupa-2.8-cp312-abi3-musllinux_1_2_ppc64le.whl", hash = "sha256:f4d01b2a08c70bbb883a9e082b6b36b89121ed5910b710f1ba11c73295ff4fba", size = 1425721, upload-time = "2026-04-15T20:06:18.009Z" }, + { url = "https://files.pythonhosted.org/packages/ed/c1/359f767c4ae024be30d909fe8a9f0e9af266bad47ce2bd2ed248fb986fcf/lupa-2.8-cp312-abi3-musllinux_1_2_riscv64.whl", hash = "sha256:7f210d5a8353e510ea1199c42cf3cbdd630553bf2bc8fb4c00fea06fdec7c798", size = 1253258, upload-time = "2026-04-15T20:06:21.17Z" }, + { url = "https://files.pythonhosted.org/packages/17/52/473f11790c261fd02bbf318a546fe040e9ec9f677181272fa78d3b4112a4/lupa-2.8-cp312-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:4f81a02806e7c7ad26d8c6fa222c8bef1b0c1b124347c879be880b41339d41e4", size = 2395272, upload-time = "2026-04-15T20:06:24.137Z" }, + { url = "https://files.pythonhosted.org/packages/94/bf/75c8795655a8836eab6a11a630352c4b7c5dc5c54d075077bc9bffdeee45/lupa-2.8-cp312-abi3-win32.whl", hash = "sha256:360056453a7a4eaa4ac5a204c31a5a014b1eb2ee5490603234d2ba831684f1f2", size = 1606136, upload-time = "2026-04-15T20:06:27.815Z" }, + { url = "https://files.pythonhosted.org/packages/d8/29/11a2cdd612b6f55e506292dfb6ba343216e80a693e7fe3f876ef204ce9c6/lupa-2.8-cp312-abi3-win_arm64.whl", hash = "sha256:1628371c6592a6d5650497a9e31fb2bb3a7e9883c1f301d1111265e484045af9", size = 1364495, upload-time = "2026-04-15T20:06:30.254Z" }, + { url = "https://files.pythonhosted.org/packages/4d/17/fa834b6b09ad17e7df5d0f7715d64877a125a3776ada689751a1f9dc2959/lupa-2.8-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:450650f91c48c2415b0d59ab3abfcfda3b6efb5b858205f4d4bda8ad141fa529", size = 1190111, upload-time = "2026-04-15T20:06:32.84Z" }, + { url = "https://files.pythonhosted.org/packages/ab/43/45589901b7d1a0e3a9d91d19a311fb6a56924e8571536c3f2212160fd953/lupa-2.8-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:27044f3363047f946b3d3aab9157cbd172b3538ada9ec1baef43432bf7d03a78", size = 1812999, upload-time = "2026-04-15T20:06:35.664Z" }, + { url = "https://files.pythonhosted.org/packages/a1/ac/4ade7d15ff5c61758d7943ac6f0a496bf1cc65b6c09f842b52a0702e664c/lupa-2.8-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:8cf4f064a0e5531afce2d7d750120c10c10f9529139af6ca6150d13151034398", size = 2368731, upload-time = "2026-04-15T20:06:37.959Z" }, + { url = "https://files.pythonhosted.org/packages/0c/27/05f950d15b8ab120b39c43588b438ff3ace70c1b1b0225a960393a497483/lupa-2.8-cp312-cp312-win_amd64.whl", hash = "sha256:281bedc5deb92d31e649a3552edd662449365a635904fa4d5cb4509c7245e34e", size = 1941809, upload-time = "2026-04-15T20:06:40.302Z" }, + { url = "https://files.pythonhosted.org/packages/a6/3f/19f83c3a0c84dc8bea8a58e7416dca6a3ede662c33c8d1ec758e5afc754a/lupa-2.8-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:45fc9da0145ecb0083ef5ff9975116cc784bd0258bdc2bd131ba15483ce18398", size = 1201203, upload-time = "2026-04-15T20:06:42.169Z" }, + { url = "https://files.pythonhosted.org/packages/89/0f/a14f0073f09610158038582e230618a48c14da6bd88185289461aa4cb854/lupa-2.8-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:58e18afed57955b41130e269c78f53d4123ab86e236b53816f4cbffa25cb5d30", size = 1806210, upload-time = "2026-04-15T20:06:45.486Z" }, + { url = "https://files.pythonhosted.org/packages/2f/14/48fff156c63a136001a7620878af7d31aa07e66b495ed621e3eddd73c294/lupa-2.8-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:fc47f536ac13a79cef47d29a2b205576a22841f042a2bcec1676b95806e7706a", size = 2359005, upload-time = "2026-04-15T20:06:47.819Z" }, + { url = "https://files.pythonhosted.org/packages/fe/18/3ac638ec90edf178242b8a2b2f00f8adae694248c03a26341ef941bb746e/lupa-2.8-cp313-cp313-win_amd64.whl", hash = "sha256:ce9404c661dbac65cc9bed351ad45e797af93d30d70be309a3fa8209ac86d93b", size = 1936754, upload-time = "2026-04-15T20:06:50.448Z" }, + { url = "https://files.pythonhosted.org/packages/b0/ef/5ee5fed6ea7459a671196359ce04bfeeaf26be1dac8ff24bf28e5c7a6e81/lupa-2.8-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:348c3f8ecabb6324dcbc05c2740d762ef8fcec7b06c79e45262ab97a217684e3", size = 1209388, upload-time = "2026-04-15T20:06:53.022Z" }, + { url = "https://files.pythonhosted.org/packages/6e/b1/67a940d5542cb0384b443fe951b5a83ea9340d1333a733a258fdd1c619ba/lupa-2.8-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:951496471056061598a7d1729a6cdf48d662fec777a9f2d8aa5a1e62fd30e5a5", size = 1826821, upload-time = "2026-04-15T20:06:55.699Z" }, + { url = "https://files.pythonhosted.org/packages/a1/a2/b354e5ba3b911ec50686003dc8897e892b9e8c5c036b33219b03d54c4daf/lupa-2.8-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:a591b9947ca347b41a63370e121d6e2b1458fe6dde9ae065029ec10a37f25ff4", size = 2366893, upload-time = "2026-04-15T20:06:58.9Z" }, + { url = "https://files.pythonhosted.org/packages/8e/52/d76066401f29539df5352f70ecded66576f32933b6045cd0bfc56cb770b9/lupa-2.8-cp314-cp314-win_amd64.whl", hash = "sha256:3903c9cf628dae2f56405503247b77a61a3a61bd2dda470e336950c74776d55d", size = 1994716, upload-time = "2026-04-15T20:07:19.194Z" }, + { url = "https://files.pythonhosted.org/packages/c3/bd/3efc437a4361c16d25e66478c50357c9a8e8ecfb718fe749eb9ca3176ef6/lupa-2.8-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:f711a8ab0486b9ac6fdda94a22ddcfbc9f0d4a27e3a8cf1bf79c6e48b33017c1", size = 1251217, upload-time = "2026-04-15T20:07:01.64Z" }, + { url = "https://files.pythonhosted.org/packages/ea/f4/2e9f8ecbaca854bfdf14af8a9b505ec0cbc640377b3b218921594b7563cd/lupa-2.8-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:dc51250e76367a3e27fcd01dc769b9bfcbbc34f48df48dde53d6af6e75b7eaa5", size = 1814701, upload-time = "2026-04-15T20:07:04.149Z" }, + { url = "https://files.pythonhosted.org/packages/ba/53/4000b1acaa8b1f3827fcff0cfcdff44d3befddda42cab7e685a49689b5a1/lupa-2.8-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f8a22088a552828958603323f0a5c4b3e11e03b75d0bf4c965ef879de9b60a8d", size = 2348414, upload-time = "2026-04-15T20:07:07.285Z" }, + { url = "https://files.pythonhosted.org/packages/d5/78/26ee48d3890cddf03cefb65f433e3492759c0b3c0582180755bddbaab7bd/lupa-2.8-cp314-cp314t-win32.whl", hash = "sha256:4f7c553c1d8cfffbe85d81daef730d12cae4b6002d457542914da0ac8a1145b3", size = 1831611, upload-time = "2026-04-15T20:07:09.752Z" }, + { url = "https://files.pythonhosted.org/packages/3c/d1/4a5cc64a3cad22821ae4c3f7a90456a08ca19457d8354f4abf46ad03c7e8/lupa-2.8-cp314-cp314t-win_amd64.whl", hash = "sha256:d8766aff03a78c80ad2d188a8bdb216de5ec838359cd87e05bbdfa56394a6105", size = 2209250, upload-time = "2026-04-15T20:07:11.906Z" }, + { url = "https://files.pythonhosted.org/packages/37/7c/cdcb654daf668192aaf36b0aeb94f2281dad092aaa5003688691131736ea/lupa-2.8-cp314-cp314t-win_arm64.whl", hash = "sha256:91d622777febda3ab1bed1d45295f2f32a4680c7b3d7caf8c669998ed5c44118", size = 1126735, upload-time = "2026-04-15T20:07:15.434Z" }, + { url = "https://files.pythonhosted.org/packages/1d/44/de1961ad38e17cd326a53c246c7e3b91178ed578f4cf22ffcd5e7e11b041/lupa-2.8-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:b036738282a5acd2e71fdddb317c9df8b87c1673aa57f403d05fcc2be8abc4ba", size = 1186020, upload-time = "2026-04-15T20:07:35.017Z" }, + { url = "https://files.pythonhosted.org/packages/13/c2/276f0b9dc8bcc5a8a58af5316dfa0e6f56be3613dd6dbcc8d3d2cb6559ba/lupa-2.8-cp39-abi3-manylinux2010_i686.manylinux_2_12_i686.manylinux_2_28_i686.whl", hash = "sha256:ac6b6e8d0e617e26a98cbb44880bcd75de5d32b3ad7b3b3793583909292b47ed", size = 1468944, upload-time = "2026-04-15T20:07:37.782Z" }, + { url = "https://files.pythonhosted.org/packages/63/38/52934e52a5180dc6425d20284d004fe4b27a4f9171a82dc99fb67af250bf/lupa-2.8-cp39-abi3-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:ba3a7dd839f90c3d2e53bebe3c192b1f3f9fd720a6781256405123211fd0dce6", size = 1172998, upload-time = "2026-04-15T20:07:40.812Z" }, + { url = "https://files.pythonhosted.org/packages/c7/82/76b3809bd0839d9b3b4ec58d06591e08f17337b6d9576877cb9d48b34e94/lupa-2.8-cp39-abi3-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:d7edb13a7a5250b5c6c22d1495d9e842b5c9fc5081c8fe6b5efe2112fe3e41f9", size = 1449975, upload-time = "2026-04-15T20:07:44.262Z" }, + { url = "https://files.pythonhosted.org/packages/16/07/2f89d54f747c67c23b4b9ae4aa8c8dd06bb409155dedcf406157f2736b66/lupa-2.8-cp39-abi3-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:891f72e0bffbed1e4175f975aeb2a083956586a100066525e1be485f617f7b25", size = 1281944, upload-time = "2026-04-15T20:07:46.458Z" }, + { url = "https://files.pythonhosted.org/packages/e7/bd/7375d2b0fcae79d806baf52a76f26c96964593f58e1372d13ae5ac09c676/lupa-2.8-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:a295f87b5b7ebbfd5191932e8cb0e51df3c7769101ac6b6c7d7c9fb27bfd1307", size = 1910455, upload-time = "2026-04-15T20:07:49.75Z" }, + { url = "https://files.pythonhosted.org/packages/8b/0c/8abb3bc0e08b311fc01db05b6e9f9ff31a8f65e4fc3f0aeb05cfef75c8ac/lupa-2.8-cp39-abi3-musllinux_1_2_armv7l.whl", hash = "sha256:4fe5d7a810b64ea8511eb885fc8cdde042ee5ff7b7d08ae78f32449756acb177", size = 1155548, upload-time = "2026-04-15T20:07:52.657Z" }, + { url = "https://files.pythonhosted.org/packages/80/2e/9eeecd3f493099721c1d3f31beeca23a4237db1a54223684df4dc96aa1bd/lupa-2.8-cp39-abi3-musllinux_1_2_i686.whl", hash = "sha256:bfc470012ef66ad064c7bd77416af03a3452ef630b04b9012595ea13f2e54518", size = 1489232, upload-time = "2026-04-15T20:07:54.92Z" }, + { url = "https://files.pythonhosted.org/packages/c3/13/731c99dc2e7652ae818a6de45bdf0142049f7cb566049061c898355f1891/lupa-2.8-cp39-abi3-musllinux_1_2_ppc64le.whl", hash = "sha256:250e035fdaffe8c87093e3ebc206ac29a26131b1568ea711d780c26001ce96e7", size = 1466321, upload-time = "2026-04-15T20:07:57.627Z" }, + { url = "https://files.pythonhosted.org/packages/de/71/3ad8cc4fc05a77dc0d3f7079348bd1cad4675a0d14c24f8e6a3ce5f008f7/lupa-2.8-cp39-abi3-musllinux_1_2_riscv64.whl", hash = "sha256:b9bddb09acfffb4f828f790f444b11dc0cca591afea1a244d9329eea2d20c003", size = 1288577, upload-time = "2026-04-15T20:07:59.913Z" }, + { url = "https://files.pythonhosted.org/packages/d8/b2/1175f6d0aa7b68627fbe2f58bd1e8bea36a89d10dfd67671d2b024c96162/lupa-2.8-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:2e64acbbd47e9b82a64405a39e0d2b36a5a7dad8ab41c0f3437f572f7d282ba3", size = 2444866, upload-time = "2026-04-15T20:08:02.753Z" }, + { url = "https://files.pythonhosted.org/packages/92/f7/e78df680c7a0ea452daac07467ca188d63c2c00ca1c884c0a50e27eb83b5/lupa-2.8-pp311-pypy311_pp73-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:32e4e5103bbddcdd2458fb2ccae6c8ba11c9997c711d7e379e0d45551d109c76", size = 1778509, upload-time = "2026-04-15T20:08:21.784Z" }, + { url = "https://files.pythonhosted.org/packages/e6/23/0e53cabb16b2a8aa9cf1fde499c097d8942c5dab709fc8e921f3b824b18b/lupa-2.8-pp311-pypy311_pp73-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:7667001804657496dee9feced2daae5000b4604a3218dd8e6b7b754982ba88b8", size = 2300480, upload-time = "2026-04-15T20:08:24.394Z" }, + { url = "https://files.pythonhosted.org/packages/7e/85/0271227eab939921a12ebba5d17aa4cd18346aa534ca7f5da09cd0b63dd4/lupa-2.8-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:86f6f668966965b15247dc32d064cfe7be67b71e584ccfacbe2f637575296878", size = 1847445, upload-time = "2026-04-15T20:08:27.031Z" }, +] + [[package]] name = "mako" version = "1.3.12" From 45cd83c9c6d5c0158fe9cec55931e84bc43f566f Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 11 Sep 2026 11:58:29 +0200 Subject: [PATCH 35/90] fix(realtime): normalize offer media types and close parsed forms --- .../proxy/realtime_endpoints/call_sessions.py | 20 ++++- litellm/proxy/realtime_endpoints/endpoints.py | 7 +- .../realtime_endpoints/test_call_sessions.py | 84 ++++++++++++++++++- 3 files changed, 102 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 8ea329945c5..c0aa0dcbe52 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -12,6 +12,7 @@ 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 from litellm._logging import verbose_proxy_logger @@ -44,6 +45,9 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth, ) from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.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 ) @@ -209,7 +213,14 @@ def decode_call(token: str, authorization: str) -> CodexRealtimeCall: async def read_codex_offer(request: Request) -> CodexRealtimeOffer: - if request.headers.get("content-type", "").startswith("multipart/form-data"): + 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 + 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", "{}")))}) @@ -257,8 +268,11 @@ async def process_codex_request( async def create_codex_realtime_call(request: Request) -> Response: - with isolated_request_stash(): - return await _create_codex_realtime_call(request) + 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: diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index 32c6abdb7b6..82041b4274f 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, @@ -375,7 +378,7 @@ async def proxy_realtime_calls( request: Request, fastapi_response: Response, ) -> Response: - if request.headers.get("content-type", "").split(";", 1)[0] in ("application/json", "multipart/form-data"): + 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) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index d65556ebd39..d730add23b2 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -11,15 +11,75 @@ 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("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"]) -async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypatch, multipart, policy): +@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 + from fastapi import Request, Response import litellm from litellm.exceptions import BudgetExceededError @@ -27,6 +87,7 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa 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 = ( @@ -36,8 +97,16 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa ) 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( @@ -48,7 +117,7 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa "query_string": b"model=query-decoy&policy=keep", "client": ("127.0.0.7", 1234), "headers": [ - *((key.lower(), value) for key, value in outbound.headers.raw), + (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"), @@ -59,12 +128,17 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa 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 @@ -95,8 +169,10 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa 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 codex.create_codex_realtime_call(request) + await proxy_realtime_calls(request, Response()) if policy == "personal_models": assert "user not allowed to access model" in str(denied.value) assert "forbidden-voice" in str(denied.value) From 7d39e9802e98741871882d1ebd6711e70a008281 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 11 Sep 2026 13:26:31 +0200 Subject: [PATCH 36/90] fix(realtime): limit ChatGPT header access to its provider --- litellm/realtime_api/main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 0b8a32dcfc1..dd9b5083ee4 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -411,7 +411,7 @@ async def _arealtime( ProviderConfigManager.get_provider_realtime_handler( LlmProviders(_custom_llm_provider), litellm_params, websocket.headers, headers ) - if _custom_llm_provider in LlmProviders._member_map_.values() + if _custom_llm_provider == LlmProviders.CHATGPT else None ) if provider_handler is not None: From dbbc6b84c220ed353565cf462a7887f70ed185b1 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 11 Sep 2026 14:09:14 +0200 Subject: [PATCH 37/90] fix(realtime): defer provider headers and test Lua with Redis --- .github/workflows/_test-unit-base.yml | 11 ++ .github/workflows/test-unit.yml | 1 + litellm/realtime_api/main.py | 4 +- litellm/utils.py | 4 +- pyproject.toml | 2 +- .../local_testing/test_realtime_call_redis.py | 102 ++++++++++++++++++ .../hooks/test_parallel_request_limiter.py | 47 -------- .../hooks/test_parallel_request_limiter_v3.py | 22 ---- tests/test_litellm/realtime_api/test_main.py | 37 +++++++ uv.lock | 72 +------------ 10 files changed, 158 insertions(+), 144 deletions(-) create mode 100644 tests/local_testing/test_realtime_call_redis.py diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index 62790e23143..4d939b3e838 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -68,6 +68,16 @@ jobs: pull-requests: read outputs: decision: ${{ steps.changes.outputs.decision }} + 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 @@ -137,6 +147,7 @@ jobs: RERUNS: ${{ inputs.reruns }} DIST: ${{ inputs.dist }} COVERAGE_CORE: sysmon + LITELLM_TEST_REDIS_PORT: ${{ job.services.redis.ports['6379'] }} run: | if [ "${WORKERS}" = "0" ]; then uv run --no-sync pytest ${TEST_PATH:?} \ diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index f55c87c2ae5..cdbbe28f1cd 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -133,6 +133,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/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index dd9b5083ee4..b4dc87b77ad 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -409,9 +409,9 @@ async def _arealtime( ) provider_handler: Final = ( ProviderConfigManager.get_provider_realtime_handler( - LlmProviders(_custom_llm_provider), litellm_params, websocket.headers, headers + LlmProviders(_custom_llm_provider), litellm_params, lambda: websocket.headers, headers ) - if _custom_llm_provider == LlmProviders.CHATGPT + if _custom_llm_provider in LlmProviders._member_map_.values() else None ) if provider_handler is not None: diff --git a/litellm/utils.py b/litellm/utils.py index fa07897fa8f..7075d8776c3 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9282,13 +9282,13 @@ class ProviderConfigManager: def get_provider_realtime_handler( provider: LlmProviders, params: GenericLiteLLMParams, - headers: Mapping[str, str], + 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, headers, extra_headers) + return ChatGPTRealtime(params, get_headers(), extra_headers) return None @staticmethod diff --git a/pyproject.toml b/pyproject.toml index b9db915a890..d33d693f794 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -207,7 +207,7 @@ dev = [ "opentelemetry-instrumentation-fastapi==0.49b0", "langfuse==2.59.7", "fastapi-offline==1.7.6", - "fakeredis[lua]==2.34.1", + "fakeredis==2.34.1", "pytest-rerunfailures==15.1", "pytest-cov==5.0.0", "parameterized==0.9.0", 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/hooks/test_parallel_request_limiter.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py index 360da3fefd6..589c8d1bd06 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -7,8 +7,6 @@ from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch import pytest -import pytest_asyncio -from fakeredis import FakeAsyncRedis, FakeRedis, FakeServer from litellm.caching.caching import DualCache from litellm.caching.redis_cache import RedisCache @@ -21,51 +19,6 @@ from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage -@pytest_asyncio.fixture(loop_scope="function") -async def isolated_legacy_redis(): - server = FakeServer() - client = FakeRedis(server=server) - async with FakeAsyncRedis(server=server) as async_client: - with ( - patch("redis.Redis", autospec=True, return_value=client), - patch("redis.asyncio.BlockingConnectionPool", autospec=True, return_value=async_client.connection_pool), - patch("redis.asyncio.Redis", autospec=True, return_value=async_client), - ): - yield RedisCache(host="fake-legacy-redis", namespace="legacy-test") - client.close() - - -@pytest.mark.asyncio -async def test_concurrent_realtime_releases_update_redis_without_lost_decrement(isolated_legacy_redis): - remote = isolated_legacy_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_release_preserves_newer_local_admission_while_redis_finishes(): started, finish = asyncio.Event(), asyncio.Event() 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 017ab26fcd4..14d078820af 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 @@ -141,28 +141,6 @@ async def test_realtime_lease_renewal_preserves_quota_past_ttl_and_does_not_resu await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") -@pytest.mark.asyncio -async def test_realtime_lease_redis_renewal_is_atomic_and_does_not_resurrect(): - from fakeredis import FakeAsyncRedis - from litellm.proxy.hooks.parallel_request_limiter_v3 import PARALLEL_RENEW_SCRIPT - - async with FakeAsyncRedis() as client: - now = (await client.time())[0] - await client.zadd("first", {"owner": now - 10, "other": now}) - await client.zadd("second", {"owner": now - PARALLEL_REQUEST_SLOT_TTL_SECONDS}) - renew = client.register_script(PARALLEL_RENEW_SCRIPT) - assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] - assert await client.zscore("first", "owner") == now - 10 - await client.zadd("second", {"owner": now - 10}) - assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [1] - assert await client.zscore("first", "owner") >= now - assert await client.ttl("first") > PARALLEL_REQUEST_SLOT_TTL_SECONDS - 10 - await client.zrem("second", "owner") - assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] - assert await client.zscore("second", "owner") is None - assert await client.zscore("first", "other") == now - - @pytest.fixture def time_controller(monkeypatch): controller = TimeController() diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index 3ab843dcea7..a4e8204c86f 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -1,4 +1,5 @@ import asyncio +import json import time from types import TracebackType from typing import Final @@ -17,6 +18,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"} diff --git a/uv.lock b/uv.lock index 5ef4b4ccb5d..cedc505acdf 100644 --- a/uv.lock +++ b/uv.lock @@ -1875,11 +1875,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/49/b5/82f89307d0d769cd9bf46a54fb9136be08e4e57c5570ae421db4c9a2ba62/fakeredis-2.34.1-py3-none-any.whl", hash = "sha256:0107ec99d48913e7eec2a5e3e2403d1bd5f8aa6489d1a634571b975289c48f12", size = 122160, upload-time = "2026-02-25T13:17:49.701Z" }, ] -[package.optional-dependencies] -lua = [ - { name = "lupa" }, -] - [[package]] name = "fastapi" version = "0.136.3" @@ -4525,7 +4520,7 @@ dev = [ { name = "basedpyright" }, { name = "botocore-stubs" }, { name = "diff-cover" }, - { name = "fakeredis", extra = ["lua"] }, + { name = "fakeredis" }, { name = "fastapi-offline" }, { name = "hypothesis" }, { name = "keyring" }, @@ -4713,7 +4708,7 @@ dev = [ { name = "basedpyright", specifier = "==1.39.7" }, { name = "botocore-stubs", specifier = "==1.43.14" }, { name = "diff-cover", specifier = "==9.7.2" }, - { name = "fakeredis", extras = ["lua"], specifier = "==2.34.1" }, + { name = "fakeredis", specifier = "==2.34.1" }, { name = "fastapi-offline", specifier = "==1.7.6" }, { name = "hypothesis", specifier = "==6.165.10" }, { name = "keyring", specifier = "==25.7.0" }, @@ -5008,69 +5003,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/db/a4/441aee36c6f6b249823d20fd91f9be9ab89d7c5a8ae542a4a4ca6d342d56/lxml-6.1.1-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:ed21202aec73cda4d55d1ce57b389aadb90ffb044e6cd1080b8347efe1b1ec84", size = 3508989, upload-time = "2026-05-18T19:18:38.158Z" }, ] -[[package]] -name = "lupa" -version = "2.8" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/c3/a6/0f869fbb07c393f15473b1eefefb7b5bec162fb7481803d040ed4dc46002/lupa-2.8.tar.gz", hash = "sha256:d8022641b9ec8ecf2c5ecbe9f47e5a70e0b87c4b5ae921b92cb02a638e0acd08", size = 6156370, upload-time = "2026-04-15T20:08:30.534Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/09/21/9be4516ddd22f8eadba336d9ba065d17d79108465ae1b7f71424ab99b9d0/lupa-2.8-cp310-abi3-win32.whl", hash = "sha256:c2a5fd15dc62374e1661a55f01744c9ec1c56f291ba4a0749d3af2174556e78f", size = 1594887, upload-time = "2026-04-15T20:05:23.377Z" }, - { url = "https://files.pythonhosted.org/packages/2d/99/1557c9685d7034d9ce8dd2b54c40a26d6deb7c67c1fdb5c801abd1a02c3f/lupa-2.8-cp310-abi3-win_arm64.whl", hash = "sha256:9e304fb1c50cf23fd8882afbe1aa87525ef8a72667bcab3b37b2bbb2bc542269", size = 1371742, upload-time = "2026-04-15T20:05:27.417Z" }, - { url = "https://files.pythonhosted.org/packages/1c/34/05ce4745b191633f90ff1ab50f1a19a37da282bb0a41fb500d9157fc9b8f/lupa-2.8-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:97bd01e90b8031e56a5fd5bb70605aea09f1dba675c1140308a52780f93d06f1", size = 1202714, upload-time = "2026-04-15T20:05:31.088Z" }, - { url = "https://files.pythonhosted.org/packages/7d/d2/f70fdbeec2d4c69ee6a469e6cddde9635fff4af4e13fb652e6a1229eef51/lupa-2.8-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:0b5ebe1a13c45767919c86750b84fe2da9f6288b6f3cea4ce7660bb2abc9d921", size = 1857453, upload-time = "2026-04-15T20:05:34.611Z" }, - { url = "https://files.pythonhosted.org/packages/97/dc/6fcda0e36e75eb6cb98dc9190fa4737d727eeae29e58f892980b2c96b656/lupa-2.8-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:097e7d0f1719a88020b67c82e05d53d7973c166952393afcecfd8434c7e19a15", size = 2408890, upload-time = "2026-04-15T20:05:37.994Z" }, - { url = "https://files.pythonhosted.org/packages/58/29/7ea176eac3c1dac83d059762daa875ad1390decc0bf2c3b4c7bbfc1f1665/lupa-2.8-cp310-cp310-win_amd64.whl", hash = "sha256:7bb223ee8f72d0dc076b0d65296ee72f1c69450f9d2fed5315f7707d98c4a03d", size = 1910396, upload-time = "2026-04-15T20:05:41.163Z" }, - { url = "https://files.pythonhosted.org/packages/b7/0a/5a740717f27aa77481e6a61b97cf79d1e0c1ede729b1268caacded915326/lupa-2.8-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:b12e43c1fb787189dfc28cd604aef0baa2cb95e27da19498d520361d0ace070a", size = 1202376, upload-time = "2026-04-15T20:05:44.049Z" }, - { url = "https://files.pythonhosted.org/packages/1b/75/6b64d0098c64275a801896cb7a6a30e7e653d25fa102c64e747292afcdbb/lupa-2.8-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:f6f603391dffb256e36a79fd2044084d5f4b8a0a4c0e5ad291cd3ab3aaf1fd0a", size = 1839271, upload-time = "2026-04-15T20:05:47.399Z" }, - { url = "https://files.pythonhosted.org/packages/7b/2f/0d4f00563046ff616ef6a421f8b776a5ffb327f7b32ed69e856d52b917a8/lupa-2.8-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:9f6f41c91366e7d0d474f87d81c1274af861f40812bf729c9f97ab4c8f3c7ac8", size = 2376251, upload-time = "2026-04-15T20:05:49.891Z" }, - { url = "https://files.pythonhosted.org/packages/4c/8e/caa83237f427d9e85b7f02c816e7270c9c9571dec1673e06b0180402f70e/lupa-2.8-cp311-cp311-win_amd64.whl", hash = "sha256:f5a6af145b0ea818f01d27bfe2583a4b538570bef61d22c8773e0eccf011234c", size = 1923488, upload-time = "2026-04-15T20:05:52.954Z" }, - { url = "https://files.pythonhosted.org/packages/ad/0b/368f2f0bc750b25c69d4563e44f677925ab5dd3d2887f9b0c15465d21a2a/lupa-2.8-cp312-abi3-macosx_10_13_x86_64.whl", hash = "sha256:f4342f4de76ae7ce2ab0672d36003bdb7e1a33252f293b569298ddd792e70e33", size = 1194056, upload-time = "2026-04-15T20:05:55.794Z" }, - { url = "https://files.pythonhosted.org/packages/5b/0f/c89eb8dd36fdea4e50ae3f7f5275bea3b0cc5d4057b8ee7b3bbc78010422/lupa-2.8-cp312-abi3-manylinux2010_i686.manylinux_2_12_i686.manylinux_2_28_i686.whl", hash = "sha256:4203fa1659315e939a5304e75001b8cc14234fb3cbb3ed86c049b0cc5d90fcee", size = 1434278, upload-time = "2026-04-15T20:05:57.94Z" }, - { url = "https://files.pythonhosted.org/packages/47/30/c3b4d2cd8733621b404b8a4214e5f852955c4ba632546dc84123bea9ee89/lupa-2.8-cp312-abi3-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:81f2d843ce668b653146c007467570210ae44be51dac6926666c51d49536f307", size = 1150068, upload-time = "2026-04-15T20:06:01.04Z" }, - { url = "https://files.pythonhosted.org/packages/8d/d2/bac12c398519efafc6af84be1974edd0d7a4895fb4735b5c8d615d298595/lupa-2.8-cp312-abi3-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:d3d0cde2c77588d1c60875a4f34f059513476c6e1775351897195b51e0f3df08", size = 1409532, upload-time = "2026-04-15T20:06:03.592Z" }, - { url = "https://files.pythonhosted.org/packages/9c/6a/18b52e11962014026e07813530b0b108ee8bc0a2a13ef0eaea5d41dce023/lupa-2.8-cp312-abi3-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:9e0d11b8f3a8dac6413f704fef7161d048bb10c58bdac6cbffa5e60efa56e9a3", size = 1242687, upload-time = "2026-04-15T20:06:06.863Z" }, - { url = "https://files.pythonhosted.org/packages/b3/8e/7fd4eb049875f61429b96780d2eae4700f0e78fe0a52db8edb231b1cd09f/lupa-2.8-cp312-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:54cff414f21f8cd8c6be4aae52541f3b9cd39602b59e3a3db9b5c9f9f674ff18", size = 1856038, upload-time = "2026-04-15T20:06:09.358Z" }, - { url = "https://files.pythonhosted.org/packages/e9/f9/37ad9d2773d30f2931890d310a4bdce28d45484206e6f48bc18b0325eabd/lupa-2.8-cp312-abi3-musllinux_1_2_armv7l.whl", hash = "sha256:24b4d8af5558e549b70daf1547f5c1c1d664ecea9fc790f83efe5d75e9a93797", size = 1128982, upload-time = "2026-04-15T20:06:12.312Z" }, - { url = "https://files.pythonhosted.org/packages/57/31/c0fd7984c24844ea79caa45c0235f61a06b38fd69a839f6c62770f8d684a/lupa-2.8-cp312-abi3-musllinux_1_2_i686.whl", hash = "sha256:ce86dff1ee7f7cf45f5622065ae991949dd7bb1703581cbc58a630137bb7ccf9", size = 1457594, upload-time = "2026-04-15T20:06:15.881Z" }, - { url = "https://files.pythonhosted.org/packages/11/f5/a28e411be30ec1bf0db1eb0c087eebc73be9e7a1adcfe6ac209861ccc446/lupa-2.8-cp312-abi3-musllinux_1_2_ppc64le.whl", hash = "sha256:f4d01b2a08c70bbb883a9e082b6b36b89121ed5910b710f1ba11c73295ff4fba", size = 1425721, upload-time = "2026-04-15T20:06:18.009Z" }, - { url = "https://files.pythonhosted.org/packages/ed/c1/359f767c4ae024be30d909fe8a9f0e9af266bad47ce2bd2ed248fb986fcf/lupa-2.8-cp312-abi3-musllinux_1_2_riscv64.whl", hash = "sha256:7f210d5a8353e510ea1199c42cf3cbdd630553bf2bc8fb4c00fea06fdec7c798", size = 1253258, upload-time = "2026-04-15T20:06:21.17Z" }, - { url = "https://files.pythonhosted.org/packages/17/52/473f11790c261fd02bbf318a546fe040e9ec9f677181272fa78d3b4112a4/lupa-2.8-cp312-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:4f81a02806e7c7ad26d8c6fa222c8bef1b0c1b124347c879be880b41339d41e4", size = 2395272, upload-time = "2026-04-15T20:06:24.137Z" }, - { url = "https://files.pythonhosted.org/packages/94/bf/75c8795655a8836eab6a11a630352c4b7c5dc5c54d075077bc9bffdeee45/lupa-2.8-cp312-abi3-win32.whl", hash = "sha256:360056453a7a4eaa4ac5a204c31a5a014b1eb2ee5490603234d2ba831684f1f2", size = 1606136, upload-time = "2026-04-15T20:06:27.815Z" }, - { url = "https://files.pythonhosted.org/packages/d8/29/11a2cdd612b6f55e506292dfb6ba343216e80a693e7fe3f876ef204ce9c6/lupa-2.8-cp312-abi3-win_arm64.whl", hash = "sha256:1628371c6592a6d5650497a9e31fb2bb3a7e9883c1f301d1111265e484045af9", size = 1364495, upload-time = "2026-04-15T20:06:30.254Z" }, - { url = "https://files.pythonhosted.org/packages/4d/17/fa834b6b09ad17e7df5d0f7715d64877a125a3776ada689751a1f9dc2959/lupa-2.8-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:450650f91c48c2415b0d59ab3abfcfda3b6efb5b858205f4d4bda8ad141fa529", size = 1190111, upload-time = "2026-04-15T20:06:32.84Z" }, - { url = "https://files.pythonhosted.org/packages/ab/43/45589901b7d1a0e3a9d91d19a311fb6a56924e8571536c3f2212160fd953/lupa-2.8-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:27044f3363047f946b3d3aab9157cbd172b3538ada9ec1baef43432bf7d03a78", size = 1812999, upload-time = "2026-04-15T20:06:35.664Z" }, - { url = "https://files.pythonhosted.org/packages/a1/ac/4ade7d15ff5c61758d7943ac6f0a496bf1cc65b6c09f842b52a0702e664c/lupa-2.8-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:8cf4f064a0e5531afce2d7d750120c10c10f9529139af6ca6150d13151034398", size = 2368731, upload-time = "2026-04-15T20:06:37.959Z" }, - { url = "https://files.pythonhosted.org/packages/0c/27/05f950d15b8ab120b39c43588b438ff3ace70c1b1b0225a960393a497483/lupa-2.8-cp312-cp312-win_amd64.whl", hash = "sha256:281bedc5deb92d31e649a3552edd662449365a635904fa4d5cb4509c7245e34e", size = 1941809, upload-time = "2026-04-15T20:06:40.302Z" }, - { url = "https://files.pythonhosted.org/packages/a6/3f/19f83c3a0c84dc8bea8a58e7416dca6a3ede662c33c8d1ec758e5afc754a/lupa-2.8-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:45fc9da0145ecb0083ef5ff9975116cc784bd0258bdc2bd131ba15483ce18398", size = 1201203, upload-time = "2026-04-15T20:06:42.169Z" }, - { url = "https://files.pythonhosted.org/packages/89/0f/a14f0073f09610158038582e230618a48c14da6bd88185289461aa4cb854/lupa-2.8-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:58e18afed57955b41130e269c78f53d4123ab86e236b53816f4cbffa25cb5d30", size = 1806210, upload-time = "2026-04-15T20:06:45.486Z" }, - { url = "https://files.pythonhosted.org/packages/2f/14/48fff156c63a136001a7620878af7d31aa07e66b495ed621e3eddd73c294/lupa-2.8-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:fc47f536ac13a79cef47d29a2b205576a22841f042a2bcec1676b95806e7706a", size = 2359005, upload-time = "2026-04-15T20:06:47.819Z" }, - { url = "https://files.pythonhosted.org/packages/fe/18/3ac638ec90edf178242b8a2b2f00f8adae694248c03a26341ef941bb746e/lupa-2.8-cp313-cp313-win_amd64.whl", hash = "sha256:ce9404c661dbac65cc9bed351ad45e797af93d30d70be309a3fa8209ac86d93b", size = 1936754, upload-time = "2026-04-15T20:06:50.448Z" }, - { url = "https://files.pythonhosted.org/packages/b0/ef/5ee5fed6ea7459a671196359ce04bfeeaf26be1dac8ff24bf28e5c7a6e81/lupa-2.8-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:348c3f8ecabb6324dcbc05c2740d762ef8fcec7b06c79e45262ab97a217684e3", size = 1209388, upload-time = "2026-04-15T20:06:53.022Z" }, - { url = "https://files.pythonhosted.org/packages/6e/b1/67a940d5542cb0384b443fe951b5a83ea9340d1333a733a258fdd1c619ba/lupa-2.8-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:951496471056061598a7d1729a6cdf48d662fec777a9f2d8aa5a1e62fd30e5a5", size = 1826821, upload-time = "2026-04-15T20:06:55.699Z" }, - { url = "https://files.pythonhosted.org/packages/a1/a2/b354e5ba3b911ec50686003dc8897e892b9e8c5c036b33219b03d54c4daf/lupa-2.8-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:a591b9947ca347b41a63370e121d6e2b1458fe6dde9ae065029ec10a37f25ff4", size = 2366893, upload-time = "2026-04-15T20:06:58.9Z" }, - { url = "https://files.pythonhosted.org/packages/8e/52/d76066401f29539df5352f70ecded66576f32933b6045cd0bfc56cb770b9/lupa-2.8-cp314-cp314-win_amd64.whl", hash = "sha256:3903c9cf628dae2f56405503247b77a61a3a61bd2dda470e336950c74776d55d", size = 1994716, upload-time = "2026-04-15T20:07:19.194Z" }, - { url = "https://files.pythonhosted.org/packages/c3/bd/3efc437a4361c16d25e66478c50357c9a8e8ecfb718fe749eb9ca3176ef6/lupa-2.8-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:f711a8ab0486b9ac6fdda94a22ddcfbc9f0d4a27e3a8cf1bf79c6e48b33017c1", size = 1251217, upload-time = "2026-04-15T20:07:01.64Z" }, - { url = "https://files.pythonhosted.org/packages/ea/f4/2e9f8ecbaca854bfdf14af8a9b505ec0cbc640377b3b218921594b7563cd/lupa-2.8-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:dc51250e76367a3e27fcd01dc769b9bfcbbc34f48df48dde53d6af6e75b7eaa5", size = 1814701, upload-time = "2026-04-15T20:07:04.149Z" }, - { url = "https://files.pythonhosted.org/packages/ba/53/4000b1acaa8b1f3827fcff0cfcdff44d3befddda42cab7e685a49689b5a1/lupa-2.8-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f8a22088a552828958603323f0a5c4b3e11e03b75d0bf4c965ef879de9b60a8d", size = 2348414, upload-time = "2026-04-15T20:07:07.285Z" }, - { url = "https://files.pythonhosted.org/packages/d5/78/26ee48d3890cddf03cefb65f433e3492759c0b3c0582180755bddbaab7bd/lupa-2.8-cp314-cp314t-win32.whl", hash = "sha256:4f7c553c1d8cfffbe85d81daef730d12cae4b6002d457542914da0ac8a1145b3", size = 1831611, upload-time = "2026-04-15T20:07:09.752Z" }, - { url = "https://files.pythonhosted.org/packages/3c/d1/4a5cc64a3cad22821ae4c3f7a90456a08ca19457d8354f4abf46ad03c7e8/lupa-2.8-cp314-cp314t-win_amd64.whl", hash = "sha256:d8766aff03a78c80ad2d188a8bdb216de5ec838359cd87e05bbdfa56394a6105", size = 2209250, upload-time = "2026-04-15T20:07:11.906Z" }, - { url = "https://files.pythonhosted.org/packages/37/7c/cdcb654daf668192aaf36b0aeb94f2281dad092aaa5003688691131736ea/lupa-2.8-cp314-cp314t-win_arm64.whl", hash = "sha256:91d622777febda3ab1bed1d45295f2f32a4680c7b3d7caf8c669998ed5c44118", size = 1126735, upload-time = "2026-04-15T20:07:15.434Z" }, - { url = "https://files.pythonhosted.org/packages/1d/44/de1961ad38e17cd326a53c246c7e3b91178ed578f4cf22ffcd5e7e11b041/lupa-2.8-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:b036738282a5acd2e71fdddb317c9df8b87c1673aa57f403d05fcc2be8abc4ba", size = 1186020, upload-time = "2026-04-15T20:07:35.017Z" }, - { url = "https://files.pythonhosted.org/packages/13/c2/276f0b9dc8bcc5a8a58af5316dfa0e6f56be3613dd6dbcc8d3d2cb6559ba/lupa-2.8-cp39-abi3-manylinux2010_i686.manylinux_2_12_i686.manylinux_2_28_i686.whl", hash = "sha256:ac6b6e8d0e617e26a98cbb44880bcd75de5d32b3ad7b3b3793583909292b47ed", size = 1468944, upload-time = "2026-04-15T20:07:37.782Z" }, - { url = "https://files.pythonhosted.org/packages/63/38/52934e52a5180dc6425d20284d004fe4b27a4f9171a82dc99fb67af250bf/lupa-2.8-cp39-abi3-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:ba3a7dd839f90c3d2e53bebe3c192b1f3f9fd720a6781256405123211fd0dce6", size = 1172998, upload-time = "2026-04-15T20:07:40.812Z" }, - { url = "https://files.pythonhosted.org/packages/c7/82/76b3809bd0839d9b3b4ec58d06591e08f17337b6d9576877cb9d48b34e94/lupa-2.8-cp39-abi3-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:d7edb13a7a5250b5c6c22d1495d9e842b5c9fc5081c8fe6b5efe2112fe3e41f9", size = 1449975, upload-time = "2026-04-15T20:07:44.262Z" }, - { url = "https://files.pythonhosted.org/packages/16/07/2f89d54f747c67c23b4b9ae4aa8c8dd06bb409155dedcf406157f2736b66/lupa-2.8-cp39-abi3-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:891f72e0bffbed1e4175f975aeb2a083956586a100066525e1be485f617f7b25", size = 1281944, upload-time = "2026-04-15T20:07:46.458Z" }, - { url = "https://files.pythonhosted.org/packages/e7/bd/7375d2b0fcae79d806baf52a76f26c96964593f58e1372d13ae5ac09c676/lupa-2.8-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:a295f87b5b7ebbfd5191932e8cb0e51df3c7769101ac6b6c7d7c9fb27bfd1307", size = 1910455, upload-time = "2026-04-15T20:07:49.75Z" }, - { url = "https://files.pythonhosted.org/packages/8b/0c/8abb3bc0e08b311fc01db05b6e9f9ff31a8f65e4fc3f0aeb05cfef75c8ac/lupa-2.8-cp39-abi3-musllinux_1_2_armv7l.whl", hash = "sha256:4fe5d7a810b64ea8511eb885fc8cdde042ee5ff7b7d08ae78f32449756acb177", size = 1155548, upload-time = "2026-04-15T20:07:52.657Z" }, - { url = "https://files.pythonhosted.org/packages/80/2e/9eeecd3f493099721c1d3f31beeca23a4237db1a54223684df4dc96aa1bd/lupa-2.8-cp39-abi3-musllinux_1_2_i686.whl", hash = "sha256:bfc470012ef66ad064c7bd77416af03a3452ef630b04b9012595ea13f2e54518", size = 1489232, upload-time = "2026-04-15T20:07:54.92Z" }, - { url = "https://files.pythonhosted.org/packages/c3/13/731c99dc2e7652ae818a6de45bdf0142049f7cb566049061c898355f1891/lupa-2.8-cp39-abi3-musllinux_1_2_ppc64le.whl", hash = "sha256:250e035fdaffe8c87093e3ebc206ac29a26131b1568ea711d780c26001ce96e7", size = 1466321, upload-time = "2026-04-15T20:07:57.627Z" }, - { url = "https://files.pythonhosted.org/packages/de/71/3ad8cc4fc05a77dc0d3f7079348bd1cad4675a0d14c24f8e6a3ce5f008f7/lupa-2.8-cp39-abi3-musllinux_1_2_riscv64.whl", hash = "sha256:b9bddb09acfffb4f828f790f444b11dc0cca591afea1a244d9329eea2d20c003", size = 1288577, upload-time = "2026-04-15T20:07:59.913Z" }, - { url = "https://files.pythonhosted.org/packages/d8/b2/1175f6d0aa7b68627fbe2f58bd1e8bea36a89d10dfd67671d2b024c96162/lupa-2.8-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:2e64acbbd47e9b82a64405a39e0d2b36a5a7dad8ab41c0f3437f572f7d282ba3", size = 2444866, upload-time = "2026-04-15T20:08:02.753Z" }, - { url = "https://files.pythonhosted.org/packages/92/f7/e78df680c7a0ea452daac07467ca188d63c2c00ca1c884c0a50e27eb83b5/lupa-2.8-pp311-pypy311_pp73-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:32e4e5103bbddcdd2458fb2ccae6c8ba11c9997c711d7e379e0d45551d109c76", size = 1778509, upload-time = "2026-04-15T20:08:21.784Z" }, - { url = "https://files.pythonhosted.org/packages/e6/23/0e53cabb16b2a8aa9cf1fde499c097d8942c5dab709fc8e921f3b824b18b/lupa-2.8-pp311-pypy311_pp73-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:7667001804657496dee9feced2daae5000b4604a3218dd8e6b7b754982ba88b8", size = 2300480, upload-time = "2026-04-15T20:08:24.394Z" }, - { url = "https://files.pythonhosted.org/packages/7e/85/0271227eab939921a12ebba5d17aa4cd18346aa534ca7f5da09cd0b63dd4/lupa-2.8-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:86f6f668966965b15247dc32d064cfe7be67b71e584ccfacbe2f637575296878", size = 1847445, upload-time = "2026-04-15T20:08:27.031Z" }, -] - [[package]] name = "mako" version = "1.3.12" From afa0b23dc20a3bfbf32728ab59b6f4a05350889a Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 11 Sep 2026 15:00:22 +0200 Subject: [PATCH 38/90] fix(realtime): bound offers and preserve duration billing --- litellm/cost_calculator.py | 2 - .../proxy/realtime_endpoints/call_sessions.py | 24 ++++ .../realtime_endpoints/test_call_sessions.py | 104 ++++++++++++++++++ tests/test_litellm/test_cost_calculator.py | 47 ++++++-- 4 files changed, 167 insertions(+), 10 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 094b6d45a9a..15060004945 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -2601,8 +2601,6 @@ def handle_live_session_duration_cost( custom_llm_provider: str, litellm_model_name: str, ) -> float: - if any(event.get("type") == "response.done" for event in results): - return 0.0 terminal: Final = next((event for event in reversed(results) if event.get("type") == "session.closed"), None) if terminal is None: return 0.0 diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index bac2abd6415..8989fadfead 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -223,7 +223,31 @@ def decode_call(token: str, authorization: str) -> CodexRealtimeCall: return call +MAX_REALTIME_OFFER_BYTES: Final = 8 * 1024 * 1024 + + +async def _cache_bounded_offer_body(request: Request) -> None: + try: + if int(request.headers.get("content-length", "")) > MAX_REALTIME_OFFER_BYTES: + raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit") + except ValueError: + pass + 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(): diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index d730add23b2..dffd4500ed7 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -11,6 +11,110 @@ 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): diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 5886ba8b85a..9f9db7588c6 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -4827,12 +4827,43 @@ def test_live_terminal_is_not_counted_twice(monkeypatch): assert handle_realtime_stream_cost_calculation( [_live_terminal_event(), _live_terminal_event()], Usage(), "chatgpt", "live-priced-test" ) == pytest.approx(0.1) - assert ( - handle_realtime_stream_cost_calculation( - [{"type": "response.done", "response": {"usage": {}}}, _live_terminal_event()], - Usage(), - "chatgpt", - "live-priced-test", - ) - == 0 + + +@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)) From eea527e9b18489ebdaec17c017b4a3b1fd3dd9ac Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sat, 12 Sep 2026 21:24:33 +0200 Subject: [PATCH 39/90] fix(images): support GPT Image 2.5 quality and preserve JSON edits --- litellm/images/main.py | 2 + .../litellm_core_utils/llm_cost_calc/utils.py | 26 +++++- ...odel_prices_and_context_window_backup.json | 16 ++-- litellm/types/images/main.py | 2 +- litellm/types/llms/openai.py | 2 + litellm/types/utils.py | 4 + litellm/utils.py | 1 + model_prices_and_context_window.json | 16 ++-- .../test_litellm/llms/chatgpt/test_images.py | 80 +++++++++++++++++++ .../test_gpt_image_cost_calculator.py | 39 ++++++++- 10 files changed, 168 insertions(+), 20 deletions(-) diff --git a/litellm/images/main.py b/litellm/images/main.py index 11df9728ede..662903ee35e 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -869,6 +869,8 @@ def image_edit( non_default_params, extra_body if isinstance(extra_body, dict) else None, ) + if image_edit_provider_config.use_multipart_form_data() + else {**non_default_params, **(extra_body if isinstance(extra_body, dict) else {})} ) # Pre Call logging diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index e5977ca4156..341bb0235d9 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -1645,7 +1645,31 @@ def calculate_image_response_cost_from_usage( usage=normalized_usage, custom_llm_provider=custom_llm_provider, ) - 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 + model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider) + cached_text: Final = _get_token_detail_value(cached_details, "text_tokens") or 0 + cached_image: Final = _get_token_detail_value(cached_details, "image_tokens") or 0 + input_text_tokens: Final = _get_token_detail_value(input_tokens_details, "text_tokens") or 0 + input_image_tokens: Final = _get_token_detail_value(input_tokens_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 = model_info.get("input_cost_per_token") or 0.0 + image_rate: Final = model_info.get("input_cost_per_image_token") + cache_text_rate: Final = model_info.get("cache_read_input_token_cost") + cache_image_rate: Final = 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/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 50d0ed52def..74f56f700c0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -29797,6 +29797,7 @@ }, "gpt-image-2.5-flare": { "cache_read_input_token_cost": 1.25e-06, + "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -29807,11 +29808,11 @@ "/v1/images/edits" ], "supports_vision": true, - "supports_pdf_input": true, - "source": "https://developers.openai.com/api/docs/pricing" + "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-flare" }, "gpt-image-2.5-flare-2026-09-08": { "cache_read_input_token_cost": 1.25e-06, + "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -29822,11 +29823,11 @@ "/v1/images/edits" ], "supports_vision": true, - "supports_pdf_input": true, - "source": "https://developers.openai.com/api/docs/pricing" + "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-flare" }, "gpt-image-2.5-sunburst": { "cache_read_input_token_cost": 1.25e-06, + "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -29837,11 +29838,11 @@ "/v1/images/edits" ], "supports_vision": true, - "supports_pdf_input": true, - "source": "https://developers.openai.com/api/docs/pricing" + "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-sunburst" }, "gpt-image-2.5-sunburst-2026-09-08": { "cache_read_input_token_cost": 1.25e-06, + "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -29852,8 +29853,7 @@ "/v1/images/edits" ], "supports_vision": true, - "supports_pdf_input": true, - "source": "https://developers.openai.com/api/docs/pricing" + "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-sunburst" }, "low/1024-x-1024/gpt-image-1.5": { "deprecation_date": "2026-12-01", 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/llms/openai.py b/litellm/types/llms/openai.py index 83eb3c4aa2a..47668823b47 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2273,6 +2273,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/utils.py b/litellm/types/utils.py index 58b940227f8..70b7d603d3f 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -250,6 +250,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): cache_creation_input_token_cost_priority: float | None # OpenAI priority service tier pricing cache_creation_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing cache_read_input_token_cost: float | None + cache_read_input_image_token_cost: ReadOnly[float | None] cache_read_input_token_cost_flex: float | None # OpenAI flex service tier pricing cache_read_input_token_cost_priority: float | None # OpenAI priority service tier pricing cache_read_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing @@ -3804,6 +3805,9 @@ all_litellm_params = ( "enable_tag_filtering", "enable_json_schema_validation", "use_xai_oauth", + "chatgpt_auth_profile", + "chatgpt_token_dir", + "chatgpt_auth_file", "auto_router_config_path", "auto_router_config", "auto_router_default_model", diff --git a/litellm/utils.py b/litellm/utils.py index 7075d8776c3..6ba76281b44 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5897,6 +5897,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_image=_model_info.get("input_cost_per_image", None), input_cost_per_audio_per_second=_model_info.get("input_cost_per_audio_per_second", None), diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 50d0ed52def..74f56f700c0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -29797,6 +29797,7 @@ }, "gpt-image-2.5-flare": { "cache_read_input_token_cost": 1.25e-06, + "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -29807,11 +29808,11 @@ "/v1/images/edits" ], "supports_vision": true, - "supports_pdf_input": true, - "source": "https://developers.openai.com/api/docs/pricing" + "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-flare" }, "gpt-image-2.5-flare-2026-09-08": { "cache_read_input_token_cost": 1.25e-06, + "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -29822,11 +29823,11 @@ "/v1/images/edits" ], "supports_vision": true, - "supports_pdf_input": true, - "source": "https://developers.openai.com/api/docs/pricing" + "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-flare" }, "gpt-image-2.5-sunburst": { "cache_read_input_token_cost": 1.25e-06, + "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -29837,11 +29838,11 @@ "/v1/images/edits" ], "supports_vision": true, - "supports_pdf_input": true, - "source": "https://developers.openai.com/api/docs/pricing" + "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-sunburst" }, "gpt-image-2.5-sunburst-2026-09-08": { "cache_read_input_token_cost": 1.25e-06, + "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -29852,8 +29853,7 @@ "/v1/images/edits" ], "supports_vision": true, - "supports_pdf_input": true, - "source": "https://developers.openai.com/api/docs/pricing" + "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-sunburst" }, "low/1024-x-1024/gpt-image-1.5": { "deprecation_date": "2026-12-01", diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 2c51b932e3e..7cd8d727a94 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -1,4 +1,6 @@ import base64 +import json +from typing import Final import httpx import pytest @@ -6,9 +8,87 @@ 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 = [] diff --git a/tests/test_litellm/test_gpt_image_cost_calculator.py b/tests/test_litellm/test_gpt_image_cost_calculator.py index 86a721f8743..36dcc921e22 100644 --- a/tests/test_litellm/test_gpt_image_cost_calculator.py +++ b/tests/test_litellm/test_gpt_image_cost_calculator.py @@ -10,15 +10,15 @@ gpt-image-1 uses token-based pricing: - Image Output: $40.00/1M tokens """ - +from typing import Final import pytest import litellm from litellm.types.utils import ( CompletionTokensDetailsWrapper, - ImageResponse, ImageObject, + ImageResponse, ImageUsage, ImageUsageInputTokensDetails, PromptTokensDetailsWrapper, @@ -42,6 +42,41 @@ 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 From 5db17b52a230acdc53b062441c0063b5cebdd7cf Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sun, 13 Sep 2026 10:19:46 +0200 Subject: [PATCH 40/90] fix(images): sync cached image pricing schema --- model_prices_and_context_window.schema.json | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index 7ed1e7e568b..b1a9f211732 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -133,6 +133,10 @@ "type": "number", "minimum": 0 }, + "cache_read_input_image_token_cost": { + "type": "number", + "minimum": 0 + }, "cache_read_input_token_cost": { "type": "number", "minimum": 0, From 0a5fba363639220467c95fb027e679dc424eecee Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sun, 13 Sep 2026 10:34:54 +0200 Subject: [PATCH 41/90] chore: sync API schemas after upstream merge --- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 -- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 7a110eff080..4c4f2164a59 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -18968,7 +18968,7 @@ } } }, - "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " + "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" }, "500": { "content": { diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ac34de29cca..c764270c07e 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -16801,7 +16801,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) @@ -16907,7 +16906,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) From 9af75b23e1bd5d6bbe8c7dd08b4d9a8740ed7170 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sun, 13 Sep 2026 10:40:12 +0200 Subject: [PATCH 42/90] refactor: reuse Live sideband authentication dependency --- litellm/proxy/proxy_server.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4046ccf8886..6c17701e961 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11918,11 +11918,14 @@ async def _reject_realtime_session( await _release_realtime_budget_reservation(user_api_key_dict) +_CODEX_LIVE_AUTH_DEPENDENCY: Final = Depends(user_api_key_auth_websocket) + + @app.websocket("/v1/live/{call_id}") async def codex_live_sideband_endpoint( websocket: WebSocket, call_id: str, - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), + user_api_key_dict: UserAPIKeyAuth = _CODEX_LIVE_AUTH_DEPENDENCY, ) -> None: from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband From 9ef11699e52bd5c68d7d8e8096f9efb115f983ac Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sun, 13 Sep 2026 11:04:09 +0200 Subject: [PATCH 43/90] fix: complete image pricing metadata and stabilize OpenAPI descriptions --- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- litellm/proxy/common_utils/swagger_utils.py | 3 ++- litellm/types/utils.py | 1 + tests/test_litellm/test_utils.py | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 ++++ 5 files changed, 9 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 4c4f2164a59..3be39265de7 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -18968,7 +18968,7 @@ } } }, - "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" + "description": "Unified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values." }, "500": { "content": { diff --git a/litellm/proxy/common_utils/swagger_utils.py b/litellm/proxy/common_utils/swagger_utils.py index 2609a98a997..4b28a4e2fd4 100644 --- a/litellm/proxy/common_utils/swagger_utils.py +++ b/litellm/proxy/common_utils/swagger_utils.py @@ -1,3 +1,4 @@ +import inspect from typing import Any, Final from pydantic import BaseModel, Field @@ -35,7 +36,7 @@ def get_status_code(exception): ERROR_RESPONSES: Final = { get_status_code(exception): { "model": ErrorResponse, - "description": exception.__doc__ or exception.__name__, + "description": inspect.cleandoc(exception.__doc__ or exception.__name__), } for exception in LITELLM_EXCEPTION_TYPES } diff --git a/litellm/types/utils.py b/litellm/types/utils.py index b784f066241..ff617162656 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3560,6 +3560,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams): cache_read_input_token_cost_above_272k_tokens_priority: float | None = None cache_read_input_token_cost_above_272k_tokens_flex: float | None = None cache_read_input_audio_token_cost: float | None = None + cache_read_input_image_token_cost: float | None = None input_cost_per_character_above_128k_tokens: float | None = None input_cost_per_audio_token: float | None = None input_cost_per_token_cache_hit: float | None = None diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index b3bb5f4f4de..75cf5c3ffb1 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -1432,6 +1432,7 @@ def test_openai_models_in_model_info(monkeypatch): 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) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c764270c07e..97545f327a5 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -29729,6 +29729,8 @@ export interface components { cache_creation_input_token_cost_ultrafast?: number | null; /** Cache Read Input Audio Token Cost */ cache_read_input_audio_token_cost?: number | null; + /** Cache Read Input Image Token Cost */ + cache_read_input_image_token_cost?: number | null; /** Cache Read Input Token Cost */ cache_read_input_token_cost?: number | null; /** Cache Read Input Token Cost Above 200K Tokens */ @@ -39943,6 +39945,8 @@ export interface components { cache_creation_input_token_cost_ultrafast?: number | null; /** Cache Read Input Audio Token Cost */ cache_read_input_audio_token_cost?: number | null; + /** Cache Read Input Image Token Cost */ + cache_read_input_image_token_cost?: number | null; /** Cache Read Input Token Cost */ cache_read_input_token_cost?: number | null; /** Cache Read Input Token Cost Above 200K Tokens */ From be9a191cac0e4b88e2623937d35ff6a411915f24 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sun, 13 Sep 2026 11:33:50 +0200 Subject: [PATCH 44/90] test: allow cached image prices in legacy catalog schema --- tests/test_litellm/test_utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 75cf5c3ffb1..d38d57be38e 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -995,6 +995,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "type": "number" }, "cache_read_input_audio_token_cost": {"type": "number"}, + "cache_read_input_image_token_cost": {"type": "number"}, "audio_transcription_config": {"type": "string"}, "deprecation_date": {"type": "string"}, "input_cost_per_audio_per_second": {"type": "number"}, From fbda29e1b71f1051f1e990da2d3f84c13f264760 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 17 Sep 2026 08:34:37 +0200 Subject: [PATCH 45/90] feat(chatgpt): add public Live routes and Codex setup guide --- docs/my-website/docs/providers/chatgpt.md | 121 + litellm/cost_calculator.py | 126 +- .../litellm_core_utils/realtime_streaming.py | 15 +- litellm/llms/chatgpt/live.py | 193 ++ litellm/proxy/_lazy_features.py | 7 +- litellm/proxy/_lazy_openapi_snapshot.json | 2003 +++++++++++++++++ litellm/proxy/_types.py | 27 + litellm/proxy/auth/auth_utils.py | 15 +- litellm/proxy/proxy_server.py | 6 + .../proxy/realtime_endpoints/call_sessions.py | 12 +- .../realtime_endpoints/call_supervision.py | 22 +- litellm/proxy/realtime_endpoints/endpoints.py | 9 + litellm/proxy/realtime_endpoints/live.py | 1070 +++++++++ litellm/types/llms/openai.py | 8 +- litellm/types/realtime.py | 12 +- .../test_realtime_streaming.py | 40 + tests/test_litellm/llms/chatgpt/test_live.py | 194 ++ .../proxy/auth/test_auth_utils.py | 23 +- .../realtime_endpoints/test_call_sessions.py | 15 +- .../test_call_supervision.py | 60 +- .../proxy/realtime_endpoints/test_live.py | 930 ++++++++ .../test_realtime_webrtc_endpoints.py | 111 +- .../proxy/test_live_route_registration.py | 57 + tests/test_litellm/test_cost_calculator.py | 225 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 1101 ++++++++- 25 files changed, 6340 insertions(+), 62 deletions(-) create mode 100644 docs/my-website/docs/providers/chatgpt.md create mode 100644 litellm/llms/chatgpt/live.py create mode 100644 litellm/proxy/realtime_endpoints/live.py create mode 100644 tests/test_litellm/llms/chatgpt/test_live.py create mode 100644 tests/test_litellm/proxy/realtime_endpoints/test_live.py create mode 100644 tests/test_litellm/proxy/test_live_route_registration.py diff --git a/docs/my-website/docs/providers/chatgpt.md b/docs/my-website/docs/providers/chatgpt.md new file mode 100644 index 00000000000..bc971be087e --- /dev/null +++ b/docs/my-website/docs/providers/chatgpt.md @@ -0,0 +1,121 @@ +# ChatGPT, Codex, and GPT-Live + +The proxy exposes the public GPT-Live session routes and retains the Codex-compatible `POST /live` route. Clients authenticate to LiteLLM with a LiteLLM virtual key. A `chatgpt` deployment uses one proxy-wide ChatGPT OAuth record, while an `openai` deployment uses its configured OpenAI API key. Do not send an upstream OAuth token as the proxy key + +## Configure LiteLLM deployments + +Codex text, image, and voice requests need separate LiteLLM aliases because their backends have different capabilities. The canonical aliases below keep the primary Qwen model separate from the ChatGPT OAuth models + +```yaml +model_list: + - model_name: qwen3.8-flash-next-codex + litellm_params: + model: openai/qwen3.8-flash-next + api_base: https://qwen.example.com/v1 + api_key: os.environ/QWEN_API_KEY + + - model_name: gpt-image-2 + litellm_params: + model: chatgpt/gpt-image-2 + + - model_name: gpt-image-2.5-flare + litellm_params: + model: chatgpt/gpt-image-2.5-flare + + - model_name: gpt-image-2.5-sunburst + litellm_params: + model: chatgpt/gpt-image-2.5-sunburst + + - model_name: gpt-realtime-1.5 + litellm_params: + model: chatgpt/gpt-realtime-1.5 + + - model_name: gpt-live-1-codex + litellm_params: + model: chatgpt/gpt-live-1-codex +``` + +If clients use shorter local aliases, publish separate aliases such as `qwen-codex`, `images`, and `voice` that point to the corresponding deployments. Authorize the exact alias sent by the client in the key or team policy, and for a restricted team member set `allowed_models` to contain the aliases used for text, image, or voice requests. Selecting the Qwen alias does not give it image-generation or voice capabilities. The Qwen endpoint in this example uses the OpenAI-compatible adapter and must expose `/v1/responses`; a native vLLM deployment can use `hosted_vllm/qwen3.8-flash-next` when that endpoint is available. LiteLLM does not promise provider-specific tool, reasoning, or stream behavior parity. For a deployment using the public OpenAI API instead, configure `model: openai/gpt-live-1` and `api_key: os.environ/OPENAI_API_KEY` under its own alias. The public API documentation uses `gpt-live-1`; the Codex alias and its backend capabilities are separate + +### Configure global ChatGPT OAuth + +The ChatGPT provider reads one auth file for the LiteLLM process. Set these environment variables before starting the proxy when the default location is not suitable + +```bash +export CHATGPT_TOKEN_DIR=/var/lib/litellm/chatgpt +export CHATGPT_AUTH_FILE=auth.json +``` + +The defaults are `~/.config/litellm/chatgpt` and `auth.json`. Persist the directory and complete the provider's device-code OAuth flow. The provider refreshes the stored record when it expires. All `chatgpt/...` deployments in that process share this record. Live rejects per-deployment `chatgpt_auth_profile`, `chatgpt_token_dir`, and `chatgpt_auth_file` overrides + +### Image generation and editing + +Use the `gpt-image-2` alias for both `/v1/images/generations` and `/v1/images/edits`. Keep the exact image aliases requested by your Codex version; a shorter `images` alias only works for clients configured to request it. Codex sends image requests to its active provider's `base_url`, with no separate image URL override. LiteLLM then selects the image deployment independently of the primary text model. ChatGPT image generation and editing require a valid global ChatGPT OAuth login, but a successful login does not establish that the account or backend supports every image operation. Image editing accepts JSON reference images and multipart files, but not masks. Image 2.5 aliases preserve the requested model name; an accepted name does not prove which backend model executed + +## Configure Codex through LiteLLM + +Codex sends its Responses requests to LiteLLM's `/v1/responses` endpoint. Point the Codex provider at the proxy and select the primary Qwen alias (or the shorter alias you published) + +```toml +model = "qwen3.8-flash-next-codex" +model_provider = "litellm" +experimental_realtime_ws_base_url = "https://litellm.example.com/v1" +experimental_realtime_webrtc_call_base_url = "https://litellm.example.com/v1" + +[model_providers.litellm] +name = "LiteLLM" +base_url = "https://litellm.example.com/v1" +wire_api = "responses" +requires_openai_auth = true +experimental_bearer_token = "" +``` + +Keep Codex signed in with ChatGPT for its client-side capability checks. `experimental_bearer_token` is the LiteLLM virtual key issued by the proxy. It must never contain the upstream ChatGPT OAuth access or refresh token. `requires_openai_auth = true` enables the Codex OpenAI-auth capability path while LiteLLM remains responsible for the upstream provider credentials + +## Configure voice routing + +The two realtime settings in the TOML example are root-level Codex settings, not fields inside `[model_providers.litellm]`. `experimental_realtime_ws_base_url` routes the Realtime WebSocket and its sideband through LiteLLM. `experimental_realtime_webrtc_call_base_url` is optional and separately routes HTTP WebRTC call creation. The optional root setting `experimental_realtime_ws_model` overrides the voice model; leave it unset to retain your client's default. An override must match the active protocol: `gpt-realtime-1.5` for legacy Realtime v1/v2 or `gpt-live-1-codex` for frameless Live v3. The base URLs end at `/v1`; LiteLLM adds the protocol-specific path + +| Voice operation | LiteLLM path | +| --- | --- | +| Legacy Realtime WebSocket | `/v1/realtime` | +| HTTP WebRTC call creation | `/v1/realtime/calls` | +| Frameless Live v3 signaling | `/v1/live` and `/v1/live/{call_id}` | +| Public Live session APIs | `/v1/live/sessions...` | + +WebSocket authentication uses `Authorization: Bearer ` by default. If the proxy sets `litellm_key_header_name`, send the virtual key in that configured header instead. Voice is experimental and pending retest: an observed mobile `POST /live` returned 201, but its sideband used the default `api.openai.com` and returned 404. The WebSocket base URL above addresses that routing gap; full bidirectional voice is not verified + +## Public Live routes + +Use the proxy host in place of `api.openai.com`. Send the configured LiteLLM alias in `session.model` when creating a session, or in the first `session.start` event for a primary WebSocket. Keep the returned session ID unchanged for subsequent operations + +| Method | Path | Request and response | +| --- | --- | --- | +| POST | `/v1/live/sessions` | JSON `session` and `transport: {type: "webrtc", sdp: ""}`; returns 201 JSON with `session.id` and `transport.sdp` | +| POST | `/v1/live/sessions/{session_id}/fork` | JSON WebRTC `transport` and optional `session` overrides; returns 200 JSON with the new session ID and SDP answer | +| GET | `/v1/live/sessions/{session_id}/content` | Downloads stored recording content without converting it to JSON | +| POST | `/v1/live/sessions/{session_id}/accept` | JSON `session` with `type: "live"` and model; successful SIP acceptance returns an empty body | +| POST | `/v1/live/sessions/{session_id}/reject` | JSON with required integer `status_code` from 300 through 699 | +| POST | `/v1/live/sessions/{session_id}/refer` | JSON with `target_uri` for the SIP destination | +| POST | `/v1/live/sessions/{session_id}/hangup` | No request body | +| WebSocket | `/v1/live/sessions` | Start with `session.start`, then wait for `session.started` before sending audio or commands | +| WebSocket | `/v1/live/sessions/{session_id}/attach` | Attach to an existing session; do not send `session.start` or input audio | +| WebSocket | `/v1/live/sessions/{session_id}/fork` | Start with `session.start` and a required `session` overrides object, which may be empty | + +Public WebRTC creation uses JSON, not the multipart or raw SDP formats used by the Codex compatibility route. `POST /live` and its existing aliases remain available for Codex clients using that format. WebRTC audio travels on media tracks; its data channel carries Live JSON events. Primary WebSocket audio uses base64 chunks in `session.input_audio.append` and `session.output_audio.delta` + +The proxy preserves Live event payloads, including nested Responses events inside `response.event`, rather than translating them into Realtime events. Session routing rewrites the configured model alias to the selected upstream model. Audio, transcript, delegation and usage events retain their upstream format. Send `session.close` and wait for `session.closed` to obtain final usage; a disconnected socket alone does not confirm successful finalization + +## Availability and verification + +Route support does not establish that every configured backend or account supports every operation. The official API describes project API-key authentication; it does not guarantee equivalent capabilities for ChatGPT OAuth. An OAuth request reaching SDP validation proves only that the request reached that validation step. It does not prove a working audio session, recording, fork or SIP call. The routes listed here have not all been tested against a real upstream service + +Session controls require a session known to the proxy and owned by the authenticated caller. Incoming SIP calls originate upstream. A proxy administrator can accept or reject the raw ID from a verified incoming-call webhook by supplying `x-litellm-live-model` with an alias that resolves to exactly one deployment. Successful acceptance returns the proxy-owned handle in `x-litellm-live-session-id`, preserving the API's empty response body. Use that handle for subsequent controls. Ordinary virtual keys cannot enroll arbitrary upstream session IDs; a trusted webhook-to-owner enrollment flow is still required for those keys + +Live duration uses cumulative `usage.seconds`; legacy Codex milliseconds remain supported. WebRTC initialization has a 15-second minimum credited against running duration, not added to it. Nested terminal Responses usage is charged separately using its backend model and deduplicated by response ID. A failed observation connection cannot establish complete usage. Managed delegation also depends on receiving its backend usage events; the upstream sideband does not replay events emitted before attachment + +Managed Responses delegation authorizes its backend model as well as the voice model. Use client delegation when a key's per-model budgets or token/request limits require admission checks for each backend invocation. Restricted-model WebRTC keys must explicitly exclude `session.update` from frontend client events when using managed delegation, because that data channel connects directly to OpenAI and could otherwise change the backend model outside the proxy's checks + +For both HTTP and WebSocket forks, restricted-model keys must explicitly set `session.delegation` to `{ "type": "client" }` or `{ "type": "responses", "responses": { "model": "authorized-backend" } }`. Empty overrides cannot safely authorize an inherited backend: the session handle records startup configuration, while later updates may have changed the upstream model. Upstream rules still determine which delegation overrides a source session permits + +See the official [Live overview](https://developers.openai.com/api/docs/guides/live), [Live API reference](https://developers.openai.com/api/reference/resources/live), [session management](https://developers.openai.com/api/docs/guides/live-conversations), [WebRTC guide](https://developers.openai.com/api/docs/guides/voice-webrtc?api=live), [WebSocket guide](https://developers.openai.com/api/docs/guides/voice-websockets?api=live), [server controls](https://developers.openai.com/api/docs/guides/voice-server-controls?api=live) and [SIP guide](https://developers.openai.com/api/docs/guides/voice-sip?api=live) for the upstream contract. The voice guides also contain Realtime tabs with different routes and formats diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 87807d83efd..000b1b9de8a 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, Field, ValidationError +from pydantic import BaseModel, ValidationError from typing_extensions import ReadOnly, TypedDict import litellm @@ -2471,6 +2471,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 @@ -2621,7 +2628,13 @@ 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( + [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, ) @@ -2639,7 +2652,13 @@ def handle_realtime_stream_cost_calculation( custom_llm_provider=custom_llm_provider, litellm_model_name=litellm_model_name, ) - total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost + live_audio_cost + 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, @@ -2649,7 +2668,11 @@ def handle_realtime_stream_cost_calculation( total_cost_usd_dollar=total_cost, 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)) + for name, cost in ( + ("transcription_cost", transcription_cost), + ("live_audio_cost", live_audio_cost), + ("live_backend_cost", backend_cost), + ) if cost > 0 } or None, @@ -2659,12 +2682,61 @@ def handle_realtime_stream_cost_calculation( return total_cost -class _LiveSessionDurationUsage(BaseModel): - audio_duration_ms: float = Field(strict=True, ge=0, allow_inf_nan=False) +class _LiveBackendEvent(BaseModel): + type: str + response: object = None -class _LiveSessionClosedEvent(BaseModel): - usage: _LiveSessionDurationUsage +class _LiveBackendEnvelope(BaseModel): + event: _LiveBackendEvent + + +def _live_backend_responses( + results: OpenAIRealtimeStreamList, logging_obj: LitellmLoggingObject | None = None +) -> tuple[ResponsesAPIResponse, ...]: + responses: Final = { + 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( @@ -2672,18 +2744,40 @@ def handle_live_session_duration_cost( custom_llm_provider: str, litellm_model_name: str, ) -> float: - terminal: Final = next((event for event in reversed(results) if event.get("type") == "session.closed"), None) - if terminal is None: - return 0.0 - try: - usage: Final = _LiveSessionClosedEvent.model_validate(terminal).usage - except ValidationError: - return 0.0 + 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: return 0.0 - return usage.audio_duration_ms / 1000 * (model_info.get("input_cost_per_second") or 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(Mapping[str, object], event.get("usage")) + return usage.duration / (1 if "seconds" in raw_usage else 1000) def handle_realtime_transcription_cost_calculation( diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 70ba3341ac9..dfbc89ce193 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -15,6 +15,7 @@ from litellm._logging import redact_internal_details_from_client_message, verbos from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.types.llms.openai import ( + OpenAILiveResponseEvent, OpenAIRealtimeEvents, OpenAIRealtimeOutputItemDone, OpenAIRealtimeResponseDelta, @@ -148,6 +149,7 @@ class RealTimeStreaming: logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER, *, account_usage: bool = True, + live_initialization_seconds: float = 0, ): self.websocket: _ClientWebSocket = websocket self.backend_ws = backend_ws @@ -155,6 +157,10 @@ class RealTimeStreaming: 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( + {"type": "litellm.live.initialization", "usage": {"seconds": live_initialization_seconds}} + ) self._backend_sent_frames: bool = False self.input_message: dict = {} self.input_messages: list[dict[str, str]] = [] @@ -266,9 +272,16 @@ class RealTimeStreaming: else: message_obj = cast(dict[str, Any], json.loads(cast(str, message))) self._collect_tool_calls_from_response_done(cast(dict, message_obj)) - if message_obj.get("type") == "session.closed" and isinstance(message_obj.get("usage"), dict): + 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"): + self.messages.append(TypeAdapter(OpenAILiveResponseEvent).validate_python(message_obj)) + return if not self._should_store_message(message_obj): return try: diff --git a/litellm/llms/chatgpt/live.py b/litellm/llms/chatgpt/live.py new file mode 100644 index 00000000000..696ac00f118 --- /dev/null +++ b/litellm/llms/chatgpt/live.py @@ -0,0 +1,193 @@ +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 ( + 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") # rebind-ok: validate successive decoding layers iteratively + if decoded == candidate: + return + candidate = decoded # rebind-ok: each percent-decoding pass reduces the input length + + +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_client: httpx.AsyncClient | None = None, + ) -> None: + self.deployment = deployment + self._http_client = http_client + 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) + client: Final = ( + self._http_client + or get_async_httpx_client( + llm_provider=LlmProviders.CHATGPT if self.deployment.provider == "chatgpt" else LlmProviders.OPENAI + ).client + ) + return await client.request( + method, + url, + headers=MappingProxyType({**self._headers, "content-type": "application/json"}), + json=dict(body) if body is not None else None, # mutable-ok: JSON encoder requires a concrete dict + timeout=60, + follow_redirects=False, + ) + + 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/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index dd1180b30ad..172be2250fe 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -213,10 +213,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 3be39265de7..3fbb0e48a0b 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -15890,6 +15890,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": { @@ -19222,6 +20060,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", @@ -37024,6 +38140,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": { @@ -37080,10 +38209,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", @@ -37228,6 +38953,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 7b111c1513e..a757a02c6ec 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -412,6 +412,33 @@ class LiteLLMRoutes(enum.Enum): "/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_utils.py b/litellm/proxy/auth/auth_utils.py index 2cb739878bd..1b096e60b09 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1705,6 +1705,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", @@ -1866,15 +1867,15 @@ def _extract_model_candidates_from_request( uses_completion_model_sources: Final = _route_matches_any_marker( route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS ) - session: Final[object] = ( - request_data.get("session") - if _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS) - else None - ) + 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 ( - _route_matches_any_marker(route=route, markers=("/realtime/calls",)) + uses_session_model + and not _route_matches_any_marker(route=route, markers=("/realtime/client_secrets",)) and isinstance(session_model, str) and session_model ): @@ -1885,7 +1886,7 @@ def _extract_model_candidates_from_request( _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): + if uses_session_model: _append_model_candidates(candidates, 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/proxy_server.py b/litellm/proxy/proxy_server.py index 6c17701e961..d51574de8b8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11920,7 +11920,12 @@ async def _reject_realtime_session( _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, @@ -11934,6 +11939,7 @@ async def codex_live_sideband_endpoint( @app.websocket("/v1/live") @app.websocket("/live") +@app.websocket("/openai/v1/live") @app.websocket("/openai/v1/realtime") @app.websocket("/v1/realtime") @app.websocket("/realtime") diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 8989fadfead..d1908919503 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -356,7 +356,13 @@ async def _create_codex_realtime_call(request: Request) -> Response: valid_token=auth, llm_router=server.llm_router, ) - data: Final = build_call_request(offer, request.query_params, request.headers) + 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 @@ -409,7 +415,9 @@ async def _create_codex_realtime_call(request: Request) -> Response: response.content, status_code=response.status_code, media_type="application/sdp", - headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), + headers=MappingProxyType( + {"Location": f"/v1/live/{token}" if live_signaling else f"/v1/realtime/calls/{token}"} + ), ) finally: try: diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index f511d943c3a..3f9ed521624 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -3,7 +3,7 @@ from collections.abc import AsyncIterator, Awaitable, Callable from contextlib import suppress from typing import Final, Protocol -from pydantic import BaseModel, Field, ValidationError +from pydantic import BaseModel, ValidationError from websockets.exceptions import ConnectionClosedOK from litellm._logging import verbose_proxy_logger @@ -16,6 +16,7 @@ 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): @@ -34,14 +35,6 @@ class _ObserverEvent(BaseModel): type: str -class _LiveDurationUsage(BaseModel): - audio_duration_ms: float = Field(strict=True, ge=0, allow_inf_nan=False) - - -class _LiveTerminalEvent(BaseModel): - usage: _LiveDurationUsage - - class CallSupervisor: def __init__( self, @@ -57,6 +50,7 @@ class CallSupervisor: 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: @@ -75,7 +69,9 @@ class CallSupervisor: self._terminal_usage_required = terminal_usage_required self._ready = asyncio.Event() self._stop = asyncio.Event() - self._started = False + self._started = connected_ready + if connected_ready: + self._ready.set() self._terminal = False self._terminal_usage_valid = False self._close_confirmed = False @@ -128,7 +124,7 @@ class CallSupervisor: if event.type == "session.closed": self._terminal = True try: - _LiveTerminalEvent.model_validate_json(message) + LiveSessionUsageEvent.model_validate_json(message) except ValidationError: self._terminal_usage_valid = False else: @@ -206,7 +202,9 @@ class CallSupervisor: await asyncio.wait_for( self._stream.log_messages(wait_for_dispatch=True), timeout=self._logging_timeout ) - self._accounting_complete = True + 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: diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index 82041b4274f..c1da14cb667 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -362,6 +362,15 @@ async def create_realtime_client_secret( return RealtimeClientSecretResponse(**upstream_json) +@router.post("/v1/live", tags=["realtime"]) +@router.post("/live", tags=["realtime"]) +@router.post("/openai/v1/live", tags=["realtime"]) +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"], diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py new file mode 100644 index 00000000000..77e732f65fa --- /dev/null +++ b/litellm/proxy/realtime_endpoints/live.py @@ -0,0 +1,1070 @@ +import asyncio +import base64 +import hashlib +import json +import time +from collections.abc import AsyncGenerator, Mapping +from contextlib import asynccontextmanager, nullcontext +from types import MappingProxyType +from typing import TYPE_CHECKING, Final + +import httpx +from fastapi import APIRouter, HTTPException, Request, Response, WebSocket, WebSocketDisconnect +from pydantic import BaseModel, Field, JsonValue, TypeAdapter +from starlette.types import Message + +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.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import ( + can_key_call_resolved_model, # pyright: ignore[reportUnknownVariableType] # legacy authorization accepts untyped deployment lists +) +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.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 +) + +_routes: Final = APIRouter() +_JSON: Final = TypeAdapter[JsonValue](JsonValue) +_EMPTY: Final[Mapping[str, JsonValue]] = MappingProxyType({}) +_MAPPING: Final = TypeAdapter(Mapping[str, object]) +_OBJECT: Final = TypeAdapter(Mapping[str, JsonValue]) +_DEPLOYMENT: Final = TypeAdapter(LiveDeployment) +_PREFIX: Final = "live_litellm_" + + +def _json_value(value: object) -> JsonValue: + if isinstance(value, Mapping): + entries: Final = _MAPPING.validate_python(value) + return {key: _json_value(item) for key, item in entries.items()} # mutable-ok: JSON wire objects require dicts + if isinstance(value, (tuple, list)): + items: Final = TypeAdapter(tuple[object, ...]).validate_python(value) + return [_json_value(item) for item in items] # mutable-ok: JSON wire arrays require lists + return _JSON.validate_python(value) + + +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()).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]) -> Request: + async def receive() -> Message: + return _mutable( + MappingProxyType({"type": "http.request", "body": _encode_json(body).encode(), "more_body": False}) + ) + + return Request(_mutable(MappingProxyType({**source.scope, "type": "http", "method": "POST"})), 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 _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, + ) + + +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 _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})))) + + +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: + values: Final = _object( + auth.model_dump( + include=MappingProxyType( + { + name: True + for name in ( + "model_max_budget", + "user_model_max_budget", + "end_user_model_max_budget", + "rpm_limit_per_model", + "tpm_limit_per_model", + "rpm_limit", + "tpm_limit", + "team_rpm_limit", + "team_tpm_limit", + "user_rpm_limit", + "user_tpm_limit", + "team_metadata", + "metadata", + "organization_metadata", + "project_metadata", + ) + } + ) + ) + ) + + def constrained(value: JsonValue | Mapping[str, JsonValue]) -> bool: + if not isinstance(value, Mapping): + return False + return any( + bool(item) + if any(marker in key for marker in ("rpm_limit", "tpm_limit", "model_max_budget")) + else constrained(item) + for key, item in value.items() + ) + + return constrained(values) + + +def _restricted_models(auth: UserAPIKeyAuth) -> bool: + key_models: Final = TypeAdapter(tuple[str, ...]).validate_python(getattr(auth, "models", ()) or ()) + team_models: Final = TypeAdapter(tuple[str, ...]).validate_python(getattr(auth, "team_models", ()) or ()) + return any( + models and "*" not in models and "all-proxy-models" not in models for models in (key_models, team_models) + ) + + +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 + if _managed_constraints(auth): + raise HTTPException( + 400, + "Managed Live delegation cannot enforce configured backend model budgets or rate limits; use client delegation", + ) + if not isinstance(responses, dict): + if _restricted_models(auth): + raise HTTPException(400, "Restricted keys require an explicit authorized delegation.responses.model") + return + model: Final = responses.get("model") + 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 "session.update" 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 (_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: + 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) + auth: Final = await _auth( + _request(request, _EMPTY if token else MappingProxyType({**body, "model": _session_model(body)})) + ) + 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 = _pinned(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() + ) # rebind-ok: each received event has a new value + 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 = _pinned(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: + pass + except Exception: + try: + await websocket.close(code=1011, reason="Live upstream connection failed") + except RuntimeError: + pass + 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/types/llms/openai.py b/litellm/types/llms/openai.py index f448244706c..c5e9fe360f0 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2037,10 +2037,15 @@ class OpenAIRealtimeStreamResponseBaseObject(TypedDict): class OpenAIRealtimeSessionClosed(TypedDict): - type: ReadOnly[Literal["session.closed"]] + 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"]] @@ -2295,6 +2300,7 @@ class OpenAIRealtimeEventTypes(Enum): OpenAIRealtimeEvents = ( OpenAIRealtimeStreamResponseBaseObject | OpenAIRealtimeSessionClosed + | OpenAILiveResponseEvent | OpenAIRealtimeStreamSessionEvents | OpenAIRealtimeStreamResponseOutputItemAdded | OpenAIRealtimeResponseContentPartAdded diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index 4c3df385ab9..b0655f36fad 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: ( diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 7c514db3e38..b3a41d6cf81 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3590,3 +3590,43 @@ 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", + "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("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/test_litellm/llms/chatgpt/test_live.py b/tests/test_litellm/llms/chatgpt/test_live.py new file mode 100644 index 00000000000..2e1a35473ef --- /dev/null +++ b/tests/test_litellm/llms/chatgpt/test_live.py @@ -0,0 +1,194 @@ +import json + +import httpx +import pytest + +from litellm.llms.chatgpt.live import LiveDeployment, LiveTransport, live_session_path + + +@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 httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + 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_client=client, + ) + 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 client.is_closed + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["fork", "accept", "reject", "refer", "hangup", "content"]) +async def test_live_all_http_operations(operation): + 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(204) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client) + 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 == 204 + + +@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 httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client) + 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 httpx.AsyncClient(transport=httpx.MockTransport(respond), follow_redirects=True) as client: + transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client) + 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 diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 8f41fc22d71..fc45ca7c699 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -572,7 +572,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", @@ -966,11 +970,24 @@ def test_get_model_from_request_extracts_realtime_session_model(route, encoded): @pytest.mark.parametrize("session", ['{"model":"actual-voice"}', {"model": "actual-voice"}]) -def test_realtime_calls_auth_uses_executed_session_model_despite_decoys(session): +@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="/v1/realtime/calls", + route=route, request_query_params={"model": "query-decoy"}, request_headers={"x-litellm-model": "header-decoy"}, ) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index dffd4500ed7..cf3be1e40b8 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -178,7 +178,9 @@ async def test_mixed_case_offer_preserves_boundary_metadata_and_closes_extra_fil @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): +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 @@ -557,8 +559,9 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id @pytest.mark.parametrize("multipart", [False, True]) @pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom"]) @pytest.mark.parametrize("signaling_credential", ["authorization", "api-key", "x-litellm-api-key", "mixed"]) +@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 + monkeypatch, multipart, credential, signaling_credential, signaling_path ): import json from unittest.mock import AsyncMock @@ -596,7 +599,7 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers( { "type": "http", "method": "POST", - "path": "/v1/realtime/calls", + "path": signaling_path, "scheme": "http", "server": ("localhost", 80), "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", @@ -672,6 +675,8 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers( 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") @@ -1147,7 +1152,9 @@ async def test_signaling_rejection_after_admission_refunds_parallel_slot(monkeyp 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) + 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") diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 6603b4a0cdb..4455ef25cd7 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -187,7 +187,8 @@ class Sink: @pytest.mark.parametrize( "duration,valid", [(0, True), (1000, True), (None, False), (-1, False), (True, False), ("1000", False)] ) -async def test_live_terminal_requires_valid_duration_for_accounting(monkeypatch, duration, valid): +@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() @@ -201,7 +202,7 @@ async def test_live_terminal_requires_valid_duration_for_accounting(monkeypatch, await socket.messages.put({"type": "session.started"}) await supervisor.start() await socket.messages.put( - {"type": "session.closed", **({"usage": {"audio_duration_ms": duration}} if duration is not None else {})} + {"type": "session.closed", **({"usage": {duration_field: duration}} if duration is not None else {})} ) await supervisor.wait() close.assert_not_awaited() @@ -695,3 +696,58 @@ async def test_ga_observer_disconnect_requires_confirmed_hangup(closure, hangup_ 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() 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..a1bc9106f84 --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -0,0 +1,930 @@ +import json +import time +from contextlib import asynccontextmanager +from types import MappingProxyType, SimpleNamespace +from unittest.mock import AsyncMock, Mock + +import httpx +import pytest +from fastapi import FastAPI, HTTPException, Request +from fastapi.testclient import TestClient + +from litellm.llms.chatgpt.live import LiveDeployment +from litellm.proxy._types import 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 handle(owner="owner"): + return live._new_handle( + "sess_upstream", + "voice", + LiveDeployment(model="gpt-live"), + 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_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): + 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 = [] + + async def authenticate(request): + authenticated_bodies.append(await request.json()) + 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) + 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, + ) + + +@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_fork_preserves_empty_overrides_and_pins_source_deployment(route_client): + source = handle() + source = source.model_copy(update={"deployment": {**source.deployment, "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() + + +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}, + {"model_max_budget": {"backend": 1}}, + {"team_tpm_limit": 10}, + {"team_metadata": {"model_rpm_limit": {"backend": 1}}}, + ], +) +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 +async def test_restricted_webrtc_cannot_change_managed_model_outside_proxy(monkeypatch): + monkeypatch.setattr(live, "_authorize", AsyncMock()) + body = { + "session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}, + "transport": {"type": "webrtc", "sdp": "offer"}, + } + auth = UserAPIKeyAuth(api_key="owner", models=["voice", "backend"]) + 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_fork_can_explicitly_select_client_delegation(route_client): + route_client.auth.models = ["voice"] + body = {"session": {"delegation": {"type": "client"}}} + token = live.encode_session(handle()) + 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.parametrize("protocol", ["http", "websocket"]) +@pytest.mark.parametrize("limits", [{"rpm_limit": 10}, {"tpm_limit": 100}, {"model_max_budget": {"backend": 1}}]) +def test_fork_with_new_limits_cannot_trust_old_client_policy(route_client, protocol, limits): + from starlette.websockets import WebSocketDisconnect + + # An unrestricted source could have switched to managed delegation after its handle was issued. + for key, value in limits.items(): + setattr(route_client.auth, key, value) + token = live.encode_session(handle()) + path = f"/v1/live/sessions/{token}/fork" + if protocol == "http": + response = route_client.client.post(path, json={"session": {}}) + 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": {}}) + with pytest.raises(WebSocketDisconnect) as rejected: + ws.receive_json() + assert rejected.value.code == 1008 + route_client.factory.assert_not_called() + + +@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 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 82f2ef097aa..3fdcf3d4b52 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.""" 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..d280b9b1a30 --- /dev/null +++ b/tests/test_litellm/proxy/test_live_route_registration.py @@ -0,0 +1,57 @@ +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", {}) + 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/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 6c4260d5d3e..0364e609a30 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -1,4 +1,3 @@ - import json from pathlib import Path from typing import Final @@ -1822,7 +1821,6 @@ def test_azure_ai_cache_cost_calculation(_local_model_cost_map): ), 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", @@ -1889,6 +1887,7 @@ def test_azure_gpt_5_6_rates_match_azure_price_page(_local_model_cost_map, model 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): """ Regression for https://github.com/BerriAI/litellm/issues/34393: two Vertex @@ -4221,6 +4220,8 @@ def test_completion_cost_together_metadata_only_model_still_uses_size_bucket(_lo ) 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. @@ -4864,13 +4865,16 @@ def test_live_terminal_duration_uses_configured_second_price(monkeypatch, rate, ) == pytest.approx(expected) -def test_live_terminal_duration_honors_deployment_override(monkeypatch): +@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(), [_live_terminal_event()]) + 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", @@ -5170,3 +5174,216 @@ def test_completion_cost_ocr_ignores_deployment_pricing_without_custom_pricing_f litellm_logging_obj=logging_obj, ) 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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 97545f327a5..1e845ad8847 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -8372,7 +8372,7 @@ export interface paths { * WebSocket: realtime_websocket_endpoint * @description WebSocket connection endpoint */ - get: operations["websocket_realtime_websocket_endpoint_get_4"]; + get: operations["websocket_realtime_websocket_endpoint_get_5"]; put?: never; post?: never; delete?: never; @@ -8381,6 +8381,145 @@ export interface paths { 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; @@ -9780,6 +9919,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; @@ -18410,7 +18708,7 @@ export interface paths { * WebSocket: realtime_websocket_endpoint * @description WebSocket connection endpoint */ - get: operations["websocket_realtime_websocket_endpoint_get_5"]; + get: operations["websocket_realtime_websocket_endpoint_get_6"]; put?: never; post?: never; delete?: never; @@ -18430,7 +18728,7 @@ export interface paths { * WebSocket: codex_live_sideband_endpoint * @description WebSocket connection endpoint */ - get: operations["websocket_codex_live_sideband_endpoint"]; + get: operations["websocket_codex_live_sideband_endpoint_get"]; put?: never; post?: never; delete?: never; @@ -18439,6 +18737,125 @@ export interface paths { 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; @@ -51059,7 +51476,7 @@ export interface operations { }; }; }; - websocket_realtime_websocket_endpoint_get_4: { + websocket_realtime_websocket_endpoint_get_5: { parameters: { query?: never; header?: never; @@ -51077,6 +51494,230 @@ export interface operations { }; }; }; + 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; @@ -53252,6 +53893,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; @@ -63397,7 +64280,7 @@ export interface operations { }; }; }; - websocket_realtime_websocket_endpoint_get_5: { + websocket_realtime_websocket_endpoint_get_6: { parameters: { query?: never; header?: never; @@ -63415,7 +64298,7 @@ export interface operations { }; }; }; - websocket_codex_live_sideband_endpoint: { + websocket_codex_live_sideband_endpoint_get: { parameters: { query?: never; header?: never; @@ -63433,6 +64316,212 @@ export interface operations { }; }; }; + 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; From 153ece276972ae4faa05bb54ab1996e1061fab4c Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 17 Sep 2026 11:10:57 +0200 Subject: [PATCH 46/90] fix(live): enforce delegated model grants and managed budgets --- docs/my-website/docs/providers/chatgpt.md | 4 +- litellm/proxy/auth/auth_checks.py | 21 +- litellm/proxy/realtime_endpoints/live.py | 523 +++++++++++++++-- ...test_batch_embed_content_transformation.py | 1 + .../proxy/realtime_endpoints/test_live.py | 536 +++++++++++++++++- 5 files changed, 1005 insertions(+), 80 deletions(-) diff --git a/docs/my-website/docs/providers/chatgpt.md b/docs/my-website/docs/providers/chatgpt.md index bc971be087e..025223632c9 100644 --- a/docs/my-website/docs/providers/chatgpt.md +++ b/docs/my-website/docs/providers/chatgpt.md @@ -114,8 +114,8 @@ Session controls require a session known to the proxy and owned by the authentic Live duration uses cumulative `usage.seconds`; legacy Codex milliseconds remain supported. WebRTC initialization has a 15-second minimum credited against running duration, not added to it. Nested terminal Responses usage is charged separately using its backend model and deduplicated by response ID. A failed observation connection cannot establish complete usage. Managed delegation also depends on receiving its backend usage events; the upstream sideband does not replay events emitted before attachment -Managed Responses delegation authorizes its backend model as well as the voice model. Use client delegation when a key's per-model budgets or token/request limits require admission checks for each backend invocation. Restricted-model WebRTC keys must explicitly exclude `session.update` from frontend client events when using managed delegation, because that data channel connects directly to OpenAI and could otherwise change the backend model outside the proxy's checks +Managed Responses delegation authorizes its backend model as well as the voice model. Use client delegation when budgets or request/token limits apply to the key, user, team, project, organization, team member, end user, or a model access group, since each backend invocation needs its own admission check. Keys scoped to access groups, projects, users, organizations, or teams are treated as model-restricted even if the key's own model list is empty. Managed WebRTC sessions with model restrictions must explicitly exclude `session.update` and wildcard events from frontend client events. That data channel connects directly to OpenAI and could otherwise change the backend model outside the proxy's checks. Client delegation does not need this restriction: the delegation type cannot change after startup or on a fork. Sparse sideband updates may omit the backend model to retain its current value -For both HTTP and WebSocket forks, restricted-model keys must explicitly set `session.delegation` to `{ "type": "client" }` or `{ "type": "responses", "responses": { "model": "authorized-backend" } }`. Empty overrides cannot safely authorize an inherited backend: the session handle records startup configuration, while later updates may have changed the upstream model. Upstream rules still determine which delegation overrides a source session permits +For both HTTP and WebSocket forks of managed sessions, restricted-model keys must explicitly provide an authorized `session.delegation.responses.model`. Empty overrides cannot safely authorize an inherited managed backend: the session handle records startup configuration, while later updates may have changed the upstream model. Client-delegation forks can use empty overrides because the delegation type is immutable See the official [Live overview](https://developers.openai.com/api/docs/guides/live), [Live API reference](https://developers.openai.com/api/reference/resources/live), [session management](https://developers.openai.com/api/docs/guides/live-conversations), [WebRTC guide](https://developers.openai.com/api/docs/guides/voice-webrtc?api=live), [WebSocket guide](https://developers.openai.com/api/docs/guides/voice-websockets?api=live), [server controls](https://developers.openai.com/api/docs/guides/voice-server-controls?api=live) and [SIP guide](https://developers.openai.com/api/docs/guides/voice-sip?api=live) for the upstream contract. The voice guides also contain Realtime tabs with different routes and formats diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 576585ee9a3..0766c320d04 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2149,6 +2149,7 @@ 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 = False, ) -> Optional["LiteLLM_TeamMembership"]: """ Returns team membership object if user is member of team. @@ -2202,6 +2203,8 @@ async def get_team_membership( user_id, team_id, ) + if raise_on_error: + raise return None @@ -4122,6 +4125,8 @@ async def _team_member_granted_models( prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, + *, + strict_grant_lookup: bool = False, ) -> Sequence[str]: """The member's own ``allowed_models`` scope; empty when the member is not narrowed below the team.""" if team_object is None or valid_token.user_id is None: @@ -4133,6 +4138,7 @@ async def _team_member_granted_models( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + raise_on_error=strict_grant_lookup, ) return () if team_membership is None else _member_allowed_models(team_membership) @@ -4143,6 +4149,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) @@ -4158,6 +4166,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 () @@ -4169,6 +4179,8 @@ async def _granted_model_lists( prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, + *, + strict_grant_lookup: bool = False, ) -> tuple[Sequence[str], ...]: """One model allowlist per level that participates in authorizing the request.""" return ( @@ -4180,6 +4192,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, ), project_object.models if project_object is not None else (), await _org_granted_models( @@ -4188,6 +4201,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, ), ) @@ -4274,6 +4288,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, ) -> tuple[str, ...]: """ The budgeted model access groups that authorized this request, sorted and deduplicated. @@ -4289,7 +4305,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 () @@ -4319,6 +4337,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, ) for granted_model in granted_models ) diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 77e732f65fa..ec9d3db173e 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -19,12 +19,32 @@ if TYPE_CHECKING: 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.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.models.budget import LiteLLM_BudgetTable +from litellm.models.team import LiteLLM_TeamTable +from litellm.models.team_membership import LiteLLM_TeamMembership +from litellm.proxy._types import ( + LiteLLM_ProjectTableCachedObj, + LiteLLM_TeamTableCachedObj, + LitellmUserRoles, + UserAPIKeyAuth, +) from litellm.proxy.auth.auth_checks import ( 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_org_object, + 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.cache_pydantic_utils import CacheCodec 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 ( + NO_TEAM_MEMBERSHIP_SENTINEL, + model_access_group_cache_key, + team_membership_reservation_cache_key, +) from litellm.proxy.hooks.parallel_request_limiter import ( _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # limiter class is the existing hook identity ) @@ -38,6 +58,11 @@ from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, 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.project_repository import ProjectRepository +from litellm.repositories.table_repositories import ModelAccessGroupBudgetRepository, TeamMembershipRepository +from litellm.repositories.team_repository import TeamRepository +from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget _routes: Final = APIRouter() _JSON: Final = TypeAdapter[JsonValue](JsonValue) @@ -45,17 +70,45 @@ _EMPTY: Final[Mapping[str, JsonValue]] = MappingProxyType({}) _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: - if isinstance(value, Mapping): - entries: Final = _MAPPING.validate_python(value) - return {key: _json_value(item) for key, item in entries.items()} # mutable-ok: JSON wire objects require dicts - if isinstance(value, (tuple, list)): - items: Final = TypeAdapter(tuple[object, ...]).validate_python(value) - return [_json_value(item) for item in items] # mutable-ok: JSON wire arrays require lists - return _JSON.validate_python(value) + 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 # rebind-ok: each visited input produces a new JSON value + if isinstance(source, Mapping): + entries: Mapping[str, object] = _MAPPING.validate_python( + source + ) # rebind-ok: entries belong to the current node + 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 + ) # rebind-ok: items belong to the current node + array: list[JsonValue] = [None] * len(items) # mutable-ok: JSON output; # rebind-ok: per-node buffer + 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]: @@ -182,6 +235,24 @@ def _session_model(body: Mapping[str, JsonValue], fallback: str | None = None) - 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 @@ -194,6 +265,38 @@ async def _authorize(model: str, auth: UserAPIKeyAuth) -> None: 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 @@ -379,54 +482,362 @@ def _session_policy(body: Mapping[str, JsonValue], source: LiveHandle | None) -> def _managed_constraints(auth: UserAPIKeyAuth) -> bool: - values: Final = _object( - auth.model_dump( - include=MappingProxyType( - { - name: True - for name in ( - "model_max_budget", - "user_model_max_budget", - "end_user_model_max_budget", - "rpm_limit_per_model", - "tpm_limit_per_model", - "rpm_limit", - "tpm_limit", - "team_rpm_limit", - "team_tpm_limit", - "user_rpm_limit", - "user_tpm_limit", - "team_metadata", - "metadata", - "organization_metadata", - "project_metadata", - ) - } - ) + 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_budget, + auth.team_max_budget, + auth.user_max_budget, + auth.end_user_max_budget, + auth.organization_max_budget, ) + ): + 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 - def constrained(value: JsonValue | Mapping[str, JsonValue]) -> bool: - if not isinstance(value, Mapping): - return False - return any( - bool(item) - if any(marker in key for marker in ("rpm_limit", "tpm_limit", "model_max_budget")) - else constrained(item) - for key, item in value.items() + 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() # mutable-ok: cycle guard for hook-provided metadata + while pending: + current: object = pending.pop() # rebind-ok: advance the explicit metadata traversal stack + 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 + ) # rebind-ok: entries belong to the current metadata node + 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 - return constrained(values) + +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: - key_models: Final = TypeAdapter(tuple[str, ...]).validate_python(getattr(auth, "models", ()) or ()) - team_models: Final = TypeAdapter(tuple[str, ...]).validate_python(getattr(auth, "team_models", ()) or ()) - return any( - models and "*" not in models and "all-proxy-models" not in models for models in (key_models, team_models) + 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 + membership_key: Final = team_membership_reservation_cache_key(user_id=auth.user_id, team_id=auth.team_id) + membership_cached_raw: Final[object] = _OBJECT_VALUE.validate_python( + await server.user_api_key_cache.async_get_cache(key=membership_key) + ) + membership_cached: Final = ( + CacheCodec.deserialize(membership_cached_raw, model_type=LiteLLM_TeamMembership) + if membership_cached_raw is not None and membership_cached_raw != NO_TEAM_MEMBERSHIP_SENTINEL + else None + ) + if membership_cached is not None or membership_cached_raw == NO_TEAM_MEMBERSHIP_SENTINEL: + return membership_cached + return await TeamMembershipRepository(server.prisma_client).table.find_unique( + where={ # mutable-ok: Prisma serializes query filters from concrete dictionaries + "user_id_team_id": { # mutable-ok: Prisma serializes nested filters from concrete dictionaries + "user_id": auth.user_id, + "team_id": auth.team_id, + } + }, + include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries + ) + + +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_from_cache: Final = await server.user_api_key_cache.async_get_cache( + key=f"team_id:{auth.team_id}", model_type=LiteLLM_TeamTableCachedObj + ) + if team_from_cache is not None: + return team_from_cache + return await TeamRepository(server.prisma_client).find_by_id(auth.team_id, id_field="team_id") + + +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 + + default_cached: Final = await server.user_api_key_cache.async_get_cache( + key=f"team_member_default_budget:{default_id}", model_type=LiteLLM_BudgetTable + ) + if default_cached is not None: + return default_cached + return await BudgetRepository(server.prisma_client).find_by_id(default_id, id_field="budget_id") + + +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 + project_from_cache: Final = await server.user_api_key_cache.async_get_cache( + key=f"project_id:{auth.project_id}", model_type=LiteLLM_ProjectTableCachedObj + ) + if project_from_cache is not None: + return project_from_cache + project_row: Final = await ProjectRepository(server.prisma_client).table.find_unique( + where={"project_id": auth.project_id}, # mutable-ok: Prisma serializes concrete query dictionaries + include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries + ) + if project_row is None: + return None + return LiteLLM_ProjectTableCachedObj.model_validate(project_row.model_dump()) + + +async def _live_project_budget_configured(auth: UserAPIKeyAuth, project: LiteLLM_ProjectTableCachedObj | None) -> bool: + if project is None: + return False + from litellm.proxy import proxy_server as server + + project_budget: Final = getattr(project, "litellm_budget_table", None) + project_budget_id: Final = getattr(project, "budget_id", None) + project_budget_from_db: Final = ( + await BudgetRepository(server.prisma_client).find_by_id(project_budget_id, id_field="budget_id") + if project_budget is None and isinstance(project_budget_id, str) + else None + ) + if _live_budget_configured(project_budget or project_budget_from_db, 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}))) + + +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 + cached_values: Final = await asyncio.gather( + *( + server.user_api_key_cache.async_get_cache( + key=model_access_group_cache_key(group), model_type=ModelAccessGroupBudget + ) + for group in matched_groups + ) + ) + cached_groups: Final = tuple(zip(matched_groups, cached_values)) + uncached_groups: Final = tuple(group for group, budget in cached_groups if budget is None) + named_group_rows: Final = ( + await ModelAccessGroupBudgetRepository(server.prisma_client).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 + "in": uncached_groups, + } + }, + include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries + ) + if uncached_groups + else () + ) + return any( + _live_budget_configured( + budget + if budget is not None + else next( + ( + getattr(row, "litellm_budget_table", None) + for row in named_group_rows + if getattr(row, "access_group_name", None) == group + ), + None, + ), + zero_is_limit=False, + ) + for group, budget in cached_groups + ) + + +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): @@ -437,26 +848,32 @@ async def _authorize_delegation(body: Mapping[str, JsonValue], auth: UserAPIKeyA responses: Final = delegation.get("responses") if delegation.get("type") != "responses" and not isinstance(responses, dict): return - if _managed_constraints(auth): + 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 backend model budgets or rate limits; use client delegation", + "Managed Live delegation cannot enforce configured budgets or rate limits; use client delegation", ) if not isinstance(responses, dict): - if _restricted_models(auth): + if body.get("type") != "session.update" and _restricted_models(auth): raise HTTPException(400, "Restricted keys require an explicit authorized delegation.responses.model") return - model: Final = responses.get("model") 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 "session.update" in events: + 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" ) @@ -465,7 +882,11 @@ async def _authorize_delegation(body: Mapping[str, JsonValue], auth: UserAPIKeyA async def _authorize_fork_policy( body: Mapping[str, JsonValue], source: LiveHandle | None, auth: UserAPIKeyAuth ) -> None: - if source is not None and (_restricted_models(auth) or _managed_constraints(auth)): + 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) diff --git a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py index 86b3f0976ab..fb5ec9b11dc 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py @@ -299,6 +299,7 @@ class TestProcessResponse: ) +@pytest.mark.usefixtures("local_model_cost_map") class TestProcessEmbedContentResponseUsage: """Gemini Embedding 2 embedContent usageMetadata must drive spend. diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index a1bc9106f84..39592e8ef3d 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -8,8 +8,15 @@ 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 UserAPIKeyAuth from litellm.proxy.realtime_endpoints import live @@ -78,6 +85,30 @@ def test_handle_serializes_mappingproxy_without_losing_pinned_deployment(): 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", @@ -322,9 +353,19 @@ async def test_websocket_delegation_model_update_is_authorized_before_forwarding "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): @@ -336,13 +377,173 @@ async def test_managed_delegation_fails_closed_for_unenforceable_constraints(lim @pytest.mark.asyncio -async def test_restricted_webrtc_cannot_change_managed_model_outside_proxy(monkeypatch): +@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.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=SimpleNamespace(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", SimpleNamespace(async_get_cache=AsyncMock(return_value=None)) + ) + 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", SimpleNamespace(async_get_cache=AsyncMock(return_value=None)) + ) + 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(limits): + assert live._managed_constraints(UserAPIKeyAuth(api_key="owner", **limits)) is False + + +def test_managed_constraints_terminates_on_cyclic_metadata_without_a_limit(): + 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", models=["voice", "backend"]) + auth = UserAPIKeyAuth(api_key="owner", **scope) with pytest.raises(HTTPException) as rejected: await live._authorize_delegation(body, auth) assert rejected.value.status_code == 400 @@ -465,9 +666,7 @@ async def test_inherited_managed_fork_cannot_bypass_new_key_constraints(): @pytest.mark.parametrize("protocol", ["http", "websocket"]) -@pytest.mark.parametrize( - "startup_policy", [{}, {"delegation": {"type": "responses", "responses": {"model": "allowed"}}}] -) +@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 @@ -504,9 +703,9 @@ async def test_explicit_fork_backend_is_authorized_even_when_startup_policy_diff assert rejected.value.status_code == 403 -def test_restricted_fork_can_explicitly_select_client_delegation(route_client): +def test_restricted_client_fork_can_inherit_delegation(route_client): route_client.auth.models = ["voice"] - body = {"session": {"delegation": {"type": "client"}}} + body = {"session": {}} token = live.encode_session(handle()) route_client.transport.request.return_value = httpx.Response( 200, json={"session": {"id": "sess_fork"}, "transport": {"type": "webrtc", "sdp": "answer"}} @@ -517,26 +716,13 @@ def test_restricted_fork_can_explicitly_select_client_delegation(route_client): route_client.transport.request.assert_awaited_once_with("POST", "live/sessions/sess_upstream/fork", body=body) -@pytest.mark.parametrize("protocol", ["http", "websocket"]) +@pytest.mark.asyncio @pytest.mark.parametrize("limits", [{"rpm_limit": 10}, {"tpm_limit": 100}, {"model_max_budget": {"backend": 1}}]) -def test_fork_with_new_limits_cannot_trust_old_client_policy(route_client, protocol, limits): - from starlette.websockets import WebSocketDisconnect - - # An unrestricted source could have switched to managed delegation after its handle was issued. - for key, value in limits.items(): - setattr(route_client.auth, key, value) - token = live.encode_session(handle()) - path = f"/v1/live/sessions/{token}/fork" - if protocol == "http": - response = route_client.client.post(path, json={"session": {}}) - 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": {}}) - with pytest.raises(WebSocketDisconnect) as rejected: - ws.receive_json() - assert rejected.value.code == 1008 - route_client.factory.assert_not_called() +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 @@ -928,3 +1114,301 @@ async def test_live_observer_becomes_ready_without_session_started_event(monkeyp 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(Exception, match="user can only access"): + await live._authorize("backend", auth) + + 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(Exception, match="org can only access"): + await live._authorize("backend", auth) + + 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(): + 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 = 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) + + assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member")) is True + + +@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 = 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) + + 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", + 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 = 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(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", litellm_budget_table=SimpleNamespace(max_budget=1)), + SimpleNamespace( + access_group_name="backend-group", litellm_budget_table=SimpleNamespace(max_budget=backend_budget) + ), + ] + db = SimpleNamespace( + litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(return_value=rows)), + ) + 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, "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 + + +@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="team", organization_id="org", models=["*"]) + group = SimpleNamespace( + access_group_name="backend-group", + 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 = 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( + 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="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 From ccefa07cea1189d50c1492ab37def4c88a870fef Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 17 Sep 2026 11:23:37 +0200 Subject: [PATCH 47/90] test: pin base-model capability regression to bundled catalog --- .../test_get_supported_openai_params.py | 48 +++++-------------- 1 file changed, 12 insertions(+), 36 deletions(-) diff --git a/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py b/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py index 722818598af..2b5d1b6e5e4 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py +++ b/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py @@ -1,4 +1,3 @@ - import pytest @@ -33,9 +32,7 @@ def test_base_model_label_alone_lacks_bedrock_tools(): """The label by itself does not advertise tools; this is what made the union necessary. Guards against the discrepancy disappearing (and the regression test above silently passing for the wrong reason).""" - params = get_supported_openai_params( - model=BEDROCK_LABEL, custom_llm_provider="bedrock" - ) + params = get_supported_openai_params(model=BEDROCK_LABEL, custom_llm_provider="bedrock") assert params is not None assert "tools" not in params @@ -46,14 +43,8 @@ def test_base_model_is_additive_not_replacement(): Bedrock: real id supports ``tools`` but not the label's reasoning hint; the union must contain the real model's ``tools`` regardless of the label being a subset.""" - real_only = set( - get_supported_openai_params( - model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock" - ) - ) - label_only = set( - get_supported_openai_params(model=BEDROCK_LABEL, custom_llm_provider="bedrock") - ) + real_only = set(get_supported_openai_params(model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock")) + label_only = set(get_supported_openai_params(model=BEDROCK_LABEL, custom_llm_provider="bedrock")) combined = set( get_supported_openai_params( model=BEDROCK_REAL_MODEL, @@ -67,17 +58,14 @@ 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-3.1-pro`` isn't in the cost map so it advertises no reasoning support, + ``gemini-3.1-pro`` 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-3.1-pro", custom_llm_provider="gemini" - ) - ) + real_only = set(get_supported_openai_params(model="gemini-3.1-pro", custom_llm_provider="gemini")) assert "reasoning_effort" not in real_only combined = set( @@ -93,21 +81,15 @@ def test_base_model_adds_capabilities_the_real_model_lacks(): def test_no_base_model_is_unchanged(): """Omitting ``base_model`` must resolve purely from ``model``.""" - with_none = get_supported_openai_params( - model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock", base_model=None - ) - plain = get_supported_openai_params( - model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock" - ) + with_none = get_supported_openai_params(model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock", base_model=None) + plain = get_supported_openai_params(model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock") assert with_none == plain def test_base_model_equal_to_model_is_unchanged(): """A ``base_model`` identical to ``model`` must not double-resolve or reorder.""" - plain = get_supported_openai_params( - model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock" - ) + plain = get_supported_openai_params(model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock") same = get_supported_openai_params( model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock", @@ -152,14 +134,10 @@ def test_bedrock_converse_alias_resolves_like_bedrock(): params saw no Bedrock capabilities for a Converse model invoked via the alias.""" anthropic_model = "bedrock/converse/us.anthropic.claude-sonnet-4-6" - via_alias = get_supported_openai_params( - model=anthropic_model, custom_llm_provider="bedrock_converse" - ) + via_alias = get_supported_openai_params(model=anthropic_model, custom_llm_provider="bedrock_converse") assert via_alias is not None - assert via_alias == get_supported_openai_params( - model=anthropic_model, custom_llm_provider="bedrock" - ) + assert via_alias == get_supported_openai_params(model=anthropic_model, custom_llm_provider="bedrock") assert "web_search_options" not in via_alias assert "tools" in via_alias @@ -167,9 +145,7 @@ def test_bedrock_converse_alias_resolves_like_bedrock(): def test_bedrock_converse_alias_keeps_nova_web_search_options(): """Nova on the ``bedrock_converse`` alias still advertises web_search_options, proving the alias routes through the model-aware config rather than a blanket Bedrock default.""" - nova_params = get_supported_openai_params( - model="amazon.nova-pro-v1:0", custom_llm_provider="bedrock_converse" - ) + nova_params = get_supported_openai_params(model="amazon.nova-pro-v1:0", custom_llm_provider="bedrock_converse") assert nova_params is not None assert "web_search_options" in nova_params From fbace1f2c6ed99fe510247e9bca26d184d4e84ba Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 05:47:34 +0200 Subject: [PATCH 48/90] test: cover ChatGPT realtime and image routes --- .../llm_cost_calc/test_llm_cost_calc_utils.py | 35 + .../test_realtime_streaming.py | 16 + tests/test_litellm/llms/chatgpt/test_codex.py | 6 + .../test_litellm/llms/chatgpt/test_images.py | 54 ++ tests/test_litellm/llms/chatgpt/test_live.py | 28 + .../llms/chatgpt/test_realtime.py | 19 + .../custom_httpx/test_llm_http_handler.py | 61 ++ .../realtime/test_openai_realtime_handler.py | 9 + .../proxy/auth/test_user_api_key_auth.py | 64 +- .../hooks/test_parallel_request_limiter.py | 48 ++ .../hooks/test_parallel_request_limiter_v3.py | 46 ++ .../realtime_endpoints/test_call_sessions.py | 340 ++++++++- .../test_call_supervision.py | 105 +++ .../proxy/realtime_endpoints/test_live.py | 644 ++++++++++++++++++ tests/test_litellm/test_cost_calculator.py | 24 + 15 files changed, 1497 insertions(+), 2 deletions(-) diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index cbe6fe198c9..465c9dba64a 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -4410,6 +4410,41 @@ def test_image_response_input_image_tokens_priced_at_image_rate(details_as_dict) expected = 19 * 5e-6 + 512 * 8e-6 + 158 * 3e-5 assert cost is not None 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/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index b3a41d6cf81..045af34ec8a 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3425,6 +3425,22 @@ async def test_live_attachment_does_not_dispatch_duplicate_usage(): 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): diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 35606931e24..59107d9acff 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -55,3 +55,9 @@ def test_signaling_preserves_selected_model_for_sideband(extra_query): 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/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 7cd8d727a94..3d966688395 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -227,3 +227,57 @@ def test_image_routes_use_configured_gateway(monkeypatch, env_name, api_base, tm 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/test_litellm/llms/chatgpt/test_live.py b/tests/test_litellm/llms/chatgpt/test_live.py index 2e1a35473ef..765b1723a5f 100644 --- a/tests/test_litellm/llms/chatgpt/test_live.py +++ b/tests/test_litellm/llms/chatgpt/test_live.py @@ -192,3 +192,31 @@ async def test_live_websocket_does_not_redirect_credentials(): 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 httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + transport = LiveTransport( + LiveDeployment("deployment-model", provider="openai", api_key="deployment-key", api_base=api_base), + {}, + http_client=client, + ) + with pytest.raises(ValueError, match="Invalid Live API base"): + await transport.request("POST", "live/sessions", {}) + assert requests == [] diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 2e3d091c088..4f367dd8229 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -475,3 +475,22 @@ async def test_supervisor_connection_preserves_call_routing(model, chatgpt_token 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/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index bed15399fbb..2f58155eb72 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -25,6 +25,7 @@ from litellm.llms.base_llm.search.transformation import BaseSearchConfig, Search from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS from litellm.llms.brave.search.transformation import BraveSearchConfig from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig +from litellm.llms.base_llm.image_generation.transformation import BaseImageGenerationConfig from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import ( BaseLLMHTTPHandler, @@ -3772,3 +3773,63 @@ async def test_realtime_http_sessions_preserve_provider_identity( 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" diff --git a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py index 4221954d787..2cde9acffc3 100644 --- a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py +++ b/tests/test_litellm/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/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index b19d7b1d192..18e4b4f98c0 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 @@ -364,7 +364,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: @@ -7755,3 +7755,65 @@ async def test_claude_view_never_reinterprets_explicit_names(monkeypatch, layer) assert data["model"] == ("foo" if layer == "unclaimed" else encoded) await _normalize_claude_model(data, token, request, "/v1/messages") 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": ""}) 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 589c8d1bd06..3b0ce1d4c7b 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -178,3 +178,51 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ litellm_parent_otel_span=None, ) 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 14d078820af..5c6a61cbc3f 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 @@ -6105,3 +6105,49 @@ async def test_an_open_circuit_breaker_reads_the_sliding_window_locally_without_ assert isinstance(values, list) assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == [] 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. diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index cf3be1e40b8..cd139563dc2 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -894,7 +894,11 @@ async def test_supervisor_policy_failure_hangs_up_before_releasing(monkeypatch): @pytest.mark.asyncio @pytest.mark.parametrize("hangup_fails", [False, True]) -async def test_supervisor_constructor_failure_closes_effective_connection(monkeypatch, hangup_fails, caplog): +@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 @@ -906,8 +910,12 @@ async def test_supervisor_constructor_failure_closes_effective_connection(monkey 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) @@ -946,6 +954,14 @@ async def test_supervisor_constructor_failure_closes_effective_connection(monkey 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 @@ -1312,3 +1328,325 @@ async def test_signaling_settles_tokens_once_with_isolated_sdk_callbacks(monkeyp 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 index 4455ef25cd7..e7581faab0d 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -751,3 +751,108 @@ async def test_live_missing_backend_accounting_invalidates_budget_after_dispatch 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 index 39592e8ef3d..4cc8f9dd16a 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -1412,3 +1412,647 @@ async def test_managed_budget_fails_closed_when_member_scope_lookup_fails_after_ ) 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", + SimpleNamespace( + async_get_cache=AsyncMock(return_value=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", + SimpleNamespace(async_get_cache=AsyncMock(return_value=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=["*"]) + monkeypatch.setattr( + proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=team)) + ) + assert await live._live_team(UserAPIKeyAuth(api_key="owner", team_id="team")) is team + + +@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) + monkeypatch.setattr( + proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=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 + + +@pytest.mark.asyncio +async def test_live_project_uses_cache_or_reports_missing_row(monkeypatch): + from litellm.proxy import proxy_server + + project = SimpleNamespace(project_id="project-1") + monkeypatch.setattr( + proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=project)) + ) + auth = UserAPIKeyAuth(api_key="owner", project_id="project-1") + assert await live._live_project(auth) is project + + monkeypatch.setattr( + proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=None)) + ) + monkeypatch.setattr( + live, + "ProjectRepository", + lambda client: SimpleNamespace(table=SimpleNamespace(find_unique=AsyncMock(return_value=None))), + ) + 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()) + 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/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 0364e609a30..a756224aa8c 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -5387,3 +5387,27 @@ def test_live_missing_backend_price_preserves_duration_and_marks_accounting_inco 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 From 7a26d0213a67dc9e5501a75bf0343ac085ffbfd1 Mon Sep 17 00:00:00 2001 From: kerry Date: Thu, 17 Sep 2026 23:56:05 +0000 Subject: [PATCH 49/90] build(deps): bump soupsieve to 2.9.2 to clear the osv-scan advisories Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- uv.lock | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/uv.lock b/uv.lock index eb4cdef76f1..e55968f0a35 100644 --- a/uv.lock +++ b/uv.lock @@ -9080,11 +9080,11 @@ wheels = [ [[package]] name = "soupsieve" -version = "2.8.4" +version = "2.9.2" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/47/2c/0a5f6f8ee0d5589e48c7640213ed5175d52cf540a06725b628cc1a45d6ce/soupsieve-2.8.4.tar.gz", hash = "sha256:e121fd02e975c695e4e9e8774a5ee35d74714b59307868dcc5319ad2d9e3328e", size = 121110, upload-time = "2026-05-24T13:55:57.154Z" } +sdist = { url = "https://files.pythonhosted.org/packages/69/99/a6ca3beb3ccacb41fb3321d8a60e5566f9e6467601ef8eba6a17e1b89778/soupsieve-2.9.2.tar.gz", hash = "sha256:4a55d8cf158a9c2e587fa4922f1bbb91d68ac829e2d6f25403a85747c71daf74", size = 122445, upload-time = "2026-08-07T00:57:24.801Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5e/f5/0c41cb68dcae6b7de4fac4188a3a9589e21fb31df21ea3a2e888db95e6c9/soupsieve-2.8.4-py3-none-any.whl", hash = "sha256:e7e6b0769c8f51ed59acab6e994b00621096cfb1c640a7509295987388fbaf65", size = 37304, upload-time = "2026-05-24T13:55:55.406Z" }, + { url = "https://files.pythonhosted.org/packages/eb/dc/ad025c1ee131eba60c69f4dd5779b18fcf1e6b21a343e2162a84d5d133c7/soupsieve-2.9.2-py3-none-any.whl", hash = "sha256:8089a26fd974ca7a1f30276d3d8492ab266ab15af581642dfe8aa162e0c1c823", size = 37370, upload-time = "2026-08-07T00:57:23.524Z" }, ] [[package]] From 865c759fc3adb52e58190cc299575649438514a3 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 06:10:56 +0200 Subject: [PATCH 50/90] Revert "build(deps): bump soupsieve to 2.9.2 to clear the osv-scan advisories" This reverts commit 7a26d0213a67dc9e5501a75bf0343ac085ffbfd1. --- uv.lock | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/uv.lock b/uv.lock index e55968f0a35..eb4cdef76f1 100644 --- a/uv.lock +++ b/uv.lock @@ -9080,11 +9080,11 @@ wheels = [ [[package]] name = "soupsieve" -version = "2.9.2" +version = "2.8.4" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/69/99/a6ca3beb3ccacb41fb3321d8a60e5566f9e6467601ef8eba6a17e1b89778/soupsieve-2.9.2.tar.gz", hash = "sha256:4a55d8cf158a9c2e587fa4922f1bbb91d68ac829e2d6f25403a85747c71daf74", size = 122445, upload-time = "2026-08-07T00:57:24.801Z" } +sdist = { url = "https://files.pythonhosted.org/packages/47/2c/0a5f6f8ee0d5589e48c7640213ed5175d52cf540a06725b628cc1a45d6ce/soupsieve-2.8.4.tar.gz", hash = "sha256:e121fd02e975c695e4e9e8774a5ee35d74714b59307868dcc5319ad2d9e3328e", size = 121110, upload-time = "2026-05-24T13:55:57.154Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/eb/dc/ad025c1ee131eba60c69f4dd5779b18fcf1e6b21a343e2162a84d5d133c7/soupsieve-2.9.2-py3-none-any.whl", hash = "sha256:8089a26fd974ca7a1f30276d3d8492ab266ab15af581642dfe8aa162e0c1c823", size = 37370, upload-time = "2026-08-07T00:57:23.524Z" }, + { url = "https://files.pythonhosted.org/packages/5e/f5/0c41cb68dcae6b7de4fac4188a3a9589e21fb31df21ea3a2e888db95e6c9/soupsieve-2.8.4-py3-none-any.whl", hash = "sha256:e7e6b0769c8f51ed59acab6e994b00621096cfb1c640a7509295987388fbaf65", size = 37304, upload-time = "2026-05-24T13:55:55.406Z" }, ] [[package]] From de113fb1fa1eda7f1039240f4a4b4105a668c25e Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 13:05:33 +0200 Subject: [PATCH 51/90] ci: preserve JWT complexity budget --- litellm/proxy/auth/handle_jwt.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 6a28cd7ff99..e14d84c9a1b 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -2380,7 +2380,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], From e6ab17cc482024d178ae03d36851ce43a8835ae8 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 13:25:36 +0200 Subject: [PATCH 52/90] fix: document realtime mutable payloads --- litellm/cost_calculator.py | 6 ++++-- litellm/images/main.py | 9 ++++++++- litellm/litellm_core_utils/realtime_streaming.py | 7 ++++++- litellm/proxy/proxy_server.py | 4 +++- litellm/proxy/realtime_endpoints/endpoints.py | 6 +++--- 5 files changed, 24 insertions(+), 8 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 4f04fc471b5..80f4d1f2133 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -2768,7 +2768,9 @@ def handle_realtime_stream_cost_calculation( potential_model_names=potential_model_names, combined_usage_object=( RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results( - [event for event in results if event.get("type") != "response.event"] + [ # 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 @@ -2832,7 +2834,7 @@ class _LiveBackendEnvelope(BaseModel): def _live_backend_responses( results: OpenAIRealtimeStreamList, logging_obj: LitellmLoggingObject | None = None ) -> tuple[ResponsesAPIResponse, ...]: - responses: Final = { + responses: Final = { # mutable-ok: deduplicate terminal backend responses by response id response.id: response for result in results if result.get("type") == "response.event" diff --git a/litellm/images/main.py b/litellm/images/main.py index 662903ee35e..64bc5d9b382 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -870,7 +870,14 @@ def image_edit( extra_body if isinstance(extra_body, dict) else None, ) if image_edit_provider_config.use_multipart_form_data() - else {**non_default_params, **(extra_body if isinstance(extra_body, dict) else {})} + 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 diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index f487e7513cf..850cfee2923 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -157,7 +157,12 @@ class RealTimeStreaming: self.messages: list[OpenAIRealtimeEvents] = [] if account_usage and live_initialization_seconds > 0: self.messages.append( - {"type": "litellm.live.initialization", "usage": {"seconds": live_initialization_seconds}} + { # 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 = {} diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2f3210dad0d..25354b9fe40 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -12176,7 +12176,9 @@ async def realtime_websocket_endpoint( # Only use explicit parameters, not all query params query_params: Final = cast( RealtimeQueryParams, - dict(_realtime_query_params_template(model, intent) + ((("call_id", call_id),) if call_id is not None else ())), + 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] = { diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index c1da14cb667..d66976d3e6f 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -362,9 +362,9 @@ async def create_realtime_client_secret( return RealtimeClientSecretResponse(**upstream_json) -@router.post("/v1/live", tags=["realtime"]) -@router.post("/live", tags=["realtime"]) -@router.post("/openai/v1/live", tags=["realtime"]) +@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 From 04f0d4b545a03694a1b6945b135e59feec1329c0 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 13:39:28 +0200 Subject: [PATCH 53/90] test: assert sanitized live access errors --- .../proxy/realtime_endpoints/test_call_sessions.py | 5 +++-- tests/test_litellm/proxy/realtime_endpoints/test_live.py | 8 +++++--- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index cd139563dc2..89e9c4efe6d 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -280,8 +280,9 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream( with pytest.raises(ProxyException) as denied: await proxy_realtime_calls(request, Response()) if policy == "personal_models": - assert "user not allowed to access model" in str(denied.value) - assert "forbidden-voice" in str(denied.value) + 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": diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index 4cc8f9dd16a..d4df56345cd 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -17,7 +17,7 @@ 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 UserAPIKeyAuth +from litellm.proxy._types import ModelAccessDeniedProxyException, UserAPIKeyAuth from litellm.proxy.realtime_endpoints import live @@ -1132,8 +1132,9 @@ async def test_authorize_enforces_authoritative_personal_user_models(monkeypatch auth = UserAPIKeyAuth(api_key="owner", models=[], user_id="user-only") await live._authorize("voice", auth) - with pytest.raises(Exception, match="user can only access"): + 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 @@ -1153,8 +1154,9 @@ async def test_authorize_enforces_authoritative_organization_models(monkeypatch, auth = UserAPIKeyAuth(api_key="owner", models=[], org_id="org-only") await live._authorize("voice", auth) - with pytest.raises(Exception, match="org can only access"): + 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 From 7f3455272e20f5621851d32bcf8135aa25255f5c Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 14:13:11 +0200 Subject: [PATCH 54/90] ci: avoid CodeQL log injection query overflow --- .github/codeql/codeql-config.yml | 5 +++++ 1 file changed, 5 insertions(+) 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. From a4accccac3604f729229e6d95a763c20031c731f Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 14:39:09 +0200 Subject: [PATCH 55/90] fix: address CodeQL realtime auth findings --- litellm/proxy/_types.py | 4 ++-- litellm/proxy/auth/auth_checks.py | 4 +++- litellm/proxy/realtime_endpoints/call_sessions.py | 10 ++++++++-- litellm/proxy/realtime_endpoints/live.py | 4 +++- litellm/proxy/utils.py | 4 ++-- tests/proxy_unit_tests/test_auth_checks.py | 2 ++ 6 files changed, 20 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 38fbe55d85a..741128b2769 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -264,8 +264,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 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 93a37f4a641..d82d69edf2a 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -6024,7 +6024,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/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index d1908919503..cf4937f3252 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -218,7 +218,10 @@ def decode_call(token: str, authorization: str) -> CodexRealtimeCall: 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()).hexdigest(): + 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 @@ -231,6 +234,7 @@ async def _cache_bounded_offer_body(request: Request) -> None: if int(request.headers.get("content-length", "")) > MAX_REALTIME_OFFER_BYTES: raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit") except ValueError: + # A missing or non-numeric content length is checked while streaming below. pass if hasattr(request, "_body"): if len(request._body) > MAX_REALTIME_OFFER_BYTES: # pyright: ignore[reportPrivateUsage] # validate Starlette's cached body without consuming it again @@ -396,7 +400,9 @@ async def _create_codex_realtime_call(request: Request) -> Response: call: Final = parse_call_response( response, alias=model, - owner=hashlib.sha256(f"Bearer {owner_key}".encode()).hexdigest(), + owner=hashlib.sha256( + f"Bearer {owner_key}".encode(), usedforsecurity=False + ).hexdigest(), expires_at=time.time() + 3600, ) except ValueError as exc: diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index ec9d3db173e..8f68e870497 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -185,7 +185,7 @@ def rewrite_session_ids(value: JsonValue | Mapping[str, JsonValue], raw_id: str, 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()).hexdigest() + return hashlib.sha256(auth.api_key.encode(), usedforsecurity=False).hexdigest() async def _auth(request: Request) -> UserAPIKeyAuth: @@ -1475,11 +1475,13 @@ async def websocket_live_session(websocket: WebSocket, session_id: str | None = try: await websocket.close(code=1008, reason="Live session rejected") except RuntimeError: + # The peer may have closed the socket before the rejection response. pass except Exception: try: await websocket.close(code=1011, reason="Live upstream connection failed") except RuntimeError: + # The peer may have closed the socket before the failure response. pass finally: if state.connection is not None: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 6cfacd02e86..dc0ba6bfbd1 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4291,7 +4291,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 @@ -6718,7 +6718,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/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 2538556d3b5..a5f00b01e0a 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/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) ], From 5d36f155b65e72869684d5980ed29e9f1bba7c1f Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 14:50:16 +0200 Subject: [PATCH 56/90] fix: suppress intentional ownership hash alerts --- litellm/proxy/_types.py | 1 + litellm/proxy/realtime_endpoints/call_sessions.py | 2 ++ litellm/proxy/realtime_endpoints/live.py | 1 + litellm/proxy/utils.py | 2 ++ 4 files changed, 6 insertions(+) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 741128b2769..f9045d20765 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -265,6 +265,7 @@ def hash_token(token: str): import hashlib # This digest is an opaque lookup identifier, not a password hash. + # codeql[py/weak-sensitive-data-hashing] hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index cf4937f3252..803c0998e5e 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -220,6 +220,7 @@ def decode_call(token: str, authorization: str) -> CodexRealtimeCall: raise HTTPException(403, "Invalid realtime call") from exc if ( call.expires_at < time.time() + # codeql[py/weak-sensitive-data-hashing] or call.owner != hashlib.sha256(authorization.encode(), usedforsecurity=False).hexdigest() ): raise HTTPException(403, "Invalid or expired realtime call") @@ -400,6 +401,7 @@ async def _create_codex_realtime_call(request: Request) -> Response: call: Final = parse_call_response( response, alias=model, + # codeql[py/weak-sensitive-data-hashing] owner=hashlib.sha256( f"Bearer {owner_key}".encode(), usedforsecurity=False ).hexdigest(), diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 8f68e870497..634425c029c 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -185,6 +185,7 @@ def rewrite_session_ids(value: JsonValue | Mapping[str, JsonValue], raw_id: str, def _owner(auth: UserAPIKeyAuth) -> str: if not auth.api_key: raise HTTPException(403, "Live sessions require an authenticated API key") + # codeql[py/weak-sensitive-data-hashing] return hashlib.sha256(auth.api_key.encode(), usedforsecurity=False).hexdigest() diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index dc0ba6bfbd1..6f4ef72f3cf 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4291,6 +4291,7 @@ class PrismaClient: def hash_token(self, token: str): # Hash the string using SHA-256 + # codeql[py/weak-sensitive-data-hashing] hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token @@ -6718,6 +6719,7 @@ def hash_token(token: str): import hashlib # Hash the string using SHA-256 + # codeql[py/weak-sensitive-data-hashing] hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token From 75e88ba9ada857e7b5b59842a5dca24d74b909ac Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 15:00:08 +0200 Subject: [PATCH 57/90] ci: filter intentional ownership hash alerts --- .github/workflows/codeql.yml | 5 +++++ litellm/proxy/_types.py | 1 - litellm/proxy/realtime_endpoints/call_sessions.py | 2 -- litellm/proxy/realtime_endpoints/live.py | 1 - litellm/proxy/utils.py | 2 -- 5 files changed, 5 insertions(+), 6 deletions(-) diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index d3a165a11da..83a7dbeb9bb 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -71,8 +71,13 @@ 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/_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/litellm/proxy/_types.py b/litellm/proxy/_types.py index f9045d20765..741128b2769 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -265,7 +265,6 @@ def hash_token(token: str): import hashlib # This digest is an opaque lookup identifier, not a password hash. - # codeql[py/weak-sensitive-data-hashing] hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 803c0998e5e..cf4937f3252 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -220,7 +220,6 @@ def decode_call(token: str, authorization: str) -> CodexRealtimeCall: raise HTTPException(403, "Invalid realtime call") from exc if ( call.expires_at < time.time() - # codeql[py/weak-sensitive-data-hashing] or call.owner != hashlib.sha256(authorization.encode(), usedforsecurity=False).hexdigest() ): raise HTTPException(403, "Invalid or expired realtime call") @@ -401,7 +400,6 @@ async def _create_codex_realtime_call(request: Request) -> Response: call: Final = parse_call_response( response, alias=model, - # codeql[py/weak-sensitive-data-hashing] owner=hashlib.sha256( f"Bearer {owner_key}".encode(), usedforsecurity=False ).hexdigest(), diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 634425c029c..8f68e870497 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -185,7 +185,6 @@ def rewrite_session_ids(value: JsonValue | Mapping[str, JsonValue], raw_id: str, def _owner(auth: UserAPIKeyAuth) -> str: if not auth.api_key: raise HTTPException(403, "Live sessions require an authenticated API key") - # codeql[py/weak-sensitive-data-hashing] return hashlib.sha256(auth.api_key.encode(), usedforsecurity=False).hexdigest() diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 6f4ef72f3cf..dc0ba6bfbd1 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4291,7 +4291,6 @@ class PrismaClient: def hash_token(self, token: str): # Hash the string using SHA-256 - # codeql[py/weak-sensitive-data-hashing] hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token @@ -6719,7 +6718,6 @@ def hash_token(token: str): import hashlib # Hash the string using SHA-256 - # codeql[py/weak-sensitive-data-hashing] hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest() return hashed_token From 3f36fe64204a9d0244394be085606a0e1a2a8836 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 15:37:03 +0200 Subject: [PATCH 58/90] fix(realtime): reject stale live deployment forks --- .../proxy/realtime_endpoints/call_sessions.py | 4 +- litellm/proxy/realtime_endpoints/live.py | 50 ++++++++++++- .../proxy/realtime_endpoints/test_live.py | 72 +++++++++++++++---- 3 files changed, 109 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index cf4937f3252..0827d4b5057 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -400,9 +400,7 @@ async def _create_codex_realtime_call(request: Request) -> Response: call: Final = parse_call_response( response, alias=model, - owner=hashlib.sha256( - f"Bearer {owner_key}".encode(), usedforsecurity=False - ).hexdigest(), + owner=hashlib.sha256(f"Bearer {owner_key}".encode(), usedforsecurity=False).hexdigest(), expires_at=time.time() + 3600, ) except ValueError as exc: diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 8f68e870497..663b31cc5fb 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -344,6 +344,48 @@ 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, @@ -1170,7 +1212,9 @@ async def _create(request: Request, token: str | None = None) -> Response: _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 = _pinned(source) if source else await _deployment(model, prepared.processed) + 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( @@ -1412,7 +1456,9 @@ async def websocket_live_session(websocket: WebSocket, session_id: str | None = ) else: await _authorize_fork_policy(_processed_body(first, prepared.processed), source, ownership.auth) - deployment: Final = _pinned(source) if source else await _deployment(model, prepared.processed) + 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 diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index d4df56345cd..950dd384054 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -26,11 +26,14 @@ def encryption_key(monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", "test-live-encryption-key") -def handle(owner="owner"): +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(model="gpt-live"), + LiveDeployment(**deployment), UserAPIKeyAuth(api_key=owner), None, ) @@ -128,6 +131,8 @@ def test_only_protocol_session_ids_are_rewritten_and_application_values_survive( @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( @@ -153,6 +158,21 @@ def route_client(monkeypatch): 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() @@ -198,8 +218,7 @@ def test_create_preserves_configuration_and_returns_owned_json_session(route_cli def test_fork_preserves_empty_overrides_and_pins_source_deployment(route_client): - source = handle() - source = source.model_copy(update={"deployment": {**source.deployment, "model_id": "deployment-a"}}) + 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) @@ -210,6 +229,39 @@ def test_fork_preserves_empty_overrides_and_pins_source_deployment(route_client) 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"}}) @@ -706,7 +758,7 @@ async def test_explicit_fork_backend_is_authorized_even_when_startup_policy_diff def test_restricted_client_fork_can_inherit_delegation(route_client): route_client.auth.models = ["voice"] body = {"session": {}} - token = live.encode_session(handle()) + 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"}} ) @@ -1942,9 +1994,7 @@ def test_websocket_requires_api_key_then_session_start(route_client): ws.receive_json() assert anonymous.value.code == 1008 - with route_client.client.websocket_connect( - "/v1/live/sessions", headers={"Authorization": "Bearer owner"} - ) as ws: + 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() @@ -1969,7 +2019,7 @@ async def test_attached_socket_authorizes_policy_and_reuses_source_without_start monkeypatch.setattr(live, "RealTimeStreaming", AttachedStream) authorize = AsyncMock() monkeypatch.setattr(live, "_authorize_delegation", authorize) - token = live.encode_session(handle()) + token = live.encode_session(handle(model_id="deployment-a")) inbound = iter([{"type": "websocket.connect"}, {"type": "websocket.disconnect", "code": 1000}]) sent: list = [] @@ -2012,9 +2062,7 @@ 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: + 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() From e0726859ef028e3c90821f3881edd7e2c86e12c1 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 15:58:54 +0200 Subject: [PATCH 59/90] test(proxy): isolate live route scheduler --- tests/test_litellm/proxy/test_live_route_registration.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/test_litellm/proxy/test_live_route_registration.py b/tests/test_litellm/proxy/test_live_route_registration.py index d280b9b1a30..80f3a5a9478 100644 --- a/tests/test_litellm/proxy/test_live_route_registration.py +++ b/tests/test_litellm/proxy/test_live_route_registration.py @@ -50,6 +50,8 @@ def test_public_live_websockets_reach_live_auth_before_legacy_sideband(monkeypat authenticate = AsyncMock(side_effect=HTTPException(403, "Live authentication rejected")) monkeypatch.setattr(live, "_auth", authenticate) monkeypatch.setattr(proxy_server, "general_settings", {}) + # 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"}): From 16a19f16c2f941e2f36bf8fdfb1233a123e8abf6 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 18 Sep 2026 20:42:12 +0200 Subject: [PATCH 60/90] fix(realtime): remove empty exception handlers --- litellm/proxy/realtime_endpoints/call_sessions.py | 8 +++++--- litellm/proxy/realtime_endpoints/live.py | 4 ++-- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 0827d4b5057..f2fb8e365e9 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -230,12 +230,14 @@ MAX_REALTIME_OFFER_BYTES: Final = 8 * 1024 * 1024 async def _cache_bounded_offer_body(request: Request) -> None: + content_length: int | None try: - if int(request.headers.get("content-length", "")) > MAX_REALTIME_OFFER_BYTES: - raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit") + content_length = int(request.headers.get("content-length", "")) except ValueError: # A missing or non-numeric content length is checked while streaming below. - pass + 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") diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 663b31cc5fb..7d75c2100fa 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -1522,13 +1522,13 @@ async def websocket_live_session(websocket: WebSocket, session_id: str | None = await websocket.close(code=1008, reason="Live session rejected") except RuntimeError: # The peer may have closed the socket before the rejection response. - pass + return except Exception: try: await websocket.close(code=1011, reason="Live upstream connection failed") except RuntimeError: # The peer may have closed the socket before the failure response. - pass + return finally: if state.connection is not None: await state.connection.close() From 0174b0a1f66a7650750bed676cf6259882a23f81 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sat, 19 Sep 2026 00:04:20 +0200 Subject: [PATCH 61/90] fix(chatgpt): repair image edits and checks after upstream merge --- litellm/cost_calculator.py | 4 +- litellm/images/main.py | 13 +++--- tests/test_litellm/proxy/test_proxy_server.py | 42 +++++++++---------- 3 files changed, 30 insertions(+), 29 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 80f4d1f2133..0aa041e0cba 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -2916,7 +2916,9 @@ def _live_duration_seconds(event: Mapping[str, object]) -> float | None: usage: Final = LiveSessionUsageEvent.model_validate(event).usage except ValidationError: return None - raw_usage: Final = cast(Mapping[str, object], event.get("usage")) + 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) diff --git a/litellm/images/main.py b/litellm/images/main.py index 530e11f62fd..c47bcfd3dce 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -864,11 +864,14 @@ 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, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 139fe2cce15..fdbaa94bf05 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -14026,32 +14026,28 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the @pytest.mark.asyncio -async def test_login_throttle_settings_are_not_hot_applied_from_the_database(): - """LIT-5285: a stored sign-in limit does not take effect on a live worker. - - _update_general_settings copies an allowlist of keys out of the DB row on every config - poll. Adding these to it would let a stored value outrank config.yaml without a restart, - so an operator locked out by a bad value could not fix it by editing YAML and restarting. - """ +async def test_login_throttle_config_settings_override_database(monkeypatch): import litellm.proxy.proxy_server as ps from litellm.proxy.proxy_server import ProxyConfig - original = dict(ps.general_settings) - try: - ps.general_settings.clear() - await ProxyConfig()._update_general_settings( - db_general_settings={ - "max_failed_login_attempts_per_source": 999, - "failed_login_window_seconds": 1, - "failed_login_block_seconds": 1, - } - ) - assert "max_failed_login_attempts_per_source" not in ps.general_settings - assert "failed_login_window_seconds" not in ps.general_settings - assert "failed_login_block_seconds" not in ps.general_settings - finally: - ps.general_settings.clear() - ps.general_settings.update(original) + 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 From fc305aa2a1a272c91e0df7cc0ce3d779d17b488b Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Tue, 22 Sep 2026 13:30:19 +0200 Subject: [PATCH 62/90] fix(chatgpt): keep membership reads fail-closed and clear the CI-only gates Four required checks were red on the merge commit. None of them came from a conflict hunk: each is a place where the merged tree disagrees with the code around it, and two only misbehave inside CI's environment. - `litellm/proxy/auth/auth_checks.py`: this PR gave `get_team_membership` a `raise_on_error` flag and defaulted it to False, so every caller -- all seven of them pre-existing -- started treating a failed membership read as "no membership row". Upstream main has no such flag and lets the exception surface, so the default quietly inverted fail-closed authorization into a grant: a Redis or database outage handed a team member the team's model scope and budget instead of rejecting the request. Default it to True, which covers the member budget check, the member model-scope check, JWT team resolution, and the compact summary gate, and name the one caller that genuinely wants the other behavior: `_team_member_granted_models` outside strict mode only attributes grants, so an unreadable member scope still degrades to "no member-level scope". - `tests/test_litellm/proxy/test_live_route_registration.py`: the websocket test runs the proxy lifespan, and the boot check now refuses an unset or weak master key before the app serves a single request, so the request never reached a Live route and all nine parametrizations failed where CI exports no key. Set one in the test rather than inherit whatever the environment has. A real key, not the weak-key opt-out, keeps the assertion honest: with a key in place the legacy sideband dependency rejects `Bearer test`, so awaiting `_auth` still proves the Live routes authenticate first. - `litellm/proxy/realtime_endpoints/call_supervision.py`: observer cleanup gathered a tuple whose length depended on a branch, which no `gather` overload can bind. That was the one `reportCallIssue` this PR added over the basedpyright budget ceiling. Two explicit branches gather the same tasks with the same cancellation and drain semantics. Correction to the merge commit message: it says the config-over-database login throttle test was dropped in favour of upstream's; both tests were actually kept, and both pass. Validation on this tree: 4546 passed, 1 skipped for the whole `proxy-auth` shard (auth, hooks, policy engine, client, realtime call redis) under `TZ=UTC`, which is the suite CI runs; the six membership fail-closed tests that were red now pass; `test_live_route_registration.py` passes 30 with the environment master key set, unset, and set to a rejected weak key; 546 passed across `tests/test_litellm/proxy/realtime_endpoints` and `tests/unit/realtime_api` after the gather change; `basedpyright` reports no `reportCallIssue` in `call_supervision.py`, and every remaining `reportCallIssue` in a file this PR touches sits on a line `git blame` traces to upstream main, so the rule returns to its base count of 113; `ruff check` and `ruff format --check` clean on the changed files. Local-only note, unrelated to this PR: `tests/test_litellm/proxy/hooks/test_batch_rate_limiter.py::test_cumulative_batch_tokens_over_tpd_returns_429_with_remaining_daily_window` fails under `TZ=Europe/Brussels` and passes under `TZ=UTC`. Upstream added it in `438d46cb50` and CI runs in UTC, so it is not affected by anything here. --- litellm/proxy/auth/auth_checks.py | 11 ++++++++++- litellm/proxy/realtime_endpoints/call_supervision.py | 10 ++++++---- .../proxy/test_live_route_registration.py | 5 +++++ 3 files changed, 21 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 2cabef85398..6526c70488e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2338,12 +2338,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 = False, + 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 @@ -4514,6 +4520,9 @@ 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) diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index 3f9ed521624..710220b17d2 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -186,10 +186,12 @@ class CallSupervisor: reader.cancel() if lease_failed is not None: lease_failed.cancel() - await asyncio.gather( - *((reader, stopped, lease_failed) if lease_failed is not None else (reader, stopped)), - return_exceptions=True, - ) + # 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(): diff --git a/tests/test_litellm/proxy/test_live_route_registration.py b/tests/test_litellm/proxy/test_live_route_registration.py index 80f3a5a9478..30d58d8a6e2 100644 --- a/tests/test_litellm/proxy/test_live_route_registration.py +++ b/tests/test_litellm/proxy/test_live_route_registration.py @@ -50,6 +50,11 @@ def test_public_live_websockets_reach_live_auth_before_legacy_sideband(monkeypat 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: From 35c00b7093897f227e7ccfccf25fc2ef0dcddad8 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Tue, 22 Sep 2026 13:55:22 +0200 Subject: [PATCH 63/90] test(proxy): stop the proxy budget from leaking between proxy tests `proxy-infra / Run tests` failed on four parametrizations of `test_real_proxy_child_auth_privacy_and_body_policy` -- the four that expect the child request to succeed. The traceback ends in upstream code this PR does not touch: `_user_api_key_auth_builder` calls `_fetch_global_spend_with_event_coordination`, whose loader reads `prisma_client.db.litellm_usertable`, and the fixture installs `object()` as the prisma client, so the read raises `'object' object has no attribute 'db'` and authentication returns 401. That block only runs when `litellm.max_budget > 0`, and the fixture never sets it, so the value came from a test that ran earlier on the same xdist worker. `test_add_proxy_budget_to_db_only_creates_user_no_keys` and `test_add_proxy_budget_to_db_backfills_budget_reset_at` assign `litellm.max_budget = 100.0` and `litellm.budget_duration` directly and never restore them, and which tests share a worker changes from run to run, which is why the same commit was green for these four cases and red for them later. Two small test fixes, no production change: - the two budget startup tests now set those globals through `monkeypatch`, so they are undone at teardown like the rest of the file's budget tests already do; - the compaction fixture pins `litellm.max_budget` to 0.0 beside the other pins, so it asserts its own precondition instead of depending on suite order. Validation: `test_native_compaction.py` plus both budget tests -- 16 passed; `ruff check --config ruff-tests.toml` clean on both files, and `ruff format --diff` leaves the added lines alone (the format check already reported this file before the change). --- .../proxy/test_native_compaction.py | 6 ++++++ tests/test_litellm/proxy/test_proxy_server.py | 17 ++++++++++------- 2 files changed, 16 insertions(+), 7 deletions(-) diff --git a/tests/test_litellm/proxy/test_native_compaction.py b/tests/test_litellm/proxy/test_native_compaction.py index d24624aacc5..b8062c6a4de 100644 --- a/tests/test_litellm/proxy/test_native_compaction.py +++ b/tests/test_litellm/proxy/test_native_compaction.py @@ -7,6 +7,7 @@ import pytest from fastapi import FastAPI, Request from pydantic import TypeAdapter +import litellm from litellm.caching.caching import DualCache from litellm.exceptions import BadRequestError from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy @@ -118,6 +119,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 84189cd3305..6d43870819d 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3200,7 +3200,7 @@ def test_normalize_datetime_for_sorting(): @pytest.mark.asyncio -async def test_add_proxy_budget_to_db_only_creates_user_no_keys(): +async def test_add_proxy_budget_to_db_only_creates_user_no_keys(monkeypatch: pytest.MonkeyPatch): """ Test that _add_proxy_budget_to_db only creates a user and no keys are added. @@ -3217,9 +3217,12 @@ async def test_add_proxy_budget_to_db_only_creates_user_no_keys(): import litellm from litellm.proxy.proxy_server import ProxyStartupEvent - # Set up required litellm settings - litellm.budget_duration = "30d" - litellm.max_budget = 100.0 + # 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) litellm_proxy_budget_name = "litellm-proxy-budget" @@ -3258,7 +3261,7 @@ async def test_add_proxy_budget_to_db_only_creates_user_no_keys(): @pytest.mark.asyncio -async def test_add_proxy_budget_to_db_backfills_budget_reset_at(): +async def test_add_proxy_budget_to_db_backfills_budget_reset_at(monkeypatch: pytest.MonkeyPatch): """ Test that _upsert_proxy_budget_with_reset_at_backfill issues a conditional update_many with `WHERE budget_reset_at IS NULL` to backfill the column on @@ -3276,8 +3279,8 @@ async def test_add_proxy_budget_to_db_backfills_budget_reset_at(): import litellm from litellm.proxy.proxy_server import ProxyStartupEvent - litellm.budget_duration = "30d" - litellm.max_budget = 100.0 + monkeypatch.setattr(litellm, "budget_duration", "30d") + monkeypatch.setattr(litellm, "max_budget", 100.0) litellm_proxy_budget_name = "litellm-proxy-budget" mock_prisma = MagicMock() From 1480b6b4888501f88f327d5d3e0ad7bde1c3cc15 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Tue, 22 Sep 2026 15:37:03 +0200 Subject: [PATCH 64/90] fix(realtime): count concurrency limits as Live managed constraints `_managed_constraints()` listed every rpm, tpm, and budget limit a key can carry and skipped `max_parallel_requests`, so a key with a configured concurrency limit could still choose managed Responses delegation. Those delegated invocations run upstream and never pass limiter admission, so the key ran concurrent backend calls past the cap it advertises. An admin-configured `global_max_parallel_requests` left the same hole for every key. Managed delegation now rejects either parallel limit being active, the way it already rejects the key's rpm, tpm, and budget limits, and the way `_live_budget_configured` already treats a budget table's `max_parallel_requests`. The tests that assert delegation stays unmanaged pin `general_settings`, so they no longer depend on what an earlier test left there. Validation: 550 passed across `tests/test_litellm/proxy/realtime_endpoints` and `tests/unit/realtime_api`, up from 546 with the four new cases; `basedpyright` reports 0 errors on `live.py`; `ruff check`, `ruff format --check`, and `scripts/test_quality_gate.py --base HEAD` clean. --- litellm/proxy/realtime_endpoints/live.py | 14 ++++++ .../proxy/realtime_endpoints/test_live.py | 45 +++++++++++++++++-- 2 files changed, 56 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 7d75c2100fa..cc1f932a1b7 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -524,6 +524,8 @@ def _session_policy(body: Mapping[str, JsonValue], source: LiveHandle | None) -> def _managed_constraints(auth: UserAPIKeyAuth) -> bool: + from litellm.proxy import proxy_server as server + if any( value is not None for value in ( @@ -539,6 +541,7 @@ def _managed_constraints(auth: UserAPIKeyAuth) -> bool: 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, @@ -548,6 +551,17 @@ def _managed_constraints(auth: UserAPIKeyAuth) -> bool: ): 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)), diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index 950dd384054..a395735ad12 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -465,6 +465,36 @@ 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", @@ -540,11 +570,17 @@ def test_managed_constraints_fails_closed_after_metadata_node_limit(): {"metadata": {"nested": [{"model_max_budget": {}}]}}, ], ) -def test_empty_model_limit_maps_do_not_mark_delegation_as_managed(limits): +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(): +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") @@ -1247,7 +1283,10 @@ async def test_sparse_responses_update_without_model_remains_valid_for_user_scop assert result is None -def test_managed_constraints_uses_exact_metadata_keys(): +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 From 68b573ea8d4632da87c4e1bfe030f53894da7ea9 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 25 Sep 2026 17:49:45 +0200 Subject: [PATCH 65/90] ci(realtime): clear the repo gates surfaced by the upstream merge Four repository gates flagged the merged tree; each fix is the narrowest one that keeps the PR behavior and the budgets untouched. - type-discipline LIT013 (upstream's new dead-suppression rule) flagged eleven `*-ok` reason comments in the Live/realtime modules that no longer suppress any rule: removed those comments only, keeping every suppression the checker still applies. budget_ratchet_check confirms no budget file moved. - type-discipline LIT014 (upstream's new one-for/one-if comprehension cap) flagged the double-for generator that flattened the Redis cluster acquire arguments; replaced it with an explicit loop appending the same (limit, ttl, slot) triple per key, identical wire payload. - basedpyright reportGeneralTypeIssues sat at 102/101: the cached-rate branch of calculate_image_response_cost_from_usage redeclared the `model_info` parameter as a `Final` local; the local is now `catalog_model_info`, and the four rate lookups read the same get_model_info entry as before. - upstream's pre_call_hook now always forwards skip_guardrails; the PR's PolicyHook test double accepts that keyword and still pins internal_realtime_observer in its assertions. Validation on this exact tree: 550 passed in the realtime/Live suites, 474 passed across the three cost-calculator suites, 715 passed for the v3 limiter plus the two auth files that exercise it, and the test-quality, ruff-strict, budget-ratchet, type-discipline, and basedpyright gates all pass with --base upstream/main. --- .../litellm_core_utils/llm_cost_calc/utils.py | 10 +++++----- litellm/llms/chatgpt/live.py | 4 ++-- litellm/llms/chatgpt/realtime.py | 2 +- .../proxy/hooks/parallel_request_limiter_v3.py | 12 ++++-------- litellm/proxy/realtime_endpoints/live.py | 16 ++++++++-------- .../realtime_endpoints/test_call_sessions.py | 4 +++- 6 files changed, 23 insertions(+), 25 deletions(-) diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index f8004b87062..bfab2d8f2df 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -1860,17 +1860,17 @@ def calculate_image_response_cost_from_usage( ) if cached_details is None: return prompt_cost + completion_cost - model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider) + catalog_model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider) cached_text: Final = _get_token_detail_value(cached_details, "text_tokens") or 0 cached_image: Final = _get_token_detail_value(cached_details, "image_tokens") or 0 input_text_tokens: Final = _get_token_detail_value(input_tokens_details, "text_tokens") or 0 input_image_tokens: Final = _get_token_detail_value(input_tokens_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 = model_info.get("input_cost_per_token") or 0.0 - image_rate: Final = model_info.get("input_cost_per_image_token") - cache_text_rate: Final = model_info.get("cache_read_input_token_cost") - cache_image_rate: Final = model_info.get("cache_read_input_image_token_cost") + 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) diff --git a/litellm/llms/chatgpt/live.py b/litellm/llms/chatgpt/live.py index 696ac00f118..01781b4c813 100644 --- a/litellm/llms/chatgpt/live.py +++ b/litellm/llms/chatgpt/live.py @@ -53,10 +53,10 @@ def _validate_session_id(session_id: str) -> None: 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") # rebind-ok: validate successive decoding layers iteratively + decoded: str = unquote(candidate, errors="strict") if decoded == candidate: return - candidate = decoded # rebind-ok: each percent-decoding pass reduces the input length + candidate = decoded def _validate_path(path: str) -> None: diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 6c275a1dec7..9fbf5d8ca57 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -245,7 +245,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): self, response: Response, model: str, model_id: str | None, headers: Mapping[str, object] | None ) -> Response: response.extensions["chatgpt_realtime"] = ( - MappingProxyType( # rebind-ok: HTTPX response extensions carry provider routing metadata + MappingProxyType( { "model": model, "model_id": model_id, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 77233d98553..65610e46e81 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -1733,15 +1733,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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( - arg - for key in keys - for arg in (by_key[key]["limit"], PARALLEL_REQUEST_SLOT_TTL_SECONDS, slot_id) - ), - ), + 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) diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index cc1f932a1b7..7ce17601ba7 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -85,18 +85,18 @@ def _json_value(value: object) -> JsonValue: source, parent, key, depth = pending.pop() if depth > 256: raise ValueError("Live JSON nesting exceeds the supported depth") - converted: JsonValue # rebind-ok: each visited input produces a new JSON value + converted: JsonValue if isinstance(source, Mapping): entries: Mapping[str, object] = _MAPPING.validate_python( source - ) # rebind-ok: entries belong to the current node + ) 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 - ) # rebind-ok: items belong to the current node - array: list[JsonValue] = [None] * len(items) # mutable-ok: JSON output; # rebind-ok: per-node buffer + ) + array: list[JsonValue] = [None] * len(items) pending.extend((item, array, index, depth + 1) for index, item in enumerate(items)) converted = array else: @@ -583,9 +583,9 @@ def _managed_constraints(auth: UserAPIKeyAuth) -> bool: ) if isinstance(value, (Mapping, list, tuple)) ] - visited: Final[set[int]] = set() # mutable-ok: cycle guard for hook-provided metadata + visited: Final[set[int]] = set() while pending: - current: object = pending.pop() # rebind-ok: advance the explicit metadata traversal stack + current: object = pending.pop() if id(current) in visited: continue visited.add(id(current)) @@ -594,7 +594,7 @@ def _managed_constraints(auth: UserAPIKeyAuth) -> bool: if isinstance(current, Mapping): entries: Mapping[str, object] = _MAPPING.validate_python( current - ) # rebind-ok: entries belong to the current metadata node + ) for key, item in entries.items(): if key in ( "rpm_limit", @@ -1411,7 +1411,7 @@ async def _wait_started( while True: event: Mapping[str, JsonValue] = _OBJECT.validate_json( await connection.recv() - ) # rebind-ok: each received event has a new value + ) if event.get("type") == "session.started": return event if startup is not None: diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 89e9c4efe6d..ca993f778f4 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -301,7 +301,9 @@ async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type, from litellm.proxy.realtime_endpoints.call_sessions import process_codex_request class PolicyHook: - async def pre_call_hook(self, user_api_key_dict, data, call_type, *, internal_realtime_observer=False): + 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") From 8b7fe007e1fc6f762400238fc9068955816ea29d Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 25 Sep 2026 19:55:06 +0200 Subject: [PATCH 66/90] fix(ci): repair PR 40366 lint and catalog gates --- .../crates/model-catalog/src/model_info.rs | 24 +++++++++++++++++++ litellm/experimental_mcp_client/client.py | 1 + litellm/llms/chatgpt/realtime.py | 18 +++++++------- ...odel_prices_and_context_window_backup.json | 4 ---- .../hooks/parallel_request_limiter_v3.py | 4 +--- litellm/proxy/realtime_endpoints/live.py | 16 ++++--------- model_prices_and_context_window.json | 4 ---- .../mcp_server/test_mcp_client_unit.py | 5 +++- 8 files changed, 42 insertions(+), 34 deletions(-) diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index 4a56e1112d1..b695fa18437 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -450,6 +450,30 @@ pub struct ModelInfo { pub output_cost_per_character_above_128k_tokens: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub output_cost_per_image: Option, + #[serde( + default, + rename = "output_cost_per_image_0.5K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_0_5k: Option, + #[serde( + default, + rename = "output_cost_per_image_1K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_1k: Option, + #[serde( + default, + rename = "output_cost_per_image_2K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_2k: Option, + #[serde( + default, + rename = "output_cost_per_image_4K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_4k: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub output_cost_per_image_1024: Option, #[serde(default, skip_serializing_if = "Option::is_none")] diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 1206f9abcbd..668bf104c84 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -869,6 +869,7 @@ class MCPClient: name=call_tool_request_params.name, arguments=call_tool_request_params.arguments, progress_callback=on_progress, + allow_input_required=False, ) try: diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 9fbf5d8ca57..f6aa04fc29d 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -244,16 +244,14 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): 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), - } - ) + 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 diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4b340963acf..10933090885 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -32254,7 +32254,6 @@ "gpt-image-2.5-flare": { "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, - "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -32270,7 +32269,6 @@ "gpt-image-2.5-flare-2026-09-08": { "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, - "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -32286,7 +32284,6 @@ "gpt-image-2.5-sunburst": { "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, - "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -32302,7 +32299,6 @@ "gpt-image-2.5-sunburst-2026-09-08": { "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, - "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 65610e46e81..ce6720291d0 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -1736,9 +1736,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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)), - ) + (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( diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 7ce17601ba7..39787c0a64f 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -87,15 +87,11 @@ def _json_value(value: object) -> JsonValue: raise ValueError("Live JSON nesting exceeds the supported depth") converted: JsonValue if isinstance(source, Mapping): - entries: Mapping[str, object] = _MAPPING.validate_python( - source - ) + 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 - ) + 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 @@ -592,9 +588,7 @@ def _managed_constraints(auth: UserAPIKeyAuth) -> bool: if len(visited) > 4096: return True if isinstance(current, Mapping): - entries: Mapping[str, object] = _MAPPING.validate_python( - current - ) + entries: Mapping[str, object] = _MAPPING.validate_python(current) for key, item in entries.items(): if key in ( "rpm_limit", @@ -1409,9 +1403,7 @@ async def _wait_started( ) -> Mapping[str, JsonValue]: async def receive_started() -> Mapping[str, JsonValue]: while True: - event: Mapping[str, JsonValue] = _OBJECT.validate_json( - await connection.recv() - ) + event: Mapping[str, JsonValue] = _OBJECT.validate_json(await connection.recv()) if event.get("type") == "session.started": return event if startup is not None: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4b340963acf..10933090885 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -32254,7 +32254,6 @@ "gpt-image-2.5-flare": { "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, - "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -32270,7 +32269,6 @@ "gpt-image-2.5-flare-2026-09-08": { "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, - "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -32286,7 +32284,6 @@ "gpt-image-2.5-sunburst": { "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, - "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", @@ -32302,7 +32299,6 @@ "gpt-image-2.5-sunburst-2026-09-08": { "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, - "cache_read_input_image_token_cost": 2e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py index 6438525706a..02f0c517cd9 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py @@ -289,7 +289,10 @@ class TestMCPClientUnitTests: assert result == mock_result mock_session_instance.initialize.assert_called_once() mock_session_instance.call_tool.assert_called_once_with( - name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY + name="test_tool", + arguments={"arg1": "value1"}, + progress_callback=ANY, + allow_input_required=False, ) From 6b5d2f093032809eeef246814a495cae095dba2b Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 25 Sep 2026 22:20:16 +0200 Subject: [PATCH 67/90] test(types): register ChatGPT connection params in the owned inventory The upstream inventory test now enumerates every owned litellm_params name, so the three ChatGPT connection options carried by this branch must be listed there. Without them the responses-caching-types shard fails on test_all_litellm_params_is_exactly_the_owned_inventory and its concatenation counterpart. --- tests/unit/types/test_litellm_params.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index e421321aaaa..ec194e41613 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", From 55c045fa86e6ce9756809977c27b18ad5611d21f Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 25 Sep 2026 22:44:02 +0200 Subject: [PATCH 68/90] fix(images): keep supports_pdf_input on the GPT Image 2.5 catalog rows An earlier commit on this branch rewrote those four model rows and dropped supports_pdf_input, so the branch silently removed a shipped capability from gpt-image-2.5-flare, its dated snapshot, gpt-image-2.5-sunburst and its dated snapshot, and repointed their source to a model page instead of the pricing page main uses. Both catalog copies stay byte-identical and the branch diff is now purely additive: the new chatgpt/gpt-live-1-codex realtime row. --- litellm/model_prices_and_context_window_backup.json | 12 ++++++++---- model_prices_and_context_window.json | 12 ++++++++---- 2 files changed, 16 insertions(+), 8 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 10933090885..897a54c9cb5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -32264,7 +32264,8 @@ "/v1/images/edits" ], "supports_vision": true, - "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-flare" + "supports_pdf_input": true, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-flare-2026-09-08": { "cache_read_input_image_token_cost": 2e-06, @@ -32279,7 +32280,8 @@ "/v1/images/edits" ], "supports_vision": true, - "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-flare" + "supports_pdf_input": true, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-sunburst": { "cache_read_input_image_token_cost": 2e-06, @@ -32294,7 +32296,8 @@ "/v1/images/edits" ], "supports_vision": true, - "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-sunburst" + "supports_pdf_input": true, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-sunburst-2026-09-08": { "cache_read_input_image_token_cost": 2e-06, @@ -32309,7 +32312,8 @@ "/v1/images/edits" ], "supports_vision": true, - "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-sunburst" + "supports_pdf_input": true, + "source": "https://developers.openai.com/api/docs/pricing" }, "low/1024-x-1024/gpt-image-1.5": { "deprecation_date": "2026-12-01", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 10933090885..897a54c9cb5 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -32264,7 +32264,8 @@ "/v1/images/edits" ], "supports_vision": true, - "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-flare" + "supports_pdf_input": true, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-flare-2026-09-08": { "cache_read_input_image_token_cost": 2e-06, @@ -32279,7 +32280,8 @@ "/v1/images/edits" ], "supports_vision": true, - "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-flare" + "supports_pdf_input": true, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-sunburst": { "cache_read_input_image_token_cost": 2e-06, @@ -32294,7 +32296,8 @@ "/v1/images/edits" ], "supports_vision": true, - "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-sunburst" + "supports_pdf_input": true, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-sunburst-2026-09-08": { "cache_read_input_image_token_cost": 2e-06, @@ -32309,7 +32312,8 @@ "/v1/images/edits" ], "supports_vision": true, - "source": "https://developers.openai.com/api/docs/models/gpt-image-2.5-sunburst" + "supports_pdf_input": true, + "source": "https://developers.openai.com/api/docs/pricing" }, "low/1024-x-1024/gpt-image-1.5": { "deprecation_date": "2026-12-01", From 5a3816a5ad7e920188f29436129559f12d5b5c92 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 25 Sep 2026 22:50:42 +0200 Subject: [PATCH 69/90] test(ci): drop the repeated UNIT_FLAG key in the shard paths test Upstream #43186 moved this file with "UNIT_FLAG" set twice in the same env literal, which fails F601 in the test-tree ruff config and turns the lint job red for every pull request that touches tests. Removing the duplicate keeps the effective environment identical. --- tests/unit/test_unit_shard_missing_paths.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/unit/test_unit_shard_missing_paths.py b/tests/unit/test_unit_shard_missing_paths.py index b528a75d0df..0360a227142 100644 --- a/tests/unit/test_unit_shard_missing_paths.py +++ b/tests/unit/test_unit_shard_missing_paths.py @@ -40,7 +40,6 @@ def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.Compl "TEST_PATH": test_path, "UNIT_FLAG": "", "WORKERS": workers, - "UNIT_FLAG": "", }, capture_output=True, text=True, From 8e30dd1be52d660ea9aa4660846f0baeb2978f4e Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sat, 26 Sep 2026 00:30:39 +0200 Subject: [PATCH 70/90] fix(live): serve authorization reads from cached auth helpers Live session authorization read the team-membership, team, project, and model-access-group budget tables through its own cache probe followed by a direct Prisma call, and never wrote the result back. Every signaling request therefore paid up to six database reads, and a caller that was not a team member re-queried the same absent row on every request. Route all six reads through the auth_checks helpers the chat path already uses, so Live shares their cache keys, TTLs, invalidation broadcasts, single-flight load and negative sentinel with the rest of the proxy: - membership through get_team_membership with raise_on_error left at its default, which also caches an absent row as NO_TEAM_MEMBERSHIP_SENTINEL - team through get_team_object, keeping a missing team row as no team rather than propagating the helper's HTTP 404 - the team's default member budget through get_team_member_default_budget - project through get_project_object, with no separate budget read: the helper includes the joined budget row, as the chat project checks assume - delegated model access group budgets through get_model_access_group_budgets_batch, so a group counts through max_budget the way the chat path enforces it instead of through rpm/tpm The gate stays conservative: any exception still raises HTTP 503 rather than letting managed Live delegation run with configured budgets unsupervised. --- litellm/proxy/realtime_endpoints/live.py | 138 +++++----------- .../proxy/realtime_endpoints/test_live.py | 148 ++++++++++++++---- 2 files changed, 152 insertions(+), 134 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 39787c0a64f..ec7f8f3ed0d 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -21,10 +21,8 @@ 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.models.team_membership import LiteLLM_TeamMembership from litellm.proxy._types import ( LiteLLM_ProjectTableCachedObj, - LiteLLM_TeamTableCachedObj, LitellmUserRoles, UserAPIKeyAuth, ) @@ -33,18 +31,16 @@ from litellm.proxy.auth.auth_checks import ( can_org_access_model, can_user_call_model, collect_matched_model_access_groups, + get_model_access_group_budgets_batch, get_org_object, + get_project_object, + get_team_member_default_budget, + 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.cache_pydantic_utils import CacheCodec 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 ( - NO_TEAM_MEMBERSHIP_SENTINEL, - model_access_group_cache_key, - team_membership_reservation_cache_key, -) from litellm.proxy.hooks.parallel_request_limiter import ( _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # limiter class is the existing hook identity ) @@ -58,11 +54,6 @@ from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, 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.project_repository import ProjectRepository -from litellm.repositories.table_repositories import ModelAccessGroupBudgetRepository, TeamMembershipRepository -from litellm.repositories.team_repository import TeamRepository -from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget _routes: Final = APIRouter() _JSON: Final = TypeAdapter[JsonValue](JsonValue) @@ -685,25 +676,12 @@ async def _live_team_membership(auth: UserAPIKeyAuth) -> object | None: if auth.team_id is None or auth.user_id is None: return None - membership_key: Final = team_membership_reservation_cache_key(user_id=auth.user_id, team_id=auth.team_id) - membership_cached_raw: Final[object] = _OBJECT_VALUE.validate_python( - await server.user_api_key_cache.async_get_cache(key=membership_key) - ) - membership_cached: Final = ( - CacheCodec.deserialize(membership_cached_raw, model_type=LiteLLM_TeamMembership) - if membership_cached_raw is not None and membership_cached_raw != NO_TEAM_MEMBERSHIP_SENTINEL - else None - ) - if membership_cached is not None or membership_cached_raw == NO_TEAM_MEMBERSHIP_SENTINEL: - return membership_cached - return await TeamMembershipRepository(server.prisma_client).table.find_unique( - where={ # mutable-ok: Prisma serializes query filters from concrete dictionaries - "user_id_team_id": { # mutable-ok: Prisma serializes nested filters from concrete dictionaries - "user_id": auth.user_id, - "team_id": auth.team_id, - } - }, - include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries + 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, ) @@ -712,12 +690,17 @@ async def _live_team(auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None: if auth.team_id is None: return None - team_from_cache: Final = await server.user_api_key_cache.async_get_cache( - key=f"team_id:{auth.team_id}", model_type=LiteLLM_TeamTableCachedObj - ) - if team_from_cache is not None: - return team_from_cache - return await TeamRepository(server.prisma_client).find_by_id(auth.team_id, id_field="team_id") + try: + return await get_team_object( + 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, + ) + except HTTPException as exc: + if exc.status_code == 404: + return None + raise def _live_team_budget_configured(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | None) -> bool: @@ -748,12 +731,12 @@ async def _live_default_budget(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | N return None from litellm.proxy import proxy_server as server - default_cached: Final = await server.user_api_key_cache.async_get_cache( - key=f"team_member_default_budget:{default_id}", model_type=LiteLLM_BudgetTable + # Like chat auth, a failed default-budget read returns None; membership errors still fail closed above. + return await get_team_member_default_budget( + default_id, + server.prisma_client, + server.user_api_key_cache, ) - if default_cached is not None: - return default_cached - return await BudgetRepository(server.prisma_client).find_by_id(default_id, id_field="budget_id") async def _live_project(auth: UserAPIKeyAuth) -> LiteLLM_ProjectTableCachedObj | None: @@ -761,33 +744,19 @@ async def _live_project(auth: UserAPIKeyAuth) -> LiteLLM_ProjectTableCachedObj | if auth.project_id is None: return None - project_from_cache: Final = await server.user_api_key_cache.async_get_cache( - key=f"project_id:{auth.project_id}", model_type=LiteLLM_ProjectTableCachedObj + 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, ) - if project_from_cache is not None: - return project_from_cache - project_row: Final = await ProjectRepository(server.prisma_client).table.find_unique( - where={"project_id": auth.project_id}, # mutable-ok: Prisma serializes concrete query dictionaries - include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries - ) - if project_row is None: - return None - return LiteLLM_ProjectTableCachedObj.model_validate(project_row.model_dump()) async def _live_project_budget_configured(auth: UserAPIKeyAuth, project: LiteLLM_ProjectTableCachedObj | None) -> bool: if project is None: return False - from litellm.proxy import proxy_server as server - project_budget: Final = getattr(project, "litellm_budget_table", None) - project_budget_id: Final = getattr(project, "budget_id", None) - project_budget_from_db: Final = ( - await BudgetRepository(server.prisma_client).find_by_id(project_budget_id, id_field="budget_id") - if project_budget is None and isinstance(project_budget_id, str) - else None - ) - if _live_budget_configured(project_budget or project_budget_from_db, zero_is_limit=True): + 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) @@ -823,44 +792,13 @@ async def _live_model_group_budget_configured( ) if not matched_groups: return False - cached_values: Final = await asyncio.gather( - *( - server.user_api_key_cache.async_get_cache( - key=model_access_group_cache_key(group), model_type=ModelAccessGroupBudget - ) - for group in matched_groups - ) - ) - cached_groups: Final = tuple(zip(matched_groups, cached_values)) - uncached_groups: Final = tuple(group for group, budget in cached_groups if budget is None) - named_group_rows: Final = ( - await ModelAccessGroupBudgetRepository(server.prisma_client).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 - "in": uncached_groups, - } - }, - include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries - ) - if uncached_groups - else () - ) - return any( - _live_budget_configured( - budget - if budget is not None - else next( - ( - getattr(row, "litellm_budget_table", None) - for row in named_group_rows - if getattr(row, "access_group_name", None) == group - ), - None, - ), - zero_is_limit=False, - ) - for group, budget in cached_groups + budgets: Final = await get_model_access_group_budgets_batch( + matched_groups, + server.prisma_client, + server.user_api_key_cache, ) + # Match chat auth: group budget rows contribute max_budget, not rpm/tpm, to this gate. + return any(_live_budget_configured(budget, zero_is_limit=False) for budget in budgets.values()) async def _managed_member_budget(auth: UserAPIKeyAuth, model: str | None = None) -> bool: diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index a395735ad12..7dddd2c215d 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -26,6 +26,18 @@ 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 + + return SimpleNamespace(async_get_cache=AsyncMock(side_effect=get), async_set_cache=AsyncMock(side_effect=set)) + + def handle(owner="owner", model_id=None): deployment = {"model": "gpt-live", "provider": "openai"} if model_id is not None: @@ -509,7 +521,9 @@ async def test_managed_delegation_checks_authoritative_member_and_default_budget from litellm.proxy import proxy_server membership = AsyncMock( - return_value=SimpleNamespace(litellm_budget_table=LiteLLM_BudgetTable(max_budget=member_limit)) + 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), @@ -519,9 +533,7 @@ async def test_managed_delegation_checks_authoritative_member_and_default_budget ), ) monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) - monkeypatch.setattr( - proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=None)) - ) + 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"}}}} @@ -544,9 +556,7 @@ async def test_managed_delegation_rejects_unverifiable_member_budget_but_allows_ 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", SimpleNamespace(async_get_cache=AsyncMock(return_value=None)) - ) + 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: @@ -1309,11 +1319,35 @@ async def test_managed_budget_reads_authoritative_member_budget(monkeypatch): litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(side_effect=find_unique)), litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=None)), ) - cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()) + 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 @@ -1324,7 +1358,7 @@ async def test_managed_budget_fails_closed_when_membership_repository_is_unreada db = SimpleNamespace( litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(side_effect=failure)), ) - cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()) + cache = _auth_cache() monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) @@ -1346,6 +1380,7 @@ async def test_managed_budget_checks_project_team_and_model_group_tables(monkeyp 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), ) @@ -1362,7 +1397,7 @@ async def test_managed_budget_checks_project_team_and_model_group_tables(monkeyp litellm_projecttable=SimpleNamespace(find_unique=AsyncMock(side_effect=find_project)), litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(side_effect=find_groups)), ) - cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()) + 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",))) @@ -1383,15 +1418,19 @@ async def test_managed_budget_uses_the_delegated_model_group(monkeypatch, backen from litellm.proxy import proxy_server rows = [ - SimpleNamespace(access_group_name="voice-group", litellm_budget_table=SimpleNamespace(max_budget=1)), + SimpleNamespace(access_group_name="voice-group", spend=0, litellm_budget_table=SimpleNamespace(max_budget=1)), SimpleNamespace( - access_group_name="backend-group", litellm_budget_table=SimpleNamespace(max_budget=backend_budget) + 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(return_value=rows)), + litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(side_effect=find_group_budgets)), ) - cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()) + 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()) @@ -1399,6 +1438,9 @@ async def test_managed_budget_uses_the_delegated_model_group(monkeypatch, backen 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"] == "model_access_group:backend-group" @pytest.mark.asyncio @@ -1406,19 +1448,19 @@ async def test_responses_delegation_fails_closed_when_inherited_org_lookup_fails from litellm.proxy import proxy_server from litellm.proxy.auth import auth_checks - team = LiteLLM_TeamTable(team_id="team", organization_id="org", models=["*"]) + 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 = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()) + 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, "proxy_logging_obj", None) monkeypatch.setattr( proxy_server, "llm_router", @@ -1434,7 +1476,7 @@ async def test_responses_delegation_fails_closed_when_inherited_org_lookup_fails 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"), + UserAPIKeyAuth(api_key="owner", models=["*"], team_id="org-lookup-team"), ) assert rejected.value.status_code == 503 @@ -1694,9 +1736,7 @@ async def test_live_team_membership_prefers_reservation_cache_and_sentinel(monke monkeypatch.setattr( proxy_server, "user_api_key_cache", - SimpleNamespace( - async_get_cache=AsyncMock(return_value=CacheCodec.serialize(membership, model_type=LiteLLM_TeamMembership)) - ), + _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" @@ -1704,7 +1744,7 @@ async def test_live_team_membership_prefers_reservation_cache_and_sentinel(monke monkeypatch.setattr( proxy_server, "user_api_key_cache", - SimpleNamespace(async_get_cache=AsyncMock(return_value=NO_TEAM_MEMBERSHIP_SENTINEL)), + _auth_cache({"team_membership:user:team": NO_TEAM_MEMBERSHIP_SENTINEL}), ) assert await live._live_team_membership(auth) is None @@ -1714,10 +1754,36 @@ 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, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=team)) + 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( @@ -1755,33 +1821,47 @@ 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, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=budget)) + 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 = SimpleNamespace(project_id="project-1") + project = LiteLLM_ProjectTable(project_id="project-1") + project_lookup = AsyncMock(return_value=project) monkeypatch.setattr( - proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=project)) + 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() - monkeypatch.setattr( - proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=None)) - ) - monkeypatch.setattr( - live, - "ProjectRepository", - lambda client: SimpleNamespace(table=SimpleNamespace(find_unique=AsyncMock(return_value=None))), - ) + 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 From 1a08de0af32f2025ad2804896ec5f0042d98a164 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sat, 26 Sep 2026 01:36:18 +0200 Subject: [PATCH 71/90] fix(live): keep budget reads loud and complete in the delegation gate The cached auth helpers cover the same rows this gate needs, but two of them answer a question this gate asks differently. get_team_object reports every failed database read as an HTTP 404, and the 404 mapping turned that into "no team"; get_team_member_default_budget returns None when its read raises. During an outage both convert a configured limit into an absent one, and the gate answers "no limits" by allowing managed Live delegation, so unsupervised spend becomes reachable through a failing database. The group-budget batch helper also flattens its row to spend and max_budget, dropping a linked group's rpm and tpm limits, and no other path in the proxy enforces those, so this gate was their only consumer. Read the team and default-budget rows through the proxy cache under the keys and TTL the rest of the proxy already uses, storing the entry on a miss and letting read errors reach the 503 handler, and fetch each group's linked budget in one batched query cached per group. The cached team entry keeps last_refreshed_at, so it is the same object the chat path writes. --- litellm/proxy/realtime_endpoints/live.py | 157 +++++++++++++++--- .../proxy/realtime_endpoints/test_live.py | 89 +++++++++- 2 files changed, 219 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index ec7f8f3ed0d..70ed3a6422f 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -3,10 +3,10 @@ import base64 import hashlib import json import time -from collections.abc import AsyncGenerator, Mapping +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping from contextlib import asynccontextmanager, nullcontext from types import MappingProxyType -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, TypeVar import httpx from fastapi import APIRouter, HTTPException, Request, Response, WebSocket, WebSocketDisconnect @@ -23,6 +23,7 @@ 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, ) @@ -31,16 +32,15 @@ from litellm.proxy.auth.auth_checks import ( can_org_access_model, can_user_call_model, collect_matched_model_access_groups, - get_model_access_group_budgets_batch, get_org_object, get_project_object, - get_team_member_default_budget, 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 from litellm.proxy.hooks.parallel_request_limiter import ( _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # limiter class is the existing hook identity ) @@ -54,10 +54,15 @@ from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, 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.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) +_LIVE_GROUP_LIMITS_CACHE_PREFIX: Final = "live:model_access_group_limits:" _MAPPING: Final = TypeAdapter(Mapping[str, object]) _OBJECT: Final = TypeAdapter(Mapping[str, JsonValue]) _DEPLOYMENT: Final = TypeAdapter(LiveDeployment) @@ -685,22 +690,57 @@ async def _live_team_membership(auth: UserAPIKeyAuth) -> object | None: ) +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 - try: - return await get_team_object( - 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, - ) - except HTTPException as exc: - if exc.status_code == 404: + team_id: Final = auth.team_id + + async def load() -> LiteLLM_TeamTableCachedObj | None: + row: Final = await TeamRepository(server.prisma_client).find_by_id(team_id, id_field="team_id") + if row is None: return None - raise + team: Final = LiteLLM_TeamTableCachedObj.model_validate(row.model_dump()) + team.last_refreshed_at = time.time() + return team + + return await _live_cached_object( + key=f"team_id:{team_id}", + model_type=LiteLLM_TeamTableCachedObj, + load=load, + ) def _live_team_budget_configured(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | None) -> bool: @@ -731,11 +771,13 @@ async def _live_default_budget(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | N return None from litellm.proxy import proxy_server as server - # Like chat auth, a failed default-budget read returns None; membership errors still fail closed above. - return await get_team_member_default_budget( - default_id, - server.prisma_client, - server.user_api_key_cache, + 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, ) @@ -769,6 +811,72 @@ async def _live_project_budget_configured(auth: UserAPIKeyAuth, project: LiteLLM 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 one query and cache one entry per group.""" + if not groups: + return () + from litellm.proxy import proxy_server as server + + rows: Final = await ModelAccessGroupBudgetRepository(server.prisma_client).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 + "in": list(groups), + } + }, + 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=f"{_LIVE_GROUP_LIMITS_CACHE_PREFIX}{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=f"{_LIVE_GROUP_LIMITS_CACHE_PREFIX}{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, @@ -792,13 +900,10 @@ async def _live_model_group_budget_configured( ) if not matched_groups: return False - budgets: Final = await get_model_access_group_budgets_batch( - matched_groups, - server.prisma_client, - server.user_api_key_cache, - ) - # Match chat auth: group budget rows contribute max_budget, not rpm/tpm, to this gate. - return any(_live_budget_configured(budget, zero_is_limit=False) for budget in budgets.values()) + # 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: diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index 7dddd2c215d..1d864fa2919 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -1440,7 +1440,94 @@ async def test_managed_budget_uses_the_delegated_model_group(monkeypatch, backen 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"] == "model_access_group:backend-group" + 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_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 From 03e541824e305fe7ebe140ba49f2ddb3d29d807d Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sat, 26 Sep 2026 01:46:46 +0200 Subject: [PATCH 72/90] docs: move the Codex gateway guide to litellm-docs The repository rule is that LiteLLM product documentation lives in litellm-docs rather than here, and the review on this pull request enforces it. The guide is now a purely additive section of the existing provider page in BerriAI/litellm-docs#1735, so the ChatGPT page keeps its single home and this branch stops shipping a second copy. --- docs/my-website/docs/providers/chatgpt.md | 121 ---------------------- 1 file changed, 121 deletions(-) delete mode 100644 docs/my-website/docs/providers/chatgpt.md diff --git a/docs/my-website/docs/providers/chatgpt.md b/docs/my-website/docs/providers/chatgpt.md deleted file mode 100644 index 025223632c9..00000000000 --- a/docs/my-website/docs/providers/chatgpt.md +++ /dev/null @@ -1,121 +0,0 @@ -# ChatGPT, Codex, and GPT-Live - -The proxy exposes the public GPT-Live session routes and retains the Codex-compatible `POST /live` route. Clients authenticate to LiteLLM with a LiteLLM virtual key. A `chatgpt` deployment uses one proxy-wide ChatGPT OAuth record, while an `openai` deployment uses its configured OpenAI API key. Do not send an upstream OAuth token as the proxy key - -## Configure LiteLLM deployments - -Codex text, image, and voice requests need separate LiteLLM aliases because their backends have different capabilities. The canonical aliases below keep the primary Qwen model separate from the ChatGPT OAuth models - -```yaml -model_list: - - model_name: qwen3.8-flash-next-codex - litellm_params: - model: openai/qwen3.8-flash-next - api_base: https://qwen.example.com/v1 - api_key: os.environ/QWEN_API_KEY - - - model_name: gpt-image-2 - litellm_params: - model: chatgpt/gpt-image-2 - - - model_name: gpt-image-2.5-flare - litellm_params: - model: chatgpt/gpt-image-2.5-flare - - - model_name: gpt-image-2.5-sunburst - litellm_params: - model: chatgpt/gpt-image-2.5-sunburst - - - model_name: gpt-realtime-1.5 - litellm_params: - model: chatgpt/gpt-realtime-1.5 - - - model_name: gpt-live-1-codex - litellm_params: - model: chatgpt/gpt-live-1-codex -``` - -If clients use shorter local aliases, publish separate aliases such as `qwen-codex`, `images`, and `voice` that point to the corresponding deployments. Authorize the exact alias sent by the client in the key or team policy, and for a restricted team member set `allowed_models` to contain the aliases used for text, image, or voice requests. Selecting the Qwen alias does not give it image-generation or voice capabilities. The Qwen endpoint in this example uses the OpenAI-compatible adapter and must expose `/v1/responses`; a native vLLM deployment can use `hosted_vllm/qwen3.8-flash-next` when that endpoint is available. LiteLLM does not promise provider-specific tool, reasoning, or stream behavior parity. For a deployment using the public OpenAI API instead, configure `model: openai/gpt-live-1` and `api_key: os.environ/OPENAI_API_KEY` under its own alias. The public API documentation uses `gpt-live-1`; the Codex alias and its backend capabilities are separate - -### Configure global ChatGPT OAuth - -The ChatGPT provider reads one auth file for the LiteLLM process. Set these environment variables before starting the proxy when the default location is not suitable - -```bash -export CHATGPT_TOKEN_DIR=/var/lib/litellm/chatgpt -export CHATGPT_AUTH_FILE=auth.json -``` - -The defaults are `~/.config/litellm/chatgpt` and `auth.json`. Persist the directory and complete the provider's device-code OAuth flow. The provider refreshes the stored record when it expires. All `chatgpt/...` deployments in that process share this record. Live rejects per-deployment `chatgpt_auth_profile`, `chatgpt_token_dir`, and `chatgpt_auth_file` overrides - -### Image generation and editing - -Use the `gpt-image-2` alias for both `/v1/images/generations` and `/v1/images/edits`. Keep the exact image aliases requested by your Codex version; a shorter `images` alias only works for clients configured to request it. Codex sends image requests to its active provider's `base_url`, with no separate image URL override. LiteLLM then selects the image deployment independently of the primary text model. ChatGPT image generation and editing require a valid global ChatGPT OAuth login, but a successful login does not establish that the account or backend supports every image operation. Image editing accepts JSON reference images and multipart files, but not masks. Image 2.5 aliases preserve the requested model name; an accepted name does not prove which backend model executed - -## Configure Codex through LiteLLM - -Codex sends its Responses requests to LiteLLM's `/v1/responses` endpoint. Point the Codex provider at the proxy and select the primary Qwen alias (or the shorter alias you published) - -```toml -model = "qwen3.8-flash-next-codex" -model_provider = "litellm" -experimental_realtime_ws_base_url = "https://litellm.example.com/v1" -experimental_realtime_webrtc_call_base_url = "https://litellm.example.com/v1" - -[model_providers.litellm] -name = "LiteLLM" -base_url = "https://litellm.example.com/v1" -wire_api = "responses" -requires_openai_auth = true -experimental_bearer_token = "" -``` - -Keep Codex signed in with ChatGPT for its client-side capability checks. `experimental_bearer_token` is the LiteLLM virtual key issued by the proxy. It must never contain the upstream ChatGPT OAuth access or refresh token. `requires_openai_auth = true` enables the Codex OpenAI-auth capability path while LiteLLM remains responsible for the upstream provider credentials - -## Configure voice routing - -The two realtime settings in the TOML example are root-level Codex settings, not fields inside `[model_providers.litellm]`. `experimental_realtime_ws_base_url` routes the Realtime WebSocket and its sideband through LiteLLM. `experimental_realtime_webrtc_call_base_url` is optional and separately routes HTTP WebRTC call creation. The optional root setting `experimental_realtime_ws_model` overrides the voice model; leave it unset to retain your client's default. An override must match the active protocol: `gpt-realtime-1.5` for legacy Realtime v1/v2 or `gpt-live-1-codex` for frameless Live v3. The base URLs end at `/v1`; LiteLLM adds the protocol-specific path - -| Voice operation | LiteLLM path | -| --- | --- | -| Legacy Realtime WebSocket | `/v1/realtime` | -| HTTP WebRTC call creation | `/v1/realtime/calls` | -| Frameless Live v3 signaling | `/v1/live` and `/v1/live/{call_id}` | -| Public Live session APIs | `/v1/live/sessions...` | - -WebSocket authentication uses `Authorization: Bearer ` by default. If the proxy sets `litellm_key_header_name`, send the virtual key in that configured header instead. Voice is experimental and pending retest: an observed mobile `POST /live` returned 201, but its sideband used the default `api.openai.com` and returned 404. The WebSocket base URL above addresses that routing gap; full bidirectional voice is not verified - -## Public Live routes - -Use the proxy host in place of `api.openai.com`. Send the configured LiteLLM alias in `session.model` when creating a session, or in the first `session.start` event for a primary WebSocket. Keep the returned session ID unchanged for subsequent operations - -| Method | Path | Request and response | -| --- | --- | --- | -| POST | `/v1/live/sessions` | JSON `session` and `transport: {type: "webrtc", sdp: ""}`; returns 201 JSON with `session.id` and `transport.sdp` | -| POST | `/v1/live/sessions/{session_id}/fork` | JSON WebRTC `transport` and optional `session` overrides; returns 200 JSON with the new session ID and SDP answer | -| GET | `/v1/live/sessions/{session_id}/content` | Downloads stored recording content without converting it to JSON | -| POST | `/v1/live/sessions/{session_id}/accept` | JSON `session` with `type: "live"` and model; successful SIP acceptance returns an empty body | -| POST | `/v1/live/sessions/{session_id}/reject` | JSON with required integer `status_code` from 300 through 699 | -| POST | `/v1/live/sessions/{session_id}/refer` | JSON with `target_uri` for the SIP destination | -| POST | `/v1/live/sessions/{session_id}/hangup` | No request body | -| WebSocket | `/v1/live/sessions` | Start with `session.start`, then wait for `session.started` before sending audio or commands | -| WebSocket | `/v1/live/sessions/{session_id}/attach` | Attach to an existing session; do not send `session.start` or input audio | -| WebSocket | `/v1/live/sessions/{session_id}/fork` | Start with `session.start` and a required `session` overrides object, which may be empty | - -Public WebRTC creation uses JSON, not the multipart or raw SDP formats used by the Codex compatibility route. `POST /live` and its existing aliases remain available for Codex clients using that format. WebRTC audio travels on media tracks; its data channel carries Live JSON events. Primary WebSocket audio uses base64 chunks in `session.input_audio.append` and `session.output_audio.delta` - -The proxy preserves Live event payloads, including nested Responses events inside `response.event`, rather than translating them into Realtime events. Session routing rewrites the configured model alias to the selected upstream model. Audio, transcript, delegation and usage events retain their upstream format. Send `session.close` and wait for `session.closed` to obtain final usage; a disconnected socket alone does not confirm successful finalization - -## Availability and verification - -Route support does not establish that every configured backend or account supports every operation. The official API describes project API-key authentication; it does not guarantee equivalent capabilities for ChatGPT OAuth. An OAuth request reaching SDP validation proves only that the request reached that validation step. It does not prove a working audio session, recording, fork or SIP call. The routes listed here have not all been tested against a real upstream service - -Session controls require a session known to the proxy and owned by the authenticated caller. Incoming SIP calls originate upstream. A proxy administrator can accept or reject the raw ID from a verified incoming-call webhook by supplying `x-litellm-live-model` with an alias that resolves to exactly one deployment. Successful acceptance returns the proxy-owned handle in `x-litellm-live-session-id`, preserving the API's empty response body. Use that handle for subsequent controls. Ordinary virtual keys cannot enroll arbitrary upstream session IDs; a trusted webhook-to-owner enrollment flow is still required for those keys - -Live duration uses cumulative `usage.seconds`; legacy Codex milliseconds remain supported. WebRTC initialization has a 15-second minimum credited against running duration, not added to it. Nested terminal Responses usage is charged separately using its backend model and deduplicated by response ID. A failed observation connection cannot establish complete usage. Managed delegation also depends on receiving its backend usage events; the upstream sideband does not replay events emitted before attachment - -Managed Responses delegation authorizes its backend model as well as the voice model. Use client delegation when budgets or request/token limits apply to the key, user, team, project, organization, team member, end user, or a model access group, since each backend invocation needs its own admission check. Keys scoped to access groups, projects, users, organizations, or teams are treated as model-restricted even if the key's own model list is empty. Managed WebRTC sessions with model restrictions must explicitly exclude `session.update` and wildcard events from frontend client events. That data channel connects directly to OpenAI and could otherwise change the backend model outside the proxy's checks. Client delegation does not need this restriction: the delegation type cannot change after startup or on a fork. Sparse sideband updates may omit the backend model to retain its current value - -For both HTTP and WebSocket forks of managed sessions, restricted-model keys must explicitly provide an authorized `session.delegation.responses.model`. Empty overrides cannot safely authorize an inherited managed backend: the session handle records startup configuration, while later updates may have changed the upstream model. Client-delegation forks can use empty overrides because the delegation type is immutable - -See the official [Live overview](https://developers.openai.com/api/docs/guides/live), [Live API reference](https://developers.openai.com/api/reference/resources/live), [session management](https://developers.openai.com/api/docs/guides/live-conversations), [WebRTC guide](https://developers.openai.com/api/docs/guides/voice-webrtc?api=live), [WebSocket guide](https://developers.openai.com/api/docs/guides/voice-websockets?api=live), [server controls](https://developers.openai.com/api/docs/guides/voice-server-controls?api=live) and [SIP guide](https://developers.openai.com/api/docs/guides/voice-sip?api=live) for the upstream contract. The voice guides also contain Realtime tabs with different routes and formats From c88ced1cae067cf6f2851f652fd468813528e1b7 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Sat, 26 Sep 2026 02:27:05 +0200 Subject: [PATCH 73/90] fix(live): keep the shared team entry complete and evict the Live group key Both gaps came from writing shared cache entries without the obligations that come with them. _live_team wrote team_id:{id} from a bare table row, while the chat path caches that key with the object_permission relation loaded. A request whose team read hit the entry Live had written would have seen a team stripped of the permissions it was about to enforce. The read now mirrors _get_team_object_from_user_api_key_cache, including its swallow-and-log degradation when the permission itself is unreadable, and the write goes through _cache_team_object so the alias-keyed entry is invalidated as usual. The Live group-limits entry was never evicted. Management writes already call _evict_model_access_group_cache_keys, which only knew the flattened entry and the registry, so a raised or lowered rpm or tpm limit kept deciding managed delegation until the entry's TTL expired. The key now comes from live_model_access_group_limits_cache_key, next to the other auth keys that must not drift, and is evicted alongside them. --- .../proxy/common_utils/user_api_key_cache.py | 11 ++++ ...model_access_group_management_endpoints.py | 10 +++- litellm/proxy/realtime_endpoints/live.py | 58 ++++++++++++++----- .../test_access_group_management.py | 12 ++-- .../proxy/realtime_endpoints/test_live.py | 26 +++++++++ 5 files changed, 95 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 61d7078ae4c..3fa26d3506d 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -313,6 +313,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/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/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 70ed3a6422f..596e1e1bb9c 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -13,6 +13,8 @@ from fastapi import APIRouter, HTTPException, Request, Response, WebSocket, WebS 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 @@ -28,10 +30,13 @@ from litellm.proxy._types import ( 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, @@ -40,7 +45,10 @@ from litellm.proxy.auth.auth_checks import ( ) 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 +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 ) @@ -62,7 +70,6 @@ _routes: Final = APIRouter() _JSON: Final = TypeAdapter[JsonValue](JsonValue) _EMPTY: Final[Mapping[str, JsonValue]] = MappingProxyType({}) _CACHEABLE_MODEL = TypeVar("_CACHEABLE_MODEL", bound=BaseModel) -_LIVE_GROUP_LIMITS_CACHE_PREFIX: Final = "live:model_access_group_limits:" _MAPPING: Final = TypeAdapter(Mapping[str, object]) _OBJECT: Final = TypeAdapter(Mapping[str, JsonValue]) _DEPLOYMENT: Final = TypeAdapter(LiveDeployment) @@ -727,20 +734,39 @@ async def _live_team(auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None: if auth.team_id is None: return None team_id: Final = auth.team_id - - async def load() -> LiteLLM_TeamTableCachedObj | None: - 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()) - team.last_refreshed_at = time.time() - return team - - return await _live_cached_object( + cached: Final = await _get_team_object_from_cache( key=f"team_id:{team_id}", - model_type=LiteLLM_TeamTableCachedObj, - load=load, + 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: @@ -848,7 +874,7 @@ async def _live_fetch_group_limits(groups: tuple[str, ...]) -> tuple[LiteLLM_Bud await asyncio.gather( *( server.user_api_key_cache.async_set_cache( - key=f"{_LIVE_GROUP_LIMITS_CACHE_PREFIX}{group}", + 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), @@ -866,7 +892,7 @@ async def _live_model_group_limits(groups: tuple[str, ...]) -> tuple[LiteLLM_Bud cached: Final = await asyncio.gather( *( server.user_api_key_cache.async_get_cache( - key=f"{_LIVE_GROUP_LIMITS_CACHE_PREFIX}{group}", + key=live_model_access_group_limits_cache_key(group), model_type=LiteLLM_BudgetTable, ) for group in groups 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_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index 1d864fa2919..768412f7570 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -1483,6 +1483,32 @@ async def test_managed_budget_fails_closed_when_the_team_row_is_unreadable(monke 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 From 8c93bf3ee44377f0462cb4aa7e523da4344639e1 Mon Sep 17 00:00:00 2001 From: Jordi Ibanez Date: Wed, 30 Sep 2026 12:23:14 +0200 Subject: [PATCH 74/90] fix(live): send Live HTTP through the shared async handler Greptile flagged that `LiveTransport.request` reached past LiteLLM's HTTP handler onto `AsyncHTTPHandler.client` to call `request` directly, which is the custom request path the repository rules ask us to avoid. The GET and POST verbs Live uses now go through `AsyncHTTPHandler.get` and `AsyncHTTPHandler.post`, with redirect-following turned off on the shared handler itself rather than per call, so the connection pool stays shared. The handler raises `MaskedHTTPStatusError` for a non-2xx POST, so that error is turned back into its `httpx.Response`: the proxy keeps answering with the upstream status, body and `x-request-id`, which the parametrised regressions now assert for 204, 404 and 503 on every Live operation. `http_client` became `http_handler` on the constructor, the only caller being the proxy, which never passed it. --- litellm/llms/chatgpt/live.py | 33 +++++----- tests/test_litellm/llms/chatgpt/test_live.py | 67 ++++++++++++++++---- 2 files changed, 71 insertions(+), 29 deletions(-) diff --git a/litellm/llms/chatgpt/live.py b/litellm/llms/chatgpt/live.py index 01781b4c813..23ea0e77630 100644 --- a/litellm/llms/chatgpt/live.py +++ b/litellm/llms/chatgpt/live.py @@ -17,6 +17,7 @@ from litellm.llms.chatgpt.realtime import ( 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, ) @@ -84,10 +85,10 @@ class LiveTransport: deployment: LiveDeployment, inbound_headers: Mapping[str, str], *, - http_client: httpx.AsyncClient | None = None, + http_handler: AsyncHTTPHandler | None = None, ) -> None: self.deployment = deployment - self._http_client = http_client + self._http_handler = http_handler params: Final = GenericLiteLLMParams.model_validate( MappingProxyType( { @@ -156,20 +157,22 @@ class LiveTransport: ): raise ValueError("Invalid Live HTTP operation") url: Final = self._url(path, query, websocket=False) - client: Final = ( - self._http_client - or get_async_httpx_client( - llm_provider=LlmProviders.CHATGPT if self.deployment.provider == "chatgpt" else LlmProviders.OPENAI - ).client - ) - return await client.request( - method, - url, - headers=MappingProxyType({**self._headers, "content-type": "application/json"}), - json=dict(body) if body is not None else None, # mutable-ok: JSON encoder requires a concrete dict - timeout=60, - follow_redirects=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 diff --git a/tests/test_litellm/llms/chatgpt/test_live.py b/tests/test_litellm/llms/chatgpt/test_live.py index 765b1723a5f..998639349a6 100644 --- a/tests/test_litellm/llms/chatgpt/test_live.py +++ b/tests/test_litellm/llms/chatgpt/test_live.py @@ -1,9 +1,21 @@ 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 @@ -32,7 +44,7 @@ async def test_live_request_preserves_payload_status_and_selected_credentials(pr 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 httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + async with live_handler(respond) as handler: transport = LiveTransport( LiveDeployment( model="deployment-model", @@ -43,7 +55,7 @@ async def test_live_request_preserves_payload_status_and_selected_credentials(pr extra_query={"gateway": "trusted", "tag": ("a +/&", "b"), "model": "bad", "session_id": "bad"}, ), {"Authorization": "Bearer proxy-key", "Cookie": "private", "OpenAI-Beta": "feature=v1"}, - http_client=client, + http_handler=handler, ) response = await transport.request( "POST", @@ -54,26 +66,53 @@ async def test_live_request_preserves_payload_status_and_selected_credentials(pr assert response.status_code == status assert response.json() == {"result": "upstream"} assert response.headers["x-request-id"] == "provider-id" - assert not client.is_closed + assert not handler.client.is_closed @pytest.mark.asyncio @pytest.mark.parametrize("operation", ["fork", "accept", "reject", "refer", "hangup", "content"]) -async def test_live_all_http_operations(operation): +@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(204) + return httpx.Response(status) - async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: - transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client) + 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 == 204 + 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 @@ -149,8 +188,8 @@ async def test_live_preserves_opaque_session_ids(session_id): assert request.url.params == httpx.QueryParams() return httpx.Response(200, json={"session_id": session_id}) - async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: - transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client) + 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 @@ -161,8 +200,8 @@ async def test_live_does_not_redirect_credentials(): assert request.url.host == "api.openai.com" return httpx.Response(307, headers={"location": "https://elsewhere.example/collect"}) - async with httpx.AsyncClient(transport=httpx.MockTransport(respond), follow_redirects=True) as client: - transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client) + 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 @@ -211,11 +250,11 @@ async def test_live_rejects_invalid_api_base_before_network(api_base): requests.append(request) return httpx.Response(200, json={}) - async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + async with live_handler(respond) as handler: transport = LiveTransport( LiveDeployment("deployment-model", provider="openai", api_key="deployment-key", api_base=api_base), {}, - http_client=client, + http_handler=handler, ) with pytest.raises(ValueError, match="Invalid Live API base"): await transport.request("POST", "live/sessions", {}) From 443a59f69f1c236acdbea87c96108216f9acdc2a Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 30 Sep 2026 13:44:48 +0200 Subject: [PATCH 75/90] test(unit): move the ChatGPT provider tests into the sharded unit tree Upstream emptied tests/test_litellm/llms and provider tests now live under tests/unit/llms, which the llm-other-providers shard already selects. The shard-coverage and ci-coverage guards fail while this PR keeps five files in the abandoned tree. --- tests/{test_litellm => unit}/llms/chatgpt/conftest.py | 0 tests/{test_litellm => unit}/llms/chatgpt/test_codex.py | 0 tests/{test_litellm => unit}/llms/chatgpt/test_images.py | 0 tests/{test_litellm => unit}/llms/chatgpt/test_live.py | 0 tests/{test_litellm => unit}/llms/chatgpt/test_realtime.py | 0 5 files changed, 0 insertions(+), 0 deletions(-) rename tests/{test_litellm => unit}/llms/chatgpt/conftest.py (100%) rename tests/{test_litellm => unit}/llms/chatgpt/test_codex.py (100%) rename tests/{test_litellm => unit}/llms/chatgpt/test_images.py (100%) rename tests/{test_litellm => unit}/llms/chatgpt/test_live.py (100%) rename tests/{test_litellm => unit}/llms/chatgpt/test_realtime.py (100%) diff --git a/tests/test_litellm/llms/chatgpt/conftest.py b/tests/unit/llms/chatgpt/conftest.py similarity index 100% rename from tests/test_litellm/llms/chatgpt/conftest.py rename to tests/unit/llms/chatgpt/conftest.py diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/unit/llms/chatgpt/test_codex.py similarity index 100% rename from tests/test_litellm/llms/chatgpt/test_codex.py rename to tests/unit/llms/chatgpt/test_codex.py diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/unit/llms/chatgpt/test_images.py similarity index 100% rename from tests/test_litellm/llms/chatgpt/test_images.py rename to tests/unit/llms/chatgpt/test_images.py diff --git a/tests/test_litellm/llms/chatgpt/test_live.py b/tests/unit/llms/chatgpt/test_live.py similarity index 100% rename from tests/test_litellm/llms/chatgpt/test_live.py rename to tests/unit/llms/chatgpt/test_live.py diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/unit/llms/chatgpt/test_realtime.py similarity index 100% rename from tests/test_litellm/llms/chatgpt/test_realtime.py rename to tests/unit/llms/chatgpt/test_realtime.py From d15284227c8ae2a3ab964e73f1b35a690123fd94 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 30 Sep 2026 13:45:07 +0200 Subject: [PATCH 76/90] fix(catalog): drop the duplicate image cost fields from the Rust model info Upstream added the same four resolution-tier image cost fields this branch added, in a different position, so the upstream sync kept both copies and rustc rejected the struct with E0124 on every cargo job. The kept block is byte-identical, so the crate now matches upstream exactly. --- .../crates/model-catalog/src/model_info.rs | 20 ------------------- 1 file changed, 20 deletions(-) diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index 98baec62360..380f6713d7a 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -301,26 +301,6 @@ pub struct ModelInfo { pub output_cost_per_character_above_128k_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image: Option, - #[serde( - rename = "output_cost_per_image_0.5K", - skip_serializing_if = "Option::is_none" - )] - pub output_cost_per_image_0_5k: Option, - #[serde( - rename = "output_cost_per_image_1K", - skip_serializing_if = "Option::is_none" - )] - pub output_cost_per_image_1k: Option, - #[serde( - rename = "output_cost_per_image_2K", - skip_serializing_if = "Option::is_none" - )] - pub output_cost_per_image_2k: Option, - #[serde( - rename = "output_cost_per_image_4K", - skip_serializing_if = "Option::is_none" - )] - pub output_cost_per_image_4k: Option, #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_1024: Option, #[serde(skip_serializing_if = "Option::is_none")] From 5110d619012f2d610beaac5c16da8f7daf37fa97 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 30 Sep 2026 13:45:07 +0200 Subject: [PATCH 77/90] fix(live): bound the group-limit lookup and import Final in the cache test check_unbounded_in_lists flagged the access-group budget query, whose `in` list is only as small as the groups attached to a key: the names are now sliced by IN_LIST_CHUNK_SIZE, because find_many_in cannot carry the budget join. The Live team-cache test also used `Final` without importing it, which broke the test-tree ruff run. --- litellm/proxy/realtime_endpoints/live.py | 27 +++++++++++++------ .../proxy/realtime_endpoints/test_live.py | 1 + 2 files changed, 20 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index efb3e9dcb6d..783d6b1beb2 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -63,6 +63,7 @@ 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 @@ -856,18 +857,28 @@ def _live_group_limits(row: object) -> LiteLLM_BudgetTable: async def _live_fetch_group_limits(groups: tuple[str, ...]) -> tuple[LiteLLM_BudgetTable, ...]: - """Fetch the linked budget of each group in one query and cache one entry per group.""" + """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 - rows: Final = await ModelAccessGroupBudgetRepository(server.prisma_client).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 - "in": list(groups), - } - }, - include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries + # `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 + 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) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index e49d45b050a..3d42175a4c3 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -2,6 +2,7 @@ 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 From a3f10ec09caf0a7e2fcfc3012d65c70924489ef8 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 30 Sep 2026 14:00:51 +0200 Subject: [PATCH 78/90] fix(live): keep Live authentication on the model it dispatches Veria flagged three ways a Live request could be judged against a model other than the one it runs: - the synthetic authentication request copied `parsed_body` from its source scope, so a reauthentication re-read the `{}` body cached by the first authentication and the per-model budget gate saw no model at all; - a key with `budget_fallbacks` could be rerouted by the budget check while `_create()` still dispatched the requested session model, so those requests now carry `litellm_pinned_realtime_model`, the marker upstream uses to make that fallback fail closed; - `image_edit()` spread caller-controlled passthrough parameters over the authenticated model, letting a restricted key name another image model in `extra_body`; the authenticated model now wins in both branches. --- litellm/llms/chatgpt/images.py | 4 +- litellm/proxy/realtime_endpoints/live.py | 25 ++++++++--- .../proxy/realtime_endpoints/test_live.py | 44 +++++++++++++++++++ tests/unit/llms/chatgpt/test_images.py | 22 ++++++++++ 4 files changed, 87 insertions(+), 8 deletions(-) diff --git a/litellm/llms/chatgpt/images.py b/litellm/llms/chatgpt/images.py index 8366e5016aa..554373aa64e 100644 --- a/litellm/llms/chatgpt/images.py +++ b/litellm/llms/chatgpt/images.py @@ -136,9 +136,9 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig): if not 1 <= len(validated) <= 5: raise ValueError("images must contain between 1 and 5 reference images") return { # mutable-ok: JSON request serialization - "model": model, "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), }, () @@ -147,8 +147,8 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig): if not 1 <= len(encoded) <= 5: raise ValueError("images must contain between 1 and 5 reference images") return { # mutable-ok: JSON request serialization - "model": model, "prompt": prompt, **image_edit_optional_request_params, + "model": model, # the authenticated alias wins over passthrough fields "images": encoded, }, () diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index 783d6b1beb2..e6fb8f7bc6d 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -214,13 +214,26 @@ async def _body(request: Request) -> Mapping[str, JsonValue]: raise HTTPException(400, "Expected a JSON object") from exc -def _request(source: Request | WebSocket, body: Mapping[str, JsonValue]) -> Request: +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}) ) - return Request(_mutable(MappingProxyType({**source.scope, "type": "http", "method": "POST"})), receive=receive) + 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: @@ -449,7 +462,7 @@ async def _budget_scope(auth: UserAPIKeyAuth) -> AsyncGenerator[_BudgetOwnership 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})))) + ownership.replace_auth(await _auth(_request(request, MappingProxyType({**body, "model": model}), model))) def _policy_object(value: object) -> Mapping[str, JsonValue]: @@ -1288,9 +1301,9 @@ def _response(response: httpx.Response, handle: LiveHandle | None = None) -> Res async def _create(request: Request, token: str | None = None) -> Response: body: Final = await _body(request) - auth: Final = await _auth( - _request(request, _EMPTY if token else MappingProxyType({**body, "model": _session_model(body)})) - ) + 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) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index 3d42175a4c3..034a90965d0 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -165,9 +165,11 @@ def route_client(monkeypatch): 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 @@ -205,6 +207,7 @@ def route_client(monkeypatch): auth=auth, factory=factory, bodies=authenticated_bodies, + scopes=authenticated_scopes, ) @@ -237,6 +240,47 @@ def test_create_preserves_configuration_and_returns_owned_json_session(route_cli 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) diff --git a/tests/unit/llms/chatgpt/test_images.py b/tests/unit/llms/chatgpt/test_images.py index 3d966688395..147c2d9702e 100644 --- a/tests/unit/llms/chatgpt/test_images.py +++ b/tests/unit/llms/chatgpt/test_images.py @@ -201,6 +201,28 @@ async def test_async_codex_edit_without_multipart_image(chatgpt_tokens): 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" From 06b2d76c24ec0d0a36d5636a2b1caf5983b12442 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 30 Sep 2026 14:08:29 +0200 Subject: [PATCH 79/90] chore(live): record why the group-limit comprehension spans two loops --- litellm/proxy/realtime_endpoints/live.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/proxy/realtime_endpoints/live.py b/litellm/proxy/realtime_endpoints/live.py index e6fb8f7bc6d..26724a5146d 100644 --- a/litellm/proxy/realtime_endpoints/live.py +++ b/litellm/proxy/realtime_endpoints/live.py @@ -881,6 +881,7 @@ async def _live_fetch_group_limits(groups: tuple[str, ...]) -> tuple[LiteLLM_Bud 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 From 2a67112e6ea5c3ef7e38b9f68a875e4fd63dfd35 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 30 Sep 2026 16:01:09 +0200 Subject: [PATCH 80/90] test(catalog): validate ultrafast pricing fields in the strict schema --- tests/unit/test_utils.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 5a7d525f83b..457654b93be 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -761,12 +761,14 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "supports_computer_use": {"type": "boolean"}, "cache_creation_input_audio_token_cost": {"type": "number"}, "cache_creation_input_token_cost": {"type": "number"}, + "cache_creation_input_token_cost_ultrafast": {"type": "number"}, "cache_creation_input_token_cost_above_1hr": {"type": "number"}, "cache_creation_input_token_cost_above_32k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_128k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_200k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_256k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens": {"type": "number"}, + "cache_creation_input_token_cost_above_272k_tokens_ultrafast": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens_flex": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens_priority": {"type": "number"}, "cache_creation_input_token_cost_above_200k_tokens_batches": {"type": "number"}, @@ -775,12 +777,14 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_creation_input_token_cost_flex": {"type": "number"}, "cache_creation_input_token_cost_priority": {"type": "number"}, "cache_read_input_token_cost": {"type": "number"}, + "cache_read_input_token_cost_ultrafast": {"type": "number"}, "cache_read_input_token_cost_above_32k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_128k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_200k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_200k_tokens_batches": {"type": "number"}, "cache_read_input_token_cost_above_256k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_272k_tokens": {"type": "number"}, + "cache_read_input_token_cost_above_272k_tokens_ultrafast": {"type": "number"}, "cache_read_input_token_cost_above_272k_tokens_flex": {"type": "number"}, "cache_read_input_token_cost_above_512k_tokens": {"type": "number"}, "cache_read_input_token_cost_batches": {"type": "number"}, @@ -805,6 +809,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"}, @@ -835,6 +840,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cost_per_second": {"type": "number"}, "input_cost_per_second": {"type": "number"}, "input_cost_per_token": {"type": "number"}, + "input_cost_per_token_ultrafast": {"type": "number"}, "input_cost_per_token_above_128k_tokens": {"type": "number"}, "input_cost_per_audio_token_batches": {"type": "number"}, "input_cost_per_image_token_batches": {"type": "number"}, @@ -904,12 +910,14 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "output_cost_per_second_1080p": {"type": "number"}, "output_cost_per_second_4k": {"type": "number"}, "output_cost_per_token": {"type": "number"}, + "output_cost_per_token_ultrafast": {"type": "number"}, "output_cost_per_token_above_32k_tokens": {"type": "number"}, "output_cost_per_token_above_128k_tokens": {"type": "number"}, "output_cost_per_token_above_200k_tokens": {"type": "number"}, "output_cost_per_token_above_200k_tokens_batches": {"type": "number"}, "output_cost_per_token_above_256k_tokens": {"type": "number"}, "output_cost_per_token_above_272k_tokens": {"type": "number"}, + "output_cost_per_token_above_272k_tokens_ultrafast": {"type": "number"}, "output_cost_per_token_above_512k_tokens": {"type": "number"}, "output_cost_per_token_batches": {"type": "number"}, "output_cost_per_reasoning_token": {"type": "number"}, From 94a49293ca5c7a078f161dc4fee23ced4179dc36 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 30 Sep 2026 16:03:12 +0200 Subject: [PATCH 81/90] test(interactions): follow operation schemas and declared path parameters --- .../interactions/test_openapi_compliance.py | 176 ++++++++++++++---- 1 file changed, 137 insertions(+), 39 deletions(-) diff --git a/tests/unit/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py index d3f1183cea6..8e460cc7c38 100644 --- a/tests/unit/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -9,7 +9,8 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v import json import os -from typing import Any, Dict +import re +from typing import Any, Dict, Final from unittest.mock import MagicMock, patch import httpx @@ -44,6 +45,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.""" @@ -60,12 +111,15 @@ class TestRequestCompliance: """Tests that our request bodies match the OpenAPI spec.""" def test_create_model_interaction_request_schema(self, spec_dict): - """Verify CreateModelInteractionParams schema fields.""" - schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] + """Verify the model request schema declared by POST /interactions.""" + schema = _model_request_schema(spec_dict) # Required fields per spec assert "model" in schema["required"] - assert "input" in schema["required"] + 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 = [ @@ -88,13 +142,8 @@ class TestRequestCompliance: def test_input_types_match_spec(self, spec_dict): """Verify input field supports string, Content, Content[], Turn[].""" - schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] - 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 @@ -295,45 +344,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.""" - paths = spec_dict["paths"] - - get_path = None - for path, methods in paths.items(): - if "{id}" in path and "interactions" in path and "get" in methods: - get_path = path - break - - 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.""" - paths = spec_dict["paths"] - - delete_path = None - for path, methods in paths.items(): - if "{id}" in path and "interactions" in path and "delete" in methods: - delete_path = path - break - - 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 From 6e94f53ab38683cf07dc50d4ae56fbfc51edb48f Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 00:44:38 +0200 Subject: [PATCH 82/90] fix: preserve CI route decisions and native trace signatures --- backend/routes/allowlist.py | 1 + litellm-rust/crates/python-bridge/src/routes/traces.rs | 1 + .../proxy/agent_endpoints/auth/test_managed_authorization.py | 5 ++++- tests/test_litellm_rust/test_traces.py | 4 ++-- 4 files changed, 8 insertions(+), 3 deletions(-) diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 232561dd154..5c280c22d80 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/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 6a18273ed4c..a17d465dcca 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -33,6 +33,7 @@ pub struct NativeTraceStorage { #[pymethods] impl NativeTraceStorage { #[new] + #[pyo3(signature = (database, url, reader_url=None))] fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult { litellm_traces::schema_statements(&database, 1, 1).map_err(map_error)?; Ok(Self { 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_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index 447ca2ce4bb..9e6e2e5b0cf 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -17,10 +17,10 @@ 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("SELECT {trace_id:String} AS trace_id", {"trace_id": "trace-1"})) + response: Final = json.loads(await storage.query("SELECT {trace_id:String} AS trace_id", {"trace_id": "trace-1"})) request: Final = recording_server.requests[0] parameters: Final = parse_qs(urlsplit(request.path).query) - assert rows == [{"trace_id": "trace-1"}] + assert response["data"] == [{"trace_id": "trace-1"}] assert request.raw_body == b"SELECT {trace_id:String} AS trace_id" assert parameters["database"] == ["trace_test"] assert parameters["param_trace_id"] == ["trace-1"] From 571635a43ed46a8721ab9e34d105cdf08558f35f Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 01:53:23 +0200 Subject: [PATCH 83/90] fix: align native traces and batch fixtures with upstream --- .../crates/python-bridge/src/routes/traces.rs | 4 +-- .../proxy/batches_endpoints/test_endpoints.py | 14 ++++++++ tests/test_litellm_rust/test_traces.py | 36 ++++++++++++++----- 3 files changed, 44 insertions(+), 10 deletions(-) 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/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 2d597abf3b8..8e3d25552e0 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -1038,6 +1038,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_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index 9e6e2e5b0cf..9d2348c9291 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -1,6 +1,7 @@ import base64 import gzip import json +import time from typing import Final from urllib.parse import parse_qs, urlsplit @@ -17,11 +18,14 @@ 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") - response: Final = json.loads(await storage.query("SELECT {trace_id:String} AS trace_id", {"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 response["data"] == [{"trace_id": "trace-1"}] - assert request.raw_body == b"SELECT {trace_id:String} AS trace_id" + assert b"o.TraceId = {trace_id:String}" in request.raw_body + assert b"trace-1" not in request.raw_body assert parameters["database"] == ["trace_test"] assert parameters["param_trace_id"] == ["trace-1"] assert parameters["readonly"] == ["1"] @@ -35,6 +39,13 @@ 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", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""}) + + +@pytest.mark.asyncio +async def test_trace_reader_rejects_arbitrary_sql() -> None: + storage: Final = NativeTraceStorage("trace_test", "http://localhost:8123", "http://localhost:8123") + with pytest.raises(ValueError, match="unknown ClickHouse read query"): await storage.query("SELECT 1", {}) @@ -52,7 +63,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")) @@ -64,20 +77,27 @@ 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 async def test_insert_encodes_and_sends_rows(recording_server: RecordingServer) -> None: recording_server.enqueue(ResponseSpec(body="")) storage: Final = NativeTraceStorage("trace_test", recording_server.base_url) + started_ms: Final = time.time_ns() // 1_000_000 await storage.insert_rows("otel_traces", [{"Timestamp": 1_234_567_890, "Input": "hello"}]) + finished_ms: Final = time.time_ns() // 1_000_000 request: Final = recording_server.requests[0] - assert json.loads(gzip.decompress(request.raw_body)) == { + row: Final = json.loads(gzip.decompress(request.raw_body)) + assert started_ms <= row["EngineReceivedMs"] <= finished_ms + assert {key: value for key, value in row.items() if key != "EngineReceivedMs"} == { "Input": "hello", "Timestamp": "1970-01-01T00:00:01.23456789Z", } - 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" From 1c3f1bb742883078e49b3b8bdc34fe8928fac5e9 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 03:20:34 +0200 Subject: [PATCH 84/90] fix: update vulnerable GitPython and Tornado dependencies --- pyproject.toml | 3 ++- uv.lock | 31 ++++++++++++++++--------------- 2 files changed, 18 insertions(+), 16 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index a81c75c2e0b..71b6ea2ee30 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -335,7 +335,8 @@ exclude = [ [tool.uv] constraint-dependencies = [ - "tornado>=6.5.8", + "tornado>=6.5.10", + "gitpython>=3.1.62", "aiohttp>=3.14.2,<4.0", "packaging>=24.0", "soupsieve>=2.8.4", diff --git a/uv.lock b/uv.lock index 2d31641a5e3..7d62a0ab8c2 100644 --- a/uv.lock +++ b/uv.lock @@ -21,11 +21,12 @@ members = [ ] constraints = [ { name = "aiohttp", specifier = ">=3.14.2,<4.0" }, + { name = "gitpython", specifier = ">=3.1.62" }, { name = "httplib2", specifier = ">=0.32.0" }, { name = "packaging", specifier = ">=24.0" }, { name = "setuptools", specifier = ">=83.0.0" }, { name = "soupsieve", specifier = ">=2.8.4" }, - { name = "tornado", specifier = ">=6.5.8" }, + { name = "tornado", specifier = ">=6.5.10" }, ] overrides = [ { name = "cryptography", specifier = ">=50.0.0,<51.0" }, @@ -2375,14 +2376,14 @@ wheels = [ [[package]] name = "gitpython" -version = "3.1.61" +version = "3.1.62" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "gitdb" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/6f/61/3285044215fb596bf093e39ccb96ece0a1076a8ca57a61e069a6a33cdb1b/gitpython-3.1.61.tar.gz", hash = "sha256:f51c24d8c0f733a195447385f5774a5dfe8767f5acfd7994a33755644c6ecc95", size = 231680, upload-time = "2026-08-28T11:01:13.761Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e0/db/3ca813cbacb23ab6fe46ff38a9b5ef8e73e970c8051f2ce903aacafe0446/gitpython-3.1.62.tar.gz", hash = "sha256:1791de66309bc0c7cfca40bf8d2e3de7ca091cbf94e6051be1ad0722c61062af", size = 231728, upload-time = "2026-09-07T02:57:21.155Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/6f/5e/49cc172da4d0578644ba37cec5cb365b1fefc603b26edea9bcac1c7f830a/gitpython-3.1.61-py3-none-any.whl", hash = "sha256:8ab28c9da863cdd9e7d7694ec46cf3e6c9a12d8a30a1acd3447aec11975d530c", size = 222118, upload-time = "2026-08-28T11:01:12.262Z" }, + { url = "https://files.pythonhosted.org/packages/d6/0b/29d7965215f8ef830a7ca1f42997fe13e5693d85e9edb18f938d063ef5f2/gitpython-3.1.62-py3-none-any.whl", hash = "sha256:7002251225e10e29d2e1f49e6532613fe5d5d9f0b6f1f02997a52b38fe56899e", size = 222753, upload-time = "2026-09-07T02:57:19.762Z" }, ] [[package]] @@ -9834,19 +9835,19 @@ wheels = [ [[package]] name = "tornado" -version = "6.5.8" +version = "6.5.10" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/10/d3/343e5bb989d6515b1646cf3d40135d73f3d5e45339bded401b56cdac24dd/tornado-6.5.8.tar.gz", hash = "sha256:9452e1b208a8bd771e2cb1f2ff564985b9b214bdebbe622793e1799e0a6bd23f", size = 520493, upload-time = "2026-08-07T02:12:42.971Z" } +sdist = { url = "https://files.pythonhosted.org/packages/06/61/53d562a57b28c08eda40b258c0f975e360541943ad7c7bef897a40caafda/tornado-6.5.10.tar.gz", hash = "sha256:a6b1ccd08c04b4a06fb5aeb381be99de5ad1e5375c1785e31d78c880feb57687", size = 537910, upload-time = "2026-09-15T13:47:48.73Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f2/d5/007086fd8df5489338e204f65adce33fd4f21a4999dbb2b9cff2f897b5f4/tornado-6.5.8-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:cc6aa787d7cfab7c3d35189dc7a56fbd2399a569624c730c6b55b3d6531d0403", size = 449487, upload-time = "2026-08-07T02:12:28.682Z" }, - { url = "https://files.pythonhosted.org/packages/70/c8/5a24a99495903f594f6a199dd7beead1cbc0a13e2cb9102727bcaaf2a997/tornado-6.5.8-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:9715b5eb79735b2bcd454ce216a9275b7c0470e64ea1bf5742f78b2f72b26eeb", size = 447649, upload-time = "2026-08-07T02:12:30.306Z" }, - { url = "https://files.pythonhosted.org/packages/6e/de/f2e733f386b85962d1b1dc82cd63d169b5b4580062b35397eac9244a41fe/tornado-6.5.8-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:547d63f450d570c14fe0e8db2cfb14c9bbd1c2503b4a6612586267955aa47b58", size = 450707, upload-time = "2026-08-07T02:12:31.95Z" }, - { url = "https://files.pythonhosted.org/packages/0b/94/20efeee9a01c141e9ac47c397f81679dfda24b32768fc4fff24e76d36c2c/tornado-6.5.8-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:7e2360a0ffbe145eca8af0b19cb7203d79b1a98dd4cccdd6b368f6f49c2e3808", size = 451677, upload-time = "2026-08-07T02:12:33.512Z" }, - { url = "https://files.pythonhosted.org/packages/42/ec/a96ccb8ccf0de2b7bc2c5fa1608a4803735018242e90c4882365a9fd418f/tornado-6.5.8-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:5d242290bdf7ab3151bc1065fdd75c0dcc21cbc7b49f22a4c56329c2d6566d22", size = 451510, upload-time = "2026-08-07T02:12:35.346Z" }, - { url = "https://files.pythonhosted.org/packages/29/b5/93185859245ad3f00e62175f29607346788b696369347f0146e0421286bb/tornado-6.5.8-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:7b94ff0e128fe0542f3bd331fb44d06260fc4ac16881545159f34ef08aad4195", size = 450917, upload-time = "2026-08-07T02:12:36.963Z" }, - { url = "https://files.pythonhosted.org/packages/97/cf/fe33cf062834487d34d1559746a4a12521033c22645b6d74d4bca702e018/tornado-6.5.8-cp39-abi3-win32.whl", hash = "sha256:67832909c4779c64942380cb5f044a5c6163d00831472d80e25e115de9917836", size = 451952, upload-time = "2026-08-07T02:12:38.512Z" }, - { url = "https://files.pythonhosted.org/packages/cb/e1/468ad54333e92ccb62627e62cb88e5fc14a2171daa67ed47b1b8542d5b86/tornado-6.5.8-cp39-abi3-win_amd64.whl", hash = "sha256:11881db6b7c168494be2c2d12e65931451bdf7ee718535418ae1d8855dd5a0ee", size = 452391, upload-time = "2026-08-07T02:12:39.971Z" }, - { url = "https://files.pythonhosted.org/packages/ad/3e/cd5e4f06e34cde33b8ef66cf36aa2b5ad46354cc1af7d2136bbe365fee1d/tornado-6.5.8-cp39-abi3-win_arm64.whl", hash = "sha256:68a7468c7e289f8514d7d664101753903217eff1bb6822c6b5994a0b5f5bcb26", size = 451411, upload-time = "2026-08-07T02:12:41.469Z" }, + { url = "https://files.pythonhosted.org/packages/cd/5b/ff5fc58fa2427c30dea74c90053f4fc5eda1e7f3833ed3ecc7147fe2b311/tornado-6.5.10-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:9261783640e23258694a9ff0795df430a5a7b0a651d3dd53dd0969ad6be16da7", size = 465883, upload-time = "2026-09-15T13:47:35.463Z" }, + { url = "https://files.pythonhosted.org/packages/ad/f5/cd7be26c34a3315532f3aef5f092465da8f59c334dd439d3c14aaef16461/tornado-6.5.10-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:83e6cf438b106c6b3852d70960967bb1b70c87438050dca0981e4b9aa751a4c1", size = 464046, upload-time = "2026-09-15T13:47:37.178Z" }, + { url = "https://files.pythonhosted.org/packages/60/33/df6d7d04854a58619f8349a51e3edb138324130a7562b0bb21f115bb940f/tornado-6.5.10-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:bdf942448169e5336451d0494d7e3d81cfa726d5aa312affdc4682dd62a62f6d", size = 467096, upload-time = "2026-09-15T13:47:38.559Z" }, + { url = "https://files.pythonhosted.org/packages/29/17/cc35dff68272d685cffd8600ffafbd8067e7d05e7348d9f80caddffbbd5f/tornado-6.5.10-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:69acca6501eed74582b76dbbceee2a91613f54728e3e418346000d7103101676", size = 468067, upload-time = "2026-09-15T13:47:40.085Z" }, + { url = "https://files.pythonhosted.org/packages/c3/01/6e5349b4e1a53a4b4972a6716785e1fe7407f312063c3972690af8ff301b/tornado-6.5.10-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:66aaa3f57d30c6e6becee83ff28055d5930ac724214bde99393eefda83d5e015", size = 467901, upload-time = "2026-09-15T13:47:41.576Z" }, + { url = "https://files.pythonhosted.org/packages/28/5e/b4facf94370dba006819c8d304376f8b9fbec6b935b5e51bf45823a9790b/tornado-6.5.10-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:4bd192b959f9128fb99b8898148070ba4574c9589b78bce42d1851131fe85828", size = 467308, upload-time = "2026-09-15T13:47:43.145Z" }, + { url = "https://files.pythonhosted.org/packages/56/ae/047938e828cafc8eca4c908fafb6588fee944e3af39a0af9d7b602499ae5/tornado-6.5.10-cp39-abi3-win32.whl", hash = "sha256:302eb1e0e3e159314eb591920529fdea80acca92df5510a2cec5bbd4f099ec72", size = 468387, upload-time = "2026-09-15T13:47:44.556Z" }, + { url = "https://files.pythonhosted.org/packages/d8/d4/5901517f05affd752490f6a654ba31b7474664e8dd80bd045a00c220bd88/tornado-6.5.10-cp39-abi3-win_amd64.whl", hash = "sha256:37ae8f150cecfdbf747fc4e12f5e9a97ecd8cf1d4cdb3f119e2de84b11196918", size = 468828, upload-time = "2026-09-15T13:47:45.961Z" }, + { url = "https://files.pythonhosted.org/packages/f3/1a/fd497f3a7f7b74bb04f4b94536b5c9f80742b5d50501fd27977652ddec16/tornado-6.5.10-cp39-abi3-win_arm64.whl", hash = "sha256:ce045d3c298fddd30e89a2777f97039d1b641eb9518ac7b26a4721903539c694", size = 467847, upload-time = "2026-09-15T13:47:47.283Z" }, ] [[package]] From 2fead4f6aa0064434ac649c0ea8a9f22686e1be3 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 04:39:11 +0200 Subject: [PATCH 85/90] fix(live): filter terminal response content in realtime logs --- .../litellm_core_utils/realtime_streaming.py | 24 +++++++++++- .../test_realtime_streaming.py | 38 +++++++++++++++++++ 2 files changed, 61 insertions(+), 1 deletion(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 1713bc11102..98dd7abf903 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -283,7 +283,29 @@ class RealTimeStreaming: 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"): - self.messages.append(TypeAdapter(OpenAILiveResponseEvent).validate_python(message_obj)) + response: Final = nested.get("response") + # Retain billing evidence even when response content is excluded from logging. + stored: Final = ( + message_obj + if self._should_store_message(message_obj) + else { + "type": "response.event", + "event": { + "type": nested["type"], + "response": { + **{ + key: value + for key, value in response.items() + if key in ("id", "created_at", "model", "usage", "service_tier") + }, + "output": [], + } + if isinstance(response, Mapping) + else None, + }, + } + ) + self.messages.append(TypeAdapter(OpenAILiveResponseEvent).validate_python(stored)) return if not self._should_store_message(message_obj): return diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index d74f3fcf89b..fbd3639eac0 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -3640,6 +3640,7 @@ def test_public_live_accounting_survives_filtered_logging(monkeypatch): "response": { "id": "resp_one", "model": "gpt-backend", + "output": [], "usage": {"total_tokens": 12}, }, }, @@ -3654,6 +3655,43 @@ def test_public_live_accounting_survives_filtered_logging(monkeypatch): 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( From 6c66e0f9dcc20c99ad6de7839c7b2f58abc3e562 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 04:39:11 +0200 Subject: [PATCH 86/90] fix: keep public fork dependencies aligned with upstream --- pyproject.toml | 3 +-- uv.lock | 31 +++++++++++++++---------------- 2 files changed, 16 insertions(+), 18 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 71b6ea2ee30..a81c75c2e0b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -335,8 +335,7 @@ exclude = [ [tool.uv] constraint-dependencies = [ - "tornado>=6.5.10", - "gitpython>=3.1.62", + "tornado>=6.5.8", "aiohttp>=3.14.2,<4.0", "packaging>=24.0", "soupsieve>=2.8.4", diff --git a/uv.lock b/uv.lock index 7d62a0ab8c2..2d31641a5e3 100644 --- a/uv.lock +++ b/uv.lock @@ -21,12 +21,11 @@ members = [ ] constraints = [ { name = "aiohttp", specifier = ">=3.14.2,<4.0" }, - { name = "gitpython", specifier = ">=3.1.62" }, { name = "httplib2", specifier = ">=0.32.0" }, { name = "packaging", specifier = ">=24.0" }, { name = "setuptools", specifier = ">=83.0.0" }, { name = "soupsieve", specifier = ">=2.8.4" }, - { name = "tornado", specifier = ">=6.5.10" }, + { name = "tornado", specifier = ">=6.5.8" }, ] overrides = [ { name = "cryptography", specifier = ">=50.0.0,<51.0" }, @@ -2376,14 +2375,14 @@ wheels = [ [[package]] name = "gitpython" -version = "3.1.62" +version = "3.1.61" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "gitdb" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/e0/db/3ca813cbacb23ab6fe46ff38a9b5ef8e73e970c8051f2ce903aacafe0446/gitpython-3.1.62.tar.gz", hash = "sha256:1791de66309bc0c7cfca40bf8d2e3de7ca091cbf94e6051be1ad0722c61062af", size = 231728, upload-time = "2026-09-07T02:57:21.155Z" } +sdist = { url = "https://files.pythonhosted.org/packages/6f/61/3285044215fb596bf093e39ccb96ece0a1076a8ca57a61e069a6a33cdb1b/gitpython-3.1.61.tar.gz", hash = "sha256:f51c24d8c0f733a195447385f5774a5dfe8767f5acfd7994a33755644c6ecc95", size = 231680, upload-time = "2026-08-28T11:01:13.761Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d6/0b/29d7965215f8ef830a7ca1f42997fe13e5693d85e9edb18f938d063ef5f2/gitpython-3.1.62-py3-none-any.whl", hash = "sha256:7002251225e10e29d2e1f49e6532613fe5d5d9f0b6f1f02997a52b38fe56899e", size = 222753, upload-time = "2026-09-07T02:57:19.762Z" }, + { url = "https://files.pythonhosted.org/packages/6f/5e/49cc172da4d0578644ba37cec5cb365b1fefc603b26edea9bcac1c7f830a/gitpython-3.1.61-py3-none-any.whl", hash = "sha256:8ab28c9da863cdd9e7d7694ec46cf3e6c9a12d8a30a1acd3447aec11975d530c", size = 222118, upload-time = "2026-08-28T11:01:12.262Z" }, ] [[package]] @@ -9835,19 +9834,19 @@ wheels = [ [[package]] name = "tornado" -version = "6.5.10" +version = "6.5.8" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/06/61/53d562a57b28c08eda40b258c0f975e360541943ad7c7bef897a40caafda/tornado-6.5.10.tar.gz", hash = "sha256:a6b1ccd08c04b4a06fb5aeb381be99de5ad1e5375c1785e31d78c880feb57687", size = 537910, upload-time = "2026-09-15T13:47:48.73Z" } +sdist = { url = "https://files.pythonhosted.org/packages/10/d3/343e5bb989d6515b1646cf3d40135d73f3d5e45339bded401b56cdac24dd/tornado-6.5.8.tar.gz", hash = "sha256:9452e1b208a8bd771e2cb1f2ff564985b9b214bdebbe622793e1799e0a6bd23f", size = 520493, upload-time = "2026-08-07T02:12:42.971Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/cd/5b/ff5fc58fa2427c30dea74c90053f4fc5eda1e7f3833ed3ecc7147fe2b311/tornado-6.5.10-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:9261783640e23258694a9ff0795df430a5a7b0a651d3dd53dd0969ad6be16da7", size = 465883, upload-time = "2026-09-15T13:47:35.463Z" }, - { url = "https://files.pythonhosted.org/packages/ad/f5/cd7be26c34a3315532f3aef5f092465da8f59c334dd439d3c14aaef16461/tornado-6.5.10-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:83e6cf438b106c6b3852d70960967bb1b70c87438050dca0981e4b9aa751a4c1", size = 464046, upload-time = "2026-09-15T13:47:37.178Z" }, - { url = "https://files.pythonhosted.org/packages/60/33/df6d7d04854a58619f8349a51e3edb138324130a7562b0bb21f115bb940f/tornado-6.5.10-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:bdf942448169e5336451d0494d7e3d81cfa726d5aa312affdc4682dd62a62f6d", size = 467096, upload-time = "2026-09-15T13:47:38.559Z" }, - { url = "https://files.pythonhosted.org/packages/29/17/cc35dff68272d685cffd8600ffafbd8067e7d05e7348d9f80caddffbbd5f/tornado-6.5.10-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:69acca6501eed74582b76dbbceee2a91613f54728e3e418346000d7103101676", size = 468067, upload-time = "2026-09-15T13:47:40.085Z" }, - { url = "https://files.pythonhosted.org/packages/c3/01/6e5349b4e1a53a4b4972a6716785e1fe7407f312063c3972690af8ff301b/tornado-6.5.10-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:66aaa3f57d30c6e6becee83ff28055d5930ac724214bde99393eefda83d5e015", size = 467901, upload-time = "2026-09-15T13:47:41.576Z" }, - { url = "https://files.pythonhosted.org/packages/28/5e/b4facf94370dba006819c8d304376f8b9fbec6b935b5e51bf45823a9790b/tornado-6.5.10-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:4bd192b959f9128fb99b8898148070ba4574c9589b78bce42d1851131fe85828", size = 467308, upload-time = "2026-09-15T13:47:43.145Z" }, - { url = "https://files.pythonhosted.org/packages/56/ae/047938e828cafc8eca4c908fafb6588fee944e3af39a0af9d7b602499ae5/tornado-6.5.10-cp39-abi3-win32.whl", hash = "sha256:302eb1e0e3e159314eb591920529fdea80acca92df5510a2cec5bbd4f099ec72", size = 468387, upload-time = "2026-09-15T13:47:44.556Z" }, - { url = "https://files.pythonhosted.org/packages/d8/d4/5901517f05affd752490f6a654ba31b7474664e8dd80bd045a00c220bd88/tornado-6.5.10-cp39-abi3-win_amd64.whl", hash = "sha256:37ae8f150cecfdbf747fc4e12f5e9a97ecd8cf1d4cdb3f119e2de84b11196918", size = 468828, upload-time = "2026-09-15T13:47:45.961Z" }, - { url = "https://files.pythonhosted.org/packages/f3/1a/fd497f3a7f7b74bb04f4b94536b5c9f80742b5d50501fd27977652ddec16/tornado-6.5.10-cp39-abi3-win_arm64.whl", hash = "sha256:ce045d3c298fddd30e89a2777f97039d1b641eb9518ac7b26a4721903539c694", size = 467847, upload-time = "2026-09-15T13:47:47.283Z" }, + { url = "https://files.pythonhosted.org/packages/f2/d5/007086fd8df5489338e204f65adce33fd4f21a4999dbb2b9cff2f897b5f4/tornado-6.5.8-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:cc6aa787d7cfab7c3d35189dc7a56fbd2399a569624c730c6b55b3d6531d0403", size = 449487, upload-time = "2026-08-07T02:12:28.682Z" }, + { url = "https://files.pythonhosted.org/packages/70/c8/5a24a99495903f594f6a199dd7beead1cbc0a13e2cb9102727bcaaf2a997/tornado-6.5.8-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:9715b5eb79735b2bcd454ce216a9275b7c0470e64ea1bf5742f78b2f72b26eeb", size = 447649, upload-time = "2026-08-07T02:12:30.306Z" }, + { url = "https://files.pythonhosted.org/packages/6e/de/f2e733f386b85962d1b1dc82cd63d169b5b4580062b35397eac9244a41fe/tornado-6.5.8-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:547d63f450d570c14fe0e8db2cfb14c9bbd1c2503b4a6612586267955aa47b58", size = 450707, upload-time = "2026-08-07T02:12:31.95Z" }, + { url = "https://files.pythonhosted.org/packages/0b/94/20efeee9a01c141e9ac47c397f81679dfda24b32768fc4fff24e76d36c2c/tornado-6.5.8-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:7e2360a0ffbe145eca8af0b19cb7203d79b1a98dd4cccdd6b368f6f49c2e3808", size = 451677, upload-time = "2026-08-07T02:12:33.512Z" }, + { url = "https://files.pythonhosted.org/packages/42/ec/a96ccb8ccf0de2b7bc2c5fa1608a4803735018242e90c4882365a9fd418f/tornado-6.5.8-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:5d242290bdf7ab3151bc1065fdd75c0dcc21cbc7b49f22a4c56329c2d6566d22", size = 451510, upload-time = "2026-08-07T02:12:35.346Z" }, + { url = "https://files.pythonhosted.org/packages/29/b5/93185859245ad3f00e62175f29607346788b696369347f0146e0421286bb/tornado-6.5.8-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:7b94ff0e128fe0542f3bd331fb44d06260fc4ac16881545159f34ef08aad4195", size = 450917, upload-time = "2026-08-07T02:12:36.963Z" }, + { url = "https://files.pythonhosted.org/packages/97/cf/fe33cf062834487d34d1559746a4a12521033c22645b6d74d4bca702e018/tornado-6.5.8-cp39-abi3-win32.whl", hash = "sha256:67832909c4779c64942380cb5f044a5c6163d00831472d80e25e115de9917836", size = 451952, upload-time = "2026-08-07T02:12:38.512Z" }, + { url = "https://files.pythonhosted.org/packages/cb/e1/468ad54333e92ccb62627e62cb88e5fc14a2171daa67ed47b1b8542d5b86/tornado-6.5.8-cp39-abi3-win_amd64.whl", hash = "sha256:11881db6b7c168494be2c2d12e65931451bdf7ee718535418ae1d8855dd5a0ee", size = 452391, upload-time = "2026-08-07T02:12:39.971Z" }, + { url = "https://files.pythonhosted.org/packages/ad/3e/cd5e4f06e34cde33b8ef66cf36aa2b5ad46354cc1af7d2136bbe365fee1d/tornado-6.5.8-cp39-abi3-win_arm64.whl", hash = "sha256:68a7468c7e289f8514d7d664101753903217eff1bb6822c6b5994a0b5f5bcb26", size = 451411, upload-time = "2026-08-07T02:12:41.469Z" }, ] [[package]] From 82d2d6dedf7e63d6300d957befa60a1402e73f24 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 05:14:05 +0200 Subject: [PATCH 87/90] style: format OpenAPI compliance tests --- .../interactions/test_openapi_compliance.py | 33 +++++++------------ 1 file changed, 12 insertions(+), 21 deletions(-) diff --git a/tests/unit/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py index 209ee6ee94c..b373fd95b28 100644 --- a/tests/unit/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -33,8 +33,7 @@ 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}" ) @@ -173,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 @@ -239,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] @@ -309,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") @@ -330,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)") @@ -447,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]}...") From 90dabda182562d6c208c177537bc0981ff44bb90 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 10:55:16 +0200 Subject: [PATCH 88/90] fix: validate Codex route inputs for type gates --- litellm/images/main.py | 5 ++- .../litellm_core_utils/llm_cost_calc/utils.py | 12 +++-- .../litellm_core_utils/realtime_streaming.py | 45 ++++++++++--------- litellm/llms/chatgpt/codex.py | 5 ++- litellm/proxy/auth/auth_utils.py | 2 +- .../proxy/hooks/parallel_request_limiter.py | 2 +- .../hooks/parallel_request_limiter_v3.py | 3 +- .../proxy/realtime_endpoints/call_sessions.py | 26 ++++++----- litellm/realtime_api/main.py | 18 +++++--- litellm/types/llms/openai.py | 2 +- 10 files changed, 74 insertions(+), 46 deletions(-) diff --git a/litellm/images/main.py b/litellm/images/main.py index 13941bf6faa..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 @@ -398,7 +399,9 @@ def image_generation( model=model, prompt=prompt, image_generation_provider_config=image_generation_config, - extra_headers=extra_headers, + 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, diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 0299bdfbf73..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 @@ -1867,10 +1868,13 @@ def calculate_image_response_cost_from_usage( 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) - cached_text: Final = _get_token_detail_value(cached_details, "text_tokens") or 0 - cached_image: Final = _get_token_detail_value(cached_details, "image_tokens") or 0 - input_text_tokens: Final = _get_token_detail_value(input_tokens_details, "text_tokens") or 0 - input_image_tokens: Final = _get_token_detail_value(input_tokens_details, "image_tokens") or 0 + 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 diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 98dd7abf903..64481d66059 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -5,6 +5,7 @@ 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 @@ -284,27 +285,31 @@ class RealTimeStreaming: nested: Final = message_obj["event"] if nested.get("type") in ("response.completed", "response.incomplete", "response.failed"): response: Final = nested.get("response") - # Retain billing evidence even when response content is excluded from logging. - stored: Final = ( - message_obj - if self._should_store_message(message_obj) - else { - "type": "response.event", - "event": { - "type": nested["type"], - "response": { - **{ - key: value - for key, value in response.items() - if key in ("id", "created_at", "model", "usage", "service_tier") - }, - "output": [], - } - if isinstance(response, Mapping) - else None, - }, - } + 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): diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index ea736c97a0d..92595890c86 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -3,7 +3,7 @@ from typing import Final from urllib.parse import urlsplit import httpx -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, TypeAdapter from typing_extensions import ReadOnly, TypedDict from litellm.types.realtime import RealtimeQueryParams, RealtimeSessionConfig @@ -69,7 +69,8 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire if not routing_data: raise ValueError("Direct call signaling requires a ChatGPT deployment") routing: Final = ChatGPTCallRouting.model_validate(routing_data) - call_id: Final = urlsplit(response.headers.get("location", "")).path.rstrip("/").rsplit("/", 1)[-1] + 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, diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 10d17f124e8..f026241d017 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1921,7 +1921,7 @@ def _extract_model_candidates_from_request( if uses_body_target_model_sources or not body_model: _append_model_candidates(candidates, request_data.get("target_model_names")) if uses_session_model: - _append_model_candidates(candidates, 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/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index feced0b8359..ddf571b6098 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -110,7 +110,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): 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).validate_python(raw) + current: Final = TypeAdapter[Mapping[str, int] | None](Mapping[str, int] | None).validate_python(raw) updated: Final = ( { # mutable-ok: shared cache counter dict **current, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index d89f30b164b..5ab3cadb66f 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4253,7 +4253,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): effective_descriptors: Final = ( tuple(_without_parallel_limit(descriptor) for descriptor in descriptors) - if call_type == "_arealtime" and is_realtime_call_attachment(data.get("websocket")) + if call_type == "_arealtime" + and is_realtime_call_attachment(TypeAdapter[object](object).validate_python(data.get("websocket"))) else descriptors ) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index f2fb8e365e9..997fd0e38a6 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -112,7 +112,9 @@ async def _start_codex_supervisor( "extra_headers": MappingProxyType( { **configured_realtime_headers( - TypeAdapter(Mapping[str, object] | None).validate_python(processed.get("extra_headers")) + TypeAdapter[Mapping[str, object] | None](Mapping[str, object] | None).validate_python( + processed.get("extra_headers") + ) ), **configured_realtime_headers(call.extra_headers), } @@ -290,11 +292,11 @@ async def process_codex_request( proxy_logging_obj=server.proxy_logging_obj, proxy_config=server.proxy_config, llm_router=server.llm_router, - user_model=server.user_model, - user_temperature=server.user_temperature, + 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=server.user_api_base, + user_api_base=TypeAdapter[str | None](str | None).validate_python(server.user_api_base), model=model, route_type=route_type, **( @@ -358,7 +360,9 @@ async def _create_codex_realtime_call(request: Request) -> Response: try: await can_key_call_resolved_model( model=model, - llm_model_list=server.llm_model_list, + llm_model_list=TypeAdapter[tuple[object, ...] | None](tuple[object, ...] | None).validate_python( + server.llm_model_list + ), valid_token=auth, llm_router=server.llm_router, ) @@ -388,7 +392,7 @@ async def _create_codex_realtime_call(request: Request) -> Response: data=processed, route_type="arealtime_calls", llm_router=server.llm_router, - user_model=server.user_model, + user_model=TypeAdapter[str | None](str | None).validate_python(server.user_model), ) try: response: Final = await result @@ -462,7 +466,9 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP call: Final = decode_call(token, f"Bearer {api_key}") await can_key_call_resolved_model( model=call.alias, - llm_model_list=server.llm_model_list, + llm_model_list=TypeAdapter[tuple[object, ...] | None](tuple[object, ...] | None).validate_python( + server.llm_model_list + ), valid_token=auth, llm_router=server.llm_router, ) @@ -530,9 +536,9 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP "extra_headers": MappingProxyType( { **configured_realtime_headers( - TypeAdapter(Mapping[str, object] | None).validate_python( - processed.get("extra_headers") - ) + TypeAdapter[Mapping[str, object] | None]( + Mapping[str, object] | None + ).validate_python(processed.get("extra_headers")) ), **configured_realtime_headers(call.extra_headers), } diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 10d33e99525..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, @@ -292,10 +294,13 @@ async def arealtime_calls( ) 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(kwargs.get("extra_headers")) + provider_config.get_realtime_calls_extra_headers(supplied_headers) if provider_config is not None - else kwargs.get("extra_headers") + else supplied_headers ) litellm_logging_obj.update_from_kwargs( kwargs=kwargs, @@ -379,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: @@ -428,6 +435,7 @@ async def _arealtime( 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, @@ -436,7 +444,7 @@ async def _arealtime( api_key=api_key, timeout=timeout, query_params=query_params, - user_api_key_dict=kwargs.get("user_api_key_dict"), + user_api_key_dict=user_api_key_dict, litellm_metadata=_build_litellm_metadata(kwargs), ) elif provider_config is not None: diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 6199f56d3a9..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) From 5387efcb33f6c5cca4c507176a0d61f21b8e3420 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 11:32:13 +0200 Subject: [PATCH 89/90] fix: reject legacy SDP credentials for ChatGPT OAuth --- litellm/llms/chatgpt/realtime.py | 7 ++ litellm/proxy/realtime_endpoints/endpoints.py | 6 ++ .../test_realtime_webrtc_endpoints.py | 99 ++++++++++++++++++- tests/unit/llms/chatgpt/test_realtime.py | 23 +++++ 4 files changed, 134 insertions(+), 1 deletion(-) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index f6aa04fc29d..a5aad4ccb8d 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -7,6 +7,7 @@ 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 @@ -258,6 +259,12 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): 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( diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index d66976d3e6f..1e40c41e9d6 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -427,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=json.dumps({"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/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py index 4dfc119ce8e..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 @@ -521,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) @@ -538,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/unit/llms/chatgpt/test_realtime.py b/tests/unit/llms/chatgpt/test_realtime.py index 4f367dd8229..327bd957c4e 100644 --- a/tests/unit/llms/chatgpt/test_realtime.py +++ b/tests/unit/llms/chatgpt/test_realtime.py @@ -287,6 +287,29 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap 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 = [] From 966b047dd5ff7e67b6afe4156172c1e3d53ba0fa Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 1 Oct 2026 11:37:01 +0200 Subject: [PATCH 90/90] fix: keep legacy SDP rejection response immutable --- litellm/proxy/realtime_endpoints/endpoints.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index 1e40c41e9d6..59c3f59e427 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -429,7 +429,7 @@ async def proxy_realtime_calls( decoded_payload: Final = _decode_realtime_token_payload(decrypted_token_value) if decoded_payload is None and decrypted_token_value.lstrip().startswith(("{", "[")): return Response( - content=json.dumps({"error": "Invalid or expired token"}), + content='{"error":"Invalid or expired token"}', status_code=http_status.HTTP_401_UNAUTHORIZED, media_type="application/json", )