This commit is contained in:
Jordi Ibáñez 2026-10-01 09:44:56 +00:00 • committed by GitHub
commit 047350c725
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
90 changed files with 16708 additions and 982 deletions

View file

@ -12,6 +12,11 @@ queries:
query-filters:
- exclude:
id: py/clear-text-logging-sensitive-data # CWE-312
# CodeQL 2.27.0 exceeds its 2 GiB result-set limit while evaluating the
# repository-wide log-injection data-flow query. Keep the remaining Python
# security-and-quality queries enabled until the upstream query scales.
- exclude:
id: py/log-injection # CWE-117
- exclude:
id: py/polynomial-redos # CWE-730
# Import resolution confuses stdlib types with management_endpoints/types.py.

View file

@ -103,6 +103,16 @@ jobs:
outputs:
decision: ${{ steps.changes.outputs.decision }}
has-coverage: ${{ steps.tests.outputs.has-coverage }}
services:
redis:
image: ${{ inputs.artifact-name == 'proxy-auth' && 'redis:8.2.9-alpine@sha256:30abb90e62f14b737010746def3ba99cc79fe19dcdb3d37b41f21fc62e7da19d' || '' }}
ports:
- '127.0.0.1::6379'
options: >-
--health-cmd "redis-cli ping"
--health-interval 2s
--health-timeout 2s
--health-retries 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
@ -183,6 +193,7 @@ jobs:
TEST_TIMEOUT_SECONDS: ${{ inputs.test-timeout-seconds }}
DIST: ${{ inputs.dist }}
COVERAGE_CORE: sysmon
LITELLM_TEST_REDIS_PORT: ${{ job.services.redis.ports['6379'] }}
run: |
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
selection="${TEST_PATH}"

View file

@ -76,9 +76,14 @@ jobs:
if: matrix.language == 'python'
uses: advanced-security/filter-sarif@2da736ff05ef065cb2894ac6892e47b5eac2c3c0 # v1.1
with:
# These SHA-256 digests are opaque ownership/cache identifiers, not password hashes.
patterns: |
-litellm/llms/oci/common_utils.py:py/weak-sensitive-data-hashing
-litellm/proxy/auth/password_policy.py:py/weak-sensitive-data-hashing
-litellm/proxy/_types.py:py/weak-sensitive-data-hashing
-litellm/proxy/realtime_endpoints/call_sessions.py:py/weak-sensitive-data-hashing
-litellm/proxy/realtime_endpoints/live.py:py/weak-sensitive-data-hashing
-litellm/proxy/utils.py:py/weak-sensitive-data-hashing
input: sarif-results/python.sarif
output: sarif-results/python.sarif

View file

@ -123,6 +123,7 @@ jobs:
tests/test_litellm/proxy/hooks
tests/test_litellm/proxy/policy_engine
tests/test_litellm/proxy/client
tests/local_testing/test_realtime_call_redis.py
workers: 2
reruns: 2
timeout-minutes: 20

View file

@ -60,6 +60,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
# Tools / agents (registry & policy admin)
"/v1/tool/",
"/v1/agents",
"/v1/traces",
# Guardrails admin
"/v2/guardrails/",
# MCP server admin + BYOK OAuth flow (UI-initiated) + dynamic per-server endpoints

View file

@ -114,6 +114,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/{provider}/",
"/toolset/",
# Realtime / streaming
"/v1/live",
"/live",
"/v1/realtime",
"/realtime",
# Health & ops

View file

@ -120,13 +120,13 @@ impl NativeTraceStorage {
fn query<'py>(
&self,
py: Python<'py>,
query: &str,
sql: &str,
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] parameters: BTreeMap<
String,
Parameter,
>,
) -> PyResult<Bound<'py, PyAny>> {
let query = ReadQuery::parse(query).map_err(map_error)?;
let query = ReadQuery::parse(sql).map_err(map_error)?;
let connection = self.reader.clone().ok_or_else(|| {
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
})?;

View file

@ -8,7 +8,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from httpx import Response
from pydantic import BaseModel
from pydantic import BaseModel, ValidationError
from typing_extensions import ReadOnly, TypedDict
import litellm
@ -2796,6 +2796,13 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
result["response"].get("usage", {})
)
usage_objects.append(usage_object)
usage_objects.extend(
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( # pyright: ignore[reportPrivateUsage] # reuse the existing Responses usage conversion for nested Live events
response.usage.model_dump()
)
for response in _live_backend_responses(results)
if response.usage is not None
)
return usage_objects
@staticmethod
@ -2978,7 +2985,15 @@ def handle_realtime_stream_cost_calculation(
potential_model_names.append(litellm_model_name)
input_cost_per_token, output_cost_per_token = _first_priced_realtime_token_costs(
potential_model_names=potential_model_names,
combined_usage_object=combined_usage_object,
combined_usage_object=(
RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(
[ # mutable-ok: collector requires a concrete event list
event for event in results if event.get("type") != "response.event"
]
)
if any(event.get("type") == "response.event" for event in results)
else combined_usage_object
),
custom_llm_provider=custom_llm_provider,
data_residency=data_residency,
)
@ -2992,7 +3007,18 @@ def handle_realtime_stream_cost_calculation(
if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results)
else 0.0
)
total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost
live_audio_cost: Final = handle_live_session_duration_cost(
results=results,
custom_llm_provider=custom_llm_provider,
litellm_model_name=litellm_model_name,
)
backend_cost: Final = sum(
_live_backend_response_cost(response, litellm_logging_obj)
for response in _live_backend_responses(results, litellm_logging_obj)
)
total_cost: Final = (
input_cost_per_token + output_cost_per_token + transcription_cost + live_audio_cost + backend_cost
)
_store_cost_breakdown_in_logging_obj(
litellm_logging_obj=litellm_logging_obj,
@ -3000,13 +3026,122 @@ def handle_realtime_stream_cost_calculation(
completion_tokens_cost_usd_dollar=output_cost_per_token,
cost_for_built_in_tools_cost_usd_dollar=0.0,
total_cost_usd_dollar=total_cost,
additional_costs={"transcription_cost": transcription_cost} if transcription_cost > 0 else None,
additional_costs={ # mutable-ok: logging cost breakdown requires a concrete dict
name: cost
for name, cost in (
("transcription_cost", transcription_cost),
("live_audio_cost", live_audio_cost),
("live_backend_cost", backend_cost),
)
if cost > 0
}
or None,
data_residency=data_residency,
)
return total_cost
class _LiveBackendEvent(BaseModel):
type: str
response: object = None
class _LiveBackendEnvelope(BaseModel):
event: _LiveBackendEvent
def _live_backend_responses(
results: OpenAIRealtimeStreamList, logging_obj: LitellmLoggingObject | None = None
) -> tuple[ResponsesAPIResponse, ...]:
responses: Final = { # mutable-ok: deduplicate terminal backend responses by response id
response.id: response
for result in results
if result.get("type") == "response.event"
and (response := _live_backend_response(result, logging_obj)) is not None
}
return tuple(responses.values())
def _mark_live_backend_accounting_incomplete(logging_obj: LitellmLoggingObject | None) -> None:
verbose_logger.warning("Live backend accounting incomplete: missing valid terminal usage or model pricing")
if logging_obj is not None:
logging_obj.model_call_details["realtime_backend_accounting_incomplete"] = True
def _live_backend_response(
result: Mapping[str, object], logging_obj: LitellmLoggingObject | None
) -> ResponsesAPIResponse | None:
try:
event: Final = _LiveBackendEnvelope.model_validate(result).event
except ValidationError:
return None
if event.type not in ("response.completed", "response.incomplete", "response.failed"):
return None
try:
response: Final = ResponsesAPIResponse.model_validate(event.response)
except ValidationError:
_mark_live_backend_accounting_incomplete(logging_obj)
return None
if response.usage is None:
_mark_live_backend_accounting_incomplete(logging_obj)
return None
return response
def _live_backend_response_cost(response: ResponsesAPIResponse, logging_obj: LitellmLoggingObject | None) -> float:
try:
return completion_cost(
completion_response=response, model=response.model, custom_llm_provider="openai", call_type="aresponses"
)
except Exception: # noqa: BLE001 # preserve measured voice cost when backend pricing cannot be resolved
_mark_live_backend_accounting_incomplete(logging_obj)
return 0.0
def handle_live_session_duration_cost(
results: OpenAIRealtimeStreamList,
custom_llm_provider: str,
litellm_model_name: str,
) -> float:
seconds: Final = max(
(
duration
for event in results
if event.get("type") in ("session.closed", "session.usage.updated")
and (duration := _live_duration_seconds(event)) is not None
),
default=0.0,
)
initialization_seconds: Final = max(
(
duration
for event in results
if event.get("type") == "litellm.live.initialization"
and (duration := _live_duration_seconds(event)) is not None
),
default=0.0,
)
try:
model_info: Final = litellm.get_model_info(model=litellm_model_name, custom_llm_provider=custom_llm_provider)
except Exception: # noqa: BLE001 # an unknown model simply has no per-second price to bill
return 0.0
return max(seconds, initialization_seconds) * (model_info.get("input_cost_per_second") or 0.0)
def _live_duration_seconds(event: Mapping[str, object]) -> float | None:
from litellm.types.realtime import LiveSessionUsageEvent
try:
usage: Final = LiveSessionUsageEvent.model_validate(event).usage
except ValidationError:
return None
raw_usage: Final = cast( # cast-ok: LiveSessionUsageEvent validated the usage mapping above
Mapping[str, object], event.get("usage")
)
return usage.duration / (1 if "seconds" in raw_usage else 1000)
def handle_realtime_transcription_cost_calculation(
results: OpenAIRealtimeStreamList,
custom_llm_provider: str,

View file

@ -9,6 +9,7 @@ if TYPE_CHECKING:
from litellm.images.utils import ImageEditRequestUtils
import httpx
from pydantic import TypeAdapter
import litellm
@ -372,6 +373,7 @@ def image_generation(
# Providers using llm_http_handler
#########################################################
elif custom_llm_provider in (
litellm.LlmProviders.CHATGPT,
litellm.LlmProviders.RECRAFT,
litellm.LlmProviders.AIML,
litellm.LlmProviders.GEMINI,
@ -397,6 +399,9 @@ def image_generation(
model=model,
prompt=prompt,
image_generation_provider_config=image_generation_config,
extra_headers=TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(
extra_headers
),
image_generation_optional_request_params=optional_params,
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params_dict,
@ -854,16 +859,28 @@ def image_edit(
additional_drop_params=kwargs.get("additional_drop_params"),
)
if image_edit_provider_config.use_multipart_form_data() and (
custom_llm_provider == "openai"
or custom_llm_provider == "azure"
or custom_llm_provider in litellm.openai_compatible_providers
):
if (
image_edit_provider_config.use_multipart_form_data()
and (
custom_llm_provider == "openai"
or custom_llm_provider == "azure"
or custom_llm_provider in litellm.openai_compatible_providers
)
) or custom_llm_provider == litellm.LlmProviders.CHATGPT:
image_edit_request_params.update(
flatten_form_field_values(
non_default_params,
extra_body if isinstance(extra_body, dict) else None,
)
if image_edit_provider_config.use_multipart_form_data()
else { # mutable-ok: image provider update requires a concrete request-parameter dict
**non_default_params,
**(
extra_body
if isinstance(extra_body, dict)
else {} # mutable-ok: empty fallback is consumed immediately
),
}
)
# Pre Call logging
@ -962,9 +979,9 @@ def image_edit(
@client
async def aimage_edit(
image: FileTypes | list[FileTypes],
model: str,
prompt: str,
image: FileTypes | list[FileTypes] | None = None,
model: str = "",
prompt: str = "",
mask: str | None = None,
n: int | None = None,
quality: str | ImageGenerationRequestQuality | None = None,
@ -1002,11 +1019,9 @@ async def aimage_edit(
model=model, api_base=local_vars.get("base_url", None)
)
images: Final = image if isinstance(image, list) else [image]
func: Final = partial(
image_edit,
image=images,
image=image,
prompt=prompt,
mask=mask,
model=model,

View file

@ -9,6 +9,7 @@ from types import MappingProxyType
from typing import Final, Literal, TypedDict, cast
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from pydantic import TypeAdapter
from typing_extensions import ReadOnly
import litellm
@ -1859,7 +1860,34 @@ def calculate_image_response_cost_from_usage(
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
return prompt_cost + completion_cost
cached_details: Final = (
input_tokens_details.get("cached_tokens_details")
if isinstance(input_tokens_details, dict)
else getattr(input_tokens_details, "cached_tokens_details", None)
)
if cached_details is None:
return prompt_cost + completion_cost
catalog_model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
details_adapter: Final = TypeAdapter[object](object)
cached_token_details: Final = details_adapter.validate_python(cached_details)
input_token_details: Final = details_adapter.validate_python(input_tokens_details)
cached_text: Final = _get_token_detail_value(cached_token_details, "text_tokens") or 0
cached_image: Final = _get_token_detail_value(cached_token_details, "image_tokens") or 0
input_text_tokens: Final = _get_token_detail_value(input_token_details, "text_tokens") or 0
input_image_tokens: Final = _get_token_detail_value(input_token_details, "image_tokens") or 0
if not (0 <= cached_text <= input_text_tokens and 0 <= cached_image <= input_image_tokens):
raise ValueError("Image cached token counts exceed their input modality counts")
text_rate: Final = catalog_model_info.get("input_cost_per_token") or 0.0
image_rate: Final = catalog_model_info.get("input_cost_per_image_token")
cache_text_rate: Final = catalog_model_info.get("cache_read_input_token_cost")
cache_image_rate: Final = catalog_model_info.get("cache_read_input_image_token_cost")
text_savings: Final = cached_text * (text_rate - cache_text_rate) if cache_text_rate is not None else 0.0
image_savings: Final = (
cached_image * ((image_rate if image_rate is not None else text_rate) - cache_image_rate)
if cache_image_rate is not None
else 0.0
)
return prompt_cost + completion_cost - text_savings - image_savings
def calculate_image_response_web_search_cost(

View file

@ -1,11 +1,14 @@
import asyncio
import json
import traceback
from collections.abc import Coroutine, Mapping, Sequence
from collections.abc import Awaitable, Callable, Coroutine, Mapping, Sequence
from contextvars import ContextVar
from dataclasses import dataclass
from enum import Enum, auto
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast
from pydantic import TypeAdapter
from typing_extensions import ReadOnly
import litellm
@ -14,9 +17,11 @@ from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESS
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend
from litellm.types.llms.openai import (
OpenAILiveResponseEvent,
OpenAIRealtimeEvents,
OpenAIRealtimeOutputItemDone,
OpenAIRealtimeResponseDelta,
OpenAIRealtimeSessionClosed,
OpenAIRealtimeStreamResponseBaseObject,
OpenAIRealtimeStreamSessionEvents,
)
@ -25,6 +30,10 @@ from litellm.types.realtime import ALL_DELTA_TYPES
from .litellm_logging import Logging as LiteLLMLogging
from .realtime_errors import client_close_code, realtime_error_event, websocket_close_reason
realtime_attachment_cleanup: Final[ContextVar[Callable[[], Awaitable[None]] | None]] = ContextVar(
"realtime_attachment_cleanup", default=None
)
if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection
from websockets.exceptions import ConnectionClosed
@ -137,12 +146,25 @@ class RealTimeStreaming:
force_transcription_model: str | None = None,
event_normalizer: RealtimeEventNormalizer | None = None,
logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER,
*,
account_usage: bool = True,
live_initialization_seconds: float = 0,
):
self.websocket: _ClientWebSocket = websocket
self.backend_ws = backend_ws
self.logging_obj = logging_obj
self._logging_worker = logging_worker
self._account_usage = account_usage
self.messages: list[OpenAIRealtimeEvents] = []
if account_usage and live_initialization_seconds > 0:
self.messages.append(
{ # mutable-ok: initialization event is appended to the mutable event history
"type": "litellm.live.initialization",
"usage": { # mutable-ok: usage payload is consumed as part of the typed event
"seconds": live_initialization_seconds,
},
}
)
self._backend_sent_frames: bool = False
self.input_message: dict = {}
self.input_messages: list[dict[str, str]] = []
@ -254,6 +276,42 @@ class RealTimeStreaming:
else:
message_obj = cast(dict[str, Any], json.loads(cast(str, message)))
self._collect_tool_calls_from_response_done(cast(dict, message_obj))
if message_obj.get("type") in ("session.closed", "session.usage.updated") and isinstance(
message_obj.get("usage"), dict
):
self.messages.append(TypeAdapter(OpenAIRealtimeSessionClosed).validate_python(message_obj))
return
if message_obj.get("type") == "response.event" and isinstance(message_obj.get("event"), dict):
nested: Final = message_obj["event"]
if nested.get("type") in ("response.completed", "response.incomplete", "response.failed"):
response: Final = nested.get("response")
response_mapping: Final = (
TypeAdapter(Mapping[str, object]).validate_python(response)
if isinstance(response, Mapping)
else None
)
# Retain billing evidence even when response content is excluded from logging.
filtered: Final[OpenAILiveResponseEvent] = {
"type": "response.event",
"event": {
"type": nested["type"],
"response": {
**MappingProxyType(
{
key: value
for key, value in response_mapping.items()
if key in ("id", "created_at", "model", "usage", "service_tier")
}
),
"output": TypeAdapter(list[object]).validate_python(()),
}
if response_mapping is not None
else None,
},
}
stored: Final = message_obj if self._should_store_message(message_obj) else filtered
self.messages.append(TypeAdapter(OpenAILiveResponseEvent).validate_python(stored))
return
if not self._should_store_message(message_obj):
return
try:
@ -408,8 +466,10 @@ class RealTimeStreaming:
if self.logging_obj:
self.logging_obj.pre_call(input=message, api_key="")
async def log_messages(self):
async def log_messages(self, *, wait_for_dispatch: bool = False):
"""Log messages in list"""
if not self._account_usage:
return
if self.logging_obj:
if self.input_messages:
self.logging_obj.model_call_details["messages"] = self.input_messages
@ -419,9 +479,12 @@ class RealTimeStreaming:
# Route through the bounded logging worker (per-coroutine timeout +
# concurrency cap) instead of a bare create_task, so a slow callback
# can't leave suspended tasks pinning each call's response in memory.
self._logging_worker.ensure_initialized_and_enqueue(
self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
)
if wait_for_dispatch:
await self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
else:
self._logging_worker.ensure_initialized_and_enqueue(
self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
)
self.logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
async def _send_to_backend(self, message: str) -> bool:
@ -1573,7 +1636,12 @@ class RealTimeStreaming:
finally:
forward_task.cancel()
client_task.cancel()
await asyncio.gather(forward_task, client_task, return_exceptions=True)
try:
await asyncio.gather(forward_task, client_task, return_exceptions=True)
finally:
cleanup: Final = realtime_attachment_cleanup.get()
if not self._account_usage and cleanup is not None:
await cleanup()
async def _close_client(self, close: BackendClose) -> None:
redacted_message: Final = redact_internal_details_from_client_message(close.message)

View file

@ -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 #
# ------------------------------------------------------------------ #

View file

@ -49,8 +49,9 @@ class Authenticator:
self.auth_file = os.path.join(self.token_dir, os.getenv("CHATGPT_AUTH_FILE", "auth.json"))
self._ensure_token_dir()
def get_api_base(self) -> str:
return os.getenv("CHATGPT_API_BASE") or os.getenv("OPENAI_CHATGPT_API_BASE") or CHATGPT_API_BASE
@staticmethod
def get_api_base(default_base: str = CHATGPT_API_BASE) -> str:
return os.getenv("CHATGPT_API_BASE") or os.getenv("OPENAI_CHATGPT_API_BASE") or default_base
def get_access_token(self) -> str:
auth_data: Final = self._read_auth_file()

View file

@ -23,8 +23,8 @@ class ChatGPTConfig(OpenAIConfig):
super().__init__()
self.authenticator = Authenticator()
def api_base_without_login(self) -> str:
return self.authenticator.get_api_base()
def api_base_without_login(self, api_base: str | None = None) -> str:
return api_base or self.authenticator.get_api_base()
def _get_openai_compatible_provider_info(
self,
@ -33,7 +33,7 @@ class ChatGPTConfig(OpenAIConfig):
api_key: str | None,
custom_llm_provider: str,
) -> tuple[str | None, str | None, str]:
dynamic_api_base: Final = self.api_base_without_login()
dynamic_api_base: Final = self.api_base_without_login(api_base)
try:
dynamic_api_key: Final = self.authenticator.get_access_token()
except GetAccessTokenError as e:

View file

@ -0,0 +1,95 @@
from collections.abc import Mapping
from typing import Final
from urllib.parse import urlsplit
import httpx
from pydantic import BaseModel, Field, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm.types.realtime import RealtimeQueryParams, RealtimeSessionConfig
class CodexRealtimeOffer(BaseModel):
sdp: str = Field(min_length=1)
session: RealtimeSessionConfig
class CodexRealtimeCall(BaseModel):
call_id: str = Field(pattern=r"^rtc_[A-Za-z0-9_-]+$")
model: str
model_id: str | None = None
alias: str
api_base: str | None = None
extra_headers: Mapping[str, str] | None = None
extra_query: Mapping[str, str | tuple[str, ...]] | None = None
usage_supervised: bool = False
parallel_reserved: bool = False
owner: str
expires_at: float
class ChatGPTCallRouting(BaseModel):
model: str
model_id: str | None = None
api_base: str | None = None
extra_headers: Mapping[str, str] | None = None
extra_query: Mapping[str, str | tuple[str, ...]] | None = None
class CodexSidebandRequest(TypedDict):
api_base: ReadOnly[str | None]
model: ReadOnly[str]
chatgpt_realtime_call_id: ReadOnly[str]
query_params: ReadOnly[RealtimeQueryParams]
extra_headers: ReadOnly[Mapping[str, str] | None]
extra_query: ReadOnly[Mapping[str, str | tuple[str, ...]] | None]
def build_call_request(
offer: CodexRealtimeOffer, query: Mapping[str, str], headers: Mapping[str, str]
) -> dict[str, object]: # mutable-ok: proxy processor enriches the request dictionary
return { # mutable-ok: proxy processor enriches the request dictionary
"model": offer.session.model,
"sdp_body": offer.sdp.encode(),
"session": offer.session.model_dump(exclude_none=True),
"openai_ephemeral_key": "",
"chatgpt_realtime_client_query": { # mutable-ok: router request parameters
key: value for key, value in query.items() if key in ("intent", "architecture")
},
"chatgpt_realtime_client_headers": { # mutable-ok: router request headers
key: value
for key, value in headers.items()
if key in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation")
},
}
def parse_call_response(response: httpx.Response, alias: str, owner: str, expires_at: float) -> CodexRealtimeCall:
routing_data: Final = response.extensions.get("chatgpt_realtime")
if not routing_data:
raise ValueError("Direct call signaling requires a ChatGPT deployment")
routing: Final = ChatGPTCallRouting.model_validate(routing_data)
location: Final = TypeAdapter(str).validate_python(response.headers.get("location", ""))
call_id: Final[str] = urlsplit(location).path.rstrip("/").rsplit("/", 1)[-1]
return CodexRealtimeCall(
call_id=call_id,
model=routing.model,
model_id=routing.model_id,
alias=alias,
owner=owner,
expires_at=expires_at,
api_base=routing.api_base,
extra_headers=routing.extra_headers,
extra_query=routing.extra_query,
)
def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest:
return CodexSidebandRequest(
api_base=call.api_base,
model=f"chatgpt/{call.model}",
chatgpt_realtime_call_id=call.call_id,
query_params=RealtimeQueryParams(model=call.model),
extra_headers=call.extra_headers,
extra_query=call.extra_query,
)

View file

@ -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,

View file

@ -0,0 +1,154 @@
import base64
import os
from collections.abc import Mapping, Sequence
from pathlib import Path
from types import MappingProxyType
from typing import Final
from httpx._types import FileTypes as HTTPFileTypes
from httpx._types import RequestFiles
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
from litellm.images.utils import ImageEditRequestUtils
from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig
from litellm.llms.openai.image_generation.gpt_transformation import GPTImageGenerationConfig
from litellm.types.llms.openai import AllMessageValues, FileTypes
from litellm.types.router import GenericLiteLLMParams
from .authenticator import Authenticator
from .common_utils import without_oauth_identity_headers
from .responses.transformation import ChatGPTResponsesAPIConfig
class ReferenceImage(BaseModel):
model_config = ConfigDict(extra="forbid")
image_url: str = Field(pattern=r"^(data:image/(png|jpeg|webp);base64,|https://)")
def encode_reference(
file: HTTPFileTypes | FileTypes,
) -> dict[str, str]: # mutable-ok: image handler requires dictionaries
content: Final = file[1] if isinstance(file, tuple) else file
raw: Final = (
Path(os.fsdecode(content)).read_bytes()
if isinstance(content, os.PathLike)
else content.encode()
if isinstance(content, str)
else content
if isinstance(content, bytes)
else content.read()
)
content_type: Final = (
file[2]
if isinstance(file, tuple) and len(file) >= 3 and file[2]
else ImageEditRequestUtils.get_image_content_type(raw)
)
if content_type not in ("image/png", "image/jpeg", "image/webp"):
raise ValueError("Reference images must be PNG, JPEG, or WEBP")
return { # mutable-ok: JSON request serialization
"image_url": f"data:{content_type};base64," + base64.b64encode(raw).decode("ascii")
}
def image_headers(
headers: Mapping[str, object], model: str, params: Mapping[str, object]
) -> dict[str, object]: # mutable-ok: image handler requires dictionaries
auth_headers: Final = ChatGPTResponsesAPIConfig().validate_environment(
headers={}, # mutable-ok: Responses adapter header contract
model=model,
litellm_params=GenericLiteLLMParams.model_validate(params),
)
return { # mutable-ok: image handler requires dictionaries
**without_oauth_identity_headers(headers),
**auth_headers,
"accept": "application/json",
}
class ChatGPTImageGenerationConfig(GPTImageGenerationConfig):
def validate_environment(
self,
headers: Mapping[str, object],
model: str,
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
) -> dict[str, object]: # mutable-ok: image handler requires dictionaries
return image_headers(headers, model, litellm_params)
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
stream: bool | None = None,
) -> str:
return f"{(api_base or Authenticator.get_api_base()).rstrip('/')}/images/generations"
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
headers: Mapping[str, object],
) -> dict[str, object]: # mutable-ok: image handler requires dictionaries
return {"model": model, "prompt": prompt, **optional_params} # mutable-ok: JSON request serialization
class ChatGPTImageEditConfig(OpenAIImageEditConfig):
def validate_environment(
self,
headers: Mapping[str, object],
model: str,
api_key: str | None = None,
litellm_params: Mapping[str, object] | None = None,
api_base: str | None = None,
) -> dict[str, object]: # mutable-ok: image handler requires dictionaries
return image_headers(headers, model, litellm_params or MappingProxyType({}))
def get_complete_url(self, model: str, api_base: str | None, litellm_params: Mapping[str, object]) -> str:
return f"{(api_base or Authenticator.get_api_base()).rstrip('/')}/images/edits"
def use_multipart_form_data(self) -> bool:
return False
def transform_image_edit_request(
self,
model: str,
prompt: str | None,
image: FileTypes | Sequence[FileTypes] | None,
image_edit_optional_request_params: Mapping[str, object],
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, object],
) -> tuple[dict[str, object], RequestFiles]: # mutable-ok: image handler requires dictionaries
if image_edit_optional_request_params.get("mask") is not None:
raise ValueError("ChatGPT image editing does not support masks")
references: Final = getattr(litellm_params, "images", None)
if references is not None:
if image:
raise ValueError("Specify only one of image or images")
validated: Final = TypeAdapter(tuple[ReferenceImage, ...]).validate_python(references)
if not 1 <= len(validated) <= 5:
raise ValueError("images must contain between 1 and 5 reference images")
return { # mutable-ok: JSON request serialization
"prompt": prompt,
**image_edit_optional_request_params,
"model": model, # the authenticated alias wins over passthrough fields
"images": tuple(item.model_dump() for item in validated),
}, ()
inputs: Final = tuple(image) if isinstance(image, list) else (image,) if image is not None else ()
encoded: Final = tuple(encode_reference(file) for file in inputs)
if not 1 <= len(encoded) <= 5:
raise ValueError("images must contain between 1 and 5 reference images")
return { # mutable-ok: JSON request serialization
"prompt": prompt,
**image_edit_optional_request_params,
"model": model, # the authenticated alias wins over passthrough fields
"images": encoded,
}, ()

View file

@ -0,0 +1,196 @@
import re
from collections.abc import Mapping
from dataclasses import dataclass, field
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, TypeAlias
from unicodedata import category
from urllib.parse import quote, unquote
import httpx
from pydantic import JsonValue
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.llms.chatgpt.realtime import (
ChatGPTRealtime,
configured_realtime_headers,
configured_realtime_query,
realtime_headers,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # shared client has legacy untyped optional params
get_shared_realtime_ssl_context,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection
LiveQuery: TypeAlias = Mapping[str, str | int | float | bool | None | tuple[str | int | float | bool | None, ...]]
LiveBody: TypeAlias = Mapping[str, JsonValue]
LiveOperation: TypeAlias = Literal["fork", "accept", "reject", "refer", "hangup", "content", "attach"]
_PATH: Final = re.compile(r"live/sessions(?:/([^/]+)/(fork|accept|reject|refer|hangup|content|attach))?\Z")
_ROUTING_QUERY: Final = frozenset(("model", "session_id", "call_id", "api_key", "api_base", "authorization"))
@dataclass(frozen=True, slots=True)
class LiveDeployment:
model: str
model_id: str | None = None
provider: Literal["chatgpt", "openai"] = "chatgpt"
api_base: str | None = None
api_key: str | None = field(default=None, repr=False)
extra_headers: Mapping[str, str] = field(default_factory=lambda: MappingProxyType({}), repr=False)
extra_query: LiveQuery = field(default_factory=lambda: MappingProxyType({}), repr=False)
def _validate_session_id(session_id: str) -> None:
candidate: str = session_id # rebind-ok: inspect every decoding layer without recursive stack exhaustion
while True:
if (
not candidate
or candidate in (".", "..")
or any(char in ("/", "\\") or category(char) in ("Cc", "Cs") for char in candidate)
):
raise ValueError("Invalid Live session ID")
decoded: str = unquote(candidate, errors="strict")
if decoded == candidate:
return
candidate = decoded
def _validate_path(path: str) -> None:
match: Final = _PATH.fullmatch(path)
if match is None:
raise ValueError("Invalid Live endpoint")
if match.group(1) is None:
return
session_id: Final = unquote(match.group(1), errors="strict")
if quote(session_id, safe="") != match.group(1):
raise ValueError("Noncanonical Live session path")
_validate_session_id(session_id)
def live_session_path(session_id: str, operation: LiveOperation) -> str:
_validate_session_id(session_id)
path: Final = f"live/sessions/{quote(session_id, safe='')}/{operation}"
_validate_path(path)
return path
class LiveTransport:
def __init__(
self,
deployment: LiveDeployment,
inbound_headers: Mapping[str, str],
*,
http_handler: AsyncHTTPHandler | None = None,
) -> None:
self.deployment = deployment
self._http_handler = http_handler
params: Final = GenericLiteLLMParams.model_validate(
MappingProxyType(
{
"api_base": deployment.api_base,
"extra_query": deployment.extra_query,
}
)
)
self._query = configured_realtime_query(params)
self._headers = (
realtime_headers(params, inbound_headers, deployment.extra_headers)
if deployment.provider == "chatgpt"
else MappingProxyType(
{
**MappingProxyType(
{
key.lower(): value
for key, value in inbound_headers.items()
if key.lower() in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation")
}
),
**configured_realtime_headers(deployment.extra_headers),
"authorization": f"Bearer {deployment.api_key or ''}",
}
)
)
def _url(self, path: str, query: LiveQuery | None, *, websocket: bool) -> str:
_validate_path(path)
base: Final = httpx.URL(
ChatGPTRealtime.get_api_base(self.deployment.api_base)
if self.deployment.provider == "chatgpt"
else self.deployment.api_base or "https://api.openai.com/v1"
)
if base.scheme not in ("https", "http", "wss", "ws") or not base.host or base.userinfo or base.fragment:
raise ValueError("Invalid Live API base")
merged: Final = base.params.merge(query or MappingProxyType({})).merge(self._query)
safe_query: Final = tuple(
(key, value) for key, value in merged.multi_items() if key.lower() not in _ROUTING_QUERY
)
return str(
base.copy_with(
scheme=("wss" if base.scheme in ("https", "wss") else "ws")
if websocket
else ("https" if base.scheme in ("https", "wss") else "http"),
path=f"{base.path.rstrip('/')}/{path}",
params=safe_query,
)
)
async def request(
self,
method: str,
path: str,
body: LiveBody | None = None,
query: LiveQuery | None = None,
) -> httpx.Response:
if (method, path.rsplit("/", 1)[-1]) not in (
("POST", "sessions"),
("POST", "fork"),
("POST", "accept"),
("POST", "reject"),
("POST", "refer"),
("POST", "hangup"),
("GET", "content"),
):
raise ValueError("Invalid Live HTTP operation")
url: Final = self._url(path, query, websocket=False)
handler: Final = self._http_handler or get_async_httpx_client(
llm_provider=LlmProviders.CHATGPT if self.deployment.provider == "chatgpt" else LlmProviders.OPENAI,
params={"follow_redirects": False},
)
headers: Final = {**self._headers, "content-type": "application/json"}
if method == "GET":
return await handler.get(url, headers=headers, timeout=60, follow_redirects=False)
try:
return await handler.post(
url,
headers=headers,
json=dict(body) if body is not None else None, # mutable-ok: JSON encoder requires a concrete dict
timeout=60,
)
except httpx.HTTPStatusError as error:
return error.response
async def connect(self, path: str, query: LiveQuery | None = None) -> "ClientConnection":
import websockets
class DirectConnect(websockets.connect):
def process_redirect(self, exc: Exception) -> Exception:
return exc
if path != "live/sessions" and path.rsplit("/", 1)[-1] not in ("attach", "fork"):
raise ValueError("Invalid Live WebSocket operation")
url: Final = self._url(path, query, websocket=True)
ssl_context: Final = get_shared_realtime_ssl_context() if url.startswith("wss://") else None
return await DirectConnect(
url,
additional_headers=self._headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
max_queue=16,
ssl=True if ssl_context is False else ssl_context,
open_timeout=20,
close_timeout=10,
)

View file

@ -0,0 +1,290 @@
from collections.abc import Mapping
from enum import Enum, auto
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from httpx import URL, QueryParams, Response
from pydantic import TypeAdapter
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.exceptions import AuthenticationError
from litellm.llms.openai.realtime.handler import OpenAIRealtime
from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import get_model_info
from .authenticator import Authenticator
from .common_utils import without_oauth_identity_headers
from .responses.transformation import ChatGPTResponsesAPIConfig
if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection
class CallAccounting(Enum):
SUPERVISED = auto()
def accounts_for_call_usage(params: GenericLiteLLMParams) -> bool:
return getattr(params, "chatgpt_call_accounting", None) is not CallAccounting.SUPERVISED
def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping[str, str]:
validated: Final = TypeAdapter(Mapping[str, str]).validate_python(
without_oauth_identity_headers(headers or MappingProxyType({}))
)
return MappingProxyType({key.lower(): value for key, value in validated.items()})
def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str | tuple[str, ...]]:
inbound: Final = TypeAdapter(Mapping[str, str]).validate_python(
getattr(params, "chatgpt_realtime_client_query", None) or MappingProxyType({})
)
configured: Final = TypeAdapter(
Mapping[str, str | int | float | bool | None | tuple[str | int | float | bool | None, ...]]
).validate_python(getattr(params, "extra_query", None) or MappingProxyType({}))
merged: Final = QueryParams(
tuple((key, value) for key, value in inbound.items() if key in ("intent", "architecture"))
).merge(configured)
return MappingProxyType(
{key: merged[key] if len(merged.get_list(key)) == 1 else tuple(merged.get_list(key)) for key in merged}
)
def realtime_call_headers(params: GenericLiteLLMParams) -> dict[str, str]: # mutable-ok: HTTP handler header contract
inbound: Final = TypeAdapter(Mapping[str, str]).validate_python(
getattr(params, "chatgpt_realtime_client_headers", None) or MappingProxyType({})
)
configured: Final = TypeAdapter(Mapping[str, object]).validate_python(
getattr(params, "extra_headers", None) or MappingProxyType({})
)
return { # mutable-ok: HTTP handler header contract
**MappingProxyType(
{
key.lower(): value
for key, value in inbound.items()
if key.lower() in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation")
}
),
**configured_realtime_headers(configured),
}
def realtime_headers(
params: GenericLiteLLMParams, headers: Mapping[str, str], extra_headers: Mapping[str, object] | None = None
) -> dict[str, str]: # mutable-ok: HTTP handler header contract
forwarded: Final = MappingProxyType(
{
key.lower(): value
for key, value in headers.items()
if key.lower() in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation")
}
)
return { # mutable-ok: HTTP handler updates headers
**ChatGPTResponsesAPIConfig().validate_environment(
headers={}, # mutable-ok: Responses adapter header contract
model="",
litellm_params=params,
),
**forwarded,
**configured_realtime_headers(extra_headers),
}
def realtime_endpoint(model: str) -> str:
try:
model_info: Final = get_model_info(model, custom_llm_provider="chatgpt")
except Exception: # noqa: BLE001 # get_model_info raises bare Exception for unmapped models
return "realtime"
return "live" if "/v1/live" in (model_info.get("supported_endpoints") or ()) else "realtime"
class ChatGPTRealtime(OpenAIRealtime):
async def open_call_connection(self, model: str, api_base: str) -> "ClientConnection":
import websockets
url: Final = self._construct_url(api_base, RealtimeQueryParams(model=model))
return await websockets.connect(
url,
additional_headers=self._profile_headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=self._get_ssl_config(url),
open_timeout=20,
)
async def close_call(self, connection: "ClientConnection", model: str, api_base: str) -> None:
from websockets.exceptions import ConnectionClosed
if realtime_endpoint(model) == "live":
try:
await connection.send('{"type":"session.close"}')
return
except (ConnectionClosed, OSError):
await self.hangup_call(api_base)
return
await self.hangup_call(api_base)
async def hangup_call(self, api_base: str) -> None:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.utils import LlmProviders
base: Final = URL(api_base)
url: Final = base.copy_with(
scheme="https" if base.scheme in ("https", "wss") else "http",
path=f"{base.path.rstrip('/')}/realtime/calls/{self._call_id}/hangup",
params=tuple(
(key, value)
for key, value in QueryParams(self._extra_query).multi_items()
if key not in ("model", "call_id")
),
)
client: Final = get_async_httpx_client(llm_provider=LlmProviders.CHATGPT)
response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10)
response.raise_for_status()
@staticmethod
def get_api_base(api_base: str | None = None) -> str:
return api_base or Authenticator.get_api_base(default_base="https://api.openai.com/v1")
def __init__(
self,
params: GenericLiteLLMParams,
headers: Mapping[str, str],
extra_headers: Mapping[str, object] | None = None,
) -> None:
super().__init__()
self._profile_headers = realtime_headers(params, headers, extra_headers)
self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None))
self._extra_query = configured_realtime_query(params)
self._account_usage = accounts_for_call_usage(params)
def _get_default_api_base(self) -> str:
return self.get_api_base()
def _resolve_api_key(self, api_key: str | None) -> str:
return "chatgpt-oauth"
def _accounts_for_call_usage(self) -> bool:
return self._account_usage
def _get_additional_headers(
self, api_key: str, *, openai_beta_realtime: bool = False
) -> dict[str, str]: # mutable-ok: HTTP handler header contract
return { # mutable-ok: HTTP handler updates headers
**(MappingProxyType({"OpenAI-Beta": "realtime=v1"}) if openai_beta_realtime else MappingProxyType({})),
**self._profile_headers,
}
def _construct_url(self, api_base: str, query_params: RealtimeQueryParams) -> str:
base: Final = URL(api_base)
endpoint: Final = realtime_endpoint(query_params.get("model", ""))
if self._call_id:
gateway_query: Final = tuple(
(key, value)
for key, value in QueryParams(self._extra_query).multi_items()
if key not in ("model", "call_id")
)
return str(
base.copy_with(
scheme="wss" if base.scheme in ("https", "wss") else "ws",
path=f"{base.path.rstrip('/')}/{endpoint}/{self._call_id}"
if endpoint == "live"
else f"{base.path.rstrip('/')}/realtime",
params=gateway_query + (() if endpoint == "live" else (("call_id", self._call_id),)),
)
)
return str(
base.copy_with(
scheme="wss" if base.scheme in ("https", "wss") else "ws",
path=f"{base.path.rstrip('/')}/{endpoint}",
params=QueryParams(TypeAdapter(Mapping[str, str | None]).validate_python(query_params)).merge(
tuple(
(key, value)
for key, value in QueryParams(self._extra_query).multi_items()
if key not in ("model", "call_id")
)
),
)
)
class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig):
realtime_calls_json: Final = True
def __init__(self, params: GenericLiteLLMParams, use_codex_backend: bool = True) -> None:
self._params = params
self._use_codex_backend = use_codex_backend
def get_api_base(
self,
api_base: str | None,
**kwargs: object, # kwargs-ok: provider interface accepts optional credentials
) -> str:
return api_base or (Authenticator.get_api_base() if self._use_codex_backend else ChatGPTRealtime.get_api_base())
def resolve_api_base(self, api_base: str | None, dynamic_api_base: str | None) -> str:
return self.get_api_base(api_base)
def get_realtime_calls_extra_headers(
self, headers: dict[str, object] | None
) -> dict[str, object]: # mutable-ok: shared HTTP handler accepts a mutable header dictionary
return {**realtime_call_headers(self._params)} # mutable-ok: shared HTTP header contract
def get_api_key(
self,
api_key: str | None,
**kwargs: object, # kwargs-ok: provider interface accepts optional credentials
) -> str:
return "chatgpt-oauth"
def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
query: Final = configured_realtime_query(self._params)
return str(URL(f"{self.get_api_base(api_base).rstrip('/')}/realtime/calls", params=query))
def transform_realtime_calls_response(
self, response: Response, model: str, model_id: str | None, headers: Mapping[str, object] | None
) -> Response:
response.extensions["chatgpt_realtime"] = MappingProxyType(
{
"model": model,
"model_id": model_id,
"api_base": ChatGPTRealtime.get_api_base(self._params.api_base),
"extra_headers": configured_realtime_headers(headers),
"extra_query": configured_realtime_query(self._params),
}
)
return response
def get_realtime_calls_headers(
self, ephemeral_key: str
) -> dict[str, str]: # mutable-ok: HTTP handler header contract
if ephemeral_key:
raise AuthenticationError(
message="ChatGPT realtime calls require an authenticated JSON or multipart offer",
llm_provider="chatgpt",
model="",
)
return realtime_headers(self._params, MappingProxyType({}))
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
api_key: str | None = None,
) -> dict[str, str]: # mutable-ok: HTTP handler header contract
return { # mutable-ok: HTTP handler updates headers
**realtime_headers(self._params, headers),
"Content-Type": "application/json",
}
def get_complete_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
return f"{self.get_api_base(api_base).rstrip('/')}/realtime/client_secrets"
def get_transcription_session_url(
self,
api_base: str | None,
model: str,
api_version: str | None = None,
) -> str:
return f"{self.get_api_base(api_base).rstrip('/')}/realtime/transcription_sessions"

View file

@ -106,6 +106,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
"reasoning",
"previous_response_id",
"truncation",
"text",
}
return {k: v for k, v in request.items() if k in allowed_keys}

