mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge 966b047dd5 into 0980f756bd
This commit is contained in:
commit
047350c725
90 changed files with 16708 additions and 982 deletions
5
.github/codeql/codeql-config.yml
vendored
5
.github/codeql/codeql-config.yml
vendored
|
|
@ -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.
|
||||
|
|
|
|||
11
.github/workflows/_test-unit-base.yml
vendored
11
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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}"
|
||||
|
|
|
|||
5
.github/workflows/codeql.yml
vendored
5
.github/workflows/codeql.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -114,6 +114,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/{provider}/",
|
||||
"/toolset/",
|
||||
# Realtime / streaming
|
||||
"/v1/live",
|
||||
"/live",
|
||||
"/v1/realtime",
|
||||
"/realtime",
|
||||
# Health & ops
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
})?;
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
95
litellm/llms/chatgpt/codex.py
Normal file
95
litellm/llms/chatgpt/codex.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
154
litellm/llms/chatgpt/images.py
Normal file
154
litellm/llms/chatgpt/images.py
Normal 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,
|
||||
}, ()
|
||||
196
litellm/llms/chatgpt/live.py
Normal file
196
litellm/llms/chatgpt/live.py
Normal 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,
|
||||
)
|
||||
290
litellm/llms/chatgpt/realtime.py
Normal file
290
litellm/llms/chatgpt/realtime.py
Normal 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"
|
||||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
74
litellm/proxy/hooks/realtime_call_lease.py
Normal file
74
litellm/proxy/hooks/realtime_call_lease.py
Normal 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()
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
562
litellm/proxy/realtime_endpoints/call_sessions.py
Normal file
562
litellm/proxy/realtime_endpoints/call_sessions.py
Normal 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)
|
||||
270
litellm/proxy/realtime_endpoints/call_supervision.py
Normal file
270
litellm/proxy/realtime_endpoints/call_supervision.py
Normal 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()
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
1639
litellm/proxy/realtime_endpoints/live.py
Normal file
1639
litellm/proxy/realtime_endpoints/live.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
102
tests/local_testing/test_realtime_call_redis.py
Normal file
102
tests/local_testing/test_realtime_call_redis.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
85
tests/test_litellm/proxy/hooks/test_realtime_call_lease.py
Normal file
85
tests/test_litellm/proxy/hooks/test_realtime_call_lease.py
Normal 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()
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
1655
tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py
Normal file
1655
tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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)
|
||||
2392
tests/test_litellm/proxy/realtime_endpoints/test_live.py
Normal file
2392
tests/test_litellm/proxy/realtime_endpoints/test_live.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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 (
|
||||
|
|
|
|||
64
tests/test_litellm/proxy/test_live_route_registration.py
Normal file
64
tests/test_litellm/proxy/test_live_route_registration.py
Normal 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()
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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]}...")
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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}}]
|
||||
|
|
|
|||
22
tests/unit/llms/chatgpt/conftest.py
Normal file
22
tests/unit/llms/chatgpt/conftest.py
Normal 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)
|
||||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
63
tests/unit/llms/chatgpt/test_codex.py
Normal file
63
tests/unit/llms/chatgpt/test_codex.py
Normal 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)
|
||||
305
tests/unit/llms/chatgpt/test_images.py
Normal file
305
tests/unit/llms/chatgpt/test_images.py
Normal 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(), {}
|
||||
)
|
||||
261
tests/unit/llms/chatgpt/test_live.py
Normal file
261
tests/unit/llms/chatgpt/test_live.py
Normal 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 == []
|
||||
519
tests/unit/llms/chatgpt/test_realtime.py
Normal file
519
tests/unit/llms/chatgpt/test_realtime.py
Normal 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")
|
||||
|
|
@ -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__()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -308,6 +308,7 @@ class TestProcessResponse:
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("local_model_cost_map")
|
||||
class TestProcessEmbedContentResponseUsage:
|
||||
"""Gemini Embedding 2 embedContent usageMetadata must drive spend.
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
],
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
1203
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
1203
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue