Merge remote-tracking branch 'upstream/litellm_internal_staging' into fix/bedrock-guardrail-image-input

This commit is contained in:
feng.tsai 2026-09-01 17:30:26 +08:00
commit a2d98df942
142 changed files with 7046 additions and 1589 deletions

77
.github/workflows/test-redis-compat.yml vendored Normal file
View file

@ -0,0 +1,77 @@
name: "Unit Tests: Redis Client Version Compatibility"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "litellm/_redis.py"
- "litellm/_redis_credential_provider.py"
- "tests/test_litellm/test_redis.py"
- "tests/test_litellm/caching/test_redis_connection_pool.py"
- ".github/workflows/test-redis-compat.yml"
- "pyproject.toml"
- "uv.lock"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
redis-compat:
name: "redis-py ${{ matrix.redis-version }}"
runs-on: ubuntu-latest
timeout-minutes: 15
strategy:
fail-fast: false
matrix:
# 5.3.1 is the version pinned in uv.lock (redisvl caps it below 6); the
# newer legs prove the inspect.signature introspection in litellm/_redis.py
# keeps extracting kwargs on the redis-py releases people actually run now.
# Only the exact release 6.0.0 is skipped: rq (pulled by the proxy extra)
# specifies `redis != 6`, which excludes 6.0.0 alone, so 6.4.0 stands in
# for the 6.x line.
redis-version: ["5.3.1", "6.4.0", "7.4.1", "8.0.1"]
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Pin redis-py to the matrix version
env:
REDIS_VERSION: ${{ matrix.redis-version }}
run: |
uv pip install "redis==${REDIS_VERSION:?}"
uv run --no-sync python -c "import redis; assert redis.__version__ == '${REDIS_VERSION:?}', redis.__version__; print('redis-py', redis.__version__)"
- name: Run redis unit tests
run: |
uv run --no-sync pytest \
tests/test_litellm/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
--tb=short -vv \
--reruns 2 \
--reruns-delay 1 \
--durations=20

View file

@ -1,9 +1,9 @@
{
"reportAny": {
"limit": 16171
"limit": 14765
},
"reportArgumentType": {
"limit": 2224
"limit": 2216
},
"reportAssignmentType": {
"limit": 319
@ -24,7 +24,7 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 5199
"limit": 4493
},
"reportFunctionMemberAccess": {
"limit": 7
@ -42,7 +42,7 @@
"limit": 12
},
"reportIndexIssue": {
"limit": 35
"limit": 25
},
"reportInvalidTypeForm": {
"limit": 34
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5611
"limit": 5607
},
"reportMissingTypeArgument": {
"limit": 15348
"limit": 15310
},
"reportMissingTypeStubs": {
"limit": 40
@ -105,13 +105,13 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38465
"limit": 38368
},
"reportUnknownParameterType": {
"limit": 19663
"limit": 19633
},
"reportUnknownVariableType": {
"limit": 30064
"limit": 29908
},
"reportUnnecessaryCast": {
"limit": 111
@ -141,6 +141,6 @@
"limit": 543
},
"reportUnusedVariable": {
"limit": 139
"limit": 137
}
}

View file

@ -5,7 +5,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t
from dataclasses import replace as dataclasses_replace
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Dict, Final, List, Literal, Optional, Tuple, cast
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, cast
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -87,7 +87,7 @@ class CheckBatchCost:
return
self.batch_processed_support_confirmed = True
async def _get_user_info(self, batch_id: str, user_id: Optional[str]) -> Dict[str, Any]:
async def _get_user_info(self, batch_id: str, user_id: Optional[str]) -> dict[str, str | None]:
"""
Look up user email and key alias by user_id for enriching the S3 callback metadata.
Returns a dict with user_api_key_user_email and user_api_key_alias (both may be None).
@ -97,8 +97,10 @@ class CheckBatchCost:
if not user_id:
return {}
try:
user_row = await self.prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_id}
user_row: prisma_models.LiteLLM_UserTable | None = (
await self.prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_id}
)
)
if user_row is None:
return {}
@ -115,8 +117,10 @@ class CheckBatchCost:
if not api_key:
return None
try:
key_row = await self.prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": api_key}
key_row: prisma_models.LiteLLM_VerificationToken | None = (
await self.prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": api_key}
)
)
return getattr(key_row, "key_alias", None) if key_row is not None else None
except Exception as e:
@ -128,8 +132,10 @@ class CheckBatchCost:
if not team_id:
return None
try:
team_row = await self.prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
team_row: prisma_models.LiteLLM_TeamTable | None = (
await self.prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
)
)
return getattr(team_row, "team_alias", None) if team_row is not None else None
except Exception as e:
@ -138,7 +144,7 @@ class CheckBatchCost:
async def _build_creator_attribution_metadata(
self, job: "LiteLLM_ManagedObjectTable", batch_id: str
) -> Dict[str, Any]:
) -> dict[str, object]:
"""
Rebuild the spend-tracking metadata for the key, team, and tags that created the
batch so the batch-cost spend log is attributed the same way a non-batch request
@ -152,7 +158,7 @@ class CheckBatchCost:
team_id = getattr(job, "team_id", None)
request_tags = getattr(job, "request_tags", None)
metadata: Dict[str, Any] = {
metadata: dict[str, object] = {
"user_api_key_user_id": job.created_by,
"user_api_key": api_key,
"user_api_key_team_id": team_id,

View file

@ -182,6 +182,10 @@ class _ManagedObjectTableActions(Protocol):
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
class _SchedulerWithJobLookup(Protocol):
def get_job(self, job_id: str) -> object: ...
class _CursorPageArgs(TypedDict, total=False):
cursor: Mapping[str, str]
skip: int
@ -853,7 +857,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
file_ids.append(file_id)
return file_ids
def get_file_ids_from_responses_input(self, input: Union[str, List[Dict[str, Any]]]) -> List[str]:
def get_file_ids_from_responses_input(self, input: Union[str, List[Dict[str, object]]]) -> List[str]:
"""
Gets file ids from responses API input.
@ -878,7 +882,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Check for direct input_file type
if item.get("type") == "input_file":
file_id = item.get("file_id")
if file_id:
if isinstance(file_id, str) and file_id:
file_ids.append(file_id)
# Check for input_file in content array
@ -887,7 +891,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
for content_item in content:
if isinstance(content_item, dict) and content_item.get("type") == "input_file":
file_id = content_item.get("file_id")
if file_id:
if isinstance(file_id, str) and file_id:
file_ids.append(file_id)
return file_ids
@ -1227,7 +1231,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Handle both output_file_id and error_file_id
for file_attr in ["output_file_id", "error_file_id"]:
file_id_value = getattr(response, file_attr, None)
file_id_value: str | None = getattr(response, file_attr, None)
if file_id_value and model_id:
decoded_output_file_id = _is_base64_encoded_unified_file_id(file_id_value)
if decoded_output_file_id and "llm_output_file_id," in decoded_output_file_id:
@ -1496,7 +1500,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
import litellm.proxy.proxy_server as proxy_server_module
# Check if the scheduler has the batch cost checking job registered
scheduler = getattr(proxy_server_module, "scheduler", None)
scheduler: Final[_SchedulerWithJobLookup | None] = getattr(proxy_server_module, "scheduler", None)
if scheduler is None:
return False
@ -1542,7 +1546,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
MAX_MATCHES_TO_RETURN = 10
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
batches = await _managed_object_table(self.prisma_client).find_many(
where={
"file_purpose": "batch",
"batch_processed": False,
@ -1552,11 +1556,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
order={"created_at": "desc"},
)
referencing_batches = []
referencing_batches: Final[list[dict[str, object]]] = []
for batch in batches:
try:
# Parse the batch file_object to check for file references
batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object
decoded_file_object = _decode_json_blob(batch.file_object)
batch_data: Mapping[str, object] = (
decoded_file_object if isinstance(decoded_file_object, Mapping) else {}
)
# Extract file IDs from batch
# Batches typically reference the unified file ID in input_file_id

View file

@ -0,0 +1,3 @@
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN IF NOT EXISTS "router_names" TEXT[] NOT NULL DEFAULT ARRAY[]::TEXT[];
ALTER TABLE "LiteLLM_ShadowEvalAttempt" ADD COLUMN IF NOT EXISTS "router_name" TEXT;

View file

@ -1531,7 +1531,8 @@ model LiteLLM_ShadowEvalJob {
group_id String // legs of one job share this; the API's job id
target_type String @default("key") // key | team | user
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
router_name String // the auto-router under evaluation, in either direction
router_name String // first (often only) auto-router under evaluation; router_names is the full set
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String
@ -1555,6 +1556,7 @@ model LiteLLM_ShadowEvalAttempt {
job_id String
request_id String // the judged real request
outcome String // real | shadow | tie | error
router_name String? // the arm this verdict scores; NULL on legacy rows, meaning the job's own router
tier String? // router's tier for the prompt, when classified
real_model String?
shadow_model String?

View file

@ -13,6 +13,7 @@ import json
# s/o [@Frank Colson](https://www.linkedin.com/in/frank-colson-422b9b183/) for this redis implementation
import os
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import Final
from urllib.parse import urlsplit, urlunsplit
@ -38,9 +39,25 @@ from ._logging import verbose_logger
AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default"
def _get_redis_kwargs():
arg_spec: Final = inspect.getfullargspec(redis.Redis)
def _unwrapped_init_args(cls: type) -> frozenset[str]:
"""Every parameter on a single class's own ``__init__``, decorator-unwrapped.
Unlike ``_init_arg_names`` below, this does not walk the MRO: ``redis.Redis``
and ``redis.RedisCluster`` (sync and async) each declare every real
constructor parameter directly on their own ``__init__``, so MRO-walking is
unnecessary — and it actively breaks the several tests here that mock the
class with ``patch(..., autospec=True)``, since ``inspect.getmro`` needs a
real ``__mro__`` that an autospec'd stand-in for a class does not provide.
Still unwraps first: redis-py >= 7.4 decorates these ``__init__``s with
``@deprecated_args`` too, which the same class of bug as ``_init_arg_names``
would otherwise silently empty this allowlist through (see its docstring).
"""
spec: Final = inspect.getfullargspec(inspect.unwrap(cls.__init__))
return frozenset(spec.args + spec.kwonlyargs)
def _get_redis_kwargs():
# Only allow primitive arguments
exclude_args: Final = {
"self",
@ -60,7 +77,7 @@ def _get_redis_kwargs():
"azure_client_secret",
}
available_args: Final = {x for x in arg_spec.args if x not in exclude_args} | include_args
available_args: Final = {x for x in _unwrapped_init_args(redis.Redis) if x not in exclude_args} | include_args
return available_args
@ -120,15 +137,23 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]:
return tuple(x for x in _init_arg_names(connection_cls) if x not in exclude_args) + include_args
def _get_redis_cluster_kwargs(client=None):
def _get_redis_cluster_kwargs(client: type | None = None):
"""Config kwargs the target cluster client's constructor actually accepts.
Defaults to the sync ``redis.RedisCluster``, but the async cluster client
(``redis.asyncio.cluster.RedisCluster``) declares connection settings such as
``decode_responses`` on its own constructor, where the sync class takes them
through ``**kwargs`` and so never names them in its signature. Introspecting
only the sync class regardless of which client is actually built silently
drops those for every async cluster caller.
"""
if client is None:
client = redis.Redis.from_url
arg_spec: Final = inspect.getfullargspec(redis.RedisCluster)
client = redis.RedisCluster
# Only allow primitive arguments
exclude_args: Final = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"}
available_args = {x for x in arg_spec.args if x not in exclude_args}
available_args = {x for x in _unwrapped_init_args(client) if x not in exclude_args}
available_args |= {
"password",
"username",
@ -161,6 +186,79 @@ def _get_redis_env_kwarg_mapping():
return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs() if x not in exclude_from_environment}
def _str_to_bool(value: str) -> bool:
return value.lower() in ("true", "1", "yes")
def _coerce_redis_kwargs_types(
redis_kwargs: Mapping[str, object],
client: type | tuple[type, ...] = redis.Redis,
) -> dict[str, object]: # mutable-ok: a caller mutates the returned kwargs before constructing its client
"""Coerces string values to the numeric/boolean type ``client``'s constructor
declares for that parameter. ``client`` may be a tuple of client classes; a
parameter's type is taken from the first signature that declares it, which
lets cluster callers coerce cluster-only kwargs such as
``cluster_error_retry_attempts`` alongside the shared connection kwargs.
Environment variables are always strings, and Helm ``--set`` stringifies values
too, so a config value like ``health_check_interval`` or ``socket_timeout``
can arrive as ``"30"``/``"5.5"`` rather than a real number. redis-py's own
connection-health-check arithmetic (``loop.time() + self.health_check_interval``)
then raises ``TypeError`` on every Redis operation instead of connecting.
``max_connections``, ``socket_timeout``, and ``socket_connect_timeout`` use an
explicit target type rather than the parameter's own signature default: redis-py
8.x changed the timeout defaults from ``None`` to int ``5``, so inferring the
type from the default would make a fractional ``"5.5"`` fail ``int()`` and get
silently dropped on 8.x while working on older versions. ``socket_keepalive``
is explicit too: its signature default is ``None``, which carries no type to
infer from, and leaving it a string makes ``"false"`` truthy.
"""
signatures: Final = tuple(inspect.signature(c) for c in (client if isinstance(client, tuple) else (client,)))
explicit_param_types: Final = MappingProxyType(
{
"max_connections": int,
"socket_timeout": float,
"socket_connect_timeout": float,
"socket_keepalive": bool,
}
)
result: Final = dict(redis_kwargs) # mutable-ok: per-key try/except coercion below needs to drop individual keys
for key, value in redis_kwargs.items():
if not isinstance(value, str):
continue
param = next((sig.parameters[key] for sig in signatures if key in sig.parameters), None)
if param is None:
continue
explicit_type = explicit_param_types.get(key)
if explicit_type is bool:
result[key] = _str_to_bool(value)
continue
if explicit_type is not None:
try:
result[key] = explicit_type(value)
except (ValueError, TypeError):
del result[key]
continue
default: object = param.default # pyright: ignore[reportAny] # inspect.Parameter.default is stubbed as Any
if default is inspect.Parameter.empty:
continue
# bool must be checked before int, since bool subclasses int
if isinstance(default, bool):
result[key] = _str_to_bool(value)
elif isinstance(default, int):
try:
result[key] = int(value)
except (ValueError, TypeError):
del result[key]
elif isinstance(default, float):
try:
result[key] = float(value)
except (ValueError, TypeError):
del result[key]
return result
def _redis_kwargs_from_environment():
mapping: Final = _get_redis_env_kwarg_mapping()
@ -505,7 +603,12 @@ def _get_redis_client_logic(**env_overrides):
raise ValueError("Either 'host' or 'url' must be specified for redis.")
# litellm.print_verbose(f"redis_kwargs: {redis_kwargs}")
return redis_kwargs
coercion_client: Final = (
(redis.Redis, redis.RedisCluster, async_redis.RedisCluster)
if redis_kwargs.get("startup_nodes")
else redis.Redis
)
return _coerce_redis_kwargs_types(redis_kwargs, client=coercion_client)
def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
@ -657,7 +760,9 @@ def get_redis_client(**env_overrides):
if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs:
return _init_redis_sentinel(redis_kwargs)
return redis.Redis(**redis_kwargs)
return redis.Redis( # pyright: ignore[reportCallIssue] # object-valued kwargs match no overload statically
**redis_kwargs, # pyright: ignore[reportArgumentType] # allow-listed and coerced against this signature
)
def get_redis_async_client(
@ -669,7 +774,7 @@ def get_redis_async_client(
if "startup_nodes" in redis_kwargs:
from redis.cluster import ClusterNode
args = _get_redis_cluster_kwargs()
args = _get_redis_cluster_kwargs(async_redis.RedisCluster)
cluster_kwargs: Final = {}
for arg in redis_kwargs:
if arg in args:

View file

@ -18,7 +18,7 @@ import time
from collections.abc import Awaitable, Callable, Sequence
from contextvars import ContextVar
from datetime import timedelta
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
@ -58,6 +58,26 @@ else:
Span = Any
class _AsyncRedisCommands(Protocol):
"""Async redis commands this cache issues.
redis-py's type stubs omit these methods on RedisCluster, so the union returned by
init_async_client() is untyped at every call site without this protocol.
"""
def ping(self) -> Awaitable[bool]: ...
def delete(self, *names: str) -> Awaitable[int]: ...
def ttl(self, name: str) -> Awaitable[int]: ...
def rpush(self, name: str, *values: str | bytes | float) -> Awaitable[int]: ...
def lpop(self, name: str, count: int | None = None) -> Awaitable[object]: ...
def pipeline(self, transaction: bool = True) -> "Pipeline[bytes]": ...
def _get_call_stack_info(num_frames: int = 2) -> str:
"""
Get the function names from the previous 1-2 functions in the call stack.
@ -429,6 +449,9 @@ class RedisCache(BaseCache):
self.redis_async_client = redis_async_client
return redis_async_client
def _async_commands(self) -> _AsyncRedisCommands:
return self.init_async_client()
def check_and_fix_namespace(self, key: str) -> str:
"""
Make sure each key starts with the given namespace
@ -1055,19 +1078,17 @@ class RedisCache(BaseCache):
await self.async_set_cache_pipeline(self.redis_batch_writing_buffer)
self.redis_batch_writing_buffer = []
def _get_cache_logic(self, cached_response: Any):
def _get_cache_logic(self, cached_response: bytes | str | None):
"""
Common 'get_cache_logic' across sync + async redis client implementations
"""
if cached_response is None:
return cached_response
# cached_response is in `b{} convert it to ModelResponse
cached_response = cached_response.decode("utf-8") # Convert bytes to string
return None
decoded: Final = cached_response.decode("utf-8") if isinstance(cached_response, bytes) else cached_response
try:
cached_response = json.loads(cached_response) # Convert string to dictionary
return json.loads(decoded)
except Exception:
cached_response = ast.literal_eval(cached_response)
return cached_response
return ast.literal_eval(decoded)
def get_cache(self, key, parent_otel_span: Span | None = None, **kwargs):
try:
@ -1314,8 +1335,7 @@ class RedisCache(BaseCache):
raise e
async def ping(self) -> bool:
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ping`
_redis_client: Final[Any] = self.init_async_client()
_redis_client: Final = self._async_commands()
start_time: Final = time.time()
print_verbose("Pinging Async Redis Cache")
try:
@ -1349,8 +1369,7 @@ class RedisCache(BaseCache):
@_redis_circuit_breaker_guard
async def delete_cache_keys(self, keys):
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
_redis_client: Final[Any] = self.init_async_client()
_redis_client: Final = self._async_commands()
keys = [self.check_and_fix_namespace(key=key) for key in keys]
# keys is a list, unpack it so it gets passed as individual elements to delete
await _redis_client.delete(*keys)
@ -1415,8 +1434,7 @@ class RedisCache(BaseCache):
@_redis_circuit_breaker_guard
async def async_delete_cache(self, key: str):
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
_redis_client: Final[Any] = self.init_async_client()
_redis_client: Final = self._async_commands()
key = self.check_and_fix_namespace(key=key)
# keys is str
return await _redis_client.delete(key)
@ -1523,8 +1541,7 @@ class RedisCache(BaseCache):
Redis ref: https://redis.io/docs/latest/commands/ttl/
"""
try:
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ttl`
_redis_client: Final[Any] = self.init_async_client()
_redis_client: Final = self._async_commands()
key = self.check_and_fix_namespace(key=key)
ttl: Final = await _redis_client.ttl(key)
if ttl <= -1: # -1 means the key does not exist, -2 key does not exist
@ -1554,7 +1571,7 @@ class RedisCache(BaseCache):
Returns:
int: The length of the list after the push operation
"""
_redis_client: Final[Any] = self.init_async_client()
_redis_client: Final = self._async_commands()
key = self.check_and_fix_namespace(key=key)
start_time: Final = time.time()
try:
@ -1621,7 +1638,7 @@ class RedisCache(BaseCache):
if len(rpush_list) == 0:
return []
_redis_client: Final[Any] = self.init_async_client()
_redis_client: Final = self._async_commands()
start_time: Final = time.time()
try:
@ -1678,7 +1695,7 @@ class RedisCache(BaseCache):
parent_otel_span: Span | None = None,
**kwargs,
) -> Any | list[Any]:
_redis_client: Final[Any] = self.init_async_client()
_redis_client: Final = self._async_commands()
key = self.check_and_fix_namespace(key=key)
start_time: Final = time.time()
print_verbose(f"LPOP from Redis list: key: {key}, count: {count}")
@ -1810,7 +1827,7 @@ class RedisCache(BaseCache):
if len(lpop_list) == 0:
return []
_redis_client: Final[Any] = self.init_async_client()
_redis_client: Final = self._async_commands()
start_time: Final = time.time()
try:

View file

@ -212,7 +212,8 @@ def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _Ch
LiteLLMCompletionResponsesConfig,
)
is_custom: Final = item.get("type") == "custom_tool_call"
item_type: Final[object] = item.get("type")
is_custom: Final = item_type == "custom_tool_call"
arguments: Final = (item.get("input") if is_custom else item.get("arguments")) or ""
name: Final = item.get("name") or ("custom_tool" if is_custom else "")
function_chunk: Final = ChatCompletionToolCallFunctionChunk(name=name, arguments=arguments)
@ -222,7 +223,7 @@ def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _Ch
function=function_chunk,
index=index,
)
raw_provider_fields: Final = item.get("provider_specific_fields")
raw_provider_fields: Final[object] = item.get("provider_specific_fields")
if isinstance(raw_provider_fields, dict):
provider_specific_fields = raw_provider_fields
elif raw_provider_fields and hasattr(raw_provider_fields, "__dict__"):
@ -507,7 +508,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
def _merge_responses_api_request_into_request_data(
self,
request_data: dict[str, Any],
request_data: dict[str, object],
responses_api_request: "ResponsesAPIOptionalRequestParams",
instructions: str | None,
) -> None:

View file

@ -2268,6 +2268,10 @@ def batch_cost_calculator(
return total_prompt_cost, total_completion_cost
def _attribute_value(obj: object, name: str) -> object:
return getattr(obj, name)
def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> list[str]:
field_names: Final = list(type(prompt_tokens_details).model_fields)
if getattr(prompt_tokens_details, "cache_write_tokens", None) is None:
@ -2293,7 +2297,7 @@ class BaseTokenUsageProcessor:
for usage in usage_objects:
# Handle direct attributes by checking what exists in the model
for attr in dir(usage):
if not attr.startswith("_") and not callable(getattr(usage, attr)):
if not attr.startswith("_") and not callable(_attribute_value(usage, attr)):
current_val = getattr(combined, attr, 0)
new_val = getattr(usage, attr, 0)
if (
@ -2313,7 +2317,7 @@ class BaseTokenUsageProcessor:
if (
hasattr(usage.prompt_tokens_details, attr)
and not attr.startswith("_")
and not callable(getattr(usage.prompt_tokens_details, attr))
and not callable(_attribute_value(usage.prompt_tokens_details, attr))
):
current_val = getattr(combined.prompt_tokens_details, attr, 0) or 0
new_val = getattr(usage.prompt_tokens_details, attr, 0) or 0
@ -2332,7 +2336,9 @@ class BaseTokenUsageProcessor:
# Check what keys exist in the model's completion_tokens_details
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
for attr in type(usage.completion_tokens_details).model_fields:
if not attr.startswith("_") and not callable(getattr(usage.completion_tokens_details, attr)):
if not attr.startswith("_") and not callable(
_attribute_value(usage.completion_tokens_details, attr)
):
current_val = getattr(combined.completion_tokens_details, attr, 0) or 0
new_val = getattr(usage.completion_tokens_details, attr, 0) or 0
if isinstance(new_val, (int, float)):

View file

@ -115,9 +115,11 @@ class SpeechToCompletionBridgeHandler:
**request_data,
)
requested_response_format: Final = optional_params.get("response_format")
if isinstance(result, ModelResponse):
return self.transformation_handler.transform_response(
model_response=result,
response_format=requested_response_format if isinstance(requested_response_format, str) else None,
)
else:
raise Exception(f"Unmapped response type. Got type: {type(result)}")

View file

@ -21,6 +21,8 @@ def _completion_response_cost(model_response: "ModelResponse") -> float | None:
GEMINI_TTS_CHAT_AUDIO_FORMAT: Final = "pcm16"
GEMINI_TTS_RAW_RESPONSE_FORMAT: Final = "pcm"
GEMINI_TTS_SUPPORTED_RESPONSE_FORMATS: Final = frozenset({"wav", GEMINI_TTS_RAW_RESPONSE_FORMAT})
class ChatAudioParam(TypedDict):
@ -29,6 +31,26 @@ class ChatAudioParam(TypedDict):
class SpeechToCompletionBridgeTransformationHandler:
def _validate_response_format(
self, model: str, custom_llm_provider: str, optional_params: Mapping[str, object]
) -> None:
if not self._is_gemini_tts_model(model):
return
response_format: Final = optional_params.get("response_format")
if not isinstance(response_format, str) or response_format in GEMINI_TTS_SUPPORTED_RESPONSE_FORMATS:
return
from litellm.exceptions import BadRequestError
supported: Final = ", ".join(sorted(GEMINI_TTS_SUPPORTED_RESPONSE_FORMATS))
raise BadRequestError(
message=(
f"Gemini TTS only produces raw PCM16 audio, so response_format='{response_format}'"
f" is not supported. Supported response formats: {supported}."
),
model=model,
llm_provider=custom_llm_provider,
)
def _chat_completion_params(self, optional_params: Mapping[str, object]) -> Mapping[str, object]:
return MappingProxyType(
{
@ -67,6 +89,7 @@ class SpeechToCompletionBridgeTransformationHandler:
litellm_logging_obj: "LiteLLMLoggingObj",
custom_llm_provider: str,
) -> dict:
self._validate_response_format(model, custom_llm_provider, optional_params)
user_message: Final[ChatCompletionUserMessage] = {"role": "user", "content": input}
return_kwargs: Final = {
"model": model,
@ -125,7 +148,14 @@ class SpeechToCompletionBridgeTransformationHandler:
"""Check if the model is a Gemini TTS model that returns PCM16 data."""
return "gemini" in model.lower() and ("tts" in model.lower() or "preview-tts" in model.lower())
def transform_response(self, model_response: "ModelResponse") -> "HttpxBinaryResponseContent":
def _gemini_tts_response_body(self, decoded_audio: bytes, response_format: str | None) -> tuple[bytes, str]:
if response_format == GEMINI_TTS_RAW_RESPONSE_FORMAT:
return decoded_audio, "audio/pcm"
return self._convert_pcm16_to_wav(decoded_audio), "audio/wav"
def transform_response(
self, model_response: "ModelResponse", response_format: str | None
) -> "HttpxBinaryResponseContent":
import base64
import httpx
@ -136,23 +166,17 @@ class SpeechToCompletionBridgeTransformationHandler:
audio_part: Final = cast(Choices, model_response.choices[0]).message.audio
if audio_part is None:
raise ValueError("No audio part found in the response")
audio_content: Final = audio_part.data
decoded_audio: Final = base64.b64decode(audio_part.data)
# Decode base64 to get binary content
binary_data = base64.b64decode(audio_content)
# Check if this is a Gemini TTS model that returns raw PCM16 data
model: Final = getattr(model_response, "model", "")
headers: Final = {}
if self._is_gemini_tts_model(model):
# Convert PCM16 to WAV format for proper audio file playback
binary_data = self._convert_pcm16_to_wav(binary_data)
headers["Content-Type"] = "audio/wav"
else:
headers["Content-Type"] = "audio/mpeg"
# Create an httpx.Response object
response: Final = httpx.Response(status_code=200, content=binary_data, headers=headers)
content, content_type = (
self._gemini_tts_response_body(decoded_audio, response_format)
if self._is_gemini_tts_model(model)
else (decoded_audio, "audio/mpeg")
)
response: Final = httpx.Response(
status_code=200, content=content, headers=MappingProxyType({"Content-Type": content_type})
)
binary_response: Final = HttpxBinaryResponseContent(response)
binary_response.set_response_cost(_completion_response_cost(model_response))
return binary_response

View file

@ -722,7 +722,7 @@ class GoogleGenAIAdapter:
)
for tool_call in tool_calls:
if not hasattr(tool_call, "function"):
if not hasattr(tool_call, "function") or isinstance(tool_call, ChatCompletionDeltaCustomToolCall):
continue
# 3. Use `index` as the primary key for accumulation

View file

@ -3,10 +3,12 @@ Arize Phoenix prompt manager that integrates with LiteLLM's prompt management sy
Fetches prompt versions from Arize Phoenix and provides workspace-based access control.
"""
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Any, Final, cast
from jinja2 import DictLoader, select_autoescape
from jinja2.sandbox import ImmutableSandboxedEnvironment
from typing_extensions import ReadOnly, TypedDict
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.integrations.prompt_management_base import (
@ -20,6 +22,31 @@ from litellm.types.utils import StandardCallbackDynamicParams
from .arize_phoenix_client import ArizePhoenixClient
class ArizePhoenixContentPart(TypedDict, total=False):
type: ReadOnly[str]
text: ReadOnly[str]
class ArizePhoenixTemplateMessage(TypedDict, total=False):
role: ReadOnly[str]
content: ReadOnly[Sequence[ArizePhoenixContentPart]]
class ArizePhoenixTemplateBody(TypedDict, total=False):
messages: ReadOnly[Sequence[ArizePhoenixTemplateMessage]]
class ArizePhoenixPromptMetadata(TypedDict):
model_name: ReadOnly[str | None]
model_provider: ReadOnly[str | None]
description: ReadOnly[str]
template_type: ReadOnly[str | None]
template_format: ReadOnly[str]
invocation_parameters: ReadOnly[Mapping[str, Mapping[str, object]]]
temperature: ReadOnly[float | None]
max_tokens: ReadOnly[int | None]
class ArizePhoenixPromptTemplate:
"""
Represents a prompt template loaded from Arize Phoenix.
@ -28,10 +55,10 @@ class ArizePhoenixPromptTemplate:
def __init__(
self,
template_id: str,
messages: list[dict[str, Any]],
metadata: dict[str, Any],
messages: Sequence[ArizePhoenixTemplateMessage],
metadata: ArizePhoenixPromptMetadata,
model: str | None = None,
):
) -> None:
self.template_id = template_id
self.messages = messages
self.metadata = metadata
@ -43,7 +70,7 @@ class ArizePhoenixPromptTemplate:
self.description = metadata.get("description", "")
self.template_format = metadata.get("template_format", "MUSTACHE")
def __repr__(self):
def __repr__(self) -> str:
return f"ArizePhoenixPromptTemplate(id='{self.template_id}', model='{self.model}')"
@ -109,7 +136,7 @@ class ArizePhoenixTemplateManager:
def _parse_prompt_data(self, data: dict[str, Any], prompt_version_id: str) -> ArizePhoenixPromptTemplate:
"""Parse Arize Phoenix prompt data and extract messages and metadata."""
template_data: Final = data.get("template", {})
template_data: Final[ArizePhoenixTemplateBody] = data.get("template", {})
messages: Final = template_data.get("messages", [])
# Extract invocation parameters
@ -129,7 +156,7 @@ class ArizePhoenixTemplateManager:
break
# Build metadata dictionary
metadata: Final = {
metadata: Final[ArizePhoenixPromptMetadata] = {
"model_name": data.get("model_name"),
"model_provider": data.get("model_provider"),
"description": data.get("description", ""),
@ -146,7 +173,9 @@ class ArizePhoenixTemplateManager:
metadata=metadata,
)
def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> list[AllMessageValues]:
def render_template(
self, template_id: str, variables: Mapping[str, object] | None = None
) -> list[AllMessageValues]:
"""Render a template with the given variables and return formatted messages."""
if template_id not in self.prompts:
raise ValueError(f"Template '{template_id}' not found")
@ -174,7 +203,9 @@ class ArizePhoenixTemplateManager:
# Combine rendered content
final_content = " ".join(rendered_content_parts)
rendered_messages.append({"role": role, "content": final_content})
rendered_messages.append(
cast("AllMessageValues", {"role": role, "content": final_content}) # cast-ok: Phoenix roles are OpenAI
)
return rendered_messages
@ -243,8 +274,8 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
def get_prompt_template(
self,
prompt_id: str,
prompt_variables: dict[str, Any] | None = None,
) -> tuple[list[AllMessageValues], dict[str, Any]]:
prompt_variables: Mapping[str, object] | None = None,
) -> tuple[list[AllMessageValues], dict[str, object]]:
"""
Get a prompt template and render it with variables.
@ -263,7 +294,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
rendered_messages: Final = self.prompt_manager.render_template(prompt_id, prompt_variables or {})
# Extract metadata
metadata: Final = {
metadata: Final[dict[str, object]] = {
"model": template.model,
"temperature": template.temperature,
"max_tokens": template.max_tokens,
@ -271,7 +302,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
# Add additional invocation parameters
invocation_params: Final = template.invocation_parameters
provider_params = {}
provider_params: Mapping[str, object] = {}
if "openai" in invocation_params:
provider_params = invocation_params["openai"]
@ -289,12 +320,12 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
self,
user_id: str | None,
messages: list[AllMessageValues],
function_call: dict[str, Any] | str | None = None,
litellm_params: dict[str, Any] | None = None,
function_call: dict[str, object] | str | None = None,
litellm_params: dict[str, object] | None = None,
prompt_id: str | None = None,
prompt_variables: dict[str, Any] | None = None,
prompt_variables: dict[str, object] | None = None,
**kwargs,
) -> tuple[list[AllMessageValues], dict[str, Any] | None]:
) -> tuple[list[AllMessageValues], dict[str, object] | None]:
"""
Pre-call hook that processes the prompt template before making the LLM call.
"""
@ -335,9 +366,9 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
except Exception as e:
# Log error but don't fail the call
import litellm
from litellm._logging import verbose_proxy_logger
litellm._logging.verbose_proxy_logger.error("Error in Arize Phoenix prompt pre_call_hook: %s", e)
verbose_proxy_logger.error("Error in Arize Phoenix prompt pre_call_hook: %s", e)
return messages, litellm_params
def get_available_prompts(self) -> list[str]:
@ -393,7 +424,8 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
rendered_messages, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)
# Extract model from metadata (if specified)
template_model: Final = prompt_metadata.get("model")
raw_template_model: Final = prompt_metadata.get("model")
template_model: Final = raw_template_model if isinstance(raw_template_model, str) else None
# Extract optional parameters from metadata
optional_params: Final = {}

View file

@ -31,6 +31,9 @@ if TYPE_CHECKING:
from litellm.caching.caching import DualCache
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.anthropic_messages.transformation import (
BaseAnthropicMessagesConfig,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp import (
MCPPostCallResponseObject,
@ -39,7 +42,7 @@ if TYPE_CHECKING:
)
from litellm.types.router import PreRoutingHookResponse
Span = _Span | Any
Span = _Span
else:
Span = Any
LiteLLMLoggingObj = Any
@ -268,7 +271,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
) -> list[dict]:
return healthy_deployments
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
async def async_pre_call_deployment_hook(
self, kwargs: dict[str, object], call_type: CallTypes | None
) -> dict | None:
"""
Allow modifying the request just before it's sent to the deployment.
@ -344,9 +349,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
async def async_post_call_streaming_deployment_hook(
self,
request_data: dict,
response_chunk: Any,
response_chunk: object,
call_type: CallTypes | None,
) -> Any | None:
) -> object | None:
"""
Allow modifying streaming chunks just before they're returned to the user.
@ -378,7 +383,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
"""
def translate_completion_output_params_streaming(
self, completion_stream: Any
self, completion_stream: object
) -> AdapterCompletionStreamWrapper | None:
"""
Translates the streaming chunk, from the OpenAI format to the custom format.
@ -418,9 +423,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
response: Any,
response: object,
request_headers: dict[str, str] | None = None,
litellm_call_info: dict[str, Any] | None = None,
litellm_call_info: dict[str, object] | None = None,
) -> dict[str, str] | None:
"""
Called after an LLM API call (success or failure) to allow injecting custom HTTP response headers.
@ -471,11 +476,11 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
) -> Any:
pass
async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]:
async def async_logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
"""For masking logged request/response. Return a modified version of the request/result."""
return kwargs, result
def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]:
def logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
"""For masking logged request/response. Return a modified version of the request/result."""
return kwargs, result
@ -581,7 +586,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
async def async_should_run_agentic_loop(
self,
response: Any,
response: object,
model: str,
messages: list[dict],
tools: list[dict] | None,
@ -642,8 +647,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
tools: dict,
model: str,
messages: list[dict],
response: Any,
anthropic_messages_provider_config: Any,
response: object,
anthropic_messages_provider_config: "BaseAnthropicMessagesConfig | None",
anthropic_messages_optional_request_params: dict,
logging_obj: "LiteLLMLoggingObj",
stream: bool,
@ -711,8 +716,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
tools: dict,
model: str,
messages: list[dict],
response: Any,
anthropic_messages_provider_config: Any,
response: object,
anthropic_messages_provider_config: "BaseAnthropicMessagesConfig | None",
anthropic_messages_optional_request_params: dict,
logging_obj: "LiteLLMLoggingObj",
stream: bool,
@ -728,7 +733,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
async def async_post_agentic_loop_response_hook(
self,
response: Any,
response: object,
plan: AgenticLoopPlan,
kwargs: dict,
) -> Any:
@ -767,7 +772,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
async def async_should_run_chat_completion_agentic_loop(
self,
response: Any,
response: object,
model: str,
messages: list[dict],
tools: list[dict] | None,
@ -785,12 +790,12 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
tools: dict,
model: str,
messages: list[dict],
response: Any,
response: object,
optional_params: dict,
logging_obj: "LiteLLMLoggingObj",
stream: bool,
kwargs: dict,
) -> Any:
) -> object:
"""
Hook to execute chat completion agentic loop based on context from should_run hook.
"""
@ -800,7 +805,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
tools: dict,
model: str,
messages: list[dict],
response: Any,
response: object,
optional_params: dict,
logging_obj: "LiteLLMLoggingObj",
stream: bool,
@ -1056,7 +1061,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
def _redact_base64(
self,
value: Any,
value: object,
depth: int = 0,
max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER,
) -> object:
@ -1079,7 +1084,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
return value
def _should_keep_content(self, content: Any) -> bool:
def _should_keep_content(self, content: object) -> bool:
"""Return True if this content item should be retained."""
if not isinstance(content, dict):
return True

View file

@ -2,10 +2,12 @@
GitLab prompt manager with configurable prompts folder.
"""
from typing import TYPE_CHECKING, Any, Final
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, TypeVar
from jinja2 import DictLoader, select_autoescape
from jinja2.sandbox import ImmutableSandboxedEnvironment
from typing_extensions import ReadOnly, TypedDict
from litellm.integrations.custom_prompt_management import CustomPromptManagement
@ -24,6 +26,19 @@ from litellm.types.utils import StandardCallbackDynamicParams
GITLAB_PREFIX: Final = "gitlab::"
_ResponseT = TypeVar("_ResponseT")
class GitLabCachedPrompt(TypedDict):
id: ReadOnly[str]
path: ReadOnly[str]
content: ReadOnly[str]
metadata: ReadOnly[Mapping[str, object]]
model: ReadOnly[str | None]
temperature: ReadOnly[float | None]
max_tokens: ReadOnly[int | None]
optional_params: ReadOnly[Mapping[str, object]]
def encode_prompt_id(raw_id: str) -> str:
"""Convert GitLab path IDs like 'invoice/extract' → 'gitlab::invoice::extract'"""
@ -206,7 +221,7 @@ class GitLabTemplateManager:
result[key] = value.strip("\"'")
return result
def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> str:
def render_template(self, template_id: str, variables: Mapping[str, object] | None = None) -> str:
if template_id not in self.prompts:
raise ValueError(f"Template '{template_id}' not found")
template: Final = self.prompts[template_id]
@ -313,7 +328,7 @@ class GitLabPromptManager(CustomPromptManagement):
def get_prompt_template(
self,
prompt_id: str,
prompt_variables: dict[str, Any] | None = None,
prompt_variables: Mapping[str, object] | None = None,
*,
ref: str | None = None,
) -> tuple[str, dict[str, Any]]:
@ -338,13 +353,13 @@ class GitLabPromptManager(CustomPromptManagement):
self,
user_id: str | None,
messages: list[AllMessageValues],
function_call: dict[str, Any] | str | None = None,
litellm_params: dict[str, Any] | None = None,
function_call: Mapping[str, object] | str | None = None,
litellm_params: dict[str, object] | None = None,
prompt_id: str | None = None,
prompt_variables: dict[str, Any] | None = None,
prompt_variables: Mapping[str, object] | None = None,
prompt_version: str | None = None,
**kwargs,
) -> tuple[list[AllMessageValues], dict[str, Any] | None]:
) -> tuple[list[AllMessageValues], dict[str, object] | None]:
if not prompt_id:
return messages, litellm_params
try:
@ -377,9 +392,9 @@ class GitLabPromptManager(CustomPromptManagement):
return final_messages, litellm_params
except Exception as e:
import litellm
from litellm._logging import verbose_proxy_logger
litellm._logging.verbose_proxy_logger.error("Error in GitLab prompt pre_call_hook: %s", e)
verbose_proxy_logger.error("Error in GitLab prompt pre_call_hook: %s", e)
return messages, litellm_params
def _parse_prompt_to_messages(self, prompt_content: str) -> list[AllMessageValues]:
@ -435,14 +450,14 @@ class GitLabPromptManager(CustomPromptManagement):
def post_call_hook(
self,
user_id: str | None,
response: Any,
response: _ResponseT,
input_messages: list[AllMessageValues],
function_call: dict[str, Any] | str | None = None,
litellm_params: dict[str, Any] | None = None,
function_call: Mapping[str, object] | str | None = None,
litellm_params: Mapping[str, object] | None = None,
prompt_id: str | None = None,
prompt_variables: dict[str, Any] | None = None,
prompt_variables: Mapping[str, object] | None = None,
**kwargs,
) -> Any:
) -> _ResponseT:
return response
def get_available_prompts(self) -> list[str]:
@ -498,7 +513,7 @@ class GitLabPromptManager(CustomPromptManagement):
messages: Final = self._parse_prompt_to_messages(rendered_prompt)
template_model: Final = prompt_metadata.get("model")
optional_params: Final[dict[str, Any]] = {}
optional_params: Final[dict[str, object]] = {}
for param in [
"temperature",
"max_tokens",
@ -658,14 +673,14 @@ class GitLabPromptCache:
self.template_manager: GitLabTemplateManager = self.prompt_manager.prompt_manager
# In-memory stores
self._by_file: dict[str, dict[str, Any]] = {}
self._by_id: dict[str, dict[str, Any]] = {}
self._by_file: dict[str, GitLabCachedPrompt] = {}
self._by_id: dict[str, GitLabCachedPrompt] = {}
# -------------------------
# Public API
# -------------------------
def load_all(self, *, recursive: bool = True) -> dict[str, dict[str, Any]]:
def load_all(self, *, recursive: bool = True) -> dict[str, GitLabCachedPrompt]:
"""
Scan GitLab for all .prompt files under prompts_path, load and parse each,
and return the mapping of repo file path -> JSON-like dict.
@ -695,7 +710,7 @@ class GitLabPromptCache:
return self._by_id
def reload(self, *, recursive: bool = True) -> dict[str, dict[str, Any]]:
def reload(self, *, recursive: bool = True) -> dict[str, GitLabCachedPrompt]:
"""Clear the cache and re-load from GitLab."""
self._by_file.clear()
self._by_id.clear()
@ -709,11 +724,11 @@ class GitLabPromptCache:
"""Return the template IDs (relative to prompts_path, without extension) currently cached."""
return list(self._by_id.keys())
def get_by_file(self, file_path: str) -> dict[str, Any] | None:
def get_by_file(self, file_path: str) -> GitLabCachedPrompt | None:
"""Get a cached prompt JSON by repo file path."""
return self._by_file.get(file_path)
def get_by_id(self, prompt_id: str) -> dict[str, Any] | None:
def get_by_id(self, prompt_id: str) -> GitLabCachedPrompt | None:
"""Get a cached prompt JSON by prompt ID (relative to prompts_path)."""
if prompt_id in self._by_id:
return self._by_id[prompt_id]
@ -728,7 +743,7 @@ class GitLabPromptCache:
# Internals
# -------------------------
def _template_to_json(self, prompt_id: str, tmpl: GitLabPromptTemplate) -> dict[str, Any]:
def _template_to_json(self, prompt_id: str, tmpl: GitLabPromptTemplate) -> GitLabCachedPrompt:
"""
Normalize a GitLabPromptTemplate into a JSON-like dict that is easy to serialize.
"""

View file

@ -12,7 +12,10 @@ For batching specific details see CustomBatchLogger class
import asyncio
import atexit
import os
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Final
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_logger
from litellm._uuid import uuid
@ -34,6 +37,21 @@ from litellm.types.integrations.posthog import (
from litellm.types.utils import StandardCallbackDynamicParams, StandardLoggingPayload
class PostHogBatchPayload(TypedDict):
api_key: ReadOnly[str]
batch: ReadOnly[Sequence[PostHogEventPayload]]
class PostHogLiteLLMParams(TypedDict, total=False):
metadata: ReadOnly[Mapping[str, object]]
class PostHogLogKwargs(TypedDict, total=False):
standard_logging_object: ReadOnly[StandardLoggingPayload]
standard_callback_dynamic_params: ReadOnly[StandardCallbackDynamicParams]
litellm_params: ReadOnly[PostHogLiteLLMParams]
class PostHogLogger(CustomBatchLogger):
def __init__(self, **kwargs):
"""
@ -137,7 +155,7 @@ class PostHogLogger(CustomBatchLogger):
if len(self.log_queue) >= self.batch_size:
await self.flush_queue()
def create_posthog_event_payload(self, kwargs: dict[str, Any]) -> PostHogEventPayload:
def create_posthog_event_payload(self, kwargs: PostHogLogKwargs) -> PostHogEventPayload:
"""
Helper function to create a PostHog event payload for logging
@ -171,11 +189,11 @@ class PostHogLogger(CustomBatchLogger):
def _create_posthog_properties(
self,
standard_logging_object: StandardLoggingPayload,
kwargs: dict[str, Any],
kwargs: PostHogLogKwargs,
event_name: str,
) -> dict[str, Any]:
) -> dict[str, object]:
"""Create PostHog properties following LLM Analytics spec"""
properties: Final = {}
properties: Final[dict[str, object]] = {}
# Core model information
properties["$ai_model"] = self._safe_get(standard_logging_object, "model", "")
@ -211,16 +229,19 @@ class PostHogLogger(CustomBatchLogger):
properties["$ai_error"] = error_str
# Add trace properties
self._add_trace_properties(properties, kwargs)
self._add_trace_properties(properties, standard_logging_object, kwargs)
# Add custom metadata fields
self._add_custom_metadata_properties(properties, kwargs)
return properties
def _add_trace_properties(self, properties: dict[str, Any], kwargs: dict[str, Any]):
standard_logging_object: Final = self._safe_get(kwargs, "standard_logging_object", {})
def _add_trace_properties(
self,
properties: dict[str, object],
standard_logging_object: StandardLoggingPayload,
kwargs: PostHogLogKwargs,
) -> None:
trace_id: Final = self._safe_get(standard_logging_object, "trace_id", self._safe_uuid())
properties["$ai_trace_id"] = trace_id
@ -232,7 +253,7 @@ class PostHogLogger(CustomBatchLogger):
if parent_id:
properties["$ai_parent_id"] = parent_id
def _add_custom_metadata_properties(self, properties: dict[str, Any], kwargs: dict[str, Any]):
def _add_custom_metadata_properties(self, properties: dict[str, object], kwargs: PostHogLogKwargs) -> None:
"""Add custom metadata fields to PostHog properties"""
metadata: Final = self._extract_metadata(kwargs)
if not isinstance(metadata, dict):
@ -277,7 +298,7 @@ class PostHogLogger(CustomBatchLogger):
if key not in litellm_internal_fields:
properties[key] = value
def _get_distinct_id(self, standard_logging_object: StandardLoggingPayload, kwargs: dict[str, Any]) -> str:
def _get_distinct_id(self, standard_logging_object: StandardLoggingPayload, kwargs: PostHogLogKwargs) -> str:
metadata: Final = self._extract_metadata(kwargs)
user_id: Final = self._safe_get(metadata, "user_id")
if user_id:
@ -291,7 +312,7 @@ class PostHogLogger(CustomBatchLogger):
return self._safe_uuid()
def _get_credentials_for_request(self, kwargs: dict[str, Any]) -> tuple[str | None, str | None]:
def _get_credentials_for_request(self, kwargs: PostHogLogKwargs) -> tuple[str | None, str | None]:
"""
Get PostHog credentials for this request.
@ -334,7 +355,7 @@ class PostHogLogger(CustomBatchLogger):
verbose_logger.debug("[POSTHOG MOCK] Mock mode enabled - API calls will be intercepted")
# Group events by credentials for batch sending
batches_by_credentials: Final[dict[tuple[str, str], list]] = {}
batches_by_credentials: Final[dict[tuple[str, str], list[PostHogEventPayload]]] = {}
for item in self.log_queue:
key = (item["api_key"], item["api_url"])
if key not in batches_by_credentials:
@ -380,18 +401,19 @@ class PostHogLogger(CustomBatchLogger):
verbose_logger.error("PostHog: Failed to initialize async components: %s", e)
raise
def _extract_metadata(self, kwargs: dict[str, Any]) -> dict[str, Any]:
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
return litellm_params.get("metadata", {}) or {}
def _extract_metadata(self, kwargs: PostHogLogKwargs) -> Mapping[str, object]:
litellm_params: Final[PostHogLiteLLMParams] = kwargs.get("litellm_params", {}) or {}
metadata: Final[Mapping[str, object]] = litellm_params.get("metadata", {}) or {}
return metadata
def _safe_uuid(self) -> str:
return str(uuid.uuid4())
def _create_posthog_payload(self, events: list, api_key: str) -> dict[str, Any]:
def _create_posthog_payload(self, events: Sequence[PostHogEventPayload], api_key: str) -> PostHogBatchPayload:
return {"api_key": api_key, "batch": events}
def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any:
if obj is None or not hasattr(obj, "get"):
def _safe_get(self, obj: Mapping[str, object] | None, key: str, default: object = None) -> object:
if not isinstance(obj, Mapping):
return default
return obj.get(key, default)
@ -412,7 +434,7 @@ class PostHogLogger(CustomBatchLogger):
try:
# Group events by credentials (same logic as async_send_batch)
batches_by_credentials: Final[dict[tuple[str, str], list]] = {}
batches_by_credentials: Final[dict[tuple[str, str], list[PostHogEventPayload]]] = {}
for item in self.log_queue:
key = (item["api_key"], item["api_url"])
if key not in batches_by_credentials:

View file

@ -1,8 +1,11 @@
"""Shadow Eval Logger: samples a shadowed key's successful LLM requests (chat completions,
Anthropic Messages, and Responses API surfaces, each normalized to chat shape), duplicates
each against the job's other arm in a detached task (the auto-router for a forward job, the
fixed baseline model for a reverse one), blind-judges real vs shadow, and appends one
``LiteLLM_ShadowEvalAttempt`` row (verdict or error) as the feature's only hot-path write.
each through every shadow arm in one detached task (each candidate auto-router for a
forward job, the fixed baseline model for a reverse one), blind-judges real vs each arm,
and appends one ``LiteLLM_ShadowEvalAttempt`` row per arm (verdict or error) as the
feature's only hot-path write. A multi-router job's arms therefore score the identical
sampled requests against the identical real responses, which is what makes their win
rates comparable head-to-head.
Counts, status, and spend derive from those rows at read time, so nothing can disagree
across pods or stop races; the hook reads active jobs through a short-TTL cache."""
@ -498,12 +501,16 @@ def _decision_classifier_cost(metadata: Mapping[str, object]) -> float:
return float(raw) if isinstance(raw, (int, float)) else 0.0
def _request_was_routed_by(request_metadata: Mapping[str, object], router_name: str) -> bool:
"""Whether the router under evaluation served this request, which is what decides
the direction it belongs to. A forward job skips its own router's traffic, since
duplicating it would compare the router to itself: guaranteed ties, judge spend for
zero information. A reverse job samples exactly that traffic and nothing else."""
return _routing_decision(request_metadata).get("router_model_name") == router_name
def _direction_admits(request_metadata: Mapping[str, object], job: "ActiveShadowEvalJob") -> bool:
"""Whether this request belongs to the job's direction. A forward job skips traffic
any of its candidate routers served: duplicating a router's own request compares it
to itself (guaranteed ties), and judging a sibling against another candidate's live
response would score candidates against each other instead of against the incumbent.
A reverse job samples exactly its one router's traffic and nothing else."""
routed_by: Final = _routing_decision(request_metadata).get("router_model_name")
if job.direction == "reverse":
return routed_by == job.router_name
return routed_by not in job.arm_router_names
@dataclass(frozen=True, slots=True)
@ -546,6 +553,7 @@ class ActiveShadowEvalJob(BaseModel):
id: str
router_name: str
router_names: tuple[str, ...] = ()
direction: ShadowEvalDirection = "forward"
baseline_model: str | None = None
shadow_percentage: float
@ -567,12 +575,25 @@ class ActiveShadowEvalJob(BaseModel):
raise ValueError("baseline_model is set for exactly the reverse jobs")
return self
@model_validator(mode="after")
def _reverse_evaluates_one_router(self) -> "ActiveShadowEvalJob":
"""A reverse row naming several routers is unsamplable (there is no one traffic
slice they share) and fails closed."""
if self.direction == "reverse" and len(self.arm_router_names) > 1:
raise ValueError("a reverse job evaluates exactly one router")
return self
@property
def shadow_target(self) -> str:
"""The model the duplicated arm calls: the router itself for a forward job, the
fixed baseline for a reverse one. Total because the validator above pins
def arm_router_names(self) -> tuple[str, ...]:
"""The job's full router set; rows from before router_names existed hold it in
router_name alone. The one place that reading lives on the sampling side."""
return self.router_names or (self.router_name,)
def arm_target(self, arm_router: str) -> str:
"""The model one duplicated arm calls: the candidate router itself for a forward
job, the fixed baseline for a reverse one. Total because the validator above pins
baseline_model to reverse jobs and only those."""
return self.baseline_model or self.router_name
return self.baseline_model or arm_router
def _as_active_job(record: object, attempts: int, spend: float) -> ActiveShadowEvalJob | None:
@ -696,7 +717,7 @@ class ShadowEvalLogger(CustomLogger):
now >= job.ends_at
or job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns
or (job.max_budget is not None and job.spend >= job.max_budget)
or _request_was_routed_by(request_metadata, job.router_name) != (job.direction == "reverse")
or not _direction_admits(request_metadata, job)
):
continue
if not _sample_hits(request_id, job.id, job.shadow_percentage):
@ -773,7 +794,10 @@ class ShadowEvalLogger(CustomLogger):
if self._inflight_shadow_tasks >= _MAX_CONCURRENT_SHADOW_TASKS:
self._record_funnel(job.id, "shed")
continue
self._job_starts[job.id] = self._job_starts.get(job.id, 0) + 1
# One start writes one attempt row per arm, and max_turns is a row
# ceiling, so admission must pre-count every arm or a multi-router
# job overshoots the valve N-fold within a cache generation.
self._job_starts[job.id] = self._job_starts.get(job.id, 0) + len(job.arm_router_names)
self._inflight_shadow_tasks += 1
asyncio.create_task(
self._run_shadow_eval(
@ -812,32 +836,74 @@ class ShadowEvalLogger(CustomLogger):
shadow_params: Mapping[str, object],
parent_metadata: Mapping[str, object],
) -> None:
"""Budget gate -> shadow call -> blind judge -> one attempt row, and every exit
in exactly one coverage bucket: the gates that decline to spend on an admitted
sample (no DB to record into, an over-budget key, an unverifiable or exhausted
eval budget) count it withheld, so eligible traffic still reconciles as
not_sampled + unjudgeable + shed + withheld + attempt rows. The prisma gate sits
above the dispatch so no provider spend happens without a place to record the
outcome, and the budget read lives here rather than in the success hook."""
"""Budget gates once per sampled request, then every router arm in turn: shadow
call -> blind judge -> one attempt row stamped with the arm. The gates that
decline to spend on an admitted sample (no DB to record into, an over-budget key,
an unverifiable or exhausted eval budget) count the REQUEST withheld before any
arm runs, so funnel counters stay per-request and a leg's eligible traffic still
reconciles as not_sampled + unjudgeable + shed + withheld + sampled requests,
where each sampled request writes one attempt row per arm. A budget crossed
mid-loop lets the remaining arms overshoot by one round, the same class of
overshoot as the samples already in flight when the cap is crossed. The prisma
gate sits above the dispatch so no provider spend happens without a place to
record the outcome, and the budget read lives here rather than in the success
hook."""
prisma: Final = self._prisma_provider()
if prisma is None:
self._record_funnel(job.id, "withheld")
return
if await _key_or_team_is_over_budget(parent_metadata):
self._record_funnel(job.id, "withheld")
return
if job.max_budget is not None:
try:
spend: Final = await self._read_job_spend(_job_spend_counter_key(job.id), job.spend, job.max_budget)
except Exception as e: # noqa: BLE001 # unverifiable budget: skip the sample rather than spend on it
verbose_logger.warning("shadow_eval: budget unverifiable for %s, sample skipped: %s", job.id, e)
self._record_funnel(job.id, "withheld")
return
if spend >= job.max_budget:
self._record_funnel(job.id, "withheld")
return
for arm_router in job.arm_router_names:
await self._run_shadow_arm(
prisma=prisma,
job=job,
arm_router=arm_router,
request_id=request_id,
messages=messages,
real_text=real_text,
real_model=real_model,
real_cost=real_cost,
real_classifier_cost=real_classifier_cost,
real_cache_hit=real_cache_hit,
control_tier=control_tier,
shadow_params=shadow_params,
parent_metadata=parent_metadata,
)
async def _run_shadow_arm(
self,
prisma: "PrismaClient",
job: ActiveShadowEvalJob,
arm_router: str,
request_id: str,
messages: Sequence[Mapping[str, object]],
real_text: str,
real_model: str,
real_cost: float,
real_classifier_cost: float,
real_cache_hit: bool,
control_tier: str | None,
shadow_params: Mapping[str, object],
parent_metadata: Mapping[str, object],
) -> None:
"""One arm's pipeline: shadow call -> blind judge -> one attempt row, every exit
recording this arm's outcome, so one arm's fault never silences a sibling arm."""
try:
if prisma is None:
self._record_funnel(job.id, "withheld")
return
if await _key_or_team_is_over_budget(parent_metadata):
self._record_funnel(job.id, "withheld")
return
if job.max_budget is not None:
try:
spend: Final = await self._read_job_spend(_job_spend_counter_key(job.id), job.spend, job.max_budget)
except Exception as e: # noqa: BLE001 # unverifiable budget: skip the sample rather than spend on it
verbose_logger.warning("shadow_eval: budget unverifiable for %s, sample skipped: %s", job.id, e)
self._record_funnel(job.id, "withheld")
return
if spend >= job.max_budget:
self._record_funnel(job.id, "withheld")
return
shadow: Final = await self._call_router_shadow(job.shadow_target, messages, shadow_params, parent_metadata)
shadow: Final = await self._call_router_shadow(
job.arm_target(arm_router), messages, shadow_params, parent_metadata
)
except Exception as e: # noqa: BLE001 # detached task: nothing billed yet, record and never raise
verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e)
await self._record_attempt(
@ -845,6 +911,7 @@ class ShadowEvalLogger(CustomLogger):
job,
request_id,
control_tier,
router_name=arm_router,
outcome="error",
error=f"pipeline error: {e}",
real_cost=real_cost,
@ -858,6 +925,7 @@ class ShadowEvalLogger(CustomLogger):
job,
request_id,
control_tier,
router_name=arm_router,
outcome="error",
error=shadow.error,
shadow_cost=shadow.cost,
@ -882,6 +950,7 @@ class ShadowEvalLogger(CustomLogger):
job,
request_id,
control_tier,
router_name=arm_router,
outcome="error",
error=verdict.error,
shadow=shadow,
@ -898,6 +967,7 @@ class ShadowEvalLogger(CustomLogger):
job,
request_id,
control_tier,
router_name=arm_router,
outcome=verdict.preference,
shadow=shadow,
real_model=real_model,
@ -916,6 +986,7 @@ class ShadowEvalLogger(CustomLogger):
job,
request_id,
control_tier,
router_name=arm_router,
outcome="error",
error=f"pipeline error: {e}",
shadow=shadow,
@ -933,6 +1004,7 @@ class ShadowEvalLogger(CustomLogger):
request_id: str,
control_tier: str | None,
*,
router_name: str,
outcome: str,
real_cost: float,
real_classifier_cost: float,
@ -955,6 +1027,7 @@ class ShadowEvalLogger(CustomLogger):
data={ # mutable-ok: Prisma payload
"job_id": job.id,
"request_id": request_id,
"router_name": router_name,
"outcome": outcome,
"tier": control_tier if job.direction == "reverse" else (shadow.tier if shadow else None),
"real_model": real_model or None,

View file

@ -416,15 +416,25 @@ class WebSearchInterceptionLogger(CustomLogger):
if not tools:
return None
if call_type in (CallTypes.responses, CallTypes.aresponses):
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
# Check if any tool is a web search tool (native or already LiteLLM standard)
has_websearch: Final = any(is_web_search_tool(t) for t in tools)
is_responses_call: Final = call_type in (CallTypes.responses, CallTypes.aresponses)
has_websearch: Final = (
any(is_web_search_tool_responses(tool) for tool in tools)
if is_responses_call
else any(is_web_search_tool(tool) for tool in tools)
)
if not has_websearch:
return None
if self.search_tool_name:
try:
from litellm.proxy.proxy_server import llm_router
except ImportError:
llm_router = None
self._select_search_tool_from_router(llm_router=llm_router)
if is_responses_call:
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
verbose_logger.debug("WebSearchInterception: Converting native web_search tools to LiteLLM standard")
# If the client sent an Anthropic-native web_search_* tool, mark the
@ -1631,9 +1641,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return None
def _select_search_tool_from_router(self, llm_router: object) -> "_SearchToolConfig | None":
if llm_router is None or not hasattr(llm_router, "search_tools"):
return None
search_tools: Final = list(getattr(llm_router, "search_tools") or [])
search_tools: Final = list(getattr(llm_router, "search_tools", []) or [])
return self._select_search_tool_from_list(search_tools=search_tools, source="router")
def _select_search_tool_from_list(
@ -1643,20 +1651,26 @@ class WebSearchInterceptionLogger(CustomLogger):
) -> "_SearchToolConfig | None":
if self.search_tool_name:
matching_tools = [tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name]
if matching_tools:
search_provider = (matching_tools[0].get("litellm_params", {}) or {}).get("search_provider")
verbose_logger.debug(
"WebSearchInterception: Found search tool '%s' from %s with provider '%s'",
self.search_tool_name,
source,
search_provider,
if not matching_tools:
raise ValueError(f"Configured search tool '{self.search_tool_name}' was not found")
selected_tool: Final = matching_tools[0]
litellm_params: Final = selected_tool.get("litellm_params")
selected_search_provider: Final = (
litellm_params.get("search_provider") if isinstance(litellm_params, Mapping) else None
)
if not isinstance(selected_search_provider, str) or not selected_search_provider.strip():
raise ValueError(
f"Configured search tool '{self.search_tool_name}' does not define a valid search provider"
)
return matching_tools[0]
verbose_logger.debug(
"WebSearchInterception: Search tool '%s' not found in %s, falling back to first available or perplexity",
"WebSearchInterception: Found search tool '%s' from %s with provider '%s'",
self.search_tool_name,
source,
selected_search_provider,
)
return selected_tool
if search_tools:
first_tool: Final = search_tools[0]

View file

@ -7,7 +7,13 @@ import os
from dataclasses import dataclass
from typing import Final
from litellm.types.files import get_file_mime_type_from_extension
from litellm.types.files import (
AUDIO_FILE_TYPES,
FILE_EXTENSIONS,
FILE_MIME_TYPES,
FileType,
get_file_mime_type_from_extension,
)
from litellm.types.utils import FileTypes
@ -323,3 +329,75 @@ def calculate_request_duration(file: FileTypes) -> float | None:
except Exception:
# Silently fail if duration extraction fails
return None
DEFAULT_SPEECH_MEDIA_TYPE: Final = "audio/mpeg"
def _speech_media_type_for_response_format(response_format: str) -> str | None:
file_type: Final = next(
(candidate for candidate, extensions in FILE_EXTENSIONS.items() if response_format.lower() in extensions),
None,
)
if file_type is None or file_type not in AUDIO_FILE_TYPES:
return None
return FILE_MIME_TYPES[file_type]
def resolve_speech_media_type(upstream_content_type: str | None, response_format: str | None) -> str:
upstream_media_type: Final = (upstream_content_type or "").split(";", 1)[0].strip().lower()
if upstream_media_type.startswith("audio/"):
return upstream_media_type
requested_media_type: Final = (
None if response_format is None else _speech_media_type_for_response_format(response_format)
)
return requested_media_type or DEFAULT_SPEECH_MEDIA_TYPE
_OGG_OPUS_HEAD_WINDOW: Final = 64
_ADTS_SYNC_AND_LAYER_MASK: Final = 0xF6
_ADTS_SYNC_AND_LAYER: Final = 0xF0
_ADTS_SAMPLE_RATE_INDEX_LIMIT: Final = 13
_MPEG_SYNC_MASK: Final = 0xE0
_MPEG_LAYER_MASK: Final = 0x06
_MPEG_RESERVED_VERSION: Final = 0x01
_MPEG_INVALID_BITRATE_INDEX: Final = 0x0F
_MPEG_RESERVED_SAMPLE_RATE_INDEX: Final = 0x03
def _adts_aac_frame_media_type(header: bytes) -> str | None:
sample_rate_index: Final = (header[2] >> 2) & 0x0F
return FILE_MIME_TYPES[FileType.AAC] if sample_rate_index < _ADTS_SAMPLE_RATE_INDEX_LIMIT else None
def _mpeg_audio_frame_media_type(header: bytes) -> str | None:
version: Final = (header[1] >> 3) & 0x03
layer: Final = header[1] & _MPEG_LAYER_MASK
bitrate_index: Final = header[2] >> 4
sample_rate_index: Final = (header[2] >> 2) & 0x03
if (
(header[1] & _MPEG_SYNC_MASK) != _MPEG_SYNC_MASK
or version == _MPEG_RESERVED_VERSION
or layer == 0
or bitrate_index == _MPEG_INVALID_BITRATE_INDEX
or sample_rate_index == _MPEG_RESERVED_SAMPLE_RATE_INDEX
):
return None
return FILE_MIME_TYPES[FileType.MP3]
def speech_media_type_from_audio_bytes(audio: bytes) -> str | None:
if audio[:4] == b"RIFF" and audio[8:12] == b"WAVE":
return FILE_MIME_TYPES[FileType.WAV]
if audio[:4] == b"fLaC":
return FILE_MIME_TYPES[FileType.FLAC]
if audio[:4] == b"OggS":
is_opus: Final = b"OpusHead" in audio[:_OGG_OPUS_HEAD_WINDOW]
return FILE_MIME_TYPES[FileType.OPUS if is_opus else FileType.OGG]
if audio[:3] == b"ID3":
return FILE_MIME_TYPES[FileType.MP3]
if len(audio) < 3 or audio[0] != 0xFF:
return None
if (audio[1] & _ADTS_SYNC_AND_LAYER_MASK) == _ADTS_SYNC_AND_LAYER:
return _adts_aac_frame_media_type(audio)
return _mpeg_audio_frame_media_type(audio)

View file

@ -3,7 +3,7 @@ Helper utilities for tracking the cost of built-in tools.
"""
from collections.abc import Mapping
from typing import Any, Final, Literal
from typing import Final, Literal
import litellm
from litellm.constants import OPENAI_FILE_SEARCH_COST_PER_1K_CALLS
@ -16,6 +16,7 @@ from litellm.types.llms.openai import (
WebSearchOptions,
)
from litellm.types.utils import (
ChatCompletionAnnotation,
Message,
ModelInfo,
ModelResponse,
@ -49,7 +50,7 @@ class StandardBuiltInToolCostTracking:
@staticmethod
def get_cost_for_built_in_tools(
model: str,
response_object: Any,
response_object: object,
usage: Usage | None = None,
custom_llm_provider: str | None = None,
standard_built_in_tools_params: StandardBuiltInToolsParams | None = None,
@ -201,8 +202,7 @@ class StandardBuiltInToolCostTracking:
model_info: Final = StandardBuiltInToolCostTracking._safe_get_model_info(
model=model, custom_llm_provider=custom_llm_provider
)
file_search_raw: Final[Any] = standard_built_in_tools_params.get("file_search", {})
file_search_usage: Final[FileSearchTool | None] = FileSearchTool(**file_search_raw) if file_search_raw else None
file_search_usage: Final[FileSearchTool | None] = standard_built_in_tools_params.get("file_search") or None
# Convert model_info to dict and extract usage parameters
model_info_dict: Final = dict(model_info) if model_info is not None else None
@ -245,7 +245,7 @@ class StandardBuiltInToolCostTracking:
@staticmethod
def _extract_file_search_params(
file_search_usage: Any,
file_search_usage: object,
) -> tuple[float | None, float | None]:
"""Extract and convert file search parameters safely."""
storage_gb = None
@ -335,7 +335,7 @@ class StandardBuiltInToolCostTracking:
@staticmethod
def _extract_token_counts(
computer_use_usage: Any,
computer_use_usage: object,
) -> tuple[int | None, int | None]:
"""Extract and convert token counts safely."""
input_tokens = None
@ -351,9 +351,9 @@ class StandardBuiltInToolCostTracking:
return input_tokens, output_tokens
@staticmethod
def _safe_convert_to_int(value: Any) -> int | None:
def _safe_convert_to_int(value: object) -> int | None:
"""Safely convert a value to int."""
if value is not None:
if isinstance(value, (int, float, str)):
try:
return int(value)
except (TypeError, ValueError):
@ -381,7 +381,7 @@ class StandardBuiltInToolCostTracking:
return usage.model_copy(update={"server_tool_use": server_tool_use})
@staticmethod
def response_object_includes_web_search_call(response_object: Any, usage: Usage | None = None) -> bool:
def response_object_includes_web_search_call(response_object: object, usage: Usage | None = None) -> bool:
"""
Check if the response object includes a web search call.
@ -446,7 +446,7 @@ class StandardBuiltInToolCostTracking:
@staticmethod
def response_object_includes_file_search_call(
response_object: Any,
response_object: object,
) -> bool:
"""
Check if the response object includes a file search call.
@ -477,11 +477,11 @@ class StandardBuiltInToolCostTracking:
message: Message | None = getattr(choice, "message", None)
if message is None:
continue
if annotations := getattr(message, "annotations", None):
if len(annotations) > 0:
for annotation in annotations:
if annotation.get("type", None) == annotation_type:
return True
annotations: list[ChatCompletionAnnotation] | None = getattr(message, "annotations", None)
if annotations:
for annotation in annotations:
if annotation.get("type", None) == annotation_type:
return True
return False
@staticmethod
@ -522,10 +522,8 @@ class StandardBuiltInToolCostTracking:
if model_info is None:
return 0.0
search_context_raw: Final[Any] = model_info.get("search_context_cost_per_query", {})
search_context_pricing: Final[SearchContextCostPerQuery] = (
SearchContextCostPerQuery(**search_context_raw) if search_context_raw else SearchContextCostPerQuery()
)
search_context_raw: Final = model_info.get("search_context_cost_per_query")
search_context_pricing: Final[SearchContextCostPerQuery] = search_context_raw or SearchContextCostPerQuery()
if web_search_options.get("search_context_size", None) == "low":
return search_context_pricing.get("search_context_size_low", 0.0)
elif web_search_options.get("search_context_size", None) == "medium":
@ -545,10 +543,8 @@ class StandardBuiltInToolCostTracking:
"""
if model_info is None:
return 0.0
search_context_raw: Final[Any] = model_info.get("search_context_cost_per_query", {}) or {}
search_context_pricing: Final[SearchContextCostPerQuery] = (
SearchContextCostPerQuery(**search_context_raw) if search_context_raw else SearchContextCostPerQuery()
)
search_context_raw: Final = model_info.get("search_context_cost_per_query")
search_context_pricing: Final[SearchContextCostPerQuery] = search_context_raw or SearchContextCostPerQuery()
return search_context_pricing.get("search_context_size_medium", 0.0)
@staticmethod
@ -714,7 +710,7 @@ class StandardBuiltInToolCostTracking:
response_object: ModelResponse,
) -> bool:
for _choice in response_object.choices:
message = getattr(_choice, "message", None)
message: Message | None = getattr(_choice, "message", None)
if (
message is not None
and hasattr(message, "annotations")

View file

@ -205,6 +205,41 @@ def is_non_content_values_set(message: AllMessageValues) -> bool:
return any(message.get(key, None) is not None for key in message if key not in ignore_keys)
_IMAGE_CONTENT_PART_TYPES: Final = frozenset({"image_url", "input_image", "image"})
_IMAGE_SCAN_MAX_DEPTH: Final = 4
def _content_parts_contain_image(parts: Sequence[object]) -> bool:
"""Depth-bounded frontier walk over nested content lists, iterative because the repo bans
recursion; an Anthropic tool_result nests its image parts exactly one level down."""
frontier = parts # rebind-ok: depth-bounded frontier walk
for _ in range(_IMAGE_SCAN_MAX_DEPTH):
if any(isinstance(part, Mapping) and part.get("type") in _IMAGE_CONTENT_PART_TYPES for part in frontier):
return True
frontier = tuple( # rebind-ok: depth-bounded frontier walk
nested
for part in frontier
if isinstance(part, Mapping)
for content in (part.get("content"),)
if isinstance(content, list)
for nested in content
)
if not frontier:
return False
return False
def request_contains_image_content(messages: Sequence[Mapping[str, object]]) -> bool:
"""Whether any message carries an image content part, across the dialects that reach
pre-routing hooks untranslated: chat-completions ``image_url``, Responses ``input_image``,
and Anthropic Messages ``image``, including images nested inside ``tool_result`` blocks."""
return any(
isinstance(content, list) and _content_parts_contain_image(content)
for message in messages
for content in (message.get("content"),)
)
def _audio_or_image_in_message_content(message: AllMessageValues) -> bool:
"""
Checks if message content contains an image or audio
@ -520,10 +555,10 @@ def update_messages_with_model_file_ids(
def update_responses_input_with_model_file_ids(
input: Any,
input: object,
model_id: str | None = None,
model_file_id_mapping: dict[str, dict[str, str]] | None = None,
) -> str | list[dict[str, Any]]:
) -> object:
"""
Updates responses API input with provider-specific file IDs.
File IDs are always inside the content array, not as direct input_file items.
@ -604,8 +639,8 @@ def update_responses_input_with_model_file_ids(
def _decode_vector_store_ids_in_tools(
tools: list[dict[str, Any]] | None,
) -> list[dict[str, Any]] | None:
tools: list[dict[str, object]] | None,
) -> list[dict[str, object]] | None:
"""
Decodes unified (LiteLLM-managed) vector_store_ids in file_search tools to
provider-native IDs. Non-unified IDs are passed through unchanged.
@ -657,10 +692,10 @@ def _decode_vector_store_ids_in_tools(
def update_responses_tools_with_model_file_ids(
tools: list[dict[str, Any]] | None,
tools: list[dict[str, object]] | None,
model_id: str | None = None,
model_file_id_mapping: dict[str, dict[str, str]] | None = None,
) -> list[dict[str, Any]] | None:
) -> list[dict[str, object]] | None:
"""
Updates responses API tools with provider-specific file IDs.
@ -853,7 +888,7 @@ def extract_file_data(file_data: FileTypes) -> ExtractedFileData:
# ---------------------------------------------------------------------------
def _estimate_json_bytes(obj: Any) -> int:
def _estimate_json_bytes(obj: object) -> int:
"""Estimate the JSON-serialised byte size of ``obj`` without materialising
JSON. Walks iteratively (no recursion stack risk).
@ -1944,7 +1979,7 @@ def drop_tool_reference_parts_from_tool_messages(
return [_drop_tool_reference_parts(message) for message in messages] # mutable-ok: pipelines mutate message lists
def _attempt_json_repair(s: str) -> Any | None:
def _attempt_json_repair(s: str) -> object | None:
"""
Attempt to repair truncated JSON produced by LLM tool calls.
@ -2060,7 +2095,7 @@ def parse_tool_call_arguments(
raise ValueError(error_message) from original_error
def split_concatenated_json_objects(raw: str) -> list[dict[str, Any]]:
def split_concatenated_json_objects(raw: str) -> list[dict[str, object]]:
"""
Split a string that contains one or more concatenated JSON objects into
a list of parsed dicts.
@ -2096,7 +2131,7 @@ def split_concatenated_json_objects(raw: str) -> list[dict[str, Any]]:
return []
decoder: Final = json.JSONDecoder()
results: Final[list[dict[str, Any]]] = []
results: Final[list[dict[str, object]]] = []
idx = 0
length: Final = len(raw)

View file

@ -4,8 +4,9 @@ import base64
import io
import struct
from collections.abc import Callable, Iterable, Mapping, Sequence
from typing import Any, Final, Literal, cast
from typing import Final, Literal, cast
import httpx
import tiktoken
import litellm
@ -171,6 +172,10 @@ def calculate_tiles_needed(
return total_tiles
def _unpack_ints(fmt: str, buffer: bytes) -> tuple[int, ...]:
return struct.unpack(fmt, buffer)
def get_image_type(image_data: bytes) -> str | None:
"""take an image (really only the first ~100 bytes max are needed)
and return 'png' 'gif' 'jpeg' 'webp' 'heic' or None. method added to
@ -210,9 +215,9 @@ def get_image_dimensions(
if data.startswith(("http://", "https://")):
try:
client: Final = _get_httpx_client()
response: Final = safe_get(client, data)
response: Final[httpx.Response] = safe_get(client, data)
max_bytes: Final = int(MAX_IMAGE_URL_DOWNLOAD_SIZE_MB * 1024 * 1024)
content_length: Final = response.headers.get("Content-Length")
content_length: Final[str | None] = response.headers.get("Content-Length")
if content_length is not None and int(content_length) > max_bytes:
pass # skip download; img_data stays None
else:
@ -229,10 +234,10 @@ def get_image_dimensions(
img_type: Final = get_image_type(img_data)
if img_type == "png":
w, h = struct.unpack(">LL", img_data[16:24])
w, h = _unpack_ints(">LL", img_data[16:24])
return w, h
elif img_type == "gif":
w, h = struct.unpack("<HH", img_data[6:10])
w, h = _unpack_ints("<HH", img_data[6:10])
return w, h
elif img_type == "jpeg":
with io.BytesIO(img_data) as fhandle:
@ -245,25 +250,25 @@ def get_image_dimensions(
while ord(byte) == 0xFF:
byte = fhandle.read(1)
ftype = ord(byte)
size = struct.unpack(">H", fhandle.read(2))[0] - 2
size = _unpack_ints(">H", fhandle.read(2))[0] - 2
fhandle.seek(1, 1)
h, w = struct.unpack(">HH", fhandle.read(4))
h, w = _unpack_ints(">HH", fhandle.read(4))
return w, h
elif img_type == "webp":
# For WebP, the dimensions are stored at different offsets depending on the format
# Check for VP8X (extended format)
if img_data[12:16] == b"VP8X":
w = struct.unpack("<I", img_data[24:27] + b"\x00")[0] + 1
h = struct.unpack("<I", img_data[27:30] + b"\x00")[0] + 1
w = _unpack_ints("<I", img_data[24:27] + b"\x00")[0] + 1
h = _unpack_ints("<I", img_data[27:30] + b"\x00")[0] + 1
return w, h
# Check for VP8 (lossy format)
elif img_data[12:16] == b"VP8 ":
w = struct.unpack("<H", img_data[26:28])[0] & 0x3FFF
h = struct.unpack("<H", img_data[28:30])[0] & 0x3FFF
w = _unpack_ints("<H", img_data[26:28])[0] & 0x3FFF
h = _unpack_ints("<H", img_data[28:30])[0] & 0x3FFF
return w, h
# Check for VP8L (lossless format)
elif img_data[12:16] == b"VP8L":
bits: Final = struct.unpack("<I", img_data[21:25])[0]
bits: Final = _unpack_ints("<I", img_data[21:25])[0]
w = (bits & 0x3FFF) + 1
h = ((bits >> 14) & 0x3FFF) + 1
return w, h
@ -420,8 +425,8 @@ def token_counter(
def _count_function_call_tokens(
key: str,
value: Any,
message: Mapping[str, Any],
value: object,
message: Mapping[str, object],
count_function: TokenCounterFunction,
) -> int:
"""
@ -587,7 +592,7 @@ def _fix_model_name(model: str) -> str:
def _count_image_tokens(
image_url: Any,
image_url: object,
use_default_image_token_count: bool,
) -> int:
"""
@ -627,7 +632,7 @@ def _count_image_tokens(
raise ValueError(f"Invalid image_url type: {type(image_url).__name__}. Expected str or dict with 'url' field.")
def _validate_anthropic_content(content: Mapping[str, Any]) -> type:
def _validate_anthropic_content(content: Mapping[str, object]) -> type:
"""
Validate and determine which Anthropic TypedDict applies.
@ -642,7 +647,7 @@ def _validate_anthropic_content(content: Mapping[str, Any]) -> type:
"tool_result": AnthropicMessagesToolResultParam,
}
expected_cls: Final = mapping.get(content_type)
expected_cls: Final = mapping.get(content_type) if isinstance(content_type, str) else None
if expected_cls is None:
raise ValueError(f"Unknown Anthropic content type: '{content_type}'")
@ -714,7 +719,7 @@ def _count_file_tokens(
def _count_anthropic_content(
content: Mapping[str, Any],
content: Mapping[str, object],
count_function: TokenCounterFunction,
use_default_image_token_count: bool,
default_token_count: int | None,
@ -729,7 +734,7 @@ def _count_anthropic_content(
avoiding hardcoded field names.
"""
typeddict_cls: Final = _validate_anthropic_content(content)
type_hints: Final = getattr(typeddict_cls, "__annotations__", {})
type_hints: Final[Mapping[str, object]] = getattr(typeddict_cls, "__annotations__", {})
tokens = 0
# Fields to skip (metadata/identifiers that don't contribute to prompt tokens)

View file

@ -100,16 +100,6 @@ InputWriteBackTarget = (
)
class _SSEDelta(TypedDict, total=False):
type: ReadOnly[str]
text: ReadOnly[str]
stop_reason: ReadOnly[str | None]
class _SSEEventData(TypedDict, total=False):
delta: ReadOnly[_SSEDelta]
def _as_str_mapping(value: Mapping[str, object]) -> Mapping[str, object]:
return value
@ -157,6 +147,16 @@ class ExtractedInput:
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
class _AnthropicSSEDelta(TypedDict, total=False):
type: ReadOnly[str]
text: ReadOnly[str]
stop_reason: ReadOnly[str | None]
class _AnthropicSSEEvent(TypedDict, total=False):
delta: ReadOnly[_AnthropicSSEDelta]
class AnthropicMessagesHandler(BaseTranslation):
"""Process Anthropic messages with guardrails.
@ -1247,8 +1247,8 @@ class AnthropicMessagesHandler(BaseTranslation):
# Only process content_block_delta events
if event_type == "content_block_delta" and data_line:
try:
data: _SSEEventData = json.loads(data_line)
delta = data.get("delta", {})
data: _AnthropicSSEEvent = json.loads(data_line)
delta: _AnthropicSSEDelta = data.get("delta", {})
if delta.get("type") == "text_delta":
text += delta.get("text", "")
except json.JSONDecodeError:
@ -1310,9 +1310,9 @@ class AnthropicMessagesHandler(BaseTranslation):
# Check for message_delta event with stop_reason
if event_type == "message_delta" and data_line:
try:
data: _SSEEventData = json.loads(data_line)
delta = data.get("delta", {})
stop_reason = delta.get("stop_reason")
data: _AnthropicSSEEvent = json.loads(data_line)
delta: _AnthropicSSEDelta = data.get("delta", {})
stop_reason: str | None = delta.get("stop_reason")
if stop_reason is not None:
return True
except json.JSONDecodeError:

View file

@ -66,6 +66,10 @@ if TYPE_CHECKING:
from litellm.llms.base_llm.chat.transformation import BaseConfig
def _loads_stream_chunk(payload: str) -> dict[str, object]:
return json.loads(payload)
async def make_call(
client: AsyncHTTPHandler | None,
api_base: str,
@ -78,7 +82,7 @@ async def make_call(
json_mode: bool,
speed: str | None = None,
tool_name_reverse_map: dict[str, str] | None = None,
) -> tuple[Any, httpx.Headers]:
) -> tuple["ModelResponseIterator", httpx.Headers]:
if client is None:
client = litellm.module_level_aclient
@ -93,7 +97,7 @@ async def make_call(
)
except httpx.HTTPStatusError as e:
error_headers = getattr(e, "headers", None)
error_response: Final = getattr(e, "response", None)
error_response: Final[object] = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
raise AnthropicError(
@ -138,7 +142,7 @@ def make_sync_call(
json_mode: bool,
speed: str | None = None,
tool_name_reverse_map: dict[str, str] | None = None,
) -> tuple[Any, httpx.Headers]:
) -> tuple["ModelResponseIterator", httpx.Headers]:
if client is None:
client = litellm.module_level_client # re-use a module level client
@ -153,7 +157,7 @@ def make_sync_call(
)
except httpx.HTTPStatusError as e:
error_headers = getattr(e, "headers", None)
error_response: Final = getattr(e, "response", None)
error_response: Final[object] = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
raise AnthropicError(
@ -292,7 +296,7 @@ class AnthropicChatCompletion(BaseLLM):
status_code: Final = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_text = getattr(e, "text", str(e))
error_response: Final = getattr(e, "response", None)
error_response: Final[object] = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
if error_response and hasattr(error_response, "text"):
@ -593,7 +597,7 @@ class AnthropicChatCompletion(BaseLLM):
status_code: Final = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_text = getattr(e, "text", str(e))
error_response: Final = getattr(e, "response", None)
error_response: Final[object] = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
if error_response and hasattr(error_response, "text"):
@ -664,10 +668,10 @@ class ModelResponseIterator:
# Accumulate web_search_tool_result blocks for multi-turn reconstruction
# See: https://github.com/BerriAI/litellm/issues/17737
self.web_search_results: list[dict[str, Any]] = []
self.web_search_results: list[dict[str, object]] = []
# Accumulate compaction blocks for multi-turn reconstruction
self.compaction_blocks: list[dict[str, Any]] = []
self.compaction_blocks: list[dict[str, object]] = []
# Accumulate streamed thinking text so final usage can split reasoning
# tokens from regular output tokens.
@ -727,7 +731,7 @@ class ModelResponseIterator:
str,
ChatCompletionToolCallChunk | None,
list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock],
dict[str, Any],
dict[str, object],
str | None,
]:
"""
@ -735,7 +739,7 @@ class ModelResponseIterator:
"""
text = ""
tool_use: ChatCompletionToolCallChunk | None = None
provider_specific_fields: Final = {}
provider_specific_fields: Final[dict[str, object]] = {}
reasoning_content: str | None = None
content_block: Final = ContentBlockDelta(**chunk)
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = []
@ -809,8 +813,8 @@ class ModelResponseIterator:
def _handle_redacted_thinking_content(
self,
content_block_start: ContentBlockStart,
provider_specific_fields: dict[str, Any],
) -> tuple[list[ChatCompletionRedactedThinkingBlock], dict[str, Any]]:
provider_specific_fields: dict[str, object],
) -> tuple[list[ChatCompletionRedactedThinkingBlock], dict[str, object]]:
"""
Handle the redacted thinking content
"""
@ -878,7 +882,7 @@ class ModelResponseIterator:
tool_use: ChatCompletionToolCallChunk | None = None
finish_reason = ""
usage: Usage | None = None
provider_specific_fields: dict[str, Any] = {}
provider_specific_fields: dict[str, object] = {}
reasoning_content: str | None = None
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
@ -1212,7 +1216,7 @@ class ModelResponseIterator:
# Try to parse as valid JSON first
try:
data_json: Final = json.loads(data_str)
data_json: Final = _loads_stream_chunk(data_str)
return self.chunk_parser(chunk=data_json)
except json.JSONDecodeError:
# Switch to accumulation mode and start accumulating
@ -1330,7 +1334,7 @@ class ModelResponseIterator:
str_line = str_line[index:]
if str_line.startswith("data:"):
data_json: Final = json.loads(str_line[5:])
data_json: Final = _loads_stream_chunk(str_line[5:])
return self.chunk_parser(chunk=data_json)
else:
return ModelResponseStream(id=self.response_id)

View file

@ -865,13 +865,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
f"Failed to fetch models from Anthropic. Status code: {response.status_code}, Response: {response.text}"
)
models: Final = response.json()["data"]
models: Final[Sequence[Mapping[str, str]]] = response.json()["data"]
litellm_model_names: Final = []
for model in models:
stripped_model_name = model["id"]
litellm_model_name = "anthropic/" + stripped_model_name
litellm_model_names.append(litellm_model_name)
litellm_model_names: Final = ["anthropic/" + model["id"] for model in models]
return litellm_model_names
def get_token_counter(self) -> BaseTokenCounter | None:
@ -1077,7 +1073,7 @@ def strip_empty_content_blocks_from_anthropic_messages(
return out
def _is_empty_text_block(block: Any) -> bool:
def _is_empty_text_block(block: object) -> bool:
if not isinstance(block, dict) or block.get("type") != "text":
return False
text: Final = block.get("text")
@ -1131,7 +1127,7 @@ def normalize_anthropic_tool_use_id(raw_id: str) -> str:
return sanitized or "tool_use_id"
def _sanitize_tool_use_id_content_block(block: Any) -> Any:
def _sanitize_tool_use_id_content_block(block: object) -> object:
if not isinstance(block, dict):
return block
block_type: Final = block.get("type")

View file

@ -18,6 +18,24 @@ TOOL_NAME_PREFIX_LENGTH: Final = OPENAI_MAX_TOOL_NAME_LENGTH - TOOL_NAME_HASH_LE
PROVIDERS_PROXYING_AN_UNKNOWN_BACKEND: Final = frozenset({"litellm_proxy"})
def _optional_attr(source: object, name: str) -> object:
return getattr(source, name, None)
def _as_string_mapping(value: object) -> Mapping[str, object] | None:
if isinstance(value, Mapping):
return value
return None
def _thought_signature(provider_specific_fields: object) -> str | None:
fields: Final = _as_string_mapping(provider_specific_fields)
if fields is None:
return None
signature: Final = fields.get("thought_signature")
return signature if isinstance(signature, str) else None
_ANTHROPIC_TOOL_SCHEMA_KEYS: Final = frozenset(
{"name", "type", "input_schema", "description", "cache_control", "strict"}
)
@ -56,7 +74,7 @@ def truncate_tool_name(name: str) -> str:
def create_tool_name_mapping(
tools: list[dict[str, Any]],
tools: Sequence[Mapping[str, object]],
) -> dict[str, str]:
"""
Create a mapping of truncated tool names to original names.
@ -70,6 +88,8 @@ def create_tool_name_mapping(
mapping: Final[dict[str, str]] = {}
for tool in tools:
original_name = tool.get("name", "")
if not isinstance(original_name, str):
continue
truncated_name = truncate_tool_name(original_name)
if truncated_name != original_name:
mapping[truncated_name] = original_name
@ -286,44 +306,44 @@ class LiteLLMAnthropicMessagesAdapter:
### FOR [BETA] `/v1/messages` endpoint support
def _extract_signature_from_tool_call(self, tool_call: Any) -> str | None:
def _extract_signature_from_tool_call(self, tool_call: object) -> str | None:
"""
Extract signature from a tool call's provider_specific_fields.
Only checks provider_specific_fields, not thinking blocks.
"""
signature = None
fields: Final = _optional_attr(tool_call, "provider_specific_fields")
if fields:
return _thought_signature(fields)
if hasattr(tool_call, "provider_specific_fields") and tool_call.provider_specific_fields:
if "thought_signature" in tool_call.provider_specific_fields:
signature = tool_call.provider_specific_fields["thought_signature"]
elif hasattr(tool_call.function, "provider_specific_fields") and tool_call.function.provider_specific_fields:
if "thought_signature" in tool_call.function.provider_specific_fields:
signature = tool_call.function.provider_specific_fields["thought_signature"]
function_fields: Final = _optional_attr(_optional_attr(tool_call, "function"), "provider_specific_fields")
if function_fields:
return _thought_signature(function_fields)
return signature
return None
def _extract_signature_from_tool_use_content(self, content: dict[str, Any]) -> str | None:
def _extract_signature_from_tool_use_content(self, content: Mapping[str, object]) -> str | None:
"""
Extract signature from a tool_use content block's provider_specific_fields.
"""
provider_specific_fields: Final = content.get("provider_specific_fields", {})
provider_specific_fields: Final = _as_string_mapping(content.get("provider_specific_fields", {}))
if provider_specific_fields:
return provider_specific_fields.get("signature")
signature: Final = provider_specific_fields.get("signature")
return signature if isinstance(signature, str) else None
return None
def _add_cache_control_if_applicable(
self,
source: Any,
target: Any,
source: object,
target: object,
model: str | None,
) -> None:
"""
Extract cache_control from source and add to target if it should be preserved.
This method accepts Any type to support both regular dicts and TypedDict objects.
TypedDict objects (like ChatCompletionTextObject, ChatCompletionImageObject, etc.)
are dicts at runtime but have specific types at type-check time. Using Any allows
this method to work with both while maintaining runtime correctness.
This method accepts an unconstrained type to support both regular dicts and
TypedDict objects. TypedDict objects (like ChatCompletionTextObject,
ChatCompletionImageObject, etc.) are dicts at runtime but have specific types at
type-check time, so the widest parameter type works with both.
Args:
source: Dict or TypedDict containing potential cache_control field
@ -751,7 +771,7 @@ class LiteLLMAnthropicMessagesAdapter:
return new_tools, tool_name_mapping
def translate_anthropic_output_format_to_openai(self, output_format: Any) -> dict[str, object] | None:
def translate_anthropic_output_format_to_openai(self, output_format: object) -> dict[str, object] | None:
"""
Translate Anthropic's output_format to OpenAI's response_format.
@ -1366,7 +1386,7 @@ class LiteLLMAnthropicMessagesAdapter:
@classmethod
def _first_positive_prompt_tokens_detail_value(cls, usage: Usage, field_names: tuple[str, ...]) -> int:
prompt_tokens_details: Final = getattr(usage, "prompt_tokens_details", None)
prompt_tokens_details: Final = _optional_attr(usage, "prompt_tokens_details")
if prompt_tokens_details is None:
return 0
@ -1374,7 +1394,7 @@ class LiteLLMAnthropicMessagesAdapter:
if isinstance(prompt_tokens_details, dict):
value = cls._positive_int(prompt_tokens_details.get(field_name))
else:
value = cls._positive_int(getattr(prompt_tokens_details, field_name, None))
value = cls._positive_int(_optional_attr(prompt_tokens_details, field_name))
if value > 0:
return value
return 0

View file

@ -14,7 +14,7 @@ Mirrors Anthropic's native ``compact_20260112`` for non-Anthropic providers:
import re
from collections.abc import Awaitable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeVar, Union, cast
from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeVar, Union, cast
from typing_extensions import NotRequired, ReadOnly, TypedDict, Unpack
@ -232,7 +232,7 @@ async def _check_summary_model_access(
key_models: Final = list(getattr(user_api_key_auth, "models", None) or [])
team_id: Final[str | None] = getattr(user_api_key_auth, "team_id", None)
team_model_aliases: Final = getattr(user_api_key_auth, "team_model_aliases", None)
team_model_aliases: Final[dict[str, str] | None] = getattr(user_api_key_auth, "team_model_aliases", None)
team_models: Final = list(getattr(user_api_key_auth, "team_models", None) or [])
user_id: Final[str | None] = getattr(user_api_key_auth, "user_id", None)
project_id: Final[str | None] = getattr(user_api_key_auth, "project_id", None)
@ -443,7 +443,9 @@ async def _check_summary_model_budget(
)
return False
end_user_model_max_budget: Final = getattr(user_api_key_auth, "end_user_model_max_budget", None)
end_user_model_max_budget: Final[dict[str, object] | None] = getattr(
user_api_key_auth, "end_user_model_max_budget", None
)
end_user_id: Final[str | None] = getattr(user_api_key_auth, "end_user_id", None)
if isinstance(end_user_model_max_budget, dict) and end_user_model_max_budget and end_user_id is not None:
try:
@ -854,8 +856,8 @@ def _extract_summary_text(raw: str | None) -> str | None:
def _system_to_openai_message(
system: str | list[dict[str, Any]] | None,
) -> Mapping[str, object] | None:
system: str | list[dict[str, object]] | None,
) -> dict[str, object] | None:
"""Translate Anthropic-shaped ``system`` to an OpenAI system message.
Accepts a bare string or a list of Anthropic content blocks; returns
@ -866,10 +868,10 @@ def _system_to_openai_message(
if isinstance(system, str):
return {"role": "system", "content": system} if system else None
if isinstance(system, list):
parts: Final[tuple[str, ...]] = tuple(
parts: Final[list[object]] = [
block.get("text", "") for block in system if isinstance(block, dict) and block.get("type") == "text"
)
joined: Final = "\n\n".join(part for part in parts if part)
]
joined: Final = "\n\n".join(part for part in parts if isinstance(part, str) and part)
return {"role": "system", "content": joined} if joined else None
return None
@ -951,7 +953,7 @@ async def _call_summary_model(
summary_model: str,
summary_messages: Sequence[Mapping[str, object]],
metadata: Mapping[str, object],
llm_router: object,
llm_router: Optional["Router"],
allowed_model_region: str | None = None,
max_tokens: int = COMPACT_SUMMARY_MAX_TOKENS,
) -> Union["ModelResponse", "CustomStreamWrapper"]:
@ -1036,10 +1038,9 @@ def _extract_usage(response: object) -> tuple[int, int]:
usage: Final[object] = getattr(response, "usage", None)
if usage is None:
return 0, 0
return (
int(getattr(usage, "prompt_tokens", 0) or 0),
int(getattr(usage, "completion_tokens", 0) or 0),
)
prompt_tokens: Final[int | None] = getattr(usage, "prompt_tokens", 0)
completion_tokens: Final[int | None] = getattr(usage, "completion_tokens", 0)
return int(prompt_tokens or 0), int(completion_tokens or 0)
def apply_client_compaction_block_history(

View file

@ -179,14 +179,14 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
)
@staticmethod
def _assistant_block_group_key(indexed_block: tuple[int, Mapping[str, Any]]) -> str:
def _assistant_block_group_key(indexed_block: tuple[int, Mapping[str, object]]) -> str:
"""Group a run of consecutive thinking blocks together; keep every other block alone."""
index, block = indexed_block
return "thinking" if block.get("type") == "thinking" else f"block:{index}"
@classmethod
def _assistant_group_to_input_item(
cls, group: tuple[Mapping[str, Any], ...]
cls, group: tuple[Mapping[str, object], ...]
) -> dict[str, Any] | None: # mutable-ok: API message payload
first: Final = group[0]
btype: Final = first.get("type")
@ -206,7 +206,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
def translate_messages_to_responses_input(
self,
messages: list[AllAnthropicPassThroughMessageValues],
) -> list[dict[str, Any]]:
) -> list[dict[str, object]]:
"""
Convert Anthropic messages list to Responses API `input` items.
@ -220,7 +220,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
assistant thinking -> reasoning
assistant tool_use -> function_call
"""
input_items: Final[list[dict[str, Any]]] = []
input_items: Final[list[dict[str, object]]] = []
for m in messages:
if m["role"] == "system":
@ -248,7 +248,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
}
)
elif isinstance(content, list):
user_parts: list[dict[str, Any]] = []
user_parts: list[Mapping[str, object]] = []
tool_image_parts: list[dict[str, Any]] = [] # mutable-ok: json content parts
for block in content:
if not isinstance(block, dict):
@ -379,9 +379,9 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
def translate_tools_to_responses_api(
self,
tools: list[AllAnthropicToolsValues],
) -> list[dict[str, Any]]:
) -> list[dict[str, object]]:
"""Convert Anthropic tool definitions to Responses API function tools."""
result: Final[list[dict[str, Any]]] = []
result: Final[list[dict[str, object]]] = []
for tool in tools:
tool_dict = cast(dict[str, Any], tool)
tool_type = tool_dict.get("type", "")
@ -392,7 +392,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
continue
# Responses turns strict mode on when `strict` is omitted, silently rewriting
# `required` to every property. Anthropic tools are non-strict unless asked.
func_tool: dict[str, Any] = {
func_tool: dict[str, object] = {
"type": "function",
"name": tool_name,
"strict": bool(tool_dict.get("strict")),
@ -407,7 +407,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
@staticmethod
def translate_tool_choice_to_responses_api(
tool_choice: AnthropicMessagesToolChoice,
) -> str | dict[str, Any]:
) -> str | dict[str, object]:
"""Convert Anthropic tool_choice to Responses API tool_choice."""
tc_type: Final = tool_choice.get("type")
if tc_type == "any":
@ -420,8 +420,8 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
@staticmethod
def translate_context_management_to_responses_api(
context_management: dict[str, Any],
) -> list[dict[str, Any]] | None:
context_management: dict[str, object],
) -> list[dict[str, object]] | None:
"""
Convert Anthropic context_management dict to OpenAI Responses API array format.
@ -435,13 +435,13 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
if not isinstance(edits, list):
return None
result: Final[list[dict[str, Any]]] = []
result: Final[list[dict[str, object]]] = []
for edit in edits:
if not isinstance(edit, dict):
continue
edit_type = edit.get("type", "")
if edit_type == "compact_20260112":
entry: dict[str, Any] = {"type": "compaction"}
entry: dict[str, object] = {"type": "compaction"}
trigger = edit.get("trigger")
if isinstance(trigger, dict) and trigger.get("value") is not None:
entry["compact_threshold"] = int(trigger["value"])
@ -451,9 +451,9 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
@staticmethod
def translate_thinking_to_reasoning(
thinking: dict[str, Any],
output_config: dict[str, Any] | None = None,
) -> dict[str, Any] | None:
thinking: dict[str, object],
output_config: dict[str, object] | None = None,
) -> dict[str, object] | None:
"""
Convert Anthropic thinking param to Responses API reasoning param.
@ -473,12 +473,14 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
if isinstance(output_config, dict) and output_config.get("effort"):
effort = output_config["effort"]
elif thinking_type == "enabled":
effort = reasoning_effort_from_thinking_budget(thinking.get("budget_tokens", 0))
raw_budget: Final = thinking.get("budget_tokens", 0)
budget_tokens: Final = int(raw_budget) if isinstance(raw_budget, (int, float)) else 0
effort = reasoning_effort_from_thinking_budget(budget_tokens)
else:
return None
auto_summary: Final = is_reasoning_auto_summary_enabled()
result: Final[dict[str, Any]] = {"effort": effort}
result: Final[dict[str, object]] = {"effort": effort}
summary: Final = thinking.get("summary")
if summary:
result["summary"] = summary
@ -570,7 +572,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
# output_format / output_config.format -> text format
# output_format: {"type": "json_schema", "schema": {...}}
# output_config: {"format": {"type": "json_schema", "schema": {...}}}
output_format: Any = anthropic_request.get("output_format")
output_format: object = anthropic_request.get("output_format")
output_config = anthropic_request.get("output_config")
if not isinstance(output_format, dict) and isinstance(output_config, dict):
output_format = output_config.get("format")
@ -620,7 +622,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
ResponseReasoningItem,
)
content: Final[list[dict[str, Any]]] = []
content: Final[list[dict[str, object]]] = []
stop_reason: AnthropicFinishReason = "end_turn"
for item in response.output:

View file

@ -5,7 +5,8 @@
import base64
import json
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Final, Generic, TypeVar, cast
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Generic, Protocol, TypeVar, cast, runtime_checkable
from litellm import verbose_logger
from litellm.llms.base_llm.managed_resources.isolation import (
@ -38,6 +39,30 @@ else:
ResourceObjectType = TypeVar("ResourceObjectType")
@runtime_checkable
class _HasIdentifier(Protocol):
id: str
class _ManagedResourceRecord(Protocol[ResourceObjectType]):
unified_resource_id: str
resource_object: ResourceObjectType
def model_dump(self) -> dict[str, object]: ...
class _ManagedResourceTable(Protocol[ResourceObjectType]):
async def create(self, *, data: Mapping[str, object]) -> object: ...
async def find_first(self, *, where: Mapping[str, object]) -> _ManagedResourceRecord[ResourceObjectType] | None: ...
async def find_many(
self, *, where: Mapping[str, object], take: int, order: Mapping[str, str]
) -> list[_ManagedResourceRecord[ResourceObjectType]]: ...
async def delete(self, *, where: Mapping[str, object]) -> object: ...
class BaseManagedResource(ABC, Generic[ResourceObjectType]):
"""
Base class for managing resources with target_model_names support.
@ -64,6 +89,9 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
self.internal_usage_cache = internal_usage_cache
self.prisma_client = prisma_client
def _resource_table(self) -> _ManagedResourceTable[ResourceObjectType]:
return getattr(self.prisma_client.db, self.table_name)
# ============================================================================
# ABSTRACT METHODS
# ============================================================================
@ -137,7 +165,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
litellm_parent_otel_span: Span | None,
model_mappings: dict[str, str],
user_api_key_dict: UserAPIKeyAuth,
additional_db_fields: dict[str, Any] | None = None,
additional_db_fields: Mapping[str, object] | None = None,
) -> None:
"""
Store unified resource ID with model mappings in cache and database.
@ -153,7 +181,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
verbose_logger.info("Storing LiteLLM Managed %s with id=%s in cache", self.resource_type, unified_resource_id)
# Prepare cache data
cache_data: Final = {
cache_data: Final[dict[str, object]] = {
"unified_resource_id": unified_resource_id,
"resource_object": resource_object,
"model_mappings": model_mappings,
@ -176,7 +204,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
)
# Prepare database data
db_data: Final = {
db_data: Final[dict[str, object]] = {
"unified_resource_id": unified_resource_id,
"model_mappings": json.dumps(model_mappings),
"flat_model_resource_ids": list(model_mappings.values()),
@ -205,7 +233,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
db_data.update(additional_db_fields)
# Store in database
table: Final = getattr(self.prisma_client.db, self.table_name)
table: Final = self._resource_table()
result: Final = await table.create(data=db_data)
verbose_logger.debug(
@ -240,7 +268,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
return result
# Check database
table: Final = getattr(self.prisma_client.db, self.table_name)
table: Final = self._resource_table()
db_object: Final = await table.find_first(where={"unified_resource_id": unified_resource_id})
if db_object:
@ -264,7 +292,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
The deleted resource object or None if not found
"""
# Get old value from database
table: Final = getattr(self.prisma_client.db, self.table_name)
table: Final = self._resource_table()
initial_value: Final = await table.find_first(where={"unified_resource_id": unified_resource_id})
if initial_value is None:
@ -515,7 +543,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
user_api_key_dict: UserAPIKeyAuth,
limit: int | None = None,
after: str | None = None,
additional_filters: dict[str, Any] | None = None,
additional_filters: Mapping[str, object] | None = None,
) -> dict[str, Any]:
"""
List resources created by a user.
@ -533,7 +561,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
if owner_filter is None:
return build_list_page([])
where_clause: Final[dict[str, Any]] = {**owner_filter}
where_clause: Final[dict[str, object]] = {**owner_filter}
if after:
where_clause["id"] = {"gt": after}
@ -544,14 +572,14 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
# Fetch resources
fetch_limit: Final = limit or 20
table: Final = getattr(self.prisma_client.db, self.table_name)
table: Final = self._resource_table()
resources: Final = await table.find_many(
where=where_clause,
take=fetch_limit,
order={"created_at": "desc"},
)
resource_objects: Final[list[Any]] = []
resource_objects: Final[list[object]] = []
for resource in resources:
try:
# Stop once we have enough
@ -559,12 +587,13 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
break
# Parse resource object
resource_data = resource.resource_object
if isinstance(resource_data, str):
resource_data = json.loads(resource_data)
stored_resource = resource.resource_object
resource_data: object = (
json.loads(stored_resource) if isinstance(stored_resource, str) else stored_resource
)
# Set unified ID
if hasattr(resource_data, "id"):
if isinstance(resource_data, _HasIdentifier):
resource_data.id = resource.unified_resource_id
elif isinstance(resource_data, dict):
resource_data["id"] = resource.unified_resource_id

View file

@ -1487,6 +1487,7 @@ class CommonBatchFilesUtils:
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
)
# Prepare the request data

View file

@ -113,6 +113,7 @@ class BedrockFilesHandler(BaseAWSLLM):
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
)
# Create S3 client

View file

@ -146,6 +146,7 @@ class _BedrockS3RequestParams(BaseModel):
aws_role_name: str | None = None
aws_web_identity_token: str | None = None
aws_sts_endpoint: str | None = None
aws_external_id: str | None = None
s3_region_name: str | None = None
s3_endpoint_url: str | None = None
@ -1029,6 +1030,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
)
# Calculate SHA256 hash of the content (REQUIRED for S3)
@ -1290,6 +1292,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
aws_role_name=request_params.aws_role_name,
aws_web_identity_token=request_params.aws_web_identity_token,
aws_sts_endpoint=request_params.aws_sts_endpoint,
aws_external_id=request_params.aws_external_id,
)
empty_body_hash: Final = hashlib.sha256(b"").hexdigest()

View file

@ -2,7 +2,7 @@ import base64
import datetime
import json
import math
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from typing import Any, Final
import httpx
@ -128,24 +128,35 @@ def is_gemini_image_model(model: str) -> bool:
return "gemini" in base_model
def _parse_image_config_string(raw_image_config: str, model: str) -> object:
try:
return json.loads(raw_image_config)
except json.JSONDecodeError as exc:
raise litellm.UnsupportedParamsError(
model=model,
message="`imageConfig` must be valid JSON when provided as a string.",
) from exc
def map_openai_image_params_to_gemini(
params: dict[str, Any],
params: Mapping[str, object],
model: str,
supported_params: Sequence[str],
optional_params: dict[str, Any] | None = None,
optional_params: Mapping[str, object] | None = None,
parse_image_config_string: bool = False,
) -> dict[str, Any]:
optional_params = optional_params or {}
) -> dict[str, object]:
already_mapped: Final[Mapping[str, object]] = optional_params or {}
filtered_params: Final = {key: value for key, value in params.items() if key in supported_params}
mapped_params: Final[dict[str, Any]] = {}
mapped_params: Final[dict[str, object]] = {}
if "n" in filtered_params and "n" not in optional_params:
if "n" in filtered_params and "n" not in already_mapped:
mapped_params["sampleCount"] = filtered_params["n"]
if "size" in filtered_params and "size" not in optional_params:
size_param: Final = filtered_params.get("size")
if isinstance(size_param, str) and "size" not in already_mapped:
image_config: Final = map_openai_size_to_gemini_image_config(
filtered_params["size"],
size_param,
model,
)
if image_config is not None:
@ -156,33 +167,30 @@ def map_openai_image_params_to_gemini(
if "imageSize" in image_config:
mapped_params["imageSize"] = image_config["imageSize"]
image_config_param = filtered_params.get("imageConfig")
if isinstance(image_config_param, str) and parse_image_config_string:
try:
image_config_param = json.loads(image_config_param)
except json.JSONDecodeError as exc:
raise litellm.UnsupportedParamsError(
model=model,
message="`imageConfig` must be valid JSON when provided as a string.",
) from exc
raw_image_config: Final = filtered_params.get("imageConfig")
image_config_param: Final[object] = (
_parse_image_config_string(raw_image_config, model)
if isinstance(raw_image_config, str) and parse_image_config_string
else raw_image_config
)
if isinstance(image_config_param, dict):
mapped_params["imageConfig"] = image_config_param
for key, value in filtered_params.items():
if key not in ("n", "size", "imageConfig", "tools", "web_search_options") and key not in optional_params:
if key not in ("n", "size", "imageConfig", "tools", "web_search_options") and key not in already_mapped:
mapped_params[key] = value
return mapped_params
def _dedupe_gemini_search_tools(tools: list[dict[str, Any]]) -> list[dict[str, Any]]:
def _dedupe_gemini_search_tools(tools: list[dict[str, object]]) -> list[dict[str, object]]:
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
search_tool_keys: Final = VertexGeminiConfig._search_tool_keys()
seen_search_keys: Final[set[str]] = set()
deduped_tools: Final[list[dict[str, Any]]] = []
deduped_tools: Final[list[dict[str, object]]] = []
for tool in tools:
if not isinstance(tool, dict):
@ -203,7 +211,7 @@ def _dedupe_gemini_search_tools(tools: list[dict[str, Any]]) -> list[dict[str, A
return deduped_tools
def _has_gemini_search_tool(tools: list[Any]) -> bool:
def _has_gemini_search_tool(tools: list[object]) -> bool:
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
@ -213,9 +221,9 @@ def _has_gemini_search_tool(tools: list[Any]) -> bool:
def map_gemini_image_tools_params(
non_default_params: dict[str, Any],
mapped_params: dict[str, Any],
) -> dict[str, Any]:
non_default_params: Mapping[str, object],
mapped_params: Mapping[str, object],
) -> dict[str, object]:
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
@ -239,21 +247,24 @@ def map_gemini_image_tools_params(
gemini_config._drop_search_tools_mixed_with_functions(result)
if isinstance(result.get("tools"), list):
result["tools"] = _dedupe_gemini_search_tools(result["tools"])
resolved_tools: Final = result.get("tools")
if isinstance(resolved_tools, list):
result["tools"] = _dedupe_gemini_search_tools(resolved_tools)
return result
def get_gemini_image_web_search_requests(
response_data: dict[str, Any],
response_data: Mapping[str, object],
) -> int | None:
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
grounding_metadata: Final[list[dict[str, Any]]] = []
for candidate in response_data.get("candidates", []):
raw_candidates: Final = response_data.get("candidates")
candidates: Final[list[object]] = raw_candidates if isinstance(raw_candidates, list) else []
grounding_metadata: Final[list[dict[str, object]]] = []
for candidate in candidates:
if not isinstance(candidate, dict):
continue
candidate_grounding = candidate.get("groundingMetadata")
@ -267,13 +278,14 @@ def get_gemini_image_web_search_requests(
def get_gemini_image_generation_config(
model: str,
optional_params: dict[str, Any],
) -> dict[str, Any]:
generation_config: Final[dict[str, Any]] = {"response_modalities": ["IMAGE", "TEXT"]}
optional_params: Mapping[str, object],
) -> dict[str, object]:
generation_config: Final[dict[str, object]] = {"response_modalities": ["IMAGE", "TEXT"]}
image_config: Final[dict[str, Any]] = {}
if isinstance(optional_params.get("imageConfig"), dict):
image_config.update(optional_params["imageConfig"])
raw_image_config: Final = optional_params.get("imageConfig")
image_config: Final[dict[str, object]] = {}
if isinstance(raw_image_config, dict):
image_config.update(raw_image_config)
if not supports_gemini_image_size(model):
image_config.pop("imageSize", None)
@ -398,7 +410,7 @@ class GeminiModelInfo(BaseLLMModelInfo):
f"Failed to fetch models from Gemini. Status code: {response.status_code}, Response: {response.json()}"
)
models: Final = response.json()["models"]
models: Final[list[dict[str, str]]] = response.json()["models"]
litellm_model_names: Final = self.process_model_name(models)
return litellm_model_names
@ -473,12 +485,12 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
async def count_tokens(
self,
model_to_use: str,
messages: list[dict[str, Any]] | None,
contents: list[dict[str, Any]] | None,
messages: list[dict[str, object]] | None,
contents: list[dict[str, object]] | None,
deployment: dict[str, Any] | None = None,
request_model: str = "",
tools: list[dict[str, Any]] | None = None,
system: Any | None = None,
tools: list[dict[str, object]] | None = None,
system: object | None = None,
) -> TokenCountResponse | None:
import copy

View file

@ -5,11 +5,13 @@ For vertex ai, check out the vertex_ai/files/handler.py file.
"""
import time
from typing import Any, Final, Literal
from collections.abc import Mapping
from typing import Final, Literal, TypedDict
from urllib.parse import urlparse
import httpx
from openai.types.file_deleted import FileDeleted
from typing_extensions import ReadOnly, Required
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
@ -18,7 +20,6 @@ from litellm.llms.base_llm.files.transformation import (
BaseFilesConfig,
LiteLLMLoggingObj,
)
from litellm.types.llms.gemini import GeminiCreateFilesResponseObject
from litellm.types.llms.openai import (
AllMessageValues,
CreateFileRequest,
@ -31,6 +32,25 @@ from litellm.types.utils import LlmProviders
from ..common_utils import GeminiModelInfo
class _GeminiFileMetadata(TypedDict, total=False):
name: ReadOnly[str]
uri: ReadOnly[Required[str]]
displayName: ReadOnly[Required[str]]
mimeType: ReadOnly[str]
sizeBytes: ReadOnly[Required[str]]
createTime: ReadOnly[Required[str]]
updateTime: ReadOnly[str]
expirationTime: ReadOnly[str]
sha256Hash: ReadOnly[str]
state: ReadOnly[str]
source: ReadOnly[str]
error: ReadOnly[Mapping[str, object]]
class _GeminiCreateFileResponse(TypedDict):
file: ReadOnly[_GeminiFileMetadata]
class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
def __init__(self):
pass
@ -41,14 +61,14 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
def validate_environment(
self,
headers: dict[Any, Any],
headers: dict[str, str],
model: str,
messages: list[AllMessageValues],
optional_params: dict[Any, Any],
litellm_params: dict[Any, Any],
optional_params: dict[str, object],
litellm_params: dict[str, object],
api_key: str | None = None,
api_base: str | None = None,
) -> dict[Any, Any]:
) -> dict[str, str]:
"""
Validate environment and add Gemini API key to headers.
Google AI Studio uses x-goog-api-key header for authentication.
@ -164,9 +184,9 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
Transform Gemini's file upload response into OpenAI-style FileObject
"""
try:
response_json: Final = raw_response.json()
response_json: Final[_GeminiCreateFileResponse] = raw_response.json()
response_object: Final = GeminiCreateFilesResponseObject(**response_json.get("file", {}))
response_object: Final = response_json["file"]
# Extract file information from Gemini response
@ -262,7 +282,7 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
"""
try:
verbose_logger.debug("Retrieve file response: %s", raw_response.text)
response_json: Final = raw_response.json()
response_json: Final[_GeminiFileMetadata] = raw_response.json()
verbose_logger.debug("Response JSON: %s", response_json)
# Map Gemini state to OpenAI status
gemini_state: Final = response_json.get("state", "STATE_UNSPECIFIED")

View file

@ -7,6 +7,8 @@ from collections import OrderedDict
from collections.abc import Mapping, Sequence
from typing import Any, Final, cast
from typing_extensions import ReadOnly, Required, TypedDict
import litellm
from litellm import verbose_logger
from litellm._uuid import uuid
@ -96,6 +98,23 @@ def _gemini_live_speech_config(voice: object) -> Mapping[str, object] | None:
return VertexGeminiConfig()._map_audio_params({"voice": voice})
class _GeminiLiveSetupEnvelope(TypedDict, total=False):
setup: ReadOnly[BidiGenerateContentSetup]
class _OpenAIRealtimeClientEvent(TypedDict, total=False):
type: ReadOnly[str]
audio: ReadOnly[Required[str]]
session: ReadOnly[dict[str, object]]
item: ReadOnly[dict[str, object]]
def _parse_setup(session_configuration_request: str) -> BidiGenerateContentSetup:
envelope: Final[_GeminiLiveSetupEnvelope] = json.loads(session_configuration_request)
empty_setup: Final[BidiGenerateContentSetup] = {}
return envelope.get("setup", empty_setup)
# Google bills Live transcription at an estimated 25 audio tokens/sec of input and
# 175 text tokens/min of output (ai.google.dev/gemini-api/docs/pricing).
GEMINI_LIVE_TRANSCRIBE_AUDIO_TOKENS_PER_SECOND: Final = 25
@ -130,7 +149,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
return True
@staticmethod
def _usage_detail_alias(details: Any, defaults: dict[str, int]) -> dict[str, Any]:
def _usage_detail_alias(details: Mapping[str, int | None] | None, defaults: dict[str, int]) -> dict[str, int]:
if not isinstance(details, dict):
return dict(defaults)
return {
@ -139,7 +158,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
}
@staticmethod
def _add_pipecat_usage_detail_aliases(usage_dict: dict[str, Any]) -> dict[str, Any]:
def _add_pipecat_usage_detail_aliases(usage_dict: dict[str, Any]) -> dict[str, object]:
usage_dict.setdefault(
"input_token_details",
GeminiRealtimeConfig._usage_detail_alias(
@ -222,8 +241,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
if not session_configuration_request:
return False
try:
setup: Final = json.loads(session_configuration_request).get("setup", {})
automatic_detection: Final = setup.get("realtimeInputConfig", {}).get("automaticActivityDetection", {})
setup: Final = _parse_setup(session_configuration_request)
automatic_detection: Final[object] = setup.get("realtimeInputConfig", {}).get(
"automaticActivityDetection", {}
)
return isinstance(automatic_detection, dict) and automatic_detection.get("disabled") is True
except (json.JSONDecodeError, TypeError, AttributeError):
return False
@ -406,7 +427,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
return "TEXT" if GeminiRealtimeConfig._is_text_only_live_model(model) else "AUDIO"
@staticmethod
def _coerce_response_modalities(model: str, modalities: Sequence[Any]) -> tuple[str, ...]:
def _coerce_response_modalities(model: str, modalities: Sequence[object]) -> tuple[str, ...]:
"""Swap responseModalities a Live model cannot produce: TEXT to AUDIO for
audio-only models, AUDIO to TEXT for text-only ones (e.g. transcribe-live)."""
normalized: Final = tuple(
@ -431,7 +452,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
def _handle_session_update(
self,
json_message: dict,
json_message: _OpenAIRealtimeClientEvent,
model: str,
session_configuration_request: str | None,
) -> list[str]:
@ -445,7 +466,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
with a 1007, tearing the session down). To carry tools/instructions, send
them on the first session.update before any conversation content.
"""
session_payload = json_message.get("session") or {}
empty_session: Final[dict[str, object]] = {}
session_payload = json_message.get("session") or empty_session
# Normalize GA-remapped fields (``output_modalities``,
# nested ``audio.input.transcription``,
# ``audio.input.turn_detection``) back to their flat beta keys so
@ -486,14 +508,15 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
verbose_logger.debug("Gemini Realtime: Ignoring session.update (setup already sent)")
return []
def _handle_conversation_item(self, json_message: dict) -> list[str]:
def _handle_conversation_item(self, json_message: _OpenAIRealtimeClientEvent) -> list[str]:
"""
Handle conversation.item.create for user text or function call output.
Converts OpenAI format to Gemini's clientContent (for user text) or
toolResponse (for function outputs).
"""
item: Final = json_message.get("item", {})
empty_item: Final[dict[str, object]] = {}
item: Final = json_message.get("item", empty_item)
item_type: Final = item.get("type")
if item_type == "function_call_output":
@ -524,7 +547,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
call_id,
)
function_response: Final[dict[str, Any]] = {"response": output_dict}
function_response: Final[dict[str, object]] = {"response": output_dict}
if self._include_function_response_id() and call_id:
function_response["id"] = call_id
if function_name:
@ -559,7 +582,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
) -> list[str]:
realtime_input_dict: BidiGenerateContentRealtimeInput = {}
try:
json_message: Final = json.loads(message)
json_message: Final[_OpenAIRealtimeClientEvent] = json.loads(message)
except json.JSONDecodeError:
if isinstance(message, bytes):
message_str = message.decode("utf-8", errors="replace")
@ -610,9 +633,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
session_configuration_request: str | None = None,
) -> OpenAIRealtimeStreamSessionEvents:
if session_configuration_request:
session_configuration_request_dict: BidiGenerateContentSetup = json.loads(
session_configuration_request
).get("setup", {})
session_configuration_request_dict: BidiGenerateContentSetup = _parse_setup(session_configuration_request)
else:
session_configuration_request_dict = {}
@ -663,7 +684,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
session_configuration_request_dict: BidiGenerateContentSetup = {}
if session_configuration_request is not None:
try:
session_configuration_request_dict = json.loads(session_configuration_request).get("setup", {})
session_configuration_request_dict = _parse_setup(session_configuration_request)
except json.JSONDecodeError:
session_configuration_request_dict = {}
generation_config: Final = session_configuration_request_dict.get("generationConfig", {})
@ -931,9 +952,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
return events
@staticmethod
def get_nested_value(obj: dict, path: str) -> Any:
def get_nested_value(obj: dict, path: str) -> object | None:
keys: Final = path.split(".")
current = obj
current: object = obj
for key in keys:
if isinstance(current, dict) and key in current:
current = current[key]
@ -1011,9 +1032,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
current_response_id = f"resp_{uuid.uuid4()}"
if session_configuration_request:
session_configuration_request_dict: BidiGenerateContentSetup = json.loads(
session_configuration_request
).get("setup", {})
session_configuration_request_dict: BidiGenerateContentSetup = _parse_setup(session_configuration_request)
else:
session_configuration_request_dict = {}
@ -1337,7 +1356,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
session_setup: BidiGenerateContentSetup = {}
if session_configuration_request is not None:
try:
session_setup = json.loads(session_configuration_request).get("setup", {})
session_setup = _parse_setup(session_configuration_request)
except (json.JSONDecodeError, TypeError):
session_setup = {}
tool_call_generation_config = session_setup.get("generationConfig", {}) or {}

View file

@ -13,55 +13,109 @@ Generated files are returned directly in the response - no separate storage need
import base64
import json
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from enum import Enum
from typing import Any, Final, Protocol
from typing import Any, Final, Protocol, TypedDict
from typing_extensions import NotRequired, ReadOnly, TypedDict
from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
class _ToolCallFunction(Protocol):
"""Function payload of an assistant tool call."""
name: str | None
arguments: str
class _ToolParameterSchema(TypedDict, total=False):
type: ReadOnly[str]
description: ReadOnly[str]
class _ToolCall(Protocol):
"""Tool call requested by the assistant on a chat completion choice."""
id: str
function: _ToolCallFunction
class _ToolArgumentSchema(TypedDict, total=False):
type: ReadOnly[str]
properties: ReadOnly[Mapping[str, _ToolParameterSchema]]
required: ReadOnly[Sequence[str]]
class _AssistantMessage(Protocol):
"""Assistant message carried by a chat completion choice."""
content: str | None
tool_calls: Sequence[_ToolCall] | None
class _OpenAIToolFunction(TypedDict, total=False):
name: ReadOnly[str]
description: ReadOnly[str]
parameters: ReadOnly[_ToolArgumentSchema]
class _CompletionChoice(Protocol):
"""Single choice of a chat completion response."""
finish_reason: str
message: _AssistantMessage
class _OpenAIToolSpec(TypedDict, total=False):
type: ReadOnly[str]
function: ReadOnly[_OpenAIToolFunction]
class _SandboxFile(TypedDict):
"""File generated inside the sandbox during a code execution run."""
class _AnthropicToolSpec(TypedDict, total=False):
name: ReadOnly[str]
description: ReadOnly[str]
input_schema: ReadOnly[_ToolArgumentSchema]
class _CodeExecutionArguments(TypedDict, total=False):
code: ReadOnly[str]
class _GeneratedFile(TypedDict, total=False):
name: ReadOnly[str]
mime_type: ReadOnly[str]
content_base64: ReadOnly[str]
size: ReadOnly[int]
class _SandboxGeneratedFile(TypedDict):
name: ReadOnly[str]
mime_type: ReadOnly[str]
content_base64: ReadOnly[str]
class _CodeExecutionArguments(TypedDict):
"""Arguments the model passes to the `litellm_code_execution` tool."""
class _SandboxExecutionResult(TypedDict):
success: ReadOnly[bool]
output: ReadOnly[str]
error: ReadOnly[str]
files: ReadOnly[Sequence[_SandboxGeneratedFile]]
code: NotRequired[ReadOnly[str]]
class _ExecutionResult(TypedDict, total=False):
iteration: ReadOnly[int]
success: ReadOnly[bool]
output: ReadOnly[str]
error: ReadOnly[str]
files: ReadOnly[Sequence[str]]
class _ToolCallFunction(Protocol):
name: str
arguments: str
class _ToolCall(Protocol):
id: str
function: _ToolCallFunction
class _AssistantMessage(Protocol):
content: str | None
tool_calls: Sequence[_ToolCall] | None
class _ResponseChoice(Protocol):
message: _AssistantMessage
finish_reason: str | None
class _CompletionResponse(Protocol):
choices: Sequence[_ResponseChoice]
class _CodeExecutionOutcome(TypedDict, total=False):
response: ReadOnly[_CompletionResponse | None]
files: ReadOnly[Sequence[_GeneratedFile]]
execution_results: ReadOnly[Sequence[_ExecutionResult]]
messages: ReadOnly[Sequence[dict[str, object]]]
max_iterations_reached: ReadOnly[bool]
def _parse_code_execution_arguments(serialized_arguments: str) -> _CodeExecutionArguments:
return json.loads(serialized_arguments)
class LiteLLMInternalTools(str, Enum):
@ -75,7 +129,7 @@ class LiteLLMInternalTools(str, Enum):
CODE_EXECUTION = "litellm_code_execution"
def get_litellm_code_execution_tool() -> dict[str, object]:
def get_litellm_code_execution_tool() -> _OpenAIToolSpec:
"""
Returns the litellm_code_execution tool definition in OpenAI format.
@ -96,7 +150,7 @@ def get_litellm_code_execution_tool() -> dict[str, object]:
}
def get_litellm_code_execution_tool_anthropic() -> dict[str, object]:
def get_litellm_code_execution_tool_anthropic() -> _AnthropicToolSpec:
"""
Returns the litellm_code_execution tool definition in Anthropic/messages API format.
@ -143,12 +197,12 @@ class CodeExecutionHandler:
async def execute_with_code_execution(
self,
model: str,
messages: list[dict],
tools: list[dict],
messages: list[dict[str, object]],
tools: list[_OpenAIToolSpec],
skill_files: dict[str, bytes],
skill_id: str | None = None,
**kwargs,
) -> dict[str, object]:
) -> _CodeExecutionOutcome:
"""
Execute an LLM call with automatic code execution handling.
@ -179,8 +233,8 @@ class CodeExecutionHandler:
)
current_messages: Final = list(messages)
generated_files: Final[list[dict[str, object]]] = [] # Files returned directly
execution_results: Final[list[dict[str, object]]] = []
generated_files: Final[list[_GeneratedFile]] = [] # Files returned directly
execution_results: Final[list[_ExecutionResult]] = []
executor: Final = SkillsSandboxExecutor(timeout=self.sandbox_timeout)
response: Any = None # Initialize to avoid possibly unbound error
@ -196,9 +250,9 @@ class CodeExecutionHandler:
**kwargs,
)
choice: _CompletionChoice = response.choices[0]
choice: _ResponseChoice = response.choices[0]
assistant_message = choice.message
stop_reason: str = choice.finish_reason
stop_reason = choice.finish_reason
# Build assistant message for conversation history
assistant_msg_dict: dict[str, object] = {
@ -236,19 +290,19 @@ class CodeExecutionHandler:
if tool_name == LiteLLMInternalTools.CODE_EXECUTION.value:
# Execute code in sandbox
try:
args: _CodeExecutionArguments = json.loads(tool_call.function.arguments)
code: str = args.get("code", "")
args = _parse_code_execution_arguments(tool_call.function.arguments)
code = args.get("code", "")
verbose_logger.debug("CodeExecutionHandler: Executing code (%s chars)", len(code))
exec_result = executor.execute(
exec_result: _SandboxExecutionResult = executor.execute(
code=code,
skill_files=skill_files,
)
verbose_logger.debug("CodeExecutionHandler: Execution result: %s", exec_result)
sandbox_files: Sequence[_SandboxFile] = exec_result["files"]
sandbox_files: Sequence[_SandboxGeneratedFile] = exec_result["files"]
execution_results.append(
{
@ -326,7 +380,7 @@ class CodeExecutionHandler:
}
def has_code_execution_tool(tools: list[dict] | None) -> bool:
def has_code_execution_tool(tools: list[_OpenAIToolSpec] | None) -> bool:
"""Check if litellm_code_execution tool is in the tools list."""
if not tools:
return False
@ -337,7 +391,7 @@ def has_code_execution_tool(tools: list[dict] | None) -> bool:
return False
def add_code_execution_tool(tools: list[dict] | None) -> list[dict]:
def add_code_execution_tool(tools: list[_OpenAIToolSpec] | None) -> list[_OpenAIToolSpec]:
"""Add litellm_code_execution tool if not already present."""
tools = tools or []
if not has_code_execution_tool(tools):

View file

@ -16,7 +16,7 @@ import io
import os
import tempfile
from dataclasses import dataclass
from typing import Any, Final, cast
from typing import Final, Protocol, cast
from litellm.llms.nvidia_riva.audio_transcription.transformation import (
RIVA_TARGET_NUM_CHANNELS,
@ -24,10 +24,30 @@ from litellm.llms.nvidia_riva.audio_transcription.transformation import (
)
from litellm.llms.nvidia_riva.common_utils import NvidiaRivaException
# Keep this as Any: the module intentionally avoids importing numpy at module
# import time (optional dependency), and project-wide mypy config evaluates this
# file in contexts where conditional type aliases can degrade to "FloatArray?".
FloatArray = Any
class FloatArray(Protocol):
"""Structural view of the ``numpy.ndarray`` surface this module relies on."""
@property
def ndim(self) -> int: ...
@property
def shape(self) -> tuple[int, ...]: ...
@property
def size(self) -> int: ...
def mean(self, axis: int) -> "FloatArray": ...
def ravel(self) -> "FloatArray": ...
def astype(self, dtype: object) -> "FloatArray": ...
def tobytes(self) -> bytes: ...
def __getitem__(self, key: object) -> "FloatArray": ...
def __mul__(self, other: float) -> "FloatArray": ...
_INSTALL_HINT = "Install Riva STT extras to enable automatic audio resampling: `pip install 'litellm[stt-nvidia-riva]'`"

View file

@ -5,10 +5,11 @@ import os
import re
from dataclasses import dataclass
from email.utils import formatdate
from typing import Any, Final, Protocol
from typing import Final, Protocol
from urllib.parse import urlparse
import httpx
from pydantic import JsonValue
from litellm.llms.base_llm.chat.transformation import BaseLLMException
@ -64,7 +65,7 @@ class OCISignerProtocol(Protocol):
See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html
"""
def do_request_sign(self, request: Any, *, enforce_content_headers: bool = False) -> None:
def do_request_sign(self, request: "OCIRequestWrapper", *, enforce_content_headers: bool = False) -> None:
pass
@ -113,7 +114,7 @@ def build_signature_string(method: str, path: str, headers: dict, signed_headers
return "\n".join(lines)
def load_private_key_from_str(key_str: str) -> Any:
def load_private_key_from_str(key_str: str) -> "rsa.RSAPrivateKey":
_require_cryptography()
key: Final = serialization.load_pem_private_key(
key_str.encode("utf-8"),
@ -124,7 +125,7 @@ def load_private_key_from_str(key_str: str) -> Any:
return key
def load_private_key_from_file(file_path: str) -> Any:
def load_private_key_from_file(file_path: str) -> "rsa.RSAPrivateKey":
"""Loads a private key from a file path."""
try:
with open(file_path, "r", encoding="utf-8") as f:
@ -421,16 +422,17 @@ OCI_JSON_TO_PYTHON_TYPES: Final[dict[str, str]] = {
}
def resolve_oci_schema_refs(schema: dict[str, Any]) -> dict[str, Any]:
def resolve_oci_schema_refs(schema: JsonValue) -> JsonValue:
"""Inline all ``$ref``/``$defs`` references — OCI does not support JSON Schema ``$ref``."""
defs: Final = schema.get("$defs", {})
resolving_stack: Final[set] = set()
raw_defs: Final = schema.get("$defs") if isinstance(schema, dict) else None
defs: Final[dict[str, JsonValue]] = raw_defs if isinstance(raw_defs, dict) else {}
resolving_stack: Final[set[str]] = set()
def _resolve(obj: Any) -> Any:
def _resolve(obj: JsonValue) -> JsonValue:
if isinstance(obj, dict):
if "$ref" in obj:
ref: Final = obj["$ref"]
if ref.startswith("#/$defs/"):
ref: Final = obj.get("$ref")
if ref is not None:
if isinstance(ref, str) and ref.startswith("#/$defs/"):
key: Final = ref.split("/")[-1]
if key in resolving_stack:
return {"type": "object"} # break cycles
@ -451,7 +453,7 @@ def resolve_oci_schema_refs(schema: dict[str, Any]) -> dict[str, Any]:
return resolved
def resolve_oci_schema_anyof(obj: Any) -> Any:
def resolve_oci_schema_anyof(obj: JsonValue) -> JsonValue:
"""Resolve Pydantic v2 ``Optional[T]`` → ``anyOf`` patterns.
Pydantic v2 emits ``{"anyOf": [{"type": "T"}, {"type": "null"}]}`` for
@ -459,10 +461,13 @@ def resolve_oci_schema_anyof(obj: Any) -> Any:
first non-null branch and merge top-level metadata into it.
"""
if isinstance(obj, dict):
if "anyOf" in obj and "type" not in obj:
non_null: Final = [t for t in obj["anyOf"] if not (isinstance(t, dict) and t.get("type") == "null")]
raw_any_of: Final = obj.get("anyOf")
if raw_any_of is not None and "type" not in obj:
branches: Final = raw_any_of if isinstance(raw_any_of, list) else []
non_null: Final = [t for t in branches if not (isinstance(t, dict) and t.get("type") == "null")]
if non_null:
resolved: Final = {**obj, **non_null[0]}
first: Final = non_null[0]
resolved: Final[dict[str, JsonValue]] = {**obj, **first} if isinstance(first, dict) else {**obj}
resolved.pop("anyOf", None)
return resolve_oci_schema_anyof(resolved)
return {k: resolve_oci_schema_anyof(v) for k, v in obj.items()}
@ -471,7 +476,7 @@ def resolve_oci_schema_anyof(obj: Any) -> Any:
return obj
def sanitize_oci_schema(schema: Any) -> Any:
def sanitize_oci_schema(schema: JsonValue) -> JsonValue:
"""Recursively remove OCI-incompatible fields from a JSON schema.
Strips ``title`` keys, removes ``None``-valued ``default`` entries,
@ -483,7 +488,7 @@ def sanitize_oci_schema(schema: Any) -> Any:
if not isinstance(schema, dict):
return schema
sanitized: Final[dict[str, Any]] = {}
sanitized: Final[dict[str, JsonValue]] = {}
for key, value in schema.items():
if key == "title":
continue
@ -513,7 +518,7 @@ def sanitize_oci_schema(schema: Any) -> Any:
return sanitized
def enrich_cohere_param_description(description: str, param_schema: dict[str, Any]) -> str:
def enrich_cohere_param_description(description: str, param_schema: dict[str, JsonValue]) -> str:
"""Embed schema constraints into a Cohere parameter description.
``CohereParameterDefinition`` only has ``type``, ``description``, and

View file

@ -170,16 +170,20 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
if model != "gpt-3.5-turbo-16k" and model != "gpt-4": # gpt-4 does not support 'response_format'
model_specific_params.append("response_format")
# Normalize model name for responses API (e.g., "responses/gpt-4.1" -> "gpt-4.1")
model_for_check: Final = model.split("responses/", 1)[1] if "responses/" in model else model
if (
model_for_check in litellm.open_ai_chat_completion_models
) or model_for_check in litellm.open_ai_text_completion_models:
if OpenAIGPTConfig.is_openai_catalog_model(model):
model_specific_params.append(
"user"
) # user is not a param supported by all openai-compatible endpoints - e.g. azure ai
return base_params + model_specific_params
@staticmethod
def is_openai_catalog_model(model: str) -> bool:
model_for_check: Final = model.split("responses/", 1)[1] if "responses/" in model else model
return (
model_for_check in litellm.open_ai_chat_completion_models
or model_for_check in litellm.open_ai_text_completion_models
)
def _map_openai_params(
self,
non_default_params: dict,
@ -755,6 +759,14 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
)
class OpenAIUnknownModelConfig(OpenAIGPTConfig):
"""A model the openai provider does not recognize is typically a LiteLLM proxy alias, so
forward reasoning_effort and let the server decide whether it is supported."""
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: inherited contract
return super().get_supported_openai_params(model) + ["reasoning_effort"] # mutable-ok: inherited contract
class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
def _map_reasoning_to_reasoning_content(self, choices: list) -> list:
"""

View file

@ -155,10 +155,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> ContainerObject:
"""Transform the OpenAI container creation response."""
response_data: Final[OpenAIContainerPayload] = raw_response.json()
# Transform the response data
container_obj: Final = ContainerObject(**response_data)
container_obj: Final = ContainerObject.model_validate(raw_response.json())
# Add cost for container creation (OpenAI containers are code interpreter sessions)
# https://platform.openai.com/docs/pricing
@ -215,10 +212,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> ContainerListResponse:
"""Transform the OpenAI container list response."""
response_data: Final[OpenAIContainerListPayload] = raw_response.json()
# Transform the response data
container_list: Final = ContainerListResponse(**response_data)
container_list: Final = ContainerListResponse.model_validate(raw_response.json())
return container_list
@ -235,7 +229,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}")
# No additional data needed for GET request
data: Final[dict[str, object]] = {}
data: Final[dict[str, str]] = {}
return url, data
@ -245,9 +239,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> ContainerObject:
"""Transform the OpenAI container retrieve response."""
response_data: Final[OpenAIContainerPayload] = raw_response.json()
# Transform the response data
container_obj: Final = ContainerObject(**response_data)
container_obj: Final = ContainerObject.model_validate(raw_response.json())
return container_obj
@ -268,7 +260,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}")
# No data needed for DELETE request
data: Final[dict[str, object]] = {}
data: Final[dict[str, str]] = {}
return url, data
@ -278,10 +270,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> DeleteContainerResult:
"""Transform the OpenAI container delete response."""
response_data: Final[OpenAIContainerDeletedPayload] = raw_response.json()
# Transform the response data
delete_result: Final = DeleteContainerResult(**response_data)
delete_result: Final = DeleteContainerResult.model_validate(raw_response.json())
return delete_result
@ -326,10 +315,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> ContainerFileListResponse:
"""Transform the OpenAI container file list response."""
response_data: Final[OpenAIContainerFileListPayload] = raw_response.json()
# Transform the response data
file_list: Final = ContainerFileListResponse(**response_data)
file_list: Final = ContainerFileListResponse.model_validate(raw_response.json())
return file_list
@ -352,7 +338,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}/files/{encoded_file_id}/content")
# No query parameters needed
params: Final[dict[str, object]] = {}
params: Final[dict[str, str]] = {}
return url, params

View file

@ -43,6 +43,7 @@ from litellm.utils import (
from ...types.llms.openai import *
from ..base import BaseLLM
from .chat.gpt_5_transformation import OpenAIGPT5Config
from .chat.gpt_transformation import OpenAIGPTConfig, OpenAIUnknownModelConfig
from .chat.o_series_transformation import OpenAIOSeriesConfig
from .common_utils import (
BaseOpenAILLM,
@ -189,7 +190,12 @@ class OpenAIConfig(BaseConfig):
elif litellm.openAIGPTAudioConfig.is_model_gpt_audio_model(model=model):
return litellm.openAIGPTAudioConfig.get_supported_openai_params(model=model)
else:
return litellm.openAIGPTConfig.get_supported_openai_params(model=model)
return self._gpt_config_for_model(model).get_supported_openai_params(model=model)
def _gpt_config_for_model(self, model: str) -> OpenAIGPTConfig:
if type(self) is OpenAIConfig and not OpenAIGPTConfig.is_openai_catalog_model(model):
return OpenAIUnknownModelConfig()
return litellm.openAIGPTConfig
def _map_openai_params(self, non_default_params: dict, optional_params: dict, model: str) -> dict:
supported_openai_params: Final = self.get_supported_openai_params(model)
@ -231,7 +237,7 @@ class OpenAIConfig(BaseConfig):
drop_params=drop_params,
)
return litellm.openAIGPTConfig.map_openai_params(
return self._gpt_config_for_model(model).map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,

View file

@ -6,10 +6,11 @@ Maps OpenAI TTS spec to RunwayML Text-to-Speech API
import asyncio
import time
from collections.abc import Coroutine
from typing import TYPE_CHECKING, Any, Final, Union
from collections.abc import Coroutine, Sequence
from typing import TYPE_CHECKING, Any, Final, TypedDict, Union
import httpx
from typing_extensions import ReadOnly
import litellm
from litellm._logging import verbose_logger
@ -31,6 +32,14 @@ else:
HttpxBinaryResponseContent = Any
class _RunwayTtsTaskResponse(TypedDict, total=False):
id: ReadOnly[str]
status: ReadOnly[str]
output: ReadOnly[Sequence[object]]
failure: ReadOnly[str]
failureCode: ReadOnly[str]
class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
"""
Configuration for RunwayML Text-to-Speech
@ -64,7 +73,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
litellm_params_dict: dict,
logging_obj: "LiteLLMLoggingObj",
timeout: float | httpx.Timeout,
extra_headers: dict[str, Any] | None,
extra_headers: dict[str, object] | None,
base_llm_http_handler: Any,
aspeech: bool,
api_base: str | None,
@ -72,7 +81,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
**kwargs: Any,
) -> Union[
"HttpxBinaryResponseContent",
Coroutine[Any, Any, "HttpxBinaryResponseContent"],
Coroutine[object, object, "HttpxBinaryResponseContent"],
]:
"""
Dispatch method to handle RunwayML TTS requests
@ -242,7 +251,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
raise TimeoutError(f"RunwayML TTS task polling timed out after {timeout_secs} seconds")
@staticmethod
def _check_task_status(response_data: dict[str, Any]) -> str:
def _check_task_status(response_data: _RunwayTtsTaskResponse) -> str:
"""
Check RunwayML task status from response.
@ -314,7 +323,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
response = client.get(url=task_url, headers=headers)
response.raise_for_status()
response_data = response.json()
response_data: _RunwayTtsTaskResponse = response.json()
# Check task status
status = self._check_task_status(response_data=response_data)
@ -362,7 +371,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
response = await client.get(url=task_url, headers=headers)
response.raise_for_status()
response_data = response.json()
response_data: _RunwayTtsTaskResponse = response.json()
# Check task status
status = self._check_task_status(response_data=response_data)
@ -453,7 +462,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
from litellm.types.llms.openai import HttpxBinaryResponseContent
try:
response_data: Final = raw_response.json()
response_data: Final[_RunwayTtsTaskResponse] = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error parsing RunwayML TTS response: {e}",
@ -483,7 +492,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
)
# Get the completed task data
task_data: Final = polled_response.json()
task_data: Final[_RunwayTtsTaskResponse] = polled_response.json()
verbose_logger.debug("RunwayML TTS polling complete, downloading audio")
@ -522,7 +531,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
from litellm.types.llms.openai import HttpxBinaryResponseContent
try:
response_data: Final = raw_response.json()
response_data: Final[_RunwayTtsTaskResponse] = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error parsing RunwayML TTS response: {e}",
@ -552,7 +561,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
)
# Get the completed task data
task_data: Final = polled_response.json()
task_data: Final[_RunwayTtsTaskResponse] = polled_response.json()
verbose_logger.debug("RunwayML TTS polling complete (async), downloading audio")

View file

@ -1,7 +1,7 @@
import re
from copy import deepcopy
from enum import Enum
from typing import Any, Final, Literal, get_type_hints
from typing import Any, Final, Literal, cast, get_type_hints
import httpx
@ -31,7 +31,7 @@ class VertexAIError(BaseLLMException):
super().__init__(message=message, status_code=status_code, headers=headers)
def redact_vertex_ai_metadata_from_logged_object(obj: Any) -> None:
def redact_vertex_ai_metadata_from_logged_object(obj: object) -> None:
if isinstance(obj, dict):
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
if field in obj:
@ -651,7 +651,7 @@ def _build_json_schema(parameters: dict) -> dict:
return parameters
def _filter_anyof_fields(schema_dict: dict[str, Any]) -> dict[str, Any]:
def _filter_anyof_fields(schema_dict: dict[str, object]) -> dict[str, object]:
"""
When anyof is present, only keep the anyof field and its contents - otherwise VertexAI will throw an error - https://github.com/BerriAI/litellm/issues/11164
Filter out other fields in the same dict.
@ -704,7 +704,7 @@ def process_items(schema, depth=0):
process_items(item, depth + 1)
def set_schema_property_ordering(schema: dict[str, Any], depth: int = 0) -> dict[str, Any]:
def set_schema_property_ordering(schema: dict[str, object], depth: int = 0) -> dict[str, object]:
"""
vertex ai and generativeai apis order output of fields alphabetically, unless you specify the order.
python dicts retain order, so we just use that. Note that this field only applies to structured outputs, and not tools.
@ -724,14 +724,16 @@ def set_schema_property_ordering(schema: dict[str, Any], depth: int = 0) -> dict
# retain propertyOrdering as an escape hatch if user already specifies it
if "propertyOrdering" not in schema:
schema["propertyOrdering"] = [k for k, v in schema["properties"].items()]
for k, v in schema["properties"].items():
set_schema_property_ordering(v, depth + 1)
if "items" in schema:
set_schema_property_ordering(schema["items"], depth + 1)
for v in schema["properties"].values():
if isinstance(v, dict):
set_schema_property_ordering(cast("dict[str, object]", v), depth + 1) # cast-ok: JSON Schema child
items: Final = schema.get("items")
if isinstance(items, dict):
set_schema_property_ordering(cast("dict[str, object]", items), depth + 1) # cast-ok: JSON Schema child
return schema
def filter_schema_fields(schema_dict: dict[str, Any], valid_fields: set[str], processed=None) -> dict[str, Any]:
def filter_schema_fields(schema_dict: dict[str, object], valid_fields: set[str], processed=None) -> dict[str, object]:
"""
Recursively filter a schema dictionary to keep only valid fields.
"""
@ -905,7 +907,7 @@ def _convert_schema_types(schema, depth=0):
"maxProperties",
}
any_of: Final[list[dict[str, Any]]] = []
any_of: Final[list[dict[str, object]]] = []
for t in type_val:
if not isinstance(t, str):
continue
@ -916,7 +918,7 @@ def _convert_schema_types(schema, depth=0):
# For object/array types, include type-specific fields
if t in ("object", "array"):
item_schema = {"type": t}
item_schema: dict[str, object] = {"type": t}
# Move type-specific fields into this anyOf item
for field in type_specific_fields:
if field in schema:
@ -1110,11 +1112,11 @@ class VertexAITokenCounter(BaseTokenCounter):
self,
model_to_use: str,
messages: list[dict[str, Any]] | None,
contents: list[dict[str, Any]] | None,
contents: list[dict[str, object]] | None,
deployment: dict[str, Any] | None = None,
request_model: str = "",
tools: list[dict[str, Any]] | None = None,
system: Any | None = None,
tools: list[dict[str, object]] | None = None,
system: object | None = None,
) -> TokenCountResponse | None:
import copy
@ -1131,25 +1133,26 @@ class VertexAITokenCounter(BaseTokenCounter):
partner_models_handler: Final = VertexAIPartnerModels()
# Extract vertex-specific params from litellm_params
vertex_project = count_tokens_params_request.get("vertex_project") or count_tokens_params_request.get(
partner_litellm_params: Final[dict[str, object]] = count_tokens_params_request
vertex_project = partner_litellm_params.get("vertex_project") or partner_litellm_params.get(
"vertex_ai_project"
)
vertex_location = count_tokens_params_request.get("vertex_location") or count_tokens_params_request.get(
vertex_location = partner_litellm_params.get("vertex_location") or partner_litellm_params.get(
"vertex_ai_location"
)
# Count tokens not available on global location: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/count-tokens
vertex_location = count_tokens_params_request.get("vertex_count_tokens_location") or vertex_location
vertex_location = partner_litellm_params.get("vertex_count_tokens_location") or vertex_location
vertex_credentials: Final = count_tokens_params_request.get(
"vertex_credentials"
) or count_tokens_params_request.get("vertex_ai_credentials")
vertex_credentials: Final = partner_litellm_params.get("vertex_credentials") or partner_litellm_params.get(
"vertex_ai_credentials"
)
result = await partner_models_handler.count_tokens(
model=model_to_use,
messages=messages or [],
litellm_params=count_tokens_params_request,
litellm_params=partner_litellm_params,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_credentials=vertex_credentials,

View file

@ -7,10 +7,14 @@ Reference: https://cloud.google.com/text-to-speech/docs/reference/rest/v1/text/s
import base64
from collections.abc import Coroutine
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Union
import httpx
from litellm.litellm_core_utils.audio_utils.utils import (
speech_media_type_from_audio_bytes,
)
from litellm.llms.base_llm.text_to_speech.transformation import (
BaseTextToSpeechConfig,
TextToSpeechRequestData,
@ -457,12 +461,11 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase):
if not response_content:
raise ValueError("No audioContent in Vertex AI TTS response")
# Decode base64 to get binary content
binary_data: Final = base64.b64decode(response_content)
# Create an httpx.Response object with the binary data
media_type: Final = speech_media_type_from_audio_bytes(binary_data)
response: Final = httpx.Response(
status_code=200,
headers=None if media_type is None else MappingProxyType({"content-type": media_type}),
content=binary_data,
)

View file

@ -13,7 +13,7 @@ import json
import os
import re
import time
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence
from contextlib import asynccontextmanager
from dataclasses import dataclass, replace
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast
@ -1206,7 +1206,7 @@ def _deserialize_json_dict(data: str | _StringMap | None) -> dict[str, str] | No
return data
def _deserialize_json_list(data: Any) -> list[dict[str, Any]] | None:
def _deserialize_json_list(data: object) -> list[dict[str, Any]] | None:
"""Deserialize a JSON array stored in the DB (``env_vars`` and friends).
Returns ``None`` for empty / null / unparseable input. Accepts strings
@ -1219,7 +1219,7 @@ def _deserialize_json_list(data: Any) -> list[dict[str, Any]] | None:
return None
if isinstance(data, str):
try:
parsed: Final = json.loads(data)
parsed: Final[object] = json.loads(data)
except (json.JSONDecodeError, TypeError):
return None
data = parsed
@ -1914,7 +1914,7 @@ class MCPServerManager:
async def load_servers_from_config(
self,
mcp_servers_config: dict[str, Any],
mcp_servers_config: dict[str, MCPServerConfig],
mcp_aliases: dict[str, str] | None = None,
):
"""
@ -3068,7 +3068,7 @@ class MCPServerManager:
return {}
cache_key: Final = "toolset_perms:" + ",".join(sorted(toolset_ids))
cached: Final = await user_api_key_cache.async_get_cache(key=cache_key)
cached: Final[dict[str, list[str]] | None] = await user_api_key_cache.async_get_cache(key=cache_key)
if cached is not None:
return cached
@ -5154,7 +5154,7 @@ class MCPServerManager:
# Wrapped so the bridge runs inside the task: the caller only holds the task and
# gathers it later, so there is no other point that still sees a block here.
async def _run_during_call_hook() -> Mapping[str, Any] | None:
async def _run_during_call_hook() -> Mapping[str, object] | None:
try:
return await proxy_logging_obj.during_call_hook(
user_api_key_dict=user_api_key_auth,
@ -5656,7 +5656,7 @@ class MCPServerManager:
async def _gather_openapi_tool_tasks(
self,
tasks: list[Any],
tasks: Sequence[Awaitable[object]],
proxy_logging_obj: ProxyLogging | None,
) -> CallToolResult:
"""Await OpenAPI tool tasks and return the tool call result."""

View file

@ -994,7 +994,7 @@ def get_key_model_rpm_limit(
# 2. Check model_max_budget
if user_api_key_dict.model_max_budget:
model_rpm_limit: Final[dict[str, Any]] = {}
model_rpm_limit: Final[dict[str, int]] = {}
for model, budget in user_api_key_dict.model_max_budget.items():
if isinstance(budget, dict) and budget.get("rpm_limit") is not None:
model_rpm_limit[model] = budget["rpm_limit"]
@ -1037,7 +1037,7 @@ def get_key_model_tpm_limit(
# 2. Check model_max_budget (iterate per-model like RPM does)
if user_api_key_dict.model_max_budget:
model_tpm_limit: Final[dict[str, Any]] = {}
model_tpm_limit: Final[dict[str, int]] = {}
for model, budget in user_api_key_dict.model_max_budget.items():
if isinstance(budget, dict) and budget.get("tpm_limit") is not None:
model_tpm_limit[model] = budget["tpm_limit"]
@ -1100,7 +1100,7 @@ def _validated_output_token_estimates_per_model(raw: object) -> Mapping[str, int
def _estimated_output_tokens_from_metadata(
metadata: Mapping[str, Any] | None,
metadata: Mapping[str, object] | None,
model_name: str | None,
) -> int | None:
"""Resolve the per-model, then global, estimate out of one metadata blob.
@ -1666,7 +1666,7 @@ def _dedupe_model_candidates(candidates: list[str]) -> list[str]:
return deduped
def _get_case_insensitive_mapping_value(mapping: Mapping[str, Any] | None, key: str) -> Any:
def _get_case_insensitive_mapping_value(mapping: Mapping[str, object] | None, key: str) -> object:
if not mapping:
return None
if key in mapping:
@ -1770,8 +1770,8 @@ def _resolve_model_id_with_router(model_id: str | None, llm_router: Router | Non
def _extract_model_candidates_from_request(
request_data: dict,
route: str,
request_headers: Mapping[str, Any] | None = None,
request_query_params: Mapping[str, Any] | None = None,
request_headers: Mapping[str, object] | None = None,
request_query_params: Mapping[str, object] | None = None,
llm_router: Router | None = None,
) -> list[str]:
candidates: Final[list[str]] = []
@ -1863,8 +1863,8 @@ def request_dispatched_to_pass_through_endpoint(request: Request | None) -> bool
def get_model_from_request(
request_data: dict,
route: str,
request_headers: Mapping[str, Any] | None = None,
request_query_params: Mapping[str, Any] | None = None,
request_headers: Mapping[str, object] | None = None,
request_query_params: Mapping[str, object] | None = None,
llm_router: Router | None = None,
request: Request | None = None,
) -> str | list[str] | None:

View file

@ -533,8 +533,8 @@ def sanitize_openai_provider_metadata(
Strips LiteLLM proxy-internal tracking fields that must not be forwarded to
OpenAI batch/file APIs.
"""
if not metadata:
return metadata
if metadata is None:
return None
sanitized: Final[dict[str, str]] = {}
for key, value in metadata.items():
if key in LITELLM_PROXY_INTERNAL_METADATA_KEYS:
@ -547,7 +547,7 @@ def sanitize_openai_provider_metadata(
key,
type(value).__name__,
)
return sanitized or None
return None if metadata and not sanitized else sanitized
def add_guardrail_to_applied_guardrails_header(request_data: dict, guardrail_name: str | None):
@ -650,7 +650,7 @@ def normalize_callback_names(callbacks: Iterable[object] | None) -> list[object]
return [c.lower() if isinstance(c, str) else c for c in callbacks]
def strip_callback_config(metadata: dict[str, Any] | None) -> dict[str, Any] | None:
def strip_callback_config(metadata: dict[str, object] | None) -> dict[str, object] | None:
"""Return key/team metadata without the slots that carry callback credentials."""
if not isinstance(metadata, dict):
return metadata

View file

@ -1,8 +1,12 @@
from collections.abc import Mapping, Sequence
from typing import Any, Final
from typing import Final, TypeAlias, Union
from litellm._logging import verbose_proxy_logger
JsonValue: TypeAlias = Union["JsonObject", "JsonArray", str, int, float, bool, None]
JsonObject: TypeAlias = dict[str, JsonValue]
JsonArray: TypeAlias = list[JsonValue]
class CustomOpenAPISpec:
"""
@ -27,7 +31,20 @@ class CustomOpenAPISpec:
RESPONSES_API_PATHS = ["/v1/responses", "/responses"]
@staticmethod
def get_pydantic_schema(model_class) -> Mapping[str, object] | None:
def _as_object(node: JsonValue) -> JsonObject:
return node if isinstance(node, dict) else {}
@staticmethod
def _as_array(node: JsonValue) -> JsonArray:
return node if isinstance(node, list) else []
@staticmethod
def _components_schemas(openapi_schema: JsonObject) -> JsonObject:
components: Final = CustomOpenAPISpec._as_object(openapi_schema.setdefault("components", {}))
return CustomOpenAPISpec._as_object(components.setdefault("schemas", {}))
@staticmethod
def get_pydantic_schema(model_class) -> JsonObject | None:
"""
Get JSON schema from a Pydantic model, handling both v1 and v2 APIs.
@ -54,9 +71,7 @@ class CustomOpenAPISpec:
return None
@staticmethod
def add_schema_to_components(
openapi_schema: dict[str, Any], schema_name: str, schema_def: Mapping[str, object]
) -> None:
def add_schema_to_components(openapi_schema: JsonObject, schema_name: str, schema_def: JsonObject) -> None:
"""
Add a schema definition to the OpenAPI components/schemas section.
@ -66,16 +81,25 @@ class CustomOpenAPISpec:
schema_def: The schema definition
"""
# Ensure components/schemas structure exists
if "components" not in openapi_schema:
openapi_schema["components"] = {}
if "schemas" not in openapi_schema["components"]:
openapi_schema["components"]["schemas"] = {}
_ = CustomOpenAPISpec._components_schemas(openapi_schema)
# Add the schema
CustomOpenAPISpec._move_defs_to_components(openapi_schema, {schema_name: schema_def})
@staticmethod
def add_request_body_to_paths(openapi_schema: dict[str, Any], paths: Sequence[str], schema_ref: str) -> None:
def _expanded_request_field(field_name: str, field_def: JsonValue) -> JsonValue:
expanded: Final = CustomOpenAPISpec._rewrite_defs_refs(
CustomOpenAPISpec._expand_field_definition(CustomOpenAPISpec._as_object(field_def))
)
if field_name != "messages":
return expanded
return {
**CustomOpenAPISpec._as_object(expanded),
"example": [{"role": "user", "content": "Hello, how are you?"}],
}
@staticmethod
def add_request_body_to_paths(openapi_schema: JsonObject, paths: Sequence[str], schema_ref: str) -> None:
"""
Add request body with expanded form fields for better Swagger UI display.
This keeps the request body but expands it to show individual fields in the UI.
@ -86,54 +110,58 @@ class CustomOpenAPISpec:
schema_ref: Reference to the schema component (e.g., "#/components/schemas/ModelName")
"""
for path in paths:
if path in openapi_schema.get("paths", {}) and "post" in openapi_schema["paths"][path]:
# Get the actual schema to extract ALL field definitions
schema_name = schema_ref.split("/")[-1] # Extract "ProxyChatCompletionRequest" from the ref
actual_schema = openapi_schema.get("components", {}).get("schemas", {}).get(schema_name, {})
schema_properties = actual_schema.get("properties", {})
required_fields = actual_schema.get("required", [])
path_item = CustomOpenAPISpec._as_object(
CustomOpenAPISpec._as_object(openapi_schema.get("paths")).get(path)
)
if "post" not in path_item:
continue
# Extract $defs and add them to components/schemas
# This fixes Pydantic v2 $defs not being resolvable in Swagger/OpenAPI
if "$defs" in actual_schema:
CustomOpenAPISpec._move_defs_to_components(openapi_schema, actual_schema["$defs"])
post_operation = CustomOpenAPISpec._as_object(path_item["post"])
# Create an expanded inline schema instead of just a $ref
# This makes Swagger UI show all individual fields in the request body editor
expanded_schema = {
"type": "object",
"required": required_fields,
"properties": {},
}
# Get the actual schema to extract ALL field definitions
schema_name = schema_ref.split("/")[-1] # Extract "ProxyChatCompletionRequest" from the ref
components = CustomOpenAPISpec._as_object(openapi_schema.get("components"))
actual_schema = CustomOpenAPISpec._as_object(
CustomOpenAPISpec._as_object(components.get("schemas")).get(schema_name)
)
schema_properties = CustomOpenAPISpec._as_object(actual_schema.get("properties"))
required_fields = actual_schema.get("required", [])
# Add all properties with their full definitions
for field_name, field_def in schema_properties.items():
expanded_field = CustomOpenAPISpec._expand_field_definition(field_def)
# Extract $defs and add them to components/schemas
# This fixes Pydantic v2 $defs not being resolvable in Swagger/OpenAPI
if "$defs" in actual_schema:
CustomOpenAPISpec._move_defs_to_components(
openapi_schema, CustomOpenAPISpec._as_object(actual_schema["$defs"])
)
# Rewrite $defs references to use components/schemas instead
expanded_field = CustomOpenAPISpec._rewrite_defs_refs(expanded_field)
# Create an expanded inline schema instead of just a $ref
# This makes Swagger UI show all individual fields in the request body editor
expanded_schema: JsonObject = {
"type": "object",
"required": required_fields,
"properties": {
field_name: CustomOpenAPISpec._expanded_request_field(field_name, field_def)
for field_name, field_def in schema_properties.items()
},
}
# Add a simple example for the messages field
if field_name == "messages":
expanded_field["example"] = [{"role": "user", "content": "Hello, how are you?"}]
# Set the request body with the expanded schema
post_operation["requestBody"] = {
"required": True,
"content": {"application/json": {"schema": expanded_schema}},
}
expanded_schema["properties"][field_name] = expanded_field
# Set the request body with the expanded schema
openapi_schema["paths"][path]["post"]["requestBody"] = {
"required": True,
"content": {"application/json": {"schema": expanded_schema}},
}
# Keep any existing parameters (like path parameters) but remove conflicting query params
if "parameters" in openapi_schema["paths"][path]["post"]:
existing_params = openapi_schema["paths"][path]["post"]["parameters"]
# Only keep path parameters, remove query params that conflict with request body
filtered_params = [param for param in existing_params if param.get("in") == "path"]
openapi_schema["paths"][path]["post"]["parameters"] = filtered_params
# Keep any existing parameters (like path parameters) but remove conflicting query params
if "parameters" in post_operation:
# Only keep path parameters, remove query params that conflict with request body
post_operation["parameters"] = [
param
for param in CustomOpenAPISpec._as_array(post_operation["parameters"])
if CustomOpenAPISpec._as_object(param).get("in") == "path"
]
@staticmethod
def _move_defs_to_components(openapi_schema: dict[str, Any], defs: Mapping[str, Mapping[str, Any]]) -> None:
def _move_defs_to_components(openapi_schema: JsonObject, defs: Mapping[str, JsonValue]) -> None:
"""
Move $defs from Pydantic v2 schema to OpenAPI components/schemas.
This makes the definitions resolvable in Swagger/OpenAPI viewers.
@ -146,23 +174,31 @@ class CustomOpenAPISpec:
return
# Ensure components/schemas exists
if "components" not in openapi_schema:
openapi_schema["components"] = {}
if "schemas" not in openapi_schema["components"]:
openapi_schema["components"]["schemas"] = {}
schemas: Final = CustomOpenAPISpec._components_schemas(openapi_schema)
# Add each definition to components/schemas
for def_name, def_schema in defs.items():
# Recursively rewrite any nested $defs references within this definition
rewritten_def = CustomOpenAPISpec._rewrite_defs_refs(def_schema)
openapi_schema["components"]["schemas"][def_name] = rewritten_def
schemas[def_name] = CustomOpenAPISpec._rewrite_defs_refs(def_schema)
# If this definition also has $defs, process them recursively
if "$defs" in def_schema:
CustomOpenAPISpec._move_defs_to_components(openapi_schema, def_schema["$defs"])
def_object = CustomOpenAPISpec._as_object(def_schema)
if "$defs" in def_object:
CustomOpenAPISpec._move_defs_to_components(
openapi_schema, CustomOpenAPISpec._as_object(def_object["$defs"])
)
@staticmethod
def _rewrite_defs_refs(schema: Any) -> Any:
def _rewritten_defs_entry(key: str, value: JsonValue) -> JsonValue:
if key == "$ref" and isinstance(value, str) and value.startswith("#/$defs/"):
# Rewrite the reference to use components/schemas
def_name: Final = value.replace("#/$defs/", "")
return f"#/components/schemas/{def_name}"
# Recursively process nested structures
return CustomOpenAPISpec._rewrite_defs_refs(value)
@staticmethod
def _rewrite_defs_refs(schema: JsonValue) -> JsonValue:
"""
Recursively rewrite $ref values from #/$defs/... to #/components/schemas/...
This converts Pydantic v2 references to OpenAPI-compatible references.
@ -174,26 +210,17 @@ class CustomOpenAPISpec:
Schema with rewritten references
"""
if isinstance(schema, dict):
result: Final = {}
for key, value in schema.items():
if key == "$ref" and isinstance(value, str) and value.startswith("#/$defs/"):
# Rewrite the reference to use components/schemas
def_name = value.replace("#/$defs/", "")
result[key] = f"#/components/schemas/{def_name}"
elif key == "$defs":
# Remove $defs from the schema since they're moved to components
continue
else:
# Recursively process nested structures
result[key] = CustomOpenAPISpec._rewrite_defs_refs(value)
return result
elif isinstance(schema, list):
return {
key: CustomOpenAPISpec._rewritten_defs_entry(key, value)
for key, value in schema.items()
if key != "$defs"
}
if isinstance(schema, list):
return [CustomOpenAPISpec._rewrite_defs_refs(item) for item in schema]
else:
return schema
return schema
@staticmethod
def _extract_field_schema(field_def: dict[str, Any]) -> dict[str, Any]:
def _extract_field_schema(field_def: JsonObject) -> JsonValue:
"""
Extract a simple schema from a Pydantic field definition for parameter display.
@ -209,10 +236,10 @@ class CustomOpenAPISpec:
# Handle anyOf (Optional fields in Pydantic v2)
if "anyOf" in field_def:
any_of: Final = field_def["anyOf"]
any_of: Final = CustomOpenAPISpec._as_array(field_def["anyOf"])
# Find the non-null type
for option in any_of:
if option.get("type") != "null":
if CustomOpenAPISpec._as_object(option).get("type") != "null":
return option
# Fallback to string if all else fails
return {"type": "string"}
@ -221,7 +248,7 @@ class CustomOpenAPISpec:
return {"type": "string"}
@staticmethod
def _expand_field_definition(field_def: dict[str, object]) -> dict[str, object]:
def _expand_field_definition(field_def: JsonObject) -> JsonObject:
"""
Expand a Pydantic field definition for inline use in OpenAPI schema.
This creates a full field definition that Swagger UI can render as individual form fields.
@ -237,12 +264,12 @@ class CustomOpenAPISpec:
@staticmethod
def add_request_schema(
openapi_schema: dict[str, object],
openapi_schema: JsonObject,
model_class: type,
schema_name: str,
paths: Sequence[str],
operation_name: str,
) -> dict[str, object]:
) -> JsonObject:
"""
Generic method to add a request schema to OpenAPI specification.
@ -282,8 +309,8 @@ class CustomOpenAPISpec:
@staticmethod
def add_chat_completion_request_schema(
openapi_schema: dict[str, object],
) -> dict[str, object]:
openapi_schema: JsonObject,
) -> JsonObject:
"""
Add ProxyChatCompletionRequest schema to chat completion endpoints for documentation.
This shows the request body in Swagger without runtime validation.
@ -309,7 +336,7 @@ class CustomOpenAPISpec:
return openapi_schema
@staticmethod
def add_embedding_request_schema(openapi_schema: dict[str, object]) -> dict[str, object]:
def add_embedding_request_schema(openapi_schema: JsonObject) -> JsonObject:
"""
Add EmbeddingRequest schema to embedding endpoints for documentation.
This shows the request body in Swagger without runtime validation.
@ -336,8 +363,8 @@ class CustomOpenAPISpec:
@staticmethod
def add_responses_api_request_schema(
openapi_schema: dict[str, object],
) -> dict[str, object]:
openapi_schema: JsonObject,
) -> JsonObject:
"""
Add ResponsesAPIRequestParams schema to responses API endpoints for documentation.
This shows the request body in Swagger without runtime validation.
@ -364,8 +391,8 @@ class CustomOpenAPISpec:
@staticmethod
def add_llm_api_request_schema_body(
openapi_schema: dict[str, object],
) -> dict[str, object]:
openapi_schema: JsonObject,
) -> JsonObject:
"""
Add LLM API request schema bodies to OpenAPI specification for documentation.
@ -376,12 +403,10 @@ class CustomOpenAPISpec:
OpenAPI schema with added request body schemas
"""
# Add chat completion request schema
openapi_schema = CustomOpenAPISpec.add_chat_completion_request_schema(openapi_schema)
with_chat_completions: Final = CustomOpenAPISpec.add_chat_completion_request_schema(openapi_schema)
# Add embedding request schema
openapi_schema = CustomOpenAPISpec.add_embedding_request_schema(openapi_schema)
with_embeddings: Final = CustomOpenAPISpec.add_embedding_request_schema(with_chat_completions)
# Add responses API request schema
openapi_schema = CustomOpenAPISpec.add_responses_api_request_schema(openapi_schema)
return openapi_schema
return CustomOpenAPISpec.add_responses_api_request_schema(with_embeddings)

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from typing import Any, Final, TypeVar, cast, overload
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast, overload
from pydantic import BaseModel
@ -9,6 +9,9 @@ from litellm.caching.dual_cache import DualCache
from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
if TYPE_CHECKING:
from opentelemetry.trace import Span
T = TypeVar("T", bound=BaseModel)
@ -40,8 +43,8 @@ class UserApiKeyCache(DualCache):
@overload
def get_cache(
self,
key: object,
parent_otel_span: object = None,
key: str,
parent_otel_span: Span | None = None,
local_only: bool = False,
*,
model_type: type[T],
@ -51,8 +54,8 @@ class UserApiKeyCache(DualCache):
@overload
def get_cache(
self,
key: object,
parent_otel_span: object = None,
key: str,
parent_otel_span: Span | None = None,
local_only: bool = False,
model_type: None = None,
**kwargs: object,
@ -60,12 +63,12 @@ class UserApiKeyCache(DualCache):
def get_cache(
self,
key: object,
parent_otel_span: object = None,
key: str,
parent_otel_span: Span | None = None,
local_only: bool = False,
model_type: type[BaseModel] | None = None,
**kwargs: object,
) -> Any | BaseModel | None:
) -> object:
if model_type is None and "model_type" in kwargs:
model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
cached: Final = super().get_cache(key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs)
@ -86,8 +89,8 @@ class UserApiKeyCache(DualCache):
@overload
async def async_get_cache(
self,
key: object,
parent_otel_span: object = None,
key: str,
parent_otel_span: Span | None = None,
local_only: bool = False,
*,
model_type: type[T],
@ -97,8 +100,8 @@ class UserApiKeyCache(DualCache):
@overload
async def async_get_cache(
self,
key: object,
parent_otel_span: object = None,
key: str,
parent_otel_span: Span | None = None,
local_only: bool = False,
model_type: None = None,
**kwargs: object,
@ -106,12 +109,12 @@ class UserApiKeyCache(DualCache):
async def async_get_cache(
self,
key: object,
parent_otel_span: object = None,
key: str,
parent_otel_span: Span | None = None,
local_only: bool = False,
model_type: type[BaseModel] | None = None,
**kwargs: object,
) -> Any | BaseModel | None:
) -> object:
if model_type is None and "model_type" in kwargs:
model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
cached: Final = await super().async_get_cache(
@ -131,12 +134,12 @@ class UserApiKeyCache(DualCache):
return None
return decoded
def set_cache(self, key: object, value: object, local_only: bool = False, **kwargs: object):
def set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object):
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs)
async def async_set_cache(self, key: object, value: object, local_only: bool = False, **kwargs: object):
async def async_set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object):
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)

View file

@ -33,6 +33,9 @@ from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS
from litellm.exceptions import ModifyResponseException
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
from litellm.litellm_core_utils.litellm_logging import (
_get_masked_values, # pyright: ignore[reportPrivateUsage] # the shared header-masking helper has no public name
)
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import bedrock_guardrail_cost
from litellm.litellm_core_utils.prompt_templates.factory import BedrockImageProcessor
from litellm.litellm_core_utils.url_utils import PayloadTooLargeError, SSRFError
@ -340,7 +343,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# Resource-less, detect-only InvokeGuardrailChecks mode. Present `checks`
# routes the guardrail to InvokeGuardrailChecks; absent => ApplyGuardrail.
self.checks: dict[str, Any] | None = self._normalize_checks(checks)
self.checks: dict[str, object] | None = self._normalize_checks(checks)
# Per-check block thresholds; a score >= threshold blocks. None => the
# check is detect-only (logged, never blocks).
self.content_filter_threshold = content_filter_threshold
@ -409,7 +412,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
]
@staticmethod
def _normalize_checks(checks: BedrockChecksConfigModel | Mapping[str, object] | None) -> dict[str, Any] | None:
def _normalize_checks(checks: BedrockChecksConfigModel | Mapping[str, object] | None) -> dict[str, object] | None:
"""Normalize the configured `checks` into a plain dict for the API body.
Accepts a pydantic ``BedrockChecksConfigModel`` or a raw dict; drops None /
@ -714,7 +717,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
def _create_bedrock_output_content_request(
self,
response: Any | ModelResponse,
response: object,
messages: list[AllMessageValues] | None = None,
) -> BedrockRequest:
"""
@ -738,9 +741,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
bedrock_request["content"] = bedrock_request_content
return bedrock_request
def _build_response_content_items(
self, response: Any | ModelResponse, has_grounding: bool
) -> list[BedrockContentItem]:
def _build_response_content_items(self, response: object, has_grounding: bool) -> list[BedrockContentItem]:
"""Build content item(s) from the model response. When the request supplied
grounding, the response is qualified ``guard_content`` so Bedrock can score it.
"""
@ -764,7 +765,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
self,
source: Literal["INPUT", "OUTPUT"],
messages: list[AllMessageValues] | None = None,
response: Any | ModelResponse | None = None,
response: object | None = None,
) -> BedrockRequest:
"""
Convert the litellm messages/response to the bedrock request format.
@ -1287,7 +1288,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
async def _apply_guardrail_content_with_chunking(
self,
content: Sequence[BedrockContentItem],
base_request_data: Mapping[str, Any],
base_request_data: Mapping[str, object],
credentials: "Credentials",
aws_region_name: str,
api_key: str | None,
@ -1466,7 +1467,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
async def _post_apply_guardrail_content_with_retry(
self,
content: Sequence[BedrockContentItem],
base_request_data: Mapping[str, Any],
base_request_data: Mapping[str, object],
credentials: "Credentials",
aws_region_name: str,
api_key: str | None,
@ -1516,7 +1517,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
async def _post_apply_guardrail_content(
self,
content: Sequence[BedrockContentItem],
base_request_data: Mapping[str, Any],
base_request_data: Mapping[str, object],
credentials: "Credentials",
aws_region_name: str,
api_key: str | None,
@ -1555,11 +1556,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
aws_region_name=aws_region_name,
api_key=api_key,
)
headers_dict: Final = dict(prepared_request.headers) # mutable-ok: the masking helper requires a dict
verbose_proxy_logger.debug(
"Bedrock AI request body: %s, url %s, headers: %s",
bedrock_request_data,
prepared_request.url,
prepared_request.headers,
_get_masked_values(headers_dict),
)
httpx_response: Final = await self._sign_and_post(
@ -2256,7 +2258,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return BedrockGuardrailResponse()
credentials, aws_region_name = self._load_credentials()
body: Final[dict[str, Any]] = {"messages": checks_messages, "checks": self.checks}
body: Final[dict[str, object]] = {"messages": checks_messages, "checks": self.checks}
api_key: Final[str | None] = request_data.get("api_key") if request_data else None
prepared_request: Final = self._prepare_request(
@ -2738,7 +2740,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
guardrail_name=self.guardrail_name,
)
detail: Final[dict[str, Any]] = {
detail: Final[dict[str, object]] = {
"error": "Violated guardrail policy",
"bedrock_guardrail_response": bedrock_guardrail_output_text,
}
@ -3297,7 +3299,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return updated_messages
def _mask_content_list(
self, content_list: list[Any], masked_texts: list[str], masking_index: int
self, content_list: Sequence[object], masked_texts: list[str], masking_index: int
) -> tuple[list[Any], int]:
"""
Apply masking to a list of content items.
@ -3310,7 +3312,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
Returns:
Updated content list with masked items
"""
new_content: Final[list[dict | str]] = []
new_content: Final[list[dict[str, object] | str]] = []
for item in content_list:
if isinstance(item, dict) and "text" in item:
new_item = item.copy()
@ -3331,7 +3333,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
def _apply_masking_to_response(
self,
response: ModelResponse | Any,
response: object,
bedrock_guardrail_response: BedrockGuardrailResponse,
) -> None:
"""

View file

@ -5,8 +5,9 @@ The public guardrail class imports this private mixin from
while preserving the existing public import path.
"""
from collections.abc import Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Optional
from typing import TYPE_CHECKING, Final, Optional
from fastapi import HTTPException
@ -23,7 +24,7 @@ if TYPE_CHECKING:
from .cisco_ai_defense import _ScanContext
def _serialize_mcp_content_item(item: object) -> dict[str, Any]:
def _serialize_mcp_content_item(item: object) -> dict[str, object]:
"""Serialize an MCP content item to a JSON-friendly dict.
Handles raw dicts, MCP SDK Pydantic models, and simple ``.text`` objects.
@ -57,7 +58,7 @@ class _CiscoAIDefenseMcpMixin:
def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: ...
async def _post_inspection(self, url: str, payload: dict[str, Any], surface: str) -> dict[str, Any]: ...
async def _post_inspection(self, url: str, payload: dict[str, object], surface: str) -> dict[str, object]: ...
def _handle_api_error(
self,
@ -67,16 +68,16 @@ class _CiscoAIDefenseMcpMixin:
start_time: datetime | None = ...,
surface: str = ...,
direction: str = ...,
) -> dict[str, Any]: ...
) -> dict[str, object]: ...
def _finalize_inspection(
self,
inspect_response: dict[str, Any],
inspect_response: dict[str, object],
request_data: dict,
context: "_ScanContext",
start_time: datetime,
response_obj: object = ...,
) -> dict[str, Any]: ...
) -> dict[str, object]: ...
# ------------------------------------------------------------------
# MCP post-tool hook (dispatcher contract)
@ -95,7 +96,7 @@ class _CiscoAIDefenseMcpMixin:
if self.inspection_type != "mcp":
return None
request_data: Final[dict[str, Any]] = {}
request_data: Final[dict[str, object]] = {}
for key in (
"name",
"litellm_call_id",
@ -188,9 +189,9 @@ class _CiscoAIDefenseMcpMixin:
original_hidden: Final = getattr(original_response_obj, "hidden_params", None)
if isinstance(original_hidden, HiddenParams):
hidden_params: Any = original_hidden
hidden_params: HiddenParams = original_hidden
else:
response_cost: Final = getattr(original_hidden, "response_cost", None)
response_cost: Final[float | None] = getattr(original_hidden, "response_cost", None)
hidden_params = HiddenParams(response_cost=response_cost) if response_cost is not None else HiddenParams()
return MCPPostCallResponseObject(
@ -200,11 +201,11 @@ class _CiscoAIDefenseMcpMixin:
@staticmethod
def _replace_mcp_tool_response(response_obj: object, replacement_obj: object) -> bool:
replacement: Final = getattr(replacement_obj, "mcp_tool_call_response", None)
replacement: Final[list[object] | None] = getattr(replacement_obj, "mcp_tool_call_response", None)
if replacement is None:
return False
inner: Final = getattr(response_obj, "mcp_tool_call_response", None)
inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
if inner is not None:
if _CiscoAIDefenseMcpMixin._replace_mcp_tool_response(inner, replacement_obj):
return True
@ -276,7 +277,7 @@ class _CiscoAIDefenseMcpMixin:
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
) -> dict[str, Any]:
) -> dict[str, object]:
del user_api_key_dict # carried via logging metadata, not the wire payload
url: Final = f"{self.api_base}{self.inspect_path}"
payload: Final = self._build_mcp_request_payload(data=data)
@ -312,7 +313,7 @@ class _CiscoAIDefenseMcpMixin:
response: object,
user_api_key_dict: UserAPIKeyAuth | None = None,
redact_response_obj: object = None,
) -> dict[str, Any]:
) -> dict[str, object]:
del user_api_key_dict # carried via logging metadata, not the wire payload
url: Final = f"{self.api_base}{self.inspect_path}"
payload: Final = self._build_mcp_response_payload(
@ -349,7 +350,7 @@ class _CiscoAIDefenseMcpMixin:
def _build_mcp_request_payload(
self,
data: dict,
) -> dict[str, Any] | None:
) -> dict[str, object] | None:
"""Build the JSON-RPC ``tools/call`` envelope sent to ``/inspect/mcp``.
The Cisco AI Defense MCP inspect endpoint expects the JSON-RPC
@ -390,7 +391,7 @@ class _CiscoAIDefenseMcpMixin:
self,
request_data: dict,
response: object,
) -> dict[str, Any] | None:
) -> dict[str, object] | None:
"""Build the MCP response-inspection body sent to ``/inspect/mcp``."""
request_payload: Final = self._build_mcp_request_payload(data=request_data)
if request_payload is None:
@ -415,7 +416,7 @@ class _CiscoAIDefenseMcpMixin:
return payload
@staticmethod
def _hydrate_mcp_tool_context(request_data: dict[str, Any]) -> None:
def _hydrate_mcp_tool_context(request_data: dict[str, object]) -> None:
metadata = request_data.get("mcp_tool_call_metadata")
if metadata is None:
nested: Final = request_data.get("metadata") or request_data.get("litellm_metadata")
@ -440,7 +441,7 @@ class _CiscoAIDefenseMcpMixin:
request_data.setdefault("server_name", server_name)
@staticmethod
def _normalize_mcp_response(response: object) -> dict[str, Any] | None:
def _normalize_mcp_response(response: object) -> dict[str, object] | None:
"""Normalize an MCP tool response into a JSON-RPC envelope.
Handles JSON-RPC dicts, raw content lists, MCP SDK models, and
@ -502,10 +503,10 @@ class _CiscoAIDefenseMcpMixin:
@staticmethod
def _build_mcp_result(
content: list[Any],
content: Sequence[object],
source: object = None,
) -> dict[str, Any]:
result: Final[dict[str, Any]] = {"content": [_serialize_mcp_content_item(item) for item in content]}
) -> dict[str, object]:
result: Final[dict[str, object]] = {"content": [_serialize_mcp_content_item(item) for item in content]}
for key in ("structuredContent", "isError"):
value = source.get(key) if isinstance(source, dict) else getattr(source, key, None)
if value is not None and (key != "isError" or isinstance(value, bool)):
@ -522,7 +523,7 @@ class _CiscoAIDefenseMcpMixin:
if response_obj is None:
return False
inner: Final = getattr(response_obj, "mcp_tool_call_response", None)
inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
if inner is not None:
return _CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text)
@ -559,7 +560,7 @@ class _CiscoAIDefenseMcpMixin:
pass
elif isinstance(response_obj, dict):
result: Final = response_obj.get("result")
target: Final[dict[Any, Any]] = result if isinstance(result, dict) else response_obj
target: Final[dict[object, object]] = result if isinstance(result, dict) else response_obj
if "structuredContent" in target:
target["structuredContent"] = replacement
replaced = True
@ -567,11 +568,11 @@ class _CiscoAIDefenseMcpMixin:
return replaced
@staticmethod
def _coerce_to_content_list(response_obj: object) -> list[Any] | None:
def _coerce_to_content_list(response_obj: object) -> list[object] | None:
"""Find the MCP content list inside supported response shapes."""
if response_obj is None:
return None
inner: Final = getattr(response_obj, "mcp_tool_call_response", None)
inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
if inner is not None:
return _CiscoAIDefenseMcpMixin._coerce_to_content_list(inner)
content: Final = getattr(response_obj, "content", None)
@ -594,8 +595,8 @@ class _CiscoAIDefenseMcpMixin:
@staticmethod
def _extract_sanitized_mcp_arguments(
inspect_response: dict[str, Any],
) -> dict[str, Any] | None:
inspect_response: dict[str, object],
) -> dict[str, object] | None:
"""Pull sanitized MCP tool-call arguments off the verdict.
Cisco can return them at the top level (``params.arguments``) or

View file

@ -80,7 +80,7 @@ import jwt
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey, RSAPublicKey
from typing_extensions import NotRequired, TypedDict
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
@ -89,6 +89,7 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrail_base_init import GuardrailBaseInitKwargs
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypesLiteral
@ -107,6 +108,19 @@ class _JWTDecodeKwargs(TypedDict):
issuer: NotRequired[str]
class _DebugHeaderClaims(TypedDict, total=False):
sub: ReadOnly[object]
iss: ReadOnly[object]
exp: ReadOnly[object]
scope: ReadOnly[str]
class _SignedClaimSummary(TypedDict):
sub: ReadOnly[object]
act: ReadOnly[Mapping[str, object]]
exp: ReadOnly[object]
# Module-level singleton for the JWKS discovery endpoint to access.
_mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None
@ -265,7 +279,8 @@ class MCPJWTSigner(CustomGuardrail):
**kwargs: Any,
) -> None:
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
super().__init__(**kwargs)
base_kwargs: Final[GuardrailBaseInitKwargs] = kwargs
super().__init__(**base_kwargs)
# --- Signing key setup ---
key_material: Final = os.environ.get(self.SIGNING_KEY_ENV)
@ -677,7 +692,7 @@ class MCPJWTSigner(CustomGuardrail):
data: dict,
jwt_claims: Mapping[str, object] | None = None,
call_type: CallTypesLiteral | None = None,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Build JWT claims for the outbound MCP access token.
@ -752,7 +767,7 @@ class MCPJWTSigner(CustomGuardrail):
# ------------------------------------------------------------------
@staticmethod
def _build_debug_header(claims: dict[str, Any], kid: str) -> str:
def _build_debug_header(claims: _DebugHeaderClaims, kid: str) -> str:
"""
Build the x-litellm-mcp-debug header value.
@ -873,16 +888,18 @@ class MCPJWTSigner(CustomGuardrail):
# FR-9: Debug header
# ------------------------------------------------------------------
if self.debug_headers:
new_headers["x-litellm-mcp-debug"] = self._build_debug_header(claims, self._kid)
debug_claims: Final[_DebugHeaderClaims] = claims
new_headers["x-litellm-mcp-debug"] = self._build_debug_header(debug_claims, self._kid)
hook_data["extra_headers"] = new_headers
logged_claims: Final[_SignedClaimSummary] = claims
verbose_proxy_logger.debug(
"MCPJWTSigner: signed JWT sub=%s act=%s tool=%s exp=%d verified=%s channel=%s call_type=%s",
claims.get("sub"),
claims.get("act", {}).get("sub"),
logged_claims.get("sub"),
logged_claims.get("act", {}).get("sub"),
hook_data.get("mcp_tool_name"),
claims["exp"],
logged_claims["exp"],
jwt_claims is not None,
bool(self.channel_token_audience),
call_type,

View file

@ -24,6 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
httpxSpecialProvider,
)
from litellm.proxy.guardrails.guardrail_hooks.noma.noma import NomaBlockedMessage
from litellm.types.guardrail_base_init import GuardrailBaseInitKwargs
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus
@ -83,7 +84,8 @@ class NomaV2Guardrail(CustomGuardrail):
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
super().__init__(**kwargs)
base_kwargs: Final[GuardrailBaseInitKwargs] = kwargs
super().__init__(**base_kwargs)
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:
@ -114,7 +116,7 @@ class NomaV2Guardrail(CustomGuardrail):
return parsed.hostname == _DEFAULT_API_BASE_HOSTNAME
@staticmethod
def _get_non_empty_str(value: Any) -> str | None:
def _get_non_empty_str(value: object) -> str | None:
if not isinstance(value, str):
return None
stripped: Final = value.strip()
@ -156,7 +158,7 @@ class NomaV2Guardrail(CustomGuardrail):
else model_call_details
)
payload: Final[dict[str, Any]] = {
payload: Final[dict[str, object]] = {
"inputs": inputs,
"request_data": payload_request_data,
"input_type": input_type,
@ -324,8 +326,9 @@ class NomaV2Guardrail(CustomGuardrail):
except NomaBlockedMessage as e:
guardrail_status = "guardrail_intervened"
blocked_detail: Final[dict[str, object]] = {"error": "blocked"}
guardrail_json_response = (
response_json if isinstance(response_json, dict) else getattr(e, "detail", {"error": "blocked"})
response_json if isinstance(response_json, dict) else getattr(e, "detail", blocked_detail)
)
raise
except Exception as e:

View file

@ -11,10 +11,10 @@
import asyncio
import json
import threading
from collections.abc import AsyncGenerator, Sequence
from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Sequence
from contextlib import asynccontextmanager
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast
import aiohttp
from typing_extensions import NotRequired, ReadOnly
@ -68,6 +68,14 @@ class _PresidioAnonymizeResponse(TypedDict):
items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]]
class _JsonResponse(Protocol):
def json(self) -> Awaitable[object]: ...
async def _json_body(response: _JsonResponse) -> object:
return await response.json()
_LoopSemaphores = dict[asyncio.AbstractEventLoop, asyncio.Semaphore]
@ -389,7 +397,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
f"expected application/json Content-Type but received '{content_type}'; body: '{error_body[:200]}'"
)
analyze_results: Final = await response.json()
analyze_results: Final = await _json_body(response)
verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
# Handle error responses from Presidio (e.g., {'error': 'No text provided'})
@ -997,7 +1005,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
except Exception as e:
raise e
def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]:
def logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
from concurrent.futures import ThreadPoolExecutor
def run_in_new_loop():
@ -1025,7 +1033,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
# No running event loop, we can safely run in this thread
return run_in_new_loop()
async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]:
async def async_logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
"""
Masks the input and output before logging to langfuse, datadog, etc.
"""
@ -1092,9 +1100,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
and not isinstance(result.choices[0], StreamingChoices)
):
await self._process_response_for_pii(response=result, request_data=kwargs, mode="mask")
elif self._is_anthropic_message_response(result):
elif isinstance(result, dict) and self._is_anthropic_message_response(result):
await self._process_anthropic_response_for_pii(
response=cast(dict, result), # cast-ok: _is_anthropic_message_response narrows via isinstance
response=result,
request_data=kwargs,
mode="mask",
)
@ -1321,7 +1329,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
async def _stream_apply_output_masking(
self,
response: Any,
response: AsyncIterable[object],
request_data: dict,
) -> AsyncGenerator[ModelResponseStream | bytes, None]:
"""Apply Presidio masking to streaming output (apply_to_output=True path)."""
@ -1425,7 +1433,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
return "\n".join(result_lines).encode("utf-8")
def _unmask_responses_api_completed_chunk(self, chunk: Any, pii_tokens: dict[str, str]) -> None:
def _unmask_responses_api_completed_chunk(self, chunk: object, pii_tokens: dict[str, str]) -> None:
"""
Unmask PII tokens in-place for a ``response.completed`` Responses API event.
@ -1434,7 +1442,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
blocks; text blocks expose a ``.text`` string attribute. We walk the tree
and replace every PII token with its original value.
"""
response_obj: Final = getattr(chunk, "response", None)
response_obj: Final[object] = getattr(chunk, "response", None)
if response_obj is None:
return
@ -1450,7 +1458,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
async def _stream_pii_unmasking(
self,
response: Any,
response: AsyncIterable[object],
request_data: dict,
) -> AsyncGenerator[ModelResponseStream | bytes, None]:
"""Apply PII unmasking to streaming output (output_parse_pii=True path)."""
@ -1526,7 +1534,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response: Any,
response: AsyncIterable[object],
request_data: dict,
) -> AsyncGenerator[ModelResponseStream | bytes, None]:
"""

View file

@ -310,7 +310,7 @@ class XecGuardGuardrail(CustomGuardrail):
scan_type: str,
suppress_errors: bool = False,
) -> dict | None:
payload: Final[dict[str, Any]] = {
payload: Final[dict[str, object]] = {
"model": self.xecguard_model,
"scan_type": scan_type,
"messages": messages,
@ -385,7 +385,7 @@ class XecGuardGuardrail(CustomGuardrail):
def _build_full_history(
self,
request_data: dict,
inputs: Any,
inputs: GenericGuardrailAPIInputs,
input_type: str,
) -> list[dict]:
"""Assemble the full message list that will be sent to XecGuard.

View file

@ -5,10 +5,11 @@ Pre-call hook that filters MCP tools semantically before LLM inference.
Reduces context window size and improves tool selection accuracy.
"""
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Optional
from collections.abc import Iterable, Mapping, Sequence
from typing import TYPE_CHECKING, Final, Optional
from fastapi import HTTPException
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
@ -30,6 +31,13 @@ if TYPE_CHECKING:
from litellm.router import Router
class SemanticToolFilterConfig(TypedDict, total=False):
enabled: ReadOnly[bool]
embedding_model: ReadOnly[str]
top_k: ReadOnly[int]
similarity_threshold: ReadOnly[float]
def _truncate_csv_at_tool_name_boundary(tool_names_csv: str, max_length: int) -> str:
"""Cap a CSV of tool names to max_length, dropping any name that does not fit whole."""
if len(tool_names_csv) <= max_length:
@ -68,7 +76,7 @@ class SemanticToolFilterHook(CustomLogger):
semantic_filter.top_k,
)
def _should_expand_mcp_tools(self, tools: list[Any]) -> bool:
def _should_expand_mcp_tools(self, tools: Iterable[Mapping[str, object]]) -> bool:
"""
Check if tools contain MCP references with server_url="litellm_proxy".
@ -82,9 +90,9 @@ class SemanticToolFilterHook(CustomLogger):
async def _expand_mcp_tools(
self,
tools: list[Any],
tools: Iterable[Mapping[str, object]],
user_api_key_dict: "UserAPIKeyAuth",
) -> list[dict[str, Any]]:
) -> list[dict[str, object]]:
"""
Expand MCP references to actual tool definitions.
@ -111,7 +119,7 @@ class SemanticToolFilterHook(CustomLogger):
)
# Convert Pydantic models to dicts for compatibility
openai_tools_as_dicts: Final = []
openai_tools_as_dicts: Final[list[dict[str, object]]] = []
for tool in openai_tools:
if hasattr(tool, "model_dump"):
tool_dict = tool.model_dump(exclude_none=True)
@ -141,8 +149,8 @@ class SemanticToolFilterHook(CustomLogger):
async def _filter_expanded_tools(
self,
data: dict,
expanded_tools: list[dict[str, Any]],
) -> list[dict[str, Any]]:
expanded_tools: list[dict[str, object]],
) -> list[dict[str, object]]:
"""
Apply the semantic filter to expanded MCP tool definitions.
@ -159,7 +167,7 @@ class SemanticToolFilterHook(CustomLogger):
return await self.filter.filter_tools(query=user_query, available_tools=expanded_tools)
def _selected_tool_names(self, filtered_tools: list[dict[str, Any]]) -> list[str]:
def _selected_tool_names(self, filtered_tools: Sequence[object]) -> list[str]:
"""Names of the semantically selected tools, as produced by the MCP expansion."""
names: Final = (self.filter._extract_tool_info(tool)[0] for tool in filtered_tools)
return [name for name in names if name]
@ -217,10 +225,10 @@ class SemanticToolFilterHook(CustomLogger):
def _emit_filter_metadata(
self,
data: dict,
mcp_tools: list[object],
filtered_mcp_tools: list[object],
native_tools: list[object],
filtered_tools: list[object],
mcp_tools: Sequence[object],
filtered_mcp_tools: Sequence[object],
native_tools: Sequence[object],
filtered_tools: Sequence[object],
) -> None:
"""
Emit response-header metadata when MCP tools were filtered.
@ -252,10 +260,10 @@ class SemanticToolFilterHook(CustomLogger):
def _emit_filter_metadata_safe(
self,
data: dict,
mcp_tools: list[object],
filtered_mcp_tools: list[object],
native_tools: list[object],
filtered_tools: list[object],
mcp_tools: Sequence[object],
filtered_mcp_tools: Sequence[object],
native_tools: Sequence[object],
filtered_tools: Sequence[object],
) -> None:
"""
Emit filter metadata without letting an emission failure abort the
@ -375,7 +383,7 @@ class SemanticToolFilterHook(CustomLogger):
)
if mcp_tools:
filtered_mcp_tools = await self.filter.filter_tools(
filtered_mcp_tools: list[object] = await self.filter.filter_tools(
query=user_query,
available_tools=mcp_tools,
)
@ -419,9 +427,9 @@ class SemanticToolFilterHook(CustomLogger):
self,
data: dict,
user_api_key_dict: "UserAPIKeyAuth",
response: Any,
response: object,
request_headers: dict[str, str] | None = None,
litellm_call_info: dict[str, Any] | None = None,
litellm_call_info: dict[str, object] | None = None,
) -> dict[str, str] | None:
"""Add semantic filter stats and tool names to response headers."""
from litellm.constants import MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH
@ -446,7 +454,7 @@ class SemanticToolFilterHook(CustomLogger):
return headers
def _get_tool_names_csv(self, tools: list[Any]) -> str:
def _get_tool_names_csv(self, tools: Sequence[object]) -> str:
"""Extract tool names and return as CSV string."""
if not tools:
return ""
@ -461,7 +469,7 @@ class SemanticToolFilterHook(CustomLogger):
@staticmethod
async def initialize_from_config(
config: dict[str, Any] | None,
config: SemanticToolFilterConfig | None,
llm_router: Optional["Router"],
) -> Optional["SemanticToolFilterHook"]:
"""

View file

@ -4,9 +4,10 @@ import json
import re
import time
from collections import OrderedDict
from collections.abc import Mapping, MutableMapping
from collections.abc import Mapping, MutableMapping, Sequence
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, cast
from fastapi import HTTPException, Request
from pydantic import ValidationError as PydanticValidationError
@ -55,7 +56,7 @@ from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_head
from litellm.types.integrations.anthropic_cache_control_hook import GATEWAY_INJECTED_CACHE_METADATA_KEY
# Cache special headers as a frozenset for O(1) lookup performance
_SPECIAL_HEADERS_CACHE: Final = frozenset(v.value.lower() for v in SpecialHeaders._member_map_.values())
_SPECIAL_HEADERS_CACHE: Final = frozenset(str(v.value).lower() for v in SpecialHeaders)
_REDACTED_HEADER_VALUE: Final = "***REDACTED***"
_CREDENTIAL_HEADER_NAMES: Final = SpecialHeaders.litellm_credential_header_names() | frozenset(
@ -126,7 +127,7 @@ def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None:
_ANTHROPIC_SESSION_ID_VALUE_RE: Final = re.compile(r"^[a-zA-Z0-9_\-]+$")
def _sanitize_for_log(value: Any) -> str:
def _sanitize_for_log(value: object) -> str:
"""
Basic log sanitization helper to reduce log-injection risk.
@ -164,7 +165,7 @@ _ENABLE_TEAM_STALE_ALIAS_BYPASS: bool | None = None
if TYPE_CHECKING:
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
from litellm.types.proxy.policy_engine import PolicyMatchContext
from litellm.types.proxy.policy_engine import Policy, PolicyMatchContext
ProxyConfig = _ProxyConfig
else:
@ -328,7 +329,7 @@ _ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_overr
_URL_DESTINATION_REQUEST_FIELDS: Final = ("model", "file_id")
def _reject_url_valued_destinations(data: dict[str, Any]) -> None:
def _reject_url_valued_destinations(data: dict[str, object]) -> None:
"""Reject URL-valued ``model``/``file_id`` unless admin-allowlisted.
Some providers (HuggingFace, Oobabooga, Gemini files) accept a URL in the
@ -387,7 +388,7 @@ def _invalid_metadata_type_error(field: str, value: object) -> ProxyException:
)
def _normalized_metadata_object(field: str, value: object) -> Mapping[str, Any]:
def _normalized_metadata_object(field: str, value: object) -> Mapping[str, object]:
"""Return ``value`` as a metadata object or raise a 400 like OpenAI does.
A JSON string that parses to an object is accepted because multipart/form-data
@ -402,6 +403,23 @@ def _normalized_metadata_object(field: str, value: object) -> Mapping[str, Any]:
raise _invalid_metadata_type_error(field=field, value=value)
def _normalized_metadata_slot(
request_data: MutableMapping[str, object], metadata_variable_name: str
) -> dict[str, object]:
"""Return the request's metadata slot as a dict, normalising it in place first.
Metadata can arrive as a JSON string (multipart/form-data, ``extra_body``). Parsing it here keeps
existing entries alive through a merge instead of silently overwriting them with an empty dict.
"""
raw: Final = request_data.get(metadata_variable_name)
if isinstance(raw, dict):
return raw
parsed: Final = safe_json_loads(raw) if isinstance(raw, str) else None
normalized: Final[dict[str, object]] = parsed if isinstance(parsed, dict) else {}
request_data[metadata_variable_name] = normalized
return normalized
def _strip_untrusted_request_header_controls(
headers: Any,
*,
@ -417,7 +435,7 @@ def _strip_untrusted_request_header_controls(
headers.pop(header_name, None)
def _is_false_like(value: Any) -> bool:
def _is_false_like(value: object) -> bool:
if isinstance(value, bool):
return value is False
if isinstance(value, str):
@ -462,7 +480,7 @@ def _key_or_team_allows_client_pricing_override(
)
def _strip_client_message_redaction_opt_out(data: dict[str, Any]) -> None:
def _strip_client_message_redaction_opt_out(data: dict[str, object]) -> None:
stripped: Final[list[str]] = []
if "turn_off_message_logging" in data and _is_false_like(data["turn_off_message_logging"]):
stripped.append("turn_off_message_logging")
@ -513,7 +531,7 @@ def _strip_client_callback_credentials(
)
def _strip_client_pricing_overrides(data: dict[str, Any]) -> None:
def _strip_client_pricing_overrides(data: dict[str, object]) -> None:
"""Drop pricing overrides from the request body and any metadata variant.
Skipped only when the calling key/team carries
@ -580,9 +598,9 @@ def _get_metadata_variable_name(request: Request) -> str:
def _promoted_trace_control_fields(
requester_metadata: Mapping[str, Any],
litellm_metadata: Mapping[str, Any],
) -> tuple[tuple[str, Any], ...]:
requester_metadata: Mapping[str, object],
litellm_metadata: Mapping[str, object],
) -> tuple[tuple[str, object], ...]:
"""Return the caller's trace-control fields that ``litellm_metadata`` does not already set."""
return tuple(
(key, value)
@ -1193,7 +1211,7 @@ class LiteLLMProxyRequestSetup:
def add_litellm_data_for_backend_llm_call(
*,
headers: dict,
request_data: Mapping[str, Any],
request_data: Mapping[str, object],
user_api_key_dict: UserAPIKeyAuth,
general_settings: dict[str, Any] | None = None,
) -> LitellmDataForBackendLLMCall:
@ -1327,6 +1345,8 @@ class LiteLLMProxyRequestSetup:
def get_sanitized_user_information_from_key(
user_api_key_dict: UserAPIKeyAuth,
) -> StandardLoggingUserAPIKeyMetadata:
stripped_metadata: Final = strip_callback_config(user_api_key_dict.metadata)
auth_metadata: Final = cast("dict[str, str] | None", stripped_metadata) # cast-ok: metadata is free-form JSON
user_api_key_logged_metadata: Final = StandardLoggingUserAPIKeyMetadata(
user_api_key_hash=user_api_key_dict.api_key, # just the hashed token
user_api_key_alias=user_api_key_dict.key_alias,
@ -1349,7 +1369,7 @@ class LiteLLMProxyRequestSetup:
user_api_key_budget_reset_at=(
user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None
),
user_api_key_auth_metadata=strip_callback_config(user_api_key_dict.metadata),
user_api_key_auth_metadata=auth_metadata,
)
return user_api_key_logged_metadata
@ -1577,14 +1597,7 @@ class LiteLLMProxyRequestSetup:
return
_metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(request_data)
metadata = request_data.get(_metadata_variable_name)
if isinstance(metadata, str):
parsed: Final = safe_json_loads(metadata)
metadata = parsed if isinstance(parsed, dict) else {}
request_data[_metadata_variable_name] = metadata
elif not isinstance(metadata, dict):
metadata = {}
request_data[_metadata_variable_name] = metadata
metadata: Final = _normalized_metadata_slot(request_data, _metadata_variable_name)
existing_tags: Final = metadata.get("tags")
metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags(
@ -1636,18 +1649,7 @@ class LiteLLMProxyRequestSetup:
# from (litellm_metadata vs metadata) so the merged tags are visible
# to _tag_max_budget_check.
_metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(request_data)
metadata = request_data.get(_metadata_variable_name)
# metadata can arrive as a JSON string (multipart/form-data, extra_body).
# Parse it so existing tags survive the merge — overwriting the string
# with {} would let a caller bypass _tag_max_budget_check on an
# over-budget body tag by also sending a within-budget header tag.
if isinstance(metadata, str):
parsed: Final = safe_json_loads(metadata)
metadata = parsed if isinstance(parsed, dict) else {}
request_data[_metadata_variable_name] = metadata
elif not isinstance(metadata, dict):
metadata = {}
request_data[_metadata_variable_name] = metadata
metadata: Final = _normalized_metadata_slot(request_data, _metadata_variable_name)
existing_tags: Final = metadata.get("tags")
metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags(
@ -1787,7 +1789,7 @@ async def add_litellm_data_to_request(
# admin-injection strip below so the audit / spend-tracking consumers of
# proxy_server_request["body"] see the cleaned metadata rather than
# attacker-forged user_api_key_* fields.
_litellm_received_at: Final = getattr(request.state, "litellm_received_at", None)
_litellm_received_at: Final[datetime | None] = getattr(request.state, "litellm_received_at", None)
arrival_time: Final = _litellm_received_at.timestamp() if _litellm_received_at is not None else time.time()
data["proxy_server_request"] = {
"url": str(request.url),
@ -2472,16 +2474,16 @@ def _resolve_provider_from_deployment(
if deployment is None:
continue
litellm_params = getattr(deployment, "litellm_params", None)
litellm_params: object = getattr(deployment, "litellm_params", None)
if litellm_params is None:
continue
custom_provider = getattr(litellm_params, "custom_llm_provider", None)
if custom_provider:
if isinstance(custom_provider, str) and custom_provider:
return custom_provider
deployment_model = getattr(litellm_params, "model", "") or ""
if "/" in deployment_model:
deployment_model = getattr(litellm_params, "model", "")
if isinstance(deployment_model, str) and "/" in deployment_model:
return deployment_model.split("/", 1)[0]
return None
@ -2904,8 +2906,8 @@ def _extract_policy_id(s: str) -> str | None:
def _match_and_track_policies(
data: dict,
context: "PolicyMatchContext",
request_body_policies: Any,
policies_override: dict[str, Any] | None = None,
request_body_policies: Sequence[str],
policies_override: dict[str, "Policy"] | None = None,
) -> tuple[list[str], dict[str, str]]:
"""
Match policies via attachments and request body, track them in metadata.
@ -2963,7 +2965,7 @@ def _apply_resolved_guardrails_to_metadata(
metadata_variable_name: str,
context: "PolicyMatchContext",
policy_names: list[str] | None = None,
policies: dict[str, Any] | None = None,
policies: dict[str, "Policy"] | None = None,
) -> None:
"""Apply resolved guardrails and pipelines to request metadata."""
from litellm._logging import verbose_proxy_logger
@ -3093,7 +3095,7 @@ async def add_guardrails_from_policy_engine(
request_body_names.append(item)
# Resolve policy versions by ID from in-memory cache (populated by sync job; no DB in hot path)
merged_policies: Final[dict[str, Any]] = dict(registry.get_all_policies())
merged_policies: Final[dict[str, Policy]] = dict(registry.get_all_policies())
fetched_policy_names: Final[list[str]] = []
for policy_id in request_body_version_ids:
result = registry.get_policy_by_id_for_request(policy_id=policy_id)

View file

@ -833,7 +833,7 @@ def _judge_collisions_for_team(
return tuple(
(role, model)
for role, model in (
*_router_arm_models(llm_router, data.router_name),
*(arm for name in data.router_names for arm in _router_arm_models(llm_router, name)),
*((("baseline", data.baseline_model),) if data.baseline_model is not None else ()),
)
if judge & judge_target(llm_router, model, team_id).models
@ -904,7 +904,7 @@ class _AttemptAggRow(BaseModel):
_ATTEMPT_AGG_ROWS: Final = TypeAdapter(list[_AttemptAggRow])
_ATTEMPT_AGG_SELECT: Final = """
_ATTEMPT_AGG_COLUMNS: Final = """
COUNT(*)::int AS turn_count,
COUNT(*) FILTER (WHERE outcome = 'real')::int AS real_wins,
COUNT(*) FILTER (WHERE outcome = 'shadow')::int AS shadow_wins,
@ -913,15 +913,34 @@ _ATTEMPT_AGG_SELECT: Final = """
COALESCE(SUM(real_cost + real_classifier_cost) FILTER (WHERE real_cost IS NOT NULL AND NOT real_cache_hit), 0)::float AS real_spend,
COALESCE(SUM(shadow_cost + shadow_classifier_cost) FILTER (WHERE real_cost IS NOT NULL AND NOT real_cache_hit), 0)::float AS shadow_spend,
COUNT(*) FILTER (WHERE real_cache_hit)::int AS cache_hit_turns
"""
_ATTEMPT_AGG_SELECT: Final = (
_ATTEMPT_AGG_COLUMNS
+ """
FROM "LiteLLM_ShadowEvalAttempt"
WHERE job_id = ANY($1::text[]) AND outcome != 'error'
GROUP BY 1
"""
)
_ATTEMPT_AGG_BY_TIER_SQL: Final = "SELECT COALESCE(tier, 'UNCLASSIFIED') AS grp," + _ATTEMPT_AGG_SELECT
_ATTEMPT_AGG_BY_MODEL_SQL: Final = "SELECT COALESCE(real_model, 'unknown') AS grp," + _ATTEMPT_AGG_SELECT
_ATTEMPT_AGG_BY_LEG_SQL: Final = "SELECT job_id AS grp," + _ATTEMPT_AGG_SELECT
# Attempt rows from before arm stamping carry no router_name; they belong to the job's
# own router, which the join reads off the leg.
_ATTEMPT_AGG_BY_ROUTER_SQL: Final = (
"SELECT COALESCE(a.router_name, j.router_name) AS grp,"
+ _ATTEMPT_AGG_COLUMNS
+ """
FROM "LiteLLM_ShadowEvalAttempt" a
JOIN "LiteLLM_ShadowEvalJob" j ON j.id = a.job_id
WHERE a.job_id = ANY($1::text[]) AND a.outcome != 'error'
GROUP BY 1
"""
)
# These guards derive spend from attempt rows, the cross-pod authority; the sampler also
# reads the live counter, so admission can stop before a row-based guard would fire (safe
# direction, and mid-deploy rows from old pods price as judge-only until the deploy ends).
@ -1060,6 +1079,7 @@ class _LegRow(BaseModel):
target_type: ShadowEvalTargetType
target_id: str
router_name: str
router_names: tuple[str, ...] = ()
direction: ShadowEvalDirection
baseline_model: str | None = None
judge_model: str
@ -1071,6 +1091,12 @@ class _LegRow(BaseModel):
stopped_at: datetime | None = None
stopped_by: str | None = None
@property
def arm_router_names(self) -> tuple[str, ...]:
"""The job's full router set; rows from before router_names existed hold it in
router_name alone. The one place that reading lives on the endpoint side."""
return self.router_names or (self.router_name,)
@field_validator("created_at", "ends_at", "stopped_at")
@classmethod
def _as_aware_utc(cls, value: datetime | None) -> datetime | None:
@ -1123,7 +1149,7 @@ def _group_response(
)
for leg in sorted(legs, key=lambda leg: (leg.target_type, leg.target_id))
),
router_name=first.router_name,
router_names=first.arm_router_names,
direction=first.direction,
baseline_model=first.baseline_model,
judge_model=first.judge_model,
@ -1252,6 +1278,9 @@ async def _shadow_eval_results(
for slice in _slices(by_leg)
}
)
by_router: Final = _ATTEMPT_AGG_ROWS.validate_python(
await _query_raw(prisma_client, _ATTEMPT_AGG_BY_ROUTER_SQL, leg_ids) or ()
)
total_turns: Final = sum(r.turn_count for r in by_tier)
funnel_rows: Final = await _query_raw(prisma_client, _FUNNEL_TOTALS_SQL, leg_ids)
counted: Final = _FunnelTotalsRow.model_validate(funnel_rows[0]) if funnel_rows else None
@ -1261,6 +1290,7 @@ async def _shadow_eval_results(
result: Final = ShadowEvalResult(
by_tier=_slices(by_tier),
by_current_model=_slices(by_model),
by_router=_slices(by_router),
overall_shadow_win_rate_pct=_pct_of(sum(r.shadow_wins for r in by_tier), total_turns),
overall_tie_rate_pct=_pct_of(sum(r.ties for r in by_tier), total_turns),
sampled_real_spend=sum(r.real_spend for r in by_tier),
@ -1314,8 +1344,15 @@ async def start_shadow_eval(
_require_admin_writer(user_api_key_dict, "start a shadow eval")
if prisma_client is None:
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
if llm_router is None or not _is_configured_pre_routing_strategy(llm_router, data.router_name):
raise HTTPException(status_code=400, detail=f"'{data.router_name}' is not a configured auto-router")
unconfigured: Final = tuple(
name
for name in data.router_names
if llm_router is None or not _is_configured_pre_routing_strategy(llm_router, name)
)
if unconfigured:
raise HTTPException(
status_code=400, detail=f"Not a configured auto-router: {', '.join(repr(n) for n in unconfigured)}"
)
token_rows: Final = (
await _verification_tokens(prisma_client).find_many(
where={"token": {"in": list(data.api_key_ids)}} # mutable-ok: Prisma filter
@ -1416,7 +1453,9 @@ async def start_shadow_eval(
ends_at: Final = now + timedelta(days=data.duration_days)
shared_config: Final = { # mutable-ok: Prisma payload
"group_id": group_id,
"router_name": data.router_name,
# a pre-router_names pod samples router_name alone, so it must be a real arm
"router_name": data.router_names[0],
"router_names": list(data.router_names), # mutable-ok: Prisma payload
"direction": data.direction,
"baseline_model": data.baseline_model,
"judge_model": data.judge_model,
@ -1477,7 +1516,7 @@ async def start_shadow_eval(
)
for target_type, target_id in sorted(requested_targets)
),
router_name=data.router_name,
router_names=data.router_names,
direction=data.direction,
baseline_model=data.baseline_model,
judge_model=data.judge_model,

View file

@ -20,6 +20,7 @@ import secrets
import traceback
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeVar, cast
import fastapi
@ -111,6 +112,7 @@ from litellm.proxy.management_helpers.team_member_permission_checks import (
TeamMemberPermissionChecks,
)
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
from litellm.proxy.spend_tracking.spend_tracking_utils import _is_master_key
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
get_ui_settings_cached,
@ -736,6 +738,45 @@ def _check_allowed_routes_caller_permission(
)
_READ_ONLY_ALLOWED_ROUTES_PRESET: Final = frozenset(("info_routes",))
def _is_safe_preset_route_transition(
incoming_allowed_routes: Sequence[str] | None,
existing_allowed_routes: Sequence[str] | None,
) -> bool:
"""
True when every route on BOTH sides is a safe `key_type` preset bucket
(empty = full access, which non-admins already get from a default
`/key/generate`), with one carve-out: a read-only (`info_routes`) key
stays read-only, so widening it needs an admin. Requiring the existing
side to be a safe preset keeps an owner from clearing an admin-set
custom route restriction (LIT-4139).
"""
incoming: Final = frozenset(incoming_allowed_routes or ())
existing: Final = frozenset(existing_allowed_routes or ())
if not (incoming | existing) <= _NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS:
return False
return existing != _READ_ONLY_ALLOWED_ROUTES_PRESET or incoming == existing
def _enforce_allowed_routes_update_permission(
data: UpdateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
if _is_safe_preset_route_transition(
incoming_allowed_routes=data.allowed_routes,
existing_allowed_routes=existing_key_row.allowed_routes,
):
return
_check_allowed_routes_caller_permission(
allowed_routes=data.allowed_routes,
user_api_key_dict=user_api_key_dict,
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
)
def _check_permissions_caller_permission(
data: GenerateRequestBase,
user_api_key_dict: UserAPIKeyAuth,
@ -2520,26 +2561,34 @@ async def _validate_mcp_servers_for_key_update(
return normalized_object_permission
def _require_prisma_client(prisma_client: PrismaClient | None) -> PrismaClient:
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "Database not connected"})
return prisma_client
async def _validate_update_key_data(
data: UpdateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
user_api_key_dict: UserAPIKeyAuth,
llm_router: Router | None,
premium_user: bool,
prisma_client: Any,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
) -> None:
"""Validate permissions and constraints for key update."""
checked_prisma_client: Final = _require_prisma_client(prisma_client)
# Reject NaN/±inf spend before it can reach the DB / spend counter.
validate_finite_spend(data.spend)
validate_budget_duration(data.budget_duration)
_is_proxy_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
_check_allowed_routes_caller_permission(
allowed_routes=data.allowed_routes,
_enforce_allowed_routes_update_permission(
data=data,
existing_key_row=existing_key_row,
user_api_key_dict=user_api_key_dict,
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
)
_check_passthrough_routes_caller_permission(
data=data,
@ -2567,7 +2616,7 @@ async def _validate_update_key_data(
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
user_api_key_dict=user_api_key_dict,
route=KeyManagementRoutes.KEY_UPDATE,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
existing_key_row=existing_key_row,
user_api_key_cache=user_api_key_cache,
)
@ -2648,12 +2697,12 @@ async def _validate_update_key_data(
# _check_key_admin_access that would otherwise require team/org admin status.
_key_is_team_key: Final = getattr(existing_key_row, "team_id", None) is not None
can_skip_admin_check: Final = (caller_is_creator or _key_is_team_key) and not _is_budget_change
if (not _is_proxy_admin) and prisma_client is not None and not can_skip_admin_check:
if (not _is_proxy_admin) and not can_skip_admin_check:
hashed_key: Final = existing_key_row.token
await _check_key_admin_access(
user_api_key_dict=user_api_key_dict,
hashed_token=hashed_key,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
route=("/key/update (max_budget/spend)" if _is_budget_change else "/key/update"),
)
@ -2664,7 +2713,7 @@ async def _validate_update_key_data(
if _team_id_to_check is not None:
team_obj = await get_team_object(
team_id=_team_id_to_check,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
)
@ -2680,7 +2729,7 @@ async def _validate_update_key_data(
await _check_team_key_limits(
team_table=team_obj,
data=data,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
)
TeamMemberPermissionChecks.enforce_member_can_assign_access_groups(
@ -2695,7 +2744,7 @@ async def _validate_update_key_data(
await _check_project_key_limits(
project_id=_project_id_to_check,
data=data,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
)
@ -2710,7 +2759,7 @@ async def _validate_update_key_data(
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
)
# Check org key limits only when throughput-related fields or organization_id change
@ -2726,7 +2775,7 @@ async def _validate_update_key_data(
org_table: Final = await get_org_object(
org_id=_org_id_to_check,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
)
if org_table is None:
raise HTTPException(
@ -2736,7 +2785,7 @@ async def _validate_update_key_data(
await _check_org_key_limits(
org_table=org_table,
data=data,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
)
# if team change - check if this is possible
@ -2766,7 +2815,7 @@ async def _validate_update_key_data(
data=data,
team_obj=team_obj,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
is_proxy_admin=_is_proxy_admin,
)
@ -3578,6 +3627,63 @@ async def _build_model_max_budget_usage(
)
def _window_max_budget(window: Mapping[str, object]) -> float | None:
"""A window's max_budget as a float; None when absent or unparseable."""
value: Final = window.get("max_budget")
if not isinstance(value, (int, float, str)):
return None
try:
return float(value)
except ValueError:
return None
async def _budget_window_usage(
window: Mapping[str, object], api_key_hash: str
) -> tuple[str, Mapping[str, object]] | None:
"""
(budget_duration, usage entry) for one budget window; None when the window
has no budget_duration to key it by.
Reads the same cross-pod counter (spend:key:{hashed_token}:window:{budget_duration})
that _virtual_key_multi_budget_check enforces against, passing the same
window_duration + window_start so a stale-low counter is re-checked against
the LiteLLM_BudgetWindowSpend row instead of a spend-log aggregate.
"""
from litellm.proxy.proxy_server import get_current_spend
duration: Final = window.get("budget_duration")
if not isinstance(duration, str) or not duration:
return None
spend: Final = await get_current_spend(
counter_key=f"spend:key:{api_key_hash}:window:{duration}",
fallback_spend=0.0,
max_budget=_window_max_budget(window),
window_entity_type="Key",
window_entity_id=api_key_hash,
window_duration=duration,
window_start=get_budget_window_start(window),
)
return duration, MappingProxyType({"current_spend": round(spend, 4)})
async def _build_budget_limits_usage(
budget_limits: Sequence[object] | str | None, api_key_hash: str
) -> Mapping[str, Mapping[str, object]] | None:
"""
Current-window spend per budget window, keyed by budget_duration, reported
next to the stored budget_limits (which is returned untouched). None when
the key has no windows, so the field only appears on keys that have them.
"""
windows: Final = _budget_limit_windows(budget_limits)
if not windows:
return None
usages: Final = await asyncio.gather(
*(_budget_window_usage(window=window, api_key_hash=api_key_hash) for window in windows)
)
return MappingProxyType({duration: usage for duration, usage in (u for u in usages if u is not None)})
@router.post(
"/v2/key/info",
tags=["key management"],
@ -3620,7 +3726,6 @@ async def info_key_fn_v2(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail={"message": "Malformed request. No keys passed in."},
)
# Resolve key_aliases to tokens so we never pass token=None (unbounded query)
tokens_to_query: Final = list(data.keys) if data.keys else []
if data.key_aliases:
@ -3662,6 +3767,13 @@ async def info_key_fn_v2(
model_max_budget=model_max_budget,
user_api_key_cache=model_max_budget_limiter.dual_cache,
)
if k_token_hash:
budget_limits_usage = await _build_budget_limits_usage(
budget_limits=k_dict.get("budget_limits"),
api_key_hash=k_token_hash,
)
if budget_limits_usage is not None:
k_dict["budget_limits_usage"] = budget_limits_usage
filtered_key_info.append(k_dict)
return {"key": data.keys, "info": filtered_key_info}
@ -3698,6 +3810,10 @@ async def info_key_fn(
- model_max_budget: dict - Per-model budgets, e.g. {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}
- model_max_budget_usage: dict | None - Current-window spend per model, present only when
the key has per-model budgets
- budget_limits: list | None - Concurrent budget windows, exactly as stored
- budget_limits_usage: dict | None - Current-window spend per budget window, e.g.
{"1h": {"current_spend": 0.0009}}, present only when the key has budget windows
(read from the same cross-pod spend counter the budget enforcement uses)
- models: list - Model_name's the key is allowed to call
- tpm_limit / rpm_limit: int | None - Tokens and requests per minute limits
- metadata: dict - Metadata for the key, e.g. {"team": "core-infra"}
@ -3765,7 +3881,7 @@ async def info_key_fn(
except Exception:
# if using pydantic v1
key_info = key_info.dict() # pyright: ignore[reportDeprecated] # deliberate pydantic v1 fallback
key_token_hash: Final = key_info.pop("token")
key_token_hash: Final[str | None] = key_info.pop("token")
model_max_budget = key_info.get("model_max_budget") or {}
budget_table: Final = key_info.get("litellm_budget_table") or {}
@ -3777,6 +3893,12 @@ async def info_key_fn(
model_max_budget=model_max_budget,
user_api_key_cache=model_max_budget_limiter.dual_cache,
)
budget_limits_usage: Final = await _build_budget_limits_usage(
budget_limits=key_info.get("budget_limits"),
api_key_hash=key_token_hash,
)
if budget_limits_usage is not None:
key_info["budget_limits_usage"] = budget_limits_usage
# Attach object_permission if object_permission_id is set
key_info = await attach_object_permission_to_dict(key_info, prisma_client)
@ -5221,7 +5343,7 @@ def _validate_reset_spend_value(reset_to: object, key_in_db: LiteLLM_Verificatio
max_budget = key_in_db.max_budget
if key_in_db.litellm_budget_table is not None:
budget_max_budget: Final = getattr(key_in_db.litellm_budget_table, "max_budget", None)
budget_max_budget: Final[float | None] = getattr(key_in_db.litellm_budget_table, "max_budget", None)
if budget_max_budget is not None:
if max_budget is None or budget_max_budget < max_budget:
max_budget = budget_max_budget

View file

@ -13,7 +13,7 @@ import copy
import json
import os
from collections.abc import AsyncIterator
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from typing import TYPE_CHECKING, Final, Literal, cast
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import Response, StreamingResponse
@ -90,7 +90,7 @@ class _ApplyPoliciesResultBase(TypedDict):
class ApplyPoliciesResult(_ApplyPoliciesResultBase, total=False):
"""Result of apply_policies. agent_response set when agent_id provided."""
agent_response: Any
agent_response: object
class _ApplyPoliciesPerItemResultBase(TypedDict):
@ -103,7 +103,7 @@ class _ApplyPoliciesPerItemResultBase(TypedDict):
class ApplyPoliciesPerItemResult(_ApplyPoliciesPerItemResultBase, total=False):
"""Result for one input when using inputs_list. agent_response set when agent_id provided."""
agent_response: Any
agent_response: object
class ApplyPoliciesListResult(TypedDict):
@ -295,8 +295,8 @@ async def test_policies_and_guardrails(
from litellm.proxy.proxy_server import chat_completion, proxy_logging_obj
from litellm.proxy.utils import handle_exception_on_proxy
def _serialize_chat_response(response: Any) -> Any:
if hasattr(response, "model_dump"):
def _serialize_chat_response(response: object) -> object:
if isinstance(response, BaseModel):
return response.model_dump(exclude_unset=True)
if isinstance(response, dict):
return response
@ -306,7 +306,7 @@ async def test_policies_and_guardrails(
inputs: GenericGuardrailAPIInputs,
agent_id: str,
user_api_key_dict: UserAPIKeyAuth,
) -> Any:
) -> object:
body: Final = _chat_body_from_inputs(inputs, agent_id, data.request_data)
req: Final = _request_with_json_body(body)
resp: Final = Response()

View file

@ -5148,6 +5148,7 @@ async def list_team_v2(
# Get teams with pagination
if use_deleted_table:
# LiteLLM_DeletedTeamTable has no litellm_model_table relation, unlike below
teams = await _deleted_team_db(prisma_client).find_many(
where=where_conditions,
skip=skip,
@ -5162,6 +5163,7 @@ async def list_team_v2(
skip=skip,
take=page_size,
order=order_by if order_by else {"created_at": "desc"}, # Default sort
include=_INCLUDE_MODEL_TABLE,
)
# Get total count for pagination
total_count = await _team_db(prisma_client).count(where=where_conditions)

View file

@ -14,7 +14,7 @@ import os
import re
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Any, Final, cast
from typing import TYPE_CHECKING, Annotated, Final, cast
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
@ -69,6 +69,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
)
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
from litellm.types.utils import LlmProviders
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
from litellm.utils import ProviderConfigManager
from .passthrough_endpoint_router import PassthroughEndpointRouter
@ -121,7 +122,21 @@ def is_passthrough_request_streaming(request_body: object) -> bool:
return bool(request_body.get("stream", False))
def get_passthrough_router_request_metadata(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, Any]:
def _optional_str(value: object) -> str | None:
return value if isinstance(value, str) else None
def _string_keyed_mapping(value: object) -> Mapping[str, object] | None:
if isinstance(value, Mapping):
return value
return None
async def _json_request_body(request: Request) -> Mapping[str, object]:
return await request.json()
def get_passthrough_router_request_metadata(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, object]:
"""
Build the request metadata carrying key-level spend attribution and the
pre-call budget reservation for a router-model passthrough request.
@ -210,7 +225,7 @@ async def llm_passthrough_factory_proxy_route(
# anthropic is streaming when 'stream' = True is in the body
if request.method == "POST":
if "multipart/form-data" not in request.headers.get("content-type", ""):
_request_body = await request.json()
_request_body = await _json_request_body(request)
else:
_request_body = await get_form_data(request)
@ -383,7 +398,7 @@ async def vllm_proxy_route(
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=_safe_get_request_headers(request),
stream=request_body.get("stream", False),
stream=is_streaming_request,
content=None,
data=None,
files=None,
@ -817,7 +832,7 @@ async def handle_bedrock_passthrough_router_model(
# Use the common processing path (same as non-router models)
# This ensures all metadata, hooks, and logging are properly initialized
data: Final[dict[str, Any]] = {}
data: Final[dict[str, object]] = {}
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
data["model"] = model
@ -861,8 +876,8 @@ async def handle_bedrock_count_tokens(
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth,
request_body: dict[str, Any],
) -> dict[str, Any]:
request_body: dict[str, object],
) -> dict[str, object]:
"""
Handle AWS Bedrock CountTokens API requests.
@ -879,7 +894,7 @@ async def handle_bedrock_count_tokens(
handler: Final = BedrockCountTokensHandler()
# Extract model from request body
model: Final = request_body.get("model")
model: Final = _optional_str(request_body.get("model"))
if not model:
raise HTTPException(status_code=400, detail={"error": "Model is required in request body"})
@ -1011,7 +1026,7 @@ async def bedrock_llm_proxy_route(
"Bedrock passthrough: Using direct Bedrock model '%s' for endpoint '%s'", model, endpoint
)
data: Final[dict[str, Any]] = {}
data: Final[dict[str, object]] = {}
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
data["method"] = request.method
@ -1110,7 +1125,7 @@ async def bedrock_proxy_route(
headers: Final = {"Content-Type": "application/json"}
# Assuming the body contains JSON data, parse it
try:
data: Final = await request.json()
data: Final = await _json_request_body(request)
except Exception as e:
raise HTTPException(status_code=400, detail={"error": e})
_request: Final = AWSRequest(method="POST", url=str(updated_url), data=json.dumps(data), headers=headers)
@ -1201,7 +1216,7 @@ async def comprehend_medical_proxy_route(
)
try:
data: Final = await request.json()
data: Final = await _json_request_body(request)
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
@ -1412,7 +1427,7 @@ async def assemblyai_proxy_route(
is_streaming_request = False
# assemblyai is streaming when 'stream' = True is in the body
if request.method == "POST":
_request_body: Final = await request.json()
_request_body: Final = await _json_request_body(request)
if _request_body.get("stream"):
is_streaming_request = True
@ -1519,7 +1534,7 @@ async def azure_proxy_route(
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=_safe_get_request_headers(request),
stream=request_body.get("stream", False),
stream=is_streaming_request,
content=None,
data=None,
files=None,
@ -1606,7 +1621,7 @@ async def azure_proxy_route(
extra_headers = auth_credentials.get("headers") or {}
base_target_url = litellm_params.get("api_base")
base_target_url = _optional_str(litellm_params.get("api_base"))
if base_target_url is None:
raise Exception(f"API base not found for {part}")
return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler(
@ -1727,7 +1742,7 @@ def get_vertex_pass_through_handler(
def _override_vertex_params_from_router_credentials(
router_credentials: Any | None,
router_credentials: LiteLLM_ManagedVectorStore | None,
vertex_project: str | None,
vertex_location: str | None,
) -> tuple[str | None, str | None]:
@ -1747,14 +1762,14 @@ def _override_vertex_params_from_router_credentials(
verbose_proxy_logger.debug("Using vector store credentials to override vertex project and location")
litellm_params: Final = router_credentials.get("litellm_params", {})
litellm_params: Final = _string_keyed_mapping(router_credentials.get("litellm_params"))
if not litellm_params:
verbose_proxy_logger.warning("Vector store credentials found but litellm_params is empty")
return vertex_project, vertex_location
# Extract vertex_project and vertex_location from litellm_params
vector_store_project: Final = litellm_params.get("vertex_project")
vector_store_location: Final = litellm_params.get("vertex_location")
vector_store_project: Final = _optional_str(litellm_params.get("vertex_project"))
vector_store_location: Final = _optional_str(litellm_params.get("vertex_location"))
if vector_store_project:
verbose_proxy_logger.debug(
@ -1762,7 +1777,6 @@ def _override_vertex_params_from_router_credentials(
vertex_project,
vector_store_project,
)
vertex_project = vector_store_project
else:
verbose_proxy_logger.warning("Vector store credentials found but missing vertex_project in litellm_params")
@ -1772,11 +1786,10 @@ def _override_vertex_params_from_router_credentials(
vertex_location,
vector_store_location,
)
vertex_location = vector_store_location
else:
verbose_proxy_logger.warning("Vector store credentials found but missing vertex_location in litellm_params")
return vertex_project, vertex_location
return vector_store_project or vertex_project, vector_store_location or vertex_location
_CREDENTIALLESS_VERTEX_MISSING_CREDENTIAL_DETAIL: Final = (
@ -1884,8 +1897,8 @@ def _forwarded_headers_for_credentialless_vertex_passthrough(
async def _prepare_vertex_auth_headers(
request: Request,
vertex_credentials: Any | None,
router_credentials: Any | None,
vertex_credentials: VertexPassThroughCredentials | None,
router_credentials: LiteLLM_ManagedVectorStore | None,
vertex_project: str | None,
vertex_location: str | None,
base_target_url: str | None,
@ -1982,7 +1995,7 @@ async def _base_vertex_proxy_route(
fastapi_response: Response,
get_vertex_pass_through_handler: BaseVertexAIPassThroughHandler,
user_api_key_dict: UserAPIKeyAuth | None = None,
router_credentials: Any | None = None,
router_credentials: LiteLLM_ManagedVectorStore | None = None,
):
"""
Base function for Vertex AI passthrough routes.
@ -2152,8 +2165,6 @@ async def vertex_discovery_proxy_route(
"""
import re
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
# Extract vector store ID from endpoint if present (e.g., dataStores/test-litellm-app_1761094730750)
vector_store_credentials: LiteLLM_ManagedVectorStore | None = None
vector_store_id_match: Final = re.search(r"dataStores/([^/]+)", endpoint)
@ -3098,7 +3109,7 @@ async def watsonx_proxy_route(
is_streaming_request = False
if request.method == "POST":
if "multipart/form-data" not in request.headers.get("content-type", ""):
_request_body = await request.json()
_request_body = await _json_request_body(request)
else:
_request_body = await get_form_data(request)

View file

@ -6,9 +6,10 @@ import posixpath
import traceback
from base64 import b64encode
from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from itertools import groupby
from typing import Any, Final, TypedDict, cast
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
from urllib.parse import urlencode, urlparse
import httpx
@ -47,6 +48,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
)
from litellm.litellm_core_utils.initialize_dynamic_callback_params import validate_no_callback_env_reference
from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
@ -78,7 +80,10 @@ from litellm.proxy.common_utils.http_parsing_utils import (
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
_get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
)
from litellm.proxy.utils import normalize_route_for_root_path
from litellm.repositories.team_repository import TeamRepository
from litellm.secret_managers.main import get_secret_str
@ -90,7 +95,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
EndpointType,
PassthroughStandardLoggingPayload,
)
from litellm.types.utils import Usage
from litellm.types.utils import TRUSTED_CALLBACK_VARS_FIELD, Usage
from .streaming_handler import PassThroughStreamingHandler
from .success_handler import PassThroughEndpointLogging
@ -99,6 +104,9 @@ from .upstream_usage_headers import (
apply_upstream_reported_usage,
)
if TYPE_CHECKING:
from litellm.proxy.proxy_server import ProxyConfig
router: Final = APIRouter()
pass_through_endpoint_logging: Final = PassThroughEndpointLogging()
@ -752,6 +760,67 @@ def _build_passthrough_failure_request_payload(
return request_payload
@dataclass(frozen=True, slots=True)
class _TeamCallbackWiring:
success_callbacks: "list[str | Callable | CustomLogger] | None" = None # mutable-ok: Logging.__init__ arg
failure_callbacks: "list[str | Callable | CustomLogger] | None" = None # mutable-ok: Logging.__init__ arg
logging_kwargs: dict[str, str | dict[str, str]] | None = None # mutable-ok: Logging.__init__ arg
def _resolve_team_callback_wiring(
user_api_key_dict: UserAPIKeyAuth,
proxy_config: "ProxyConfig",
route_description: str,
) -> _TeamCallbackWiring:
"""Resolve key/team dynamic logging callbacks for a passthrough request.
Mirrors add_litellm_data_to_request: callback_vars are unpacked top-level
(read by initialize_standard_callback_dynamic_params) and also stamped on
the proxy-owned trusted-vars field (read by get_trusted_callback_params).
Fails open: a callback resolution or validation error is logged at error
level and the request proceeds without dynamic callbacks, since a broken
logging config must not fail the customer's upstream call (and the
websocket is already accepted by the time this runs on that path). The
env-reference check runs here because the deprecated callback_settings
branch skips AddTeamCallback validation, and Logging.__init__ would
otherwise reject the vars mid-request.
"""
try:
callback_settings_obj: Final = _get_dynamic_logging_metadata(
user_api_key_dict=user_api_key_dict, proxy_config=proxy_config
)
if callback_settings_obj and callback_settings_obj.callback_vars:
for (
item
) in callback_settings_obj.callback_vars.items(): # rebind-ok: dict.items iteration for env-ref validation
validate_no_callback_env_reference(item[0], item[1], source="key/team callback metadata")
except Exception: # noqa: BLE001 - a broken logging config must never fail the passthrough request
verbose_proxy_logger.exception(
"%s: failed to resolve team logging callbacks, continuing without them",
route_description,
)
return _TeamCallbackWiring()
if callback_settings_obj is None:
return _TeamCallbackWiring()
callback_vars: Final = callback_settings_obj.callback_vars
success_callbacks: Final = callback_settings_obj.success_callback
failure_callbacks: Final = callback_settings_obj.failure_callback
logging_kwargs: Final = (
None
if not callback_vars
else { # mutable-ok: Logging arg
**callback_vars,
TRUSTED_CALLBACK_VARS_FIELD: callback_vars,
}
)
return _TeamCallbackWiring(
success_callbacks=None if success_callbacks is None else [*success_callbacks], # mutable-ok: Logging arg
failure_callbacks=None if failure_callbacks is None else [*failure_callbacks], # mutable-ok: Logging arg
logging_kwargs=logging_kwargs,
)
async def _log_passthrough_upstream_failure(
response: httpx.Response,
user_api_key_dict: UserAPIKeyAuth,
@ -845,7 +914,7 @@ async def pass_through_request(
from litellm.proxy.pass_through_endpoints.passthrough_guardrails import (
PassthroughGuardrailHandler,
)
from litellm.proxy.proxy_server import proxy_logging_obj
from litellm.proxy.proxy_server import proxy_config, proxy_logging_obj
#########################################################
# Initialize variables
@ -930,6 +999,11 @@ async def pass_through_request(
# read e.g. ``chat gpt-4o`` instead of ``chat unknown``.
passthrough_model: Final = (_parsed_body.get("model") if isinstance(_parsed_body, dict) else None) or "unknown"
start_time: Final = datetime.now()
team_callbacks: Final = _resolve_team_callback_wiring(
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
route_description="pass_through_endpoint",
)
logging_obj = Logging(
model=passthrough_model,
messages=[{"role": "user", "content": safe_dumps(_parsed_body)}],
@ -938,6 +1012,9 @@ async def pass_through_request(
start_time=start_time,
litellm_call_id=litellm_call_id,
function_id="1245",
dynamic_success_callbacks=team_callbacks.success_callbacks,
dynamic_failure_callbacks=team_callbacks.failure_callbacks,
kwargs=team_callbacks.logging_kwargs,
)
# Store passthrough guardrails config on logging_obj for field targeting
@ -2022,7 +2099,7 @@ async def websocket_passthrough_request(
setup_model_rewriter: Optional rewrite of the setup frame's model before it reaches the upstream
"""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.proxy_server import proxy_logging_obj
from litellm.proxy.proxy_server import proxy_config, proxy_logging_obj
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
PassthroughStandardLoggingPayload,
)
@ -2055,6 +2132,11 @@ async def websocket_passthrough_request(
upstream_headers[header_name] = header_value
# Initialize logging object similar to HTTP passthrough
team_callbacks: Final = _resolve_team_callback_wiring(
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
route_description="websocket_passthrough",
)
logging_obj: Final = Logging(
model="unknown",
messages=[{"role": "user", "content": "WebSocket connection"}],
@ -2063,6 +2145,9 @@ async def websocket_passthrough_request(
start_time=start_time,
litellm_call_id=litellm_call_id,
function_id="websocket_passthrough",
dynamic_success_callbacks=team_callbacks.success_callbacks,
dynamic_failure_callbacks=team_callbacks.failure_callbacks,
kwargs=team_callbacks.logging_kwargs,
)
# Create passthrough logging payload
@ -3148,6 +3233,14 @@ def _get_pass_through_endpoints_from_config() -> list[PassThroughGenericEndpoint
return returned_endpoints
def _config_field_endpoints(response: ConfigFieldInfo) -> list[object] | None:
return response.field_value
def _request_app(request: Request) -> FastAPI:
return request.app
async def _get_pass_through_endpoints_from_db(
endpoint_id: str | None = None,
user_api_key_dict: UserAPIKeyAuth | None = None,
@ -3164,7 +3257,7 @@ async def _get_pass_through_endpoints_from_db(
except Exception:
return []
pass_through_endpoint_data: Final[list | None] = response.field_value
pass_through_endpoint_data: Final = _config_field_endpoints(response)
if pass_through_endpoint_data is None:
return []
@ -3327,7 +3420,7 @@ async def update_pass_through_endpoints(
detail={"error": "No pass-through endpoints found"},
)
pass_through_endpoint_data: Final[list | None] = response.field_value
pass_through_endpoint_data: Final[list | None] = _config_field_endpoints(response)
if pass_through_endpoint_data is None:
raise HTTPException(
status_code=404,
@ -3398,7 +3491,7 @@ async def update_pass_through_endpoints(
_custom_headers: dict | None = updated_endpoint.headers or {}
_custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers)
route_app: Final[FastAPI] = request.app
route_app: Final = _request_app(request)
if updated_endpoint.include_subpath:
InitPassThroughEndpointHelpers.add_subpath_route(
app=route_app,
@ -3490,7 +3583,7 @@ async def create_pass_through_endpoints(
_custom_headers: dict | None = created_endpoint.headers or {}
_custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers)
route_app: Final[FastAPI] = request.app
route_app: Final = _request_app(request)
if created_endpoint.include_subpath:
InitPassThroughEndpointHelpers.add_subpath_route(
app=route_app,
@ -3558,7 +3651,7 @@ async def delete_pass_through_endpoints(
response = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None)
## Update field by removing endpoint
pass_through_endpoint_data: Final[list | None] = response.field_value
pass_through_endpoint_data: Final[list | None] = _config_field_endpoints(response)
if response.field_value is None or pass_through_endpoint_data is None:
raise HTTPException(
status_code=400,

View file

@ -263,6 +263,7 @@ from litellm.litellm_core_utils.agentic_loop_settings import (
validated_max_agentic_loops,
)
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.audio_utils.utils import resolve_speech_media_type
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_litellm_metadata_from_kwargs,
@ -11061,15 +11062,14 @@ async def audio_speech(
if callback_headers:
custom_headers.update(callback_headers)
# Determine media type based on model type
media_type = "audio/mpeg" # Default for OpenAI TTS
request_model: Final = data.get("model", "")
if request_model:
request_model_lower: Final = request_model.lower()
if "gemini" in request_model_lower and (
"tts" in request_model_lower or "preview-tts" in request_model_lower
):
media_type = "audio/wav" # Gemini TTS returns WAV format after conversion
requested_format: Final = data.get("response_format")
upstream_content_type: Final = (
response.response.headers.get("content-type") if isinstance(response, HttpxBinaryResponseContent) else None
)
media_type: Final = resolve_speech_media_type(
upstream_content_type=upstream_content_type,
response_format=requested_format if isinstance(requested_format, str) else None,
)
return StreamingResponse(
_audio_speech_chunk_generator(response),
@ -11085,7 +11085,15 @@ async def audio_speech(
)
verbose_proxy_logger.error("litellm.proxy.proxy_server.audio_speech(): Exception occured - %s", e)
verbose_proxy_logger.debug(traceback.format_exc())
raise e
if isinstance(e, (ProxyException, HTTPException)):
raise e
raise ProxyException(
message=getattr(e, "message", f"{e}"),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
openai_code=getattr(e, "code", None),
code=getattr(e, "status_code", 500),
)
@router.post(

View file

@ -7,12 +7,14 @@ Provides:
"""
import base64
import json
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import orjson
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from fastapi.responses import ORJSONResponse, StreamingResponse
from starlette.datastructures import UploadFile
import litellm
from litellm._logging import verbose_proxy_logger
@ -45,6 +47,16 @@ if TYPE_CHECKING:
router: Final = APIRouter()
def _as_string_keyed_mapping(value: object) -> Mapping[str, object] | None:
if isinstance(value, Mapping):
return value
return None
def _response_attr(source: object, name: str) -> object:
return getattr(source, name, None)
def _raise_vector_store_scan_depth_exceeded() -> None:
raise HTTPException(
status_code=400,
@ -53,8 +65,8 @@ def _raise_vector_store_scan_depth_exceeded() -> None:
def _append_payload_to_scan_stack(
payload_stack: list[tuple[Any, int]],
value: Any,
payload_stack: list[tuple[object, int]],
value: object,
next_depth: int,
) -> None:
if isinstance(value, dict):
@ -117,7 +129,7 @@ async def _authorize_nested_vector_store_ids(
def _build_file_metadata_entry(
response: Any,
response: object,
file_data: tuple[str, bytes, str] | None = None,
file_url: str | None = None,
) -> Mapping[str, str | int | None]:
@ -135,11 +147,11 @@ def _build_file_metadata_entry(
from datetime import datetime, timezone
# Extract file_id from response
file_id = None
if hasattr(response, "get"):
file_id = response.get("file_id")
elif hasattr(response, "file_id"):
file_id = response.file_id
mapping_response: Final = _as_string_keyed_mapping(response)
raw_file_id: Final = (
mapping_response.get("file_id") if mapping_response is not None else _response_attr(response, "file_id")
)
file_id: Final = raw_file_id if isinstance(raw_file_id, str) else None
# Extract file information from file_data tuple
filename = None
@ -152,7 +164,7 @@ def _build_file_metadata_entry(
content_type = file_data[2] if len(file_data) > 2 else None
# Build file metadata entry
file_entry: Final = {
file_entry: Final[dict[str, str | int | None]] = {
"file_id": file_id,
"filename": filename,
"file_url": file_url,
@ -169,7 +181,7 @@ def _build_file_metadata_entry(
async def _save_vector_store_to_db_from_rag_ingest(
response: Any,
response: object,
ingest_options: Mapping[str, dict[str, str | None]],
prisma_client: "PrismaClient",
user_api_key_dict: UserAPIKeyAuth,
@ -197,10 +209,11 @@ async def _save_vector_store_to_db_from_rag_ingest(
)
# Handle both dict and object responses
if hasattr(response, "get"):
vector_store_id = response.get("vector_store_id")
mapping_response: Final = _as_string_keyed_mapping(response)
if mapping_response is not None:
vector_store_id = mapping_response.get("vector_store_id")
elif hasattr(response, "vector_store_id"):
vector_store_id = response.vector_store_id
vector_store_id = _response_attr(response, "vector_store_id")
else:
verbose_proxy_logger.warning("Unable to extract vector_store_id from response type: %s", type(response))
return
@ -266,14 +279,13 @@ async def _save_vector_store_to_db_from_rag_ingest(
verbose_proxy_logger.info("Vector store %s already exists, appending file to metadata", vector_store_id)
# Update existing vector store with new file
existing_metadata = existing_vector_store.vector_store_metadata or {}
if isinstance(existing_metadata, str):
import json
stored_metadata: Final = existing_vector_store.vector_store_metadata or {}
existing_metadata: dict[str, object] = (
json.loads(stored_metadata) if isinstance(stored_metadata, str) else stored_metadata
)
existing_metadata = json.loads(existing_metadata)
ingested_files: Final = existing_metadata.get("ingested_files", [])
ingested_files.append(file_entry)
previous_files: Final = existing_metadata.get("ingested_files", [])
ingested_files: Final = [*previous_files, file_entry] if isinstance(previous_files, list) else [file_entry]
existing_metadata["ingested_files"] = ingested_files
# Update the vector store
@ -340,9 +352,9 @@ async def parse_rag_ingest_request(
# Get file
file_obj = form_data.get("file")
if file_obj is not None and hasattr(file_obj, "read"):
if isinstance(file_obj, UploadFile):
file_content = await file_obj.read(MAX_UPLOAD_SIZE_BYTES + 1)
file_data = (file_obj.filename, file_content, file_obj.content_type)
file_data = (file_obj.filename or "", file_content, file_obj.content_type or "")
# Parse JSON from 'request' form field (contains full request body as JSON)
request_json_str: Final[str | bytes | None] = form_data.get("request")

View file

@ -76,17 +76,17 @@ class _StreamEventParser:
async def background_streaming_task(
polling_id: str,
data: dict,
data: dict[str, object],
polling_handler: ResponsePollingHandler,
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth,
general_settings: dict,
general_settings: dict[str, object],
llm_router: "Router | None",
proxy_config: "ProxyConfig",
proxy_logging_obj: "ProxyLogging",
select_data_generator,
user_model,
select_data_generator: Callable[..., object] | None,
user_model: str | None,
user_temperature: float | None,
user_request_timeout: float | None,
user_max_tokens: int | None,

View file

@ -1531,7 +1531,8 @@ model LiteLLM_ShadowEvalJob {
group_id String // legs of one job share this; the API's job id
target_type String @default("key") // key | team | user
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
router_name String // the auto-router under evaluation, in either direction
router_name String // first (often only) auto-router under evaluation; router_names is the full set
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String
@ -1555,6 +1556,7 @@ model LiteLLM_ShadowEvalAttempt {
job_id String
request_id String // the judged real request
outcome String // real | shadow | tie | error
router_name String? // the arm this verdict scores; NULL on legacy rows, meaning the job's own router
tier String? // router's tier for the prompt, when classified
real_model String?
shadow_model String?

View file

@ -6,7 +6,7 @@ from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import Any, Final, NoReturn, cast
from typing import Final, NoReturn, SupportsFloat, SupportsIndex, SupportsInt, cast
from fastapi import HTTPException, status
@ -35,6 +35,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.router import Router
from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget
from litellm.types.router import DeploymentTypedDict
@dataclass
@ -697,7 +698,7 @@ def _get_budget_limit_counters(
for window in budget_limits:
window_dict = _coerce_window(window)
budget_duration = window_dict.get("budget_duration")
max_budget = window_dict.get("max_budget")
max_budget = _to_float(window_dict.get("max_budget"))
if not budget_duration or max_budget is None or max_budget <= 0:
continue
window_start = get_budget_window_start(window_dict)
@ -724,18 +725,20 @@ def _get_budget_limit_counters(
return counters
def _coerce_window(window: Any) -> dict:
if isinstance(window, dict):
def _coerce_window(window: object) -> Mapping[str, object]:
if isinstance(window, Mapping):
return window
if isinstance(window, str):
try:
parsed: Final = json.loads(window)
return parsed if isinstance(parsed, dict) else {}
parsed: Final[object] = json.loads(window)
except Exception:
return {}
if hasattr(window, "model_dump"):
return window.model_dump()
return {}
return parsed if isinstance(parsed, Mapping) else {}
model_dump: Final = getattr(window, "model_dump", None)
if not callable(model_dump):
return {}
dumped: Final[object] = model_dump()
return dumped if isinstance(dumped, Mapping) else {}
async def _reserve_counter(
@ -953,7 +956,7 @@ def _get_entry_reserved_cost(entry: dict, default_reserved_cost: float) -> float
return default_reserved_cost
def get_budget_window_start(window: Any) -> datetime | None:
def get_budget_window_start(window: object) -> datetime | None:
window_dict: Final = _coerce_window(window)
budget_duration: Final = window_dict.get("budget_duration")
if budget_duration is None:
@ -971,7 +974,7 @@ def get_budget_window_start(window: Any) -> datetime | None:
return reset_at - timedelta(seconds=duration_seconds)
def _coerce_datetime(value: Any) -> datetime | None:
def _coerce_datetime(value: object) -> datetime | None:
if value is None:
return None
if isinstance(value, datetime):
@ -1245,11 +1248,11 @@ def _get_model_cost_infos(
def _deployment_tiered_pricing_table(
deployment: dict[str, Any],
deployment: DeploymentTypedDict,
llm_router: Router,
) -> list[dict] | None:
model_id: Final = deployment.get("model_info", {}).get("id")
backend_model: Final = deployment.get("litellm_params", {}).get("model")
) -> Sequence[Mapping[str, object]] | None:
model_id: Final = _get_value(_get_value(deployment, "model_info"), "id")
backend_model: Final = _get_value(_get_value(deployment, "litellm_params"), "model")
if not isinstance(model_id, str) or not isinstance(backend_model, str):
return None
deployment_model_info: Final = llm_router.get_deployment_model_info(model_id=model_id, model_name=backend_model)
@ -1414,7 +1417,7 @@ def _estimate_output_tokens(
return min(requested, model_ceiling)
def _count_text_tokens(model: str, text: Any) -> int:
def _count_text_tokens(model: str, text: object) -> int:
if text is None:
return 0
@ -1454,8 +1457,8 @@ def _is_input_only_route(route: str) -> bool:
)
def _to_float(value: Any) -> float | None:
if value is None:
def _to_float(value: object) -> float | None:
if not isinstance(value, (SupportsFloat, SupportsIndex, str, bytes, bytearray)):
return None
try:
return float(value)
@ -1463,8 +1466,8 @@ def _to_float(value: Any) -> float | None:
return None
def _to_int(value: Any) -> int | None:
if value is None:
def _to_int(value: object) -> int | None:
if not isinstance(value, (SupportsInt, SupportsIndex, str, bytes, bytearray)):
return None
try:
return int(value)
@ -1472,7 +1475,7 @@ def _to_int(value: Any) -> int | None:
return None
def _get_value(obj: Any, key: str) -> Any:
if isinstance(obj, dict):
def _get_value(obj: object, key: str) -> object:
if isinstance(obj, Mapping):
return obj.get(key)
return getattr(obj, key, None)

View file

@ -1,10 +1,10 @@
import os
import re
import secrets
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from datetime import datetime as dt
from typing import Any, Final, Literal, cast
from typing import Final, Literal, Protocol, cast, runtime_checkable
from pydantic import BaseModel
@ -222,7 +222,28 @@ def get_spend_logs_id(call_type: str, response_obj: dict, kwargs: dict) -> str |
return resolved_id
def _extract_usage_for_ocr_call(response_obj: Any, response_obj_dict: dict) -> dict:
_MISSING_ATTRIBUTE: Final = object()
def _attribute_or_missing(source: object, name: str) -> object:
return getattr(source, name, _MISSING_ATTRIBUTE)
@runtime_checkable
class _ModelDumpable(Protocol):
def model_dump(self) -> object: ...
def _dumped_usage_info(usage_info: object) -> object:
if isinstance(usage_info, _ModelDumpable):
return usage_info.model_dump()
instance_dict: Final = _attribute_or_missing(usage_info, "__dict__")
if instance_dict is not _MISSING_ATTRIBUTE:
return instance_dict
return usage_info
def _extract_usage_for_ocr_call(response_obj: object, response_obj_dict: dict) -> dict:
"""
Extract usage information for OCR/AOCR calls.
@ -243,12 +264,10 @@ def _extract_usage_for_ocr_call(response_obj: Any, response_obj_dict: dict) -> d
usage_info = response_obj_dict.get("usage_info")
# Try to extract usage_info from object attributes if not found in dict
if not usage_info and hasattr(response_obj, "usage_info"):
usage_info = response_obj.usage_info
if hasattr(usage_info, "model_dump"):
usage_info = usage_info.model_dump()
elif hasattr(usage_info, "__dict__"):
usage_info = vars(usage_info)
if not usage_info:
attribute_usage_info: Final = _attribute_or_missing(response_obj, "usage_info")
if attribute_usage_info is not _MISSING_ATTRIBUTE:
usage_info = _dumped_usage_info(attribute_usage_info)
# For OCR, we track pages instead of tokens
if usage_info is not None:
@ -620,6 +639,14 @@ def _ensure_datetime_utc(timestamp: datetime) -> datetime:
return timestamp
async def _query_raw_rows(
prisma_client: PrismaClient,
sql_query: str,
*args: object,
) -> Sequence[Mapping[str, object]] | None:
return await prisma_client.db.query_raw(sql_query, *args)
async def get_spend_by_team(
start_date: dt,
end_date: dt,
@ -681,7 +708,7 @@ async def get_spend_by_team(
group_by_day;
"""
db_response: Final = await prisma_client.db.query_raw(sql_query, start_date, end_date, team_id)
db_response: Final = await _query_raw_rows(prisma_client, sql_query, start_date, end_date, team_id)
if db_response is None:
return []
@ -756,7 +783,7 @@ async def get_spend_by_team_and_customer(
group_by_day;
"""
db_response: Final = await prisma_client.db.query_raw(sql_query, start_date, end_date, team_id, customer_id)
db_response: Final = await _query_raw_rows(prisma_client, sql_query, start_date, end_date, team_id, customer_id)
if db_response is None:
return []
@ -811,7 +838,7 @@ def _sanitize_request_body_for_spend_logs_payload(
return {}
visited.add(obj_id)
def _sanitize_value(value: Any) -> Any:
def _sanitize_value(value: object) -> object:
if isinstance(value, dict):
return _sanitize_request_body_for_spend_logs_payload(value, visited, max_string_length_prompt_in_db)
elif isinstance(value, list):
@ -1106,7 +1133,7 @@ def _sanitize_error_information_for_spend_logs(
return cast(StandardLoggingPayloadErrorInformation, sanitized)
def _convert_to_json_serializable_dict(obj: Any, visited: set | None = None, max_depth: int = 20) -> Any:
def _convert_to_json_serializable_dict(obj: object, visited: set[int] | None = None, max_depth: int = 20) -> object:
"""
Convert object to JSON-serializable dict, handling Pydantic models safely.
@ -1160,6 +1187,13 @@ def _convert_to_json_serializable_dict(obj: Any, visited: set | None = None, max
visited.remove(obj_id)
def _convert_mapping_to_json_serializable(obj: Mapping[str, object]) -> dict[str, object]:
converted: Final = _convert_to_json_serializable_dict(obj)
if isinstance(converted, dict):
return converted
return dict(obj)
def _get_proxy_server_request_for_spend_logs_payload(
metadata: dict,
litellm_params: dict,
@ -1196,7 +1230,7 @@ def _get_proxy_server_request_for_spend_logs_payload(
# If redaction is enabled, convert to serializable dict before redacting
if should_redact_message_logging(model_call_details=model_call_details):
_request_body = _convert_to_json_serializable_dict(_request_body)
_request_body = _convert_mapping_to_json_serializable(_request_body)
perform_redaction(model_call_details=_request_body, result=None)
_request_body = _sanitize_request_body_for_spend_logs_payload(_request_body)
@ -1241,7 +1275,7 @@ def _get_response_for_spend_logs_payload(
if payload is None:
return "{}"
if _should_store_prompts_and_responses_in_spend_logs():
response_obj: Any = payload.get("response")
response_obj: object = payload.get("response")
if response_obj is None:
return "{}"

View file

@ -3,10 +3,11 @@ import asyncio
import json
import os
from collections import Counter
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import (
Any,
Final,
NamedTuple,
Protocol,
cast, # noqa: TID251 # prisma types Json columns as fields.Json but de-serializes them to plain python on read
)
@ -15,6 +16,7 @@ from urllib.parse import urlparse
from fastapi import APIRouter, Body, Depends, File, HTTPException, UploadFile
from pydantic import ConfigDict, JsonValue, ValidationError, create_model
from pydantic.fields import FieldInfo
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
@ -44,6 +46,31 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
router: Final = APIRouter()
JsonSchemaItems: Final = TypedDict(
"JsonSchemaItems",
{"$ref": ReadOnly[str], "enum": ReadOnly[Sequence[JsonValue]]},
total=False,
)
class JsonSchemaNode(TypedDict, total=False):
type: ReadOnly[str]
description: ReadOnly[str]
enum: ReadOnly[Sequence[JsonValue]]
anyOf: ReadOnly[Sequence["JsonSchemaNode"]]
items: ReadOnly["JsonSchemaItems"]
properties: ReadOnly[Mapping[str, "JsonSchemaNode"]]
_EMPTY_SCHEMA_DEFS: Final[Mapping[str, "JsonSchemaNode"]] = MappingProxyType({})
class JsonSchemaPropertyEntry(TypedDict):
description: ReadOnly[str]
type: ReadOnly[str]
items: NotRequired[ReadOnly["JsonSchemaItems"]]
class _SsoSettingsMappingRow(Protocol):
@property
def sso_settings(self) -> Mapping[str, object] | None: ...
@ -157,10 +184,10 @@ class UIThemeConfig(BaseModel):
class SettingsResponse(BaseModel):
"""Base response model for settings with values and schema information"""
values: dict[str, Any]
values: dict[str, object]
"""The current configuration values"""
field_schema: dict[str, Any]
field_schema: dict[str, object]
"""Schema information including descriptions and property types for UI display"""
@ -548,6 +575,62 @@ async def delete_allowed_ip(
return {"message": f"IP {ip_address.ip} deleted successfully", "status": "success"}
def _resolve_non_null_variant(field_info: JsonSchemaNode) -> JsonSchemaNode:
"""Pydantic v2 renders Optional fields as ``anyOf: [actual_type, null]``."""
if "anyOf" not in field_info:
return field_info
return next((variant for variant in field_info["anyOf"] if variant.get("type") != "null"), field_info)
def _schema_items_entry(resolved: JsonSchemaNode, defs: Mapping[str, JsonSchemaNode]) -> "JsonSchemaItems | None":
"""Items info (including enum values) for array fields, so the UI can render a multi-select dropdown."""
if "items" not in resolved:
return None
items: Final = resolved["items"]
if "$ref" not in items:
return items
ref_def: Final = defs.get(items["$ref"].split("/")[-1])
if ref_def is None or "enum" not in ref_def:
return None
enum_items: Final[JsonSchemaItems] = {"enum": ref_def["enum"]}
return enum_items
def _schema_property_entry(field_info: JsonSchemaNode, defs: Mapping[str, JsonSchemaNode]) -> JsonSchemaPropertyEntry:
resolved: Final = _resolve_non_null_variant(field_info)
items_entry: Final = _schema_items_entry(resolved, defs)
description: Final = field_info.get("description", "")
type_name: Final = resolved.get("type", "string")
if items_entry is None:
entry: Final[JsonSchemaPropertyEntry] = {"description": description, "type": type_name}
return entry
entry_with_items: Final[JsonSchemaPropertyEntry] = {
"description": description,
"type": type_name,
"items": items_entry,
}
return entry_with_items
class _RootSchema(NamedTuple):
description: str
properties: Mapping[str, JsonSchemaNode]
nested_defs: Mapping[str, JsonSchemaNode]
defs: Mapping[str, JsonSchemaNode]
def _root_schema(settings_class: type[BaseModel]) -> _RootSchema:
from pydantic import TypeAdapter
raw_schema: Final = TypeAdapter(settings_class).json_schema(by_alias=True)
return _RootSchema(
description=raw_schema.get("description", ""),
properties=raw_schema["properties"],
nested_defs=raw_schema.get("definitions", _EMPTY_SCHEMA_DEFS),
defs=raw_schema["$defs"] if "$defs" in raw_schema else raw_schema.get("definitions", _EMPTY_SCHEMA_DEFS),
)
async def _get_settings_with_schema(
settings_key: str,
settings_class: type[BaseModel],
@ -561,69 +644,43 @@ async def _get_settings_with_schema(
settings_class: The Pydantic class to use for schema
config: The config dictionary
"""
from pydantic import TypeAdapter
litellm_settings: Final = config.get("litellm_settings", {}) or {}
settings_data: Final = litellm_settings.get(settings_key, {}) or {}
# Create the settings object
settings: Final = settings_class(**(settings_data))
# Get the schema
schema: Final = TypeAdapter(settings_class).json_schema(by_alias=True)
root_schema: Final = _root_schema(settings_class)
# Convert to dict for response
settings_dict: Final = settings.model_dump()
# Add descriptions to the response
result: Final = {
"values": settings_dict,
"field_schema": {
"description": schema.get("description", ""),
"properties": {},
},
schema_properties_out: Final[Mapping[str, JsonSchemaPropertyEntry]] = {
field_name: _schema_property_entry(field_info, root_schema.defs)
for field_name, field_info in root_schema.properties.items()
}
# Add property descriptions
defs: Final = schema.get("$defs", schema.get("definitions", {}))
for field_name, field_info in schema["properties"].items():
# For Optional fields, Pydantic v2 uses anyOf with [actual_type, null].
# Resolve the non-null variant to get the real type and items.
resolved = field_info
if "anyOf" in field_info:
for variant in field_info["anyOf"]:
if variant.get("type") != "null":
resolved = variant
break
prop_entry: dict = {
"description": field_info.get("description", ""),
"type": resolved.get("type", "string"),
}
# Pass through items info (including enum values) for array fields
# so the UI can render a multi-select dropdown
if "items" in resolved:
items = resolved["items"]
# Resolve $ref to enum definitions if needed
if "$ref" in items:
ref_name = items["$ref"].split("/")[-1]
ref_def = defs.get(ref_name, {})
if "enum" in ref_def:
prop_entry["items"] = {"enum": ref_def["enum"]}
else:
prop_entry["items"] = items
result["field_schema"]["properties"][field_name] = prop_entry
# Add nested object descriptions
for def_name, def_schema in schema.get("definitions", {}).items():
result["field_schema"][def_name] = {
nested_defs_out: Final[Mapping[str, Mapping[str, object]]] = {
def_name: {
"description": def_schema.get("description", ""),
"properties": {
prop_name: {"description": prop_info.get("description", "")}
for prop_name, prop_info in def_schema.get("properties", {}).items()
},
}
for def_name, def_schema in root_schema.nested_defs.items()
}
return result
return {
"values": settings_dict,
"field_schema": {
"description": root_schema.description,
"properties": schema_properties_out,
**nested_defs_out,
},
}
@router.get(
@ -930,32 +987,29 @@ async def get_sso_settings():
resolved: Final = resolve_sso_config(sso_db_settings, os.environ)
# Get the schema for UI display
from pydantic import TypeAdapter
schema: Final = TypeAdapter(SSOConfig).json_schema(by_alias=True)
root_schema: Final = _root_schema(SSOConfig)
# Convert to dict for response, masking OAuth client secrets so plaintext
# is never sent to the UI.
sso_dict: Final = mask_sensitive_keys(resolved.config.model_dump(), set(SSO_SECRET_FIELDS))
# Add descriptions to the response
result: Final = {
"values": sso_dict,
"provenance": resolved.provenance,
"field_schema": {
"description": schema.get("description", ""),
"properties": {},
},
}
# Add property descriptions
for field_name, field_info in schema["properties"].items():
result["field_schema"]["properties"][field_name] = {
schema_properties_out: Final[Mapping[str, Mapping[str, str]]] = {
field_name: {
"description": field_info.get("description", ""),
"type": field_info.get("type", "string"),
}
for field_name, field_info in root_schema.properties.items()
}
return result
return {
"values": sso_dict,
"provenance": resolved.provenance,
"field_schema": {
"description": root_schema.description,
"properties": schema_properties_out,
},
}
@router.patch(
@ -1309,7 +1363,7 @@ UI_SETTINGS_CACHE_KEY: Final = "ui_settings:settings_dict"
UI_SETTINGS_CACHE_TTL: Final = 600 # 10 minutes
async def get_ui_settings_cached() -> dict[str, Any]:
async def get_ui_settings_cached() -> dict[str, JsonValue]:
"""
Return the persisted UI settings dict, using DualCache for reads.

View file

@ -77,6 +77,16 @@ def _is_json_array(value: object) -> TypeIs[list[object]]: # guard-ok: trivial
return isinstance(value, list)
def _optional_str(value: object) -> str | None:
"""Keep a JSON payload entry only when it is a string, since the wire format is caller-controlled."""
return value if isinstance(value, str) else None
def _json_array_or_empty(value: object) -> Sequence[object]:
"""Narrow a JSON payload entry that the caller iterates, tolerating a missing or malformed value."""
return value if _is_json_array(value) else ()
def _is_str_mapping(value: object) -> TypeIs[dict[str, str]]: # guard-ok: verifies every value is str
return _is_json_object(value) and all(isinstance(item, str) for item in value.values())
@ -96,10 +106,6 @@ class _GetsLitellmParams(Protocol):
def __call__(self, key: str, default: Mapping[str, object], /) -> LiteLLM_Params: ...
class _PopsOptionalStr(Protocol):
def __call__(self, key: str, default: None, /) -> str | None: ...
class _UnmasksPiiText(Protocol):
def __call__(self, text: str, pii_tokens: Mapping[str, str]) -> str: ...
@ -127,10 +133,6 @@ def _typed_gets_litellm_params(fn: _GetsLitellmParams) -> _GetsLitellmParams:
return fn
def _typed_pops_optional_str(fn: _PopsOptionalStr) -> _PopsOptionalStr:
return fn
_SHOULD_STORE_RESULT_IN_CACHE_ATTR: Final = "_should_store_result_in_cache"
_UNMASK_PII_TEXT_ATTR: Final = "_unmask_pii_text"
@ -342,7 +344,7 @@ class BaseResponsesAPIStreamingIterator:
ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
):
_item: Final = getattr(openai_responses_api_chunk, "item", None)
_item: Final[object] = getattr(openai_responses_api_chunk, "item", None)
if _item is not None:
ResponsesAPIRequestUtils._encode_container_id_on_output_item(
item=_item,
@ -350,7 +352,7 @@ class BaseResponsesAPIStreamingIterator:
model_id=_stream_model_id,
)
elif _event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED:
_annotation: Final = getattr(openai_responses_api_chunk, "annotation", None)
_annotation: Final[object] = getattr(openai_responses_api_chunk, "annotation", None)
if _annotation is not None:
ResponsesAPIRequestUtils._encode_container_id_on_output_item(
item=_annotation,
@ -1310,8 +1312,7 @@ def _build_synthetic_response_events(
)
if item_type == "message":
raw_content_parts = output_item_payload.get("content")
content_parts: Sequence[object] = raw_content_parts if _is_json_array(raw_content_parts) else []
content_parts: Sequence[object] = _json_array_or_empty(output_item_payload.get("content"))
for content_index, part in enumerate(content_parts):
part_payload = _dump_response_object(part)
events.append(
@ -1359,9 +1360,8 @@ def _build_synthetic_response_events(
)
)
elif item_type == "reasoning":
raw_summary_items = output_item_payload.get("summary")
summary_items: Sequence[object] = raw_summary_items if _is_json_array(raw_summary_items) else []
for summary_index, summary in enumerate(summary_items):
summaries: Sequence[object] = _json_array_or_empty(output_item_payload.get("summary"))
for summary_index, summary in enumerate(summaries):
summary_payload = _dump_response_object(summary)
summary_text = str(summary_payload.get("text") or "")
for i in range(0, len(summary_text), chunk_size):
@ -2518,14 +2518,12 @@ class ManagedResponsesWebSocketHandler:
# reuse the router-resolved self.model; passing the alias raw to
# litellm.aresponses fails in get_llm_provider. A genuinely different
# provider-prefixed per-frame model is still honored.
requested_model: Final[str | None] = _typed_pops_optional_str(call_kwargs.pop)("model", None)
requested_model: Final[str | None] = _optional_str(call_kwargs.pop("model", None))
model: Final[str] = (
self.model if requested_model is None or requested_model == self.model_group else requested_model
)
previous_response_id: Final[str | None] = _typed_pops_optional_str(call_kwargs.pop)(
"previous_response_id", None
)
previous_response_id: Final[str | None] = _optional_str(call_kwargs.pop("previous_response_id", None))
current_messages: Final = self._input_to_messages(call_kwargs.get("input"))
# Fetch history once; reused in both _apply_history and _save_turn_history

View file

@ -154,6 +154,9 @@ model_list:
# Fallback model if tier cannot be determined
default_model: gpt-4o
# Replace a routed model that cannot take image input (default: false)
modality_routing: true
```
## Usage
@ -178,6 +181,25 @@ response = litellm.completion(
## Special Behaviors
### Modality-based capability routing
The classifier reads text alone, so a request carrying an image can classify cheap and land on a
text-only model, which rejects it with a provider 400 no fallback catches. With
`modality_routing: true`, one gate inspects every decided placement: when the routed model is
explicitly declared `supports_vision: false` (deployment `model_info` first, the model cost map
otherwise; unmapped names stay routable, and a multi-deployment group must accept on every
deployment), the request is re-placed on the nearest HIGHER tier holding a capable model, with
routing plugins still applied to the re-pick, then on `default_model` (never on plugin routers
and never for a plan-floored decision), and otherwise rejected with a clear 400 naming the
router. The walk only ever goes up, so a plan-mode floor cannot be undercut; a router whose only
vision model sits below the decided tier gets the 400 and an actionable message instead.
A same-tier re-pick keeps the decision's cause and adds `modality:image` to `signals`; a tier
change or default takeover records `cause: modality_escalation` with the displaced placement
(`modality_escalated_from:<TIER>` or `modality_displaced_default_model`). Escalations are never
pinned by session affinity, and a KEPT session pin bypasses the gate entirely: a session pinned
to a text-only model keeps it even when an image arrives.
### Heuristic-first chaining
`classifier_type: heuristic_first` runs the local scorer on every request and only calls the LLM

View file

@ -30,6 +30,7 @@ from litellm.constants import EMPTY_MAPPING, RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.utils import (
@ -738,6 +739,10 @@ def _decision_is_pinnable(decision: StandardLoggingRoutingDecision | None) -> bo
size shrinks again the moment the client compacts: pinning the escalated tier would hold the
session on the big-window model long after the oversized context that forced it is gone. The
gate re-fires per request, so leaving these unpinned costs nothing but the classifier call.
A modality escalation is transient the same way: it describes what this one call carries (an
image), not what the session's traffic looks like, and pinning it would hold every following
text turn on the vision-capable model the image forced.
"""
return decision is None or (
decision.get("cause")
@ -745,6 +750,7 @@ def _decision_is_pinnable(decision: StandardLoggingRoutingDecision | None) -> bo
"default_model_fallback",
"plan_mode",
"housekeeping",
"modality_escalation",
)
and not decision.get("context_escalated")
)
@ -2274,6 +2280,175 @@ class ComplexityRouter(CustomLogger):
return pinned_model
return self.get_model_for_tier(escalated_tier)
def _model_accepts_image_input(self, model_name: str) -> bool:
"""Whether a routed model or pool entry can serve an image request.
Resolved through the deployments that would actually serve the name; a name with no
deployment on the router is served by the SDK directly and is checked against the model
cost map itself. Only an explicit supports_vision false excludes, a deployment-level
model_info override first and the map otherwise, so unmapped custom names stay routable.
A multi-deployment group must accept on EVERY deployment: the router picks a deployment
inside the group after this gate runs, so a mixed group marked eligible could still hand
the image to its text-only member and fail with the exact 400 the gate exists to prevent.
"""
from litellm.utils import is_vision_explicitly_disabled
def deployment_accepts(deployment: Mapping[str, Any]) -> bool:
declared: Final = (deployment.get("model_info") or EMPTY_MAPPING).get("supports_vision")
if declared is not None:
return declared is True
litellm_model: Final = (deployment.get("litellm_params") or EMPTY_MAPPING).get("model") or model_name
return not is_vision_explicitly_disabled(litellm_model)
deployments: Final = self.litellm_router_instance.get_model_list(model_name=model_name)
if not deployments:
return not is_vision_explicitly_disabled(model_name)
return all(deployment_accepts(deployment) for deployment in deployments)
def _modality_eligible_models(self) -> frozenset[str]:
"""Every configured pool entry, plus default_model, that can serve an image request."""
names: Final = frozenset(entry for pool in self._tier_pools().values() for entry in pool) | frozenset(
name for name in (self.config.default_model,) if name
)
return frozenset(name for name in names if self._model_accepts_image_input(name))
async def _gate_response_modality(
self,
response: PreRoutingHookResponse,
messages: list[dict[str, Any]] | None, # mutable-ok: forwarded verbatim to the list-typed re-pick
resolved_messages: Sequence[Mapping[str, object]] | None,
request_kwargs: dict, # mutable-ok: same shape the hook receives
) -> PreRoutingHookResponse:
"""Replace a routed model that cannot accept this request's image input.
The single modality owner, applied to the decided response at the hook's exits so every
routing path is covered uniformly. A KEPT session pin is exempt by design (its cause);
replacement picks and every other path are just responses. The re-placement walks
UPWARD-ONLY from the decision's tier (so a plan-mode floor can never be undercut), picks
through `_pick_model_for_tier` so routing plugins still apply, then falls to
default_model (never on plugin routers, and never on a plan-floored decision, since
default_model carries no tier guarantee), else raises the clear 400. The rewritten
decision keeps its cause on a same-tier repick and becomes modality_escalation when the
tier moved or default_model took over, with the displaced placement in signals.
"""
decision: Final = response.routing_decision
if (
not self.config.modality_routing
or not resolved_messages
or response.model is None
or (decision is not None and decision.get("cause") == "session_affinity_pin")
or not request_contains_image_content(resolved_messages)
or self._model_accepts_image_input(response.model)
):
return response
eligible: Final = self._modality_eligible_models()
names: Final = self.config.tier_names()
pools: Final = self._tier_pools()
decided: Final = decision.get("tier") if decision is not None else None
start: Final = names.index(decided) if isinstance(decided, str) and decided in names else 0
capable: Final = next(
(name for name in names[start:] if any(entry in eligible for entry in pools.get(name, ()))), None
)
if capable is not None:
new_tier: ComplexityTier | str | None = capable if self.config.has_custom_tiers else ComplexityTier(capable)
repick_messages: Final = list(resolved_messages) # mutable-ok: the pick's param is list-typed
new_model = await self._pick_model_for_tier(
new_tier,
messages,
repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them
request_kwargs,
allowed_models=tuple(entry for entry in pools.get(capable, ()) if entry in eligible),
)
elif self._modality_default_model_usable(request_kwargs, resolved_messages, eligible):
new_tier = None
new_model = self._placed_default_model()
else:
import litellm
raise litellm.BadRequestError(
message=(
f"Auto-router {self.model_name} received a request with image input, but no model "
f"at or above the decided tier accepts images and modality_routing is enabled. "
f"Tiers checked: {', '.join(names[start:])}. Add a vision-capable model to a tier, "
f"or set a vision-capable default_model, or remove the image content."
),
model=self.model_name,
llm_provider="",
)
self._restamp_adaptive_choice(request_kwargs, response.model, new_model)
same_tier: Final = capable is not None and decided == capable
base_cause: Final = (decision.get("cause") if decision is not None else None) or "default_fallback"
displaced_default: Final = decided is None and response.model == self.config.default_model
markers: Final = (
"modality:image",
*((f"modality_escalated_from:{decided}",) if not same_tier and isinstance(decided, str) else ()),
*(("modality_displaced_default_model",) if not same_tier and displaced_default else ()),
)
old_signals: Final = tuple(decision.get("signals") or ()) if decision is not None else ()
new_decision: Final = self._build_routing_decision(
routed_model=new_model,
cause=base_cause if same_tier else "modality_escalation",
tier=new_tier,
score=decision.get("score") if decision is not None else None,
signals=(*old_signals, *markers),
matched_keyword=decision.get("matched_keyword") if decision is not None else None,
escalation_keyword=decision.get("escalation_keyword") if decision is not None else None,
escalated=bool(decision.get("escalated", False)) if decision is not None else False,
classifier_model=decision.get("classifier_model") if decision is not None else None,
classifier_cost=decision.get("classifier_cost") if decision is not None else None,
conversation_continuing=bool(decision.get("conversation_continuing", True))
if decision is not None
else True,
tier_litellm_params=self._litellm_params_for_model(new_tier, new_model),
context_escalation_original_tier=(
decision.get("context_escalation_original_tier") if decision is not None else None
),
)
from litellm.types.router import PreRoutingHookResponse as HookResponse
return HookResponse(
model=new_model,
messages=response.messages,
litellm_params=self._litellm_params_for_model(new_tier, new_model),
routing_decision=new_decision,
)
def _modality_default_model_usable(
self,
request_kwargs: Mapping[str, object],
resolved_messages: Sequence[Mapping[str, object]] | None,
eligible: frozenset[str],
) -> bool:
"""default_model may serve a gated request only when it is configured, plugin-free
(it is never checked against the plugin pipeline), capability-eligible, and the turn
carries no plan-mode sentinel. The sentinel is re-detected here rather than read off
the decision record, because the record only marks turns the floor RAISED; a sentinel
turn already at or above the floor keeps its ordinary cause, and default_model carries
no tier the floor could vouch for on any sentinel turn."""
return (
bool(self.config.default_model)
and not self.config.plugins
and self.config.default_model in eligible
and self._matched_plan_mode_signal(request_kwargs, resolved_messages) is None
)
def _placed_default_model(self) -> str:
"""The default_model behind a usable-default verdict; the raise is the type-level
proof, not a reachable path."""
model: Final = self.config.default_model
if model is None:
raise ValueError(f"Auto-router {self.model_name}: modality gate routed to an unset default_model")
return model
@staticmethod
def _restamp_adaptive_choice(request_kwargs: Mapping[str, object], old_model: str, new_model: str) -> None:
"""The adaptive feedback loop reads its chosen-model marker from request metadata; a
gate rewrite must move the marker with the model or rewards land on the displaced one."""
metadata: Final = request_kwargs.get("metadata")
if isinstance(metadata, dict) and metadata.get("adaptive_router_chosen_model") == old_model:
metadata["adaptive_router_chosen_model"] = new_model
def _lexical_tier_override(self, user_message: str) -> KeywordOverride | None:
"""When keyword_tier_rules match literally, the most-severe matched tier wins.
@ -2655,25 +2830,30 @@ class ComplexityRouter(CustomLogger):
session_tier_litellm_params: Final = self._litellm_params_for_model(routed_pin_tier, routed_model)
has_original_messages: Final = messages is not None and len(messages) > 0
return self._with_session_deployment_affinity(
PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=session_tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause=cause,
tier=routed_pin_tier,
matched_keyword=pin_plan_sentinel if plan_floored else None,
escalation_keyword=pin_escalation_keyword,
escalated=escalated,
conversation_continuing=conversation_continuing,
tier_litellm_params=session_tier_litellm_params,
context_escalation_original_tier=pin_context_original_tier,
await self._gate_response_modality(
PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=session_tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause=cause,
tier=routed_pin_tier,
matched_keyword=pin_plan_sentinel if plan_floored else None,
escalation_keyword=pin_escalation_keyword,
escalated=escalated,
conversation_continuing=conversation_continuing,
tier_litellm_params=session_tier_litellm_params,
context_escalation_original_tier=pin_context_original_tier,
),
),
messages,
resolved_messages,
request_kwargs,
)
)
response: Final = await self._classify_and_route(
routed_response: Final = await self._classify_and_route(
model=model,
request_kwargs=request_kwargs,
messages=messages,
@ -2682,6 +2862,11 @@ class ComplexityRouter(CustomLogger):
conversation_continuing=conversation_continuing,
resolved_messages=resolved_messages,
)
response: Final = (
await self._gate_response_modality(routed_response, messages, resolved_messages, request_kwargs)
if routed_response is not None
else None
)
# Sentinel presence, not the plan_mode cause, gates the pin write: a plan-mode turn
# classified at or above the floor keeps its ordinary cause, yet on an adaptive router
# the hard floor constrained its pick, so pinning it would carry a plan-mode-shaped

View file

@ -848,6 +848,18 @@ class ComplexityRouterConfig(BaseModel):
"drift plus the response tokens."
),
)
modality_routing: bool = Field(
default=False,
description=(
"Route image-bearing requests only to models that can accept image input. The "
"classifier reads text alone, so an image request whose text classifies cheap "
"otherwise lands on a text-only model and fails with a provider 400. When enabled, "
"a routed model explicitly declared supports_vision false (deployment model_info "
"or the model cost map; unmapped names stay routable) is replaced by the nearest "
"HIGHER tier holding a capable model, then default_model, else a clear 400. A kept "
"session-affinity pin still wins even when an image arrives."
),
)
# Semantic (embedding) matching for keyword_tier_rules instead of literal text matching
semantic_keyword_matching: bool = Field(

View file

@ -0,0 +1,24 @@
"""Typed view of the scalar keyword payload guardrails forward to ``CustomGuardrail.__init__``.
Guardrail subclasses collect their base-class options in ``**kwargs`` and splat them into
``super().__init__``. Declaring the payload's shape here lets the checker resolve each
forwarded argument to its real parameter type instead of ``Any``.
"""
from typing_extensions import ReadOnly, TypedDict
class GuardrailBaseInitKwargs(TypedDict, total=False):
guardrail_name: ReadOnly[str | None]
default_on: ReadOnly[bool]
mask_request_content: ReadOnly[bool]
mask_response_content: ReadOnly[bool]
violation_message_template: ReadOnly[str | None]
end_session_after_n_fails: ReadOnly[int | None]
on_violation: ReadOnly[str | None]
realtime_violation_message: ReadOnly[str | None]
on_sensitive_data: ReadOnly[str | None]
sensitive_data_route_to_model: ReadOnly[str | None]
sticky_session_routing: ReadOnly[bool]
run_in_parallel: ReadOnly[bool]
only_scan_new_messages: ReadOnly[bool]

View file

@ -251,8 +251,12 @@ DEFAULT_SHADOW_EVAL_JUDGE_MODEL: Final[str] = "anthropic/claude-sonnet-5"
# Sample-count ceiling written on every new job: a zero-cost error loop (a shadow arm that
# fails before billing) never consumes spend budget, so it must terminate on count instead.
# A multi-router job writes one attempt row per router arm, so the valve is reached
# proportionally sooner; it is a safety valve, not a sample budget.
SHADOW_EVAL_TURN_VALVE: Final[int] = 10_000
SHADOW_EVAL_MAX_ROUTERS: Final[int] = 4
class StartShadowEvalRequest(BaseModel):
"""Start duplicating one or more targets' traffic for blind comparison against an auto-router.
@ -288,7 +292,24 @@ class StartShadowEvalRequest(BaseModel):
"to across all their teams: JWT requests carrying their subject claim and virtual keys they own"
),
)
router_name: str = Field(description="The auto-router under evaluation, in either direction")
router_name: str | None = Field(
default=None,
description=(
"The auto-router under evaluation, in either direction: the single-router spelling of "
"router_names. Provide exactly one of the two fields"
),
)
router_names: tuple[str, ...] = Field(
default=(),
max_length=SHADOW_EVAL_MAX_ROUTERS,
description=(
"The auto-routers under evaluation, at most "
f"{SHADOW_EVAL_MAX_ROUTERS}. Every sampled request runs through every router listed and each "
"arm is judged independently against the same real response, so routers compare head-to-head "
"on identical traffic. More than one router requires direction 'forward'. After validation "
"this field always carries the full deduplicated set, whichever spelling the caller used"
),
)
direction: ShadowEvalDirection = Field(
default="forward",
description=(
@ -332,7 +353,8 @@ class StartShadowEvalRequest(BaseModel):
"Per-target USD budget for the eval's own overhead, the shadow-arm and judge calls, priced with "
"the same figures the spend pipeline bills. EACH scoped target samples until its recorded eval "
"spend reaches this, so a job over N targets spends at most about N times max_budget; in-flight "
"samples can overshoot the cap by one sampling cache window"
"samples can overshoot the cap by one sampling cache window. Every router arm draws from the "
"same per-target budget, so a multi-router job reaches it proportionally sooner"
),
)
@ -373,6 +395,23 @@ class StartShadowEvalRequest(BaseModel):
raise ValueError("baseline_model is only meaningful when direction is 'reverse'")
return self
@model_validator(mode="after")
def _resolve_router_set(self) -> "StartShadowEvalRequest":
"""Whichever spelling the caller used, router_names leaves validation as the full
deduplicated set, so every downstream reader consumes one field."""
if (self.router_name is None) == (not self.router_names):
raise ValueError("provide exactly one of router_name or router_names")
single: Final = () if self.router_name is None else (self.router_name,)
routers: Final = tuple(dict.fromkeys(self.router_names or single))
if not all(name.strip() for name in routers):
raise ValueError("router names must be non-empty strings")
if len(routers) > 1 and self.direction == "reverse":
raise ValueError("a reverse job evaluates one router against baseline_model; pass a single router")
# A returned model_copy is ignored on the __init__ construction path, so the
# normalization must land as a self attribute store to hold for every caller.
self.router_names = routers
return self
class ShadowEvalSlice(BaseModel):
"""Judge outcomes for one slice of a job's verdicts: a router tier, one of the
@ -428,15 +467,28 @@ class ShadowEvalResult(BaseModel):
"and in reverse the models the router itself picked"
)
)
by_router: tuple[ShadowEvalSlice, ...] = Field(
default=(),
description=(
"One slice per router arm, grouped on the router name. Every arm of a multi-router job is "
"judged against the same real responses over the same sampled requests, so these slices "
"compare routers head-to-head: like-for-like win rates and spends on identical traffic. "
"Verdicts from before arm stamping existed count toward the job's own router"
),
)
overall_shadow_win_rate_pct: float
overall_tie_rate_pct: float
sampled_real_spend: float = Field(
default=0.0,
description="USD the real arm billed across all judged turns, cache-served turns excluded",
description=(
"USD the real arm billed across all judged turns, cache-served turns excluded. A judged turn "
"is one (request, router arm) verdict, so a multi-router job counts the real response once per "
"arm it was judged against; per-router comparisons read by_router"
),
)
sampled_shadow_spend: float = Field(
default=0.0,
description="USD the shadow arm billed across the same turns, judge excluded, like for like",
description="USD the shadow arms billed across the same turns, judge excluded, like for like",
)
not_sampled_count: int | None = Field(
default=None,
@ -540,7 +592,13 @@ class ShadowEvalJobResponse(BaseModel):
min_length=1,
description="The targets whose traffic this job evaluates, and only theirs, each with its own budget",
)
router_name: str
router_names: tuple[str, ...] = Field(
min_length=1,
description=(
"Every auto-router this job runs as a shadow arm. Multi-router jobs sample one slice of "
"traffic and judge every arm against the same real responses"
),
)
direction: ShadowEvalDirection = "forward"
baseline_model: str | None = None
judge_model: str
@ -562,6 +620,13 @@ class ShadowEvalJobResponse(BaseModel):
last_error: str | None = Field(default=None, description="Most recent attempt error; detail endpoint only")
results: ShadowEvalResult | None = Field(default=None, description="Stratified verdicts; detail endpoint only")
@computed_field
@property
def router_name(self) -> str:
"""The first router, kept for callers that predate router_names; derived so the
two fields can never disagree."""
return self.router_names[0]
@computed_field
@property
def status(self) -> ShadowEvalStatus:

View file

@ -2840,6 +2840,10 @@ RoutingDecisionCause = Literal[
# never called. The matched sentinel rides in matched_keyword. Distinct from the keyword causes,
# which are operator-authored rules; these sentinels ship with the router.
"housekeeping",
# modality_routing replaced the decided placement: the request carries an image and the
# routed model does not accept image input, so the nearest higher capable tier or
# default_model served instead. The displaced placement rides in signals.
"modality_escalation",
"session_affinity_pin",
"session_affinity_escalation",
# classification_mode 'user_turn': the request is an agent loop's continuation turn (no new

View file

@ -2660,10 +2660,19 @@ def _is_explicitly_disabled_factory(model: str, custom_llm_provider: str | None,
``_supports_factory`` so caching, fallback, and normalisation improvements
apply here automatically.
"""
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
try:
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model, custom_llm_provider=custom_llm_provider
)
declared: Final = declared_authenticating_provider(model, custom_llm_provider)
if declared is not None:
model = model.removeprefix(
f"{declared}/"
) # rebind-ok: mirrors get_llm_provider's split without its OAuth flow
custom_llm_provider = declared # rebind-ok: same
else:
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model, custom_llm_provider=custom_llm_provider
)
model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)
val: Final = model_info.get(key)
if val is False:
@ -2751,6 +2760,15 @@ def supports_computer_use(model: str, custom_llm_provider: str | None = None) ->
)
def is_vision_explicitly_disabled(model: str, custom_llm_provider: str | None = None) -> bool:
"""True only when supports_vision is explicitly declared false for the model.
The opt-out mirror of :func:`supports_vision`: a missing declaration reads as not
disabled, so unknown or newly added models stay eligible for image routing.
"""
return _is_explicitly_disabled_factory(model, custom_llm_provider, "supports_vision")
def supports_vision(model: str, custom_llm_provider: str | None = None) -> bool:
"""
Check if the given model supports vision and return a boolean value.
@ -8260,10 +8278,17 @@ class ProviderConfigManager:
"""
# Handle OpenAI special cases (O-series and GPT-5 models)
if provider == LlmProviders.OPENAI:
from litellm.llms.openai.chat.gpt_transformation import (
OpenAIGPTConfig,
OpenAIUnknownModelConfig,
)
if litellm.openaiOSeriesConfig.is_model_o_series_model(model=model):
return litellm.openaiOSeriesConfig
if litellm.OpenAIGPT5Config.is_model_gpt_5_model(model=model):
return litellm.OpenAIGPT5Config()
if not OpenAIGPTConfig.is_openai_catalog_model(model):
return OpenAIUnknownModelConfig()
# Handle Azure before the generic map so base_model can be threaded through
if provider == LlmProviders.AZURE:
@ -9415,6 +9440,10 @@ class ProviderConfigManager:
return RunwayMLTextToSpeechConfig()
elif litellm.LlmProviders.VERTEX_AI == provider:
if "gemini" in model:
# Gemini TTS uses the speech_to_completion bridge, and Google Cloud TTS param
# mapping would drop response_format before the bridge sees it (LIT-6501)
return None
from litellm.llms.vertex_ai.text_to_speech.transformation import (
VertexAITextToSpeechConfig,
)

View file

@ -1,6 +1,6 @@
{
"ANN001": {
"limit": 2995
"limit": 2991
},
"ANN002": {
"limit": 71
@ -12,10 +12,10 @@
"limit": 2002
},
"ANN202": {
"limit": 845
"limit": 841
},
"ANN204": {
"limit": 698
"limit": 694
},
"ANN205": {
"limit": 112
@ -24,7 +24,7 @@
"limit": 133
},
"ANN401": {
"limit": 587
"limit": 387
},
"ASYNC230": {
"limit": 11
@ -117,7 +117,7 @@
"limit": 1
},
"PERF102": {
"limit": 23
"limit": 21
},
"PERF401": {
"limit": 12
@ -168,7 +168,7 @@
"limit": 3
},
"RET504": {
"limit": 175
"limit": 173
},
"RUF012": {
"limit": 239
@ -198,7 +198,7 @@
"limit": 56
},
"SIM102": {
"limit": 314
"limit": 310
},
"SIM103": {
"limit": 119
@ -231,7 +231,7 @@
"limit": 5
},
"TID251": {
"limit": 1108
"limit": 1084
},
"TRY002": {
"limit": 524

View file

@ -1531,7 +1531,8 @@ model LiteLLM_ShadowEvalJob {
group_id String // legs of one job share this; the API's job id
target_type String @default("key") // key | team | user
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
router_name String // the auto-router under evaluation, in either direction
router_name String // first (often only) auto-router under evaluation; router_names is the full set
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String
@ -1555,6 +1556,7 @@ model LiteLLM_ShadowEvalAttempt {
job_id String
request_id String // the judged real request
outcome String // real | shadow | tie | error
router_name String? // the arm this verdict scores; NULL on legacy rows, meaning the job's own router
tier String? // router's tier for the prompt, when classified
real_model String?
shadow_model String?

View file

@ -64,14 +64,12 @@
- {id: mgmt.budget.list_v1.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "management_v1/budgets.py:129", rationale: "Budget enumeration the Budgets page can page, sort and filter"}
- {id: mgmt.budget.list_v1.admin_only, module: mgmt, tier: P1, surface: api, assertions: [admin_only], source: "management_v1/budgets.py:129", rationale: "A caller without admin view is refused, not served an empty page"}
- {id: mgmt.callback.list.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "callback_management_endpoints.py", rationale: "Callback config (smoke)"}
- {id: mgmt.cache_settings.update.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "cache_settings_endpoints.py", rationale: "Cache config (smoke). Deliberately uncovered: the previous test read the live settings and wrote them back, which proves nothing (identical values in, so a no-op POST still passes) while being able to break the deployment. /cache/settings persists what it receives and that row outranks YAML cache_params, re-applied on a timer, so a write that omits ssl or redis_startup_nodes turns a TLS cluster into a plaintext standalone node and every later Redis call hangs. That took out 60 of 72 tests on 2026-07-25. GET cannot round-trip it either: it resolves the stored row overlaid with REDIS_* env and never reads YAML, so on a fresh deploy it cannot see YAML ssl to echo back. A safe test needs an isolated proxy, or LIT-4816 fixed so a partial write cannot downgrade transport. Do not re-add a read-then-write-back test against a shared proxy."}
- {id: mgmt.cost_tracking.estimate.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "cost_tracking_settings.py", rationale: "Cost estimate (smoke)"}
- {id: mgmt.router_settings.update.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "router_settings_endpoints.py", rationale: "Router config (smoke)"}
- {id: mgmt.jwt_key_mapping.new.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "jwt_key_mapping_endpoints.py", rationale: "JWT->key mapping (smoke)"}
- {id: mgmt.compliance.gdpr.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "compliance_endpoints.py", rationale: "GDPR ops (smoke)"}
- {id: mgmt.tool_management.list.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "tool_management_endpoints.py", rationale: "Tool inventory (smoke)"}
- {id: mgmt.fallback_management.update.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "fallback_management_endpoints.py", rationale: "Fallback config (smoke)"}
- {id: mgmt.config_override.hashicorp_vault.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "config_override_endpoints.py", rationale: "Vault integration (smoke)"}
- {id: mgmt.workflow.list.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "workflow_management_endpoints.py", rationale: "Workflow tracking (smoke)"}
- {id: mgmt.credential_migration.check.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:4252", rationale: "Encryption migration (smoke)"}
- {id: mgmt.credential.new.serves_request, module: mgmt, tier: P1, surface: api, assertions: [serves_request], source: "credential_endpoints/endpoints.py:42", rationale: "Stored credential resolves into a deployment and serves a live /messages request"}

View file

@ -7,9 +7,13 @@ so a read-back reflects the change. Router settings, which mutate global proxy
state, are exercised with a benign, self-restoring change so a shared proxy is left
as it was found.
Cache settings are deliberately not covered here; see the rationale on
mgmt.cache_settings.update.happy_path in coverage_registry/mgmt.yaml before adding
a test for that route.
Cache settings and the Vault config override are deliberately not covered here.
Both routes reconfigure the whole proxy: /cache/settings persists what it receives
into a row that outranks the YAML cache_params and is re-applied on a timer, and
/config_overrides/hashicorp_vault swaps the process-wide secret manager. Neither can
be exercised safely against the shared proxy the suites run on, so they need an
isolated proxy before a test lands. Do not add a read-then-write-back test for
either one.
"""
from __future__ import annotations

View file

@ -29,6 +29,7 @@ export const E2E_PROXY_ADMIN_USER_ID = "e2e-proxy-admin";
export const E2E_PROXY_ADMIN_EMAIL = "admin@test.local";
export const E2E_INTERNAL_USER_ID = "e2e-internal-user";
export const E2E_INTERNAL_USER_EMAIL = "internal@test.local";
export const E2E_TEAM_ADMIN_USER_ID = "e2e-team-admin";
// Key aliases for seeded test keys (match seed.sql)
export const E2E_UPDATE_LIMITS_KEY_ALIAS = "e2eUpdateLimitsKey";
@ -46,3 +47,5 @@ export const E2E_TEAM_ORG_ID = "e2e-team-org";
export const E2E_TEAM_ORG_ALIAS = "E2E Team In Org";
export const E2E_TEAM_NO_ADMIN_ID = "e2e-team-no-admin";
export const E2E_TEAM_NO_ADMIN_ALIAS = "E2E Team No Admin";
export const E2E_TEAM_KEYGEN_ID = "e2e-team-keygen";
export const E2E_TEAM_KEYGEN_ALIAS = "E2E Team Keygen";

View file

@ -29,7 +29,7 @@ INSERT INTO "LiteLLM_UserTable" ("user_id", "user_email", "user_role", "teams",
VALUES
('e2e-proxy-admin', 'admin@test.local', 'proxy_admin', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-admin-viewer', 'adminviewer@test.local', 'proxy_admin_viewer', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-internal-user', 'internal@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-org"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-internal-user', 'internal@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-org","e2e-team-keygen"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-internal-viewer', 'viewer@test.local', 'internal_user_viewer', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-team-admin', 'teamadmin@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-delete"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-invitable-user', 'invitable@test.local', 'internal_user', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
@ -63,6 +63,17 @@ INSERT INTO "LiteLLM_TeamTable" (
'[{"role":"user","user_id":"e2e-invitable-user"}]'::jsonb,
'{}'::jsonb, '{"fake-openai-gpt-4"}', 0.0, '{}'::jsonb, '{}'::jsonb, false);
INSERT INTO "LiteLLM_TeamTable" (
"team_id", "team_alias", "organization_id", "admins", "members",
"members_with_roles", "metadata", "models", "spend", "model_spend", "model_max_budget", "blocked",
"team_member_permissions"
) VALUES
('e2e-team-keygen', 'E2E Team Keygen', NULL,
'{}', '{"e2e-internal-user"}',
'[{"role":"user","user_id":"e2e-internal-user"}]'::jsonb,
'{}'::jsonb, '{"fake-openai-gpt-4"}', 0.0, '{}'::jsonb, '{}'::jsonb, false,
'{"/key/generate"}');
-- 6. Team Memberships (only user_id, team_id, spend — no created_at/updated_at)
INSERT INTO "LiteLLM_TeamMembership" ("user_id", "team_id", "spend")
VALUES
@ -72,6 +83,7 @@ VALUES
('e2e-removable-member', 'e2e-team-crud', 0.0),
('e2e-team-admin', 'e2e-team-delete', 0.0),
('e2e-internal-user', 'e2e-team-org', 0.0),
('e2e-internal-user', 'e2e-team-keygen', 0.0),
('e2e-invitable-user', 'e2e-team-no-admin', 0.0);
-- 7. Verification Tokens (API Keys)

View file

@ -84,6 +84,34 @@ export async function waitForSpendLog(
throw new Error(`spend log for request ${requestId} never appeared (last /spend/logs status ${lastStatus})`);
}
export async function waitForSpendLogByPrompt(
request: APIRequestContext,
prompt: string,
timeoutMs = 60_000,
): Promise<string> {
const deadline = Date.now() + timeoutMs;
let lastStatus = 0;
while (Date.now() < deadline) {
const res = await request.get(`${rootPath()}/spend/logs`, {
headers: { Authorization: `Bearer ${masterKey()}` },
});
lastStatus = res.status();
if (res.ok()) {
const rows: { request_id?: string; messages?: unknown; proxy_server_request?: unknown }[] = await res.json();
const row = (Array.isArray(rows) ? rows : []).find(
(candidate) =>
JSON.stringify(candidate.messages ?? "").includes(prompt) ||
JSON.stringify(candidate.proxy_server_request ?? "").includes(prompt),
);
if (row?.request_id) {
return row.request_id;
}
}
await new Promise((r) => setTimeout(r, 2_000));
}
throw new Error(`no spend log row carrying prompt ${prompt} appeared (last /spend/logs status ${lastStatus})`);
}
const isoDay = (d: Date): string => d.toISOString().slice(0, 10);
/**

View file

@ -0,0 +1,81 @@
import { test, expect } from "@playwright/test";
import { ADMIN_STORAGE_PATH, E2E_TEAM_NO_ADMIN_ID } from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage, dismissFeedbackPopup, clickTeamId } from "../../helpers/navigation";
test.describe("Guardrails", () => {
test.use({ storageState: ADMIN_STORAGE_PATH });
test("Create a Presidio guardrail, see it in team settings, and delete it", async ({ page }) => {
const guardrailName = `e2e-presidio-${Date.now()}`;
await navigateToPage(page, Page.Guardrails);
await dismissFeedbackPopup(page);
await page.getByRole("button", { name: /Add New Guardrail/i }).click();
await page.getByRole("menuitem", { name: "Add Provider Guardrail" }).click();
const dialog = page.getByRole("dialog", { name: "Create guardrail" });
await expect(dialog).toBeVisible({ timeout: 10_000 });
await dialog.getByLabel("Guardrail Name").fill(guardrailName);
const providerSelect = dialog.getByRole("combobox", { name: "Guardrail Provider" });
await providerSelect.click();
await providerSelect.fill("Presidio");
await page.getByRole("option", { name: "Presidio PII" }).click();
await dialog.getByLabel("Mode", { exact: true }).click();
await page.keyboard.type("pre_call");
await expect(page.getByRole("option", { name: "pre_call" })).toBeAttached({ timeout: 5_000 });
await page.keyboard.press("Enter");
await expect(dialog.getByText("pre_call", { exact: true })).toBeVisible({ timeout: 5_000 });
await dialog.getByText("Create guardrail", { exact: true }).click();
await dialog.getByLabel("presidio_analyzer_api_base").fill("http://127.0.0.1:9999");
await expect(dialog.getByLabel("presidio_analyzer_api_base")).toHaveValue("http://127.0.0.1:9999");
await dialog.getByLabel("presidio_anonymizer_api_base").fill("http://127.0.0.1:9999");
await expect(dialog.getByLabel("presidio_anonymizer_api_base")).toHaveValue("http://127.0.0.1:9999");
await dialog.getByRole("button", { name: "Next" }).click();
await expect(dialog.getByText("Configure PII Protection")).toBeVisible({ timeout: 10_000 });
await dialog.getByRole("button", { name: "Select All & Mask" }).click();
await dialog.getByRole("button", { name: "Create Guardrail" }).click();
await expect(page.getByText("Guardrail created successfully").first()).toBeVisible({ timeout: 15_000 });
const row = page.getByRole("row").filter({ hasText: guardrailName });
await expect(row).toHaveCount(1, { timeout: 15_000 });
await navigateToPage(page, Page.Teams);
await dismissFeedbackPopup(page);
await clickTeamId(page, E2E_TEAM_NO_ADMIN_ID);
await page.getByRole("tab", { name: "Settings" }).click();
await page.getByRole("button", { name: "Edit Settings" }).click();
const guardrailsSelect = page.getByRole("combobox", { name: "Select guardrails" });
await expect(guardrailsSelect).toBeVisible({ timeout: 10_000 });
await guardrailsSelect.click();
await guardrailsSelect.fill(guardrailName);
await expect(page.getByRole("option", { name: guardrailName })).toBeVisible({ timeout: 10_000 });
await page.keyboard.press("Escape");
await navigateToPage(page, Page.Guardrails);
await expect(row).toHaveCount(1, { timeout: 15_000 });
await row.getByRole("button", { name: "Open guardrail actions" }).click();
await page.getByRole("menuitem", { name: "Delete" }).click();
const deleteModal = page.getByRole("dialog", { name: "Delete Guardrail" });
await expect(deleteModal).toBeVisible({ timeout: 5_000 });
await deleteModal.getByRole("button", { name: "Delete", exact: true }).click();
await expect(page.getByText(`Guardrail "${guardrailName}" deleted successfully`)).toBeVisible({
timeout: 10_000,
});
await expect(row).toHaveCount(0, { timeout: 15_000 });
await page.reload();
await expect(page.getByRole("button", { name: /Add New Guardrail/i })).toBeVisible({ timeout: 20_000 });
await expect(page.getByRole("row").filter({ hasText: guardrailName })).toHaveCount(0);
});
});

View file

@ -3,10 +3,13 @@ import {
E2E_INTERNAL_USER_KEY_ALIAS,
E2E_TEAM_CRUD_ALIAS,
E2E_TEAM_CRUD_ID,
E2E_TEAM_KEYGEN_ALIAS,
INTERNAL_USER_STORAGE_PATH,
} from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage, clickTeamId } from "../../helpers/navigation";
import { CHAT_MODEL_A, MOCK_RESPONSE_TEXT, masterKey } from "../../helpers/traffic";
import { keySourceSelect, onlyVisible, openPlayground, selectModel, sendMessage } from "../../helpers/playground";
test.describe("Internal User", () => {
test.use({ storageState: INTERNAL_USER_STORAGE_PATH });
@ -37,6 +40,55 @@ test.describe("Internal User", () => {
await expect(page.getByRole("tab", { name: "Members" })).not.toBeVisible();
});
test("Internal user creates a team key and uses it in the Playground", async ({ page, request }) => {
const suffix = Date.now();
const auth = { Authorization: `Bearer ${masterKey()}` };
await navigateToPage(page, Page.ApiKeys);
await page.getByRole("button", { name: /Create New Key/i }).click();
await expect(page.getByText("Key Ownership")).toBeVisible({ timeout: 10_000 });
await expect(page.getByRole("radio", { name: "You", exact: true })).toBeVisible({ timeout: 10_000 });
await expect(page.getByRole("radio", { name: "Another User" })).toHaveCount(0);
const keyName = `e2e-internal-team-key-${suffix}`;
await page.getByLabel(/Key Name/).fill(keyName);
const teamSelect = page.getByTestId("team-dropdown").getByRole("combobox");
await teamSelect.click();
await page.keyboard.type(E2E_TEAM_KEYGEN_ALIAS);
await page.getByRole("option", { name: E2E_TEAM_KEYGEN_ALIAS }).first().click();
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "All Team Models", exact: true }).click();
await page.keyboard.press("Escape");
await page.getByRole("button", { name: "Create Key", exact: true }).click();
await expect(page.getByText("Save your Key")).toBeVisible({ timeout: 10_000 });
const apiKey = (await page.getByRole("dialog", { name: "Save your Key" }).locator("pre").innerText()).trim();
expect(apiKey).toMatch(/^sk-/);
await page.keyboard.press("Escape");
try {
await openPlayground(page);
await keySourceSelect(page).click();
await onlyVisible(page.getByRole("option", { name: "Virtual Key" })).click({ timeout: 15_000 });
const keyInput = onlyVisible(page.getByPlaceholder("Enter custom Virtual Key"));
await expect(keyInput).toBeVisible({ timeout: 10_000 });
await keyInput.fill(apiKey);
await selectModel(page, CHAT_MODEL_A);
await sendMessage(page, `internal user team key ping ${keyName}`);
await expect(page.getByText(MOCK_RESPONSE_TEXT, { exact: false }).first()).toBeVisible({ timeout: 60_000 });
} finally {
await request.post("/key/delete", { headers: auth, data: { keys: [apiKey] } });
}
});
test("Virtual Keys page does not surface litellm-dashboard team keys", async ({ page }) => {
await navigateToPage(page, Page.ApiKeys);

View file

@ -1,14 +1,13 @@
import { test, expect } from "@playwright/test";
import { INTERNAL_USER_STORAGE_PATH, E2E_TEAM_CRUD_ALIAS, E2E_TEAM_ORG_ALIAS } from "../../constants";
import {
INTERNAL_USER_STORAGE_PATH,
E2E_TEAM_CRUD_ALIAS,
E2E_TEAM_KEYGEN_ALIAS,
E2E_TEAM_ORG_ALIAS,
} from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage } from "../../helpers/navigation";
/**
* Differential partner to internalUserNoTeam.spec.ts: the seeded
* e2e-internal-user belongs to exactly two teams, so the Create Key dropdown
* must list both. Without this, the no-team spec's "zero options" assertion
* would still pass against a bug that empties the dropdown for everyone.
*/
test.describe("Internal User with team memberships", () => {
test.use({ storageState: INTERNAL_USER_STORAGE_PATH });
@ -21,10 +20,9 @@ test.describe("Internal User with team memberships", () => {
const teamSelect = page.getByTestId("team-dropdown").getByRole("combobox");
await teamSelect.click();
// Both seeded memberships render, and nothing else does — proving the
// dropdown is scoped to the user's teams rather than empty or unfiltered.
await expect(page.getByRole("option", { name: E2E_TEAM_CRUD_ALIAS })).toBeVisible({ timeout: 10_000 });
await expect(page.getByRole("option", { name: E2E_TEAM_ORG_ALIAS })).toBeVisible();
await expect(page.getByRole("option")).toHaveCount(2);
await expect(page.getByRole("option", { name: E2E_TEAM_KEYGEN_ALIAS })).toBeVisible();
await expect(page.getByRole("option")).toHaveCount(3);
});
});

View file

@ -2,7 +2,14 @@ import { test, expect, type Locator, type Page as PlaywrightPage } from "@playwr
import { ADMIN_STORAGE_PATH } from "../../constants";
import { navigateToPage, dismissFeedbackPopup } from "../../helpers/navigation";
import { Page } from "../../fixtures/pages";
import { CHAT_MODEL_A, MOCK_RESPONSE_TEXT, sendChatCompletion, waitForSpendLog } from "../../helpers/traffic";
import {
CHAT_MODEL_A,
MOCK_RESPONSE_TEXT,
sendChatCompletion,
waitForSpendLog,
waitForSpendLogByPrompt,
} from "../../helpers/traffic";
import { openPlayground, selectModel, sendMessage } from "../../helpers/playground";
/**
* Anchored to traffic this spec generates itself, with a unique prompt and end user per run, so it
@ -46,6 +53,23 @@ test.describe("Logs page", () => {
permissions: ["clipboard-read", "clipboard-write"],
});
test("a chat sent from the Playground lands in Logs with its content", async ({ page, request }) => {
const prompt = `logs-playground-prompt-${uniqueSuffix()}`;
await openPlayground(page);
await selectModel(page, CHAT_MODEL_A);
await sendMessage(page, prompt);
await expect(page.getByText(MOCK_RESPONSE_TEXT, { exact: false }).first()).toBeVisible({ timeout: 60_000 });
const requestId = await waitForSpendLogByPrompt(request, prompt);
const row = await openLogsForRequest(page, requestId);
await row.click();
const drawer = page.getByRole("dialog").first();
await expect(drawer.getByText("Request & Response")).toBeVisible({ timeout: 20_000 });
await expect(drawer.getByText(prompt, { exact: false }).first()).toBeVisible({ timeout: 20_000 });
await expect(drawer.getByText(MOCK_RESPONSE_TEXT, { exact: false }).first()).toBeVisible({ timeout: 20_000 });
});
test("a served request expands to its request and response", async ({ page, request }) => {
const prompt = `logs-detail-prompt-${uniqueSuffix()}`;
const requestId = await sendChatCompletion(request, {

View file

@ -1,7 +1,8 @@
import { test, expect } from "@playwright/test";
import { test, expect, type APIRequestContext } from "@playwright/test";
import { ADMIN_STORAGE_PATH } from "../../constants";
import { navigateToPage, dismissFeedbackPopup } from "../../helpers/navigation";
import { Page } from "../../fixtures/pages";
import { masterKey } from "../../helpers/traffic";
test.describe("AI Hub (internal admin view)", () => {
test.use({ storageState: ADMIN_STORAGE_PATH });
@ -77,4 +78,89 @@ test.describe("Public model hub (/ui/model_hub_table)", () => {
// agents/MCP servers exist, so we don't assert on them in a fresh CI run.
await expect(page.getByRole("tab", { name: "Model Hub" })).toBeVisible({ timeout: 10_000 });
});
test("Agent Hub and MCP Hub tabs render their public entries", async ({ page, request }) => {
const suffix = `${Date.now()}`;
const agentName = `e2e-public-agent-${suffix}`;
const mcpServerName = `e2e_public_mcp_${suffix}`;
const auth = { Authorization: `Bearer ${masterKey()}` };
const publicMcpServerIds = async (api: APIRequestContext): Promise<string[]> => {
const res = await api.get("/public/mcp_hub");
expect(res.ok(), `public mcp_hub read failed (${res.status()}): ${await res.text()}`).toBe(true);
const servers: { server_id: string }[] = await res.json();
return servers.map((server) => server.server_id);
};
const seedPublicEntries = async (
api: APIRequestContext,
priorMcpIds: string[],
): Promise<{ agentId: string; serverId: string }> => {
const agentRes = await api.post("/v1/agents", {
headers: auth,
data: {
agent_name: agentName,
agent_card_params: {
name: agentName,
description: "E2E public agent",
version: "1.0.0",
url: "http://127.0.0.1:9999/",
capabilities: {},
skills: [],
defaultInputModes: ["text"],
defaultOutputModes: ["text"],
},
},
});
expect(agentRes.ok(), `agent create failed (${agentRes.status()}): ${await agentRes.text()}`).toBe(true);
const agentId = (await agentRes.json()).agent_id as string;
const serverRes = await api.post("/v1/mcp/server", {
headers: auth,
data: {
server_name: mcpServerName,
url: "http://127.0.0.1:9999/mcp",
transport: "http",
description: "E2E public MCP server",
},
});
expect(serverRes.ok(), `mcp server create failed (${serverRes.status()}): ${await serverRes.text()}`).toBe(true);
const serverId = (await serverRes.json()).server_id as string;
const agentPublicRes = await api.post(`/v1/agents/${agentId}/make_public`, { headers: auth });
expect(agentPublicRes.ok(), `agent make_public failed: ${await agentPublicRes.text()}`).toBe(true);
const mcpPublicRes = await api.post("/v1/mcp/make_public", {
headers: auth,
data: { mcp_server_ids: [...priorMcpIds, serverId] },
});
expect(mcpPublicRes.ok(), `mcp make_public failed: ${await mcpPublicRes.text()}`).toBe(true);
return { agentId, serverId };
};
const priorMcpIds = await publicMcpServerIds(request);
const { agentId, serverId } = await seedPublicEntries(request, priorMcpIds);
try {
await page.goto(`/ui/model_hub_table?key=${masterKey()}`);
await dismissFeedbackPopup(page);
const agentHubTab = page.getByRole("tab", { name: "Agent Hub" });
await expect(agentHubTab).toBeVisible({ timeout: 15_000 });
await agentHubTab.click();
await expect(page.getByText("Available Agents")).toBeVisible({ timeout: 10_000 });
await expect(page.getByRole("row").filter({ hasText: agentName })).toHaveCount(1, { timeout: 10_000 });
await expect(page.getByText("E2E public agent").first()).toBeVisible();
const mcpHubTab = page.getByRole("tab", { name: "MCP Hub" });
await expect(mcpHubTab).toBeVisible();
await mcpHubTab.click();
await expect(page.getByText("Available MCP Servers")).toBeVisible({ timeout: 10_000 });
await expect(page.getByRole("row").filter({ hasText: mcpServerName })).toHaveCount(1, { timeout: 10_000 });
await expect(page.getByText("E2E public MCP server").first()).toBeVisible();
} finally {
await request.post("/v1/mcp/make_public", { headers: auth, data: { mcp_server_ids: priorMcpIds } });
await request.delete(`/v1/agents/${agentId}`, { headers: auth });
await request.delete(`/v1/mcp/server/${serverId}`, { headers: auth });
}
});
});

View file

@ -212,6 +212,116 @@ test.describe("Add Model", () => {
.toBe(true);
});
test("Add a model with a stored credential, pass Test Connect, and serve traffic", async ({ page, request }) => {
const masterKey = users[Role.ProxyAdmin].password;
const auth = { Authorization: `Bearer ${masterKey}` };
const credentialName = `e2e-cred-reuse-${Date.now()}`;
const createCred = await page.request.post("/credentials", {
headers: auth,
data: {
credential_name: credentialName,
credential_values: { api_key: "fake-key", api_base: MOCK_LLM_BASE },
credential_info: { custom_llm_provider: "openai" },
},
});
expect(createCred.ok(), `POST /credentials failed (${createCred.status()}): ${await createCred.text()}`).toBe(true);
// Multi-instance stacks propagate a new credential to the probe-serving instances on a periodic
// sync; consecutive successes guard against a load balancer alternating synced and stale replicas
let consecutiveProbeSuccesses = 0;
await expect
.poll(
async () => {
const probe = await page.request.post("/health/test_connection", {
headers: auth,
data: {
litellm_params: {
model: "openai/fake-gpt-4",
custom_llm_provider: "openai",
litellm_credential_name: credentialName,
},
model_info: {},
mode: "chat",
},
});
const healthy = probe.ok() && (await probe.json()).status === "success";
consecutiveProbeSuccesses = healthy ? consecutiveProbeSuccesses + 1 : 0;
return consecutiveProbeSuccesses;
},
{
message: `stored credential ${credentialName} never became usable for a connection test`,
timeout: 60_000,
},
)
.toBeGreaterThanOrEqual(3);
try {
await navigateToPage(page, Page.Models);
await page.getByRole("tab", { name: "Add Model" }).click();
await selectProvider(page, "OpenAI-Compatible Endpoints (Together AI, etc.)");
const publicName = `e2e-cred-model-${Date.now()}`;
uiAddedModelName = publicName;
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "Custom Model Name (Enter below)" }).click();
await page.keyboard.press("Escape");
await page.getByPlaceholder("Enter custom model name").fill(publicName);
const credentialSelect = page.getByRole("combobox", { name: "Existing Credentials" });
await credentialSelect.click();
await credentialSelect.fill(credentialName);
await page.getByRole("option", { name: credentialName, exact: true }).click();
await expect(page.locator("#api_key")).toHaveCount(0);
await expect(page.locator("#api_base")).toHaveCount(0);
await page.getByRole("button", { name: "Test Connect" }).click();
await expect(page.getByText("Connection Test Results")).toBeVisible({ timeout: 10_000 });
await expect(page.getByTestId("connection-success-msg")).toBeVisible({ timeout: 30_000 });
const resultsModal = page.getByRole("dialog", { name: "Connection Test Results" });
await resultsModal.locator('[data-slot="dialog-footer"]').getByRole("button", { name: "Close" }).click();
await expect(resultsModal).toBeHidden({ timeout: 5_000 });
const created = await captureRequestBody(page, { method: "POST", urlIncludes: "/model/new" }, async () => {
await page.getByRole("button", { name: "Add Model" }).last().click();
});
expect(created.litellm_params?.litellm_credential_name, "the picked credential goes on the wire").toBe(
credentialName,
);
expect(created.litellm_params?.api_key, "no raw api key goes on the wire").toBeUndefined();
await expect(page.getByText("created successfully")).toBeVisible({ timeout: 15_000 });
await expect
.poll(
async () => {
try {
await sendChatCompletion(request, { model: publicName, prompt: `hello via ${credentialName}` });
return true;
} catch {
return false;
}
},
{
message: `model ${publicName} added with a stored credential never served a request`,
timeout: 30_000,
},
)
.toBe(true);
} finally {
const stored = uiAddedModelName ? await findDeploymentByName(page, uiAddedModelName) : undefined;
const id = stored?.model_info?.id;
if (id) {
await page.request.post("/model/delete", { headers: auth, data: { id } });
uiAddedModelName = "";
}
await page.request.delete(`/credentials/${credentialName}`, { headers: auth });
}
});
test("Test connection with bad credentials shows failure", async ({ page }) => {
await navigateToPage(page, Page.Models);
await page.getByRole("tab", { name: "Add Model" }).click();

View file

@ -0,0 +1,72 @@
import { test, expect, type Page as PlaywrightPage } from "@playwright/test";
import { ADMIN_STORAGE_PATH, E2E_TEAM_CRUD_ID } from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage } from "../../helpers/navigation";
import { readBack } from "../../helpers/roundTrip";
import { masterKey } from "../../helpers/traffic";
type DeploymentRow = { model_name?: string };
async function findDeploymentByName(page: PlaywrightPage, modelName: string): Promise<DeploymentRow | undefined> {
const body = await readBack<{ data: DeploymentRow[] }>(page, "/v2/model/info");
return body.data.find((row) => row.model_name === modelName);
}
test.describe("Delete team model", () => {
test.use({ storageState: ADMIN_STORAGE_PATH });
test("Delete a team-scoped model and verify it leaves the team's model list", async ({ page }) => {
const modelName = `e2e-team-model-delete-${Date.now()}`;
const createResponse = await page.request.post("/model/new", {
headers: { Authorization: `Bearer ${masterKey()}` },
data: {
model_name: modelName,
litellm_params: {
model: "openai/fake-gpt-4",
api_base: `http://127.0.0.1:${process.env.MOCK_LLM_PORT ?? "8090"}/v1`,
api_key: "fake-key",
},
model_info: { team_id: E2E_TEAM_CRUD_ID },
},
});
expect(createResponse.ok(), `/model/new failed: ${createResponse.status()} ${await createResponse.text()}`).toBe(
true,
);
await expect
.poll(async () => (await findDeploymentByName(page, modelName)) !== undefined, {
message: `deployment ${modelName} never appeared in /v2/model/info after create`,
timeout: 30_000,
})
.toBe(true);
await navigateToPage(page, Page.Models);
await page.getByPlaceholder("Search model names").fill(modelName);
const row = page.getByRole("row").filter({ hasText: modelName });
await expect(row).toHaveCount(1, { timeout: 15_000 });
await expect(row.getByText(E2E_TEAM_CRUD_ID)).toBeVisible({ timeout: 10_000 });
await row.getByRole("button", { name: "Delete model" }).click();
const modal = page.getByRole("dialog", { name: "Delete Model" });
await expect(modal).toBeVisible({ timeout: 5_000 });
await expect(modal.getByText(modelName).first()).toBeVisible();
await modal.getByRole("button", { name: "Delete", exact: true }).click();
await expect(page.getByText("Model deleted successfully").first()).toBeVisible({ timeout: 10_000 });
await expect(row).toHaveCount(0, { timeout: 15_000 });
await expect
.poll(async () => await findDeploymentByName(page, modelName), {
message: `deployment ${modelName} still readable from /v2/model/info after delete`,
timeout: 15_000,
})
.toBeUndefined();
await page.reload();
await page.getByPlaceholder("Search model names").fill(modelName);
await expect(page.getByText("No models found").first()).toBeVisible({ timeout: 15_000 });
await expect(page.getByRole("row").filter({ hasText: modelName })).toHaveCount(0);
});
});

View file

@ -0,0 +1,97 @@
import { test, expect } from "@playwright/test";
import { ADMIN_STORAGE_PATH } from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage, dismissFeedbackPopup } from "../../helpers/navigation";
import { CHAT_MODEL_A, MOCK_RESPONSE_TEXT, masterKey } from "../../helpers/traffic";
test.describe("Second proxy admin", () => {
test.use({ storageState: { cookies: [], origins: [] } });
test("an invited admin can log in, mint a key, and call a model with it", async ({ page, browser, request }) => {
const suffix = Date.now();
const email = `second-admin-${suffix}@test.local`;
const password = "e2e-second-admin-password";
const auth = { Authorization: `Bearer ${masterKey()}` };
const inviteAdminUser = async (): Promise<string> => {
const adminContext = await browser.newContext({ storageState: ADMIN_STORAGE_PATH });
try {
const adminPage = await adminContext.newPage();
await navigateToPage(adminPage, Page.Users);
await dismissFeedbackPopup(adminPage);
await adminPage.getByRole("button", { name: "+ Invite User", exact: true }).click();
const dialog = adminPage.getByRole("dialog", { name: "Invite User" });
await expect(dialog).toBeVisible({ timeout: 5_000 });
await dialog.getByLabel("User Email").fill(email);
await dialog.getByLabel(/Global Proxy Role/).click();
await adminPage.getByRole("option", { name: /Admin \(All Permissions\)/ }).click();
const createdResponse = adminPage.waitForResponse(
(res) => res.url().includes("/user/new") && res.request().method() === "POST",
);
await dialog.getByRole("button", { name: "Invite User" }).click();
const createdBody = await (await createdResponse).json();
const createdUserId = (createdBody.data?.user_id ?? createdBody.user_id) as string;
expect(createdUserId, "created user id from /user/new").toBeTruthy();
await expect(adminPage.getByText("API user Created").first()).toBeVisible({ timeout: 10_000 });
return createdUserId;
} finally {
await adminContext.close();
}
};
const userId = await inviteAdminUser();
try {
const passwordRes = await request.post("/user/update", {
headers: auth,
data: { user_email: email, password },
});
expect(passwordRes.ok(), `setting password failed (${passwordRes.status()}): ${await passwordRes.text()}`).toBe(
true,
);
await page.goto("/ui/login");
await page.getByPlaceholder("Enter your username").fill(email);
await page.getByPlaceholder("Enter your password").fill(password);
await page.getByRole("button", { name: "Login", exact: true }).click();
await expect(page.locator("a", { hasText: "Virtual Keys" })).toBeVisible({ timeout: 30_000 });
await dismissFeedbackPopup(page);
await navigateToPage(page, Page.ApiKeys);
await page.getByRole("button", { name: /Create New Key/i }).click();
await expect(page.getByText("Key Ownership")).toBeVisible({ timeout: 10_000 });
await page.getByLabel(/Key Name/).fill(`e2e-second-admin-key-${suffix}`);
await page.getByRole("combobox", { name: "Select models" }).click();
await page.getByRole("option", { name: "All Proxy Models", exact: true }).click();
await page.keyboard.press("Escape");
await page.getByRole("button", { name: "Create Key", exact: true }).click();
await expect(page.getByText("Save your Key")).toBeVisible({ timeout: 10_000 });
const apiKey = (await page.getByRole("dialog", { name: "Save your Key" }).locator("pre").innerText()).trim();
expect(apiKey).toMatch(/^sk-/);
await page.keyboard.press("Escape");
const response = await page.request.post("/chat/completions", {
headers: { Authorization: `Bearer ${apiKey}` },
data: {
model: CHAT_MODEL_A,
messages: [{ role: "user", content: `second admin ping ${suffix}` }],
},
});
expect(response.status()).toBe(200);
const body = await response.json();
expect(body.choices?.[0]?.message?.content).toBe(MOCK_RESPONSE_TEXT);
} finally {
if (userId) {
await request.post("/user/delete", { headers: auth, data: { user_ids: [userId] } });
}
}
});
});

View file

@ -1,6 +1,7 @@
import { test, expect, type Page as PlaywrightPage } from "@playwright/test";
import {
E2E_INTERNAL_USER_KEY_ALIAS,
E2E_TEAM_ADMIN_USER_ID,
E2E_TEAM_CRUD_ALIAS,
E2E_TEAM_CRUD_ID,
TEAM_ADMIN_STORAGE_PATH,
@ -8,6 +9,8 @@ import {
import { Page } from "../../fixtures/pages";
import { navigateToPage, dismissFeedbackPopup, clickTeamId } from "../../helpers/navigation";
import { captureRequestBody, readBack } from "../../helpers/roundTrip";
import { CHAT_MODEL_A, masterKey } from "../../helpers/traffic";
import { keySourceSelect, modelSelect, onlyVisible, openPlayground } from "../../helpers/playground";
/**
* Every identifier a roster is addressable by. Which of user_id / user_email is populated depends on
@ -128,6 +131,91 @@ test.describe("Team Admin", () => {
.not.toContain("e2e-removable-member");
});
test("Team admin sees all team models in the Playground model dropdown", async ({ page, request }) => {
const suffix = Date.now();
const teamModelName = `e2e-team-dropdown-model-${suffix}`;
const auth = { Authorization: `Bearer ${masterKey()}` };
const teamRes = await request.post("/team/new", {
headers: auth,
data: {
team_alias: `e2e-playground-team-${suffix}`,
models: [CHAT_MODEL_A],
members_with_roles: [{ role: "admin", user_id: E2E_TEAM_ADMIN_USER_ID }],
},
});
expect(teamRes.ok(), `team create failed (${teamRes.status()}): ${await teamRes.text()}`).toBe(true);
const teamId = (await teamRes.json()).team_id as string;
try {
const modelRes = await request.post("/model/new", {
headers: auth,
data: {
model_name: teamModelName,
litellm_params: {
model: "openai/fake-gpt-4",
api_base: `http://127.0.0.1:${process.env.MOCK_LLM_PORT ?? "8090"}/v1`,
api_key: "fake-key",
},
model_info: { team_id: teamId },
},
});
expect(modelRes.ok(), `model create failed (${modelRes.status()}): ${await modelRes.text()}`).toBe(true);
const modelId = (await modelRes.json()).model_info?.id as string;
try {
const keyRes = await request.post("/key/generate", { headers: auth, data: { team_id: teamId } });
expect(keyRes.ok(), `key generate failed (${keyRes.status()}): ${await keyRes.text()}`).toBe(true);
const teamKey = (await keyRes.json()).key as string;
try {
await expect
.poll(
async () => {
const res = await request.get("/model_group/info", {
headers: { Authorization: `Bearer ${teamKey}` },
});
if (!res.ok()) return false;
const body: { data?: { model_group?: string }[] } = await res.json();
return (body.data ?? []).some((group) => group.model_group === teamModelName);
},
{
message: `model group ${teamModelName} never became visible to the team key`,
timeout: 30_000,
},
)
.toBe(true);
await openPlayground(page);
await keySourceSelect(page).click();
await onlyVisible(page.getByRole("option", { name: "Virtual Key" })).click({ timeout: 15_000 });
const keyInput = onlyVisible(page.getByPlaceholder("Enter custom Virtual Key"));
await expect(keyInput).toBeVisible({ timeout: 10_000 });
await keyInput.fill(teamKey);
const select = modelSelect(page);
await select.click();
await select.fill(teamModelName);
await expect(onlyVisible(page.getByRole("option", { name: teamModelName }))).toBeVisible({
timeout: 15_000,
});
await select.fill(CHAT_MODEL_A);
await expect(onlyVisible(page.getByRole("option", { name: CHAT_MODEL_A }))).toBeVisible({
timeout: 15_000,
});
} finally {
await request.post("/key/delete", { headers: auth, data: { keys: [teamKey] } });
}
} finally {
await request.post("/model/delete", { headers: auth, data: { id: modelId } });
}
} finally {
await request.post("/team/delete", { headers: auth, data: { team_ids: [teamId] } });
}
});
test("Team admin can create a team key with All Team Models", async ({ page }) => {
await navigateToPage(page, Page.ApiKeys);
await dismissFeedbackPopup(page);

View file

@ -1,15 +1,14 @@
"""
Regression tests for Redis connection pool leak fixes (RC1-RC5).
Tests are pure unit tests — no Redis server required.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import redis.asyncio as async_redis
from litellm._redis import get_redis_async_client, get_redis_connection_pool
from litellm._redis import (
_coerce_redis_kwargs_types,
_get_redis_client_logic,
_get_redis_env_kwarg_mapping,
get_redis_async_client,
get_redis_connection_pool,
)
def test_url_config_uses_passed_pool():
@ -60,16 +59,14 @@ def test_max_connections_url_config_string_value(monkeypatch):
assert pool.max_connections == 25
def test_max_connections_url_config_invalid_value():
"""Invalid max_connections should be silently ignored, falling back
to the pool default (50 for BlockingConnectionPool)."""
with patch("litellm._redis._get_redis_client_logic") as mock_logic:
mock_logic.return_value = {
"url": "redis://localhost:6379/0",
"max_connections": "not_a_number",
}
def test_max_connections_url_config_invalid_value(monkeypatch):
"""Invalid max_connections from an env var should be silently dropped,
falling back to the pool default (50 for BlockingConnectionPool)."""
monkeypatch.setenv("REDIS_URL", "redis://localhost:6379/0")
monkeypatch.delenv("REDIS_HOST", raising=False)
monkeypatch.setenv("REDIS_MAX_CONNECTIONS", "not_a_number")
pool = get_redis_connection_pool()
pool = get_redis_connection_pool()
# BlockingConnectionPool default is 50
assert pool.max_connections == 50
@ -128,3 +125,173 @@ async def test_disconnect_idempotent():
await cache.disconnect()
await cache.disconnect() # should not raise
def test_coerce_redis_kwargs_types_int():
"""String values for int-typed Redis params are coerced to int."""
result = _coerce_redis_kwargs_types({"health_check_interval": "30", "port": "6380", "db": "1"})
assert result["health_check_interval"] == 30
assert isinstance(result["health_check_interval"], int)
assert result["port"] == 6380
assert result["db"] == 1
def test_coerce_redis_kwargs_types_bool():
"""String values for bool-typed Redis params are coerced to bool."""
result = _coerce_redis_kwargs_types({"ssl": "true", "decode_responses": "false"})
assert result["ssl"] is True
assert result["decode_responses"] is False
def test_coerce_redis_kwargs_types_none_default_numeric():
"""String values for known None-default numeric params are coerced."""
result = _coerce_redis_kwargs_types({"max_connections": "20", "socket_timeout": "5.5"})
assert result["max_connections"] == 20
assert isinstance(result["max_connections"], int)
assert result["socket_timeout"] == 5.5
assert isinstance(result["socket_timeout"], float)
def _redis_signature_pre_8x(
socket_timeout=None,
socket_connect_timeout=None,
max_connections=None,
health_check_interval=0,
):
"""Stand-in for the redis-py <= 7.x Redis signature, where the timeout defaults are None."""
def _redis_signature_8x(
socket_timeout=5,
socket_connect_timeout=5,
max_connections=None,
health_check_interval=0,
):
"""Stand-in for the redis-py 8.x Redis signature, where the timeout defaults became int 5."""
@pytest.mark.parametrize(
"client",
[_redis_signature_pre_8x, _redis_signature_8x],
ids=["redis-py<=7.x", "redis-py-8.x"],
)
def test_coerce_fractional_socket_timeout_survives_signature_default_change(client):
"""redis-py 8.x changed socket_timeout's default from None to int 5. Deriving the
target type from the signature default made int("5.5") raise, so the key was dropped
and REDIS_SOCKET_TIMEOUT=5.5 silently disappeared on 8.x."""
result = _coerce_redis_kwargs_types(
{"socket_timeout": "5.5", "socket_connect_timeout": "2.5", "max_connections": "20"},
client=client,
)
assert result["socket_timeout"] == pytest.approx(5.5)
assert isinstance(result["socket_timeout"], float)
assert result["socket_connect_timeout"] == pytest.approx(2.5)
assert isinstance(result["socket_connect_timeout"], float)
assert result["max_connections"] == 20
assert isinstance(result["max_connections"], int)
def test_coerce_invalid_socket_timeout_is_still_dropped():
"""Garbage must not survive the explicit-type path; Redis falls back to its own default."""
result = _coerce_redis_kwargs_types({"socket_timeout": "not_a_number"}, client=_redis_signature_8x)
assert "socket_timeout" not in result
def test_coerce_redis_kwargs_types_invalid_drops_key():
"""A string that cannot be coerced to the expected numeric type is dropped."""
result = _coerce_redis_kwargs_types({"health_check_interval": "not_a_number"})
assert "health_check_interval" not in result
def test_coerce_redis_kwargs_types_non_string_unchanged():
"""Non-string values pass through without modification."""
result = _coerce_redis_kwargs_types({"health_check_interval": 30, "ssl": True})
assert result["health_check_interval"] == 30
assert result["ssl"] is True
def test_health_check_interval_from_env_is_int(monkeypatch):
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_HEALTH_CHECK_INTERVAL", "30")
pool = get_redis_connection_pool()
assert pool is not None
interval = pool.connection_kwargs.get("health_check_interval")
assert interval == 30
assert isinstance(interval, int), f"Expected int, got {type(interval)}: {interval!r}"
def _signature_without_defaults(testkey):
"""Stand-in for a client whose parameter declares no default at all."""
def _signature_with_float_default(myparam=1.0):
"""Stand-in for a client whose parameter declares a float default."""
def test_coerce_redis_kwargs_types_empty_default_param_unchanged():
"""String params whose signature entry has no default (inspect.Parameter.empty) are left as-is."""
result = _coerce_redis_kwargs_types({"testkey": "some_value"}, client=_signature_without_defaults)
assert result["testkey"] == "some_value"
assert isinstance(result["testkey"], str)
def test_coerce_redis_kwargs_types_float_valid():
"""String values for params whose signature default is a float are coerced to float."""
result = _coerce_redis_kwargs_types({"myparam": "3.14"}, client=_signature_with_float_default)
assert result["myparam"] == pytest.approx(3.14)
assert isinstance(result["myparam"], float)
def test_coerce_redis_kwargs_types_float_invalid_drops_key():
"""An unconvertible string for a float-default param is dropped from the result."""
result = _coerce_redis_kwargs_types({"myparam": "not_a_float"}, client=_signature_with_float_default)
assert "myparam" not in result
@pytest.mark.parametrize(
("raw", "expected"),
[("false", False), ("true", True), ("0", False), ("1", True)],
)
def test_coerce_socket_keepalive_string(raw, expected):
"""socket_keepalive's signature default is None, so it needs an explicit bool
coercion: a leftover "false" string is truthy and enables keepalive."""
result = _coerce_redis_kwargs_types({"socket_keepalive": raw})
assert result["socket_keepalive"] is expected
def test_get_redis_client_logic_coerces_cluster_only_kwargs(monkeypatch):
"""Cluster-only kwargs (absent from redis.Redis's signature) must still be
coerced when routing to a cluster, or Helm-stringified values reach
RedisCluster as strings."""
for envvar in (*_get_redis_env_kwarg_mapping(), "REDIS_CLUSTER_NODES", "REDIS_SENTINEL_NODES"):
monkeypatch.delenv(envvar, raising=False)
result = _get_redis_client_logic(
startup_nodes='[{"host": "localhost", "port": 7000}]',
cluster_error_retry_attempts="5",
require_full_coverage="false",
health_check_interval="30",
)
assert result["cluster_error_retry_attempts"] == 5
assert isinstance(result["cluster_error_retry_attempts"], int)
assert result["require_full_coverage"] is False
assert result["health_check_interval"] == 30
assert isinstance(result["health_check_interval"], int)
def test_get_redis_client_logic_raises_without_host_or_url(monkeypatch):
"""_get_redis_client_logic raises ValueError when neither host nor url is provided."""
for envvar in (*_get_redis_env_kwarg_mapping(), "REDIS_CLUSTER_NODES", "REDIS_SENTINEL_NODES"):
monkeypatch.delenv(envvar, raising=False)
with pytest.raises(ValueError, match="Either 'host' or 'url' must be specified for redis"):
_get_redis_client_logic()

View file

@ -1,3 +1,4 @@
import base64
from typing import Final
from unittest.mock import MagicMock
@ -8,8 +9,17 @@ from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS
from litellm.endpoints.speech.speech_to_completion_bridge.transformation import (
SpeechToCompletionBridgeTransformationHandler,
)
from litellm.types.utils import ChatCompletionAudioResponse, Choices, Message, ModelResponse
GEMINI_TTS_MODEL: Final = "gemini-3.1-flash-tts-preview"
PCM_BYTES: Final = b"\x01\x02\x03\x04" * 6
def _model_response(model: str, pcm: bytes) -> ModelResponse:
audio: Final = ChatCompletionAudioResponse(
data=base64.b64encode(pcm).decode(), expires_at=0, transcript="hello"
)
return ModelResponse(model=model, choices=[Choices(message=Message(content=None, audio=audio))])
def _bridge_request(response_format: str | None) -> dict:
@ -28,7 +38,7 @@ def _bridge_request(response_format: str | None) -> dict:
)
@pytest.mark.parametrize("response_format", ["wav", "mp3", "pcm", None])
@pytest.mark.parametrize("response_format", ["wav", "pcm", None])
def test_gemini_tts_request_keeps_speech_response_format_out_of_chat_params(response_format: str | None) -> None:
request: Final = _bridge_request(response_format)
@ -60,3 +70,48 @@ def test_non_gemini_request_forwards_speech_response_format_as_audio_format() ->
assert "response_format" not in request
assert request["audio"] == {"voice": "alloy", "format": "wav"}
@pytest.mark.parametrize("response_format", ["mp3", "flac", "opus", "aac"])
def test_gemini_tts_request_rejects_formats_gemini_cannot_produce(response_format: str) -> None:
with pytest.raises(litellm.BadRequestError) as excinfo:
_bridge_request(response_format)
assert excinfo.value.status_code == 400
assert response_format in str(excinfo.value)
assert "pcm" in str(excinfo.value)
assert "wav" in str(excinfo.value)
def test_gemini_tts_pcm_response_returns_raw_pcm_bytes() -> None:
response: Final = SpeechToCompletionBridgeTransformationHandler().transform_response(
model_response=_model_response(GEMINI_TTS_MODEL, PCM_BYTES),
response_format="pcm",
)
assert response.response.content == PCM_BYTES
assert response.response.headers["content-type"] == "audio/pcm"
@pytest.mark.parametrize("response_format", ["wav", None])
def test_gemini_tts_wav_and_default_responses_wrap_pcm_in_wav(response_format: str | None) -> None:
response: Final = SpeechToCompletionBridgeTransformationHandler().transform_response(
model_response=_model_response(GEMINI_TTS_MODEL, PCM_BYTES),
response_format=response_format,
)
body: Final = response.response.content
assert body[:4] == b"RIFF"
assert body[8:12] == b"WAVE"
assert body[44:] == PCM_BYTES
assert response.response.headers["content-type"] == "audio/wav"
def test_non_gemini_response_keeps_original_bytes_and_mpeg_content_type() -> None:
response: Final = SpeechToCompletionBridgeTransformationHandler().transform_response(
model_response=_model_response("gpt-4o-audio-preview", PCM_BYTES),
response_format="mp3",
)
assert response.response.content == PCM_BYTES
assert response.response.headers["content-type"] == "audio/mpeg"

View file

@ -65,6 +65,7 @@ def _job_record(job: ActiveShadowEvalJob, target_type="key", target_id="key-hash
target_type=target_type,
target_id=target_id,
router_name=job.router_name,
router_names=job.router_names,
direction=job.direction,
baseline_model=job.baseline_model,
shadow_percentage=job.shadow_percentage,
@ -81,6 +82,7 @@ def _router(
shadow_text="shadow answer",
judge_json='{"preference": "A", "confidence": 0.9, "reasoning": "x"}',
classifier_cost=None,
sibling_router_texts=None,
):
"""One mock router serving the shadow call first, the judge call second, told apart by
the internal-origin stamp rather than the model, since a reverse job's shadow arm names
@ -100,6 +102,15 @@ def _router(
decision["classifier_cost"] = classifier_cost
kwargs["metadata"]["routing_decision"] = decision
return {"choices": [{"message": {"content": shadow_text}}], "usage": {"completion_tokens": 5}}
if sibling_router_texts and kwargs["model"] in sibling_router_texts:
kwargs["metadata"]["routing_decision"] = {
"tier_label": "MEDIUM",
"routed_model": f"{kwargs['model']}-pick",
}
return {
"choices": [{"message": {"content": sibling_router_texts[kwargs["model"]]}}],
"usage": {"completion_tokens": 5},
}
return ModelResponse(
model=kwargs["model"],
choices=[{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": shadow_text}}],
@ -1210,29 +1221,36 @@ class TestJobValidation:
{"direction": "reverse"},
{"baseline_model": "baseline-model"},
{"direction": "sideways", "baseline_model": "baseline-model"},
{"direction": "reverse", "baseline_model": "baseline-model", "router_names": ("a", "b")},
],
ids=["reverse-without-baseline", "forward-with-baseline", "unknown-direction"],
ids=["reverse-without-baseline", "forward-with-baseline", "unknown-direction", "reverse-with-router-set"],
)
def test_unsamplable_shapes_are_rejected(self, overrides):
with pytest.raises(ValidationError):
_job(**overrides)
def test_shadow_target_follows_direction(self):
assert _job().shadow_target == "my-router"
assert _reverse_job().shadow_target == "baseline-model"
def test_arm_target_follows_direction(self):
assert _job().arm_target("my-router") == "my-router"
assert _reverse_job().arm_target("my-router") == "baseline-model"
def test_rows_from_before_router_names_carry_their_set_in_router_name(self):
assert _job().arm_router_names == ("my-router",)
assert _job(router_names=("my-router", "alt-router")).arm_router_names == ("my-router", "alt-router")
@pytest.mark.asyncio
class TestDirection:
@pytest.mark.parametrize(
"job,routed_by,sampled",
"job,routed_by,attempt_rows",
[
(_job(), None, True),
(_job(), "my-router", False),
(_job(), "other-router", True),
(_reverse_job(), "my-router", True),
(_reverse_job(), None, False),
(_reverse_job(), "other-router", False),
(_job(), None, 1),
(_job(), "my-router", 0),
(_job(), "other-router", 1),
(_reverse_job(), "my-router", 1),
(_reverse_job(), None, 0),
(_reverse_job(), "other-router", 0),
(_job(router_names=("my-router", "alt-router")), "alt-router", 0),
(_job(router_names=("my-router", "alt-router")), "other-router", 2),
],
ids=[
"forward-samples-unrouted",
@ -1241,20 +1259,24 @@ class TestDirection:
"reverse-samples-its-own-router",
"reverse-skips-unrouted",
"reverse-skips-another-router",
"forward-skips-any-candidates-own-traffic",
"forward-multi-samples-once-per-arm",
],
)
async def test_direction_decides_which_traffic_is_sampled(self, job, routed_by, sampled):
async def test_direction_decides_which_traffic_is_sampled(self, job, routed_by, attempt_rows):
"""The two directions partition the key's traffic: whatever one samples, the other
skips, so a key running both never judges the same turn twice for the same reason."""
skips, so a key running both never judges the same turn twice for the same reason.
A multi-router job extends the forward skip to every candidate: a request one
candidate served must not be judged as the incumbent against another candidate."""
prisma = _prisma()
logger = _logger(router=_router(), prisma=prisma, jobs=(job,))
logger = _logger(router=_router(sibling_router_texts={"alt-router": "alt answer"}), prisma=prisma, jobs=(job,))
await logger.async_log_success_event(
_success_kwargs(request_metadata=_routed_by(routed_by) if routed_by else {}), RESPONSE, None, None
)
await _drain(logger)
assert prisma.db.litellm_shadowevalattempt.create.await_count == int(sampled)
assert prisma.db.litellm_shadowevalattempt.create.await_count == attempt_rows
async def test_reverse_duplicates_against_the_baseline_model(self):
prisma = _prisma()
@ -1316,6 +1338,134 @@ class TestDirection:
assert logger._job_starts == {"forward-job": 1, "reverse-job": 1}
@pytest.mark.asyncio
class TestMultiRouterArms:
async def test_every_arm_judges_the_same_request_and_stamps_its_own_row(self):
"""One sampled request, one row per candidate router, both judged against the same
real response: the paired comparison that makes multi-router win rates comparable."""
prisma = _prisma()
router = _router(sibling_router_texts={"alt-router": "alt answer"})
logger = _logger(router=router, prisma=prisma)
await logger._run_shadow_eval(
job=_job(router_names=("my-router", "alt-router")),
request_id="req-1",
messages=({"role": "user", "content": "hi"},),
real_text="real answer",
real_model="claude-opus",
real_cost=0.001,
real_classifier_cost=0.0,
real_cache_hit=False,
control_tier=None,
shadow_params={},
parent_metadata={},
)
rows = [call.kwargs["data"] for call in prisma.db.litellm_shadowevalattempt.create.await_args_list]
assert [row["router_name"] for row in rows] == ["my-router", "alt-router"]
assert {row["request_id"] for row in rows} == {"req-1"}
assert [row["shadow_model"] for row in rows] == ["cheap-model", "alt-router-pick"]
assert all(row["outcome"] in ("real", "shadow", "tie") for row in rows)
assert all(row["real_cost"] == 0.001 for row in rows)
async def test_a_single_router_job_stamps_its_router_on_the_row(self):
prisma = _prisma()
logger = _logger(router=_router(), prisma=prisma)
await logger._run_shadow_eval(
job=_job(),
request_id="req-1",
messages=({"role": "user", "content": "hi"},),
real_text="real answer",
real_model="claude-opus",
real_cost=0.0,
real_classifier_cost=0.0,
real_cache_hit=False,
control_tier=None,
shadow_params={},
parent_metadata={},
)
row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]
assert row["router_name"] == "my-router"
async def test_one_arms_failure_never_silences_the_sibling(self):
prisma = _prisma()
router = _router(sibling_router_texts={"alt-router": "alt answer"})
healthy = router.acompletion.side_effect
async def first_arm_explodes(**kwargs):
if kwargs["model"] == "my-router":
raise RuntimeError("provider exploded")
return await healthy(**kwargs)
router.acompletion.side_effect = first_arm_explodes
logger = _logger(router=router, prisma=prisma)
await logger._run_shadow_eval(
job=_job(router_names=("my-router", "alt-router")),
request_id="req-1",
messages=({"role": "user", "content": "hi"},),
real_text="real answer",
real_model="claude-opus",
real_cost=0.0,
real_classifier_cost=0.0,
real_cache_hit=False,
control_tier=None,
shadow_params={},
parent_metadata={},
)
rows = [call.kwargs["data"] for call in prisma.db.litellm_shadowevalattempt.create.await_args_list]
assert [row["router_name"] for row in rows] == ["my-router", "alt-router"]
assert rows[0]["outcome"] == "error"
assert "provider exploded" in rows[0]["error"]
assert rows[1]["outcome"] in ("real", "shadow", "tie")
async def test_the_turn_valve_counts_every_arm_a_start_will_write(self):
"""max_turns is a row ceiling and one sampled request writes one row per arm, so
admission pre-counts the arms: a two-arm job with two turns of budget admits one
request, not two."""
prisma = _prisma()
router = _router(sibling_router_texts={"alt-router": "alt answer"})
logger = _logger(
router=router, prisma=prisma, jobs=(_job(router_names=("my-router", "alt-router"), max_turns=2),)
)
await logger.async_log_success_event(_success_kwargs(request_id="req-1"), RESPONSE, None, None)
await logger.async_log_success_event(_success_kwargs(request_id="req-2"), RESPONSE, None, None)
await _drain(logger)
rows = [call.kwargs["data"] for call in prisma.db.litellm_shadowevalattempt.create.await_args_list]
assert {row["request_id"] for row in rows} == {"req-1"}
assert len(rows) == 2
async def test_a_withheld_request_runs_no_arm_and_counts_once(self):
"""The budget gates run once per sampled request, before any arm: funnel counters
stay per-request, so coverage math is arm-count independent."""
prisma = _prisma()
router = _router(sibling_router_texts={"alt-router": "alt answer"})
logger = _logger(router=router, prisma=prisma)
await logger._run_shadow_eval(
job=_job(router_names=("my-router", "alt-router"), max_budget=1.0, spend=2.0),
request_id="req-1",
messages=({"role": "user", "content": "hi"},),
real_text="real answer",
real_model="claude-opus",
real_cost=0.0,
real_classifier_cost=0.0,
real_cache_hit=False,
control_tier=None,
shadow_params={},
parent_metadata={},
)
router.acompletion.assert_not_called()
prisma.db.litellm_shadowevalattempt.create.assert_not_called()
assert logger._test_funnel == [("job-1", "withheld")]
@pytest.mark.asyncio
class TestActiveJobsFailClosed:
async def test_a_row_the_sampler_cannot_read_is_dropped_not_guessed(self):

View file

@ -14,7 +14,7 @@ from litellm.integrations.websearch_interception.handler import (
)
from litellm.llms.base_llm.search.transformation import SearchResponse
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, ProxyException, UserAPIKeyAuth
from litellm.types.utils import LlmProviders
from litellm.types.utils import CallTypes, LlmProviders
def test_initialize_from_proxy_config():
@ -230,6 +230,124 @@ async def test_execute_search_passes_selected_search_tool_litellm_params(monkeyp
assert forwarded_kwargs["max_retries"] == 2
@pytest.mark.asyncio
@pytest.mark.parametrize(
("search_tools", "error"),
[
pytest.param(None, "was not found", id="router-not-configured"),
pytest.param(
[{"search_tool_name": "other-search", "litellm_params": {"search_provider": "tavily"}}],
"was not found",
id="requested-tool-not-configured",
),
pytest.param(
[{"search_tool_name": "parallel-search", "litellm_params": "not-a-mapping"}],
"does not define a valid search provider",
id="invalid-parameters",
),
pytest.param(
[{"search_tool_name": "parallel-search", "litellm_params": {}}],
"does not define a valid search provider",
id="missing-provider",
),
pytest.param(
[{"search_tool_name": "parallel-search", "litellm_params": {"search_provider": " "}}],
"does not define a valid search provider",
id="whitespace-provider",
),
pytest.param(
[{"search_tool_name": "parallel-search", "litellm_params": {"search_provider": 123}}],
"does not define a valid search provider",
id="invalid-provider",
),
],
)
async def test_execute_search_rejects_invalid_explicit_search_tool(monkeypatch, search_tools, error):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(search_tool_name="parallel-search")
router = None if search_tools is None else MagicMock(search_tools=search_tools)
mock_asearch = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(ValueError, match=f"Configured search tool 'parallel-search' {error}"):
await logger._execute_search("what is litellm")
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_execute_search_honors_explicit_parallel_search_tool(monkeypatch):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(search_tool_name="parallel-search")
router = MagicMock(
search_tools=[
{
"search_tool_name": "other-search",
"litellm_params": {"search_provider": "tavily", "api_key": "other-key"},
},
{
"search_tool_name": "parallel-search",
"litellm_params": {"search_provider": "parallel_ai", "api_key": "parallel-key"},
},
],
)
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
await logger._execute_search("what is litellm")
mock_asearch.assert_awaited_once_with(
query="what is litellm",
search_provider="parallel_ai",
api_key="parallel-key",
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("search_tools", "expected_search_kwargs"),
[
pytest.param(None, {"search_provider": "perplexity"}, id="router-not-configured"),
pytest.param(
[
{
"search_tool_name": "first-search",
"litellm_params": {"search_provider": "tavily", "api_key": "first-key"},
},
{
"search_tool_name": "parallel-search",
"litellm_params": {"search_provider": "parallel_ai", "api_key": "parallel-key"},
},
],
{"search_provider": "tavily", "api_key": "first-key"},
id="first-configured-tool",
),
],
)
async def test_execute_search_preserves_implicit_provider_selection(monkeypatch, search_tools, expected_search_kwargs):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger()
router = None if search_tools is None else MagicMock(search_tools=search_tools)
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
await logger._execute_search("what is litellm")
mock_asearch.assert_awaited_once_with(query="what is litellm", **expected_search_kwargs)
@pytest.mark.asyncio
async def test_execute_search_attributes_spend_to_the_calling_key(monkeypatch):
"""An intercepted search is billed and logged against the key that made the LLM request.
@ -397,6 +515,72 @@ async def test_execute_search_enforces_team_search_tool_permission(monkeypatch):
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("call_type", "web_search_tool"),
[
pytest.param(
CallTypes.acompletion,
{"type": "web_search_20250305", "name": "web_search"},
id="chat-completion",
),
pytest.param(CallTypes.responses, {"type": "web_search"}, id="responses"),
pytest.param(CallTypes.aresponses, {"type": "web_search"}, id="async-responses"),
pytest.param(
CallTypes.anthropic_messages,
{"type": "web_search_20250305", "name": "web_search"},
id="anthropic-messages",
),
],
)
async def test_deployment_hook_dispatcher_propagates_missing_explicit_search_tool(
monkeypatch, call_type, web_search_tool
):
import litellm
from litellm.proxy import proxy_server
from litellm.utils import async_pre_call_deployment_hook
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="parallel-search")
mock_asearch = AsyncMock()
kwargs = {
"model": "bedrock/claude-sonnet-4",
"tools": [web_search_tool],
"custom_llm_provider": "bedrock",
}
monkeypatch.setattr(
proxy_server,
"llm_router",
MagicMock(search_tools=[{"search_tool_name": "other-search", "litellm_params": {"search_provider": "tavily"}}]),
)
monkeypatch.setattr(litellm, "callbacks", [logger])
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(ValueError, match="Configured search tool 'parallel-search' was not found"):
await async_pre_call_deployment_hook(kwargs=kwargs, call_type=call_type.value)
assert kwargs["tools"] == [web_search_tool]
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_deployment_hook_skips_explicit_tool_validation_for_non_search_responses(monkeypatch):
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="parallel-search")
monkeypatch.setattr(proxy_server, "llm_router", MagicMock(search_tools=[]))
result = await logger.async_pre_call_deployment_hook(
kwargs={
"tools": [{"type": "function", "name": "calculator"}],
"custom_llm_provider": "bedrock",
},
call_type=CallTypes.aresponses,
)
assert result is None
@pytest.mark.asyncio
async def test_async_pre_call_deployment_hook_provider_from_top_level_kwargs():
"""Test that async_pre_call_deployment_hook finds custom_llm_provider at top-level kwargs.

View file

@ -1433,3 +1433,50 @@ class TestFlattenTopLevelSchemaCombinators:
flatten_top_level_schema_combinators(schema)
assert schema == snapshot
class TestRequestContainsImageContent:
"""One detector for every dialect that reaches pre-routing hooks untranslated."""
@pytest.mark.parametrize(
"part",
[
{"type": "image_url", "image_url": {"url": "data:image/png;base64,aGk="}},
{"type": "input_image", "image_url": "data:image/png;base64,aGk="},
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "aGk="}},
{
"type": "tool_result",
"tool_use_id": "tu_1",
"content": [{"type": "image", "source": {"type": "base64", "data": "aGk="}}],
},
],
)
def test_detects_every_image_dialect_including_tool_results(self, part):
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}, part]}]
assert request_contains_image_content(messages) is True
@pytest.mark.parametrize(
"messages",
[
[{"role": "user", "content": "plain string"}],
[{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
[{"role": "user", "content": [{"type": "input_audio", "input_audio": {"data": "x"}}]}],
[{"role": "user", "content": [{"type": "tool_result", "content": [{"type": "text", "text": "ok"}]}]}],
[{"role": "user", "content": None}],
[],
],
)
def test_ignores_text_audio_and_degenerate_shapes(self, messages):
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
assert request_contains_image_content(messages) is False
def test_hostile_nesting_is_depth_bounded(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
nested: dict = {"type": "image", "source": {"type": "base64", "data": "aGk="}}
for _ in range(50):
nested = {"type": "tool_result", "content": [nested]}
assert request_contains_image_content([{"role": "user", "content": [nested]}]) is False

Some files were not shown because too many files have changed in this diff Show more