View file

@ -6251,6 +6251,9 @@ class BaseLLMHTTPHandler:
Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and
header auth when available; falls back to the legacy OpenAI-style defaults.
"""
from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.OPENAI,
@ -6276,7 +6279,11 @@ class BaseLLMHTTPHandler:
}
if extra_headers:
headers.update(extra_headers)
headers.update(
without_oauth_identity_headers(extra_headers)
if isinstance(provider_config, ChatGPTRealtimeHTTPConfig)
else extra_headers
)
logging_obj.pre_call(
input=request_data,
@ -6327,6 +6334,9 @@ class BaseLLMHTTPHandler:
- sdp: the SDP offer (text)
- session: JSON string with {"type": "realtime", "model": "...", ...}
"""
from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.OPENAI,
@ -6344,11 +6354,15 @@ class BaseLLMHTTPHandler:
}
if extra_headers:
headers.update(extra_headers)
headers.update(
without_oauth_identity_headers(extra_headers)
if isinstance(provider_config, ChatGPTRealtimeHTTPConfig)
else extra_headers
)
# Build multipart form data: sdp + session JSON
session_data: Final = session_config or {}
if "type" not in session_data:
if "type" not in session_data and not getattr(provider_config, "realtime_calls_json", False):
session_data["type"] = "realtime"
if "model" not in session_data and model:
session_data["model"] = model
@ -6371,6 +6385,13 @@ class BaseLLMHTTPHandler:
)
try:
if getattr(provider_config, "realtime_calls_json", False):
return await async_httpx_client.post(
url=url,
headers=headers,
json={"sdp": sdp_text, "session": session_data}, # mutable-ok: JSON signaling payload
timeout=timeout,
)
return await async_httpx_client.post(
url=url,
headers=headers,
@ -6572,6 +6593,14 @@ class BaseLLMHTTPHandler:
raise Exception(f"Unexpected error while closing WebSocket: {close_error}")
return None
@staticmethod
def _image_extra_headers(custom_llm_provider: str, headers: Mapping[str, object]) -> Mapping[str, object]:
if custom_llm_provider == "chatgpt":
from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers
return without_oauth_identity_headers(headers)
return headers
def image_edit_handler(
self,
model: str,
@ -6628,7 +6657,7 @@ class BaseLLMHTTPHandler:
)
if extra_headers:
headers.update(extra_headers)
headers.update(self._image_extra_headers(custom_llm_provider, extra_headers))
api_base: Final = image_edit_provider_config.get_complete_url(
model=model,
@ -6729,7 +6758,7 @@ class BaseLLMHTTPHandler:
)
if extra_headers:
headers.update(extra_headers)
headers.update(self._image_extra_headers(custom_llm_provider, extra_headers))
api_base: Final = image_edit_provider_config.get_complete_url(
model=model,
@ -6848,7 +6877,7 @@ class BaseLLMHTTPHandler:
)
if extra_headers:
headers.update(extra_headers)
headers.update(self._image_extra_headers(custom_llm_provider, extra_headers))
api_base: Final = image_generation_provider_config.get_complete_url(
model=model,
@ -6956,7 +6985,7 @@ class BaseLLMHTTPHandler:
)
if extra_headers:
headers.update(extra_headers)
headers.update(self._image_extra_headers(custom_llm_provider, extra_headers))
api_base: Final = image_generation_provider_config.get_complete_url(
model=model,

View file

@ -39,6 +39,14 @@ class OpenAIRealtime(OpenAIChatCompletion):
"""
return "https://api.openai.com/"
def _resolve_api_key(self, api_key: str | None) -> str:
if api_key is None:
raise ValueError("api_key is required for OpenAI realtime calls")
return api_key
def _accounts_for_call_usage(self) -> bool:
return True
def _get_additional_headers(
self,
api_key: str,
@ -118,6 +126,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
query_params: RealtimeQueryParams | None = None,
user_api_key_dict: object | None = None,
litellm_metadata: dict | None = None,
account_usage: bool = True,
**kwargs: object,
):
import websockets
@ -125,8 +134,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
if api_base is None:
api_base = self._get_default_api_base()
if api_key is None:
raise ValueError("api_key is required for OpenAI realtime calls")
resolved_api_key: Final = self._resolve_api_key(api_key)
# Use all query params if provided, else fallback to just model
if query_params is None:
@ -144,12 +152,12 @@ class OpenAIRealtime(OpenAIChatCompletion):
"If your client expects beta event names, add 'OpenAI-Beta: realtime=v1' "
"to the WebSocket headers sent to the LiteLLM proxy."
)
headers: Final = self._get_additional_headers(api_key, openai_beta_realtime=openai_beta_realtime)
headers: Final = self._get_additional_headers(resolved_api_key, openai_beta_realtime=openai_beta_realtime)
# Log a masked request preview consistent with other endpoints.
logging_obj.pre_call(
input=None,
api_key=api_key,
api_key=resolved_api_key,
additional_args={
"api_base": url,
"headers": headers,
@ -173,6 +181,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
model if (query_params or {}).get("intent") == "transcription" else None
),
event_normalizer=self._make_event_normalizer(),
account_usage=account_usage and self._accounts_for_call_usage(),
)
await realtime_streaming.bidirectional_forward()

View file

@ -31156,6 +31156,16 @@
"max_tokens": 8191,
"mode": "embedding"
},
"chatgpt/gpt-live-1-codex": {
"litellm_provider": "chatgpt",
"mode": "realtime",
"supported_endpoints": [
"/v1/realtime/calls",
"/v1/live"
],
"supports_audio_input": true,
"supports_audio_output": true
},
"chatgpt/gpt-5.5": {
"litellm_provider": "chatgpt",
"source": "https://platform.openai.com/docs/models/gpt-5.5",

View file

@ -231,10 +231,15 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
"/watsonx/",
),
),
LazyFeature(
name="live",
module_path="litellm.proxy.realtime_endpoints.live",
path_prefixes=("/openai/v1/live/sessions", "/v1/live/sessions", "/live/sessions"),
),
LazyFeature(
name="realtime",
module_path="litellm.proxy.realtime_endpoints.endpoints",
path_prefixes=("/openai/v1/realtime", "/v1/realtime", "/realtime"),
path_prefixes=("/openai/v1/realtime", "/v1/realtime", "/realtime", "/openai/v1/live", "/v1/live", "/live"),
),
LazyFeature(
name="anthropic_passthrough",

File diff suppressed because it is too large Load diff

View file

@ -114,6 +114,10 @@ class ReconcileOutcome(NamedTuple):
live_after: frozenset[str] | None
class InternalRequestOrigin(enum.Enum):
REALTIME_OBSERVER = enum.auto()
class SupportedDBObjectType(str, enum.Enum):
"""
Supported database object types for fine-grained DB storage control.
@ -271,8 +275,8 @@ class Litellm_EntityType(enum.Enum):
def hash_token(token: str):
import hashlib
# Hash the string using SHA-256
hashed_token: Final = hashlib.sha256(token.encode()).hexdigest()
# This digest is an opaque lookup identifier, not a password hash.
hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest()
return hashed_token
@ -422,6 +426,36 @@ class LiteLLMRoutes(enum.Enum):
"/realtime?{model}",
"/v1/realtime?{model}",
"/openai/v1/realtime?{model}",
"/live",
"/v1/live",
"/v1/live/{call_id}",
"/openai/v1/live",
"/live/{call_id}",
"/openai/v1/live/{call_id}",
"/live/sessions",
"/live/sessions/{session_id}/attach",
"/live/sessions/{session_id}/fork",
"/live/sessions/{session_id}/content",
"/live/sessions/{session_id}/accept",
"/live/sessions/{session_id}/reject",
"/live/sessions/{session_id}/refer",
"/live/sessions/{session_id}/hangup",
"/v1/live/sessions",
"/v1/live/sessions/{session_id}/attach",
"/v1/live/sessions/{session_id}/fork",
"/v1/live/sessions/{session_id}/content",
"/v1/live/sessions/{session_id}/accept",
"/v1/live/sessions/{session_id}/reject",
"/v1/live/sessions/{session_id}/refer",
"/v1/live/sessions/{session_id}/hangup",
"/openai/v1/live/sessions",
"/openai/v1/live/sessions/{session_id}/attach",
"/openai/v1/live/sessions/{session_id}/fork",
"/openai/v1/live/sessions/{session_id}/content",
"/openai/v1/live/sessions/{session_id}/accept",
"/openai/v1/live/sessions/{session_id}/reject",
"/openai/v1/live/sessions/{session_id}/refer",
"/openai/v1/live/sessions/{session_id}/hangup",
# realtime (GA WebRTC HTTP routes)
"/realtime/client_secrets",
"/v1/realtime/client_secrets",

View file

@ -2398,11 +2398,18 @@ async def get_team_membership(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
raise_on_error: bool = True,
) -> Optional["LiteLLM_TeamMembership"]:
"""
Returns team membership object if user is member of team.
Do a isolated check for team membership vs. doing a combined key + team + user + team-membership check, as key might come in frequently for different users/teams. Larger call will slowdown query time. This way we get to cache the constant (key/team/user info) and only update based on the changing value (team membership).
``raise_on_error`` defaults to True because the callers that apply member-level limits -- the budget and
model-scope checks in ``common_checks``, the JWT team resolution, and the compact summary gate -- cannot
tell an absent row apart from a failed read, so swallowing an outage there hands the member whatever the
team allows. A caller that only attributes grants, and can proceed with the lists it already holds,
passes False and degrades to "no member-level scope".
"""
if user_id is None or team_id is None:
return None
@ -2416,7 +2423,17 @@ async def get_team_membership(
inflight: Final[object] = _team_membership_inflight.get(_key)
if isinstance(inflight, asyncio.Task):
return _membership_from_shared_load(await asyncio.shield(inflight))
try:
return _membership_from_shared_load(await asyncio.shield(inflight))
except Exception:
verbose_proxy_logger.exception(
"Error getting team membership for user_id: %s, team_id: %s",
user_id,
team_id,
)
if raise_on_error:
raise
return None
if prisma_client is None:
raise Exception("No db connected")
@ -2439,7 +2456,17 @@ async def get_team_membership(
_team_membership_inflight.pop(_key, None)
task.add_done_callback(_clear_inflight)
return _membership_from_shared_load(await asyncio.shield(task))
try:
return _membership_from_shared_load(await asyncio.shield(task))
except Exception:
verbose_proxy_logger.exception(
"Error getting team membership for user_id: %s, team_id: %s",
user_id,
team_id,
)
if raise_on_error:
raise
return None
def model_in_access_group(model: str, team_models: list[str] | None, llm_router: Router | None) -> bool:
@ -4702,6 +4729,8 @@ async def _team_member_granted_models(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
*,
strict_grant_lookup: bool = False,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> Sequence[str]:
@ -4716,6 +4745,10 @@ async def _team_member_granted_models(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
# Spelled out because it is the one caller that wants the opposite of the default: outside
# strict mode this walk only attributes grants, so an unreadable member scope degrades to
# "no member-level scope" instead of failing the request.
raise_on_error=strict_grant_lookup,
)
return () if team_membership is None else _member_allowed_models(team_membership)
@ -4726,6 +4759,8 @@ async def _org_granted_models(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
*,
strict_grant_lookup: bool = False,
) -> Sequence[str]:
"""The org allowlist reached through the key, or through its team when the key names no org."""
org_id: Final = valid_token.org_id or (team_object.organization_id if team_object is not None else None)
@ -4741,6 +4776,8 @@ async def _org_granted_models(
)
except Exception as e: # noqa: BLE001 # fail-safe: attribution degrades to "no org grant", it must never break auth
verbose_proxy_logger.debug("access group attribution: org lookup failed: %s", e)
if strict_grant_lookup:
raise
return ()
return org_object.models if org_object is not None else ()
@ -4752,6 +4789,8 @@ async def _granted_model_lists(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
*,
strict_grant_lookup: bool = False,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> tuple[Sequence[str], ...]:
@ -4765,6 +4804,7 @@ async def _granted_model_lists(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
strict_grant_lookup=strict_grant_lookup,
team_membership=team_membership,
team_membership_loaded=team_membership_loaded,
),
@ -4775,6 +4815,7 @@ async def _granted_model_lists(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
strict_grant_lookup=strict_grant_lookup,
),
)
@ -4861,6 +4902,8 @@ async def collect_matched_model_access_groups(
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
*,
strict_grant_lookup: bool = False,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> tuple[str, ...]:
@ -4878,7 +4921,9 @@ async def collect_matched_model_access_groups(
The whole walk is gated on the budget registry, because collecting every match costs a full scan
of each allowlist where the plain access check stops at the first hit. An empty registry means no
group carries a budget, so there is nothing to attribute and no work worth doing.
group carries a budget, so there is nothing to attribute and no work worth doing. The strict
lookup mode is reserved for enforcement paths that must not treat an unavailable inherited grant
as absent; the default remains fail-safe attribution for ordinary request telemetry.
"""
if model is None or valid_token is None or llm_router is None or prisma_client is None:
return ()
@ -4908,6 +4953,7 @@ async def collect_matched_model_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
strict_grant_lookup=strict_grant_lookup,
team_membership=team_membership,
team_membership_loaded=team_membership_loaded,
)
@ -6537,7 +6583,9 @@ def is_model_allowed_by_pattern(model: str, allowed_model_pattern: str) -> bool:
bool: True if model matches the pattern, False otherwise
"""
if "*" in allowed_model_pattern:
pattern: Final = f"^{allowed_model_pattern.replace('*', '.*')}$"
# Treat the configured model pattern as a glob; only '*' is special.
escaped_pattern: Final = re.escape(allowed_model_pattern)
pattern: Final = "^" + escaped_pattern.replace("\\*", ".*") + "$"
return bool(re.match(pattern, model))
return False

View file

@ -1710,6 +1710,7 @@ _MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS: Final = ("/evals",)
_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS: Final = (
"/realtime/client_secrets",
"/realtime/calls",
"/live/sessions",
)
_MODEL_ROUTING_ID_FIELDS: Final = (
"file_id",
@ -1900,15 +1901,27 @@ def _extract_model_candidates_from_request(
uses_completion_model_sources: Final = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS
)
uses_session_model: Final = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS
) or route.rstrip("/") in ("/live", "/v1/live", "/openai/v1/live")
session: Final[object] = request_data.get("session") if uses_session_model else None
parsed_session: Final[object] = safe_json_loads(session) if isinstance(session, str) else session
session_model: Final[object] = parsed_session.get("model") if isinstance(parsed_session, dict) else None
if (
uses_session_model
and not _route_matches_any_marker(route=route, markers=("/realtime/client_secrets",))
and isinstance(session_model, str)
and session_model
):
candidates.append(session_model)
return candidates
body_model: Final = request_data.get("model")
_append_model_candidates(candidates, body_model)
if uses_body_target_model_sources or not body_model:
_append_model_candidates(candidates, request_data.get("target_model_names"))
if _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS):
session: Final = request_data.get("session")
if isinstance(session, dict):
_append_model_candidates(candidates, session.get("model"))
if uses_session_model:
_append_model_candidates(candidates, TypeAdapter[object](object).validate_python(session_model))
if uses_completion_model_sources and isinstance(request_data.get("completion"), dict):
_append_model_candidates(candidates, request_data["completion"].get("model"))

View file

@ -2485,7 +2485,7 @@ class JWTAuthManager:
return JWTIdentity(user_id=user_id if is_admin else canonical_id, user_object=user, agent_id=agent_id)
@staticmethod
async def authorize_jwt(
async def authorize_jwt( # noqa: C901 # preserves the established JWT authorization flow split from auth_builder
api_key: str,
jwt_handler: JWTHandler,
request_data: dict[str, object],

View file

@ -452,7 +452,7 @@ async def _check_key_model_budget_with_fallback(
model=model_name,
)
except litellm.BudgetExceededError as e:
if request_data.get("model") != model_name:
if request_data.get("model") != model_name or request.scope.get("litellm_pinned_realtime_model") == model_name:
raise e
fallback_model: Final = await model_max_budget_limiter.get_fallback_model_within_budget(
user_api_key_dict=valid_token,
@ -641,6 +641,38 @@ def _apply_budget_limits_to_end_user_params(
verbose_proxy_logger.debug("Applied budget limits to end user %s", end_user_id)
def get_websocket_api_key(websocket: WebSocket) -> str | None:
from litellm.proxy.proxy_server import general_settings
custom_header: Final = general_settings.get("litellm_key_header_name")
if isinstance(custom_header, str):
if not websocket.headers.get(custom_header):
return None
request: Final = Request(
{"type": "http", "headers": websocket.scope.get("headers", [])} # mutable-ok: ASGI request scope
)
return get_api_key_from_custom_header(request, custom_header)
custom_key: Final = websocket.headers.get("x-litellm-api-key")
if custom_key is not None:
return _get_bearer_token_or_received_api_key(custom_key)
authorization: Final = websocket.headers.get("authorization")
if authorization:
if not authorization.startswith("Bearer "):
raise HTTPException(status_code=403, detail="Invalid Authorization header format")
return authorization[len("Bearer ") :].strip()
api_key: Final = websocket.headers.get("api-key")
if api_key:
return api_key
return next(
(
protocol.strip().removeprefix("openai-insecure-api-key.")
for protocol in websocket.headers.get("sec-websocket-protocol", "").split(",")
if protocol.strip().startswith("openai-insecure-api-key.")
),
None,
)
async def user_api_key_auth_websocket(websocket: WebSocket) -> UserAPIKeyAuth:
return await user_api_key_auth_websocket_for_model(websocket, model=websocket.query_params.get("model"))
@ -668,36 +700,29 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
request._url = websocket.url
async def return_body():
return _realtime_request_body(model)
request.body = return_body
authorization: Final = websocket.headers.get("authorization")
# If no Authorization header, try the api-key header
if not authorization:
api_key = websocket.headers.get("api-key")
if not api_key:
# Try extracting from WebSocket subprotocol (browser clients)
for protocol in websocket.headers.get("sec-websocket-protocol", "").split(","):
protocol = protocol.strip()
if protocol.startswith("openai-insecure-api-key."):
api_key = protocol[len("openai-insecure-api-key.") :]
break
if not api_key:
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
raise HTTPException(status_code=403, detail="No API key provided")
else:
# Extract the API key from the Bearer token
if not authorization.startswith("Bearer "):
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
raise HTTPException(status_code=403, detail="Invalid Authorization header format")
api_key = authorization[len("Bearer ") :].strip()
try:
api_key: Final = get_websocket_api_key(websocket)
except HTTPException:
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
raise
if not api_key:
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
raise HTTPException(status_code=403, detail="No API key provided")
# Call user_api_key_auth with the extracted API key
# Note: You'll need to modify this to work with WebSocket context if needed
try:
from litellm.proxy.realtime_endpoints.call_sessions import decode_call
call_token: Final = websocket.path_params.get("call_id") or websocket.query_params.get("call_id")
resolved_model: Final = decode_call(call_token, f"Bearer {api_key}").alias if call_token is not None else model
if call_token is not None:
request.scope["litellm_pinned_realtime_model"] = resolved_model
async def return_body():
return _realtime_request_body(resolved_model)
request.body = return_body
return await user_api_key_auth(request=request, api_key=f"Bearer {api_key}")
except Exception as e:
if is_invalid_virtual_key_error(e):

View file

@ -2015,7 +2015,9 @@ class ProxyBaseLLMRequestProcessing:
model: str | None = None,
llm_router: Router | None = None,
rate_limited_model: str | None = None,
*,
skip_guardrails: bool = False,
internal_realtime_observer: bool = False,
) -> tuple[dict, LiteLLMLoggingObj]:
start_time: Final = datetime.now() # start before calling guardrail hooks
@ -2209,6 +2211,11 @@ class ProxyBaseLLMRequestProcessing:
data=self.data,
call_type=route_type,
skip_guardrails=skip_guardrails,
**(
MappingProxyType({"internal_realtime_observer": True})
if internal_realtime_observer
else MappingProxyType({})
),
)
await _enforce_guardrail_added_tag_budgets(
data=self.data,

View file

@ -336,6 +336,17 @@ def model_access_group_cache_key(access_group_name: str) -> str:
return f"model_access_group:{access_group_name}"
def live_model_access_group_limits_cache_key(access_group_name: str) -> str:
"""Cache key the Live delegation gate stores one access group's full limit row under.
The gate needs the rpm and tpm columns that ``model_access_group:{name}`` flattens away, so it
keeps its own entry next to the flattened one. Any eviction of the flattened entry must clear
this key too: the gate reads cache-first, and a raised or lowered group limit left cached here
keeps permitting or refusing managed delegation until the entry's TTL expires (LIT-3803).
"""
return f"live:model_access_group_limits:{access_group_name}"
def model_access_group_registry_cache_key() -> str:
"""Cache key for the set of model access group names that have a budget row."""
return "model_access_group_registry"

View file

@ -1,9 +1,10 @@
import asyncio
import sys
from collections.abc import Mapping
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter
from typing_extensions import TypedDict
import litellm
@ -12,7 +13,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.exceptions import RateLimitType
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth
from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, InternalRequestOrigin, UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
get_key_model_rpm_limit,
get_key_model_tpm_limit,
@ -51,11 +52,81 @@ class CacheObject(TypedDict):
request_count_end_user_id: dict | None
class _RealtimeAttachmentReservations(BaseModel):
cache_keys: tuple[str, ...] = ()
global_acquired: bool = False
def acquire(self, key: str) -> None:
self.cache_keys = tuple(dict.fromkeys((*self.cache_keys, key)))
def acquire_global(self) -> None:
self.global_acquired = True
def take(self) -> tuple[tuple[str, ...], bool]:
owned: Final = (self.cache_keys, self.global_acquired)
self.cache_keys = ()
self.global_acquired = False
return owned
_RELEASE_REALTIME_COUNTER_LUA: Final = """
local raw = redis.call('GET', KEYS[1])
if not raw then return 0 end
local value = cjson.decode(raw)
value.current_requests = math.max(value.current_requests - 1, 0)
redis.call('SET', KEYS[1], cjson.encode(value), 'KEEPTTL')
return 1
"""
class _PROXY_MaxParallelRequestsHandler(CustomLogger):
# Class variables or attributes
def __init__(self, internal_usage_cache: InternalUsageCache):
self.internal_usage_cache = internal_usage_cache
def begin_realtime_attachment(self, request_data: dict[str, object]) -> None:
request_data["_legacy_realtime_attachment_reservations"] = ( # rebind-ok: request-scoped cleanup receipt
_RealtimeAttachmentReservations()
)
async def async_release_realtime_attachment(
self, request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth
) -> None:
receipt: Final = request_data.get("_legacy_realtime_attachment_reservations")
if not isinstance(receipt, _RealtimeAttachmentReservations):
return
keys, global_acquired = receipt.take()
if global_acquired:
await self.internal_usage_cache.async_increment_cache(
key="global_max_parallel_requests",
value=-1,
local_only=True,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
for key in keys:
await self._release_realtime_counter(key)
async def _release_realtime_counter(self, key: str) -> None:
local: Final = self.internal_usage_cache.dual_cache.in_memory_cache
remote: Final = self.internal_usage_cache.dual_cache.redis_cache
raw: Final[object] = local.get_cache(key)
current: Final = TypeAdapter[Mapping[str, int] | None](Mapping[str, int] | None).validate_python(raw)
updated: Final = (
{ # mutable-ok: shared cache counter dict
**current,
"current_requests": max(current["current_requests"] - 1, 0),
}
if current is not None
else None
)
if updated is not None:
local.set_cache(key, updated, ttl=60)
if remote is not None:
release: Final = remote.async_register_script(_RELEASE_REALTIME_COUNTER_LUA)
await release(keys=(key,), args=())
if local.get_cache(key) is updated:
local.delete_cache(key)
def print_verbose(self, print_statement):
try:
verbose_proxy_logger.debug(print_statement)
@ -143,6 +214,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
local_only=True,
)
receipt: Final = data.get("_legacy_realtime_attachment_reservations")
if isinstance(receipt, _RealtimeAttachmentReservations):
receipt.acquire(request_count_api_key)
return new_val
def time_to_next_minute(self) -> float:
@ -300,6 +374,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
local_only=True,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
receipt: Final = data.get("_legacy_realtime_attachment_reservations")
if isinstance(receipt, _RealtimeAttachmentReservations):
receipt.acquire_global()
_model = data.get("model", None)
current_date: Final = datetime.now().strftime("%Y-%m-%d")
@ -481,6 +558,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
values_to_update_in_cache=values_to_update_in_cache,
)
if isinstance(data.get("_legacy_realtime_attachment_reservations"), _RealtimeAttachmentReservations):
await self.internal_usage_cache.async_batch_set_cache(
cache_list=values_to_update_in_cache,
ttl=60,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
return
asyncio.create_task(
self.internal_usage_cache.async_batch_set_cache(
cache_list=values_to_update_in_cache,
@ -490,6 +574,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
)
async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time):
releases_slot: Final = kwargs.get("internal_request_origin") is not InternalRequestOrigin.REALTIME_OBSERVER
from litellm.proxy.common_utils.callback_utils import (
get_model_group_from_litellm_kwargs,
)
@ -522,7 +607,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
# Setup values
# ------------
if global_max_parallel_requests is not None:
if releases_slot and global_max_parallel_requests is not None:
# get value from cache
_key: Final = "global_max_parallel_requests"
# decrement
@ -553,13 +638,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_requests": int(releases_slot),
"current_tpm": 0,
"current_rpm": 0,
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -594,13 +679,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_requests": int(releases_slot),
"current_tpm": 0,
"current_rpm": 0,
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -620,13 +705,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_tpm": total_tokens,
"current_rpm": 1,
"current_requests": int(releases_slot),
"current_tpm": total_tokens if releases_slot else 0,
"current_rpm": int(releases_slot),
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -646,13 +731,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_tpm": total_tokens,
"current_rpm": 1,
"current_requests": int(releases_slot),
"current_tpm": total_tokens if releases_slot else 0,
"current_rpm": int(releases_slot),
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -672,13 +757,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_tpm": total_tokens,
"current_rpm": 1,
"current_requests": int(releases_slot),
"current_tpm": total_tokens if releases_slot else 0,
"current_rpm": int(releases_slot),
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -695,6 +780,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
self.print_verbose(e)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
if kwargs.get("internal_request_origin") is InternalRequestOrigin.REALTIME_OBSERVER:
return
try:
self.print_verbose("Inside Max Parallel Request Failure Hook")
litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs)

View file

@ -10,8 +10,8 @@ import itertools
import logging
import os
import uuid
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence, Set
from contextlib import asynccontextmanager
from collections.abc import AsyncGenerator, Awaitable, Callable, Generator, Mapping, Sequence, Set
from contextlib import asynccontextmanager, contextmanager
from contextvars import ContextVar
from dataclasses import dataclass, field
from datetime import datetime, timezone
@ -68,6 +68,7 @@ from litellm.proxy.hooks.batch_enqueued_tokens import (
canonical_provider_batch_id,
)
from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease, is_realtime_call_attachment
from litellm.router_utils.add_retry_fallback_headers import (
ensure_response_additional_headers,
response_has_hidden_params,
@ -396,6 +397,23 @@ end
return results
"""
PARALLEL_RENEW_SCRIPT: Final = """
local clock = redis.call('TIME')
local now = tonumber(clock[1])
local ttl = tonumber(ARGV[2])
for i = 1, #KEYS do
local score = redis.call('ZSCORE', KEYS[i], ARGV[1])
if not score or tonumber(score) <= now - ttl then
return {0}
end
end
for i = 1, #KEYS do
redis.call('ZADD', KEYS[i], 'XX', now, ARGV[1])
redis.call('EXPIRE', KEYS[i], ttl)
end
return {1}
"""
TOKEN_INCREMENT_SCRIPT: Final = """
local results = {}
@ -523,9 +541,17 @@ class ParallelRequestGauge(TypedDict):
descriptor_key: str
def _without_parallel_limit(descriptor: RateLimitDescriptor) -> RateLimitDescriptor:
rate_limit: Final[RateLimitDescriptorRateLimitObject] = {
**(descriptor.get("rate_limit") or MappingProxyType({})),
"max_parallel_requests": None,
}
return RateLimitDescriptor(key=descriptor["key"], value=descriptor["value"], rate_limit=rate_limit)
class ParallelSlotAcquisition(TypedDict):
slot_id: str
counter_keys: list[str]
counter_keys: Sequence[str]
class RateLimitStatus(TypedDict):
@ -706,6 +732,15 @@ def get_request_stash() -> RequestRateLimiterStash | None:
return _request_stash.get()
@contextmanager
def isolated_request_stash() -> Generator[None]:
token: Final = _request_stash.set(None)
try:
yield
finally:
_request_stash.reset(token)
def get_or_create_request_stash() -> RequestRateLimiterStash:
stash = _request_stash.get()
if stash is None:
@ -756,6 +791,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
parallel_acquire_script: _AsyncLuaScript | None
parallel_release_script: _AsyncLuaScript | None
parallel_count_script: _AsyncLuaScript | None
parallel_renew_script: _AsyncLuaScript | None
def __init__(
self,
@ -797,6 +833,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self.parallel_count_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
PARALLEL_COUNT_SCRIPT
)
self.parallel_renew_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
PARALLEL_RENEW_SCRIPT
)
else:
self.batch_rate_limiter_script = None
self.batch_counter_read_script = None
@ -806,6 +845,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self.parallel_acquire_script = None
self.parallel_release_script = None
self.parallel_count_script = None
self.parallel_renew_script = None
self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60))
@ -1168,7 +1208,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
async def in_memory_cache_sliding_window(
self,
keys: list[str],
keys: Sequence[str],
now_int: int,
window_size: int,
) -> CacheCounterValues:
@ -1357,7 +1397,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return (None,) * len(calls)
return tuple(batch.script(source, run, keys, args) for keys, args in calls)
def _group_keys_by_hash_tag(self, keys: list[str]) -> dict[str, list[str]]:
def _group_keys_by_hash_tag(self, keys: Sequence[str]) -> Mapping[str, Sequence[str]]:
"""
Group keys by their Redis hash tag to ensure cluster compatibility.
@ -1377,7 +1417,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
groups[slot_key].append(key)
else:
# For regular Redis, no grouping needed - process all keys together
groups[REDIS_NODE_HASHTAG_NAME] = keys
return MappingProxyType({REDIS_NODE_HASHTAG_NAME: keys})
return groups
@ -1753,12 +1793,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
gauge_keys: Final = [gauge["counter_key"] for gauge in gauges]
if self._is_redis_cluster() and self.parallel_acquire_script is not None:
return await self._check_cluster_parallel_gauges(gauges, slot_id, parent_otel_span, read_only)
if read_only:
if self.parallel_count_script is not None:
try:
raw_counts: Final[list[CacheCounterValue]] = await self.parallel_count_script(
keys=gauge_keys,
args=[PARALLEL_REQUEST_SLOT_TTL_SECONDS for _ in gauges],
args=tuple(PARALLEL_REQUEST_SLOT_TTL_SECONDS for _ in gauges),
)
counts = [max(0, int(value)) for value in raw_counts]
except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the local mirror unless fail-closed rejects
@ -1825,6 +1868,129 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
async with self._check_and_increment_lock:
return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span)
async def _check_cluster_parallel_gauges(
self,
gauges: Sequence[ParallelRequestGauge],
slot_id: str,
parent_otel_span: Span | None,
read_only: bool,
) -> RateLimitResponse:
by_key: Final = MappingProxyType(
{
gauge["counter_key"]: min(
(candidate for candidate in gauges if candidate["counter_key"] == gauge["counter_key"]),
key=lambda candidate: candidate["limit"],
)
for gauge in gauges
}
)
groups: Final = self._group_keys_by_hash_tag(tuple(by_key))
counts: Final[dict[str, int]] = {} # mutable-ok: gather independent Redis-slot results
attempted: Final[list[str]] = [] # mutable-ok: rollback includes requests whose responses were lost
try:
for keys in groups.values():
if read_only:
if self.parallel_count_script is None:
raise RuntimeError("Redis cluster parallel count script is unavailable")
counts.update(
(key, max(0, int(count)))
for key, count in zip(
keys,
await self.parallel_count_script(
keys=keys, args=tuple(PARALLEL_REQUEST_SLOT_TTL_SECONDS for _ in keys)
),
strict=True,
)
)
continue
if self.parallel_acquire_script is None:
raise RuntimeError("Redis cluster parallel acquire script is unavailable")
attempted.extend(keys)
acquire_args: list[object] = [] # mutable-ok: Redis EVAL args are flattened per slot below
for key in keys:
acquire_args.extend((by_key[key]["limit"], PARALLEL_REQUEST_SLOT_TTL_SECONDS, slot_id))
(raw,) = (await self.parallel_acquire_script(keys=keys, args=tuple(acquire_args)),)
if int(raw[0]) == 1:
await self._rollback_cluster_parallel_slots(tuple(attempted), slot_id, parent_otel_span)
return RateLimitResponse(
overall_code="OVER_LIMIT",
statuses=[ # mutable-ok: response contract requires a list
self._gauge_status(by_key[keys[int(raw[1]) - 1]], int(raw[2]), "OVER_LIMIT")
],
)
counts.update((key, int(count)) for key, count in zip(keys, raw[1:], strict=True))
for key in keys:
await self.internal_usage_cache.async_set_cache(
key=key,
value=counts[key],
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
except BaseException:
if attempted:
await self._rollback_cluster_parallel_slots(tuple(attempted), slot_id, parent_otel_span)
raise
statuses: Final = tuple(
self._gauge_status(
gauge,
counts[gauge["counter_key"]],
"OVER_LIMIT" if read_only and counts[gauge["counter_key"]] >= gauge["limit"] else "OK",
)
for gauge in gauges
)
return RateLimitResponse(
overall_code="OVER_LIMIT" if any(item["code"] == "OVER_LIMIT" for item in statuses) else "OK",
statuses=list(statuses), # mutable-ok: RateLimitResponse contract requires a list
)
async def _rollback_cluster_parallel_slots(
self, counter_keys: tuple[str, ...], slot_id: str, parent_otel_span: Span | None
) -> None:
rollback: Final = asyncio.create_task(
self._release_cluster_parallel_slots(counter_keys, slot_id, parent_otel_span)
)
cancelled = False # rebind-ok: defer repeated caller cancellation until compensation finishes
while not rollback.done():
try:
await asyncio.shield(rollback)
except asyncio.CancelledError:
cancelled = True
except Exception: # noqa: BLE001 # retrieve and report the completed task's exception below
break
try:
rollback.result()
except Exception: # noqa: BLE001 # preserve admission failure; unreachable Redis slots expire by TTL
verbose_proxy_logger.error("Could not roll back all Redis cluster parallel request slots")
if cancelled:
raise asyncio.CancelledError
async def _release_cluster_parallel_slots(
self, counter_keys: tuple[str, ...], slot_id: str, parent_otel_span: Span | None
) -> None:
first_error: Exception | None = None # rebind-ok: finish every shard before reporting the first failure
for keys in self._group_keys_by_hash_tag(counter_keys).values():
try:
if self.parallel_release_script is None:
raise RuntimeError("Redis cluster parallel release script is unavailable")
for key, count in zip(
keys,
await self.parallel_release_script(keys=keys, args=tuple(slot_id for _ in keys)),
strict=True,
):
await self.internal_usage_cache.async_set_cache(
key=key,
value=max(0, int(count)),
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
except Exception as exc: # noqa: BLE001 # one unreachable shard must not strand the other shards
if first_error is None:
first_error = exc
if first_error is not None:
raise first_error
async def _read_local_gauge_counts(
self,
gauge_keys: list[str],
@ -1894,6 +2060,68 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
statuses.append(self._gauge_status(gauge, in_flight + 1, "OK"))
return RateLimitResponse(overall_code="OK", statuses=statuses)
def transfer_realtime_call_slot(self, request_data: Mapping[str, object]) -> RealtimeCallLease | None:
call_id: Final = request_data.get("litellm_call_id")
if not isinstance(call_id, str):
return None
stash: Final = get_request_stash_for_call(call_id)
if stash is None or stash.parallel_slot is None:
return None
slot_id: Final = stash.parallel_slot["slot_id"]
counter_keys: Final = tuple(stash.parallel_slot["counter_keys"])
stash.parallel_slot = None
async def renew() -> bool:
return await self._renew_realtime_call_slot(slot_id, counter_keys)
async def release() -> None:
await self._release_parallel_request_slots(
ParallelSlotAcquisition(slot_id=slot_id, counter_keys=counter_keys)
)
return RealtimeCallLease(renew=renew, release=release)
async def _renew_realtime_call_slot(self, slot_id: str, counter_keys: tuple[str, ...]) -> bool:
if self.parallel_renew_script is not None:
try:
for keys in self._group_keys_by_hash_tag(counter_keys).values():
if tuple(
await self.parallel_renew_script(keys=keys, args=(slot_id, PARALLEL_REQUEST_SLOT_TTL_SECONDS))
) != (1,):
return False
return True
except Exception: # noqa: BLE001 # Redis ownership cannot be established by a local count mirror
return False
async with self._check_and_increment_lock:
now: Final = self._get_current_time().timestamp()
cutoff: Final = now - PARALLEL_REQUEST_SLOT_TTL_SECONDS
values: Final[tuple[ParallelGaugeCacheValue | None, ...]] = tuple(
[
await self.internal_usage_cache.async_get_cache(
key=counter_key, local_only=True, litellm_parent_otel_span=None
)
for counter_key in counter_keys
]
)
if any(
not isinstance(value, dict)
or not isinstance(score := value.get(slot_id), (int, float))
or score <= cutoff
for value in values
):
return False
for counter_key, value in zip(counter_keys, values):
if not isinstance(value, dict):
return False
await self.internal_usage_cache.async_set_cache(
key=counter_key,
value={**value, slot_id: now}, # mutable-ok: slot registry readers require a concrete dict
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
local_only=True,
litellm_parent_otel_span=None,
)
return True
async def _release_stashed_parallel_slot(
self,
stash: RequestRateLimiterStash | None,
@ -1931,11 +2159,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
slot_id: Final = acquisition["slot_id"]
if not counter_keys or not slot_id:
return
if self._is_redis_cluster() and self.parallel_release_script is not None:
await self._release_cluster_parallel_slots(tuple(counter_keys), slot_id, parent_otel_span)
return
if self.parallel_release_script is not None:
try:
raw: Final[list[CacheCounterValue]] = await self.parallel_release_script(
keys=counter_keys,
args=[slot_id for _ in counter_keys],
args=tuple(slot_id for _ in counter_keys),
)
await self._mirror_released_parallel_slots(counter_keys, raw, parent_otel_span)
return
@ -4020,6 +4251,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if d["key"] == "tag" and d["rate_limit"] is not None and d["rate_limit"].get("tokens_per_unit") is not None
)
effective_descriptors: Final = (
tuple(_without_parallel_limit(descriptor) for descriptor in descriptors)
if call_type == "_arealtime"
and is_realtime_call_attachment(TypeAdapter[object](object).validate_python(data.get("websocket")))
else descriptors
)
# Only check rate limits if we have descriptors with actual limits
if descriptors:
# First pass: RPM and max_parallel_requests sliding-window check.
@ -4038,16 +4276,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# double-charge every request.
parallel_counter_keys: Final = [
self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests")
for d in descriptors
for d in effective_descriptors
if (d.get("rate_limit") or {}).get("max_parallel_requests") is not None
]
parallel_slot_id: Final = uuid.uuid4().hex if parallel_counter_keys else None
first_pass_descriptors: Final = (
descriptors
effective_descriptors
if self.tpm_reservation_enabled
else tuple(
d for d in descriptors if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY)
d
for d in effective_descriptors
if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY)
)
)
response: Final = await self.should_rate_limit(
@ -5242,6 +5482,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
await self._release_stashed_parallel_slot(get_request_stash(), None)
async def async_release_realtime_attachment(
self, request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth
) -> None:
await self.async_post_call_failure_hook(
request_data={}, # mutable-ok: existing failure hook requires dict; attachment has no billable usage
original_exception=Exception("Realtime attachment completed"),
user_api_key_dict=user_api_key_dict,
)
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
"""
Release completed-request slots and update rate limit headers in the response.

View file

@ -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()

View file

@ -26,6 +26,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
live_model_access_group_limits_cache_key,
model_access_group_cache_key,
model_access_group_registry_cache_key,
)
@ -200,7 +201,14 @@ async def _evict_model_access_group_cache_keys(access_group: str, auth_cache: Us
)
await evict_and_broadcast(
cache_keys=(model_access_group_cache_key(access_group), model_access_group_registry_cache_key()),
cache_keys=(
model_access_group_cache_key(access_group),
# The Live delegation gate caches the same group's full limit row next to the flattened
# entry because it needs the rpm and tpm columns; leaving that entry behind keeps the
# old limit deciding managed delegation until its TTL expires.
live_model_access_group_limits_cache_key(access_group),
model_access_group_registry_cache_key(),
),
user_api_key_cache=auth_cache,
)

View file

@ -1613,6 +1613,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
except Exception as e:
verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e)
from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS
await CALL_SUPERVISORS.shutdown()
await _drain_spend_event_producer_on_shutdown()
# Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect
@ -12899,6 +12903,28 @@ async def _reject_realtime_session(
await _release_realtime_max_parallel_slot(user_api_key_dict)
_CODEX_LIVE_AUTH_DEPENDENCY: Final = Depends(user_api_key_auth_websocket)
reserve_lazy_slot(app, "live")
reserve_lazy_slot(app, "realtime")
@app.websocket("/openai/v1/live/{call_id}")
@app.websocket("/live/{call_id}")
@app.websocket("/v1/live/{call_id}")
async def codex_live_sideband_endpoint(
websocket: WebSocket,
call_id: str,
user_api_key_dict: UserAPIKeyAuth = _CODEX_LIVE_AUTH_DEPENDENCY,
) -> None:
from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband
await codex_realtime_sideband(websocket, call_id, user_api_key_dict)
@app.websocket("/v1/live")
@app.websocket("/live")
@app.websocket("/openai/v1/live")
@app.websocket("/openai/v1/realtime")
@app.websocket("/v1/realtime")
@app.websocket("/realtime")
@ -12906,12 +12932,18 @@ async def realtime_websocket_endpoint(
websocket: WebSocket,
model: str | None = fastapi.Query(None, description="The model to use for the websocket connection."),
intent: str | None = fastapi.Query(None, description="The intent of the websocket connection."),
call_id: str | None = None,
guardrails: str | None = fastapi.Query(
None,
description="Comma-separated list of guardrail names to apply to this request.",
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket),
):
if call_id is not None:
from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband
await codex_realtime_sideband(websocket, call_id, user_api_key_dict)
return
requested_protocols: Final = [
p.strip() for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",") if p.strip()
]
@ -12943,7 +12975,12 @@ async def realtime_websocket_endpoint(
await websocket.accept(**accept_kwargs)
# Only use explicit parameters, not all query params
query_params: Final = cast(RealtimeQueryParams, dict(_realtime_query_params_template(model, intent)))
query_params: Final = cast(
RealtimeQueryParams,
dict( # mutable-ok: FastAPI request query params must be materialized as a dict
_realtime_query_params_template(model, intent) + ((("call_id", call_id),) if call_id is not None else ())
),
)
data: dict[str, object] = {
"model": route_model,

View file

@ -0,0 +1,562 @@
import asyncio
import base64
import hashlib
import json
import time
from collections.abc import Awaitable, Callable, Mapping
from contextlib import AsyncExitStack, nullcontext
from contextvars import Token
from types import MappingProxyType
from typing import Final, Literal
import httpx
from fastapi import HTTPException, Request, Response, WebSocket
from pydantic import TypeAdapter
from starlette.formparsers import MultiPartException, MultiPartParser
from starlette.types import Message, Scope
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.realtime_streaming import (
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
RealTimeStreaming,
realtime_attachment_cleanup,
)
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.chatgpt.codex import (
CodexRealtimeCall,
CodexRealtimeOffer,
build_call_request,
build_sideband_request,
parse_call_response,
)
from litellm.llms.chatgpt.realtime import (
CallAccounting,
ChatGPTRealtime,
configured_realtime_headers,
realtime_endpoint,
)
from litellm.proxy._types import InternalRequestOrigin, ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
from litellm.proxy.auth.user_api_key_auth import (
get_api_key,
get_api_key_from_custom_header,
get_websocket_api_key,
user_api_key_auth,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
from litellm.proxy.common_utils.http_parsing_utils import (
_normalize_media_type, # pyright: ignore[reportPrivateUsage] # reuse the shared HTTP media-type normalization contract
)
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # existing built-in limiter has no public alias
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # existing built-in limiter has no public alias
isolated_request_stash,
)
from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease, realtime_call_attachment
from litellm.proxy.spend_tracking.budget_reservation import (
invalidate_budget_reservation_counters,
release_or_invalidate_budget_reservation,
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.router import GenericLiteLLMParams
async def supervise_codex_call(
request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth, lease: RealtimeCallLease | None = None
) -> None:
with isolated_request_stash():
await _start_codex_supervisor(request, call, auth, lease)
async def _start_codex_supervisor(
request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth, lease: RealtimeCallLease | None
) -> None:
import litellm
from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, CallSupervisor
async def receive() -> Message:
body: Final[RealtimeQueryParams] = {"model": call.alias}
message: Final[Message] = {
"type": "http.request",
"body": json.dumps(body).encode(),
"more_body": False,
}
return message
async def send(_message: Message) -> None:
return None
supervision_owned = False # rebind-ok: supervisor owns cleanup after construction
effective_handler: ChatGPTRealtime | None = None # rebind-ok: reuse hook-enriched credentials for cleanup
sockets: Final = AsyncExitStack()
try:
observer_scope: Final[Scope] = {**request.scope}
observer_request: Final = Request(observer_scope, receive=receive)
processed, logger = await process_codex_request(
observer_request,
{ # mutable-ok: common request processing enriches metadata
**build_sideband_request(call),
"model": call.alias,
},
auth,
call.alias,
"_arealtime",
internal_realtime_observer=True,
)
pinned: Final = { # mutable-ok: logging and provider parameter contract
**processed,
**build_sideband_request(call),
"extra_headers": MappingProxyType(
{
**configured_realtime_headers(
TypeAdapter[Mapping[str, object] | None](Mapping[str, object] | None).validate_python(
processed.get("extra_headers")
)
),
**configured_realtime_headers(call.extra_headers),
}
),
"litellm_metadata": { # mutable-ok: Logging.update_from_kwargs requires a dict to retain ownership metadata
**TypeAdapter(Mapping[str, object]).validate_python(
processed.get("litellm_metadata") or MappingProxyType({})
),
**(
MappingProxyType(
{
"model_info": { # mutable-ok: logging and cost callbacks require a concrete model-info dict
**litellm.get_model_info(model=call.model_id),
"id": call.model_id,
}
}
)
if call.model_id is not None
else MappingProxyType({})
),
},
}
logger.update_from_kwargs(
kwargs=pinned,
model=call.model,
user=None,
optional_params={}, # mutable-ok: logging contract
litellm_params={ # mutable-ok: Logging.update_from_kwargs pops metadata from its argument
**logger.litellm_params,
"litellm_metadata": pinned["litellm_metadata"],
"arealtime": True,
},
custom_llm_provider="chatgpt",
)
params: Final = GenericLiteLLMParams.model_validate(pinned)
handler: Final = ChatGPTRealtime(
params, request.headers, TypeAdapter(Mapping[str, object]).validate_python(pinned["extra_headers"])
)
effective_handler = handler
api_base: Final = ChatGPTRealtime.get_api_base(call.api_base)
connection: Final = await handler.open_call_connection(call.model, api_base)
sockets.push_async_callback(connection.close)
async def close_call() -> None:
await handler.close_call(connection, call.model, api_base)
async def force_close_call() -> None:
await handler.hangup_call(api_base)
frontend_scope: Final[Scope] = {**request.scope, "type": "websocket"}
frontend: Final = WebSocket(frontend_scope, receive=receive, send=send)
stream: Final = RealTimeStreaming(frontend, connection, logger, model=call.model, user_api_key_dict=auth)
supervisor: Final = CallSupervisor(
connection,
stream,
logger,
auth,
close_call,
force_close_call=force_close_call,
terminal_usage_required=realtime_endpoint(call.model) == "live",
lease=lease,
)
supervision_owned = True
sockets.pop_all()
await CALL_SUPERVISORS.start(supervisor)
except BaseException:
if not supervision_owned:
try:
fallback_handler: Final = effective_handler or ChatGPTRealtime(
GenericLiteLLMParams.model_validate(build_sideband_request(call)),
request.headers,
call.extra_headers,
)
await fallback_handler.hangup_call(ChatGPTRealtime.get_api_base(call.api_base))
except Exception: # noqa: BLE001 # preserve original failure without logging provider credentials
verbose_proxy_logger.error("Realtime startup cleanup could not confirm upstream termination")
try:
await invalidate_budget_reservation_counters(budget_reservation=auth.budget_reservation)
except Exception: # noqa: BLE001 # cleanup errors must not replace the original startup failure
verbose_proxy_logger.error("Realtime startup cleanup could not invalidate budget counters")
else:
await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation)
finally:
try:
await sockets.aclose()
except Exception: # noqa: BLE001 # socket cleanup must preserve the original startup failure
verbose_proxy_logger.error("Realtime startup cleanup could not close observer socket")
raise
def encode_call(call: CodexRealtimeCall) -> str:
encrypted: Final = encrypt_value_helper(call.model_dump_json())
return "rtc_litellm_" + base64.urlsafe_b64encode(encrypted.encode()).decode().rstrip("=")
def decode_call(token: str, authorization: str) -> CodexRealtimeCall:
try:
if not token.startswith("rtc_litellm_"):
raise ValueError("Invalid call prefix")
encoded: Final = token.removeprefix("rtc_litellm_")
encrypted: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True)
plaintext: Final = decrypt_value_helper(encrypted.decode(), key="codex_realtime_call")
call: Final = CodexRealtimeCall.model_validate_json(plaintext or "")
except (ValueError, TypeError, UnicodeError) as exc:
raise HTTPException(403, "Invalid realtime call") from exc
if (
call.expires_at < time.time()
or call.owner != hashlib.sha256(authorization.encode(), usedforsecurity=False).hexdigest()
):
raise HTTPException(403, "Invalid or expired realtime call")
return call
MAX_REALTIME_OFFER_BYTES: Final = 8 * 1024 * 1024
async def _cache_bounded_offer_body(request: Request) -> None:
content_length: int | None
try:
content_length = int(request.headers.get("content-length", ""))
except ValueError:
# A missing or non-numeric content length is checked while streaming below.
content_length = None
if content_length is not None and content_length > MAX_REALTIME_OFFER_BYTES:
raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit")
if hasattr(request, "_body"):
if len(request._body) > MAX_REALTIME_OFFER_BYTES: # pyright: ignore[reportPrivateUsage] # validate Starlette's cached body without consuming it again
raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit")
return
if request._form is not None and request._stream_consumed: # pyright: ignore[reportPrivateUsage] # a mixed-case empty form cache may leave the stream unread
return
body: Final = bytearray()
async for chunk in request.stream():
if len(body) + len(chunk) > MAX_REALTIME_OFFER_BYTES:
raise HTTPException(413, "Realtime offer exceeds the 8 MiB limit")
body.extend(chunk)
request._body = bytes(body) # pyright: ignore[reportPrivateUsage] # Starlette has no public setter for its shared body cache
async def read_codex_offer(request: Request) -> CodexRealtimeOffer:
await _cache_bounded_offer_body(request)
content_type: Final = request.headers.get("content-type", "")
if _normalize_media_type(content_type) == "multipart/form-data":
if content_type.split(";", 1)[0] != "multipart/form-data" and not await request.form():
try:
request._form = await MultiPartParser(request.headers, request.stream()).parse() # pyright: ignore[reportPrivateUsage] # Starlette exposes no setter for its shared form cache; # rebind-ok: Request.close must own and close uploaded files
request.scope.pop("parsed_body", None)
except MultiPartException as exc:
raise HTTPException(400, "Invalid realtime multipart offer") from exc
form: Final = await request.form()
return CodexRealtimeOffer.model_validate(
MappingProxyType({"sdp": form.get("sdp"), "session": json.loads(str(form.get("session", "{}")))})
)
return CodexRealtimeOffer.model_validate(await request.json())
async def process_codex_request(
request: Request,
data: dict[str, object], # mutable-ok: common request processor enriches this dictionary
auth: UserAPIKeyAuth,
model: str,
route_type: Literal["arealtime_calls", "_arealtime"],
*,
internal_realtime_observer: bool = False,
) -> tuple[dict[str, object], Logging]: # mutable-ok: common request processor returns enriched routing arguments
from litellm.proxy import proxy_server as server
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
processed, logging_obj = await processor.common_processing_pre_call_logic(
request=request,
general_settings=server.general_settings,
user_api_key_dict=auth,
version=server.version,
proxy_logging_obj=server.proxy_logging_obj,
proxy_config=server.proxy_config,
llm_router=server.llm_router,
user_model=TypeAdapter[str | None](str | None).validate_python(server.user_model),
user_temperature=TypeAdapter[float | None](float | None).validate_python(server.user_temperature),
user_request_timeout=server.user_request_timeout,
user_max_tokens=server.user_max_tokens,
user_api_base=TypeAdapter[str | None](str | None).validate_python(server.user_api_base),
model=model,
route_type=route_type,
**(
MappingProxyType({"internal_realtime_observer": True})
if internal_realtime_observer
else MappingProxyType({})
),
)
if internal_realtime_observer:
logging_obj.model_call_details["internal_request_origin"] = InternalRequestOrigin.REALTIME_OBSERVER
return processed, logging_obj
async def create_codex_realtime_call(request: Request) -> Response:
try:
with isolated_request_stash():
return await _create_codex_realtime_call(request)
finally:
await request.close()
async def _create_codex_realtime_call(request: Request) -> Response:
from litellm.proxy import proxy_server as server
try:
offer: Final = await read_codex_offer(request)
except ValueError as exc:
raise HTTPException(400, "Invalid realtime offer: expected sdp and session") from exc
model: Final = offer.session.model
if not model:
raise HTTPException(400, "session.model is required")
auth: Final = await user_api_key_auth(
request=request,
api_key=request.headers.get("authorization", ""),
azure_api_key_header=request.headers.get("api-key", ""),
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
custom_litellm_key_header=request.headers.get("x-litellm-api-key"),
)
selected_key, _ = get_api_key(
request=request,
api_key=request.headers.get("authorization", ""),
azure_api_key_header=request.headers.get("api-key", ""),
custom_litellm_key_header=request.headers.get("x-litellm-api-key"),
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
pass_through_endpoints=None,
route="/v1/realtime/calls",
)
custom_header: Final = server.general_settings.get("litellm_key_header_name")
owner_key: Final = (
get_api_key_from_custom_header(request, custom_header) if isinstance(custom_header, str) else selected_key
)
supervision_started = False # rebind-ok: transfer reservation ownership only after supervision is established
call_lease: RealtimeCallLease | None = None
lease_transferred = False # rebind-ok: failed startup leaves the signaling task responsible for its lease
preprocessing_started = False # rebind-ok: only refund reservations belonging to this signaling request
limiter: Final = server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter")
try:
await can_key_call_resolved_model(
model=model,
llm_model_list=TypeAdapter[tuple[object, ...] | None](tuple[object, ...] | None).validate_python(
server.llm_model_list
),
valid_token=auth,
llm_router=server.llm_router,
)
live_signaling: Final = request.url.path.rstrip("/") in ("/live", "/v1/live", "/openai/v1/live")
query: Final = (
MappingProxyType({"intent": "quicksilver", "architecture": "avas", **request.query_params})
if live_signaling
else request.query_params
)
data: Final = build_call_request(offer, query, request.headers)
signaling_auth: Final = auth.model_copy(update=MappingProxyType({"budget_reservation": None}))
if isinstance(limiter, _PROXY_MaxParallelRequestsHandler) and (
auth.max_parallel_requests is not None
or server.general_settings.get("global_max_parallel_requests") is not None
):
raise HTTPException(400, "Realtime calls with parallel limits require the V3 rate limiter")
preprocessing_started = True
processed, _ = await process_codex_request(request, data, signaling_auth, model, "arealtime_calls")
if isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
call_lease = limiter.transfer_realtime_call_slot(processed)
if call_lease is not None:
call_lease.start()
if not await call_lease.renew():
raise HTTPException(503, "Realtime call quota reservation was lost")
with isolated_request_stash():
result: Final = await server.route_request(
data=processed,
route_type="arealtime_calls",
llm_router=server.llm_router,
user_model=TypeAdapter[str | None](str | None).validate_python(server.user_model),
)
try:
response: Final = await result
except BaseLLMException as exc:
raise HTTPException(exc.status_code, str(exc)) from exc
if not isinstance(response, httpx.Response):
raise HTTPException(502, "Invalid realtime signaling response")
if response.is_error:
return Response(response.content, status_code=response.status_code, media_type="application/json")
try:
call: Final = parse_call_response(
response,
alias=model,
owner=hashlib.sha256(f"Bearer {owner_key}".encode(), usedforsecurity=False).hexdigest(),
expires_at=time.time() + 3600,
)
except ValueError as exc:
raise HTTPException(400, str(exc)) from exc
supervised_call: Final = call.model_copy(
update=MappingProxyType({"usage_supervised": True, "parallel_reserved": call_lease is not None})
)
token: Final = encode_call(supervised_call)
supervision_started = True
if call_lease is None:
await supervise_codex_call(request, supervised_call, auth)
else:
await supervise_codex_call(request, supervised_call, auth, call_lease)
lease_transferred = True
return Response(
response.content,
status_code=response.status_code,
media_type="application/sdp",
headers=MappingProxyType(
{"Location": f"/v1/live/{token}" if live_signaling else f"/v1/realtime/calls/{token}"}
),
)
finally:
try:
if call_lease is not None and not lease_transferred:
await call_lease.close()
finally:
try:
if preprocessing_started and isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
await asyncio.shield(
limiter.async_post_call_failure_hook(
request_data={}, # mutable-ok: existing failure-hook contract
original_exception=Exception("Realtime signaling completed without token usage"),
user_api_key_dict=auth,
)
)
finally:
if not supervision_started:
await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation)
async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAPIKeyAuth) -> None:
import litellm
from litellm.proxy import proxy_server as server
protocols: Final = tuple(
p.strip() for p in websocket.headers.get("sec-websocket-protocol", "").split(",") if p.strip()
)
logging_obj: Logging | None = None # rebind-ok: cleanup needs the logger only after pre-call succeeds
attachment_limiter: _PROXY_MaxParallelRequestsHandler | _PROXY_MaxParallelRequestsHandler_v3 | None = None
cleanup_token: Token[Callable[[], Awaitable[None]] | None] | None = None
try:
try:
api_key: Final = get_websocket_api_key(websocket)
if not api_key:
raise HTTPException(403, "No API key provided")
call: Final = decode_call(token, f"Bearer {api_key}")
await can_key_call_resolved_model(
model=call.alias,
llm_model_list=TypeAdapter[tuple[object, ...] | None](tuple[object, ...] | None).validate_python(
server.llm_model_list
),
valid_token=auth,
llm_router=server.llm_router,
)
except (HTTPException, ProxyException):
await websocket.close(code=1008, reason="Invalid realtime call")
return
async def receive() -> Message:
return { # mutable-ok: ASGI receive message
"type": "http.request",
"body": json.dumps({"model": call.alias}).encode(), # mutable-ok: JSON request serialization
"more_body": False,
}
request: Final = Request(
{ # mutable-ok: Starlette stores request state in the ASGI scope
**websocket.scope,
"type": "http",
"method": "POST",
"path": websocket.scope.get("path", "/v1/realtime"),
},
receive=receive,
)
data: Final = { # mutable-ok: common request processor enriches routing arguments
**build_sideband_request(call),
"model": call.alias,
"websocket": websocket,
"guardrails": [ # mutable-ok: guardrail processing expects a list
name.strip() for name in websocket.query_params.get("guardrails", "").split(",") if name.strip()
],
}
limiter: Final = server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter")
if call.usage_supervised and isinstance(
limiter, (_PROXY_MaxParallelRequestsHandler, _PROXY_MaxParallelRequestsHandler_v3)
):
attachment_limiter = limiter
if isinstance(limiter, _PROXY_MaxParallelRequestsHandler):
limiter.begin_realtime_attachment(data)
try:
with realtime_call_attachment(websocket) if call.parallel_reserved else nullcontext():
processed, logging_obj = await process_codex_request(request, data, auth, call.alias, "_arealtime")
except Exception: # noqa: BLE001 # custom hook exceptions must reject the connection
verbose_proxy_logger.exception("Realtime sideband pre-call rejected")
await websocket.close(code=1008, reason="Realtime pre-call rejected")
return
await websocket.accept(
subprotocol=next((p for p in protocols if not p.startswith("openai-insecure-api-key.")), None)
)
if attachment_limiter is not None:
selected_limiter: Final = attachment_limiter
async def release_attachment() -> None:
await selected_limiter.async_release_realtime_attachment(data, auth)
cleanup_token = realtime_attachment_cleanup.set(release_attachment)
await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call
model=f"chatgpt/{call.model}",
websocket=websocket,
**MappingProxyType(
{
key: value
for key, value in { # mutable-ok: retain processed metadata with pinned routing
**processed,
**build_sideband_request(call),
"extra_headers": MappingProxyType(
{
**configured_realtime_headers(
TypeAdapter[Mapping[str, object] | None](
Mapping[str, object] | None
).validate_python(processed.get("extra_headers"))
),
**configured_realtime_headers(call.extra_headers),
}
),
"websocket": websocket,
"user_api_key_dict": auth,
"chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None,
}.items()
if key not in ("model", "websocket")
}
),
)
finally:
try:
if attachment_limiter is not None:
await attachment_limiter.async_release_realtime_attachment(data, auth)
finally:
if cleanup_token is not None:
realtime_attachment_cleanup.reset(cleanup_token)
if logging_obj is None or not logging_obj.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY):
await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation)

