From bce8aa1dcc4a224a6cd38f64d8c98700c437f696 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 07:20:37 +0200 Subject: [PATCH 01/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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/38] 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))