View file

@ -0,0 +1,270 @@
import asyncio
from collections.abc import AsyncIterator, Awaitable, Callable
from contextlib import suppress
from typing import Final, Protocol
from pydantic import BaseModel, ValidationError
from websockets.exceptions import ConnectionClosedOK
from litellm._logging import verbose_proxy_logger
from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease
from litellm.proxy.spend_tracking.budget_reservation import (
invalidate_budget_reservation_counters,
release_or_invalidate_budget_reservation,
)
from litellm.types.realtime import LiveSessionUsageEvent
class ObserverSocket(Protocol):
def __aiter__(self) -> AsyncIterator[str | bytes]: ...
async def close(self) -> None: ...
class UsageSink(Protocol):
def store_message(self, message: str) -> None: ...
async def log_messages(self, *, wait_for_dispatch: bool = False) -> None: ...
class _ObserverEvent(BaseModel):
type: str
class CallSupervisor:
def __init__(
self,
upstream: ObserverSocket,
stream: UsageSink,
logging_obj: Logging,
auth: UserAPIKeyAuth,
close_call: Callable[[], Awaitable[None]],
*,
ready_timeout: float = 20,
lifetime: float = 3600,
drain_timeout: float = 5,
termination_timeout: float = 60,
logging_timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE,
terminal_usage_required: bool = True,
connected_ready: bool = False,
force_close_call: Callable[[], Awaitable[None]] | None = None,
lease: RealtimeCallLease | None = None,
) -> None:
self._upstream = upstream
self._stream = stream
self._logging = logging_obj
self._auth = auth
self._close_call = close_call
self._force_close_call = force_close_call
self._lease = lease
self._ready_timeout = ready_timeout
self._lifetime = lifetime
self._drain_timeout = drain_timeout
self._termination_timeout = termination_timeout
self._logging_timeout = logging_timeout
self._terminal_usage_required = terminal_usage_required
self._ready = asyncio.Event()
self._stop = asyncio.Event()
self._started = connected_ready
if connected_ready:
self._ready.set()
self._terminal = False
self._terminal_usage_valid = False
self._close_confirmed = False
self._accounting_complete = False
self._task: asyncio.Task[None] | None = None
async def start(self) -> None:
if self._task is not None:
raise RuntimeError("Call observer already started")
self._task = asyncio.create_task(self._run())
try:
await asyncio.wait_for(self._ready.wait(), timeout=self._ready_timeout)
if self._lease is not None and not await self._lease.renew():
raise RuntimeError("Call observer lost its quota reservation during startup")
if not self._started or self._terminal or self._task.done():
raise RuntimeError("Call observer ended before session became available")
except BaseException:
await self._close_after_failed_start()
raise
async def _close_after_failed_start(self) -> None:
cleanup: Final = asyncio.create_task(self.close())
while not cleanup.done():
with suppress(asyncio.CancelledError):
await asyncio.shield(cleanup)
cleanup.result()
async def close(self) -> None:
self._stop.set()
await self.wait()
async def wait(self) -> None:
if self._task is not None:
await asyncio.shield(self._task)
async def _read(self) -> None:
try:
await self._read_events()
except ConnectionClosedOK:
return
async def _read_events(self) -> None:
event: _ObserverEvent
async for message in self._upstream:
self._stream.store_message(message.decode("utf-8") if isinstance(message, bytes) else message)
event = _ObserverEvent.model_validate_json(message)
if event.type in ("session.started", "session.created"):
self._started = True
self._ready.set()
if event.type == "session.closed":
self._terminal = True
try:
LiveSessionUsageEvent.model_validate_json(message)
except ValidationError:
self._terminal_usage_valid = False
else:
self._terminal_usage_valid = True
return
def _usage_complete(self) -> bool:
if self._terminal_usage_required:
return self._terminal and self._terminal_usage_valid
return self._terminal or self._close_confirmed
async def _run(self) -> None:
try:
await self._observe()
finally:
try:
if self._lease is not None:
await self._lease.close()
finally:
self._ready.set()
async def _observe(self) -> None:
reader: Final = asyncio.create_task(self._read())
stopped: Final = asyncio.create_task(self._stop.wait())
lease_failed: Final = asyncio.create_task(self._lease.wait_failed()) if self._lease is not None else None
try:
await asyncio.wait(
(reader, stopped, lease_failed) if lease_failed is not None else (reader, stopped),
timeout=self._lifetime,
return_when=asyncio.FIRST_COMPLETED,
)
finally:
try:
if not self._terminal:
deadline: Final = asyncio.get_running_loop().time() + self._termination_timeout
primary_deadline: Final = (
deadline - self._termination_timeout / 2
if self._terminal_usage_required and self._force_close_call is not None
else deadline
)
try:
await asyncio.wait_for(
self._close_call(), timeout=max(0.0, primary_deadline - asyncio.get_running_loop().time())
)
self._close_confirmed = True
except Exception: # noqa: BLE001 # provider exceptions can contain credentials
verbose_proxy_logger.error("Realtime observer could not terminate upstream call")
await self._drain(reader, timeout=max(0.0, primary_deadline - asyncio.get_running_loop().time()))
if self._terminal_usage_required and not self._terminal and self._force_close_call is not None:
remaining: Final = max(0.0, deadline - asyncio.get_running_loop().time())
try:
await asyncio.wait_for(self._force_close_call(), timeout=remaining)
self._close_confirmed = True
except Exception: # noqa: BLE001 # provider exceptions can contain credentials
verbose_proxy_logger.error("Realtime observer independent hangup failed")
await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time()))
finally:
stopped.cancel()
reader.cancel()
if lease_failed is not None:
lease_failed.cancel()
# Branching rather than a conditional star-unpacked tuple: the overload solver cannot
# bind one result type across a tuple whose length depends on the branch.
if lease_failed is not None:
await asyncio.gather(reader, stopped, lease_failed, return_exceptions=True)
else:
await asyncio.gather(reader, stopped, return_exceptions=True)
with suppress(Exception):
await self._upstream.close()
if not self._usage_complete():
self._logging.model_call_details["realtime_usage_incomplete"] = True
verbose_proxy_logger.error(
"Realtime observer ended without terminal usage; recorded usage is partial"
)
try:
try:
await asyncio.wait_for(
self._stream.log_messages(wait_for_dispatch=True), timeout=self._logging_timeout
)
self._accounting_complete = not bool(
self._logging.model_call_details.get("realtime_backend_accounting_incomplete")
)
except asyncio.TimeoutError:
verbose_proxy_logger.error("Realtime observer timed out dispatching usage accounting")
finally:
if not self._accounting_complete:
self._logging.model_call_details["realtime_accounting_incomplete"] = True
if self._started and (not self._usage_complete() or not self._accounting_complete):
await invalidate_budget_reservation_counters(
budget_reservation=self._auth.budget_reservation
)
elif not self._logging.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY):
await release_or_invalidate_budget_reservation(
budget_reservation=self._auth.budget_reservation
)
finally:
self._ready.set()
async def _drain(self, reader: asyncio.Task[None], *, timeout: float | None = None) -> None:
try:
await asyncio.wait_for(
asyncio.shield(reader),
timeout=self._drain_timeout if timeout is None else min(self._drain_timeout, timeout),
)
except asyncio.TimeoutError:
if not self._usage_complete():
verbose_proxy_logger.error("Realtime observer timed out draining terminal usage")
except Exception: # noqa: BLE001 # cleanup must settle the socket even when reading or closing fails
verbose_proxy_logger.error("Realtime observer could not drain terminal usage")
return
class CallSupervisors:
def __init__(self) -> None:
self._tasks: tuple[asyncio.Task[None], ...] = ()
self._calls: tuple[CallSupervisor, ...] = ()
async def start(self, supervisor: CallSupervisor) -> None:
self._calls = (*self._calls, supervisor)
try:
await supervisor.start()
except BaseException:
self._calls = tuple(call for call in self._calls if call is not supervisor)
raise
task: Final = asyncio.create_task(self._watch(supervisor))
self._tasks = (*self._tasks, task)
async def _watch(self, supervisor: CallSupervisor) -> None:
try:
try:
await supervisor.wait()
except Exception: # noqa: BLE001 # task must be consumed without exposing provider exception payloads
verbose_proxy_logger.error("Realtime observer accounting failed")
finally:
self._calls = tuple(call for call in self._calls if call is not supervisor)
self._tasks = tuple(task for task in self._tasks if task is not asyncio.current_task())
async def shutdown(self) -> None:
await asyncio.gather(*(call.close() for call in self._calls), return_exceptions=True)
await asyncio.gather(*self._tasks, return_exceptions=True)
CALL_SUPERVISORS: Final = CallSupervisors()

View file

@ -16,7 +16,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import (
_normalize_media_type, # pyright: ignore[reportPrivateUsage] # reuse the shared HTTP media-type normalization contract
_read_request_body,
)
from litellm.proxy.common_utils.openai_error_payload import (
error_status_code,
openai_error_param,
@ -359,6 +362,15 @@ async def create_realtime_client_secret(
return RealtimeClientSecretResponse(**upstream_json)
@router.post("/v1/live", tags=["realtime"]) # mutable-ok: FastAPI route metadata uses a mutable tag list
@router.post("/live", tags=["realtime"]) # mutable-ok: FastAPI route metadata uses a mutable tag list
@router.post("/openai/v1/live", tags=["realtime"]) # mutable-ok: FastAPI route metadata uses a mutable tag list
async def proxy_live_calls(request: Request) -> Response:
from litellm.proxy.realtime_endpoints.call_sessions import create_codex_realtime_call
return await create_codex_realtime_call(request)
@router.post(
"/v1/realtime/calls",
tags=["realtime"],
@ -375,6 +387,11 @@ async def proxy_realtime_calls(
request: Request,
fastapi_response: Response,
) -> Response:
if _normalize_media_type(request.headers.get("content-type", "")) in ("application/json", "multipart/form-data"):
from litellm.proxy.realtime_endpoints.call_sessions import create_codex_realtime_call
return await create_codex_realtime_call(request)
from litellm.proxy.proxy_server import (
add_litellm_data_to_request,
general_settings,
@ -410,6 +427,12 @@ async def proxy_realtime_calls(
sdp_body: Final[bytes] = await request.body()
decoded_payload: Final = _decode_realtime_token_payload(decrypted_token_value)
if decoded_payload is None and decrypted_token_value.lstrip().startswith(("{", "[")):
return Response(
content='{"error":"Invalid or expired token"}',
status_code=http_status.HTTP_401_UNAUTHORIZED,
media_type="application/json",
)
if decoded_payload is not None:
# Check token expiry
expires_at: Final = decoded_payload.get("expires_at")

File diff suppressed because it is too large Load diff

View file

@ -2352,7 +2352,9 @@ class ProxyLogging:
data: None,
call_type: CallTypesLiteral,
guardrails_only: bool = False,
*,
skip_guardrails: bool = False,
internal_realtime_observer: bool = False,
) -> None:
pass
@ -2363,7 +2365,9 @@ class ProxyLogging:
data: dict,
call_type: CallTypesLiteral,
guardrails_only: bool = False,
*,
skip_guardrails: bool = False,
internal_realtime_observer: bool = False,
) -> dict:
pass
@ -2373,7 +2377,9 @@ class ProxyLogging:
data: dict | None,
call_type: CallTypesLiteral,
guardrails_only: bool = False,
*,
skip_guardrails: bool = False,
internal_realtime_observer: bool = False,
) -> dict | None:
"""
Allows users to modify/reject the incoming request to the proxy, without having to deal with parsing Request body.
@ -2481,6 +2487,10 @@ class ProxyLogging:
deferred_route_exc: SensitiveDataRouteException | None = None
for _callback in caps.resolved_callbacks:
if internal_realtime_observer and isinstance(
_callback, (_PROXY_MaxParallelRequestsHandler, _PROXY_MaxParallelRequestsHandler_v3)
):
continue
start_time = time.time()
try:
if isinstance(_callback, CustomGuardrail) and data is not None:
@ -4611,7 +4621,7 @@ class PrismaClient:
def hash_token(self, token: str):
# Hash the string using SHA-256
hashed_token: Final = hashlib.sha256(token.encode()).hexdigest()
hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest()
return hashed_token
@ -7115,7 +7125,7 @@ def hash_token(token: str):
import hashlib
# Hash the string using SHA-256
hashed_token: Final = hashlib.sha256(token.encode()).hexdigest()
hashed_token: Final = hashlib.sha256(token.encode(), usedforsecurity=False).hexdigest()
return hashed_token

View file

@ -6,6 +6,8 @@ from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from pydantic import TypeAdapter
import litellm
from litellm.constants import (
AZURE_OPENAI_AUDIO_PROVIDERS,
@ -89,6 +91,7 @@ def _get_realtime_http_provider_config(
dynamic_api_base: str | None,
dynamic_api_key: str | None,
litellm_params: GenericLiteLLMParams,
is_call: bool = False,
) -> tuple["BaseRealtimeHTTPConfig | None", str, str]:
"""
Return (provider_config, resolved_api_base, resolved_api_key) for the
@ -106,13 +109,15 @@ def _get_realtime_http_provider_config(
provider_config = ProviderConfigManager.get_provider_realtime_http_config(
model="",
provider=LlmProviders(custom_llm_provider),
params=litellm_params,
is_call=is_call,
)
raw_api_base: Final = dynamic_api_base or litellm_params.api_base
raw_api_key: Final = dynamic_api_key or litellm_params.api_key
if provider_config is not None:
resolved_api_base = provider_config.get_api_base(api_base=raw_api_base)
resolved_api_base = provider_config.resolve_api_base(litellm_params.api_base, dynamic_api_base)
resolved_api_key = provider_config.get_api_key(api_key=raw_api_key)
else:
# Fallback for providers without a dedicated HTTP config (treated as OpenAI-compatible).
@ -285,9 +290,18 @@ async def arealtime_calls(
dynamic_api_base=dynamic_api_base,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
is_call=True,
)
if session is not None:
session = _with_resolved_session_model(session, model_name)
supplied_headers: Final = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(
kwargs.get("extra_headers")
)
call_headers: Final = (
provider_config.get_realtime_calls_extra_headers(supplied_headers)
if provider_config is not None
else supplied_headers
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
@ -295,7 +309,7 @@ async def arealtime_calls(
litellm_params={"api_base": resolved_api_base},
custom_llm_provider=custom_llm_provider,
)
return await base_llm_http_handler.async_realtime_calls_handler(
response: Final = await base_llm_http_handler.async_realtime_calls_handler(
api_base=resolved_api_base,
openai_ephemeral_key=openai_ephemeral_key,
sdp_body=sdp_body,
@ -304,10 +318,17 @@ async def arealtime_calls(
provider_config=provider_config,
model=model_name,
session_config=session,
extra_headers=kwargs.get("extra_headers"),
extra_headers=call_headers,
client=kwargs.get("client"),
api_version=litellm_params.api_version,
)
return (
provider_config.transform_realtime_calls_response(
response, model_name, litellm_logging_obj.get_router_model_id(), call_headers
)
if provider_config is not None
else response
)
async def vertex_access_token_resolver(
@ -363,8 +384,10 @@ async def _arealtime(
For PROXY use only.
"""
headers = cast(dict | None, kwargs.get("headers"))
extra_headers: Final = cast(dict | None, kwargs.get("extra_headers"))
headers = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(kwargs.get("headers"))
extra_headers: Final = TypeAdapter[dict[str, object] | None](dict[str, object] | None).validate_python(
kwargs.get("extra_headers")
)
if headers is None:
headers = {}
if extra_headers is not None:
@ -404,7 +427,27 @@ async def _arealtime(
model=model,
provider=LlmProviders(_custom_llm_provider),
)
if provider_config is not None:
provider_handler: Final = (
ProviderConfigManager.get_provider_realtime_handler(
LlmProviders(_custom_llm_provider), litellm_params, lambda: websocket.headers, headers
)
if _custom_llm_provider in LlmProviders._member_map_.values()
else None
)
if provider_handler is not None:
user_api_key_dict: Final = TypeAdapter[object](object).validate_python(kwargs.get("user_api_key_dict"))
await provider_handler.async_realtime(
model=model,
websocket=websocket,
logging_obj=litellm_logging_obj,
api_base=api_base or None,
api_key=api_key,
timeout=timeout,
query_params=query_params,
user_api_key_dict=user_api_key_dict,
litellm_metadata=_build_litellm_metadata(kwargs),
)
elif provider_config is not None:
await base_llm_http_handler.async_realtime(
model=model,
websocket=websocket,

View file

@ -31,7 +31,7 @@ class NativeTraceStorage:
def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> Future[None]: ...
def insert_rows(self, table: str, rows: Sequence[Mapping[str, JsonValue]]) -> Future[None]: ...
def lens_query(self, name: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ...
def query(self, query: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ...
def query(self, sql: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ...
@final
class NativeDiagnosticProcessor:

View file

@ -16,7 +16,7 @@ class ImageEditOptionalRequestParams(TypedDict, total=False):
input_fidelity: Literal["high", "low"] | None
mask: str | None
n: int | None
quality: Literal["high", "medium", "low", "standard", "auto"] | None
quality: Literal["high", "medium", "low", "standard", "auto", "xhigh", "max"] | None
response_format: Literal["url", "b64_json"] | None
size: str | None
user: str | None

View file

@ -87,6 +87,11 @@ class ProviderConnection:
litellm_credential_name: str | None = None
configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None
use_xai_oauth: bool | None = None
# ChatGPT OAuth deployment options; read from litellm_params by the ChatGPT
# adapters and listed as owned so they are never swept into extra_body.
chatgpt_auth_profile: str | None = None
chatgpt_token_dir: str | None = None
chatgpt_auth_file: str | None = None
@dataclass(frozen=True, slots=True, kw_only=True)

View file

@ -93,7 +93,7 @@ from litellm.types.responses.main import (
from .base import CachedTokensDetails
FileContent = IO[bytes] | bytes | PathLike
FileContent = IO[bytes] | bytes | PathLike[str]
FileTypes = (
# file (or bytes)
@ -2081,6 +2081,16 @@ class OpenAIRealtimeStreamResponseBaseObject(TypedDict):
type: str
class OpenAIRealtimeSessionClosed(TypedDict):
type: ReadOnly[Literal["session.closed", "session.usage.updated", "litellm.live.initialization"]]
usage: ReadOnly[Mapping[str, object]]
class OpenAILiveResponseEvent(TypedDict):
type: ReadOnly[Literal["response.event"]]
event: ReadOnly[Mapping[str, object]]
class OpenAIRealtimeConversationObject(TypedDict, total=False):
id: str
object: Required[Literal["realtime.conversation"]]
@ -2341,6 +2351,8 @@ class OpenAIRealtimeEventTypes(Enum):
OpenAIRealtimeEvents = (
OpenAIRealtimeStreamResponseBaseObject
| OpenAIRealtimeSessionClosed
| OpenAILiveResponseEvent
| OpenAIRealtimeStreamSessionEvents
| OpenAIRealtimeStreamResponseOutputItemAdded
| OpenAIRealtimeResponseContentPartAdded
@ -2371,6 +2383,8 @@ class ImageGenerationRequestQuality(str, Enum):
LOW = "low"
MEDIUM = "medium"
HIGH = "high"
XHIGH = "xhigh"
MAX = "max"
AUTO = "auto"
STANDARD = "standard"
HD = "hd"

View file

@ -1,6 +1,6 @@
from typing import Any, Literal
from pydantic import BaseModel
from pydantic import AliasChoices, BaseModel, Field
from typing_extensions import ReadOnly, TypedDict
from .llms.openai import (
@ -12,6 +12,16 @@ from .llms.openai import (
ALL_DELTA_TYPES = Literal["text", "audio"]
class LiveSessionDurationUsage(BaseModel):
duration: float = Field(
strict=True, ge=0, allow_inf_nan=False, validation_alias=AliasChoices("seconds", "audio_duration_ms")
)
class LiveSessionUsageEvent(BaseModel):
usage: LiveSessionDurationUsage
class RealtimeResponseTransformInput(TypedDict):
session_configuration_request: str | None
current_output_item_id: (
@ -49,6 +59,7 @@ class RealtimeModalityResponseTransformOutput(TypedDict):
class RealtimeQueryParams(TypedDict, total=False):
model: str
intent: str | None
call_id: ReadOnly[str]
# Add more fields as needed

View file

@ -429,6 +429,7 @@ if TYPE_CHECKING:
)
from litellm.llms.cohere.common_utils import CohereModelInfo
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.openai.realtime.handler import OpenAIRealtime
from litellm.proxy._types import AllowedModelRegion
from litellm.router_utils.get_retry_from_policy import (
get_num_retries_from_retry_policy,
@ -444,7 +445,7 @@ if TYPE_CHECKING:
ChatCompletionToolCallFunctionChunk,
)
from litellm.types.rerank import RerankResponse
from litellm.types.router import LiteLLM_Params
from litellm.types.router import GenericLiteLLMParams, LiteLLM_Params
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.base_llm.completion.transformation import BaseTextCompletionConfig
@ -6199,6 +6200,7 @@ def _get_model_info_helper(
input_cost_per_second=_model_info.get("input_cost_per_second", None),
input_cost_per_audio_token=_model_info.get("input_cost_per_audio_token", None),
input_cost_per_image_token=_model_info.get("input_cost_per_image_token", None),
cache_read_input_image_token_cost=_model_info.get("cache_read_input_image_token_cost", None),
input_cost_per_video_token=_model_info.get("input_cost_per_video_token", None),
input_cost_per_audio_token_batches=_model_info.get("input_cost_per_audio_token_batches", None),
input_cost_per_image_token_batches=_model_info.get("input_cost_per_image_token_batches", None),
@ -9527,6 +9529,10 @@ class ProviderConfigManager:
model: str,
provider: LlmProviders,
) -> BaseImageGenerationConfig | None:
if LlmProviders.CHATGPT == provider:
from litellm.llms.chatgpt.images import ChatGPTImageGenerationConfig
return ChatGPTImageGenerationConfig()
if LlmProviders.OPENAI == provider:
from litellm.llms.openai.image_generation import (
get_openai_image_generation_config,
@ -9707,16 +9713,36 @@ class ProviderConfigManager:
return MetaRealtimeConfig()
return None
@staticmethod
def get_provider_realtime_handler(
provider: LlmProviders,
params: GenericLiteLLMParams,
get_headers: Callable[[], Mapping[str, str]],
extra_headers: Mapping[str, object] | None = None,
) -> OpenAIRealtime | None:
if provider == LlmProviders.CHATGPT:
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
return ChatGPTRealtime(params, get_headers(), extra_headers)
return None
@staticmethod
def get_provider_realtime_http_config(
model: str,
provider: LlmProviders,
params: GenericLiteLLMParams | None = None,
is_call: bool = False,
) -> BaseRealtimeHTTPConfig | None:
"""
Return the HTTP transformation config for realtime HTTP endpoints
(POST /realtime/client_secrets and POST /realtime/calls).
"""
if LlmProviders.CHATGPT == provider:
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
from litellm.types.router import GenericLiteLLMParams
return ChatGPTRealtimeHTTPConfig(params or GenericLiteLLMParams(), use_codex_backend=is_call)
if LlmProviders.OPENAI == provider:
from litellm.llms.openai.realtime.http_transformation import (
OpenAIRealtimeHTTPConfig,
@ -9736,6 +9762,10 @@ class ProviderConfigManager:
model: str,
provider: LlmProviders,
) -> BaseImageEditConfig | None:
if LlmProviders.CHATGPT == provider:
from litellm.llms.chatgpt.images import ChatGPTImageEditConfig
return ChatGPTImageEditConfig()
if LlmProviders.OPENAI == provider:
from litellm.llms.openai.image_edit import get_openai_image_edit_config

View file

@ -31156,6 +31156,16 @@
"max_tokens": 8191,
"mode": "embedding"
},
"chatgpt/gpt-live-1-codex": {
"litellm_provider": "chatgpt",
"mode": "realtime",
"supported_endpoints": [
"/v1/realtime/calls",
"/v1/live"
],
"supports_audio_input": true,
"supports_audio_output": true
},
"chatgpt/gpt-5.5": {
"litellm_provider": "chatgpt",
"source": "https://platform.openai.com/docs/models/gpt-5.5",

View file

@ -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

View file

@ -365,6 +365,9 @@ async def test_unknown_invocation_target_leaves_billing_unset(monkeypatch: pytes
("/v1/realtime", "GET", True),
("/v1/realtime", "POST", False),
("/v1/realtime/client_secrets", "POST", False),
("/live", "POST", False),
("/v1/live", "POST", False),
("/live/sessions/session/accept", "POST", False),
("/mcp/tools/call", "POST", True),
("/a2a/target/message/send", "POST", True),
("/v1/a2a/target/message/send", "POST", True),
@ -573,7 +576,7 @@ def test_registered_inference_routes_have_an_explicit_managed_access_decision(ro
"/videos", "/batches", "/files", "/fine_tuning", "/assistants", "/threads", "/utils/",
"/vector_stores", "/vector_store/", "/search", "/containers", "/skills", "/claude-code/",
"/interactions", "/agents", "/responses/{", "/responses/input_tokens",
"/realtime/client_secrets", "/realtime/calls", "/realtime/transcription_sessions",
"/realtime/client_secrets", "/realtime/calls", "/realtime/transcription_sessions", "/live",
)) or normalized in ("/models", "/cursor/models", "/cursor/v1/models")
concrete: Final = route.split("?")[0].replace("{model}", "model").replace("{model_name:path}", "model")
assert managed_agent_route_allowed(concrete, None) is not unsupported, route

View file

@ -828,7 +828,11 @@ def _azure_relay_router():
model_list=[
{
"model_name": "gpt",
"litellm_params": {"model": "azure_ai/gpt-5.4-mini", "api_base": "https://a.services.ai.azure.com", "api_key": "k"},
"litellm_params": {
"model": "azure_ai/gpt-5.4-mini",
"api_base": "https://a.services.ai.azure.com",
"api_key": "k",
},
},
{
"model_name": "other-group",
@ -1279,7 +1283,8 @@ def test_get_model_from_request_handles_managed_id_decoder_failures():
"/openai/v1/realtime/calls",
],
)
def test_get_model_from_request_extracts_realtime_session_model(route):
@pytest.mark.parametrize("encoded", [False, True])
def test_get_model_from_request_extracts_realtime_session_model(route, encoded):
"""The effective realtime model lives in ``session.model`` (not the
top-level ``model``). It must be surfaced so can_key_call_model() can
validate the model a restricted key is actually requesting.
@ -1289,13 +1294,49 @@ def test_get_model_from_request_extracts_realtime_session_model(route):
"""
assert (
get_model_from_request(
request_data={"session": {"type": "realtime", "model": "gpt-realtime"}},
request_data={"session": '{"model":"gpt-realtime"}' if encoded else {"model": "gpt-realtime"}},
route=route,
)
== "gpt-realtime"
)
@pytest.mark.parametrize("session", ['{"model":"actual-voice"}', {"model": "actual-voice"}])
@pytest.mark.parametrize(
"route",
[
"/v1/realtime/calls",
"/v1/live",
"/live",
"/openai/v1/live",
"/v1/live/sessions",
"/live/sessions",
"/openai/v1/live/sessions",
"/v1/live/sessions/incoming/accept",
],
)
def test_realtime_calls_auth_uses_executed_session_model_despite_decoys(session, route):
assert (
get_model_from_request(
request_data={"model": "body-decoy", "session": session},
route=route,
request_query_params={"model": "query-decoy"},
request_headers={"x-litellm-model": "header-decoy"},
)
== "actual-voice"
)
@pytest.mark.parametrize("model", ["voice,alias", " voice "])
def test_realtime_calls_auth_preserves_exact_session_model(model):
assert get_model_from_request(request_data={"session": {"model": model}}, route="/v1/realtime/calls") == model
@pytest.mark.parametrize("session", ["invalid", "null", "[]", "12", '"text"', "{}"])
def test_realtime_model_extraction_ignores_invalid_serialized_session(session):
assert get_model_from_request(request_data={"session": session}, route="/v1/realtime/calls") is None
def test_get_model_from_request_realtime_includes_top_level_and_session_model():
"""When both top-level and session model are present, both are returned so
neither path can smuggle a disallowed model past the model-access check."""

View file

@ -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"""

View file

@ -365,7 +365,7 @@ async def test_custom_auth_does_not_enforce_key_model_access_by_default():
async def test_post_custom_auth_expired_key_returns_unauthorized():
expired_token = UserAPIKeyAuth(
token="test_token",
expires=datetime.now() - timedelta(minutes=1),
expires=datetime.now(timezone.utc) - timedelta(minutes=1),
)
with pytest.raises(ProxyException) as exc_info:
@ -8456,6 +8456,52 @@ def test_user_api_key_auth_opens_a_datadog_span_for_accepted_and_rejected_keys(t
assert [span for span in report["spans"] if span == auth_span] == [auth_span, auth_span]
@pytest.mark.asyncio
@pytest.mark.parametrize("attachment", ["path", "query"])
@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom", "custom-mixed"])
@pytest.mark.parametrize("query_model", [b"", b"model=unbudgeted"])
async def test_sideband_auth_uses_encrypted_model_for_budget_checks(monkeypatch, attachment, credential, query_model):
import hashlib
import importlib
import time
from unittest.mock import AsyncMock
from fastapi import WebSocket
from litellm.llms.chatgpt.codex import CodexRealtimeCall
from litellm.proxy.realtime_endpoints.call_sessions import encode_call
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-sideband-budget-salt")
token = encode_call(CodexRealtimeCall(
call_id="rtc_test", model="gpt-live-1-codex", alias="budgeted-voice",
owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300,
))
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential.startswith("custom") else {})
seen = []
async def authenticate(request, api_key):
seen.append((await request.json(), api_key))
return "authenticated-with-model"
monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate)
websocket = WebSocket({
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/live/" + token if attachment == "path" else "/v1/realtime",
"path_params": {"call_id": token} if attachment == "path" else {},
"query_string": query_model + (b"&call_id=" + token.encode() if attachment == "query" else b""),
"headers": {
"authorization": [(b"authorization", b"Bearer owner")],
"api-key": [(b"api-key", b"owner")],
"x-litellm-api-key": [(b"x-litellm-api-key", b"owner")],
"custom": [(b"x-proxy-key", b"Bearer owner")],
"custom-mixed": [(b"x-proxy-key", b"Bearer owner"), (b"authorization", b"Bearer other-owner")],
"subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")],
}[credential],
}, AsyncMock(), AsyncMock())
assert await auth_module.user_api_key_auth_websocket(websocket) == "authenticated-with-model"
assert seen == [({"model": "budgeted-voice"}, "Bearer owner")]
@pytest.mark.asyncio
@pytest.mark.parametrize("is_proxy_admin", [False, True], ids=["standard-return", "proxy-admin-return"])
async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_admin):
@ -8572,6 +8618,87 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a
assert token.jwt_claims == {"sub": "jwt-user"}
@pytest.mark.asyncio
@pytest.mark.parametrize("attachment", ["path", "query"])
async def test_sideband_rejects_budget_fallback_before_rerouting(monkeypatch, attachment):
import hashlib
import importlib
import time
from types import SimpleNamespace
from unittest.mock import AsyncMock
from fastapi import HTTPException, WebSocket
from litellm.llms.chatgpt.codex import CodexRealtimeCall
from litellm.proxy.realtime_endpoints.call_sessions import encode_call
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-sideband-budget-salt")
token = encode_call(CodexRealtimeCall(
call_id="rtc_test", model="gpt-live-1-codex", alias="budgeted-voice",
owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300,
))
limiter = SimpleNamespace(
is_key_within_model_budget=AsyncMock(side_effect=litellm.BudgetExceededError(current_cost=2, max_budget=1)),
get_fallback_model_within_budget=AsyncMock(return_value="cheap-voice"),
)
auth = UserAPIKeyAuth(models=["budgeted-voice", "cheap-voice"])
async def authenticate(request, api_key):
data = await request.json()
await auth_module._check_key_model_budget_with_fallback(auth, limiter, data["model"], data, request)
return auth
monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate)
monkeypatch.setattr(auth_module, "can_key_call_model", AsyncMock())
send = AsyncMock()
websocket = WebSocket({
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/live/" + token if attachment == "path" else "/v1/realtime",
"path_params": {"call_id": token} if attachment == "path" else {},
"query_string": b"call_id=" + token.encode() if attachment == "query" else b"",
"headers": [(b"authorization", b"Bearer owner")],
}, AsyncMock(), send)
with pytest.raises(HTTPException) as error:
await auth_module.user_api_key_auth_websocket(websocket)
assert error.value.status_code == 403
limiter.get_fallback_model_within_budget.assert_not_awaited()
send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""})
@pytest.mark.asyncio
@pytest.mark.parametrize("custom_value", [None, b"Bearer different-owner"])
async def test_sideband_custom_header_cannot_fall_back_to_other_credentials(monkeypatch, custom_value):
import hashlib
import importlib
import time
from unittest.mock import AsyncMock
from fastapi import HTTPException, WebSocket
from litellm.proxy import proxy_server
from litellm.llms.chatgpt.codex import CodexRealtimeCall
from litellm.proxy.realtime_endpoints.call_sessions import encode_call
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-custom-header-salt")
monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"})
token = encode_call(CodexRealtimeCall(
call_id="rtc_test", model="gpt-live-1-codex", alias="voice",
owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300,
))
authenticate = AsyncMock()
monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate)
send = AsyncMock()
websocket = WebSocket({
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/live/" + token, "path_params": {"call_id": token}, "query_string": b"",
"headers": [(b"authorization", b"Bearer owner")]
+ ([(b"x-proxy-key", custom_value)] if custom_value is not None else []),
}, AsyncMock(), send)
with pytest.raises(HTTPException) as error:
await auth_module.user_api_key_auth_websocket(websocket)
assert error.value.status_code == 403
authenticate.assert_not_awaited()
send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""})
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route", ["/v1/messages", "/messages", "/v1/chat/completions", "/chat/completions", "/v1/responses", "/responses"]
@ -8641,6 +8768,67 @@ async def test_claude_view_never_reinterprets_explicit_names(monkeypatch, layer)
assert data["model"] == ("foo" if layer == "unclaimed" else encoded)
def _malformed_authorization_websocket(send):
from unittest.mock import AsyncMock
from fastapi import WebSocket
return WebSocket(
{
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/realtime", "query_string": b"",
"headers": [(b"authorization", b"Token malformed")],
},
AsyncMock(),
send,
)
@pytest.mark.parametrize("authorization_value", ["Token malformed", "bearer lowercase"])
def test_get_websocket_api_key_rejects_malformed_authorization(monkeypatch, authorization_value):
import importlib
from unittest.mock import AsyncMock
from fastapi import HTTPException, WebSocket
from litellm.proxy import proxy_server
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setattr(proxy_server, "general_settings", {})
websocket = WebSocket(
{
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/realtime", "query_string": b"",
"headers": [(b"authorization", authorization_value.encode())],
},
AsyncMock(),
AsyncMock(),
)
with pytest.raises(HTTPException) as error:
auth_module.get_websocket_api_key(websocket)
assert error.value.status_code == 403
assert error.value.detail == "Invalid Authorization header format"
@pytest.mark.asyncio
async def test_websocket_auth_closes_policy_violation_on_malformed_authorization(monkeypatch):
import importlib
from unittest.mock import AsyncMock
from fastapi import HTTPException
from litellm.proxy import proxy_server
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setattr(proxy_server, "general_settings", {})
send = AsyncMock()
websocket = _malformed_authorization_websocket(send)
with pytest.raises(HTTPException) as error:
await auth_module.user_api_key_auth_websocket(websocket)
assert error.value.status_code == 403
assert error.value.detail == "Invalid Authorization header format"
send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""})
ISSUER_ONE = "https://issuer-one.example.com"
ISSUER_TWO = "https://issuer-two.example.com"
@ -9205,6 +9393,33 @@ async def test_router_settings_model_group_alias_authorizes_target_for_team(monk
assert get_client_requested_model(request) == "AgentX-LLM"
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ["configured-voice", None])
async def test_websocket_auth_explicit_model_overrides_query(monkeypatch, model):
import importlib
from fastapi import WebSocket
from litellm.proxy import proxy_server
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setattr(proxy_server, "general_settings", {})
seen = []
async def authenticate(request, api_key):
seen.append((await request.json(), api_key))
return "authenticated"
monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate)
websocket = WebSocket({
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/realtime", "path_params": {},
"query_string": b"model=untrusted-query",
"headers": [(b"x-litellm-api-key", b"owner")],
}, AsyncMock(), AsyncMock())
assert await auth_module.user_api_key_auth_websocket_for_model(websocket, model) == "authenticated"
assert seen == [({"model": model or ""}, "Bearer owner")]
@pytest.mark.asyncio
async def test_reserve_budget_after_common_checks_hands_the_reservation_to_the_request_state():
from fastapi import Request

View file

@ -1039,6 +1039,20 @@ def _raw_batches_request(body: Dict[str, Any]) -> MagicMock:
request.headers = {"Content-Type": "application/json"}
request.client = MagicMock()
request.client.host = "127.0.0.1"
request.scope = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.3"},
"http_version": "1.1",
"method": "POST",
"scheme": "http",
"path": "/v1/batches",
"raw_path": b"/v1/batches",
"query_string": b"",
"root_path": "",
"headers": [(b"content-type", b"application/json"), (b"host", b"localhost")],
"client": ("127.0.0.1", 54321),
"server": ("localhost", 8000),
}
request.body = AsyncMock(return_value=json.dumps(body).encode())
return request

View file

@ -2,12 +2,16 @@
Unit Tests for the max parallel request limiter v1 for the proxy
"""
import asyncio
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import RedisCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler,
)
@ -15,6 +19,98 @@ from litellm.proxy.utils import InternalUsageCache, hash_token
from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage
@pytest.mark.asyncio
async def test_realtime_release_preserves_newer_local_admission_while_redis_finishes():
started, finish = asyncio.Event(), asyncio.Event()
async def release(**kwargs):
started.set()
await finish.wait()
remote = MagicMock(spec=RedisCache)
remote.async_register_script.return_value = AsyncMock(side_effect=release)
cache = DualCache(redis_cache=remote)
handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache))
await cache.async_set_cache("key", {"current_requests": 1, "current_rpm": 1, "current_tpm": 7}, local_only=True)
task = asyncio.create_task(handler._release_realtime_counter("key"))
await started.wait()
next_admission = {"current_requests": 1, "current_rpm": 2, "current_tpm": 7}
await cache.async_set_cache("key", next_admission, local_only=True)
finish.set()
await task
assert await cache.async_get_cache("key", local_only=True) == next_admission
@pytest.mark.asyncio
@pytest.mark.parametrize("reject_team", [False, True])
async def test_realtime_attachment_releases_only_acquired_legacy_slots(reject_team):
cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache))
auth = UserAPIKeyAuth(
api_key="attachment-key",
user_id="attachment-user",
team_id="attachment-team",
team_rpm_limit=0 if reject_team else 100,
max_parallel_requests=1,
end_user_id="attachment-end-user",
metadata={"model_rpm_limit": {"test-model": 100}},
)
data = {"model": "test-model", "metadata": {"global_max_parallel_requests": 10}}
minute = datetime.now().strftime("%Y-%m-%d-%H-%M")
team_key = f"attachment-team::{minute}::request_count"
await cache.async_set_cache(team_key, {"current_requests": 3, "current_tpm": 7, "current_rpm": 4})
handler.begin_realtime_attachment(data)
if reject_team:
with pytest.raises(ProxyRateLimitError, match="Rate Limit Handler"):
await handler.async_pre_call_hook(auth, cache, data, "_arealtime")
else:
await handler.async_pre_call_hook(auth, cache, data, "_arealtime")
await handler.async_release_realtime_attachment(data, auth)
await handler.async_release_realtime_attachment(data, auth)
assert await cache.async_get_cache("global_max_parallel_requests") == 0
assert await cache.async_get_cache(f"attachment-key::{minute}::request_count") == {
"current_requests": 0,
"current_tpm": 0,
"current_rpm": 1,
}
assert await cache.async_get_cache(f"attachment-user::{minute}::request_count") == {
"current_requests": 0,
"current_tpm": 0,
"current_rpm": 1,
}
assert await cache.async_get_cache(team_key) == {
"current_requests": 3,
"current_tpm": 7,
"current_rpm": 4 if reject_team else 5,
}
assert await cache.async_get_cache(f"attachment-key::test-model::{minute}::request_count") == {
"current_requests": 0,
"current_tpm": 0,
"current_rpm": 1,
}
end_user = await cache.async_get_cache(f"attachment-end-user::{minute}::request_count")
assert end_user == (None if reject_team else {"current_requests": 0, "current_tpm": 0, "current_rpm": 1})
if not reject_team:
handler.begin_realtime_attachment(data)
await handler.async_pre_call_hook(auth, cache, data, "_arealtime")
await handler.async_release_realtime_attachment(data, auth)
@pytest.mark.asyncio
async def test_realtime_attachment_rejected_before_acquisition_preserves_other_slot():
cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache))
auth = UserAPIKeyAuth(api_key="busy-key", max_parallel_requests=1)
minute = datetime.now().strftime("%Y-%m-%d-%H-%M")
key = f"busy-key::{minute}::request_count"
current = {"current_requests": 1, "current_tpm": 13, "current_rpm": 2}
await cache.async_set_cache(key, current)
data = {"model": "test-model"}
handler.begin_realtime_attachment(data)
with pytest.raises(ProxyRateLimitError, match="Rate Limit Handler"):
await handler.async_pre_call_hook(auth, cache, data, "_arealtime")
await handler.async_release_realtime_attachment(data, auth)
assert await cache.async_get_cache(key) == current
@pytest.mark.asyncio
async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token():
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
@ -62,9 +158,7 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_
team_id = "litellm-team"
end_user_id = "customer-1"
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
@ -103,7 +197,52 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_
key=f"{scope_id}::{precise_minute}::request_count",
litellm_parent_otel_span=None,
)
assert current["current_tpm"] == 50, (
f"expected 50 tokens counted for {scope_id}, "
f"got {current['current_tpm']}"
)
assert current["current_tpm"] == 50, f"expected 50 tokens counted for {scope_id}, got {current['current_tpm']}"
@pytest.mark.asyncio
async def test_realtime_attachment_release_without_receipt_never_touches_counters():
dual_cache = MagicMock()
handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache))
auth = UserAPIKeyAuth(api_key="no-receipt")
await handler.async_release_realtime_attachment({}, auth)
await handler.async_release_realtime_attachment(
{"_legacy_realtime_attachment_reservations": {"cache_keys": [], "global_acquired": True}}, auth
)
# A release without a matching begin (or with a foreign receipt shape) must not decrement anything.
assert dual_cache.mock_calls == []
@pytest.mark.asyncio
async def test_failure_event_skips_realtime_observer_without_decrementing_slots():
from datetime import datetime
from litellm.proxy._types import InternalRequestOrigin
def failure_kwargs() -> dict:
return {
"litellm_params": {"metadata": {"user_api_key": "observer-hash", "global_max_parallel_requests": 5}},
"exception": RuntimeError("backend disconnected"),
}
dual_cache = MagicMock()
dual_cache.async_get_cache = AsyncMock(return_value=None)
dual_cache.async_increment_cache = AsyncMock()
dual_cache.async_batch_set_cache = AsyncMock()
handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache))
start = datetime.now()
end = datetime.now()
kwargs = failure_kwargs()
kwargs["internal_request_origin"] = InternalRequestOrigin.REALTIME_OBSERVER
await handler.async_log_failure_event(kwargs, None, start, end)
# The observer-internal failure mirror must leave the client-facing slot untouched.
assert dual_cache.mock_calls == []
dual_cache.mock_calls.clear()
await handler.async_log_failure_event(failure_kwargs(), None, start, end)
assert dual_cache.async_increment_cache.await_count >= 1
assert any(
call.kwargs.get("key") == "global_max_parallel_requests" and call.kwargs.get("value") == -1
for call in dual_cache.async_increment_cache.await_args_list
)

View file

@ -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()

View file

@ -784,20 +784,22 @@ def _proxy_with_stubbed_reload(prisma):
def _eviction_journal(access_group):
"""Both auth cache keys, in the order a write path has to evict them."""
"""Every auth cache key that holds this group's limits, in the order a write path has to evict them."""
from litellm.proxy.common_utils.user_api_key_cache import (
live_model_access_group_limits_cache_key,
model_access_group_cache_key,
model_access_group_registry_cache_key,
)
return [
f"auth_cache.delete:{model_access_group_cache_key(access_group)}",
f"auth_cache.delete:{live_model_access_group_limits_cache_key(access_group)}",
f"auth_cache.delete:{model_access_group_registry_cache_key()}",
]
def _assert_evicted_after_write(journal, access_group, write_entry):
"""Exactly the two keys, in order, after the DB write. Deliberately not a tail slice: what
"""Exactly the cached keys, in order, after the DB write. Deliberately not a tail slice: what
has to hold is that the eviction follows the write, not that nothing follows the eviction."""
evictions = [entry for entry in journal if entry.startswith("auth_cache.delete:")]
assert evictions == _eviction_journal(access_group)
@ -1206,7 +1208,7 @@ async def test_list_access_groups_reports_a_budgetless_group_as_unbudgeted_rathe
@pytest.mark.asyncio
async def test_put_access_group_budget_evicts_both_auth_cache_keys():
async def test_put_access_group_budget_evicts_every_cached_limit_key():
"""Auth reads the per-group row and the registry of budgeted groups cache-first with no
freshness check, so a PUT that skips either eviction returns 200 and enforces nothing until
the TTL expires. Both keys, after the write."""
@ -1233,7 +1235,7 @@ async def test_put_access_group_budget_evicts_both_auth_cache_keys():
@pytest.mark.asyncio
async def test_delete_access_group_budget_evicts_both_auth_cache_keys():
async def test_delete_access_group_budget_evicts_every_cached_limit_key():
"""Clearing a budget has the same window as setting one: until both keys are dropped, auth
keeps enforcing the budget that is already gone."""
from litellm.proxy.management_endpoints.model_access_group_management_endpoints import (
@ -1252,7 +1254,7 @@ async def test_delete_access_group_budget_evicts_both_auth_cache_keys():
@pytest.mark.asyncio
async def test_deleting_the_access_group_evicts_both_auth_cache_keys():
async def test_deleting_the_access_group_evicts_every_cached_limit_key():
"""The group-delete cascade drops the budget row too, so it owes the same two evictions."""
from litellm.proxy.management_endpoints.model_access_group_management_endpoints import (
delete_access_group,

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,858 @@
import asyncio
import json
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.realtime_endpoints.call_supervision import CallSupervisor, CallSupervisors
class Socket:
def __init__(self):
self.messages = asyncio.Queue()
self.closed = False
def __aiter__(self):
return self
async def __anext__(self):
message = await self.messages.get()
if message is None:
raise StopAsyncIteration
if isinstance(message, Exception):
raise message
return json.dumps(message)
async def close(self):
self.closed = True
@pytest.mark.asyncio
@pytest.mark.parametrize("lease_lost", [False, True])
async def test_supervisor_holds_call_lease_until_terminal_accounting(lease_lost):
from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
lost = asyncio.Event()
lease = MagicMock(spec=RealtimeCallLease)
lease.wait_failed = lost.wait
async def release():
assert socket.closed
assert sink.logs == 1
lease.close = AsyncMock(side_effect=release)
async def close():
await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}})
terminate = AsyncMock(side_effect=close)
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), terminate, lease=lease)
await socket.messages.put({"type": "session.started"})
await supervisor.start()
lease.close.assert_not_awaited()
if lease_lost:
lost.set()
else:
await close()
await asyncio.wait_for(supervisor.wait(), 1)
assert terminate.await_count == int(lease_lost)
lease.close.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("stalled_step", ["close", "drain"])
async def test_live_initial_close_reserves_time_for_independent_hangup(stalled_step):
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
close_cancelled = asyncio.Event()
async def close():
if stalled_step == "close":
try:
await asyncio.Event().wait()
finally:
close_cancelled.set()
async def force_close():
await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}})
force = AsyncMock(side_effect=force_close)
sink = Sink(logger)
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close,
force_close_call=force,
drain_timeout=1,
termination_timeout=0.08,
)
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await asyncio.wait_for(supervisor.close(), timeout=0.5)
force.assert_awaited_once()
assert close_cancelled.is_set() == (stalled_step == "close")
assert any(event["type"] == "session.closed" for event in sink.events)
assert not logger.model_call_details.get("realtime_usage_incomplete")
assert sink.logs == 1
assert socket.closed
@pytest.mark.asyncio
@pytest.mark.parametrize("fallback", ["terminal", "no_terminal", "timeout"])
async def test_live_unacknowledged_close_uses_bounded_independent_hangup(monkeypatch, fallback):
from litellm.proxy.realtime_endpoints import call_supervision
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
invalidate = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
async def force_close():
if fallback == "terminal":
await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}})
elif fallback == "timeout":
await asyncio.Event().wait()
force = AsyncMock(side_effect=force_close)
close = AsyncMock()
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close,
force_close_call=force,
drain_timeout=0.01,
termination_timeout=0.08,
)
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await asyncio.wait_for(supervisor.close(), timeout=0.5)
close.assert_awaited_once()
force.assert_awaited_once()
assert socket.closed
if fallback == "terminal":
invalidate.assert_not_awaited()
assert not logger.model_call_details.get("realtime_usage_incomplete")
else:
invalidate.assert_awaited_once()
assert logger.model_call_details["realtime_usage_incomplete"] is True
@pytest.mark.asyncio
async def test_live_confirmed_terminal_does_not_force_hangup():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
async def close():
await socket.messages.put({"type": "session.closed"})
force = AsyncMock()
supervisor = CallSupervisor(socket, Sink(logger), logger, UserAPIKeyAuth(), close, force_close_call=force)
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await supervisor.close()
force.assert_not_awaited()
class Sink:
def __init__(self, logger):
self.logger = logger
self.events = []
self.logs = 0
def store_message(self, message):
self.events.append(json.loads(message))
async def log_messages(self, *, wait_for_dispatch=False):
assert wait_for_dispatch
self.logs += 1
self.logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
@pytest.mark.asyncio
@pytest.mark.parametrize(
"duration,valid", [(0, True), (1000, True), (None, False), (-1, False), (True, False), ("1000", False)]
)
@pytest.mark.parametrize("duration_field", ["audio_duration_ms", "seconds"])
async def test_live_terminal_requires_valid_duration_for_accounting(monkeypatch, duration, valid, duration_field):
from litellm.proxy.realtime_endpoints import call_supervision
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
invalidate = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
close = AsyncMock()
force = AsyncMock()
supervisor = CallSupervisor(socket, Sink(logger), logger, UserAPIKeyAuth(), close, force_close_call=force)
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await socket.messages.put(
{"type": "session.closed", **({"usage": {duration_field: duration}} if duration is not None else {})}
)
await supervisor.wait()
close.assert_not_awaited()
force.assert_not_awaited()
assert socket.closed
assert bool(logger.model_call_details.get("realtime_usage_incomplete")) is not valid
assert invalidate.await_count == (0 if valid else 1)
def fixture(*, ready_timeout=1, lifetime=1):
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
async def hangup():
assert not socket.closed
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
close_call = AsyncMock(side_effect=hangup)
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close_call,
ready_timeout=ready_timeout,
lifetime=lifetime,
drain_timeout=0.05,
)
return socket, sink, close_call, supervisor
@pytest.mark.asyncio
async def test_observer_logs_webrtc_usage_without_client_sideband():
socket, sink, close_call, supervisor = fixture()
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await socket.messages.put({"type": "response.done", "response": {"usage": {"total_tokens": 15}}})
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 19}})
await supervisor.wait()
await supervisor.close()
assert sink.logs == 1
assert sink.events[-1]["usage"]["total_tokens"] == 19
assert sink.events[1]["response"]["usage"]["total_tokens"] == 15
assert socket.closed
close_call.assert_not_awaited()
@pytest.mark.asyncio
async def test_early_upstream_eof_rejects_start():
socket, sink, close_call, supervisor = fixture()
await socket.messages.put(None)
with pytest.raises(RuntimeError, match="ended before"):
await supervisor.start()
assert socket.closed
assert sink.logs == 1
close_call.assert_awaited_once()
@pytest.mark.asyncio
async def test_cancelled_start_hangs_up_and_drains_terminal_usage():
socket, sink, close_call, supervisor = fixture()
started = asyncio.create_task(supervisor.start())
await asyncio.sleep(0)
started.cancel()
with pytest.raises(asyncio.CancelledError):
await started
close_call.assert_awaited_once()
assert socket.closed
assert sink.logs == 1
assert sink.events[-1]["usage"]["total_tokens"] == 42
@pytest.mark.asyncio
@pytest.mark.parametrize("cancel_count", [1, 2, 3])
async def test_repeated_start_cancellation_keeps_lease_until_shutdown_finishes(cancel_count):
from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease
reading = asyncio.Event()
close_entered = asyncio.Event()
allow_close = asyncio.Event()
released = asyncio.Event()
class ObservedSocket(Socket):
async def __anext__(self):
reading.set()
return await super().__anext__()
socket = ObservedSocket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
async def close_call():
close_entered.set()
await allow_close.wait()
await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}})
async def release():
released.set()
lease = RealtimeCallLease(renew=AsyncMock(return_value=True), release=release)
lease.start()
supervisor = CallSupervisor(
socket, sink, logger, UserAPIKeyAuth(), close_call, lease=lease, ready_timeout=10, termination_timeout=10
)
registry = CallSupervisors()
async def signaling():
transferred = False
try:
await registry.start(supervisor)
transferred = True
finally:
# The signaling endpoint retains lease ownership until registry startup succeeds.
if not transferred:
await lease.close()
started = asyncio.create_task(signaling())
shutdown = None
try:
await asyncio.wait_for(reading.wait(), timeout=1)
started.cancel()
await asyncio.wait_for(close_entered.wait(), timeout=1)
for _ in range(cancel_count - 1):
started.cancel()
done, _ = await asyncio.wait({started}, timeout=0.02)
assert not done
assert not released.is_set()
shutdown = asyncio.create_task(registry.shutdown())
done, _ = await asyncio.wait({started, shutdown}, timeout=0.02)
assert not done
assert not released.is_set()
assert not socket.closed
assert sink.logs == 0
allow_close.set()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(started, timeout=1)
await asyncio.wait_for(shutdown, timeout=1)
assert socket.closed
assert sink.logs == 1
assert released.is_set()
assert sink.events[-1]["usage"]["audio_duration_ms"] == 1000
finally:
allow_close.set()
await asyncio.wait_for(supervisor.wait(), timeout=1)
await asyncio.gather(started, return_exceptions=True)
if shutdown is not None:
await shutdown
await registry.shutdown()
await lease.close()
@pytest.mark.asyncio
async def test_worker_shutdown_drains_all_calls():
registry = CallSupervisors()
socket, sink, close_call, supervisor = fixture()
await socket.messages.put({"type": "session.created"})
await registry.start(supervisor)
await registry.shutdown()
await registry.shutdown()
close_call.assert_awaited_once()
assert socket.closed
assert sink.logs == 1
assert sink.events[-1]["usage"]["total_tokens"] == 42
@pytest.mark.asyncio
async def test_ready_timeout_hangs_up_before_returning_error():
socket, sink, close_call, supervisor = fixture(ready_timeout=0.01)
with pytest.raises(asyncio.TimeoutError):
await supervisor.start()
close_call.assert_awaited_once()
assert socket.closed
assert sink.logs == 1
@pytest.mark.asyncio
async def test_lifetime_limit_closes_call_and_collects_final_usage():
socket, sink, close_call, supervisor = fixture(lifetime=0.01)
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await supervisor.wait()
close_call.assert_awaited_once()
assert socket.closed
assert sink.events[-1]["usage"]["total_tokens"] == 42
@pytest.mark.asyncio
async def test_socket_eof_after_ready_still_hangs_up_provider_call(monkeypatch):
from litellm.proxy.realtime_endpoints import call_supervision
invalidate = AsyncMock()
release = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
socket, sink, close_call, supervisor = fixture()
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await socket.messages.put(None)
await supervisor.wait()
close_call.assert_awaited_once()
assert socket.closed
assert sink.logs == 1
assert sink.logger.model_call_details["realtime_usage_incomplete"] is True
invalidate.assert_awaited_once()
release.assert_not_awaited()
@pytest.mark.asyncio
async def test_observer_error_rejects_start(caplog):
socket, sink, close_call, supervisor = fixture()
await socket.messages.put(RuntimeError("private-provider-credential"))
with pytest.raises(RuntimeError, match="ended before"):
await supervisor.start()
assert socket.closed
close_call.assert_awaited_once()
assert "private-provider-credential" not in caplog.text
@pytest.mark.asyncio
async def test_failed_logging_invalidates_reservation_without_zeroing_spend(monkeypatch):
from litellm.proxy.realtime_endpoints import call_supervision
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = MagicMock()
sink.log_messages = AsyncMock(side_effect=RuntimeError("logging unavailable"))
release = AsyncMock()
invalidate = AsyncMock()
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock())
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await socket.messages.put({"type": "session.closed"})
with pytest.raises(RuntimeError, match="logging unavailable"):
await supervisor.wait()
release.assert_not_awaited()
invalidate.assert_awaited_once_with(budget_reservation=None)
assert logger.model_call_details["realtime_accounting_incomplete"] is True
assert socket.closed
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
@pytest.mark.asyncio
async def test_start_rejects_terminal_session_while_accounting_is_pending():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
dispatch_started = asyncio.Event()
allow_dispatch = asyncio.Event()
async def log_messages(*, wait_for_dispatch=False):
dispatch_started.set()
await allow_dispatch.wait()
logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
sink = MagicMock()
sink.log_messages = AsyncMock(side_effect=log_messages)
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock())
await socket.messages.put({"type": "session.created"})
await socket.messages.put({"type": "session.closed"})
startup = asyncio.create_task(supervisor.start())
try:
await asyncio.wait_for(dispatch_started.wait(), timeout=1)
finally:
allow_dispatch.set()
with pytest.raises(RuntimeError, match="ended before"):
await asyncio.wait_for(startup, timeout=1)
assert socket.closed
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
@pytest.mark.asyncio
async def test_shutdown_bounds_accounting_and_invalidates_partial_dispatch(monkeypatch):
from litellm.proxy.realtime_endpoints import call_supervision
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
dispatch_cancelled = asyncio.Event()
invalidate = AsyncMock()
release = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
async def log_messages(*, wait_for_dispatch=False):
try:
await asyncio.Event().wait()
finally:
dispatch_cancelled.set()
async def hangup():
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
sink = MagicMock()
sink.log_messages = AsyncMock(side_effect=log_messages)
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), hangup, logging_timeout=0.01)
registry = CallSupervisors()
await socket.messages.put({"type": "session.created"})
await registry.start(supervisor)
await asyncio.wait_for(registry.shutdown(), timeout=1)
assert dispatch_cancelled.is_set()
assert socket.closed
assert logger.model_call_details["realtime_accounting_incomplete"] is True
assert not logger.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY)
invalidate.assert_awaited_once_with(budget_reservation=None)
release.assert_not_awaited()
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
@pytest.mark.asyncio
async def test_shutdown_waits_for_usage_dispatch_completion():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
dispatch_started = asyncio.Event()
dispatch_complete = asyncio.Event()
dispatch_finished = asyncio.Event()
async def log_messages(*, wait_for_dispatch=False):
assert wait_for_dispatch
dispatch_started.set()
await dispatch_complete.wait()
logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
dispatch_finished.set()
async def hangup():
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
sink = MagicMock()
sink.log_messages = AsyncMock(side_effect=log_messages)
registry = CallSupervisors()
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), hangup)
await socket.messages.put({"type": "session.created"})
await registry.start(supervisor)
shutdown = asyncio.create_task(registry.shutdown())
try:
await asyncio.wait_for(dispatch_started.wait(), timeout=1)
assert not shutdown.done()
assert not dispatch_finished.is_set()
finally:
dispatch_complete.set()
await asyncio.wait_for(shutdown, timeout=1)
assert dispatch_finished.is_set()
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal_usage_required", [True, False])
async def test_confirmed_hangup_without_terminal_usage_matches_protocol(terminal_usage_required):
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
close_call = AsyncMock()
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close_call,
terminal_usage_required=terminal_usage_required,
drain_timeout=0.01,
)
await socket.messages.put({"type": "session.created"})
await supervisor.start()
await socket.messages.put({"type": "response.done", "response": {"usage": {"total_tokens": 17}}})
await supervisor.close()
assert bool(logger.model_call_details.get("realtime_usage_incomplete")) == terminal_usage_required
assert sink.events[-1]["response"]["usage"]["total_tokens"] == 17
assert sink.logs == 1
assert socket.closed
close_call.assert_awaited_once()
@pytest.mark.asyncio
async def test_shutdown_allows_hangup_longer_than_usage_drain_timeout():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
hangup_started = asyncio.Event()
allow_hangup = asyncio.Event()
hangup_finished = asyncio.Event()
async def hangup():
hangup_started.set()
await allow_hangup.wait()
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
hangup_finished.set()
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
hangup,
drain_timeout=0.01,
termination_timeout=1,
terminal_usage_required=False,
)
registry = CallSupervisors()
await socket.messages.put({"type": "session.created"})
await registry.start(supervisor)
shutdown = asyncio.create_task(registry.shutdown())
try:
await asyncio.wait_for(hangup_started.wait(), timeout=1)
await asyncio.sleep(0.04)
assert not shutdown.done()
assert not socket.closed
assert not hangup_finished.is_set()
finally:
allow_hangup.set()
await asyncio.wait_for(shutdown, timeout=1)
assert hangup_finished.is_set()
assert socket.closed
assert sink.logs == 1
assert sink.events[-1]["usage"]["total_tokens"] == 42
assert not logger.model_call_details.get("realtime_usage_incomplete")
@pytest.mark.asyncio
async def test_termination_timeout_cancels_hangup_and_finishes_cleanup():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
hangup_cancelled = asyncio.Event()
async def hangup():
try:
await asyncio.Event().wait()
finally:
hangup_cancelled.set()
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
hangup,
drain_timeout=0.01,
termination_timeout=0.02,
terminal_usage_required=False,
)
await socket.messages.put({"type": "session.created"})
await supervisor.start()
await asyncio.wait_for(supervisor.close(), timeout=1)
assert hangup_cancelled.is_set()
assert socket.closed
assert sink.logs == 1
assert logger.model_call_details["realtime_usage_incomplete"] is True
@pytest.mark.asyncio
@pytest.mark.parametrize("closure", ["eof", "normal_close", "error"])
@pytest.mark.parametrize("hangup_succeeds", [True, False])
async def test_ga_observer_disconnect_requires_confirmed_hangup(closure, hangup_succeeds):
from websockets.exceptions import ConnectionClosedOK
from websockets.frames import Close
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
close_call = AsyncMock(side_effect=None if hangup_succeeds else RuntimeError("unconfirmed hangup"))
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close_call,
terminal_usage_required=False,
drain_timeout=0.01,
)
await socket.messages.put({"type": "session.created"})
await supervisor.start()
await socket.messages.put(
None
if closure == "eof"
else ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True)
if closure == "normal_close"
else RuntimeError("observer failed")
)
await supervisor.wait()
close_call.assert_awaited_once()
assert bool(logger.model_call_details.get("realtime_usage_incomplete")) == (not hangup_succeeds)
assert sink.logs == 1
assert socket.closed
@pytest.mark.asyncio
async def test_live_connected_attach_is_ready_without_session_started_and_bills_once():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
async def close():
await socket.messages.put({"type": "session.closed", "usage": {"seconds": 30}})
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close,
connected_ready=True,
ready_timeout=0.1,
)
await supervisor.start()
await socket.messages.put({"type": "session.usage.updated", "usage": {"seconds": 15}})
await supervisor.close()
await supervisor.close()
assert sink.logs == 1
assert sink.events[-1] == {"type": "session.closed", "usage": {"seconds": 30}}
assert "realtime_usage_incomplete" not in logger.model_call_details
assert socket.closed
@pytest.mark.asyncio
async def test_live_missing_backend_accounting_invalidates_budget_after_dispatch(monkeypatch):
from litellm.proxy.realtime_endpoints import call_supervision
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
invalidate = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
class IncompleteSink(Sink):
async def log_messages(self, *, wait_for_dispatch=False):
await super().log_messages(wait_for_dispatch=wait_for_dispatch)
logger.model_call_details["realtime_backend_accounting_incomplete"] = True
sink = IncompleteSink(logger)
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock(), connected_ready=True)
await supervisor.start()
await socket.messages.put({"type": "session.closed", "usage": {"seconds": 30}})
await supervisor.wait()
assert sink.logs == 1
assert logger.model_call_details["realtime_accounting_incomplete"] is True
assert "realtime_usage_incomplete" not in logger.model_call_details
invalidate.assert_awaited_once()
@pytest.mark.asyncio
async def test_repeated_start_is_rejected_while_observer_is_running():
socket, sink, close_call, supervisor = fixture()
await socket.messages.put({"type": "session.started"})
await supervisor.start()
with pytest.raises(RuntimeError, match="Call observer already started"):
await supervisor.start()
# The rejected second start leaves the running observer untouched.
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 7}})
await supervisor.wait()
await supervisor.close()
close_call.assert_not_awaited()
assert sink.logs == 1
assert socket.closed
@pytest.mark.asyncio
async def test_start_rejects_quota_reservation_lost_during_startup():
from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
async def hangup():
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
close_call = AsyncMock(side_effect=hangup)
renew = AsyncMock(return_value=False)
release = AsyncMock()
lease = RealtimeCallLease(renew=renew, release=release, interval=3600)
lease.start()
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close_call,
ready_timeout=10,
lifetime=10,
drain_timeout=0.05,
lease=lease,
)
await socket.messages.put({"type": "session.started"})
with pytest.raises(RuntimeError, match="lost its quota reservation during startup"):
await supervisor.start()
assert socket.closed
close_call.assert_awaited_once()
assert renew.await_count >= 1
release.assert_awaited_once()
assert sink.logs == 1
@pytest.mark.asyncio
async def test_registry_watch_logs_observer_accounting_failure_without_payload(caplog, monkeypatch):
import logging
from litellm.proxy.realtime_endpoints import call_supervision
caplog.set_level(logging.ERROR, logger="LiteLLM Proxy")
class FailingSink(Sink):
async def log_messages(self, *, wait_for_dispatch=False):
raise RuntimeError("observer accounting secret-token")
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = FailingSink(logger)
async def hangup():
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
close_call = AsyncMock(side_effect=hangup)
invalidate = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close_call,
ready_timeout=5,
lifetime=5,
drain_timeout=0.05,
)
registry = CallSupervisors()
await socket.messages.put({"type": "session.created"})
await registry.start(supervisor)
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
with pytest.raises(RuntimeError, match="observer accounting secret-token"):
await supervisor.wait()
await registry.shutdown()
assert registry._calls == ()
assert registry._tasks == ()
assert socket.closed
close_call.assert_not_awaited()
invalidate.assert_awaited_once_with(budget_reservation=None)
assert logger.model_call_details["realtime_accounting_incomplete"] is True
proxy_logs = [record.getMessage() for record in caplog.records if record.name == "LiteLLM Proxy"]
assert any("Realtime observer accounting failed" in message for message in proxy_logs)
assert not any("secret-token" in message for message in proxy_logs)

File diff suppressed because it is too large Load diff

View file

@ -7,7 +7,7 @@ Tests for LiteLLM proxy realtime WebRTC HTTP endpoints:
import json
import time
from collections.abc import Awaitable
from typing import Protocol
from typing import Final, Protocol
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -15,7 +15,6 @@ import pytest
from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
@ -120,6 +119,114 @@ def proxy_app(monkeypatch):
return proxy_server.app
@pytest.mark.parametrize("path", ["/v1/live", "/live", "/openai/v1/live"])
@pytest.mark.parametrize("query", ["", "?intent=custom&architecture=custom"])
def test_live_multipart_offer_runs_authenticated_call_pipeline(
proxy_app: FastAPI, monkeypatch: pytest.MonkeyPatch, path: str, query: str
) -> None:
from litellm.proxy import proxy_server
from litellm.proxy.realtime_endpoints import call_sessions
monkeypatch.setenv("LITELLM_SALT_KEY", "test-live-signaling-salt")
session: Final = {"model": "gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}}
authenticate: Final = AsyncMock(wraps=call_sessions.user_api_key_auth)
process: Final = AsyncMock(side_effect=lambda request, data, *args: (data, MagicMock()))
supervise: Final = AsyncMock()
upstream: Final = httpx.Response(
201,
content=b"v=0\r\nanswer",
headers={"Location": "/v1/live/rtc_private"},
extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}},
)
async def route(**kwargs: object) -> Awaitable[httpx.Response]:
assert kwargs["route_type"] == "arealtime_calls"
assert isinstance(kwargs["data"], dict)
assert kwargs["data"]["sdp_body"] == b"v=0\r\noffer"
assert kwargs["data"]["session"] == session
assert kwargs["data"]["chatgpt_realtime_client_query"] == (
{"architecture": "custom", "intent": "custom"}
if query
else {"architecture": "avas", "intent": "quicksilver"}
)
async def respond() -> httpx.Response:
return upstream
return respond()
monkeypatch.setattr(call_sessions, "user_api_key_auth", authenticate)
monkeypatch.setattr(call_sessions, "process_codex_request", process)
monkeypatch.setattr(call_sessions, "supervise_codex_call", supervise)
monkeypatch.setattr(proxy_server, "route_request", route)
response: Final = TestClient(proxy_app).post(
f"{path}{query}",
headers={"Authorization": "Bearer sk-test-master-key"},
files={
"sdp": (None, "v=0\r\noffer", "application/sdp"),
"session": (None, json.dumps(session), "application/json"),
},
)
assert response.status_code == 201
assert response.content == b"v=0\r\nanswer"
assert response.headers["content-type"] == "application/sdp"
authenticate.assert_awaited_once()
process.assert_awaited_once()
assert process.await_args.args[3:] == ("gpt-live-1-codex", "arealtime_calls")
assert isinstance(process.await_args.args[2], UserAPIKeyAuth)
supervise.assert_awaited_once()
assert response.headers["location"].startswith("/v1/live/")
token: Final = response.headers["location"].rsplit("/", 1)[-1]
call: Final = call_sessions.decode_call(token, "Bearer sk-test-master-key")
assert call.call_id == "rtc_private"
assert call.alias == "gpt-live-1-codex"
assert call.usage_supervised
@pytest.mark.parametrize("path", ["/v1/live", "/live", "/openai/v1/live"])
def test_live_multipart_offer_rejects_invalid_credentials_before_routing(
proxy_app: FastAPI, monkeypatch: pytest.MonkeyPatch, path: str
) -> None:
from litellm.proxy import proxy_server
from litellm.proxy.realtime_endpoints import call_sessions
authenticate: Final = AsyncMock(wraps=call_sessions.user_api_key_auth)
route: Final = AsyncMock()
monkeypatch.setattr(call_sessions, "user_api_key_auth", authenticate)
monkeypatch.setattr(proxy_server, "route_request", route)
response: Final = TestClient(proxy_app).post(
path,
files={"sdp": (None, "v=0\r\n"), "session": (None, '{"model":"gpt-live-1-codex"}')},
)
assert response.status_code == 401
authenticate.assert_awaited_once()
route.assert_not_awaited()
def test_live_multipart_offer_rejects_model_outside_key_scope(
proxy_app: FastAPI, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy import proxy_server
from litellm.proxy.realtime_endpoints import call_sessions
authenticate: Final = AsyncMock(return_value=UserAPIKeyAuth(models=["another-model"]))
route: Final = AsyncMock()
monkeypatch.setattr(call_sessions, "user_api_key_auth", authenticate)
monkeypatch.setattr(proxy_server, "route_request", route)
response: Final = TestClient(proxy_app).post(
"/v1/live",
headers={"Authorization": "Bearer restricted-key"},
files={"sdp": (None, "v=0\r\n"), "session": (None, '{"model":"gpt-live-1-codex"}')},
)
assert response.status_code == 403
assert "gpt-live-1-codex" in response.text
authenticate.assert_awaited_once()
route.assert_not_awaited()
@pytest.fixture
def mock_route_request_client_secrets():
"""Mock route_request to return a fake upstream client_secrets response."""
@ -414,12 +521,109 @@ def test_realtime_calls_invalid_token_returns_401(proxy_app):
assert "Invalid or expired token" in response.json().get("error", "")
@pytest.mark.parametrize("handle_kind", ["codex", "live"])
@pytest.mark.parametrize("provider", ["chatgpt", "openai"])
def test_legacy_sdp_rejects_handle_ciphertext_before_oauth_dispatch(
proxy_app, monkeypatch, tmp_path, handle_kind, provider
):
import base64
import hashlib
from litellm import Router
from litellm.llms.chatgpt.codex import CodexRealtimeCall
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy import proxy_server
from litellm.proxy.realtime_endpoints.call_sessions import encode_call
from litellm.proxy.realtime_endpoints.live import LiveHandle, encode_session
monkeypatch.setenv("LITELLM_SALT_KEY", "test-sdp-handle-salt")
monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path))
monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json")
(tmp_path / "auth.json").write_text(
json.dumps(
{"access_token": "test-sdp-oauth", "account_id": "test-sdp-account", "expires_at": time.time() + 3600}
)
)
owner = hashlib.sha256(b"Bearer restricted-key", usedforsecurity=False).hexdigest()
handle = (
encode_call(
CodexRealtimeCall(
call_id="rtc_allowed",
model="gpt-live-1-codex",
alias="allowed-voice",
owner=owner,
expires_at=time.time() + 3600,
)
)
if handle_kind == "codex"
else encode_session(
LiveHandle(
session_id="live_allowed",
alias="allowed-voice",
deployment={"model": "chatgpt/gpt-live-1-codex"},
owner=owner,
expires_at=time.time() + 3600,
policy={},
)
)
)
encoded = handle.removeprefix("rtc_litellm_").removeprefix("live_litellm_")
ciphertext = base64.urlsafe_b64decode(encoded + "=" * (-len(encoded) % 4)).decode()
requests = []
def respond(request):
requests.append(request)
return httpx.Response(201, text="v=0\r\nanswer")
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
router = Router(
model_list=[
{
"model_name": "forbidden-voice",
"litellm_params": {
"model": "chatgpt/gpt-live-1-codex" if provider == "chatgpt" else "openai/gpt-realtime"
},
}
]
)
async def add_data(data, **kwargs):
return data
async def pre_call(user_api_key_dict, data, call_type):
return data
async def route(data, route_type, **kwargs):
assert route_type == "arealtime_calls"
return router.arealtime_calls(**data, client=client)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "user_model", None)
monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", add_data)
monkeypatch.setattr(proxy_server, "route_request", route)
monkeypatch.setattr(
proxy_server,
"proxy_logging_obj",
MagicMock(pre_call_hook=AsyncMock(side_effect=pre_call), post_call_failure_hook=AsyncMock()),
)
response = TestClient(proxy_app).post(
"/v1/realtime/calls?model=forbidden-voice",
headers={"Authorization": f"Bearer {ciphertext}", "Content-Type": "application/sdp"},
content=b"v=0\r\noffer",
)
assert response.status_code == 401, [(r.url.path, r.headers.get("authorization")) for r in requests]
assert not requests
@pytest.mark.asyncio
@pytest.mark.parametrize("token_format", ["versioned", "legacy"])
async def test_realtime_calls_success_with_valid_encrypted_token(
proxy_app,
mock_route_request_realtime_calls,
mock_add_litellm_data,
mock_pre_call_hook,
token_format,
):
"""POST /v1/realtime/calls returns 201 with valid encrypted token from client_secrets."""
# Build a valid encrypted token (same format as client_secrets returns)
@ -431,7 +635,7 @@ async def test_realtime_calls_success_with_valid_encrypted_token(
team_id=None,
expires_at=future_expires_at,
)
encrypted_token = encrypt_value_helper(token_payload)
encrypted_token = encrypt_value_helper(token_payload if token_format == "versioned" else "fake_upstream_epk")
client = TestClient(proxy_app)
with (

View file

@ -0,0 +1,64 @@
from unittest.mock import AsyncMock
import httpx
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
@pytest.mark.asyncio
@pytest.mark.parametrize("prefix", ["/live", "/v1/live", "/openai/v1/live"])
@pytest.mark.parametrize(
("method", "suffix"),
[
("POST", ""),
("POST", "/opaque/fork"),
("POST", "/opaque/accept"),
("POST", "/opaque/reject"),
("POST", "/opaque/refer"),
("POST", "/opaque/hangup"),
("GET", "/opaque/content"),
],
)
async def test_public_live_http_routes_reach_live_auth_before_generic_passthrough(monkeypatch, prefix, method, suffix):
from litellm.proxy import proxy_server
from litellm.proxy.realtime_endpoints import live
authenticate = AsyncMock(side_effect=HTTPException(401, "Live authentication required"))
monkeypatch.setattr(live, "_auth", authenticate)
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=proxy_server.app), base_url="http://proxy"
) as client:
response = await client.request(
method,
prefix + "/sessions" + suffix,
json={"session": {"model": "voice"}, "transport": {"type": "webrtc", "sdp": "offer"}},
)
assert response.status_code == 401
assert "Live authentication required" in response.text
authenticate.assert_awaited_once()
@pytest.mark.parametrize("prefix", ["/live", "/v1/live", "/openai/v1/live"])
@pytest.mark.parametrize("suffix", ["", "/opaque/attach", "/opaque/fork"])
def test_public_live_websockets_reach_live_auth_before_legacy_sideband(monkeypatch, prefix, suffix):
from starlette.websockets import WebSocketDisconnect
from litellm.proxy import proxy_server
from litellm.proxy.realtime_endpoints import live
authenticate = AsyncMock(side_effect=HTTPException(403, "Live authentication rejected"))
monkeypatch.setattr(live, "_auth", authenticate)
monkeypatch.setattr(proxy_server, "general_settings", {})
# TestClient runs the proxy lifespan, and the boot check refuses a weak or unset master key
# before the app serves anything. Set a safe key here instead of relying on the ambient one, so
# the request really reaches the routes: with a key in place the legacy sideband dependency
# would reject the "Bearer test" header, so awaiting _auth still proves which auth ran first.
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-live-route-registration-test-master-key")
# A previous proxy test may leave the module scheduler bound to a closed loop.
monkeypatch.setattr(proxy_server, "scheduler", None)
with TestClient(proxy_server.app) as client:
with pytest.raises(WebSocketDisconnect):
with client.websocket_connect(prefix + "/sessions" + suffix, headers={"authorization": "Bearer test"}):
pass
authenticate.assert_awaited_once()

View file

@ -121,6 +121,11 @@ async def test_real_proxy_child_auth_privacy_and_body_policy(
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "general_settings", {})
# Pin the proxy-wide budget: authentication only reads the global spend when a
# proxy max budget is configured, and that read goes through the stub prisma
# client above. A budget left set by an earlier test on the same worker would
# turn this fixture's child requests into 401s.
monkeypatch.setattr(litellm, "max_budget", 0.0)
monkeypatch.setattr(common_request_processing, "route_request", route)
with inherit_message_logging_privacy(True):
call: Final = with_proxy_compaction_executor(

View file

@ -3299,7 +3299,10 @@ async def test_add_proxy_budget_to_db_only_creates_user_no_keys(monkeypatch: pyt
import litellm
from litellm.proxy.proxy_server import ProxyStartupEvent
# Set up required litellm settings
# Set up required litellm settings. Through monkeypatch rather than plain
# assignment: `litellm.max_budget` is process-global, and any later test on
# this worker that authenticates reads the global proxy spend whenever a
# proxy budget is set, which needs a real prisma client.
monkeypatch.setattr(litellm, "budget_duration", "30d")
monkeypatch.setattr(litellm, "max_budget", 100.0)
@ -14810,6 +14813,31 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the
)
@pytest.mark.asyncio
async def test_login_throttle_config_settings_override_database(monkeypatch):
import litellm.proxy.proxy_server as ps
from litellm.proxy.proxy_server import ProxyConfig
config = ProxyConfig()
configured = {
"max_failed_login_attempts_per_source": 5,
"failed_login_window_seconds": 60,
"failed_login_block_seconds": 120,
}
config.settings.load_yaml(configured)
monkeypatch.setattr(ps, "general_settings", config.settings)
await config._update_general_settings(
db_general_settings={
"max_failed_login_attempts_per_source": 999,
"failed_login_window_seconds": 1,
"failed_login_block_seconds": 1,
}
)
for key, value in configured.items():
assert ps.general_settings[key] == value
assert config.settings.source(key) == "config"
@pytest.mark.asyncio
async def test_login_throttle_limits_from_the_config_file_outrank_the_database(monkeypatch):
import litellm.proxy.proxy_server as ps

View file

@ -1,6 +1,7 @@
import datetime as real_datetime
import smtplib
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
@ -11,15 +12,16 @@ from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.bug_report import ISSUE_URL_BASE
from litellm.proxy._types import ProxyErrorTypes, UserAPIKeyAuth
from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy
from litellm.proxy.utils import (
PrismaClient,
ProxyLogging,
get_custom_url,
handle_exception_on_proxy,
join_paths,
)
from litellm.types.guardrails import GuardrailEventHooks
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy.utils import get_custom_url, join_paths
def test_get_custom_url(monkeypatch):
monkeypatch.setenv("SERVER_ROOT_PATH", "/litellm")
custom_url = get_custom_url(request_base_url="http://0.0.0.0:4000", route="ui/")
@ -2153,6 +2155,90 @@ async def test_proxy_only_error_5xx_keeps_traceback_and_runs_sync_callbacks(monk
assert "test_proxy_utils" in captured["async_traceback"]
def test_create_model_info_response_resolves_alias_to_deployment_model():
"""A public model name that is not itself a cost-map key must not be resolved through
the fallback-generalization rules: `bedrock-claude-opus-5` matches the generic
claude-family baseline (200k/64k) by substring, while the deployment it fronts really
accepts 1M/128k. Regression for the /v1/models alias resolution introduced in v1.94.0."""
from litellm import Router
saved_model_cost = dict(litellm.model_cost)
try:
router = Router(
model_list=[
{
"model_name": "bedrock-claude-opus-5",
"litellm_params": {
"custom_llm_provider": "bedrock",
"model": "bedrock/eu.anthropic.claude-opus-5",
},
"model_info": {"base_model": "eu.anthropic.claude-opus-5"},
}
]
)
response = create_model_info_response(model_id="bedrock-claude-opus-5", provider="openai", llm_router=router)
finally:
litellm.model_cost.clear()
litellm.model_cost.update(saved_model_cost)
assert response["max_input_tokens"] == 1000000
assert response["max_output_tokens"] == 128000
def test_create_model_info_response_keeps_exact_alias_over_generalized_deployment_model():
"""Mirror of the alias bug: when the deployment points at a custom backend name that
only matches a generalization rule, the listed name's exact cost-map entry is the
better answer and must win."""
from litellm import Router
saved_model_cost = dict(litellm.model_cost)
try:
router = Router(
model_list=[
{
"model_name": "claude-opus-5",
"litellm_params": {
"custom_llm_provider": "bedrock",
"model": "bedrock/my-claude-opus-5-provisioned",
},
}
]
)
response = create_model_info_response(model_id="claude-opus-5", provider="openai", llm_router=router)
finally:
litellm.model_cost.clear()
litellm.model_cost.update(saved_model_cost)
assert response["max_input_tokens"] == 1000000
def test_create_model_info_response_falls_back_to_alias_for_opaque_deployment_name():
"""An Azure deployment named after the resource rather than the model has no cost-map
entry; the listed name still does, and must keep answering."""
from litellm import Router
saved_model_cost = dict(litellm.model_cost)
try:
router = Router(
model_list=[
{
"model_name": "gpt-4o",
"litellm_params": {"model": "azure/my-gpt4o-deployment"},
}
]
)
response = create_model_info_response(model_id="gpt-4o", provider="openai", llm_router=router)
finally:
litellm.model_cost.clear()
litellm.model_cost.update(saved_model_cost)
assert response["max_input_tokens"] == 128000
assert response["max_output_tokens"] == 16384
def test_create_model_info_response_resolves_mode_through_deployment_model():
"""`mode` is derived from the same lookup, so an aliased embedding deployment
currently reports no mode at all; it must report `embedding`."""
@ -2169,9 +2255,7 @@ def test_create_model_info_response_resolves_mode_through_deployment_model():
]
)
response = create_model_info_response(
model_id="my-embeddings", provider="openai", llm_router=router
)
response = create_model_info_response(model_id="my-embeddings", provider="openai", llm_router=router)
finally:
litellm.model_cost.clear()
litellm.model_cost.update(saved_model_cost)
@ -2289,7 +2373,9 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp
with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()):
await proxy_logging_obj.post_call_failure_hook(
request_data={"metadata": {}},
original_exception=HTTPException(status_code=400, detail="Upstream passthrough request failed with status 400"),
original_exception=HTTPException(
status_code=400, detail="Upstream passthrough request failed with status 400"
),
user_api_key_dict=UserAPIKeyAuth(),
traceback_str=upstream_traceback,
)
@ -2299,6 +2385,130 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp
assert "REDACTED" in recorder.received_traceback
@pytest.mark.asyncio
@pytest.mark.parametrize("limiter_version", [1, 3])
@pytest.mark.parametrize("limit", ["rpm_limit", "max_parallel_requests"])
async def test_internal_realtime_observer_preserves_quota_and_custom_hooks(monkeypatch, limiter_version, limit):
import asyncio
from datetime import datetime
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
_request_stash,
get_request_stash,
)
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
observed = []
class Hook(CustomLogger):
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
observed.append(call_type)
return {**data, "extra_headers": {"x-hook": "required"}}
cache = DualCache()
limiter_type = _PROXY_MaxParallelRequestsHandler if limiter_version == 1 else _PROXY_MaxParallelRequestsHandler_v3
limiter = limiter_type(InternalUsageCache(dual_cache=cache))
proxy = ProxyLogging(UserApiKeyCache())
monkeypatch.setattr(litellm, "callbacks", [limiter, Hook()])
token = _request_stash.set(None)
try:
auth = UserAPIKeyAuth(api_key="observer-quota-test", **{limit: 1})
await proxy.pre_call_hook(
auth, {"model": "voice", "litellm_call_id": "signaling", "metadata": {}}, "arealtime_calls"
)
await asyncio.sleep(0)
initial_stash = get_request_stash()
result = await proxy.pre_call_hook(
auth,
{"model": "voice", "litellm_call_id": "observer", "metadata": {}},
"_arealtime",
internal_realtime_observer=True,
)
assert result["extra_headers"] == {"x-hook": "required"}
assert observed == ["arealtime_calls", "_arealtime"]
if limiter_version == 3:
assert get_request_stash() is initial_stash
assert initial_stash.owner_litellm_call_id == "signaling"
if limit == "max_parallel_requests":
await limiter.async_log_success_event(
{
"litellm_call_id": "signaling",
"litellm_params": {"metadata": {"user_api_key": auth.api_key, "user_api_key_model_max_budget": {}}},
},
litellm.ModelResponse(usage=litellm.Usage(total_tokens=0)),
datetime.now(),
datetime.now(),
)
if limiter_version == 3:
assert initial_stash.parallel_slot is None
await proxy.pre_call_hook(
auth, {"model": "voice", "litellm_call_id": "next", "metadata": {}}, "arealtime_calls"
)
if limiter_version == 1 and limit == "max_parallel_requests":
from litellm.proxy._types import InternalRequestOrigin
await asyncio.sleep(0)
observer_kwargs = {
"internal_request_origin": InternalRequestOrigin.REALTIME_OBSERVER,
"litellm_call_id": "observer",
"litellm_params": {"metadata": {"user_api_key": auth.api_key, "user_api_key_model_max_budget": {}}},
}
await limiter.async_log_success_event(
observer_kwargs,
litellm.ModelResponse(usage=litellm.Usage(total_tokens=17)),
datetime.now(),
datetime.now(),
)
current = await limiter.internal_usage_cache.async_get_cache(
key=f"{auth.api_key}::{datetime.now():%Y-%m-%d-%H-%M}::request_count", litellm_parent_otel_span=None
)
assert current["current_requests"] == 1
assert current["current_tpm"] == 17
with pytest.raises(HTTPException) as error:
await proxy.pre_call_hook(
auth,
{"model": "voice", "litellm_call_id": "forged", "metadata": {}, "internal_realtime_observer": True},
"_arealtime",
)
assert error.value.status_code == 429
finally:
_request_stash.reset(token)
@pytest.mark.asyncio
@pytest.mark.parametrize("scope", ["key", "user", "team", "end_user"])
async def test_internal_observer_missing_legacy_counter_only_adds_usage(scope):
from datetime import datetime
from litellm.proxy._types import InternalRequestOrigin
from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
from litellm.proxy.utils import InternalUsageCache
limiter = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache=DualCache()))
metadata = {"user_api_key": "expired-key", "user_api_key_model_max_budget": {}}
if scope in ("user", "team"):
metadata[f"user_api_key_{scope}_id"] = "expired-scope"
kwargs = {
"internal_request_origin": InternalRequestOrigin.REALTIME_OBSERVER,
"litellm_params": {"metadata": metadata},
**({"user": "expired-scope"} if scope == "end_user" else {}),
}
await limiter.async_log_success_event(
kwargs, litellm.ModelResponse(usage=litellm.Usage(total_tokens=23)), datetime.now(), datetime.now()
)
identity = "expired-key" if scope == "key" else "expired-scope"
current = await limiter.internal_usage_cache.async_get_cache(
key=f"{identity}::{datetime.now():%Y-%m-%d-%H-%M}::request_count", litellm_parent_otel_span=None
)
assert current == {"current_requests": 0, "current_tpm": 23, "current_rpm": 0}
class TestPrismaClientTokenAuthBehindThePool:
"""Behind the in-container pool the supervisor renews the writer's database
token and hands the workers a loopback URL with a static password, so the

View file

@ -18,16 +18,27 @@ async def test_trace_reader_projects_connection_and_parameters(recording_server:
recording_server.enqueue(ResponseSpec(body={"data": [{"trace_id": "trace-1"}]}))
reader_url: Final = recording_server.base_url.replace("http://", "http://reader:p%40ss%2Fword%25@")
storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, reader_url + "?database=wrong")
rows: Final = json.loads(await storage.query("trace_spans", {"trace_id": "trace-1"}))
response: Final = json.loads(
await storage.query("trace_spans", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""})
)
request: Final = recording_server.requests[0]
parameters: Final = parse_qs(urlsplit(request.path).query)
assert rows == {"data": [{"trace_id": "trace-1"}]}
parameters: Final = parse_qs(urlsplit(request.path).query, keep_blank_values=True)
assert response == {"data": [{"trace_id": "trace-1"}]}
assert b"o.TraceId = {trace_id:String}" in request.raw_body
assert parameters["database"] == ["trace_test"]
assert parameters["param_trace_id"] == ["trace-1"]
assert parameters["readonly"] == ["1"]
assert "user" not in parameters
assert "password" not in parameters
assert b"trace-1" not in request.raw_body
assert parameters == {
"database": ["trace_test"],
"param_trace_id": ["trace-1"],
"param_team_ids": ["[]"],
"param_api_key_hash": [""],
"param_trace_ref": [""],
"readonly": ["1"],
"default_format": ["JSON"],
"max_execution_time": ["10"],
"max_result_rows": ["1000"],
"result_overflow_mode": ["throw"],
"wait_end_of_query": ["1"],
}
assert request.headers["authorization"] == "Basic " + base64.b64encode(b"reader:p@ss/word%").decode()
@ -36,7 +47,7 @@ async def test_trace_reader_rejects_success_status_with_embedded_error(recording
recording_server.enqueue(ResponseSpec(body={"data": [], "exception": "query failed"}))
storage: Final = NativeTraceStorage("trace_test", recording_server.base_url, recording_server.base_url)
with pytest.raises(RuntimeError, match="invalid or failed JSON"):
await storage.query("trace_spans", {})
await storage.query("trace_spans", {"trace_id": "trace-1", "team_ids": [], "api_key_hash": "", "trace_ref": ""})
@pytest.mark.asyncio
@ -61,7 +72,9 @@ async def test_schema_binding_rejects_non_positive_retention() -> None:
@pytest.mark.asyncio
async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement(recording_server: RecordingServer) -> None:
async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement(
recording_server: RecordingServer,
) -> None:
recording_server.expected_requests = 2
recording_server.enqueue(ResponseSpec(body=""))
recording_server.enqueue(ResponseSpec(status=403, body="denied"))
@ -73,9 +86,10 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement
assert recording_server.requests[0].raw_body.startswith(b"CREATE DATABASE IF NOT EXISTS")
assert recording_server.requests[1].raw_body.startswith(b"CREATE TABLE IF NOT EXISTS")
assert "readonly" not in parse_qs(urlsplit(recording_server.requests[0].path).query)
assert recording_server.requests[0].headers["authorization"] == "Basic " + base64.b64encode(
b"writer:p@ss/word%"
).decode()
assert (
recording_server.requests[0].headers["authorization"]
== "Basic " + base64.b64encode(b"writer:p@ss/word%").decode()
)
@pytest.mark.asyncio
@ -87,11 +101,14 @@ async def test_insert_encodes_and_sends_rows(recording_server: RecordingServer)
after: Final = time.time_ns() // 1_000_000
request: Final = recording_server.requests[0]
row: Final = json.loads(gzip.decompress(request.raw_body))
assert type(row["EngineReceivedMs"]) is int
assert before <= row["EngineReceivedMs"] <= after
assert row == {
"Input": "hello",
"Timestamp": "1970-01-01T00:00:01.23456789Z",
"EngineReceivedMs": row["EngineReceivedMs"],
}
assert parse_qs(urlsplit(request.path).query)["query"] == ["INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow"]
assert parse_qs(urlsplit(request.path).query)["query"] == [
"INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow"
]
assert request.headers["content-encoding"] == "gzip"

View file

@ -10,7 +10,7 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v
import json
import os
import re
from typing import Any, Dict
from typing import Any, Dict, Final
from unittest.mock import MagicMock, patch
import httpx
@ -33,30 +33,10 @@ def _load_openapi_spec_dict() -> Dict[str, Any]:
return response.json()
except Exception as e: # pragma: no cover - defensive, env-dependent
pytest.skip(
f"Skipping Google Interactions OpenAPI compliance tests - "
f"unable to load spec from {OPENAPI_SPEC_URL}: {e}"
f"Skipping Google Interactions OpenAPI compliance tests - unable to load spec from {OPENAPI_SPEC_URL}: {e}"
)
def _model_create_request_schema(spec_dict: Dict[str, Any]) -> Dict[str, Any]:
schemas = spec_dict["components"]["schemas"]
create_path = next(path for path in spec_dict["paths"] if path.endswith("/interactions"))
body_schema = spec_dict["paths"][create_path]["post"]["requestBody"]["content"]["application/json"]["schema"]
variants = [schemas[option["$ref"].split("/")[-1]] for option in body_schema.get("oneOf", []) if "$ref" in option]
return next(variant for variant in variants if "model" in variant.get("properties", {}))
def _interaction_resource_path(spec_dict: Dict[str, Any], method: str) -> str | None:
return next(
(
path
for path, methods in spec_dict["paths"].items()
if re.search(r"/interactions/\{[^}]+\}$", path) and method in methods
),
None,
)
def _declared_type_value(variant_schema: Dict[str, Any]) -> Any:
"""The single `type` value a union variant pins, whether spelled as a const or a 1-item enum."""
type_property = variant_schema.get("properties", {}).get("type", {})
@ -64,6 +44,56 @@ def _declared_type_value(variant_schema: Dict[str, Any]) -> Any:
return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None)
def _resolve_local_ref(spec_dict: dict[str, Any], schema: dict[str, Any]) -> dict[str, Any]:
"""Resolve component references used by operations, schemas, and parameters."""
if "$ref" not in schema:
return schema
reference: Final = schema["$ref"]
assert reference.startswith("#/components/"), f"Expected a local component reference: {reference}"
category, name = reference.removeprefix("#/components/").split("/")
return spec_dict["components"][category][name.replace("~1", "/").replace("~0", "~")]
def _interaction_operation(
spec_dict: dict[str, Any], method: str, *, individual: bool = False
) -> tuple[str, dict[str, Any]]:
"""Match collection or item routes exactly, independent of placeholder names."""
pattern: Final = r"(?:/[^/]+)*/interactions" + (r"/(\{[^/{}]+\})" if individual else "")
matches: Final = tuple(
(path, path_item, match)
for path, path_item in spec_dict["paths"].items()
if (match := re.fullmatch(pattern, path)) and method in path_item
)
assert len(matches) == 1, f"Expected one {method.upper()} interactions endpoint, got {matches}"
path, path_item, match = matches[0]
operation: Final = path_item[method]
if individual:
parameter_name: Final = match.group(1)[1:-1]
parameters: Final = {
(parameter["name"], parameter["in"]): parameter
for raw_parameter in (*path_item.get("parameters", ()), *operation.get("parameters", ()))
for parameter in (_resolve_local_ref(spec_dict, raw_parameter),)
}
parameter: Final = parameters.get((parameter_name, "path"))
assert parameter is not None, f"{path} must declare its interaction ID path parameter"
assert parameter.get("required") is True, f"{path} must require its interaction ID"
parameter_schema: Final = _resolve_local_ref(spec_dict, parameter["schema"])
assert parameter_schema.get("type") == "string", f"{path} must accept a string interaction ID"
return path, operation
def _model_request_schema(spec_dict: dict[str, Any]) -> dict[str, Any]:
"""Find the model variant of the JSON body declared by the create operation."""
_, operation = _interaction_operation(spec_dict, "post")
request_body: Final = _resolve_local_ref(spec_dict, operation["requestBody"])
assert request_body.get("required") is True, "Creating an interaction must require a request body"
schema: Final = _resolve_local_ref(spec_dict, request_body["content"]["application/json"]["schema"])
variants: Final = tuple(_resolve_local_ref(spec_dict, variant) for variant in schema.get("oneOf", (schema,)))
model_variants: Final = tuple(variant for variant in variants if "model" in variant.get("properties", {}))
assert len(model_variants) == 1, f"Expected one model request variant, got {model_variants}"
return model_variants[0]
@pytest.fixture(scope="module")
def spec_dict() -> Dict[str, Any]:
"""Load raw spec dict for manual validation."""
@ -80,10 +110,14 @@ class TestRequestCompliance:
"""Tests that our request bodies match the OpenAPI spec."""
def test_create_model_interaction_request_schema(self, spec_dict):
schema = _model_create_request_schema(spec_dict)
"""Verify the model request schema declared by POST /interactions."""
schema = _model_request_schema(spec_dict)
assert "model" in schema["required"]
assert "input" in schema["properties"]
for field in ("model", "input"):
assert field in schema["properties"]
assert schema["properties"][field].get("readOnly") is not True
assert _resolve_local_ref(spec_dict, schema["properties"][field]).get("readOnly") is not True
# Check our supported optional fields exist in spec
our_optional_fields = [
@ -106,13 +140,8 @@ class TestRequestCompliance:
def test_input_types_match_spec(self, spec_dict):
"""Verify input field supports string, Content, Content[], Turn[]."""
schema = _model_create_request_schema(spec_dict)
input_schema = schema["properties"]["input"]
# The input property may be inline oneOf or a $ref to InteractionsInput
if "$ref" in input_schema:
ref_name = input_schema["$ref"].split("/")[-1]
input_schema = spec_dict["components"]["schemas"][ref_name]
schema = _model_request_schema(spec_dict)
input_schema = _resolve_local_ref(spec_dict, schema["properties"]["input"])
# Should be oneOf with multiple types
assert "oneOf" in input_schema
@ -143,22 +172,18 @@ class TestRequestCompliance:
discriminator = content_schema.get("discriminator")
if discriminator is not None:
assert (
discriminator.get("propertyName") == "type"
), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'"
assert discriminator.get("propertyName") == "type", (
f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'"
)
variant_names = [
option["$ref"].split("/")[-1]
for option in content_schema.get("oneOf", [])
if "$ref" in option
option["$ref"].split("/")[-1] for option in content_schema.get("oneOf", []) if "$ref" in option
]
assert variant_names, f"Content is not a union of named variants: {content_schema}"
mapping = (discriminator or {}).get("mapping") or {}
type_values = {
variant: mapping_value
for mapping_value, ref in mapping.items()
for variant in [ref.split("/")[-1]]
variant: mapping_value for mapping_value, ref in mapping.items() for variant in [ref.split("/")[-1]]
} or {
variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {}))
for variant in variant_names
@ -209,7 +234,9 @@ class TestRequestCompliance:
for option in spec_dict["components"]["schemas"]["Step"]["oneOf"]
if "$ref" in option
}
assert {"UserInputStep", "ModelOutputStep"} <= step_variants, f"Step union is missing role steps: {step_variants}"
assert {"UserInputStep", "ModelOutputStep"} <= step_variants, (
f"Step union is missing role steps: {step_variants}"
)
for step_name, type_value in [("UserInputStep", "user_input"), ("ModelOutputStep", "model_output")]:
step_schema = spec_dict["components"]["schemas"][step_name]
@ -279,9 +306,7 @@ class TestResponseCompliance:
expected_fields = ["total_input_tokens", "total_output_tokens", "total_tokens"]
for field in expected_fields:
assert (
field in usage_schema["properties"]
), f"Usage field '{field}' not in spec"
assert field in usage_schema["properties"], f"Usage field '{field}' not in spec"
print(f"✓ Usage field '{field}' exists")
@ -300,9 +325,7 @@ class TestToolsCompliance:
"""Verify FunctionDeclaration schema for function tools."""
if "FunctionDeclaration" in spec_dict["components"]["schemas"]:
func_schema = spec_dict["components"]["schemas"]["FunctionDeclaration"]
assert "name" in func_schema.get(
"properties", {}
) or "name" in func_schema.get("required", [])
assert "name" in func_schema.get("properties", {}) or "name" in func_schema.get("required", [])
print("✓ FunctionDeclaration schema found")
else:
print("⚠ FunctionDeclaration schema not found (may be nested)")
@ -313,33 +336,94 @@ class TestEndpointCompliance:
def test_create_endpoint_exists(self, spec_dict):
"""Verify POST /interactions endpoint exists."""
paths = spec_dict["paths"]
# Find the create interactions endpoint
create_path = None
for path, methods in paths.items():
if "interactions" in path and "post" in methods:
create_path = path
break
assert create_path is not None, "POST /interactions endpoint not found"
create_path, _ = _interaction_operation(spec_dict, "post")
print(f"✓ Create endpoint: POST {create_path}")
def test_get_endpoint_exists(self, spec_dict):
"""Verify GET /interactions/{id} endpoint exists."""
get_path = _interaction_resource_path(spec_dict, "get")
assert get_path is not None, "GET /interactions/{id} endpoint not found"
get_path, _ = _interaction_operation(spec_dict, "get", individual=True)
print(f"✓ Get endpoint: GET {get_path}")
def test_delete_endpoint_exists(self, spec_dict):
"""Verify DELETE /interactions/{id} endpoint exists."""
delete_path = _interaction_resource_path(spec_dict, "delete")
assert delete_path is not None, "DELETE /interactions/{id} endpoint not found"
delete_path, _ = _interaction_operation(spec_dict, "delete", individual=True)
print(f"✓ Delete endpoint: DELETE {delete_path}")
class TestOperationResolution:
"""Keep structural resolution strict without depending on generated names."""
@pytest.mark.parametrize("as_union", [False, True])
def test_model_schema_comes_from_create_operation(self, as_union):
model_schema: Final = {"properties": {"model": {"type": "string"}}, "required": ["model"]}
reference: Final = {"$ref": "#/components/schemas/RenamedModelRequest"}
body_schema: Final = (
{"oneOf": [{"properties": {"agent": {"type": "string"}}}, reference]} if as_union else reference
)
spec: Final = {
"paths": {
"/{version}/interactions": {
"post": {
"requestBody": {"required": True, "content": {"application/json": {"schema": body_schema}}}
}
}
},
"components": {
"schemas": {"RenamedModelRequest": model_schema, "CreateModelInteractionParams": {"properties": {}}}
},
}
assert _model_request_schema(spec) is model_schema
@pytest.mark.parametrize("method,shared", [("get", False), ("delete", True)])
def test_item_route_accepts_a_renamed_declared_identifier(self, method, shared):
parameter: Final = {"name": "renamedId", "in": "path", "required": True, "schema": {"type": "string"}}
parameters: Final = [{"$ref": "#/components/parameters/Identifier"}]
operation: Final = {"parameters": [] if shared else parameters}
path: Final = "/{version}/interactions/{renamedId}"
spec: Final = {
"paths": {path: {"parameters": parameters if shared else [], method: operation}},
"components": {"parameters": {"Identifier": parameter}},
}
assert _interaction_operation(spec, method, individual=True) == (path, operation)
@pytest.mark.parametrize(
"path,parameter,error",
[
(
"/interactions/{id}/cancel",
{"required": True, "type": "string"},
"Expected one GET interactions endpoint",
),
(
"/other_interactions/{id}",
{"required": True, "type": "string"},
"Expected one GET interactions endpoint",
),
("/interactions/{id}", {"required": False, "type": "string"}, "must require its interaction ID"),
("/interactions/{id}", {"required": True, "type": "integer"}, "must accept a string interaction ID"),
],
)
def test_item_route_rejects_incompatible_contracts(self, path, parameter, error):
spec: Final = {
"paths": {
path: {
"get": {
"parameters": [
{
"name": "id",
"in": "path",
"required": parameter["required"],
"schema": {"type": parameter["type"]},
}
]
}
}
}
}
with pytest.raises(AssertionError, match=error):
_interaction_operation(spec, "get", individual=True)
if __name__ == "__main__":
# Quick manual test
import httpx
@ -356,6 +440,4 @@ if __name__ == "__main__":
if method in ["get", "post", "delete", "put", "patch"]:
print(f" {method.upper()} {path}")
print(
f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}..."
)
print(f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}...")

View file

@ -3093,6 +3093,39 @@ def test_image_response_input_image_tokens_priced_at_image_rate(details_as_dict)
assert round(cost, 12) == round(expected, 12)
@pytest.mark.parametrize("excess", ["text", "image"])
def test_image_response_cached_modality_counts_cannot_exceed_inputs(excess):
"""
A cached_tokens_details entry larger than the matching input modality count
would turn cache reads into negative savings; the calculation must reject
the inconsistent usage instead of pricing it.
"""
from unittest.mock import patch
from litellm.litellm_core_utils.llm_cost_calc.utils import (
calculate_image_response_cost_from_usage,
)
from litellm.types.utils import Usage
cached: dict = (
{"text_tokens": 11, "image_tokens": 0} if excess == "text" else {"text_tokens": 0, "image_tokens": 101}
)
image_response = ImageResponse(data=[ImageObject(b64_json="x")])
image_response.usage = Usage(
prompt_tokens=0,
completion_tokens=0,
total_tokens=212,
input_tokens=110,
input_tokens_details={"text_tokens": 10, "image_tokens": 100, "cached_tokens_details": cached},
output_tokens=102,
output_tokens_details={"image_tokens": 102, "text_tokens": 0},
)
with pytest.raises(ValueError, match="Image cached token counts exceed their input modality counts"):
calculate_image_response_cost_from_usage(
model="gpt-image-2",
image_response=image_response,
custom_llm_provider="openai",
)
GEMINI_DAY0_LAUNCH_PRICING = [
("gemini-3.6-flash", 7.5e-07, 3.75e-06, 7.5e-08),
("gemini/gemini-3.6-flash", 7.5e-07, 3.75e-06, 7.5e-08),

View file

@ -57,10 +57,11 @@ def test_base_model_is_additive_not_replacement():
assert real_only <= combined
@pytest.mark.usefixtures("local_model_cost_map")
def test_base_model_adds_capabilities_the_real_model_lacks():
"""Regression for #27717 (the behavior the union must preserve).
``gemini-exp-9999`` isn't in the cost map so it advertises no reasoning support,
``gemini-exp-9999`` isn't in the bundled cost map, so it advertises no reasoning support,
but the registered ``gemini-3.1-pro-preview`` base_model does. The hint must add
``reasoning_effort``/``thinking`` without the call erroring."""
real_only = set(get_supported_openai_params(model="gemini-exp-9999", custom_llm_provider="gemini"))

View file

@ -3399,6 +3399,82 @@ async def test_refused_session_does_not_stamp_the_reservation_ownership_marker()
assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in session.logging.model_call_details
def test_live_terminal_usage_survives_filtered_event_logging(monkeypatch):
from litellm.cost_calculator import RealtimeAPITokenUsageProcessor
def terminal():
return {"type": "session.closed", "usage": {"audio_duration_ms": 4000, "backend_model_usage": []}}
monkeypatch.setattr(litellm, "logged_real_time_event_types", [])
stream = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
event = {**terminal(), "private_transcript": "Do not retain this text"}
stream.store_message(event)
assert stream.messages == [terminal()]
usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(stream.messages)
assert usage.total_tokens == 0
@pytest.mark.asyncio
async def test_live_attachment_does_not_dispatch_duplicate_usage():
worker = MagicMock()
logger = MagicMock()
stream = RealTimeStreaming(MagicMock(), MagicMock(), logger, logging_worker=worker, account_usage=False)
stream.store_message({"type": "session.closed", "usage": {"audio_duration_ms": 4000}})
await stream.log_messages()
worker.ensure_initialized_and_enqueue.assert_not_called()
logger.dispatch_success_handlers.assert_not_called()
@pytest.mark.asyncio
async def test_log_messages_flush_awaits_dispatch_instead_of_enqueueing():
worker = MagicMock()
logger = MagicMock()
logger.model_call_details = {}
logger.dispatch_success_handlers = AsyncMock()
stream = RealTimeStreaming(MagicMock(), MagicMock(), logger, logging_worker=worker)
stream.store_message({"type": "session.created"})
await stream.log_messages(wait_for_dispatch=True)
logger.dispatch_success_handlers.assert_awaited_once_with(stream.messages, prefer_async_handlers=True)
worker.ensure_initialized_and_enqueue.assert_not_called()
assert logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] is True
@pytest.mark.asyncio
@pytest.mark.parametrize("account_usage", [False, True])
async def test_attachment_cleanup_runs_in_owning_context_only(account_usage):
from litellm.litellm_core_utils.realtime_streaming import realtime_attachment_cleanup
contexts = []
async def one(name):
task = asyncio.current_task()
callback = AsyncMock(side_effect=lambda: contexts.append((name, asyncio.current_task() is task)))
token = realtime_attachment_cleanup.set(callback)
try:
websocket = MagicMock()
websocket.receive_text = AsyncMock(side_effect=RuntimeError("disconnected"))
backend = MagicMock()
async def recv(**kwargs):
await asyncio.Event().wait()
backend.recv = recv
stream = RealTimeStreaming(websocket, backend, MagicMock(), account_usage=account_usage)
await stream.bidirectional_forward()
if account_usage:
callback.assert_not_awaited()
else:
callback.assert_awaited_once()
finally:
realtime_attachment_cleanup.reset(token)
await asyncio.gather(one("first"), one("second"))
assert sorted(contexts) == ([] if account_usage else [("first", True), ("second", True)])
assert realtime_attachment_cleanup.get() is None
@pytest.mark.asyncio
async def test_refused_session_stamps_the_failure_ownership_marker():
"""LIT-6463: the enqueued failure callback releases the key's max_parallel_requests
@ -3550,3 +3626,81 @@ async def test_provider_bytes_are_sent_raw_after_pacing():
assert [call.args[0] for call in backend_ws.send.await_args_list] == [b"\x00\x01", '{"type":"endStream"}']
provider_config.pace_backend_send.assert_awaited_once_with(b"\x00\x01")
def test_public_live_accounting_survives_filtered_logging(monkeypatch):
monkeypatch.setattr(litellm, "logged_real_time_event_types", [])
stream = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
events = [
{"type": "session.usage.updated", "usage": {"seconds": 15}},
{
"type": "response.event",
"event": {
"type": "response.completed",
"response": {
"id": "resp_one",
"model": "gpt-backend",
"output": [],
"usage": {"total_tokens": 12},
},
},
},
{"type": "session.closed", "usage": {"seconds": 30}},
]
for event in events:
stream.store_message({**event, "private_transcript": "do not retain"})
stream.store_message(
{"type": "response.event", "event": {"type": "response.output_text.delta", "delta": "private"}}
)
assert stream.messages == events
@pytest.mark.parametrize("terminal", ["response.completed", "response.incomplete", "response.failed"])
@pytest.mark.parametrize("allowed", [[], ["response.event"], "*"])
def test_live_terminal_logging_filters_content_and_preserves_accounting(
monkeypatch: pytest.MonkeyPatch, terminal: str, allowed: list[str] | str
) -> None:
from litellm.cost_calculator import _live_backend_responses
monkeypatch.setattr(litellm, "logged_real_time_event_types", allowed)
stream = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())
response = {
"id": "resp_private",
"created_at": 1,
"model": "gpt-backend",
"output": [{"type": "message", "content": [{"type": "output_text", "text": "private answer"}]}],
"instructions": "private instructions",
"metadata": {"private": "metadata"},
"usage": {"input_tokens": 20, "output_tokens": 10, "total_tokens": 30},
}
event = {"type": "response.event", "event": {"type": terminal, "response": response}}
stream.store_message(event)
stored = stream.messages[0]["event"]["response"]
if allowed:
assert stored == response
else:
assert stored == {
key: value for key, value in response.items() if key not in ("output", "instructions", "metadata")
} | {"output": []}
measured = _live_backend_responses(stream.messages)
assert len(measured) == 1
assert measured[0].id == "resp_private"
assert measured[0].model == "gpt-backend"
assert measured[0].usage.total_tokens == 30
assert response["instructions"] == "private instructions"
assert response["output"][0]["content"][0]["text"] == "private answer"
@pytest.mark.parametrize("account_usage,expected", [(True, 1), (False, 0)])
def test_live_initialization_is_retained_only_by_accounting_owner(account_usage, expected):
stream = RealTimeStreaming(
MagicMock(),
MagicMock(),
MagicMock(),
account_usage=account_usage,
live_initialization_seconds=15,
)
assert len(stream.messages) == expected
if account_usage:
assert stream.messages == [{"type": "litellm.live.initialization", "usage": {"seconds": 15}}]

View file

@ -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)

View file

@ -30,6 +30,29 @@ def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> Generator[None, Non
class TestChatGPTResponsesAPITransformation:
def test_guardian_preserves_strict_output_schema(self):
text = {
"format": {
"type": "json_schema",
"name": "review",
"strict": True,
"schema": {
"type": "object",
"properties": {"allowed": {"type": "boolean"}},
"required": ["allowed"],
"additionalProperties": False,
},
}
}
request = ChatGPTResponsesAPIConfig().transform_responses_api_request(
model="codex-auto-review",
input="Review the command pwd",
response_api_optional_request_params={"text": text},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert request["text"] == text
@pytest.mark.parametrize(
"model_name",
[

View file

@ -0,0 +1,63 @@
import hashlib
import time
import httpx
import pytest
from litellm.llms.chatgpt.codex import CodexRealtimeCall, build_sideband_request, parse_call_response
def test_encrypted_call_preserves_repeated_gateway_query(monkeypatch):
from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_call
monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-repeated-query")
authorization = "Bearer test-owner"
call = CodexRealtimeCall(
call_id="rtc_repeated",
model="gpt-live-1-codex",
alias="voice",
owner=hashlib.sha256(authorization.encode()).hexdigest(),
expires_at=time.time() + 60,
extra_query={"tag": ["alpha +/&", "beta"], "gateway": "tenant"},
)
restored = decode_call(encode_call(call), authorization)
assert restored.extra_query == {"tag": ("alpha +/&", "beta"), "gateway": "tenant"}
assert build_sideband_request(restored)["extra_query"] == restored.extra_query
@pytest.mark.parametrize("location", ["", "/v1/realtime/calls/foreign-id"])
def test_signaling_rejects_invalid_upstream_call_id(location):
response = httpx.Response(201, headers={"Location": location},
extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}})
with pytest.raises(ValueError, match="String should match pattern"):
parse_call_response(response, "voice", "owner", 1000)
@pytest.mark.parametrize("extra_query", [None, {"gateway_token": "opaque +/& value"}])
def test_signaling_preserves_selected_model_for_sideband(extra_query):
response = httpx.Response(
201,
headers={"Location": "/v1/realtime/calls/rtc_provider"},
extensions={
"chatgpt_realtime": {
"model": "gpt-live-1-codex",
"api_base": "https://voice.example/codex",
"extra_headers": {"x-gateway-route": "voice"},
**({"extra_query": extra_query} if extra_query is not None else {}),
}
},
)
call = parse_call_response(response, "voice", "owner", 1000)
request = build_sideband_request(CodexRealtimeCall.model_validate_json(call.model_dump_json(exclude_none=True)))
assert request["api_base"] == "https://voice.example/codex"
assert request["model"] == "chatgpt/gpt-live-1-codex"
assert request["chatgpt_realtime_call_id"] == "rtc_provider"
assert request["query_params"] == {"model": "gpt-live-1-codex"}
assert request["extra_headers"] == {"x-gateway-route": "voice"}
assert request["extra_query"] == extra_query
def test_signaling_requires_chatgpt_routing_extension():
response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_unrouted"})
with pytest.raises(ValueError, match="Direct call signaling requires a ChatGPT deployment"):
parse_call_response(response, "voice", "owner", 1000)

View file

@ -0,0 +1,305 @@
import base64
import json
from typing import Final
import httpx
import pytest
import litellm
from litellm.llms.chatgpt.images import ChatGPTImageEditConfig, ChatGPTImageGenerationConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.types.llms.openai import ImageGenerationRequestQuality
from litellm.types.router import GenericLiteLLMParams
@pytest.mark.parametrize(
"model,quality",
[
("gpt-image-2", ImageGenerationRequestQuality.AUTO),
("gpt-image-2.5-flare", ImageGenerationRequestQuality.XHIGH),
("gpt-image-2.5-flare", ImageGenerationRequestQuality.MAX),
("gpt-image-2.5-sunburst", ImageGenerationRequestQuality.XHIGH),
("gpt-image-2.5-sunburst", ImageGenerationRequestQuality.MAX),
],
)
@pytest.mark.parametrize("editing", [False, True])
def test_image_25_transmits_model_quality_and_transparency(model, quality, editing, chatgpt_tokens):
expected: Final = {
"model": model,
"prompt": "a red circle with transparent surroundings",
"quality": quality.value,
"background": "transparent",
"size": "2048x2048",
**({"images": [{"image_url": "data:image/png;base64,aGVsbG8="}]} if editing else {}),
}
def respond(request):
assert str(request.url) == "https://chatgpt.com/backend-api/codex/images/" + (
"edits" if editing else "generations"
)
assert request.headers["content-type"] == "application/json"
assert json.loads(request.content) == expected
return httpx.Response(
200,
json={"created": 1, "data": [{"b64_json": "aGVsbG8="}], "quality": quality.value},
)
client: Final = HTTPHandler()
client.client = httpx.Client(transport=httpx.MockTransport(respond))
operation: Final = litellm.image_edit if editing else litellm.image_generation
try:
response: Final = operation(
**{**expected, "model": "chatgpt/" + model, "quality": quality},
client=client,
chatgpt_token_dir=chatgpt_tokens,
)
assert response.data[0].b64_json == "aGVsbG8="
assert response.quality == quality.value
finally:
client.client.close()
@pytest.mark.parametrize("model", ["gpt-image-2", "gpt-image-2.5-flare", "gpt-image-2.5-sunburst"])
def test_json_edit_preserves_provider_params_and_extra_body_precedence(model, chatgpt_tokens):
references: Final = [{"image_url": "data:image/png;base64,aGVsbG8="}]
def respond(request):
assert request.headers["content-type"] == "application/json"
assert json.loads(request.content) == {
"model": model,
"prompt": "red circle",
"images": references,
"seed": 7,
"provider_options": {"steps": 30, "enabled": True},
"output_compression": 90,
}
return httpx.Response(200, json={"created": 1, "data": [{"b64_json": "aGVsbG8="}]})
with httpx.Client(transport=httpx.MockTransport(respond)) as http_client:
response: Final = litellm.image_edit(
model="chatgpt/" + model,
prompt="red circle",
images=references,
client=HTTPHandler(client=http_client),
chatgpt_token_dir=chatgpt_tokens,
seed=42,
output_compression=90,
extra_body={"seed": 7, "provider_options": {"steps": 30, "enabled": True}},
)
assert response.data[0].b64_json == "aGVsbG8="
@pytest.mark.parametrize("api_base", [None, "https://image-gateway.test"])
def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens, api_base):
requests = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"created": 1, "data": [{"b64_json": "aGVsbG8="}]})
client = HTTPHandler()
client.client = httpx.Client(transport=httpx.MockTransport(respond))
result = litellm.image_generation(
model="chatgpt/gpt-image-2",
prompt="blue circle",
api_base=api_base,
client=client,
quality="auto",
size="auto",
background="auto",
extra_headers={"x-gateway-route": "images", "aUtHoRiZaTiOn": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"},
)
assert requests[0].headers["x-gateway-route"] == "images"
assert result.data[0].b64_json == "aGVsbG8="
assert str(requests[0].url) == (api_base or "https://chatgpt.com/backend-api/codex") + "/images/generations"
assert requests[0].headers["authorization"] == "Bearer test-token-" + "default"
assert requests[0].headers["chatgpt-account-id"] == "test-account-" + "default"
assert b'"model":"gpt-image-2"' in requests[0].content
def test_codex_json_edit_survives_sdk_dispatch(chatgpt_tokens):
requests = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"created": 1, "data": [{"b64_json": "aGVsbG8="}]})
client = HTTPHandler()
client.client = httpx.Client(transport=httpx.MockTransport(respond))
references = [{"image_url": "data:image/png;base64,aGVsbG8="}]
result = litellm.image_edit(
model="chatgpt/gpt-image-2",
prompt="red circle",
extra_headers={"x-gateway-route": "images", "authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"},
images=references,
client=client,
quality="auto",
size="auto",
)
assert result.data[0].b64_json == "aGVsbG8="
assert str(requests[0].url) == "https://chatgpt.com/backend-api/codex/images/edits"
import json
assert json.loads(requests[0].content)["images"] == references
assert requests[0].headers["authorization"] == "Bearer test-token-default"
assert requests[0].headers["chatgpt-account-id"] == "test-account-default"
assert requests[0].headers["x-gateway-route"] == "images"
@pytest.mark.parametrize(
"references", [[], [{"image_url": "file:///etc/passwd"}], [{}], [{"image_url": "https://example.com/a.png"}] * 6]
)
def test_edit_rejects_invalid_references(references):
with pytest.raises(ValueError, match=r"images must contain|validation error"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2", "edit", None, {}, GenericLiteLLMParams(images=references), {}
)
def test_edit_converts_multipart_image_bytes():
data, files = ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2", "edit", b"example", {}, GenericLiteLLMParams(), {}
)
assert not files
assert base64.b64decode(data["images"][0]["image_url"].split(",", 1)[1]) == b"example"
def test_image_auth_does_not_accept_inbound_override(chatgpt_tokens):
headers = ChatGPTImageGenerationConfig().validate_environment(
{"authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, "gpt-image-2", [], {}, {"chatgpt_token_dir": chatgpt_tokens}
)
assert httpx.Headers(headers)["authorization"] == "Bearer test-token-default"
assert httpx.Headers(headers)["chatgpt-account-id"] == "test-account-default"
@pytest.mark.asyncio
async def test_async_codex_edit_without_multipart_image(chatgpt_tokens):
requests = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"created": 1, "data": [{"b64_json": "aGVsbG8="}]})
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
response = await litellm.aimage_edit(
model="chatgpt/gpt-image-2",
prompt="red circle",
extra_headers={"x-gateway-route": "images", "authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"},
client=client,
images=[{"image_url": "data:image/png;base64,aGVsbG8="}],
chatgpt_auth_profile="account3",
)
assert response.data[0].b64_json == "aGVsbG8="
assert str(requests[0].url).endswith("/codex/images/edits")
assert requests[0].headers["content-type"] == "application/json"
await client.client.aclose()
assert requests[0].headers["authorization"] == "Bearer test-token-default"
assert requests[0].headers["chatgpt-account-id"] == "test-account-default"
assert requests[0].headers["x-gateway-route"] == "images"
@pytest.mark.parametrize(
"references",
[None, [{"image_url": "data:image/png;base64,aGVsbG8="}]],
ids=["uploaded-image", "reference-images"],
)
def test_edit_keeps_the_authenticated_model_over_passthrough_fields(tmp_path, references):
image = None
if references is None:
image = tmp_path / "reference.png"
image.write_bytes(b"reference image bytes")
data, _ = ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2",
"edit",
image,
{"model": "gpt-image-2.5-flare", "size": "1024x1024"},
GenericLiteLLMParams(images=references),
{},
)
assert data["model"] == "gpt-image-2"
assert data["size"] == "1024x1024"
@pytest.mark.parametrize("as_tuple", [False, True])
def test_edit_accepts_filesystem_path(tmp_path, as_tuple):
image = tmp_path / "reference.png"
image.write_bytes(b"reference image bytes")
data, files = ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2", "edit", ("reference.png", image, "image/png") if as_tuple else image,
{}, GenericLiteLLMParams(), {}
)
assert not files
assert data["images"] == ({"image_url": "data:image/png;base64," + base64.b64encode(image.read_bytes()).decode()},)
@pytest.mark.parametrize("env_name", ["CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"])
@pytest.mark.parametrize("api_base", [None, "https://deployment.example/codex"])
def test_image_routes_use_configured_gateway(monkeypatch, env_name, api_base, tmp_path):
token_path = tmp_path / "unavailable-token-directory"
token_path.write_text("not a directory")
monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(token_path))
monkeypatch.delenv("CHATGPT_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False)
monkeypatch.setenv(env_name, "https://gateway.example/codex/")
expected = api_base or "https://gateway.example/codex"
assert ChatGPTImageGenerationConfig().get_complete_url(api_base, None, "gpt-image-2", {}, {}) == (
expected + "/images/generations"
)
assert ChatGPTImageEditConfig().get_complete_url("gpt-image-2", api_base, {}) == expected + "/images/edits"
@pytest.mark.parametrize(
"reference",
[
b"GIF89a" + b"\x00" * 32,
("reference.gif", b"hello", "image/gif"),
],
ids=["detected-gif", "declared-gif"],
)
def test_edit_rejects_non_bitmap_reference_content_type(reference):
with pytest.raises(ValueError, match="Reference images must be PNG, JPEG, or WEBP"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2", "edit", reference, {}, GenericLiteLLMParams(), {}
)
def test_edit_rejects_mask_before_any_provider_call():
with pytest.raises(ValueError, match="ChatGPT image editing does not support masks"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2",
"edit",
"data:image/png;base64,aGVsbG8=",
{"mask": "data:image/png;base64,aGVsbG8="},
GenericLiteLLMParams(),
{},
)
def test_edit_rejects_image_and_images_together():
with pytest.raises(ValueError, match="Specify only one of image or images"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2",
"edit",
"data:image/png;base64,aGVsbG8=",
{},
GenericLiteLLMParams(images=[{"image_url": "data:image/png;base64,aGVsbG8="}]),
{},
)
@pytest.mark.parametrize(
"images",
[
[],
["data:image/png;base64,aGVsbG8="] * 6,
],
ids=["zero", "six"],
)
def test_edit_enforces_one_to_five_reference_images(images):
with pytest.raises(ValueError, match="images must contain between 1 and 5 reference images"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2", "edit", images, {}, GenericLiteLLMParams(), {}
)

View file

@ -0,0 +1,261 @@
import json
from contextlib import asynccontextmanager
from unittest.mock import AsyncMock
import httpx
import pytest
from litellm.llms.chatgpt.live import LiveDeployment, LiveTransport, live_session_path
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@asynccontextmanager
async def live_handler(respond):
handler = AsyncHTTPHandler(transport=httpx.MockTransport(respond), follow_redirects=False)
try:
yield handler
finally:
await handler.close()
@pytest.mark.asyncio
@pytest.mark.parametrize("provider", ["chatgpt", "openai"])
@pytest.mark.parametrize("status", [201, 403, 429, 503])
async def test_live_request_preserves_payload_status_and_selected_credentials(provider, status, chatgpt_tokens):
payload = {
"session": {"model": "deployment-model", "tools": [{"type": "function", "name": "lookup"}]},
"transport": {"type": "webrtc", "sdp": "v=0\r\n"},
"future_option": {"nested": [True, None, 3]},
}
def respond(request):
assert request.url.path == "/custom/v1/live/sessions"
assert request.headers["authorization"] == (
"Bearer test-token-default" if provider == "chatgpt" else "Bearer deployment-key"
)
assert request.headers.get("chatgpt-account-id") == ("test-account-default" if provider == "chatgpt" else None)
assert request.headers["x-gateway"] == "configured"
assert request.headers["openai-beta"] == "feature=v1"
assert "cookie" not in request.headers
assert json.loads(request.content) == payload
assert request.url.params.get_list("tag") == ["a +/&", "b"]
assert request.url.params["gateway"] == "trusted"
assert request.url.params["cursor"] == "opaque +/&"
assert not {"model", "call_id", "session_id", "api_key"}.intersection(request.url.params)
return httpx.Response(status, json={"result": "upstream"}, headers={"x-request-id": "provider-id"})
async with live_handler(respond) as handler:
transport = LiveTransport(
LiveDeployment(
model="deployment-model",
provider=provider,
api_key="deployment-key",
api_base="https://gateway.example/custom/v1/?gateway=base",
extra_headers={"x-gateway": "configured", "Authorization": "bad", "ChatGPT-Account-Id": "bad"},
extra_query={"gateway": "trusted", "tag": ("a +/&", "b"), "model": "bad", "session_id": "bad"},
),
{"Authorization": "Bearer proxy-key", "Cookie": "private", "OpenAI-Beta": "feature=v1"},
http_handler=handler,
)
response = await transport.request(
"POST",
"live/sessions",
payload,
{"gateway": "untrusted", "cursor": "opaque +/&", "call_id": "bad", "api_key": "bad"},
)
assert response.status_code == status
assert response.json() == {"result": "upstream"}
assert response.headers["x-request-id"] == "provider-id"
assert not handler.client.is_closed
@pytest.mark.asyncio
@pytest.mark.parametrize("operation", ["fork", "accept", "reject", "refer", "hangup", "content"])
@pytest.mark.parametrize("status", [204, 404, 503])
async def test_live_all_http_operations(operation, status):
def respond(request):
assert request.url.path == f"/v1/live/sessions/sess_new-ID/{operation}"
assert request.method == ("GET" if operation == "content" else "POST")
assert request.url.params["output_format"] == "json"
return httpx.Response(status)
async with live_handler(respond) as handler:
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler)
response = await transport.request(
"GET" if operation == "content" else "POST",
live_session_path("sess_new-ID", operation),
query={"output_format": "json"},
)
assert response.status_code == status
@pytest.mark.asyncio
async def test_live_request_uses_handler_methods():
handler = AsyncMock(spec=AsyncHTTPHandler)
handler.get.return_value = httpx.Response(404)
handler.post.return_value = httpx.Response(503)
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler)
get_response = await transport.request("GET", live_session_path("sess_1", "content"))
post_response = await transport.request("POST", "live/sessions", {"transport": {"type": "webrtc"}})
assert get_response.status_code == 404
assert post_response.status_code == 503
handler.get.assert_awaited_once_with(
"https://api.openai.com/v1/live/sessions/sess_1/content",
headers={"authorization": "Bearer key", "content-type": "application/json"},
timeout=60,
follow_redirects=False,
)
handler.post.assert_awaited_once_with(
"https://api.openai.com/v1/live/sessions",
headers={"authorization": "Bearer key", "content-type": "application/json"},
json={"transport": {"type": "webrtc"}},
timeout=60,
)
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["live/sessions", "live/sessions/sess_1/attach", "live/sessions/sess_1/fork"])
async def test_live_websocket_paths_bounds_and_auth(path, chatgpt_tokens):
from websockets.asyncio.server import serve
async def observe(connection):
assert connection.request.headers["Authorization"] == "Bearer test-token-default"
await connection.send(connection.request.path)
async with serve(observe, "127.0.0.1", 0) as server:
port = server.sockets[0].getsockname()[1]
transport = LiveTransport(
LiveDeployment("model", api_base=f"http://127.0.0.1:{port}/v1", extra_query={"route": "a+&b"}),
{"Authorization": "Bearer proxy-key"},
)
connection = await transport.connect(path, {"checkpoint": "opaque+value"})
try:
received = await connection.recv()
url = httpx.URL(f"http://127.0.0.1{received}")
assert url.path == f"/v1/{path}"
assert url.params["route"] == "a+&b"
assert url.params["checkpoint"] == "opaque+value"
finally:
await connection.close()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"path",
[
"https://evil.example/live/sessions",
"//evil.example/live/sessions",
"live/sessions/../accept",
"live/sessions/sess%2Fbad/accept",
"live/sessions/sess%5Cbad/accept",
"live/sessions/sess%252Fbad/accept",
"live/sessions/%252e%252e/accept",
"live/sessions/%2E%2E/accept",
"live/sessions/sess%00bad/accept",
"live/sessions/sess%0Abad/accept",
"live/sessions/sess.foo?query/accept",
"live/sessions/sess_1/accept?url=https://evil.example",
"live/sessions/sess_1/accept#fragment",
"live/sessions/sess_1/accept\n",
],
)
async def test_live_rejects_noncanonical_paths_before_network(path):
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {})
with pytest.raises(ValueError, match=r"(?:Invalid|Noncanonical) Live"):
await transport.request("POST", path)
with pytest.raises(ValueError, match=r"(?:Invalid|Noncanonical) Live"):
await transport.connect(path)
@pytest.mark.parametrize(
"session_id",
["../x", "sess/x", "sess%2Fx", "sess\\x", "sess%255cx", ".", "..", "%252e%252e", "", "sess\n", "sess\x00"],
)
def test_live_session_ids_cannot_inject_path_or_query(session_id):
with pytest.raises(ValueError, match="Invalid Live session ID"):
live_session_path(session_id, "content")
@pytest.mark.asyncio
@pytest.mark.parametrize("session_id", ["sess.foo", "sess-\u00f1\u4e2d", "sess?x#y", "sess 50%", "x" * 1024])
async def test_live_preserves_opaque_session_ids(session_id):
from urllib.parse import quote
def respond(request):
assert request.url.raw_path == f"/v1/live/sessions/{quote(session_id, safe='')}/content".encode()
assert request.url.params == httpx.QueryParams()
return httpx.Response(200, json={"session_id": session_id})
async with live_handler(respond) as handler:
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler)
response = await transport.request("GET", live_session_path(session_id, "content"))
assert response.json()["session_id"] == session_id
@pytest.mark.asyncio
async def test_live_does_not_redirect_credentials():
def respond(request):
assert request.url.host == "api.openai.com"
return httpx.Response(307, headers={"location": "https://elsewhere.example/collect"})
async with live_handler(respond) as handler:
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler)
response = await transport.request("POST", "live/sessions", {})
assert response.status_code == 307
@pytest.mark.asyncio
async def test_live_websocket_does_not_redirect_credentials():
from websockets.asyncio.server import serve
from websockets.datastructures import Headers
from websockets.exceptions import InvalidStatus
from websockets.http11 import Response
async def unused(connection):
pytest.fail("Redirected WebSocket must never open")
def redirect(connection, request):
assert request.headers["authorization"] == "Bearer deployment-key"
return Response(307, "Temporary Redirect", Headers({"Location": "/elsewhere"}))
async with serve(unused, "127.0.0.1", 0, process_request=redirect) as server:
port = server.sockets[0].getsockname()[1]
transport = LiveTransport(
LiveDeployment(
"model", provider="openai", api_key="deployment-key", api_base=f"http://127.0.0.1:{port}/v1"
),
{},
)
with pytest.raises(InvalidStatus) as failure:
await transport.connect("live/sessions")
assert failure.value.response.status_code == 307
@pytest.mark.asyncio
@pytest.mark.parametrize(
"api_base",
[
"ftp://gateway.example/v1",
"https://user:secret@gateway.example/v1",
"https://gateway.example/v1#fragment",
"not a url",
],
)
async def test_live_rejects_invalid_api_base_before_network(api_base):
requests: list = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={})
async with live_handler(respond) as handler:
transport = LiveTransport(
LiveDeployment("deployment-model", provider="openai", api_key="deployment-key", api_base=api_base),
{},
http_handler=handler,
)
with pytest.raises(ValueError, match="Invalid Live API base"):
await transport.request("POST", "live/sessions", {})
assert requests == []

View file

@ -0,0 +1,519 @@
import json
import sys
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
import httpx
import pytest
import litellm
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.router import GenericLiteLLMParams
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", ["closed", "network"])
@pytest.mark.parametrize("hangup_status", [200, 503])
async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, chatgpt_tokens, monkeypatch):
from websockets.exceptions import ConnectionClosedOK
from websockets.frames import Close
from litellm.caching.llm_caching_handler import LLMClientCache
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache())
handler = ChatGPTRealtime(
GenericLiteLLMParams(
chatgpt_realtime_call_id="rtc_live_closed",
chatgpt_token_dir=chatgpt_tokens,
extra_query={"gateway": "tenant", "tag": ["alpha +/&", "beta"]},
),
{},
{"x-gateway-token": "test-only"},
)
connection = SimpleNamespace(
send=AsyncMock(
side_effect=(
ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True)
if failure == "closed"
else OSError("socket unavailable")
)
)
)
requests = []
def respond(request):
requests.append(request)
return httpx.Response(hangup_status)
client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
with patch("httpx.AsyncClient", return_value=client) as create_client:
for _ in range(2):
with pytest.raises(httpx.HTTPStatusError) if hangup_status == 503 else nullcontext():
await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1")
assert not client.is_closed
create_client.assert_called_once()
finally:
await client.aclose()
assert len(requests) == 2
assert requests[0].method == "POST"
assert requests[0].url.path == "/v1/realtime/calls/rtc_live_closed/hangup"
assert requests[0].url.params.get_list("tag") == ["alpha +/&", "beta"]
assert requests[0].url.params["gateway"] == "tenant"
assert requests[0].headers["x-gateway-token"] == "test-only"
assert requests[0].headers["Authorization"] == "Bearer test-token-default"
assert requests[0].extensions["timeout"]["read"] == 10
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"])
@pytest.mark.parametrize("source", ["default", "explicit", "CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"])
async def test_realtime_session_urls_honor_gateway(endpoint, source, chatgpt_tokens, monkeypatch):
monkeypatch.setenv("CHATGPT_TOKEN_DIR", chatgpt_tokens)
monkeypatch.delenv("CHATGPT_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False)
gateway = "https://voice.example/custom/v1/"
if source in ("CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"):
monkeypatch.setenv(source, gateway)
requests = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"client_secret": {"value": "test-secret"}})
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
kwargs = {"model": "chatgpt/gpt-realtime-1.5", "client": client}
if source == "explicit":
kwargs["api_base"] = gateway
try:
if endpoint == "client_secrets":
await litellm.acreate_realtime_client_secret(**kwargs)
else:
await litellm.acreate_realtime_transcription_session(**kwargs)
finally:
await client.client.aclose()
base = "https://api.openai.com/v1" if source == "default" else gateway.rstrip("/")
assert len(requests) == 1
assert str(requests[0].url) == f"{base}/realtime/{endpoint}"
assert requests[0].headers["authorization"] == "Bearer test-token-default"
@pytest.mark.asyncio
@pytest.mark.parametrize("inbound_headers", [{}, {"openai-alpha": "quicksilver=v2"}])
@pytest.mark.parametrize("model, endpoint", [("gpt-live-1-codex", "live"), ("gpt-realtime-1.5", "realtime")])
async def test_routed_call_preserves_deployment_gateway_headers(
inbound_headers, model, endpoint, chatgpt_tokens, monkeypatch
):
from litellm.llms.chatgpt.codex import (
CodexRealtimeCall,
CodexRealtimeOffer,
build_call_request,
build_sideband_request,
parse_call_response,
)
monkeypatch.setenv("CHATGPT_TOKEN_DIR", chatgpt_tokens)
requests = []
def respond(request):
requests.append(request)
return httpx.Response(201, text="v=0\r\n", headers={"location": "/v1/realtime/calls/rtc_test"})
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
router = litellm.Router(
model_list=[
{
"model_name": "voice-gateway",
"litellm_params": {
"model": f"chatgpt/{model}",
"api_base": "https://voice.example/backend-api/codex",
"extra_headers": {"x-gateway-route": "configured"},
"extra_query": {
"gateway_token": "configured",
"intent": "pinned-intent",
"count": 7,
"fraction": 1.5,
"enabled": True,
"disabled": False,
"blank": None,
"tag": ["alpha +/&", "beta"],
"empty": [],
"model": "other-model",
"call_id": "rtc_wrong",
},
},
"model_info": {"id": "selected-gateway-deployment"},
}
],
num_retries=0,
)
offer = CodexRealtimeOffer(sdp="v=0\r\n", session={"model": "voice-gateway"})
try:
response = await router.arealtime_calls(
**build_call_request(offer, {"intent": "quicksilver", "architecture": "avas"}, inbound_headers),
client=client,
)
assert requests[0].headers.get("x-gateway-route") == "configured"
assert dict(requests[0].url.params) == {
"gateway_token": "configured",
"intent": "pinned-intent",
"architecture": "avas",
"count": "7",
"fraction": "1.5",
"enabled": "true",
"disabled": "false",
"blank": "",
"tag": "alpha +/&",
"model": "other-model",
"call_id": "rtc_wrong",
}
assert requests[0].url.params.get_list("tag") == ["alpha +/&", "beta"]
assert response.extensions["chatgpt_realtime"]["extra_query"] == {
**dict(requests[0].url.params),
"tag": ("alpha +/&", "beta"),
"empty": (),
}
assert response.extensions["chatgpt_realtime"]["extra_headers"]["x-gateway-route"] == "configured"
for name, value in inbound_headers.items():
assert requests[0].headers[name] == value
call = parse_call_response(response, alias="voice-gateway", owner="test-owner", expires_at=1)
restored = CodexRealtimeCall.model_validate_json(call.model_dump_json())
assert restored.model_id == "selected-gateway-deployment"
assert restored.model == model
handler = ChatGPTRealtime(GenericLiteLLMParams.model_validate(build_sideband_request(restored)), {})
sideband_url = httpx.URL(handler._construct_url(restored.api_base, {"model": restored.model}))
assert {key: value for key, value in sideband_url.params.items() if key != "call_id"} == {
key: value for key, value in requests[0].url.params.items() if key not in ("model", "call_id")
}
assert sideband_url.params.get("call_id") == ("rtc_test" if endpoint == "realtime" else None)
assert sideband_url.params.get_list("tag") == ["alpha +/&", "beta"]
assert sideband_url.path.endswith("/realtime" if endpoint == "realtime" else "/live/rtc_test")
finally:
await client.client.aclose()
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ["gpt-realtime-1.5", "gpt-live-1-codex"])
@pytest.mark.parametrize("call_id", [None, "rtc_existing"])
async def test_websocket_forwards_configured_headers_without_client_identity(model, call_id, chatgpt_tokens):
websocket = SimpleNamespace(
headers={"authorization": "Bearer client", "cookie": "private-cookie", "openai-alpha": "client-value"},
scope={},
receive_text=AsyncMock(side_effect=RuntimeError("client disconnected")),
send_text=AsyncMock(),
close=AsyncMock(),
)
with patch("websockets.connect") as connect:
connect.return_value.__aenter__ = AsyncMock(side_effect=RuntimeError("stop before streaming"))
await litellm._arealtime(
model=f"chatgpt/{model}",
websocket=websocket,
api_base="https://voice.example/codex",
chatgpt_realtime_call_id=call_id,
query_params={"model": model, "intent": "client-intent"},
extra_query={"intent": "configured-intent", "tag": ["alpha +/&", "beta"]},
headers={"x-deployment-header": "configured"},
extra_headers={
"X-Gateway-Route": "voice",
"OpenAI-Alpha": "configured-value",
"aUtHoRiZaTiOn": "Bearer wrong",
"CHATGPT-ACCOUNT-ID": "wrong",
},
)
connect.assert_called_once()
headers = httpx.Headers(connect.call_args.kwargs["additional_headers"])
upstream_url = httpx.URL(connect.call_args.args[0])
assert upstream_url.params.get_list("intent") == ["configured-intent"]
assert upstream_url.params.get_list("tag") == ["alpha +/&", "beta"]
assert headers["x-deployment-header"] == "configured"
assert headers["x-gateway-route"] == "voice"
assert headers["openai-alpha"] == "configured-value"
assert headers["authorization"] == "Bearer test-token-default"
assert headers["chatgpt-account-id"] == "test-account-default"
assert "cookie" not in headers
@pytest.mark.asyncio
@pytest.mark.parametrize("api_base", [None, "https://voice.example/backend-api/codex"])
async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, api_base):
requests = []
def respond(request):
requests.append(request)
return httpx.Response(201, text="v=0\r\n", headers={"location": "/v1/realtime/calls/rtc_test"})
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
response = await litellm.arealtime_calls(
model="chatgpt/gpt-live-1-codex",
api_base=api_base,
openai_ephemeral_key="",
sdp_body=b"v=0\r\n",
session={"model": "chatgpt/gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}},
extra_query={"intent": "quicksilver", "architecture": "avas"},
chatgpt_realtime_client_query={"intent": "untrusted-override", "architecture": "avas", "untrusted": "bad"},
extra_headers={
"openai-alpha": "quicksilver=v2",
"x-gateway-route": "voice",
"aUtHoRiZaTiOn": "Bearer wrong",
"CHATGPT-ACCOUNT-ID": "wrong",
},
client=client,
)
assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1")
assert response.extensions["chatgpt_realtime"]["extra_headers"] == {
"openai-alpha": "quicksilver=v2",
"x-gateway-route": "voice",
}
assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com")
assert response.status_code == 201
assert response.extensions["chatgpt_realtime"]["extra_query"] == {"intent": "quicksilver", "architecture": "avas"}
assert requests[0].url.path == "/backend-api/codex/realtime/calls"
assert requests[0].url.params["architecture"] == "avas"
assert requests[0].headers["authorization"] == "Bearer test-token-" + "default"
assert requests[0].headers["chatgpt-account-id"] == "test-account-default"
assert requests[0].headers["openai-alpha"] == "quicksilver=v2"
assert requests[0].headers["x-gateway-route"] == "voice"
assert json.loads(requests[0].content) == {
"sdp": "v=0\r\n",
"session": {"model": "gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}},
}
await client.client.aclose()
@pytest.mark.asyncio
async def test_chatgpt_call_rejects_ephemeral_key_before_oauth_dispatch(chatgpt_tokens):
requests = []
def respond(request):
requests.append(request)
return httpx.Response(201, text="v=0\r\n")
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
with pytest.raises(litellm.AuthenticationError):
await litellm.arealtime_calls(
model="chatgpt/gpt-live-1-codex",
openai_ephemeral_key="legacy-ephemeral-key",
sdp_body=b"v=0\r\n",
client=client,
)
assert not requests
finally:
await client.client.aclose()
@pytest.mark.asyncio
async def test_openai_call_preserves_explicit_identity_headers():
requests = []
def respond(request):
requests.append(request)
return httpx.Response(201, text="v=0\r\n")
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
response = await litellm.arealtime_calls(
model="openai/gpt-realtime-1.5",
openai_ephemeral_key="original-key",
sdp_body=b"v=0\r\n",
extra_headers={"Authorization": "Bearer explicit-key", "chatgpt-account-id": "custom-account"},
client=client,
)
assert response.status_code == 201
assert requests[0].headers["authorization"] == "Bearer explicit-key"
assert requests[0].headers["chatgpt-account-id"] == "custom-account"
assert requests[0].headers["content-type"].startswith("multipart/form-data")
finally:
await client.client.aclose()
@pytest.mark.parametrize("model,endpoint", [("gpt-realtime-1.5", "realtime"), ("gpt-live-1-codex", "live")])
def test_realtime_uses_platform_endpoint_with_oauth_headers(model, endpoint, chatgpt_tokens, local_model_cost_map):
handler = ChatGPTRealtime(
GenericLiteLLMParams(),
{
"authorization": "Bearer proxy-key",
"openai-alpha": "quicksilver=v2",
},
)
assert handler._construct_url("https://api.openai.com/v1", {"model": model}) == (
f"wss://api.openai.com/v1/{endpoint}?model={model}"
)
headers = handler._get_additional_headers("unused")
assert headers["Authorization"] == "Bearer test-token-default"
assert "authorization" not in headers
assert headers["openai-alpha"] == "quicksilver=v2"
@pytest.mark.parametrize("endpoint", ["live", "realtime"])
def test_new_realtime_session_preserves_gateway_query(endpoint, chatgpt_tokens, local_model_cost_map):
model = "gpt-live-1-codex" if endpoint == "live" else "gpt-realtime-1.5"
handler = ChatGPTRealtime(
GenericLiteLLMParams(
chatgpt_token_dir=chatgpt_tokens,
chatgpt_realtime_client_query={"intent": "conversation", "architecture": "client-architecture"},
extra_query={
"gateway_token": "opaque +/& value",
"intent": "gateway-intent",
"architecture": "gateway-architecture",
"model": "other-model",
"call_id": "rtc_other",
},
),
{},
)
url = httpx.URL(handler._construct_url("https://gateway.example/v1", {"model": model, "intent": "query-intent"}))
assert url.path == f"/v1/{endpoint}"
assert dict(url.params) == {
"model": model,
"gateway_token": "opaque +/& value",
"intent": "gateway-intent",
"architecture": "gateway-architecture",
}
@pytest.mark.asyncio
async def test_openai_http_call_does_not_require_websockets(monkeypatch):
monkeypatch.delitem(sys.modules, "litellm.llms.chatgpt.realtime", raising=False)
for name in tuple(sys.modules):
if name == "websockets" or name.startswith("websockets."):
monkeypatch.delitem(sys.modules, name)
monkeypatch.setitem(sys.modules, "websockets", None)
requests = []
def respond(request):
requests.append(request)
return httpx.Response(201, text="v=0\r\n")
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
response = await litellm.arealtime_calls(
model="openai/gpt-realtime-1.5",
openai_ephemeral_key="test-only",
sdp_body=b"v=0\r\n",
api_key="test-only",
client=client,
)
assert response.status_code == 201
assert len(requests) == 1
assert requests[0].url.path == "/v1/realtime/calls"
finally:
await client.client.aclose()
@pytest.mark.parametrize("endpoint", ["live", "realtime"])
@pytest.mark.parametrize("call_id", [None, "rtc_metadata"])
def test_realtime_routes_new_models_using_registered_metadata(endpoint, call_id, chatgpt_tokens, local_model_cost_map):
model = "metadata-voice-model"
litellm.register_model({f"chatgpt/{model}": {
"litellm_provider": "chatgpt", "mode": "realtime", "supported_endpoints": [f"/v1/{endpoint}"]
}})
handler = ChatGPTRealtime(GenericLiteLLMParams(chatgpt_realtime_call_id=call_id), {})
expected = (
f"wss://api.openai.com/v1/{endpoint}?model={model}" if call_id is None
else f"wss://api.openai.com/v1/live/{call_id}" if endpoint == "live"
else f"wss://api.openai.com/v1/realtime?call_id={call_id}"
)
assert handler._construct_url("https://api.openai.com/v1", {"model": model}) == expected
def test_realtime_unknown_model_keeps_standard_endpoint(chatgpt_tokens, local_model_cost_map):
handler = ChatGPTRealtime(GenericLiteLLMParams(), {})
assert handler._construct_url("https://api.openai.com/v1", {"model": "unknown-voice-model"}) == (
"wss://api.openai.com/v1/realtime?model=unknown-voice-model"
)
@pytest.mark.parametrize("env_name", ["CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"])
@pytest.mark.parametrize("api_base", [None, "https://deployment.example/codex"])
def test_realtime_routes_use_configured_gateway(monkeypatch, env_name, api_base, chatgpt_tokens):
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
monkeypatch.delenv("CHATGPT_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False)
monkeypatch.setenv(env_name, "https://gateway.example/codex/")
expected = api_base or "https://gateway.example/codex"
config = ChatGPTRealtimeHTTPConfig(GenericLiteLLMParams())
assert config.get_realtime_calls_url(api_base, "gpt-live-1-codex") == expected + "/realtime/calls"
handler = ChatGPTRealtime(GenericLiteLLMParams(), {})
assert handler._construct_url(handler.get_api_base(api_base), {"model": "gpt-realtime-1.5"}) == (
expected.replace("https://", "wss://") + "/realtime?model=gpt-realtime-1.5"
)
@pytest.mark.parametrize("model,endpoint", [("gpt-live-1-codex", "live"), ("gpt-realtime-1.5", "realtime")])
def test_sideband_restores_gateway_query_without_overriding_call(model, endpoint, chatgpt_tokens):
handler = ChatGPTRealtime(
GenericLiteLLMParams(
chatgpt_realtime_call_id="rtc_selected",
extra_query={"gateway_token": "opaque +/& value", "model": "other", "call_id": "rtc_other"},
),
{},
)
url = httpx.URL(handler._construct_url("https://gateway.example/v1", {"model": model}))
assert url.params["gateway_token"] == "opaque +/& value"
assert "model" not in url.params
if endpoint == "live":
assert url.path == "/v1/live/rtc_selected"
assert "call_id" not in url.params
else:
assert url.path == "/v1/realtime"
assert url.params["call_id"] == "rtc_selected"
def test_client_cannot_forge_supervised_call_accounting(chatgpt_tokens):
from litellm.llms.chatgpt.realtime import CallAccounting, accounts_for_call_usage
assert accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting={"supervised": True}))
assert accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting="supervised"))
assert not accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting=CallAccounting.SUPERVISED))
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ["gpt-live-1-codex", "gpt-realtime-1.5"])
async def test_supervisor_connection_preserves_call_routing(model, chatgpt_tokens):
handler = ChatGPTRealtime(
GenericLiteLLMParams(
chatgpt_token_dir=chatgpt_tokens,
chatgpt_realtime_call_id="rtc_owner",
extra_query={"gateway_token": "a+b&c"},
),
{"openai-alpha": "quicksilver=v2"},
{"x-gateway-token": "configured"},
)
connection = AsyncMock()
with patch("websockets.connect", AsyncMock(return_value=connection)) as connect:
assert await handler.open_call_connection(model, "https://gateway.example/v1") is connection
url = httpx.URL(connect.call_args.args[0])
assert url.params["gateway_token"] == "a+b&c"
assert connect.call_args.kwargs["additional_headers"]["x-gateway-token"] == "configured"
assert url.path.endswith("/rtc_owner") if model == "gpt-live-1-codex" else url.params["call_id"] == "rtc_owner"
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ["gpt-live-1-codex", "gpt-realtime-1.5"])
async def test_close_call_prefers_session_close_only_for_live_models(model, chatgpt_tokens):
handler = ChatGPTRealtime(
GenericLiteLLMParams(chatgpt_realtime_call_id="rtc_close", chatgpt_token_dir=chatgpt_tokens),
{},
{},
)
connection = SimpleNamespace(send=AsyncMock())
handler.hangup_call = AsyncMock()
await handler.close_call(connection, model, "https://gateway.example/v1")
if model == "gpt-live-1-codex":
connection.send.assert_awaited_once_with('{"type":"session.close"}')
handler.hangup_call.assert_not_awaited()
else:
connection.send.assert_not_awaited()
handler.hangup_call.assert_awaited_once_with("https://gateway.example/v1")

View file

@ -3835,6 +3835,125 @@ def test_image_edit_handler_keeps_the_sync_transform():
assert response.data[0].b64_json == "sync"
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"])
@pytest.mark.parametrize("provider", ["chatgpt", "openai"])
@pytest.mark.parametrize("authorization_header", ["Authorization", "aUtHoRiZaTiOn"])
async def test_realtime_http_sessions_preserve_provider_identity(
endpoint, provider, authorization_header, tmp_path, monkeypatch
):
import time
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig
from litellm.types.router import GenericLiteLLMParams
monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path))
monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json")
(tmp_path / "auth.json").write_text(
json.dumps({"access_token": "test-resolved", "account_id": "test-selected", "expires_at": time.time() + 3600})
)
config = ChatGPTRealtimeHTTPConfig(GenericLiteLLMParams()) if provider == "chatgpt" else OpenAIRealtimeHTTPConfig()
requests = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"id": "session-test"})
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
response = await BaseLLMHTTPHandler()._async_realtime_session_post(
endpoint=endpoint,
api_base="https://gateway.example/v1",
api_key="test-openai",
request_data={"session": {"model": "gpt-realtime-1.5"}},
logging_obj=Mock(),
timeout=5,
provider_config=config,
model="gpt-realtime-1.5",
extra_headers={
authorization_header: "Bearer test-override",
"CHATGPT-ACCOUNT-ID": "test-other-account",
"x-gateway-route": "required",
},
client=client,
)
assert response.status_code == 200
assert not client.client.is_closed
finally:
await client.client.aclose()
assert len(requests) == 1
assert requests[0].url.path == f"/v1/realtime/{endpoint}"
assert requests[0].headers["x-gateway-route"] == "required"
if provider == "chatgpt":
assert requests[0].headers.get_list("authorization") == ["Bearer test-resolved"]
assert requests[0].headers.get_list("chatgpt-account-id") == ["test-selected"]
else:
assert requests[0].headers.get_list("authorization")[-1] == "Bearer test-override"
assert requests[0].headers["chatgpt-account-id"] == "test-other-account"
class _ImageGenerationRecordingConfig(BaseImageGenerationConfig):
def get_supported_openai_params(self, model):
return ["size"]
def map_openai_params(self, non_default_params, optional_params, model, drop_params):
optional_params.update(non_default_params)
return optional_params
def validate_environment(self, headers, model, messages, optional_params, litellm_params, api_key=None, api_base=None):
return {"authorization": f"Bearer {api_key}"}
def get_complete_url(self, api_base, api_key, model, optional_params, litellm_params, stream=None):
return "https://images.example/v1/generations"
def transform_image_generation_request(self, model, prompt, optional_params, litellm_params, headers):
return {"model": model, "prompt": prompt}
def transform_image_generation_response(self, model, raw_response, model_response, logging_obj, request_data, optional_params, litellm_params, encoding=None, api_key=None, json_mode=None):
return ImageResponse(data=[ImageObject(b64_json=raw_response.json()["created"])])
def test_image_extra_headers_strips_oauth_identity_only_for_chatgpt():
headers: Final = {"authorization": "Bearer oauth", "chatgpt-account-id": "acct-1", "x-router": "keep"}
assert BaseLLMHTTPHandler._image_extra_headers("openai", headers) is headers
stripped: Final = BaseLLMHTTPHandler._image_extra_headers("chatgpt", headers)
assert dict(stripped) == {"x-router": "keep"}
@pytest.mark.asyncio
async def test_async_image_generation_handler_merges_extra_headers_for_non_chatgpt():
requests: Final = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"created": "ok"})
client: Final = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
response: Final = await BaseLLMHTTPHandler().async_image_generation_handler(
model="image-model",
prompt="a red circle",
image_generation_provider_config=_ImageGenerationRecordingConfig(),
image_generation_optional_request_params={},
custom_llm_provider="openai",
litellm_params={"api_key": "sk-image"},
logging_obj=Mock(),
timeout=10,
extra_headers={"x-router-header": "routed"},
api_key="sk-image",
client=client,
)
finally:
await client.client.aclose()
assert requests[0].headers["x-router-header"] == "routed"
assert requests[0].headers["authorization"] == "Bearer sk-image"
assert requests[0].url == "https://images.example/v1/generations"
assert response.data[0].b64_json == "ok"
class _ScriptedClientWebSocket(_FakeClientWebSocket):
def __init__(self, messages: list[str], last_event_type: str) -> None:
super().__init__()

View file

@ -26,6 +26,15 @@ def test_openai_realtime_handler_url_construction(api_base):
assert "model=gpt-4o-realtime-preview-2024-10-01" in url
def test_openai_realtime_handler_requires_api_key():
from litellm.llms.openai.realtime.handler import OpenAIRealtime
handler = OpenAIRealtime()
with pytest.raises(ValueError, match="api_key is required for OpenAI realtime calls"):
handler._resolve_api_key(None)
assert handler._resolve_api_key("sk-realtime-key") == "sk-realtime-key"
def test_openai_realtime_handler_url_with_extra_params():
from litellm.llms.openai.realtime.handler import OpenAIRealtime
from litellm.types.realtime import RealtimeQueryParams

View file

@ -308,6 +308,7 @@ class TestProcessResponse:
)
@pytest.mark.usefixtures("local_model_cost_map")
class TestProcessEmbedContentResponseUsage:
"""Gemini Embedding 2 embedContent usageMetadata must drive spend.

View file

@ -547,6 +547,8 @@ async def test_virtual_key_max_budget_check(
False,
), # don't match on pattern
("openai/gpt-4o", ["openai/*"], True), # openai wildcard access
("openai/gpt+4", ["openai/gpt+*"], True), # regex metacharacters stay literal
("openai/gpttt4", ["openai/gpt+*"], False), # regex metacharacters do not overmatch
("gpt-4", ["gpt-3.5-turbo"], False), # model not in allowed list
("claude-3", [], True), # empty model list (allows all)
],

View file

@ -857,6 +857,7 @@ async def test_user_api_key_auth_websocket():
# Prepare a mock WebSocket object
mock_websocket = MagicMock(spec=WebSocket)
mock_websocket.query_params = {"model": "some_model"}
mock_websocket.path_params = {}
mock_websocket.headers = {"authorization": "Bearer some_api_key"}
# Mock the scope attribute that user_api_key_auth_websocket accesses
mock_websocket.scope = {"headers": [(b"authorization", b"Bearer some_api_key")]}
@ -880,6 +881,7 @@ async def test_user_api_key_auth_websocket():
assert request_arg.headers["authorization"] == "Bearer some_api_key"
assert mock_user_api_key_auth.call_args.kwargs["api_key"] == "Bearer some_api_key"
assert await request_arg.json() == {"model": "some_model"}
@pytest.mark.asyncio
@ -893,6 +895,7 @@ async def test_user_api_key_auth_websocket_carries_asgi_path():
mock_websocket = MagicMock(spec=WebSocket)
mock_websocket.query_params = {"model": "some_model"}
mock_websocket.path_params = {}
mock_websocket.headers = {"authorization": "Bearer some_api_key"}
mock_websocket.scope = {
"type": "websocket",

View file

@ -1,4 +1,5 @@
import asyncio
import json
import time
from types import TracebackType
from typing import Final
@ -26,6 +27,42 @@ class FakeLogging:
pass
@pytest.mark.parametrize("provider", [litellm.LlmProviders.XAI, litellm.LlmProviders.OPENAI, litellm.LlmProviders.GEMINI])
def test_realtime_handler_factory_does_not_read_headers_without_a_handler(provider):
from litellm.types.router import GenericLiteLLMParams
read_headers = MagicMock(side_effect=AssertionError("Headers must not be read"))
assert realtime_main.ProviderConfigManager.get_provider_realtime_handler(
provider, GenericLiteLLMParams(), read_headers
) is None
read_headers.assert_not_called()
def test_realtime_handler_factory_passes_actual_chatgpt_headers(tmp_path, monkeypatch):
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
from litellm.types.router import GenericLiteLLMParams
monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path))
monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json")
(tmp_path / "auth.json").write_text(
json.dumps({"access_token": "factory-test-token", "account_id": "factory-account", "expires_at": time.time() + 3600})
)
params = GenericLiteLLMParams(litellm_session_id="factory-session")
headers = {"openai-alpha": "quicksilver=v2"}
extra_headers = {"x-gateway-route": "required"}
read_headers = MagicMock(return_value=headers)
result = realtime_main.ProviderConfigManager.get_provider_realtime_handler(
litellm.LlmProviders.CHATGPT, params, read_headers, extra_headers
)
assert isinstance(result, ChatGPTRealtime)
read_headers.assert_called_once_with()
outgoing_headers = result._get_additional_headers("unused")
assert outgoing_headers["openai-alpha"] == headers["openai-alpha"]
assert outgoing_headers["x-gateway-route"] == extra_headers["x-gateway-route"]
assert outgoing_headers["session_id"] == "factory-session"
assert outgoing_headers["Authorization"] == "Bearer factory-test-token"
def test_resolves_top_level_session_model():
resolved = _with_resolved_session_model({"model": "alias/gpt-realtime"}, "gpt-realtime")
assert resolved == {"model": "gpt-realtime"}
@ -510,6 +547,32 @@ async def test_arealtime_azure_env_beta_protocol_wins_over_a_ga_client(monkeypat
)
@pytest.mark.parametrize("is_call", [False, True])
@pytest.mark.parametrize("provider", ["chatgpt", "openai", "azure"])
def test_realtime_http_provider_controls_dynamic_base_precedence(provider, is_call, monkeypatch):
from litellm.types.router import GenericLiteLLMParams
monkeypatch.delenv("CHATGPT_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False)
config, base, key = realtime_main._get_realtime_http_provider_config(
custom_llm_provider=provider,
dynamic_api_base="https://dynamic.example/v1",
dynamic_api_key="dynamic-key",
litellm_params=GenericLiteLLMParams(api_base="https://configured.example/v1"),
is_call=is_call,
)
expected_base = "https://configured.example/v1" if provider == "chatgpt" else "https://dynamic.example/v1"
assert base == expected_base
assert key == ("chatgpt-oauth" if provider == "chatgpt" else "dynamic-key")
assert config is not None
if provider == "chatgpt":
assert config.get_realtime_calls_url(base, "gpt-realtime-1.5") == expected_base + "/realtime/calls"
else:
assert config.get_realtime_calls_extra_headers({"x-gateway-route": "required"}) == {
"x-gateway-route": "required"
}
async def _vertex_provider_config_for(monkeypatch, model: str, vertex_location: str | None):
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import VertexChirpRealtimeConfig
from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig

View file

@ -1479,13 +1479,81 @@ def test_azure_ai_cache_cost_calculation(_local_model_cost_map):
print(f"Output cost: {output_cost}, Expected: {expected_output_cost}")
print(f"Total cost: {total_cost}")
assert abs(input_cost - expected_input_cost) < 1e-10, (
f"Input cost mismatch: got {input_cost}, expected {expected_input_cost}"
assert (
abs(input_cost - expected_input_cost) < 1e-10
), f"Input cost mismatch: got {input_cost}, expected {expected_input_cost}"
assert (
abs(output_cost - expected_output_cost) < 1e-10
), f"Output cost mismatch: got {output_cost}, expected {expected_output_cost}"
AZURE_GPT_5_6_MAP_KEYS = (
"azure/gpt-5.6",
"azure/gpt-5.6-sol",
"azure/gpt-5.6-terra",
"azure/gpt-5.6-luna",
"azure/us/gpt-5.6",
"azure/us/gpt-5.6-sol",
"azure/us/gpt-5.6-terra",
"azure/us/gpt-5.6-luna",
"azure/eu/gpt-5.6",
"azure/eu/gpt-5.6-sol",
"azure/eu/gpt-5.6-terra",
"azure/eu/gpt-5.6-luna",
)
def test_azure_gpt_5_6_cache_write_tokens_are_billed(_local_model_cost_map):
"""
Azure bills gpt-5.6 prompt cache writes at 1.25x the input rate on every
tier, but the azure entries carried no ``cache_creation_input_token_cost``,
so cache-write tokens were billed at the plain input rate instead.
"""
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
usage = Usage(
completion_tokens=100,
prompt_tokens=2000,
total_tokens=2100,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=0, text_tokens=687),
cache_creation_input_tokens=1313,
)
assert abs(output_cost - expected_output_cost) < 1e-10, (
f"Output cost mismatch: got {output_cost}, expected {expected_output_cost}"
input_cost, output_cost = generic_cost_per_token(
model="azure/gpt-5.6-luna", usage=usage, custom_llm_provider="azure"
)
assert input_cost == pytest.approx(687 * 2e-07 + 1313 * 2.5e-07)
assert output_cost == pytest.approx(100 * 1.2e-06)
@pytest.mark.parametrize("model", AZURE_GPT_5_6_MAP_KEYS)
def test_azure_gpt_5_6_rates_match_azure_price_page(_local_model_cost_map, model):
"""
Per the Azure OpenAI price page (rendered 2026-08-26): cache writes cost
1.25x input on every gpt-5.6 tier, and Data Zone costs 1.1x Global for
standard and priority alike (us/eu priority rates previously sat at 1.25x).
"""
entry = litellm.model_cost[model]
input_keys = [key for key in entry if key.startswith("input_cost_per_token")]
assert input_keys
for key in input_keys:
suffix = key[len("input_cost_per_token") :]
assert entry["cache_creation_input_token_cost" + suffix] == pytest.approx(entry[key] * 1.25)
zone = model.split("/")[1]
if zone in ("us", "eu"):
global_entry = litellm.model_cost["azure/" + model.split("/", 2)[2]]
prefixes = ("input_cost_per_token", "output_cost_per_token", "cache_read", "cache_creation")
token_cost_keys = [key for key in entry if key.startswith(prefixes)]
global_token_cost_keys = [key for key in global_entry if key.startswith(prefixes)]
assert len(token_cost_keys) >= 9
assert set(token_cost_keys) <= set(global_token_cost_keys)
for key in token_cost_keys:
assert entry[key] == pytest.approx(global_entry[key] * 1.1), key
def test_vertex_regional_deployment_costs_uplift_over_global(monkeypatch):
"""
@ -3647,6 +3715,101 @@ def test_combine_usage_objects_sums_mirrored_cache_write_fields_once():
assert combined_pair.prompt_tokens_details.cache_creation_tokens == 100
def test_completion_cost_prices_anthropic_shaped_cache_read_tokens(_local_model_cost_map):
"""Regression: an Anthropic /v1/messages response reports cache reads as top-level
cache_read_input_tokens with input_tokens excluding them. Reading that usage as
Responses API usage dropped the cache tokens and billed the whole prompt at the
uncached input rate, overstating spend on cache hits."""
response = {
"id": "msg_1",
"type": "message",
"role": "assistant",
"model": "gpt-5.6-sol",
"stop_reason": "end_turn",
"content": [{"type": "text", "text": "1"}],
"usage": {"input_tokens": 3, "output_tokens": 5, "cache_read_input_tokens": 4014},
}
cost = litellm.completion_cost(
completion_response=response,
model="gpt-5.6-sol",
custom_llm_provider="openai",
)
assert cost == pytest.approx(3 * 4e-6 + 4014 * 4e-7 + 5 * 2e-5, rel=1e-9)
def _together_chat_response(model: str, prompt_tokens: int, completion_tokens: int, cached_tokens: int) -> ModelResponse:
return ModelResponse(
id="chatcmpl-together-cache",
choices=[{"finish_reason": "stop", "index": 0, "message": {"content": "acknowledged", "role": "assistant"}}],
created=1756164000,
model=model,
object="chat.completion",
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
),
)
def test_completion_cost_prices_together_cached_tokens_at_cache_read_rate(_local_model_cost_map):
"""Regression: Together reports prompt_tokens_details.cached_tokens but no together_ai
registry entry carried cache_read_input_token_cost, so cache-hit tokens were priced at
0.0 and spend on cache-heavy workloads was understated."""
cost = completion_cost(
completion_response=_together_chat_response(
model="deepseek-ai/DeepSeek-V4-Flash-0731", prompt_tokens=7864, completion_tokens=16, cached_tokens=7863
),
custom_llm_provider="together_ai",
)
assert cost == pytest.approx(1 * 1.4e-07 + 7863 * 3e-08 + 16 * 2.8e-07, rel=1e-9)
def test_completion_cost_together_mapped_model_skips_size_bucket(_local_model_cost_map):
"""Regression: any together model whose name matches (\\d+b) was rewritten to a
together-ai-* size bucket before the registry lookup, so mapped models like
Muse-Glimmer-30B never used their per-model rates, cache fields included."""
cost = completion_cost(
completion_response=_together_chat_response(
model="meta-models/Muse-Glimmer-30B", prompt_tokens=63, completion_tokens=16, cached_tokens=0
),
custom_llm_provider="together_ai",
)
assert cost == pytest.approx(63 * 3.5e-07 + 16 * 1.5e-06, rel=1e-9)
def test_completion_cost_together_unmapped_model_still_uses_size_bucket(_local_model_cost_map):
cost = completion_cost(
completion_response=_together_chat_response(
model="qwen/Qwen2-72B-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0
),
custom_llm_provider="together_ai",
)
assert cost == pytest.approx((23 + 15) * 9e-07, rel=1e-9)
def test_completion_cost_together_metadata_only_model_still_uses_size_bucket(_local_model_cost_map):
assert "input_cost_per_token" not in litellm.model_cost["together_ai/togethercomputer/CodeLlama-34b-Instruct"]
cost = completion_cost(
completion_response=_together_chat_response(
model="togethercomputer/CodeLlama-34b-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0
),
custom_llm_provider="together_ai",
)
assert cost == pytest.approx((23 + 15) * 8e-07, rel=1e-9)
def test_select_model_name_strips_unregistered_alias_prefix(_local_model_cost_map):
"""A router-facing model_name alias containing "/" whose leading segment is NOT a
registered provider must not be double-prefixed into a non-existent cost key.
@ -4233,6 +4396,106 @@ def test_collect_and_combine_realtime_usage_stores_partitioned_text_tokens() ->
assert combined.completion_tokens_details.audio_tokens == 0
def _live_terminal_event(duration=4000):
return {"type": "session.closed", "usage": {"audio_duration_ms": duration, "backend_model_usage": []}}
@pytest.mark.parametrize("rate,expected", [(0.025, 0.1), (0, 0), (None, 0)])
def test_live_terminal_duration_uses_configured_second_price(monkeypatch, rate, expected):
monkeypatch.setitem(
litellm.model_cost,
"live-priced-test",
{"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": rate},
)
assert handle_realtime_stream_cost_calculation(
[_live_terminal_event()], Usage(), "chatgpt", "live-priced-test"
) == pytest.approx(expected)
@pytest.mark.parametrize("public_live", [False, True])
def test_live_terminal_duration_honors_deployment_override(monkeypatch, public_live):
monkeypatch.setitem(
litellm.model_cost,
"live-deployment-test",
{"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025},
)
result = RealtimeAPITokenUsageProcessor.create_logging_realtime_object(
Usage(), [{"type": "session.closed", "usage": {"seconds": 4}}] if public_live else [_live_terminal_event()]
)
assert completion_cost(
completion_response=result,
model="gpt-live-1",
custom_llm_provider="chatgpt",
call_type="_arealtime",
custom_pricing=True,
router_model_id="live-deployment-test",
) == pytest.approx(0.1)
@pytest.mark.parametrize("duration", [-1, True, "4000", float("inf"), float("nan"), None])
def test_live_terminal_invalid_duration_does_not_create_spend(monkeypatch, duration):
monkeypatch.setitem(
litellm.model_cost,
"live-priced-test",
{"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025},
)
assert (
handle_realtime_stream_cost_calculation(
[_live_terminal_event(duration)], Usage(), "chatgpt", "live-priced-test"
)
== 0
)
def test_live_terminal_is_not_counted_twice(monkeypatch):
monkeypatch.setitem(
litellm.model_cost,
"live-priced-test",
{"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025},
)
assert handle_realtime_stream_cost_calculation(
[_live_terminal_event(), _live_terminal_event()], Usage(), "chatgpt", "live-priced-test"
) == pytest.approx(0.1)
@pytest.mark.parametrize("with_tokens", [False, True])
@pytest.mark.parametrize("terminal_count", [1, 2])
@pytest.mark.parametrize("duration_priced", [False, True])
def test_live_terminal_with_response_done_preserves_configured_billing(
monkeypatch, with_tokens, terminal_count, duration_priced
):
monkeypatch.setitem(
litellm.model_cost,
"realtime-deployment-test",
{
"litellm_provider": "chatgpt",
"mode": "realtime",
**(
{"input_cost_per_second": 0.025}
if duration_priced
else {"input_cost_per_token": 0.001, "output_cost_per_token": 0.002}
),
},
)
events = [
{
"type": "response.done",
"response": {
"usage": ({"input_tokens": 10, "output_tokens": 5, "total_tokens": 15} if with_tokens else {})
},
},
*(_live_terminal_event() for _ in range(terminal_count)),
]
usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(events)
result = RealtimeAPITokenUsageProcessor.create_logging_realtime_object(usage, events)
assert completion_cost(
completion_response=result,
model="gpt-live-1-codex" if duration_priced else "gpt-realtime-1.5",
custom_llm_provider="chatgpt",
call_type="_arealtime",
custom_pricing=True,
router_model_id="realtime-deployment-test",
) == pytest.approx(0.1 if duration_priced else (0.02 if with_tokens else 0))
def test_realtime_combine_sums_nested_cached_tokens_details():
results: OpenAIRealtimeStreamList = [
{
@ -4561,6 +4824,241 @@ def test_completion_cost_ocr_ignores_deployment_pricing_without_custom_pricing_f
assert cost == 0.0
@pytest.mark.parametrize("terminal", [False, True])
def test_public_live_seconds_are_cumulative_and_backend_usage_is_separately_priced(monkeypatch, terminal):
monkeypatch.setitem(
litellm.model_cost,
"live-seconds-test",
{
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_second": 0.025,
"input_cost_per_token": 100,
"output_cost_per_token": 100,
},
)
monkeypatch.setitem(
litellm.model_cost,
"live-backend-test",
{
"litellm_provider": "openai",
"mode": "responses",
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
},
)
backend = {
"type": "response.event",
"event": {
"type": "response.completed",
"response": {
"id": "resp_backend",
"created_at": 1,
"model": "live-backend-test",
"output": [],
"usage": {"input_tokens": 20, "output_tokens": 10, "total_tokens": 30},
},
},
}
events = [
{"type": "session.usage.updated", "usage": {"seconds": 15}},
{"type": "session.usage.updated", "usage": {"seconds": 30}},
backend,
backend,
{"type": "session.closed" if terminal else "session.usage.updated", "usage": {"seconds": 30}},
]
usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(events)
assert usage.total_tokens == 30
assert handle_realtime_stream_cost_calculation(events, usage, "openai", "live-seconds-test") == pytest.approx(0.79)
@pytest.mark.parametrize("seconds", [-1, True, "30", float("inf"), float("nan"), None])
def test_public_live_invalid_seconds_are_not_billed(monkeypatch, seconds):
monkeypatch.setitem(
litellm.model_cost,
"live-seconds-test",
{
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_second": 0.025,
},
)
assert (
handle_realtime_stream_cost_calculation(
[{"type": "session.closed", "usage": {"seconds": seconds}}], Usage(), "openai", "live-seconds-test"
)
== 0
)
@pytest.mark.parametrize("terminal", [False, True])
def test_live_duration_does_not_regress_when_primary_and_observer_events_interleave(monkeypatch, terminal):
monkeypatch.setitem(
litellm.model_cost,
"live-interleaved-test",
{"litellm_provider": "openai", "mode": "realtime", "input_cost_per_second": 0.025},
)
events = [
{"type": "session.closed" if terminal else "session.usage.updated", "usage": {"seconds": 30}},
{"type": "session.usage.updated", "usage": {"seconds": 15}},
]
assert handle_realtime_stream_cost_calculation(events, Usage(), "openai", "live-interleaved-test") == pytest.approx(
0.75
)
@pytest.mark.parametrize("seconds,expected", [(None, 15), (0, 15), (4, 15), (15, 15), (30, 30)])
def test_live_webrtc_initialization_is_credited_against_duration(monkeypatch, seconds, expected):
monkeypatch.setitem(
litellm.model_cost,
"live-init-test",
{
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_second": 0.025,
},
)
events = [{"type": "litellm.live.initialization", "usage": {"seconds": 15}}]
if seconds is not None:
events.append({"type": "session.closed", "usage": {"seconds": seconds}})
assert handle_realtime_stream_cost_calculation(events, Usage(), "openai", "live-init-test") == pytest.approx(
expected * 0.025
)
def test_live_invalid_terminal_retains_last_reported_partial_usage(monkeypatch):
monkeypatch.setitem(
litellm.model_cost,
"live-partial-test",
{
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_second": 0.025,
},
)
events = [
{"type": "litellm.live.initialization", "usage": {"seconds": 15}},
{"type": "session.usage.updated", "usage": {"seconds": 30}},
{"type": "session.closed", "usage": {"seconds": "invalid"}},
]
assert handle_realtime_stream_cost_calculation(events, Usage(), "openai", "live-partial-test") == pytest.approx(
0.75
)
@pytest.mark.parametrize(
"nested",
[
{"type": "response.created", "response": {"model": "still-starting"}},
{"type": "response.in_progress", "response": {}},
{"type": "future.event", "response": ["unknown", "payload"]},
{"type": "response.completed", "response": {"id": "broken", "usage": "invalid"}},
],
)
def test_live_partial_or_malformed_backend_events_preserve_duration(monkeypatch, nested):
from unittest.mock import MagicMock
monkeypatch.setitem(
litellm.model_cost,
"live-resilient-test",
{
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_second": 0.025,
},
)
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
events = [
{"type": "response.event", "event": nested},
{"type": "session.closed", "usage": {"seconds": 30}},
]
usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(events)
assert usage.total_tokens == 0
assert handle_realtime_stream_cost_calculation(
events,
usage,
"openai",
"live-resilient-test",
litellm_logging_obj=logger,
) == pytest.approx(0.75)
assert bool(logger.model_call_details.get("realtime_backend_accounting_incomplete")) == (
nested["type"] == "response.completed"
)
@pytest.mark.parametrize("missing_usage_duplicate", [False, True])
def test_live_missing_backend_price_preserves_duration_and_marks_accounting_incomplete(
monkeypatch, missing_usage_duplicate
):
from unittest.mock import MagicMock
monkeypatch.setitem(
litellm.model_cost,
"live-resilient-test",
{
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_second": 0.025,
},
)
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
response = {
"id": "resp_unknown",
"created_at": 1,
"model": "unmapped-live-backend-price-test",
"output": [],
"usage": {"input_tokens": 20, "output_tokens": 10, "total_tokens": 30},
}
events = [
{"type": "response.event", "event": {"type": "response.completed", "response": response}},
*(
[
{
"type": "response.event",
"event": {"type": "response.completed", "response": {**response, "usage": None}},
}
]
if missing_usage_duplicate
else []
),
{"type": "session.closed", "usage": {"seconds": 30}},
]
usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(events)
assert usage.total_tokens == 30
assert handle_realtime_stream_cost_calculation(
events,
usage,
"openai",
"live-resilient-test",
litellm_logging_obj=logger,
) == pytest.approx(0.75)
assert logger.model_call_details["realtime_backend_accounting_incomplete"] is True
@pytest.mark.parametrize(
"envelope",
[
{"type": "response.event"},
{"type": "response.event", "event": {"response": {"id": "resp"}}},
{"type": "response.event", "event": "not-an-object"},
],
)
def test_live_backend_malformed_envelope_is_dropped_without_accounting_flag(envelope):
"""
Malformed event envelopes (missing event, missing event.type, wrong shape)
are skipped silently: unlike a terminal response.completed that fails
response validation, they must not mark the call's accounting incomplete.
"""
from unittest.mock import MagicMock
from litellm.cost_calculator import _live_backend_response
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
assert _live_backend_response(envelope, logger) is None
assert "realtime_backend_accounting_incomplete" not in logger.model_call_details
def test_completion_cost_prices_responses_websocket_turns_per_service_tier():
"""Issue #41299: a session mixing default and priority turns must price each turn at
its own returned service_tier, not the summed usage at a single tier."""

View file

@ -10,6 +10,8 @@ gpt-image-1 uses token-based pricing:
- Image Output: $40.00/1M tokens
"""
from typing import Final
import pytest
import litellm
@ -37,6 +39,141 @@ def _use_local_model_cost_map(monkeypatch):
class TestGPTImageCostCalculator:
"""Test the OpenAI gpt-image cost calculator"""
@pytest.mark.parametrize("family", ["flare", "sunburst"])
@pytest.mark.parametrize("snapshot", ["", "-2026-09-08"])
@pytest.mark.parametrize("call_type", ["image_generation", "image_edit"])
@pytest.mark.parametrize("cached_text,cached_image", [(0, 0), (50, 500)])
def test_image_25_official_prices(self, family, snapshot, call_type, cached_text, cached_image):
response: Final = ImageResponse(
created=1,
data=[],
usage={
"input_tokens": 1100,
"output_tokens": 100,
"total_tokens": 1200,
"input_tokens_details": {
"text_tokens": 100,
"image_tokens": 1000,
"cached_tokens": cached_text + cached_image,
"cached_tokens_details": {"text_tokens": cached_text, "image_tokens": cached_image},
},
},
)
cost: Final = litellm.completion_cost(
model="gpt-image-2.5-" + family + snapshot,
completion_response=response,
call_type=call_type,
custom_llm_provider="openai",
)
expected: Final = (
(100 - cached_text) * 5e-6
+ cached_text * 1.25e-6
+ (1000 - cached_image) * 8e-6
+ cached_image * 2e-6
+ 100 * 30e-6
)
assert cost == pytest.approx(expected)
def test_gpt_image_1_cost_with_text_only(self):
"""Test cost calculation with only text input tokens"""
from litellm.llms.openai.image_generation.cost_calculator import cost_calculator
usage = ImageUsage(
input_tokens=100,
output_tokens=5000,
total_tokens=5100,
input_tokens_details=ImageUsageInputTokensDetails(
text_tokens=100,
image_tokens=0,
),
)
image_response = ImageResponse(
created=1234567890,
data=[ImageObject(url="http://example.com/image.jpg")],
)
image_response.usage = usage
cost = cost_calculator(
model="gpt-image-1",
image_response=image_response,
custom_llm_provider="openai",
)
# Expected cost:
# Text input: 100 * $5/1M = 0.0005
# Image output: 5000 * $40/1M = 0.2
# Total: 0.2005
expected_cost = 0.0005 + 0.2
assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}"
def test_gpt_image_1_cost_with_image_input(self):
"""Test cost calculation with both text and image input tokens (for edits)"""
from litellm.llms.openai.image_generation.cost_calculator import cost_calculator
usage = ImageUsage(
input_tokens=600,
output_tokens=5000,
total_tokens=5600,
input_tokens_details=ImageUsageInputTokensDetails(
text_tokens=100,
image_tokens=500,
),
)
image_response = ImageResponse(
created=1234567890,
data=[ImageObject(url="http://example.com/image.jpg")],
)
image_response.usage = usage
cost = cost_calculator(
model="gpt-image-1",
image_response=image_response,
custom_llm_provider="openai",
)
# Expected cost:
# Text input: 100 * $5/1M = 0.0005
# Image input: 500 * $10/1M = 0.005
# Image output: 5000 * $40/1M = 0.2
# Total: 0.2055
expected_cost = 0.0005 + 0.005 + 0.2
assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}"
def test_gpt_image_1_mini_cost(self):
"""Test cost calculation for gpt-image-1-mini model"""
from litellm.llms.openai.image_generation.cost_calculator import cost_calculator
usage = ImageUsage(
input_tokens=100,
output_tokens=5000,
total_tokens=5100,
input_tokens_details=ImageUsageInputTokensDetails(
text_tokens=100,
image_tokens=0,
),
)
image_response = ImageResponse(
created=1234567890,
data=[ImageObject(url="http://example.com/image.jpg")],
)
image_response.usage = usage
cost = cost_calculator(
model="gpt-image-1-mini",
image_response=image_response,
custom_llm_provider="openai",
)
# Expected cost for gpt-image-1-mini:
# Text input: 100 * $2/1M = 0.0002
# Image output: 5000 * $8/1M = 0.04
# Total: 0.0402
expected_cost = 0.0002 + 0.04
assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}"
def test_gpt_image_1_cost_no_usage(self):
"""Test that cost returns 0 when no usage data is available"""
from litellm.llms.openai.image_generation.cost_calculator import cost_calculator

View file

@ -786,7 +786,6 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"cache_read_input_token_cost_above_272k_tokens_flex": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens_ultrafast": {"type": "number"},
"cache_read_input_token_cost_above_512k_tokens": {"type": "number"},
"input_cost_per_token_above_272k_tokens_ultrafast": {"type": "number"},
"cache_read_input_token_cost_batches": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens_batches": {"type": "number"},
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": {"type": "number"},
@ -809,6 +808,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"input_cost_per_token_above_200k_tokens_batches": {"type": "number"},
"input_cost_per_token_above_256k_tokens": {"type": "number"},
"input_cost_per_token_above_272k_tokens": {"type": "number"},
"input_cost_per_token_above_272k_tokens_ultrafast": {"type": "number"},
"input_cost_per_token_above_512k_tokens": {"type": "number"},
"cache_read_input_token_cost_flex": {"type": "number"},
"cache_read_input_token_cost_priority": {"type": "number"},
@ -1009,6 +1009,8 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"/v1/messages",
"/v1/images/generations",
"/v1/realtime",
"/v1/realtime/calls",
"/v1/live",
"/v1/realtime/transcription_sessions",
"/v1/images/variations",
"/v1/images/edits",
@ -1289,7 +1291,11 @@ def test_openai_models_in_model_info(monkeypatch):
model_map = litellm.model_cost
violated_models = []
for model, info in model_map.items():
if info.get("litellm_provider") == "openai" and info.get("supports_vision") is True:
if (
info.get("litellm_provider") == "openai"
and info.get("supports_vision") is True
and info.get("mode") != "image_generation"
):
if info.get("supports_pdf_input") is not True:
violated_models.append(model)
assert len(violated_models) == 0, f"The following models should support pdf input: {violated_models}"

View file

@ -85,6 +85,9 @@ CONNECTION_NAMES: Final = (
"litellm_credential_name",
"configurable_clientside_auth_params",
"use_xai_oauth",
"chatgpt_auth_profile",
"chatgpt_auth_file",
"chatgpt_token_dir",
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",

File diff suppressed because it is too large Load diff