Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_anthropic_fast_mode_cache_and_response_speed

# Conflicts:
#	tests/test_litellm/test_cost_calculator.py
This commit is contained in:
mateo-berri 2026-08-26 12:10:33 -07:00
commit ece187ea24
80 changed files with 2348 additions and 621 deletions

5
.github/mutmut-coverage.rc vendored Normal file
View file

@ -0,0 +1,5 @@
# mutmut's gather_coverage() looks covered lines up by absolute path, so the
# repo's `relative_files = true` makes every lookup miss and mutmut generates
# zero mutants. Point COVERAGE_RCFILE here for mutation runs only.
[run]
relative_files = false

View file

@ -87,11 +87,20 @@ jobs:
run: |
uv pip uninstall pytest-retry || true
# Ends before the job's own deadline so a run that outlasts the budget is
# still followed by the report and upload steps. mutmut saves after every
# mutant result, to mutants/<source path>.meta, so an interrupted run
# still scores the mutants it finished and export-cicd-stats can read
# them; a cancelled job skips those steps and publishes nothing at all.
- name: Run mutmut
timeout-minutes: 300
env:
# Make the mutants/ sandbox win over site-packages on sys.path so the
# trampolined files are imported instead of the installed copy.
PYTHONPATH: ${{ github.workspace }}/mutants
# Without this mutmut finds no covered lines and generates 0 mutants.
# See the file itself for why.
COVERAGE_RCFILE: ${{ github.workspace }}/.github/mutmut-coverage.rc
run: |
set -o pipefail
mkdir -p mutants
@ -130,6 +139,7 @@ jobs:
mutmut-run.log
mutants/mutmut-stats.json
mutants/mutmut-cicd-stats.json
mutants/**/*.meta
mutants/litellm/proxy/management_endpoints/**/*.py
if-no-files-found: warn
retention-days: 14

2
.gitignore vendored
View file

@ -3,6 +3,8 @@
tests/e2e/.fixtures/
.venv-typecheck
.venv_policy_test
.venv-mutmut
mutants/
.env
.claude
CLAUDE.local.md

View file

@ -1,6 +1,6 @@
{
"reportAny": {
"limit": 18505
"limit": 18483
},
"reportArgumentType": {
"limit": 2564
@ -24,7 +24,7 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 5976
"limit": 5960
},
"reportFunctionMemberAccess": {
"limit": 7
@ -57,7 +57,7 @@
"limit": 5659
},
"reportMissingTypeArgument": {
"limit": 15504
"limit": 15484
},
"reportMissingTypeStubs": {
"limit": 40
@ -105,13 +105,13 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38828
"limit": 38808
},
"reportUnknownParameterType": {
"limit": 19847
"limit": 19829
},
"reportUnknownVariableType": {
"limit": 30386
"limit": 30356
},
"reportUnnecessaryCast": {
"limit": 117

View file

@ -12,8 +12,9 @@ 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
from collections.abc import Callable, Mapping
from typing import Final
from urllib.parse import urlsplit, urlunsplit
import redis
import redis.asyncio as async_redis
@ -50,6 +51,7 @@ def _get_redis_kwargs():
include_args: Final = {
"url",
"redis_connect_func",
"credential_provider",
"gcp_service_account",
"gcp_ssl_ca_certs",
"azure_redis_ad_token",
@ -155,7 +157,8 @@ def _get_redis_cluster_kwargs(client=None):
def _get_redis_env_kwarg_mapping():
PREFIX: Final = "REDIS_"
return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs()}
exclude_from_environment: Final = frozenset({"credential_provider"})
return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs() if x not in exclude_from_environment}
def _redis_kwargs_from_environment():
@ -353,6 +356,12 @@ def get_redis_url_from_environment():
return f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}"
def _url_without_userinfo(url: str) -> str:
parts: Final = urlsplit(url)
netloc: Final = parts.netloc.rsplit("@", 1)[-1]
return urlunsplit((parts.scheme, netloc, parts.path, parts.query, parts.fragment))
def _get_redis_client_logic(**env_overrides):
"""
Common functionality across sync + async redis client implementations
@ -410,54 +419,58 @@ def _get_redis_client_logic(**env_overrides):
if _service_name is not None:
redis_kwargs["service_name"] = _service_name
# Handle GCP IAM authentication
_gcp_service_account: Final = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT")
_gcp_ssl_ca_certs: Final = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS")
if _gcp_service_account is not None:
verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.")
redis_kwargs["redis_connect_func"] = create_gcp_iam_redis_connect_func(
service_account=_gcp_service_account, ssl_ca_certs=_gcp_ssl_ca_certs
if redis_kwargs.get("credential_provider") is None:
# Handle GCP IAM authentication
_gcp_service_account: Final = redis_kwargs.get("gcp_service_account") or get_secret_str(
"REDIS_GCP_SERVICE_ACCOUNT"
)
# Store GCP service account in redis_connect_func for async cluster access
redis_kwargs["redis_connect_func"]._gcp_service_account = _gcp_service_account
_gcp_ssl_ca_certs: Final = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS")
# Remove GCP-specific kwargs that shouldn't be passed to Redis client
redis_kwargs.pop("gcp_service_account", None)
redis_kwargs.pop("gcp_ssl_ca_certs", None)
if _gcp_service_account is not None:
verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.")
redis_kwargs["redis_connect_func"] = create_gcp_iam_redis_connect_func(
service_account=_gcp_service_account, ssl_ca_certs=_gcp_ssl_ca_certs
)
# Store GCP service account in redis_connect_func for async cluster access
redis_kwargs["redis_connect_func"]._gcp_service_account = _gcp_service_account
# Only enable SSL if explicitly requested AND SSL CA certs are provided
if _gcp_ssl_ca_certs and redis_kwargs.get("ssl", False):
redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs
# Only enable SSL if explicitly requested AND SSL CA certs are provided
if _gcp_ssl_ca_certs and redis_kwargs.get("ssl", False):
redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs
# Handle Azure AD authentication (after GCP IAM block)
_azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN")
# Handle Azure AD authentication (after GCP IAM block)
_azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN")
_azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true"
_azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true"
if _azure_ad_enabled and _gcp_service_account is not None:
verbose_logger.warning(
"Both GCP IAM (gcp_service_account) and Azure AD (azure_redis_ad_token) are configured for Redis. "
"Using GCP IAM. Remove one to avoid misconfiguration."
)
if _azure_ad_enabled and _gcp_service_account is not None:
verbose_logger.warning(
"Both GCP IAM (gcp_service_account) and Azure AD (azure_redis_ad_token) are configured for Redis. "
"Using GCP IAM. Remove one to avoid misconfiguration."
)
if _azure_ad_enabled and _gcp_service_account is None:
_azure_client_id: Final = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID")
_azure_tenant_id: Final = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID")
_azure_client_secret: Final = redis_kwargs.get("azure_client_secret") or get_secret_str("AZURE_CLIENT_SECRET")
if _azure_ad_enabled and _gcp_service_account is None:
_azure_client_id: Final = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID")
_azure_tenant_id: Final = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID")
_azure_client_secret: Final = redis_kwargs.get("azure_client_secret") or get_secret_str(
"AZURE_CLIENT_SECRET"
)
verbose_logger.debug("Setting up Azure AD authentication for Redis.")
redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func(
azure_client_id=_azure_client_id,
azure_tenant_id=_azure_tenant_id,
azure_client_secret=_azure_client_secret,
)
# Marker for async paths to detect Azure AD auth. The live credential
# object is attached separately as `_azure_credential` by
# `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret
# are intentionally NOT exposed on the function to avoid leaking
# credentials via inspection or logging.
redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True
verbose_logger.debug("Setting up Azure AD authentication for Redis.")
redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func(
azure_client_id=_azure_client_id,
azure_tenant_id=_azure_tenant_id,
azure_client_secret=_azure_client_secret,
)
# Marker for async paths to detect Azure AD auth. The live credential
# object is attached separately as `_azure_credential` by
# `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret
# are intentionally NOT exposed on the function to avoid leaking
# credentials via inspection or logging.
redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True
redis_kwargs.pop("gcp_service_account", None)
redis_kwargs.pop("gcp_ssl_ca_certs", None)
# Always remove Azure-specific kwargs that shouldn't be passed to Redis client
redis_kwargs.pop("azure_redis_ad_token", None)
@ -465,6 +478,13 @@ def _get_redis_client_logic(**env_overrides):
redis_kwargs.pop("azure_tenant_id", None)
redis_kwargs.pop("azure_client_secret", None)
if redis_kwargs.get("credential_provider") is not None:
redis_kwargs.pop("redis_connect_func", None)
redis_kwargs.pop("username", None)
redis_kwargs.pop("password", None)
if redis_kwargs.get("url") is not None:
redis_kwargs["url"] = _url_without_userinfo(redis_kwargs["url"])
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
# Only strip host/port/db/password when not routing to a cluster.
# When startup_nodes is also present the cluster path takes priority and
@ -532,8 +552,7 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
service_name: Final = redis_kwargs.get("service_name")
connection_kwargs: Final = _get_redis_sentinel_connection_kwargs(redis_kwargs)
connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT)
sentinel_kwargs: Final = dict(connection_kwargs)
sentinel_kwargs["password"] = sentinel_password
sentinel_kwargs: Final = _sentinel_auth_kwargs(connection_kwargs, sentinel_password)
if not sentinel_nodes or not service_name:
raise ValueError("Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel.")
@ -605,7 +624,12 @@ def _async_credential_provider(redis_connect_func: object | None) -> CredentialP
def _async_auth_kwargs(redis_kwargs: dict) -> dict:
"""Swaps a connect func an async path cannot run for the equivalent credential provider,
which supersedes any static username or password redis-py would otherwise reject it with."""
credential_provider: Final = _async_credential_provider(redis_kwargs.get("redis_connect_func"))
explicit_provider: Final = redis_kwargs.get("credential_provider")
credential_provider: Final = (
explicit_provider
if explicit_provider is not None
else _async_credential_provider(redis_kwargs.get("redis_connect_func"))
)
if credential_provider is None:
return redis_kwargs
@ -738,8 +762,20 @@ def get_redis_connection_pool(
return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs)
def _redis_kwargs_for_logging(redis_kwargs: Mapping[str, object]) -> Mapping[str, object]:
return {
key: "<credential provider>"
if key == "credential_provider" and value is not None
else "<redis connect function>"
if key == "redis_connect_func" and value is not None
else value
for key, value in redis_kwargs.items()
}
def _pretty_print_redis_config(redis_kwargs: dict) -> None:
"""Pretty print the Redis configuration using rich with sensitive data masking"""
redis_kwargs_for_logging: Final = _redis_kwargs_for_logging(redis_kwargs)
try:
import logging
@ -757,7 +793,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None:
masker = SensitiveDataMasker()
# Mask sensitive data in redis_kwargs
masked_redis_kwargs = masker.mask_dict(redis_kwargs)
masked_redis_kwargs = masker.mask_dict(redis_kwargs_for_logging)
# Create main panel title
title: Final = Text("Redis Configuration", style="bold blue")
@ -820,7 +856,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None:
except ImportError:
# Fallback to simple logging if rich is not available
masker = SensitiveDataMasker()
masked_redis_kwargs = masker.mask_dict(redis_kwargs)
masked_redis_kwargs = masker.mask_dict(redis_kwargs_for_logging)
verbose_logger.info("Redis configuration: %s", masked_redis_kwargs)
except Exception as e:
verbose_logger.error("Error pretty printing Redis configuration: %s", e)

View file

@ -551,7 +551,7 @@ def _get_batch_job_usage_from_response_body(
return usage
def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[str, Any]) -> dict:
def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[str, Any]) -> Mapping[str, Any]:
"""
Get the ``result`` object from a line of an Anthropic message batch results JSONL file.
@ -563,7 +563,7 @@ def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[st
def _get_response_from_batch_job_output_file(
batch_job_output_file: Mapping[str, Any], custom_llm_provider: str = "openai"
) -> Any:
) -> Mapping[str, Any]:
"""
Get the response from the batch job output file
"""

View file

@ -175,6 +175,10 @@ _RedisCallResult = TypeVar("_RedisCallResult")
_swallowed_redis_failures: Final[ContextVar[int]] = ContextVar("litellm_swallowed_redis_failures", default=0)
def _opaque_kwarg_key(value: object) -> str:
return f"{type(value).__name__}-{id(value)}"
@functools.lru_cache(maxsize=1)
def _redis_health_error_types() -> tuple[type, ...]:
"""Exception types that mean the Redis backend itself is unhealthy.
@ -399,10 +403,9 @@ class RedisCache(BaseCache):
Generate a cache key for the async Redis client based on connection parameters.
This ensures different Redis configurations use different cached clients.
"""
# Create a stable representation of redis_kwargs for hashing
# Sort keys to ensure consistent hash regardless of parameter order
sorted_kwargs: Final = sorted(self.redis_kwargs.items())
kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True)
kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True, default=_opaque_kwarg_key)
kwargs_hash: Final = hashlib.sha256(kwargs_str.encode()).hexdigest()[:16]
return f"async-redis-client-{kwargs_hash}"
@ -1384,10 +1387,10 @@ class RedisCache(BaseCache):
dict: {"status": "success" | "failed", "message": str, "error": Optional[str]}
"""
try:
import redis.asyncio as redis_async
from .._redis import get_redis_async_client
# Create a fresh Redis client with current settings
redis_client: Final = redis_async.Redis(**self.redis_kwargs)
redis_client: Final = get_redis_async_client(**self.redis_kwargs)
# Test the connection
ping_result: Final = await redis_client.ping()

View file

@ -64,22 +64,9 @@ class RedisClusterCache(RedisCache):
dict: {"status": "success" | "failed", "message": str, "error": Optional[str]}
"""
try:
import redis.asyncio as redis_async
from redis.cluster import ClusterNode
from .._redis import get_redis_async_client
# Create ClusterNode objects from startup_nodes
cluster_kwargs: Final = self.redis_kwargs.copy()
startup_nodes: Final = cluster_kwargs.pop("startup_nodes", [])
new_startup_nodes: Final[list[ClusterNode]] = []
for item in startup_nodes:
new_startup_nodes.append(ClusterNode(**item))
# Create a fresh Redis Cluster client with current settings
redis_client: Final = redis_async.RedisCluster(
startup_nodes=new_startup_nodes,
**cluster_kwargs,
)
redis_client: Final = get_redis_async_client(**self.redis_kwargs)
# Test the connection
ping_result: Final = await redis_client.ping()

View file

@ -295,7 +295,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
async def async_post_call_failure_deployment_hook(
self,
request_data: Mapping[str, Any],
request_data: Mapping[str, object],
exception: Exception,
call_type: CallTypes | None,
fallback_depth: int | None = None,

View file

@ -15,7 +15,7 @@ from collections import OrderedDict
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Any, Final, TypeAlias
from typing import Final, TypeAlias
from urllib.parse import quote
from opentelemetry.sdk.trace import TracerProvider
@ -32,6 +32,7 @@ from litellm.integrations.otel.presets import (
dynamic_otlp_headers,
project_routing_headers,
)
from litellm.types.utils import StandardCallbackDynamicParams
# Exporter kinds that ignore headers — never rewritten with dynamic credentials.
_NON_OTLP_KINDS: Final = ("console", "in_memory", "inmemory", "memory")
@ -166,7 +167,7 @@ class TenantTracerCache:
def route_for(
self,
default: Tracer,
dynamic_params: Any,
dynamic_params: StandardCallbackDynamicParams | None,
auth_metadata: Mapping[str, str] | None = None,
) -> TenantRoute:
"""Return the tracer (and trace-detachment flag) for this request.

View file

@ -2495,12 +2495,12 @@ class PrometheusLogger(CustomLogger):
return None
def _get_user_email() -> str | None:
val = _metadata.get("user_api_key_user_email")
if val is not None:
return val
val = _litellm_params_metadata.get("user_api_key_user_email")
if val is not None:
return val
from_metadata: Final = _metadata.get("user_api_key_user_email")
if from_metadata is not None:
return from_metadata
from_params: Final = _litellm_params_metadata.get("user_api_key_user_email")
if from_params is not None:
return from_params
if user_api_key_auth is not None:
return self._safe_get(user_api_key_auth, "user_email")
return None
@ -3576,7 +3576,9 @@ class PrometheusLogger(CustomLogger):
except Exception as e:
verbose_logger.exception("Error initializing user/team count metrics: %s", e)
async def _set_key_list_budget_metrics(self, keys: list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken]):
async def _set_key_list_budget_metrics(
self, keys: list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken]
) -> None:
"""Helper function to set budget metrics for a list of keys"""
for key in keys:
if isinstance(key, UserAPIKeyAuth):

View file

@ -2,6 +2,7 @@
Helper functions for health check calls.
"""
import base64
from collections.abc import Callable
from typing import TYPE_CHECKING, Final, Literal
@ -13,6 +14,14 @@ if TYPE_CHECKING:
# Minimal PDF for health checks - base64 encoded 1-page PDF with just "test"
TEST_PDF_URL = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="
# Minimal image for health checks - base64 encoded 512x512 solid-gray PNG
TEST_IMAGE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAAgAAAAIACAIAAAB7GkOtAAAFlklEQVR42u3VMQEAAAzCMKQjHQ97l0jo0xSAlyIBgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGACAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGACAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGACAAUgAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGACAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGACAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgBgAAAYAAAGAIABAGAAABgAAAYAgAEAYAAAGAAABgCAAQBgAAAYAAAGAIABAGAAABgAADcDrctaAb6XeXAAAAAASUVORK5CYII="
def get_image_file_for_health_check() -> bytes:
"""Return the image used for health checks."""
return base64.b64decode(TEST_IMAGE_BASE64)
class HealthCheckHelpers:
@staticmethod
@ -127,6 +136,7 @@ class HealthCheckHelpers:
"audio_speech",
"audio_transcription",
"image_generation",
"image_edit",
"video_generation",
"rerank",
"realtime",
@ -185,6 +195,11 @@ class HealthCheckHelpers:
**_filter_model_params(model_params=model_params),
prompt=prompt,
),
"image_edit": lambda: litellm.aimage_edit(
**_filter_model_params(model_params=model_params),
image=get_image_file_for_health_check(),
prompt=prompt or "test",
),
"video_generation": lambda: litellm.avideo_generation(
**_filter_model_params(model_params=model_params),
prompt=prompt or "test video generation",

View file

@ -1,6 +1,6 @@
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Any
from typing import Any, Final
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
@ -39,7 +39,7 @@ class TranscriptionUsageObjectTransformation:
return None
_INTERACTIONS_MODALITY_FIELDS: Mapping[str, str] = MappingProxyType(
_INTERACTIONS_MODALITY_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{
"text": "text_tokens",
"audio": "audio_tokens",
@ -59,7 +59,7 @@ def _token_count(value: object) -> int:
def _modality_token_sums(entries: Sequence[Mapping[str, Any]]) -> Mapping[str, int]:
fields = frozenset(field for entry in entries if (field := _modality_field(entry)) is not None)
fields: Final = frozenset(field for entry in entries if (field := _modality_field(entry)) is not None)
return MappingProxyType(
{
field: sum(_token_count(entry.get("tokens")) for entry in entries if _modality_field(entry) == field)
@ -69,10 +69,13 @@ def _modality_token_sums(entries: Sequence[Mapping[str, Any]]) -> Mapping[str, i
def _google_search_query_count(usage_object: Mapping[str, Any]) -> int:
entries: Final = usage_object.get("grounding_tool_count")
if not isinstance(entries, Sequence):
return 0
return sum(
_token_count(entry.get("count"))
for entry in tuple(usage_object.get("grounding_tool_count") or ())
if isinstance(entry, Mapping) and entry.get("type") == "google_search" # pyright: ignore[reportUnnecessaryIsInstance] # provider JSON, not the empty tuple inferred from `or ()`
for entry in entries
if isinstance(entry, Mapping) and entry.get("type") == "google_search"
)
@ -112,30 +115,30 @@ class InteractionsUsageObjectTransformation:
@staticmethod
def transform_interactions_usage_object(usage_object: Mapping[str, Any]) -> Usage:
input_entries = tuple(usage_object.get("input_tokens_by_modality") or ()) + tuple(
input_entries: Final = tuple(usage_object.get("input_tokens_by_modality") or ()) + tuple(
usage_object.get("tool_use_tokens_by_modality") or ()
)
cached_sums = _modality_token_sums(tuple(usage_object.get("cached_tokens_by_modality") or ()))
output_sums = _modality_token_sums(tuple(usage_object.get("output_tokens_by_modality") or ()))
cached_sums: Final = _modality_token_sums(tuple(usage_object.get("cached_tokens_by_modality") or ()))
output_sums: Final = _modality_token_sums(tuple(usage_object.get("output_tokens_by_modality") or ()))
total_cached_tokens = _token_count(usage_object.get("total_cached_tokens"))
input_sums = _subtract_cached_from_input(
total_cached_tokens: Final = _token_count(usage_object.get("total_cached_tokens"))
input_sums: Final = _subtract_cached_from_input(
input_sums=_modality_token_sums(input_entries),
cached_sums=cached_sums,
total_cached_tokens=total_cached_tokens,
)
reasoning_tokens = _token_count(usage_object.get("total_reasoning_tokens")) or _token_count(
reasoning_tokens: Final = _token_count(usage_object.get("total_reasoning_tokens")) or _token_count(
usage_object.get("total_thought_tokens")
)
prompt_tokens = _token_count(usage_object.get("total_input_tokens")) + _token_count(
prompt_tokens: Final = _token_count(usage_object.get("total_input_tokens")) + _token_count(
usage_object.get("total_tool_use_tokens")
)
completion_tokens = _token_count(usage_object.get("total_output_tokens")) + reasoning_tokens
total_tokens = _token_count(usage_object.get("total_tokens")) or (prompt_tokens + completion_tokens)
completion_tokens: Final = _token_count(usage_object.get("total_output_tokens")) + reasoning_tokens
total_tokens: Final = _token_count(usage_object.get("total_tokens")) or (prompt_tokens + completion_tokens)
web_search_requests = _google_search_query_count(usage_object)
prompt_tokens_details = (
web_search_requests: Final = _google_search_query_count(usage_object)
prompt_tokens_details: Final = (
PromptTokensDetailsWrapper(
cached_tokens=total_cached_tokens or None,
web_search_requests=web_search_requests or None,
@ -144,7 +147,7 @@ class InteractionsUsageObjectTransformation:
if input_sums or total_cached_tokens or web_search_requests
else None
)
completion_tokens_details = (
completion_tokens_details: Final = (
CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens or None,
**output_sums,

View file

@ -511,9 +511,6 @@ def update_messages_with_model_file_ids(
if "llm_output_file_id," in unified_file_id:
provider_file_id = unified_file_id.split("llm_output_file_id,")[1].split(";")[0]
if not provider_file_id and is_model_embedded_id(file_id):
# `litellm:<raw_id>;model,<m>` encoding from the
# x-litellm-model upload path. Strip the wrapper
# so the provider sees its own ID.
provider_file_id = get_original_file_id(file_id)
file_object_file_field["file_id"] = provider_file_id or file_id
if format:
@ -588,9 +585,6 @@ def update_responses_input_with_model_file_ids(
updated_content_item["file_id"] = provider_file_id
updated_content.append(updated_content_item)
elif is_model_embedded_id(file_id):
# `litellm:<raw_id>;model,<m>` encoding from the
# x-litellm-model upload path. Strip the wrapper
# so the provider sees its own ID.
updated_content_item = content_item.copy()
updated_content_item["file_id"] = get_original_file_id(file_id)
updated_content.append(updated_content_item)

View file

@ -10,6 +10,7 @@
import asyncio
import copy
import inspect
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import litellm
@ -191,7 +192,7 @@ def _redact_standard_logging_object(model_call_details: dict):
standard_logging_object["response"] = {"text": redacted_str}
def _redact_tool_calls_dict(message: dict) -> None:
def _redact_tool_calls_dict(message: Mapping[str, object]) -> None:
"""Redact tool call / function_call arguments in a dict-form message or delta."""
tool_calls: Final = message.get("tool_calls")
if isinstance(tool_calls, list):

View file

@ -4,7 +4,7 @@ This file contains common utils for anthropic calls.
import copy
import re
from collections.abc import Mapping, Sequence
from collections.abc import Mapping, MutableMapping, Sequence
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Any, Final, Literal
@ -93,8 +93,8 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup
"""
Handle Anthropic OAuth token detection and header setup.
If an OAuth token is detected in the Authorization header, extracts it
and sets the required OAuth headers.
If an OAuth token is detected in the Authorization header (any casing),
extracts it and sets the required OAuth headers.
Args:
headers: Request headers dict
@ -104,16 +104,21 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup
Tuple of (updated headers, api_key)
"""
# Check Authorization header (passthrough / forwarded requests)
auth_header: Final = headers.get("authorization", "")
if auth_header and auth_header.startswith(f"Bearer {ANTHROPIC_OAUTH_TOKEN_PREFIX}"):
api_key = auth_header.replace("Bearer ", "")
headers.pop("x-api-key", None)
auth_header: Final = next((value for name, value in headers.items() if name.lower() == "authorization"), "")
if auth_header.startswith(f"Bearer {ANTHROPIC_OAUTH_TOKEN_PREFIX}"):
api_key = auth_header.removeprefix("Bearer ")
for name in tuple(
header_name for header_name in headers if header_name.lower() in ("x-api-key", "authorization")
):
headers.pop(name)
headers["authorization"] = auth_header
headers["anthropic-beta"] = _merge_beta_headers(headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER)
headers["anthropic-dangerous-direct-browser-access"] = "true"
return headers, api_key
# Check api_key directly (standard chat/completion flow)
if api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX):
headers.pop("x-api-key", None)
for name in tuple(header_name for header_name in headers if header_name.lower() == "x-api-key"):
headers.pop(name)
headers["authorization"] = f"Bearer {api_key}"
headers["anthropic-beta"] = _merge_beta_headers(headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER)
headers["anthropic-dangerous-direct-browser-access"] = "true"
@ -468,7 +473,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
@staticmethod
def maybe_drop_disabled_thinking(
model: str,
optional_params: dict, # mutable-ok: in-place out-param, same contract as AnthropicConfig._maybe_drop_speed_param
optional_params: MutableMapping[str, object], # mutable-ok: in-place out-param, as in _maybe_drop_speed_param
custom_llm_provider: str,
) -> None:
"""Omit ``thinking={'type': 'disabled'}`` for always-on-thinking models

View file

@ -352,8 +352,8 @@ async def _check_summary_model_budget(
)
return False
user_model_max_budget: Final = getattr(user_api_key_auth, "user_model_max_budget", None)
user_id: Final = getattr(user_api_key_auth, "user_id", None)
user_model_max_budget: Final = user_api_key_auth.user_model_max_budget
user_id: Final = user_api_key_auth.user_id
if isinstance(user_model_max_budget, dict) and user_model_max_budget and user_id is not None:
try:
await model_max_budget_limiter.is_user_within_model_budget(

View file

@ -8,6 +8,7 @@ from litellm.constants import (
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
)
from litellm.exceptions import AuthenticationError
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import verbose_logger
from litellm.llms.base_llm.anthropic_messages.transformation import (
@ -307,10 +308,20 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
# Check for Anthropic OAuth token in Authorization header
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
if "x-api-key" not in headers and "authorization" not in headers:
header_names: Final = frozenset(name.lower() for name in headers)
if "x-api-key" not in header_names and "authorization" not in header_names:
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key)
if auth_header is not None:
headers.update(auth_header)
if auth_header is None:
raise AuthenticationError(
message=(
"Missing Anthropic API Key - A call is being made to anthropic but no key is set "
"either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` "
"or `ANTHROPIC_AUTH_TOKEN` in your environment vars"
),
llm_provider=self._resolved_provider,
model=model,
)
headers.update(auth_header)
if "anthropic-version" not in headers:
headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION
if "content-type" not in headers:

View file

@ -4,6 +4,8 @@ This file contains the calling Azure OpenAI's `/openai/realtime` endpoint.
This requires websockets, and is currently only supported on LiteLLM Proxy.
"""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final, cast
from litellm._logging import _redact_string, verbose_proxy_logger
@ -30,6 +32,21 @@ async def forward_messages(client_ws: Any, backend_ws: Any):
class AzureOpenAIRealtime(AzureChatCompletion):
@staticmethod
def get_auth_headers(api_key: str | None, azure_ad_token: str | None) -> Mapping[str, str]:
"""
Build the websocket handshake auth headers, preferring a static api-key and falling back to
an Azure AD (Entra ID) bearer token. Never sends both.
"""
if api_key:
return MappingProxyType({"api-key": api_key})
if azure_ad_token:
return MappingProxyType({"Authorization": f"Bearer {azure_ad_token}"})
raise ValueError(
"Missing Azure credentials for the realtime endpoint. Set an api_key, or configure Azure AD auth "
"(azure_ad_token, tenant_id/client_id/client_secret, or a managed identity)"
)
def _construct_url(
self,
api_base: str,
@ -117,13 +134,13 @@ class AzureOpenAIRealtime(AzureChatCompletion):
query_params=query_params,
)
auth_headers: Final = self.get_auth_headers(api_key=api_key, azure_ad_token=azure_ad_token)
try:
ssl_context: Final = get_shared_realtime_ssl_context()
async with websockets.connect(
url,
additional_headers={
"api-key": api_key,
},
additional_headers=auth_headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_context,
) as backend_ws:

View file

@ -418,12 +418,16 @@ class AmazonConverseConfig(BaseConfig):
Handle the reasoning_effort parameter based on the model type.
- GPT-OSS models: passed through unchanged via additionalModelRequestFields.
- OpenAI GPT-5.x models: mapped to ``reasoning.effort`` via additionalModelRequestFields.
- Nova 2 models: transformed to reasoningConfig.
- Anthropic models: mapped to ``thinking`` (and ``output_config.effort`` on
adaptive Claude 4.6 / 4.7).
"""
if "gpt-oss" in model:
optional_params["reasoning_effort"] = reasoning_effort
elif "openai.gpt-5" in model:
reasoning: Final[BedrockConverseGptReasoningEffortBlock] = {"effort": reasoning_effort}
optional_params["reasoning"] = reasoning
elif self._is_nova_2_model(model):
reasoning_config: Final = self._transform_reasoning_effort_to_reasoning_config(reasoning_effort)
optional_params.update(reasoning_config)
@ -555,7 +559,7 @@ class AmazonConverseConfig(BaseConfig):
# only anthropic and mistral support tool choice config. otherwise (E.g. cohere) will fail the call - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html
supported_params.append("tool_choice")
if "gpt-oss" in model:
if "gpt-oss" in model or "openai.gpt-5" in model or "openai.gpt-5" in base_model:
supported_params.append("reasoning_effort")
elif self._is_nova_2_model(model):
# Nova 2 models support reasoning_effort (transformed to reasoningConfig)
@ -903,7 +907,7 @@ class AmazonConverseConfig(BaseConfig):
optional_params["_parallel_tool_use_config"] = {
"tool_choice": {"type": "auto", "disable_parallel_tool_use": not value}
}
if param == "thinking":
if param == "thinking" and "openai.gpt-5" not in model:
if (
isinstance(value, dict)
and value.get("type") == "adaptive"

View file

@ -243,7 +243,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
return remaining_input, cls._filter_unsupported_tools(hoisted_tools)
@staticmethod
def _agent_message_text(item: "Mapping[str, Any]") -> str:
def _agent_message_text(item: "Mapping[str, object]") -> str:
content: Final = item.get("content")
if not isinstance(content, list):
return ""
@ -254,7 +254,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
)
@classmethod
def _normalize_agent_message_item(cls, item: "Mapping[str, Any]") -> "_RewrittenAssistantMessageItem | None":
def _normalize_agent_message_item(cls, item: "Mapping[str, object]") -> "_RewrittenAssistantMessageItem | None":
text: Final = cls._agent_message_text(item)
if not text:
return None
@ -266,7 +266,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
return rewritten
@staticmethod
def _normalize_context_compaction_item(item: "Mapping[str, Any]") -> "_RewrittenCompactionItem | None":
def _normalize_context_compaction_item(item: "Mapping[str, object]") -> "_RewrittenCompactionItem | None":
encrypted_content: Final = item.get("encrypted_content")
if not isinstance(encrypted_content, str) or not encrypted_content:
return None
@ -274,7 +274,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
return rewritten
@staticmethod
def _normalize_local_shell_call_item(item: "Mapping[str, Any]") -> "_RewrittenFunctionCallItem | None":
def _normalize_local_shell_call_item(item: "Mapping[str, object]") -> "_RewrittenFunctionCallItem | None":
call_id: Final = item.get("call_id")
if not isinstance(call_id, str) or not call_id:
return None

View file

@ -5620,10 +5620,9 @@ class BaseLLMHTTPHandler:
kwargs=hook_kwargs,
)
except Exception as e:
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in async_should_run_agentic_loop [call_id=%s model=%s]: %s",
_call_id,
logging_obj.litellm_call_id,
model,
str(e),
)
@ -5645,10 +5644,9 @@ class BaseLLMHTTPHandler:
except AgenticLoopSafetyError as e:
if not self._can_replace_turn_with_terminal_response(stream, api_surface):
raise
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
verbose_logger.warning(
"LiteLLM.AgenticLoopRefused: ending turn [call_id=%s model=%s]: %s",
_call_id,
logging_obj.litellm_call_id,
model,
str(e),
)

View file

@ -12282,7 +12282,7 @@
},
"claude-3-haiku-20240307": {
"cache_creation_input_token_cost": 3e-07,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr": 5e-07,
"cache_read_input_token_cost": 3e-08,
"deprecation_date": "2026-04-20",
"input_cost_per_token": 2.5e-07,
@ -12301,7 +12301,7 @@
},
"claude-3-opus-20240229": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr": 3e-05,
"cache_read_input_token_cost": 1.5e-06,
"deprecation_date": "2026-01-05",
"input_cost_per_token": 1.5e-05,
@ -12515,7 +12515,10 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_output_config": true,
"prompt_cache_min_tokens": 1024
"prompt_cache_min_tokens": 1024,
"provider_specific_entry": {
"us": 1.1
}
},
"claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
@ -49469,6 +49472,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"global.openai.gpt-5.6-sol": {
@ -49494,6 +49498,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"us.openai.gpt-5.6-terra": {
@ -49519,6 +49524,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"global.openai.gpt-5.6-terra": {
@ -49544,6 +49550,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"us.openai.gpt-5.6-luna": {
@ -49569,6 +49576,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"global.openai.gpt-5.6-luna": {
@ -49594,6 +49602,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"bedrock_mantle/openai.gpt-5.5": {
@ -49601,7 +49610,7 @@
"cache_read_input_token_cost": 5.5e-07,
"output_cost_per_token": 3.3e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 272000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
@ -49628,7 +49637,7 @@
"cache_read_input_token_cost": 2.75e-07,
"output_cost_per_token": 1.65e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 272000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
@ -50821,7 +50830,10 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true
"supports_native_structured_output": true,
"provider_specific_entry": {
"us": 1.1
}
},
"claude-mythos-preview": {
"cache_creation_input_token_cost": 1.25e-05,
@ -50856,7 +50868,10 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true
"supports_native_structured_output": true,
"provider_specific_entry": {
"us": 1.1
}
},
"gemini/gemini-robotics-er-2-streaming-preview": {
"input_cost_per_audio_token": 2e-06,

View file

@ -306,15 +306,15 @@ _UpstreamGrantRejection = Literal["no_access_token", "expired_lifetime"]
- ``expired_lifetime``: the response reports a parseable, non-positive ``expires_in``, i.e. an upstream
token that is already dead, so sealing it would forward a bearer the edge cannot use
An absent or unparseable ``expires_in`` is NOT a rejection; the lifetime is merely unknown and the
envelope caps it, the by-design behaviour for an upstream that omits the field."""
envelope uses its fallback lifetime, the by-design behaviour for an upstream that omits the field."""
def _classify_upstream_lifetime(raw_expires_in: object) -> "int | Literal['unspecified', 'expired']":
"""Classify an upstream ``expires_in`` into a positive number of seconds, ``"unspecified"`` (absent
or unparseable, so the envelope caps it), or ``"expired"`` (a non-positive value the upstream reports
or unparseable, so the envelope uses its fallback), or ``"expired"`` (a non-positive value the upstream reports
as already elapsed). Telling "we do not know the lifetime" apart from "the upstream says it is
already dead" is what stops an explicitly-expired token from silently receiving the envelope's 1h
cap. The expired decision is made on the parsed numeric value, not on ``int(...)`` of it, so a
already dead" is what stops an explicitly-expired token from silently receiving the envelope's
one-hour fallback. The expired decision is made on the parsed numeric value, not on ``int(...)`` of it, so a
positive sub-second lifetime in ``(0, 1)`` is not truncated to ``0`` and misread as elapsed; the
envelope works in whole seconds, so such a lifetime clamps up to its 1s floor. ``bool`` is excluded
(an ``int`` subclass but never a real lifetime), and the conversions can raise on ``NaN`` /
@ -335,7 +335,7 @@ def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenG
"""Validate an upstream OAuth token response into a typed grant, or say why it cannot back an
envelope. Each field is isinstance-checked so nothing untyped from ``response.json()`` reaches the
grant. ``expires_in`` is read three ways (see :func:`_classify_upstream_lifetime`): an unknown
lifetime leaves the grant ``expires_in`` ``None`` for the envelope to cap, a positive value is
lifetime leaves the grant ``expires_in`` ``None`` for the envelope fallback, a positive value is
honoured, and an explicit already-elapsed value is a rejection rather than a silent fall-through to
the cap."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
@ -357,8 +357,8 @@ def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenG
token_type=token_type if isinstance(token_type, str) and token_type else "Bearer",
# The upstream refresh_token is deliberately NOT sealed: the edge never consumes it (it forwards
# only token_type + access_token), so it would be dead weight embedding a long-lived upstream
# credential in the client-held bearer, and it enlarges the envelope. Refresh support is a
# follow-up (a dedicated refresh-envelope); the client re-runs authorization_code at the cap.
# credential in the client-held bearer, and it enlarges the envelope. The dedicated refresh
# envelope carries that credential separately.
refresh_token=None,
scope=scope if isinstance(scope, str) and scope else None,
expires_in=lifetime if isinstance(lifetime, int) else None,
@ -387,6 +387,7 @@ _BridgeMintError = Literal[
"not_configured",
"no_upstream_token",
"upstream_token_expired",
"upstream_lifetime_unrepresentable",
"too_large",
]
@ -456,6 +457,12 @@ def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
"server_error",
"the upstream token response reports an already-expired lifetime",
)
case "upstream_lifetime_unrepresentable":
status, code, desc = (
502,
"server_error",
"the upstream token response reports an unrepresentable lifetime",
)
case "too_large":
status, code, desc = (
502,
@ -619,6 +626,7 @@ def _finish_bridge_mint(
build_bridge_token_response,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
EnvelopeLifetimeUnrepresentable,
SealedEnvelope,
UpstreamTokenGrant,
)
@ -627,6 +635,8 @@ def _finish_bridge_mint(
if not isinstance(grant, UpstreamTokenGrant):
return _upstream_rejection_to_mint_error(grant)
sealed: Final = build_bridge_token_response(ready.identity, grant, ready.keys, now)
if isinstance(sealed, EnvelopeLifetimeUnrepresentable):
return "upstream_lifetime_unrepresentable"
if not isinstance(sealed, SealedEnvelope):
return "too_large"
# Report expires_in from the JWT's own second-truncated exp, rounding the elapsed portion up, so the

View file

@ -92,7 +92,7 @@ def build_bridge_token_response(
The producer mirror of :func:`resolve_bridge_envelope`: a thin, pure wrapper over
:func:`mint_envelope` that returns the sealed envelope, or the mint error as a value
(an oversized grant) for the caller to map onto an OAuth error response.
for the caller to map onto an OAuth error response.
"""
return mint_envelope(identity, grant, keys, now)

View file

@ -19,17 +19,16 @@ in plaintext anywhere in the envelope.
Failures are values: :func:`open_envelope` returns one of the frozen
``EnvelopeOpenError`` variants (discriminated on ``tag``) for invalid, expired,
tampered, or undecryptable input, and :func:`mint_envelope` returns
``EnvelopeTooLarge`` for oversized grants. Error values carry tags and sizes only,
never token material.
tampered, or undecryptable input, and :func:`mint_envelope` returns a typed error
for oversized grants or an unrepresentable provider lifetime. Error values carry
tags and metadata only, never token material.
The pydantic input models reject programmer errors at construction (e.g. a
non-positive ``expires_in`` or an empty required field). :func:`open_envelope` is
additionally total over hostile, attacker-controlled input: it never raises, only
returns an ``EnvelopeOpenError``. :func:`mint_envelope` operates on a
gateway-supplied grant (an upstream IdP's UTF-8 JSON token response), so it does not
defend against non-UTF-8 field content that cannot survive JSON parsing; its only
value-typed failure is ``EnvelopeTooLarge``.
defend against non-UTF-8 field content that cannot survive JSON parsing.
"""
from __future__ import annotations
@ -57,10 +56,11 @@ ENVELOPE_ISSUER: Final = "litellm-mcp-bridge"
"""``iss`` claim stamped into every envelope and required back on open."""
MAX_ENVELOPE_TTL_SECONDS: Final = 3600
"""Hard ceiling on ACCESS envelope lifetime. ``exp`` is ``min(upstream expires_in, this cap)``
(the cap alone when the upstream omits ``expires_in``), matching the 1h lifetime of the
BYOK session bearer this module's signing approach is borrowed from: a client-held
credential should never outlive a bounded window even when the upstream token does."""
"""Fallback ACCESS envelope lifetime when the upstream omits ``expires_in``.
The historical exported name is retained for import compatibility. When the upstream
reports a positive lifetime, the envelope matches it so a renewal does not consume a
still-valid provider refresh grant."""
MAX_REFRESH_ENVELOPE_TTL_SECONDS: Final = 1209600
"""Hard ceiling on REFRESH envelope lifetime (14 days). A refresh envelope only renews the short-lived
@ -202,7 +202,15 @@ class EnvelopeTooLarge(BaseModel):
max_bytes: int
EnvelopeMintError: TypeAlias = EnvelopeTooLarge
class EnvelopeLifetimeUnrepresentable(BaseModel):
"""A positive provider lifetime cannot be represented as a Python datetime."""
model_config = ConfigDict(frozen=True)
tag: Literal["envelope_lifetime_unrepresentable"] = "envelope_lifetime_unrepresentable"
expires_in: int
EnvelopeMintError: TypeAlias = EnvelopeTooLarge | EnvelopeLifetimeUnrepresentable
class NotAnEnvelope(BaseModel):
@ -307,11 +315,17 @@ def mint_envelope(
) -> SealedEnvelope | EnvelopeMintError:
"""Seal ``grant`` for ``identity`` into a client-held envelope.
``exp`` is ``min(grant.expires_in, MAX_ENVELOPE_TTL_SECONDS)`` seconds from ``now``
(the cap alone when ``expires_in`` is absent). Returns ``EnvelopeTooLarge`` when the
serialized envelope exceeds ``MAX_ENVELOPE_BYTES``.
``exp`` is ``grant.expires_in`` seconds from ``now`` when the upstream reports a
lifetime, or ``MAX_ENVELOPE_TTL_SECONDS`` when it does not. Returns
``EnvelopeLifetimeUnrepresentable`` when that positive lifetime cannot be represented
as a Python datetime, or ``EnvelopeTooLarge`` when the serialized envelope exceeds
``MAX_ENVELOPE_BYTES``.
"""
expires_at: Final = now + timedelta(seconds=_envelope_ttl_seconds(grant.expires_in))
ttl_seconds: Final = _envelope_ttl_seconds(grant.expires_in)
try:
expires_at: Final = now + timedelta(seconds=ttl_seconds)
except OverflowError:
return EnvelopeLifetimeUnrepresentable(expires_in=ttl_seconds)
return _seal(
kind="access",
prefix=ENVELOPE_PREFIX,
@ -457,7 +471,7 @@ def _open_claims(
def _envelope_ttl_seconds(upstream_expires_in: int | None) -> int:
if upstream_expires_in is None:
return MAX_ENVELOPE_TTL_SECONDS
return min(upstream_expires_in, MAX_ENVELOPE_TTL_SECONDS)
return upstream_expires_in
def _refresh_ttl_seconds(upstream_refresh_expires_in: int | None) -> int:

View file

@ -54,9 +54,9 @@ the envelope issuer so a token of one family can never validate in the other eve
hypothetical shared signing key."""
SESSION_TTL_SECONDS: Final = 3600
"""Session ACCESS token lifetime (1h), matching the access-envelope and BYOK session bearer
windows: a client-held credential never outlives a bounded window, and each refresh
re-validates the live user before re-minting."""
"""Session ACCESS token lifetime (1h), matching the BYOK session bearer window: a
client-held credential never outlives a bounded window, and each refresh re-validates
the live user before re-minting."""
SESSION_REFRESH_TTL_SECONDS: Final = 1209600
"""Session REFRESH token lifetime (14 days), matching the refresh-envelope bound. Each

View file

@ -2819,7 +2819,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# Values stay `object` rather than BudgetConfig: this is the raw JSON column,
# and validating it here would make one malformed row fail auth outright.
# resolve_model_budget validates the single entry a request actually needs.
user_model_max_budget: dict[str, object] | None = None
user_model_max_budget: Mapping[str, object] | None = None
request_route: str | None = None
is_session_token: bool = False
# Server-only marker set exclusively by the MCP gateway admission path
@ -2997,8 +2997,8 @@ class UserInfoV2Response(LiteLLMPydanticObjectBase):
sso_user_id: str | None = None
teams: list[str] = [] # Just team IDs, not full team objects
object_permission: LiteLLM_ObjectPermissionTable | None = None
model_max_budget: dict | None = None
model_max_budget_usage: dict | None = None
model_max_budget: Mapping[str, object] | None = None
model_max_budget_usage: Mapping[str, Mapping[str, object]] | None = None
from litellm.models.config import LiteLLM_Config as LiteLLM_Config # noqa: E402

View file

@ -212,9 +212,9 @@ async def _read_user_model_max_budget(
user_id: str | None,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: object,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging,
) -> dict | None:
) -> Mapping[str, object] | None:
"""The user row's `model_max_budget`, or None when the row cannot be read.
A user whose row is missing must not be refused: this is a budget lookup,
@ -228,13 +228,13 @@ async def _read_user_model_max_budget(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
parent_otel_span=parent_otel_span, # pyright: ignore[reportArgumentType] # Span is a runtime union, not usable in an annotation here
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception as e: # noqa: BLE001 # mirrors the main path's tolerance
verbose_logger.debug("Unable to read user for the per-model budget check: %s", e)
return None
return getattr(user_obj, "model_max_budget", None)
return user_obj.model_max_budget if user_obj is not None else None
async def _check_user_model_budget(
@ -3267,8 +3267,7 @@ async def _run_post_custom_auth_checks(
# loaded the user row yet. The attach is unconditional because the post-call
# spend hook reads this field off the token: gating it on the same condition
# as enforcement would leave the user's counter uncharged whenever this
# request was not itself enforceable, which is the untracked-spend bug this
# PR exists to fix.
# request was not itself enforceable, so its spend would go untracked.
user_budget: Final = await _read_user_model_max_budget(
user_id=valid_token.user_id,
prisma_client=prisma_client,

View file

@ -4,7 +4,7 @@ import json
import logging
import math
import traceback
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
from datetime import datetime
from functools import lru_cache
from types import MappingProxyType
@ -284,7 +284,7 @@ def _deferred_stream_logging_is_armed(request_data: dict) -> bool:
)
def _assembled_model_came_from_a_later_chunk(chunks: list, assembled_model: object) -> bool:
def _assembled_model_came_from_a_later_chunk(chunks: Sequence[object], assembled_model: object) -> bool:
"""Report whether stream_chunk_builder picked a model the first chunk did not carry.
Azure Model Router puts the routed model on the chunks after the first one, and the
@ -306,7 +306,10 @@ def _assembled_model_came_from_a_later_chunk(chunks: list, assembled_model: obje
)
def _assembled_model_is_the_name_the_client_asked_for(request_data: dict, assembled_model: object) -> bool:
def _assembled_model_is_the_name_the_client_asked_for(
request_data: Mapping[str, object],
assembled_model: object,
) -> bool:
"""Report whether the assembled model is the public name the proxy stamps onto chunks.
That stamp is what leaves an unpriced alias on the partial response, so the deployment's

View file

@ -1182,7 +1182,7 @@ class ResetBudgetJob:
if not raw:
continue
row_id: str = row[source.id_column]
windows: list = raw if isinstance(raw, list) else json.loads(raw)
windows: list[dict[str, object]] = raw if isinstance(raw, list) else json.loads(raw)
changed = False
for window in windows:
counter_key = f"{source.counter_prefix}:{row_id}:window:{window['budget_duration']}"

View file

@ -183,7 +183,7 @@ async def run_with_timeout(task, timeout):
return {"error": "Timeout exceeded", "exception": timeout_exception}
def _is_strategy_router_deployment(litellm_params: dict) -> bool:
def _is_strategy_router_deployment(litellm_params: Mapping[str, object]) -> bool:
"""True for strategy-router deployments."""
model: Final[object] = litellm_params.get("model", "")
return isinstance(model, str) and classify_strategy_router_model(model) is not None

View file

@ -1796,6 +1796,7 @@ async def test_model_connection(
"audio_speech",
"audio_transcription",
"image_generation",
"image_edit",
"video_generation",
"batch",
"rerank",

View file

@ -4,7 +4,7 @@ import json
import re
import time
from collections import OrderedDict
from collections.abc import Mapping
from collections.abc import Mapping, MutableMapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
@ -1629,7 +1629,7 @@ class LiteLLMProxyRequestSetup:
def refresh_proxy_server_request_body_snapshot(
data: dict, # mutable-ok: mutates proxy_server_request.body in place on the shared request dict
data: MutableMapping[str, object],
) -> None:
"""
Re-snapshot ``data["proxy_server_request"]["body"]`` from the current state of ``data``.

View file

@ -2294,10 +2294,9 @@ async def _process_single_key_update(
prisma_client=prisma_client,
)
_existing_row_metadata: Final = getattr(existing_key_row, "metadata", None)
enforce_batch_enqueued_token_limit_is_admin_only(
data=update_key_request,
existing_metadata=_existing_row_metadata if isinstance(_existing_row_metadata, dict) else None,
existing_metadata=existing_key_row.metadata,
user_api_key_dict=user_api_key_dict,
entity="key",
)

View file

@ -1354,11 +1354,12 @@ def _completed_batch_safe_to_retire(response: "LiteLLMBatch") -> bool:
reports no successful request lines. When counts are unknown, stay eligible so
the next poller pass revisits it. (#37713)
"""
if getattr(response, "output_file_id", None) is not None:
if response.output_file_id is not None:
return True
request_counts = getattr(response, "request_counts", None)
completed = getattr(request_counts, "completed", None)
return completed == 0
request_counts = response.request_counts
if request_counts is None:
return False
return request_counts.completed == 0
async def update_batch_in_database(

View file

@ -6,6 +6,7 @@ Provider-specific Pass-Through Endpoints
Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
"""
import hmac
import json
import os
import re
@ -28,6 +29,7 @@ from litellm.constants import (
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import (
_get_bearer_token,
@ -1791,29 +1793,11 @@ _HEADERS_NEVER_FORWARDED_TO_VERTEX: Final = frozenset({"content-length", "host"}
)
_VERTEX_CALLER_KEY_HEADER_PRECEDENCE: Final = (
SpecialHeaders.custom_litellm_api_key.value.lower(),
SpecialHeaders.openai_authorization.value.lower(),
SpecialHeaders.azure_authorization.value.lower(),
SpecialHeaders.anthropic_authorization.value.lower(),
SpecialHeaders.google_ai_studio_authorization.value.lower(),
SpecialHeaders.azure_apim_authorization.value.lower(),
)
_MAPPED_ROUTE_CALLER_KEY_HEADER: Final = "litellm_user_api_key"
def _operator_configured_caller_key_header_names() -> tuple[tuple[str, ...], tuple[str, ...]]:
"""Operator-configured caller-key header names, as (override, pass_through).
``user_api_key_auth`` accepts the caller's key from two runtime-configured
header sources beyond the built-in ones, at opposite ends of its precedence.
``general_settings.litellm_key_header_name`` overrides every built-in source
(it replaces the resolved key after ``get_api_key`` runs), so it is highest
precedence. Each ``general_settings.pass_through_endpoints`` entry's
``headers.litellm_user_api_key`` is checked last inside ``get_api_key``, so it
is lowest. Google never consumes either, so both are also dropped by name.
"""
def _operator_configured_caller_key_header_names() -> tuple[str, ...]:
"""Operator-configured header names ``user_api_key_auth`` reads the caller's key from."""
from litellm.proxy.proxy_server import general_settings
custom_key_header: Final = general_settings.get("litellm_key_header_name")
@ -1829,80 +1813,53 @@ def _operator_configured_caller_key_header_names() -> tuple[tuple[str, ...], tup
if isinstance(headers, dict) and isinstance(headers.get("litellm_user_api_key"), str)
)
)
return override, pass_through
return override + pass_through
def _authenticated_caller_key_values(request: Request) -> frozenset[str]:
"""The value ``user_api_key_auth`` would accept as this caller's LiteLLM key.
The Vertex route authenticates through ``Depends(user_api_key_auth)``, which
resolves the key by precedence, matched here exactly. The ``/vertex_ai`` route
is a mapped pass-through route, so a header literally named
``litellm_user_api_key`` overrides every other source (``user_api_key_auth``
applies it last), making it highest precedence. Then an operator
``litellm_key_header_name``, then the built-in headers in ``get_api_key`` order,
then a ``pass_through_endpoints`` ``litellm_user_api_key`` header which
``get_api_key`` checks last. Some of those headers (``Authorization``,
``x-goog-api-key``) are also kept as genuine bring-your-own Google credentials,
so returning only the value that actually authenticated lets the filter strip
that value wherever it appears while leaving a real Google credential in place.
An empty set means no caller key was found, so nothing is value-stripped.
"""
incoming: Final = _safe_get_request_headers(request)
override_headers, pass_through_headers = _operator_configured_caller_key_header_names()
ordered_names: Final = (
(_MAPPED_ROUTE_CALLER_KEY_HEADER,)
+ override_headers
+ _VERTEX_CALLER_KEY_HEADER_PRECEDENCE
+ pass_through_headers
def _is_authenticated_caller_jwt(value: str, jwt_claims: Mapping[str, object]) -> bool:
"""Whether a header value is the JWT whose claims ``user_api_key_auth`` stored as ``jwt_claims``."""
presented_claims: Final = JWTHandler.get_unverified_claims(value)
if presented_claims is None:
return False
return all(
presented_claims.get(name) == claim
for name, claim in jwt_claims.items()
if name not in JWTHandler.LITELLM_INTERNAL_CLAIMS
)
present_values: Final = (incoming[name] for name in ordered_names if incoming.get(name))
authenticated_key: Final = next(
(stripped for value in present_values if (stripped := _normalize_credential_value(value))),
"",
)
return frozenset({authenticated_key}) if authenticated_key else frozenset()
def _forwarded_headers_for_credentialless_vertex_passthrough(request: Request) -> Mapping[str, str]:
"""
Header set to forward on the bring-your-own-credentials Vertex passthrough
branch, used when the proxy has no Vertex credential configured.
def _is_authenticated_caller_secret(value: str, user_api_key_dict: UserAPIKeyAuth) -> bool:
"""Whether a header value is the master key, the JWT that authenticated, or the key stored as ``api_key``."""
from litellm.proxy.proxy_server import master_key
No credential the proxy accepts for caller authentication is forwarded to
Google. ``user_api_key_auth`` reads the caller's key from every header in
``SpecialHeaders.litellm_credential_header_names()``, and Vertex only ever
authenticates with an OAuth token in ``Authorization`` or an API key in
``x-goog-api-key``. So the proxy-only auth headers Google never consumes
(everything in that set except those two, e.g. ``x-litellm-api-key`` /
``api-key`` / ``x-api-key`` / ``Ocp-Apim-Subscription-Key``, plus the mapped
pass-through ``litellm_user_api_key`` header and any operator-configured
``litellm_key_header_name`` / ``pass_through_endpoints`` key header) are dropped
by name. ``Authorization`` and ``x-goog-api-key`` may
instead carry a genuine bring-your-own Google credential, so they are kept
unless their value is the caller's authenticated LiteLLM key, which is dropped
by value (normalizing any ``Bearer`` / ``Basic`` / ``AWS4`` auth-scheme prefix
the same way authentication does). Because the value that authenticated is
resolved by the same precedence ``user_api_key_auth`` uses, a virtual key sent
only in ``x-goog-api-key`` (or in an operator-configured key header) is dropped
too, while a real Google key in ``x-goog-api-key`` alongside a virtual key in a
higher-precedence header is preserved. When neither a surviving
``Authorization`` nor ``x-goog-api-key`` remains the request is rejected so the
virtual key cannot leak upstream.
"""
normalized: Final = _normalize_credential_value(value)
if master_key is not None and hmac.compare_digest(normalized.encode(), master_key.encode()):
return True
jwt_claims: Final = user_api_key_dict.jwt_claims
if jwt_claims and _is_authenticated_caller_jwt(normalized, jwt_claims):
return True
authenticated_key: Final = user_api_key_dict.api_key
if authenticated_key is None:
return False
if master_key is None and not normalized.startswith("sk-"):
return False
stored_representation: Final = UserAPIKeyAuth._safe_hash_litellm_api_key(normalized) # pyright: ignore[reportPrivateUsage] # the exact transform auth applied when it stored api_key
return hmac.compare_digest(stored_representation.encode(), authenticated_key.encode())
def _forwarded_headers_for_credentialless_vertex_passthrough(
request: Request, user_api_key_dict: UserAPIKeyAuth
) -> Mapping[str, str]:
"""Caller headers to forward on the bring-your-own-credentials Vertex branch, minus LiteLLM secrets."""
incoming: Final = _safe_get_request_headers(request)
caller_key_values: Final = _authenticated_caller_key_values(request)
override_headers, pass_through_headers = _operator_configured_caller_key_header_names()
never_forwarded: Final = (
_HEADERS_NEVER_FORWARDED_TO_VERTEX.union((_MAPPED_ROUTE_CALLER_KEY_HEADER,))
.union(override_headers)
.union(pass_through_headers)
never_forwarded: Final = _HEADERS_NEVER_FORWARDED_TO_VERTEX.union(
(_MAPPED_ROUTE_CALLER_KEY_HEADER, *_operator_configured_caller_key_header_names())
)
forwarded: Final = MappingProxyType(
{
name: value
for name, value in incoming.items()
if name not in never_forwarded and _normalize_credential_value(value) not in caller_key_values
if name not in never_forwarded and not _is_authenticated_caller_secret(value, user_api_key_dict)
}
)
if "authorization" not in forwarded and "x-goog-api-key" not in forwarded:
@ -1918,6 +1875,7 @@ async def _prepare_vertex_auth_headers(
vertex_location: str | None,
base_target_url: str | None,
get_vertex_pass_through_handler: BaseVertexAIPassThroughHandler,
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[Mapping[str, str], str | None, bool, str | None, str | None]:
"""
Prepare authentication headers for Vertex AI pass-through requests.
@ -1930,6 +1888,8 @@ async def _prepare_vertex_auth_headers(
vertex_location: Vertex location
base_target_url: Base URL for the Vertex AI service
get_vertex_pass_through_handler: Handler for the specific Vertex AI service
user_api_key_dict: The caller's resolved authentication, so only the secret that
authenticated them is stripped on the credential-less branch
Returns:
Tuple containing:
@ -1944,7 +1904,7 @@ async def _prepare_vertex_auth_headers(
# Use headers from the incoming request if no vertex credentials are found
if (vertex_credentials is None or vertex_credentials.vertex_project is None) and router_credentials is None:
headers = _forwarded_headers_for_credentialless_vertex_passthrough(request)
headers = _forwarded_headers_for_credentialless_vertex_passthrough(request, user_api_key_dict)
headers_passed_through = True
verbose_proxy_logger.debug(
"default_vertex_config not set, forwarding caller-provided headers %s", tuple(headers.keys())
@ -2104,6 +2064,7 @@ async def _base_vertex_proxy_route(
vertex_location=vertex_location,
base_target_url=base_target_url,
get_vertex_pass_through_handler=get_vertex_pass_through_handler,
user_api_key_dict=user_api_key_dict,
)
if base_target_url is None:

View file

@ -4101,7 +4101,7 @@ def resolve_complexity_router_plugins(
complexity_router_config["classifier_plugin"] = resolved_classifier # rebind-ok: out-param, resolved in place
def validate_deployment_max_agentic_loops(model: Mapping[str, Any]) -> None:
def validate_deployment_max_agentic_loops(model: Mapping[str, object]) -> None:
"""
Reject a per-deployment `max_agentic_loops` the agentic loop cannot honor.
@ -4111,7 +4111,9 @@ def validate_deployment_max_agentic_loops(model: Mapping[str, Any]) -> None:
start. Left unchecked entirely, a `0` used to read as the default ceiling
of 3 and a non-integer failed every request to that model instead.
"""
litellm_params: Final = model.get("litellm_params") or {}
litellm_params: Final = model.get("litellm_params")
if not isinstance(litellm_params, Mapping):
return
if "max_agentic_loops" not in litellm_params:
return

View file

@ -1267,7 +1267,7 @@ def _count_input_tokens_for_models(
_INPUT_SIZE_FIELDS: Final = ("messages", "prompt", "input", "query", "documents", "tools", "tool_choice")
def _approximate_input_size(request_body: dict) -> int:
def _approximate_input_size(request_body: Mapping[str, object]) -> int:
"""Length of the request's input text, a cheap stand-in for tokenizing cost.
Every field _count_input_tokens hands the tokenizer is sized here, and

View file

@ -2,6 +2,8 @@
import asyncio
import os
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final, Literal, cast
import litellm
@ -29,6 +31,7 @@ from litellm.utils import ProviderConfigManager
from ..litellm_core_utils.get_litellm_params import get_litellm_params
from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ..llms.azure.common_utils import get_azure_ad_token
from ..llms.azure.realtime.handler import AzureOpenAIRealtime
from ..llms.bedrock.realtime.handler import BedrockRealtime
from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
@ -44,6 +47,7 @@ bedrock_realtime: Final = BedrockRealtime()
xai_realtime: Final = XAIRealtime()
vertex_llm_base: Final = VertexBase()
base_llm_http_handler = BaseLLMHTTPHandler()
_EMPTY_MODEL_PARAMS: Final[Mapping[str, Any]] = MappingProxyType({})
def _with_resolved_session_model(session: dict[str, Any], model_name: str) -> dict[str, Any]:
@ -411,13 +415,16 @@ async def _arealtime(
if realtime_protocol is None and (query_params or {}).get("intent") == "transcription":
realtime_protocol = "GA"
realtime_protocol = realtime_protocol or "beta"
resolved_azure_ad_token: Final = (
None if api_key else get_azure_ad_token(GenericLiteLLMParams(**kwargs, azure_ad_token=azure_ad_token))
)
await azure_realtime.async_realtime(
model=model,
websocket=websocket,
api_base=api_base,
api_key=api_key,
api_version=api_version,
azure_ad_token=None,
azure_ad_token=resolved_azure_ad_token,
client=None,
timeout=timeout,
logging_obj=litellm_logging_obj,
@ -550,6 +557,17 @@ async def _arealtime(
raise ValueError(f"Unsupported model: {model}")
def _realtime_health_check_auth_headers(
custom_llm_provider: str, api_key: str | None, model_params: Mapping[str, Any]
) -> Mapping[str, str | None]:
if custom_llm_provider != "azure":
return MappingProxyType({"api-key": api_key})
return azure_realtime.get_auth_headers(
api_key=api_key,
azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))),
)
async def _realtime_health_check(
model: str,
custom_llm_provider: str,
@ -578,6 +596,11 @@ async def _realtime_health_check(
import websockets
url: str | None = None
auth_headers: Final = _realtime_health_check_auth_headers(
custom_llm_provider=custom_llm_provider,
api_key=api_key,
model_params=model_params or _EMPTY_MODEL_PARAMS,
)
if custom_llm_provider == "azure":
url = azure_realtime._construct_url(
api_base=api_base or "",
@ -627,9 +650,7 @@ async def _realtime_health_check(
ssl_context = get_shared_realtime_ssl_context()
async with websockets.connect(
url,
additional_headers={
"api-key": api_key,
},
additional_headers=auth_headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_context,
):

View file

@ -32,7 +32,7 @@ class _PrismaClientView(Protocol):
class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]):
"""Repository for proxy model database operations with encryption support."""
def __init__(self, prisma_client: object, encryption_key: str | None = None):
def __init__(self, prisma_client: object, encryption_key: str | None = None) -> None:
super().__init__(prisma_client)
self._encryption_key = encryption_key

View file

@ -1300,16 +1300,14 @@ class LiteLLMCompletionResponsesConfig:
if isinstance(content, str) and content.strip():
return content
if isinstance(content, list):
text_parts: Final[list[str]] = [] # mutable-ok: text accumulator
for block in content:
if not isinstance(block, Mapping):
continue
block_type = block.get("type")
if block_type in ("encrypted_content", "redacted_thinking"):
continue
text = block.get("text")
if isinstance(text, str) and text.strip():
text_parts.append(text.strip())
text_parts: Final = tuple(
text.strip()
for block in content
if isinstance(block, Mapping)
and block.get("type") not in ("encrypted_content", "redacted_thinking")
and isinstance(text := block.get("text"), str)
and text.strip()
)
if text_parts:
return "\n".join(text_parts)
return None
@ -1325,13 +1323,11 @@ class LiteLLMCompletionResponsesConfig:
summary: Final[object] = input_item.get("summary")
if not isinstance(summary, list):
return None
text_parts: Final[list[str]] = [] # mutable-ok: text accumulator
for block in summary:
if not isinstance(block, Mapping):
continue
text = block.get("text")
if isinstance(text, str) and text.strip():
text_parts.append(text.strip())
text_parts: Final = tuple(
text.strip()
for block in summary
if isinstance(block, Mapping) and isinstance(text := block.get("text"), str) and text.strip()
)
return "\n".join(text_parts) if text_parts else None
@staticmethod

View file

@ -10697,10 +10697,9 @@ class Router:
_router_model_name: str = model_value
elif isinstance(model_value, dict):
_model_value = RouterModelGroupAliasItem(**model_value)
if _model_value["hidden"] is True:
if _model_value["hidden"] is True and model_name is None:
continue
else:
_router_model_name = _model_value["model"]
_router_model_name = _model_value["model"]
else:
continue

View file

@ -3,7 +3,7 @@ from collections.abc import Mapping, Sequence
from dataclasses import MISSING, dataclass, field, fields
from enum import Enum
from types import MappingProxyType
from typing import Any, ClassVar, Final, Literal
from typing import Any, ClassVar, Final, Literal, cast
import litellm
@ -326,21 +326,25 @@ def validate_prometheus_deployment_and_latency_caller_identity() -> str:
)
def validate_caller_identity_settings(litellm_settings: Mapping[str, Any]) -> None:
def validate_caller_identity_settings(litellm_settings: Mapping[str, object]) -> None:
"""Store the caller-identity mode from litellm_settings and validate it together
with prometheus_metrics_config, raising on an invalid value or on include_labels
that request a label the selected mode removes."""
if "prometheus_deployment_and_latency_caller_identity" not in litellm_settings:
return
litellm.prometheus_deployment_and_latency_caller_identity = litellm_settings[
"prometheus_deployment_and_latency_caller_identity"
]
litellm.prometheus_deployment_and_latency_caller_identity = (
cast( # cast-ok: validated on the next line, which raises on an invalid value
'Literal["api_key_alias", "user_email", "both"]',
litellm_settings["prometheus_deployment_and_latency_caller_identity"],
)
)
caller_identity_mode: Final = validate_prometheus_deployment_and_latency_caller_identity()
if caller_identity_mode != "user_email":
return
raw_metrics_config: Final = litellm_settings.get("prometheus_metrics_config")
conflicting_metrics: Final = tuple(
metric_name
for metric_config in (litellm_settings.get("prometheus_metrics_config") or ())
for metric_config in (raw_metrics_config if isinstance(raw_metrics_config, list) else ())
if isinstance(metric_config, dict) and "api_key_alias" in (metric_config.get("include_labels") or ())
for metric_name in (metric_config.get("metrics") or ())
if metric_name in PROMETHEUS_DEPLOYMENT_AND_LATENCY_CALLER_IDENTITY_METRICS

View file

@ -3,7 +3,7 @@ from collections.abc import Sequence
from enum import Enum
from typing import TYPE_CHECKING, Any, Final, Literal
from typing_extensions import Required, TypedDict, override
from typing_extensions import ReadOnly, Required, TypedDict, override
from .openai import ChatCompletionToolCallChunk
@ -97,6 +97,10 @@ class BedrockConverseReasoningContentBlockDelta(TypedDict, total=False):
text: str
class BedrockConverseGptReasoningEffortBlock(TypedDict):
effort: ReadOnly[str]
class GuardrailConverseTextBlock(TypedDict, total=False):
text: str

View file

@ -1292,7 +1292,7 @@ async def async_post_call_success_deployment_hook(
async def async_post_call_failure_deployment_hook(
request_data: Mapping[str, Any], exception: Exception, call_type: str
request_data: Mapping[str, object], exception: Exception, call_type: str
) -> None:
"""
Notify CustomLogger callbacks that a deployment attempt failed.

View file

@ -12282,7 +12282,7 @@
},
"claude-3-haiku-20240307": {
"cache_creation_input_token_cost": 3e-07,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr": 5e-07,
"cache_read_input_token_cost": 3e-08,
"deprecation_date": "2026-04-20",
"input_cost_per_token": 2.5e-07,
@ -12301,7 +12301,7 @@
},
"claude-3-opus-20240229": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr": 3e-05,
"cache_read_input_token_cost": 1.5e-06,
"deprecation_date": "2026-01-05",
"input_cost_per_token": 1.5e-05,
@ -12515,7 +12515,10 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_output_config": true,
"prompt_cache_min_tokens": 1024
"prompt_cache_min_tokens": 1024,
"provider_specific_entry": {
"us": 1.1
}
},
"claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
@ -49469,6 +49472,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"global.openai.gpt-5.6-sol": {
@ -49494,6 +49498,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"us.openai.gpt-5.6-terra": {
@ -49519,6 +49524,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"global.openai.gpt-5.6-terra": {
@ -49544,6 +49550,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"us.openai.gpt-5.6-luna": {
@ -49569,6 +49576,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"global.openai.gpt-5.6-luna": {
@ -49594,6 +49602,7 @@
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true
},
"bedrock_mantle/openai.gpt-5.5": {
@ -49601,7 +49610,7 @@
"cache_read_input_token_cost": 5.5e-07,
"output_cost_per_token": 3.3e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 272000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
@ -49628,7 +49637,7 @@
"cache_read_input_token_cost": 2.75e-07,
"output_cost_per_token": 1.65e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 272000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
@ -50821,7 +50830,10 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true
"supports_native_structured_output": true,
"provider_specific_entry": {
"us": 1.1
}
},
"claude-mythos-preview": {
"cache_creation_input_token_cost": 1.25e-05,
@ -50856,7 +50868,10 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true
"supports_native_structured_output": true,
"provider_specific_entry": {
"us": 1.1
}
},
"gemini/gemini-robotics-er-2-streaming-preview": {
"input_cost_per_audio_token": 2e-06,

View file

@ -12,10 +12,10 @@
"limit": 2012
},
"ANN202": {
"limit": 852
"limit": 847
},
"ANN204": {
"limit": 711
"limit": 706
},
"ANN205": {
"limit": 112
@ -24,7 +24,7 @@
"limit": 133
},
"ANN401": {
"limit": 1157
"limit": 1153
},
"ASYNC230": {
"limit": 11

View file

@ -16,9 +16,12 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.integrations.prometheus import (
DEFINED_PROMETHEUS_METRICS,
PROMETHEUS_DEPLOYMENT_AND_LATENCY_CALLER_IDENTITY_METRICS,
LabelValidationError,
PrometheusMetricLabels,
UserAPIKeyLabelNames,
UserAPIKeyLabelValues,
validate_caller_identity_settings,
validate_prometheus_deployment_and_latency_caller_identity,
)
from litellm.types.utils import StandardLoggingPayload
@ -379,23 +382,12 @@ def test_deployment_failure_email_fallbacks_reach_both_real_counters(
async def test_proxy_config_loads_caller_identity_before_initializing_callbacks(tmp_path: Path):
from litellm.proxy.proxy_server import ProxyConfig
config_path = tmp_path / "config.yaml"
config_path.write_text(
yaml.safe_dump(
{
"model_list": [
{
"model_name": "test-model",
"litellm_params": {"model": "openai/gpt-4", "api_key": "test-key"},
}
],
"litellm_settings": {
"callbacks": ["prometheus"],
"prometheus_deployment_and_latency_caller_identity": "both",
},
},
sort_keys=False,
)
config_path = _write_proxy_config(
tmp_path,
{
"callbacks": ["prometheus"],
"prometheus_deployment_and_latency_caller_identity": "both",
},
)
observed_modes: list[str] = []
@ -409,3 +401,287 @@ async def test_proxy_config_loads_caller_identity_before_initializing_callbacks(
assert observed_modes == ["both"]
assert litellm.prometheus_deployment_and_latency_caller_identity == "both"
def _identity_settings(mode: object, metrics_config: object = None) -> dict[str, object]:
settings: dict[str, object] = {"prometheus_deployment_and_latency_caller_identity": mode}
if metrics_config is not None:
settings["prometheus_metrics_config"] = metrics_config
return settings
def test_validate_mode_returns_each_accepted_value_and_defaults_to_api_key_alias(
monkeypatch: pytest.MonkeyPatch,
):
for mode in IDENTITY_MODES:
_set_caller_identity(monkeypatch, mode)
assert validate_prometheus_deployment_and_latency_caller_identity() == mode
monkeypatch.delattr(litellm, "prometheus_deployment_and_latency_caller_identity")
assert validate_prometheus_deployment_and_latency_caller_identity() == "api_key_alias"
def test_accepted_values_constant_matches_parametrized_modes():
from litellm.types.integrations.prometheus import (
PROMETHEUS_DEPLOYMENT_AND_LATENCY_CALLER_IDENTITY_VALUES,
)
assert PROMETHEUS_DEPLOYMENT_AND_LATENCY_CALLER_IDENTITY_VALUES == IDENTITY_MODES
assert len(TARGET_METRICS) == 9
@pytest.mark.parametrize(
"invalid_mode",
("user-email", "USER_EMAIL", "", None, True, 1, ["user_email"], {"mode": "user_email"}),
)
def test_validate_mode_rejects_invalid_values_and_names_accepted_ones(
monkeypatch: pytest.MonkeyPatch,
invalid_mode: object,
):
monkeypatch.setattr(litellm, "prometheus_deployment_and_latency_caller_identity", invalid_mode)
with pytest.raises(ValueError, match="prometheus_deployment_and_latency_caller_identity") as exc_info:
validate_prometheus_deployment_and_latency_caller_identity()
message = str(exc_info.value)
assert repr(invalid_mode) in message
for accepted_value in IDENTITY_MODES:
assert accepted_value in message
def test_validate_caller_identity_settings_without_key_leaves_mode_untouched(
monkeypatch: pytest.MonkeyPatch,
):
_set_caller_identity(monkeypatch, "both")
validate_caller_identity_settings({"prometheus_metrics_config": []})
assert litellm.prometheus_deployment_and_latency_caller_identity == "both"
@pytest.mark.parametrize("mode", IDENTITY_MODES)
def test_validate_caller_identity_settings_stores_each_valid_mode(mode: str):
validate_caller_identity_settings(_identity_settings(mode))
assert litellm.prometheus_deployment_and_latency_caller_identity == mode
@pytest.mark.parametrize("invalid_mode", ("user-email", None))
def test_validate_caller_identity_settings_rejects_invalid_and_null_modes(invalid_mode: object):
with pytest.raises(ValueError, match="prometheus_deployment_and_latency_caller_identity"):
validate_caller_identity_settings(_identity_settings(invalid_mode))
def test_user_email_mode_conflict_error_names_every_conflicting_metric_and_only_those():
metrics_config = [
{
"group": "non_target",
"metrics": ["litellm_overhead_with_guardrails_latency_metric"],
"include_labels": ["api_key_alias"],
},
{
"group": "target_pair",
"metrics": ["litellm_deployment_total_requests", "litellm_llm_api_latency_metric"],
"include_labels": ["api_key_alias"],
},
{
"group": "target_single",
"metrics": ["litellm_request_queue_time_seconds"],
"include_labels": ["api_key_alias"],
},
]
with pytest.raises(ValueError, match="prometheus_deployment_and_latency_caller_identity") as exc_info:
validate_caller_identity_settings(_identity_settings("user_email", metrics_config))
message = str(exc_info.value)
for conflicting_metric in (
"litellm_deployment_total_requests",
"litellm_llm_api_latency_metric",
"litellm_request_queue_time_seconds",
):
assert conflicting_metric in message
assert "litellm_overhead_with_guardrails_latency_metric" not in message
assert "prometheus_deployment_and_latency_caller_identity" in message
assert "user_email" in message
@pytest.mark.parametrize(
("mode", "metrics_config"),
(
(
"user_email",
[
{
"group": "g",
"metrics": ["litellm_deployment_total_requests"],
"include_labels": ["user_email"],
}
],
),
(
"user_email",
[
{
"group": "g",
"metrics": ["litellm_overhead_with_guardrails_latency_metric"],
"include_labels": ["api_key_alias"],
}
],
),
(
"api_key_alias",
[
{
"group": "g",
"metrics": ["litellm_deployment_total_requests"],
"include_labels": ["api_key_alias"],
}
],
),
(
"both",
[
{
"group": "g",
"metrics": ["litellm_deployment_total_requests"],
"include_labels": ["api_key_alias"],
}
],
),
("user_email", None),
("user_email", ["not-a-dict"]),
(
"user_email",
[{"group": "g", "metrics": ["litellm_deployment_total_requests"], "include_labels": None}],
),
("user_email", [{"group": "g", "metrics": None, "include_labels": ["api_key_alias"]}]),
),
)
def test_validate_caller_identity_settings_accepts_non_conflicting_configs(
mode: str,
metrics_config: object,
):
settings = _identity_settings(mode)
settings["prometheus_metrics_config"] = metrics_config
validate_caller_identity_settings(settings)
assert litellm.prometheus_deployment_and_latency_caller_identity == mode
def _write_proxy_config(tmp_path: Path, litellm_settings: dict[str, object]) -> Path:
config_path = tmp_path / "config.yaml"
config_path.write_text(
yaml.safe_dump(
{
"model_list": [
{
"model_name": "test-model",
"litellm_params": {"model": "openai/gpt-4", "api_key": "test-key"},
}
],
"litellm_settings": litellm_settings,
},
sort_keys=False,
)
)
return config_path
@pytest.mark.asyncio
@pytest.mark.parametrize(
"litellm_settings",
(
{
"callbacks": ["prometheus"],
"prometheus_deployment_and_latency_caller_identity": "user-email",
},
{
"callbacks": ["prometheus"],
"prometheus_deployment_and_latency_caller_identity": None,
},
{
"callbacks": ["prometheus"],
"prometheus_deployment_and_latency_caller_identity": "user_email",
"prometheus_metrics_config": [
{
"group": "g",
"metrics": ["litellm_deployment_total_requests"],
"include_labels": ["api_key_alias"],
}
],
},
),
ids=("typo-mode", "null-mode", "include-labels-conflict"),
)
async def test_proxy_config_fails_boot_before_callbacks_on_invalid_caller_identity_config(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
litellm_settings: dict[str, object],
):
from litellm.proxy.proxy_server import ProxyConfig
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
config_path = _write_proxy_config(tmp_path, litellm_settings)
with patch( # test-quality-ok: asserts boot fails before any callback initialization
"litellm.proxy.proxy_server.initialize_callbacks_on_proxy"
) as callback_init:
with pytest.raises(ValueError, match="prometheus_deployment_and_latency_caller_identity"):
await ProxyConfig().load_config(router=None, config_file_path=str(config_path))
callback_init.assert_not_called()
def test_failed_init_leaves_registry_clean_so_a_corrected_retry_succeeds(
monkeypatch: pytest.MonkeyPatch,
):
_set_caller_identity(monkeypatch, "user-email")
with pytest.raises(ValueError, match="prometheus_deployment_and_latency_caller_identity"):
PrometheusLogger()
assert list(REGISTRY._collector_to_names) == [] # pyright: ignore[reportPrivateUsage]
_set_caller_identity(monkeypatch, "user_email")
logger = PrometheusLogger()
assert "user_email" in logger.get_labels_for_metric("litellm_deployment_total_requests")
@pytest.mark.parametrize("invalid_label", ("api_key_alias", "user_email"))
def test_label_validation_error_names_mode_setting_for_identity_labels_on_target_metric(
monkeypatch: pytest.MonkeyPatch,
invalid_label: str,
):
_set_caller_identity(monkeypatch, "user_email")
error = LabelValidationError(
metric_name="litellm_deployment_total_requests",
invalid_labels=[invalid_label],
valid_labels=["user_email"],
)
assert "prometheus_deployment_and_latency_caller_identity='user_email'" in error.message
assert invalid_label in error.message
def test_label_validation_error_keeps_base_message_for_non_identity_cases(
monkeypatch: pytest.MonkeyPatch,
):
_set_caller_identity(monkeypatch, "user_email")
non_target_metric = LabelValidationError(
metric_name="litellm_overhead_with_guardrails_latency_metric",
invalid_labels=["api_key_alias"],
valid_labels=[],
)
non_identity_label = LabelValidationError(
metric_name="litellm_deployment_total_requests",
invalid_labels=["bogus_label"],
valid_labels=[],
)
for error in (non_target_metric, non_identity_label):
assert "caller-identity" not in error.message
assert error.message.startswith("Invalid labels for metric")

View file

@ -531,6 +531,43 @@ def test_generic_cost_per_token_bedrock_mantle_gpt56_long_context(_local_model_c
)
@pytest.mark.parametrize(
"model",
[
"bedrock_mantle/openai.gpt-5.5",
"bedrock_mantle/openai.gpt-5.4",
],
)
def test_generic_cost_per_token_bedrock_mantle_gpt55_gpt54_long_context_flat_rate(_local_model_cost_map, model):
"""Bedrock serves gpt-5.5 and gpt-5.4 up to its enforced 1,050,000-token prompt maximum and documents
no long-context tier for them, so a prompt past 272K is billed at the flat per-token rates."""
model_cost_map = litellm.model_cost[model]
assert model_cost_map["max_input_tokens"] == 1050000
assert [key for key in model_cost_map if "above_272k" in key] == []
served_prompt_tokens = 1030590
cached_tokens = 100000
completion_tokens = 1000
usage = Usage(
prompt_tokens=served_prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=served_prompt_tokens + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
)
prompt_cost, completion_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="bedrock_mantle",
)
assert round(prompt_cost, 10) == round(
model_cost_map["input_cost_per_token"] * (served_prompt_tokens - cached_tokens)
+ model_cost_map["cache_read_input_token_cost"] * cached_tokens,
10,
)
assert round(completion_cost, 10) == round(model_cost_map["output_cost_per_token"] * completion_tokens, 10)
def test_generic_cost_per_token_honors_non_standard_above_threshold():
"""Regression for #30344: get_model_info must keep arbitrary
input/output_cost_per_token_above_<N>_tokens thresholds, not only the hard-coded

View file

@ -12,6 +12,55 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS
@pytest.mark.asyncio
async def test_image_edit_health_check_handler_uses_png_and_prompt():
model_params = {"model": "openai/gpt-image-1", "api_key": "sk-test"}
mode_handlers = HealthCheckHelpers.get_mode_handlers(
model="gpt-image-1",
custom_llm_provider="openai",
model_params=model_params,
)
assert "image_edit" in mode_handlers
with patch( # test-quality-ok: the public health-check path has no dependency injection seam
"litellm.aimage_edit", new_callable=AsyncMock, return_value={}
) as mock_aimage_edit:
await mode_handlers["image_edit"]()
await HealthCheckHelpers.get_mode_handlers(
model="gpt-image-1",
custom_llm_provider="openai",
model_params=model_params,
prompt="edit this image",
)["image_edit"]()
assert mock_aimage_edit.call_count == 2
default_call = mock_aimage_edit.call_args_list[0].kwargs
explicit_call = mock_aimage_edit.call_args_list[1].kwargs
assert default_call["model"] == "openai/gpt-image-1"
assert default_call["prompt"] == "test"
assert explicit_call["prompt"] == "edit this image"
image = default_call["image"]
assert isinstance(image, bytes)
assert image.startswith(b"\x89PNG")
assert int.from_bytes(image[16:20], "big") == 512
assert int.from_bytes(image[20:24], "big") == 512
@pytest.mark.asyncio
async def test_ahealth_check_supports_image_edit_mode():
with patch( # test-quality-ok: the public health-check path has no dependency injection seam
"litellm.aimage_edit", new_callable=AsyncMock, return_value={}
):
result = await ahealth_check(
{"model": "gpt-image-1", "api_key": "sk-test"},
mode="image_edit",
)
assert "error" not in result
assert "Mode image_edit not supported" not in str(result)
def test_update_model_params_with_health_check_tracking_information():
"""Test _update_model_params_with_health_check_tracking_information adds required tracking info."""
initial_model_params = {"model": "gpt-3.5-turbo", "api_key": "test_key"}

View file

@ -1092,6 +1092,7 @@ def test_anthropic_messages_validate_adds_beta_header():
messages=[{"role": "user", "content": [{"type": "text", "text": "Hi"}]}],
optional_params={"context_management": _sample_context_management_payload()},
litellm_params={},
api_key="fake-anthropic-key",
)
assert headers["anthropic-beta"] == "context-management-2025-06-27"

View file

@ -679,7 +679,7 @@ def test_translate_anthropic_to_openai_orders_top_level_and_midturn_system():
def _translate_with_metadata(
model: str, metadata: dict[str, Any], custom_llm_provider: str | None
model: str, metadata: dict[str, str], custom_llm_provider: str | None
) -> dict[str, Any]:
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
anthropic_message_request={

View file

@ -28,6 +28,7 @@ def test_messages_drop_params_strips_speed_for_unsupported_models():
messages=[{"role": "user", "content": "Hello"}],
optional_params=dict(optional_params),
litellm_params={},
api_key="fake-anthropic-key",
)
result = config.transform_anthropic_messages_request(
model="claude-sonnet-4-6",
@ -60,6 +61,7 @@ def test_messages_drop_params_keeps_speed_for_supporting_models():
messages=[{"role": "user", "content": "Hello"}],
optional_params=dict(optional_params),
litellm_params={},
api_key="fake-anthropic-key",
)
result = config.transform_anthropic_messages_request(
model="claude-opus-4-6",

View file

@ -18,9 +18,7 @@ from unittest.mock import patch
import pytest
sys.path.insert(
0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))
)
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../..")))
# Fake tokens for testing (not real secrets)
FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef"
@ -31,21 +29,37 @@ FAKE_AUTH_TOKEN = "sk-ant-aut01-fake-auth-token-for-testing-123456789"
class TestOptionallyHandleAnthropicOAuth:
"""Tests for optionally_handle_anthropic_oauth function."""
def test_oauth_token_in_authorization_header(self):
@pytest.mark.parametrize("header_name", ["authorization", "Authorization", "AUTHORIZATION"])
def test_oauth_token_in_authorization_header(self, header_name):
"""OAuth token in Authorization header should be detected and headers set correctly."""
from litellm.llms.anthropic.common_utils import (
optionally_handle_anthropic_oauth,
)
headers = {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
updated_headers, extracted_api_key = optionally_handle_anthropic_oauth(
headers, None
)
headers = {header_name: f"Bearer {FAKE_OAUTH_TOKEN}"}
updated_headers, extracted_api_key = optionally_handle_anthropic_oauth(headers, None)
assert extracted_api_key == FAKE_OAUTH_TOKEN
assert updated_headers["anthropic-beta"] == "oauth-2025-04-20"
assert updated_headers["anthropic-dangerous-direct-browser-access"] == "true"
assert "x-api-key" not in updated_headers
assert [name for name in updated_headers if name.lower() == "authorization"] == ["authorization"]
assert updated_headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
@pytest.mark.parametrize("api_key_header_name", ["x-api-key", "X-Api-Key"])
def test_oauth_removes_x_api_key_any_casing(self, api_key_header_name):
"""When OAuth wins, a client x-api-key header is removed whatever its casing."""
from litellm.llms.anthropic.common_utils import (
optionally_handle_anthropic_oauth,
)
headers = {api_key_header_name: FAKE_REGULAR_KEY, "Authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
updated_headers, extracted_api_key = optionally_handle_anthropic_oauth(headers, None)
assert extracted_api_key == FAKE_OAUTH_TOKEN
assert [name for name in updated_headers if name.lower() == "x-api-key"] == []
assert [name for name in updated_headers if name.lower() == "authorization"] == ["authorization"]
assert updated_headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
def test_oauth_token_in_api_key_directly(self):
"""OAuth token passed as api_key should set Authorization: Bearer header."""
@ -54,9 +68,7 @@ class TestOptionallyHandleAnthropicOAuth:
)
headers = {}
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(
headers, FAKE_OAUTH_TOKEN
)
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(headers, FAKE_OAUTH_TOKEN)
assert returned_api_key == FAKE_OAUTH_TOKEN
assert updated_headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
@ -71,9 +83,7 @@ class TestOptionallyHandleAnthropicOAuth:
)
headers = {"x-api-key": FAKE_OAUTH_TOKEN}
updated_headers, _ = optionally_handle_anthropic_oauth(
headers, FAKE_OAUTH_TOKEN
)
updated_headers, _ = optionally_handle_anthropic_oauth(headers, FAKE_OAUTH_TOKEN)
assert "x-api-key" not in updated_headers
assert updated_headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
@ -85,9 +95,7 @@ class TestOptionallyHandleAnthropicOAuth:
)
headers = {}
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(
headers, FAKE_REGULAR_KEY
)
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(headers, FAKE_REGULAR_KEY)
assert returned_api_key == FAKE_REGULAR_KEY
assert "authorization" not in updated_headers
@ -101,9 +109,7 @@ class TestOptionallyHandleAnthropicOAuth:
)
headers = {"authorization": f"Bearer {FAKE_REGULAR_KEY}"}
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(
headers, FAKE_REGULAR_KEY
)
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(headers, FAKE_REGULAR_KEY)
assert returned_api_key == FAKE_REGULAR_KEY
assert "anthropic-dangerous-direct-browser-access" not in updated_headers
@ -115,9 +121,7 @@ class TestOptionallyHandleAnthropicOAuth:
)
headers = {}
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(
headers, None
)
updated_headers, returned_api_key = optionally_handle_anthropic_oauth(headers, None)
assert returned_api_key is None
assert "authorization" not in updated_headers
@ -539,16 +543,12 @@ class TestProxyOAuthHeaderForwarding:
)
# Should preserve OAuth even with flag=False
cleaned_without_flag = clean_headers(
raw_headers, forward_llm_provider_auth_headers=False
)
cleaned_without_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=False)
assert "authorization" in cleaned_without_flag
assert cleaned_without_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
# Should also preserve OAuth with flag=True
cleaned_with_flag = clean_headers(
raw_headers, forward_llm_provider_auth_headers=True
)
cleaned_with_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=True)
assert "authorization" in cleaned_with_flag
assert cleaned_with_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
@ -932,9 +932,7 @@ class TestValidateEnvironmentAuthToken:
config = AnthropicModelInfo()
with mock_patch.dict("os.environ", {}, clear=True):
with pytest.raises(
Exception, match=r"ANTHROPIC_API_KEY.*ANTHROPIC_AUTH_TOKEN"
):
with pytest.raises(Exception, match=r"ANTHROPIC_API_KEY.*ANTHROPIC_AUTH_TOKEN"):
config.validate_environment(
headers={},
model="claude-sonnet-4-5-20250929",
@ -980,9 +978,7 @@ class TestGetAuthToken:
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ", {"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN}, clear=True
):
with mock_patch.dict("os.environ", {"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN}, clear=True):
assert AnthropicModelInfo.get_auth_token() == FAKE_AUTH_TOKEN
def test_returns_none_when_not_set(self):
@ -1106,7 +1102,9 @@ class TestGetAuthHeader:
"""Non-standard API key and custom api_base returns Bearer when use_bearer_for_custom_base=True."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
result = AnthropicModelInfo.get_auth_header(api_key="my-custom-key", api_base="https://custom-gateway.com", use_bearer_for_custom_base=True)
result = AnthropicModelInfo.get_auth_header(
api_key="my-custom-key", api_base="https://custom-gateway.com", use_bearer_for_custom_base=True
)
assert result == {"authorization": "Bearer my-custom-key"}
def test_custom_api_base_get_auth_header_uses_x_api_key_when_standard(self):
@ -1124,10 +1122,7 @@ class TestGetApiBaseFallbackChain:
"""Explicit api_base param takes precedence over all env vars."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert (
AnthropicModelInfo.get_api_base("https://explicit.example.com")
== "https://explicit.example.com"
)
assert AnthropicModelInfo.get_api_base("https://explicit.example.com") == "https://explicit.example.com"
def test_defaults_to_anthropic_api(self):
"""get_api_base returns the default Anthropic API base when no env vars are set."""
@ -1180,9 +1175,7 @@ class TestPassthroughAuthToken:
)
config = AnthropicMessagesConfig()
with mock_patch.dict(
"os.environ", {"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN}, clear=True
):
with mock_patch.dict("os.environ", {"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN}, clear=True):
updated_headers, _ = config.validate_anthropic_messages_environment(
headers={},
model="claude-sonnet-4-5-20250929",
@ -1227,6 +1220,52 @@ class TestPassthroughAuthToken:
assert updated_headers["x-api-key"] == FAKE_REGULAR_KEY
assert "authorization" not in updated_headers
def test_passthrough_missing_credentials_raises_authentication_error(self):
"""Passthrough endpoint should raise locally instead of forwarding an unauthenticated request."""
from unittest.mock import patch as mock_patch
import litellm
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
config = AnthropicMessagesConfig()
with mock_patch.dict("os.environ", {}, clear=True):
with pytest.raises(litellm.AuthenticationError, match="Missing Anthropic API Key"):
config.validate_anthropic_messages_environment(
headers={},
model="claude-sonnet-4-5-20250929",
messages=[{"role": "user", "content": "Hello"}],
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
@pytest.mark.parametrize("header_name", ["x-api-key", "X-Api-Key", "X-API-KEY"])
def test_passthrough_client_x_api_key_header_is_kept(self, header_name):
"""A client-forwarded x-api-key header, whatever its casing, should satisfy validation without env credentials."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
config = AnthropicMessagesConfig()
with mock_patch.dict("os.environ", {}, clear=True):
updated_headers, _ = config.validate_anthropic_messages_environment(
headers={header_name: FAKE_REGULAR_KEY},
model="claude-sonnet-4-5-20250929",
messages=[{"role": "user", "content": "Hello"}],
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
assert [name for name in updated_headers if name.lower() == "x-api-key"] == [header_name]
assert updated_headers[header_name] == FAKE_REGULAR_KEY
def test_passthrough_get_complete_url_honours_base_url_env(self):
"""get_complete_url should use ANTHROPIC_BASE_URL when api_base is None."""
from unittest.mock import patch as mock_patch
@ -1290,14 +1329,8 @@ class TestAnthropicThinkingSignatureSelfHeal:
)
assert is_anthropic_invalid_thinking_signature_error("") is False
assert (
is_anthropic_invalid_thinking_signature_error("rate limit exceeded")
is False
)
assert (
is_anthropic_invalid_thinking_signature_error("invalid_request_error: model not found")
is False
)
assert is_anthropic_invalid_thinking_signature_error("rate limit exceeded") is False
assert is_anthropic_invalid_thinking_signature_error("invalid_request_error: model not found") is False
assert is_anthropic_invalid_thinking_signature_error("thinking signature is malformed") is False
def test_strip_thinking_blocks_from_anthropic_messages(self):
@ -1688,10 +1721,7 @@ class TestAnthropicThinkingSignatureSelfHeal:
base = "call_abc123"
sig = "CiIBDDnWx+/a=="
assert (
normalize_anthropic_tool_use_id(f"{base}{THOUGHT_SIGNATURE_SEPARATOR}{sig}")
== base
)
assert normalize_anthropic_tool_use_id(f"{base}{THOUGHT_SIGNATURE_SEPARATOR}{sig}") == base
def test_anthropic_messages_config_http_retry_helpers(self):
import httpx
@ -1715,15 +1745,11 @@ class TestAnthropicThinkingSignatureSelfHeal:
resp_bad = httpx.Response(400, request=req, text="rate limit exceeded")
err_bad = httpx.HTTPStatusError("bad", request=req, response=resp_bad)
assert (
config.should_retry_anthropic_messages_on_http_error(err_bad, {}) is False
)
assert config.should_retry_anthropic_messages_on_http_error(err_bad, {}) is False
resp_500 = httpx.Response(500, request=req, text=err_text)
err_500 = httpx.HTTPStatusError("bad", request=req, response=resp_500)
assert (
config.should_retry_anthropic_messages_on_http_error(err_500, {}) is False
)
assert config.should_retry_anthropic_messages_on_http_error(err_500, {}) is False
data = {
"model": "claude-sonnet-4-20250514",
@ -1746,7 +1772,6 @@ class TestAnthropicThinkingSignatureSelfHeal:
assert data["messages"] == []
class TestClaudeOpus48AdaptiveThinking:
"""Opus 4.8 requires adaptive thinking (``thinking.type='adaptive'`` +
``output_config.effort``). Detection is driven by the
@ -1776,9 +1801,7 @@ class TestClaudeOpus48AdaptiveThinking:
assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True
def test_resolver_reads_flag_through_bedrock_invoke_prefix(
self, local_model_cost_map
):
def test_resolver_reads_flag_through_bedrock_invoke_prefix(self, local_model_cost_map):
"""The resolver fix: ``bedrock/invoke/...`` resolves to the flagged
Bedrock entry. Pure ``_supports_factory`` without prefix-stripping
returns False here, which is why the data-only fix alone was not enough."""
@ -1828,9 +1851,7 @@ class TestClaudeOpus48AdaptiveThinking:
"claude-sonnet-4.6",
],
)
def test_adaptive_thinking_detected_for_opus_4_6_4_7_and_sonnet_4_6(
self, local_model_cost_map, model
):
def test_adaptive_thinking_detected_for_opus_4_6_4_7_and_sonnet_4_6(self, local_model_cost_map, model):
"""Opus 4.6/4.7 and Sonnet 4.6 carry the ``supports_adaptive_thinking`` flag,
so detection holds purely from the cost map with no version-rule
fallback. Each alias form the Bedrock/anthropic paths see resolves to a flagged
@ -1850,9 +1871,7 @@ class TestClaudeOpus48AdaptiveThinking:
"claude-fable-preview",
],
)
def test_unmapped_aliases_without_parseable_version_stay_non_adaptive(
self, local_model_cost_map, model
):
def test_unmapped_aliases_without_parseable_version_stay_non_adaptive(self, local_model_cost_map, model):
"""An alias absent from the map, not matched by any ``fallback_generalizations``
rule, and without any parseable family version stays non-adaptive. ``fable``
without a major version matches neither the core-family 4.6+ gate nor the
@ -1878,9 +1897,7 @@ class TestClaudeOpus48AdaptiveThinking:
"us.anthropic.claude-fable-5-preview",
],
)
def test_adaptive_thinking_version_fallback_for_unmapped_high_versions(
self, local_model_cost_map, model
):
def test_adaptive_thinking_version_fallback_for_unmapped_high_versions(self, local_model_cost_map, model):
"""Provider-prefixed or suffixed Claude names that resolve to no mapped entry
still resolve to adaptive when the id carries claude-<family>- at version 4.6
or higher, bare 5+ majors included. The version gate is the declarative
@ -1901,9 +1918,7 @@ class TestClaudeOpus48AdaptiveThinking:
"us.anthropic.claude-opus-4-20250514",
],
)
def test_adaptive_thinking_not_detected_for_unmapped_low_versions(
self, local_model_cost_map, model
):
def test_adaptive_thinking_not_detected_for_unmapped_low_versions(self, local_model_cost_map, model):
"""Unmapped Claude names below 4.6 stay non-adaptive through the declarative path.
The eight-digit dated Opus 4.0 id (``...-4-20250514``) is the date-safety case: the
version rule caps the minor at two digits, so the date is not misread as a >= 4.6
@ -1942,14 +1957,11 @@ class TestDefaultSuffixAdaptiveThinking:
"vertex_ai/claude-fable-5@default",
],
)
def test_default_suffix_models_are_adaptive_thinking(
self, local_model_cost_map, model: str
) -> None:
def test_default_suffix_models_are_adaptive_thinking(self, local_model_cost_map, model: str) -> None:
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True, (
f"{model} not classified as adaptive thinking. "
"Check _model_map_lookup_candidates strips @default suffix."
f"{model} not classified as adaptive thinking. Check _model_map_lookup_candidates strips @default suffix."
)
@pytest.mark.parametrize(
@ -1959,15 +1971,11 @@ class TestDefaultSuffixAdaptiveThinking:
("vertex_ai/claude-sonnet-4-6@default", "claude-sonnet-4-6"),
],
)
def test_lookup_candidates_include_bare_name(
self, model: str, expected_bare: str
) -> None:
def test_lookup_candidates_include_bare_name(self, model: str, expected_bare: str) -> None:
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
candidates = AnthropicModelInfo._model_map_lookup_candidates(model)
assert expected_bare in candidates, (
f"Expected '{expected_bare}' in candidates for '{model}', got: {candidates}"
)
assert expected_bare in candidates, f"Expected '{expected_bare}' in candidates for '{model}', got: {candidates}"
class TestCapabilityProbeUsesCallerProvider:
@ -1980,42 +1988,27 @@ class TestCapabilityProbeUsesCallerProvider:
BEDROCK_MODEL = "global.anthropic.claude-opus-4-8"
def test_exact_bedrock_entry_flag_is_authoritative_for_bedrock_caller(
self, local_model_cost_map, monkeypatch
):
def test_exact_bedrock_entry_flag_is_authoritative_for_bedrock_caller(self, local_model_cost_map, monkeypatch):
import litellm
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert (
AnthropicModelInfo._is_adaptive_thinking_model(self.BEDROCK_MODEL, "bedrock")
is True
)
assert AnthropicModelInfo._is_adaptive_thinking_model(self.BEDROCK_MODEL, "bedrock") is True
monkeypatch.setitem(
litellm.model_cost[self.BEDROCK_MODEL], "supports_adaptive_thinking", False
)
monkeypatch.setitem(litellm.model_cost[self.BEDROCK_MODEL], "supports_adaptive_thinking", False)
litellm.get_model_info.cache_clear()
assert (
AnthropicModelInfo._is_adaptive_thinking_model(self.BEDROCK_MODEL, "bedrock")
is False
)
assert AnthropicModelInfo._is_adaptive_thinking_model(self.BEDROCK_MODEL, "bedrock") is False
def test_native_anthropic_probe_still_reads_anthropic_entry(
self, local_model_cost_map, monkeypatch
):
def test_native_anthropic_probe_still_reads_anthropic_entry(self, local_model_cost_map, monkeypatch):
import litellm
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
monkeypatch.setitem(
litellm.model_cost[self.BEDROCK_MODEL], "supports_adaptive_thinking", False
)
monkeypatch.setitem(litellm.model_cost[self.BEDROCK_MODEL], "supports_adaptive_thinking", False)
litellm.get_model_info.cache_clear()
assert (
AnthropicModelInfo._is_adaptive_thinking_model("claude-opus-4-8", "anthropic")
is True
)
assert AnthropicModelInfo._is_adaptive_thinking_model("claude-opus-4-8", "anthropic") is True
def test_create_anthropic_model_list_response_shape():
from litellm.llms.anthropic.common_utils import (
create_anthropic_model_list_response,
@ -2100,4 +2093,4 @@ def test_create_anthropic_model_list_response_empty():
assert response["data"] == []
assert response["has_more"] is False
assert response["first_id"] is None
assert response["last_id"] is None
assert response["last_id"] is None

View file

@ -559,3 +559,219 @@ async def test_async_realtime_default_maintains_backwards_compatibility():
mock_realtime_streaming.call_args.kwargs["backend_uses_beta_protocol"]
is True
)
class _DummyAsyncContextManager:
def __init__(self, value):
self.value = value
async def __aenter__(self):
return self.value
async def __aexit__(self, exc_type, exc, tb):
return None
@pytest.mark.asyncio
async def test_async_realtime_uses_bearer_token_when_no_api_key():
"""
Entra ID-only Azure realtime deployments have no static api-key, so the handshake must
authenticate with `Authorization: Bearer <azure_ad_token>` and must not send `api-key`.
Regression test for https://github.com/BerriAI/litellm/issues/34654
"""
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
handler = AzureOpenAIRealtime()
mock_backend_ws = AsyncMock()
with (
patch(
"websockets.connect",
return_value=_DummyAsyncContextManager(mock_backend_ws),
) as mock_ws_connect,
patch( # test-quality-ok: handler owns the streaming loop, only the handshake headers are under test
"litellm.llms.azure.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming,
):
mock_realtime_streaming.return_value.bidirectional_forward = AsyncMock()
await handler.async_realtime(
model="gpt-realtime-whisper",
websocket=AsyncMock(),
logging_obj=MagicMock(),
api_base="https://my-endpoint.openai.azure.com",
api_key=None,
api_version="2024-10-01-preview",
azure_ad_token="my-entra-token",
)
headers = mock_ws_connect.call_args.kwargs["additional_headers"]
assert headers == {"Authorization": "Bearer my-entra-token"}
def test_get_auth_headers_prefers_api_key_and_never_sends_both():
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
assert AzureOpenAIRealtime.get_auth_headers(api_key="test-key", azure_ad_token="my-entra-token") == {
"api-key": "test-key"
}
def test_get_auth_headers_without_credentials_raises():
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
with pytest.raises(ValueError, match="Missing Azure credentials"):
AzureOpenAIRealtime.get_auth_headers(api_key=None, azure_ad_token=None)
@pytest.mark.asyncio
async def test_arealtime_resolves_azure_ad_token_when_no_api_key(monkeypatch):
"""
`_arealtime` must resolve an Azure AD token (managed identity, service principal, etc.)
and forward it to the handler when the deployment has no api_key.
Regression test for https://github.com/BerriAI/litellm/issues/34654
"""
from litellm.realtime_api import main as realtime_main
mock_async_realtime = AsyncMock()
monkeypatch.setattr(realtime_main, "azure_realtime", MagicMock(async_realtime=mock_async_realtime))
monkeypatch.setattr(
realtime_main,
"get_llm_provider",
lambda model, api_base=None, api_key=None: (
"gpt-realtime-whisper",
"azure",
None,
"https://my-endpoint.openai.azure.com",
),
)
monkeypatch.delenv("AZURE_API_KEY", raising=False)
captured_params = {}
def fake_get_azure_ad_token(litellm_params):
captured_params["tenant_id"] = litellm_params.get("tenant_id")
return "my-entra-token"
monkeypatch.setattr(realtime_main, "get_azure_ad_token", fake_get_azure_ad_token)
await realtime_main._arealtime(
model="azure/gpt-realtime-whisper",
websocket=MagicMock(),
api_version="2024-10-01-preview",
litellm_logging_obj=MagicMock(),
tenant_id="my-tenant",
client_id="my-client",
client_secret="my-secret",
)
assert mock_async_realtime.call_args.kwargs["azure_ad_token"] == "my-entra-token"
assert captured_params["tenant_id"] == "my-tenant"
@pytest.mark.asyncio
async def test_arealtime_does_not_resolve_azure_ad_token_when_api_key_present(monkeypatch):
from litellm.realtime_api import main as realtime_main
mock_async_realtime = AsyncMock()
monkeypatch.setattr(realtime_main, "azure_realtime", MagicMock(async_realtime=mock_async_realtime))
monkeypatch.setattr(
realtime_main,
"get_llm_provider",
lambda model, api_base=None, api_key=None: (
"gpt-realtime-whisper",
"azure",
"test-key",
"https://my-endpoint.openai.azure.com",
),
)
def fail_get_azure_ad_token(litellm_params):
raise AssertionError("should not resolve an AD token when an api_key is configured")
monkeypatch.setattr(realtime_main, "get_azure_ad_token", fail_get_azure_ad_token)
await realtime_main._arealtime(
model="azure/gpt-realtime-whisper",
websocket=MagicMock(),
api_key="test-key",
api_version="2024-10-01-preview",
litellm_logging_obj=MagicMock(),
)
assert mock_async_realtime.call_args.kwargs["azure_ad_token"] is None
@pytest.mark.asyncio
async def test_realtime_health_check_uses_bearer_token_when_no_api_key(monkeypatch):
"""
An Entra ID-only realtime deployment must also pass its realtime health check.
Regression test for https://github.com/BerriAI/litellm/issues/34654
"""
from litellm.realtime_api import main as realtime_main
connect_calls = []
monkeypatch.setattr(
realtime_main,
"get_azure_ad_token",
lambda litellm_params: "my-entra-token",
)
def fake_connect(url, **kwargs):
connect_calls.append(kwargs)
return _DummyAsyncContextManager(MagicMock())
monkeypatch.setattr("websockets.connect", fake_connect)
assert (
await realtime_main._realtime_health_check(
model="gpt-realtime-whisper",
custom_llm_provider="azure",
api_key=None,
api_base="https://my-endpoint.openai.azure.com",
api_version="2024-10-01-preview",
model_params={"tenant_id": "my-tenant"},
)
is True
)
assert connect_calls[0]["additional_headers"] == {"Authorization": "Bearer my-entra-token"}
@pytest.mark.asyncio
async def test_arealtime_forwards_deployment_azure_ad_token(monkeypatch):
"""
The router binds a deployment's `azure_ad_token` to `_arealtime`'s named parameter rather than
**kwargs, so it must still reach the handler.
Regression test for https://github.com/BerriAI/litellm/issues/34654
"""
from litellm.realtime_api import main as realtime_main
mock_async_realtime = AsyncMock()
monkeypatch.setattr(realtime_main, "azure_realtime", MagicMock(async_realtime=mock_async_realtime))
monkeypatch.setattr(
realtime_main,
"get_llm_provider",
lambda model, api_base=None, api_key=None: (
"gpt-realtime-whisper",
"azure",
None,
"https://my-endpoint.openai.azure.com",
),
)
monkeypatch.delenv("AZURE_API_KEY", raising=False)
monkeypatch.setattr(realtime_main.litellm, "api_key", None)
await realtime_main._arealtime(
model="azure/gpt-realtime-whisper",
websocket=MagicMock(),
api_version="2024-10-01-preview",
azure_ad_token="deployment-entra-token",
litellm_logging_obj=MagicMock(),
)
assert mock_async_realtime.call_args.kwargs["azure_ad_token"] == "deployment-entra-token"

View file

@ -284,6 +284,71 @@ def test_reasoning_with_forced_tool_choice_switches_to_auto():
assert optional_params["tool_choice"] == {"auto": {}}
@pytest.mark.parametrize(
"model",
[
"us.openai.gpt-5.6-sol",
"global.openai.gpt-5.6-terra",
"bedrock/converse/us.openai.gpt-5.6-luna",
],
)
def test_reasoning_effort_maps_to_reasoning_effort_for_openai_gpt5_converse(model, local_model_cost_map):
"""OpenAI GPT-5.x on Bedrock Converse routes reasoning_effort to
``additionalModelRequestFields.reasoning.effort`` rather than Anthropic ``thinking``."""
config = AmazonConverseConfig()
assert "reasoning_effort" in config.get_supported_openai_params(model=model)
optional_params = config.map_openai_params(
non_default_params={"reasoning_effort": "high"},
optional_params={},
model=model,
drop_params=False,
)
assert optional_params["reasoning"] == {"effort": "high"}
assert "thinking" not in optional_params
assert "reasoning_effort" not in optional_params
_, additional_request_params, _, _ = config._prepare_request_params(optional_params, model)
assert additional_request_params["reasoning"] == {"effort": "high"}
assert "thinking" not in additional_request_params
@pytest.mark.parametrize(
"model",
[
"us.openai.gpt-5.6-sol",
"bedrock/converse/global.openai.gpt-5.6-luna",
],
)
def test_openai_gpt5_converse_never_forwards_thinking(model, local_model_cost_map):
"""GPT-5.x on Converse must never send Anthropic ``thinking``/``output_config`` (Bedrock rejects them).
Regression: ``thinking`` is not advertised as supported, and even when supplied alongside
``reasoning_effort`` in either order it never survives into the request."""
config = AmazonConverseConfig()
supported = config.get_supported_openai_params(model=model)
assert "thinking" not in supported
assert "output_config" not in supported
thinking_block = {"type": "enabled", "budget_tokens": 2048}
for non_default_params in (
{"reasoning_effort": "high", "thinking": thinking_block},
{"thinking": thinking_block, "reasoning_effort": "high"},
):
optional_params = config.map_openai_params(
non_default_params=dict(non_default_params),
optional_params={},
model=model,
drop_params=False,
)
_, additional_request_params, _, _ = config._prepare_request_params(optional_params, model)
assert additional_request_params["reasoning"] == {"effort": "high"}
assert "thinking" not in additional_request_params
@pytest.mark.parametrize(
"model",
[

View file

@ -293,15 +293,16 @@ def test_bedrock_gpt_5_6_advertises_only_converse_supported_features(
@pytest.mark.parametrize("profile", GPT_5_6_PROFILES, ids=lambda p: p.model_id)
def test_bedrock_gpt_5_6_offers_tools_but_not_reasoning(profile, local_model_cost_map):
"""Converse rejects the Anthropic-shaped thinking block LiteLLM emits for
reasoning_effort, so neither reasoning param may be offered yet, while the tool
params these models do accept must be."""
def test_bedrock_gpt_5_6_offers_tools_and_reasoning_effort_but_not_thinking(profile, local_model_cost_map):
"""GPT-5.x on Converse maps reasoning_effort to reasoning.effort, so reasoning_effort
is offered while the Anthropic-only thinking/output_config are not, alongside the tool
params these models accept."""
supported = AmazonConverseConfig().get_supported_openai_params(
model=f"bedrock/{profile.model_id}"
)
assert "tools" in supported
assert "tool_choice" in supported
assert "reasoning_effort" not in supported
assert "reasoning_effort" in supported
assert "thinking" not in supported
assert "output_config" not in supported

View file

@ -1673,7 +1673,7 @@ class TestBedrockMantleResponsesPricing:
assert info["input_cost_per_token"] == pytest.approx(5.5e-06)
assert info["output_cost_per_token"] == pytest.approx(3.3e-05)
assert info["cache_read_input_token_cost"] == pytest.approx(5.5e-07)
assert info["max_input_tokens"] == 272000
assert info["max_input_tokens"] == 1050000
def test_gpt_5_4_pricing_and_mode(self, local_cost_map):
info = litellm.get_model_info("bedrock_mantle/openai.gpt-5.4")
@ -1681,7 +1681,7 @@ class TestBedrockMantleResponsesPricing:
assert info["input_cost_per_token"] == pytest.approx(2.75e-06)
assert info["output_cost_per_token"] == pytest.approx(1.65e-05)
assert info["cache_read_input_token_cost"] == pytest.approx(2.75e-07)
assert info["max_input_tokens"] == 272000
assert info["max_input_tokens"] == 1050000
@pytest.mark.parametrize(
"model, input_cost, cache_creation_cost, cache_read_cost, output_cost",
@ -1753,13 +1753,14 @@ def _repo_cost_map(map_name: str) -> dict[str, dict[str, object]]:
return json.loads(paths[map_name].read_text())
class TestGpt56MantleRegistryEntries:
"""Locks the gpt-5.6 frontier entries to Bedrock Mantle's live behavior.
class TestMantleGptRegistryEntries:
"""Locks the OpenAI GPT entries to Bedrock Mantle's live behavior.
Mantle enforces a 1,050,000-token prompt maximum for gpt-5.6 sol/terra/luna
(oversize requests 400 with "prompt tokens (N) exceed model maximum
(1050000)", and a 1,030,590-token request completes), matching the OpenAI
Bedrock guide. mode must stay "responses": Mantle's native
and for gpt-5.5 and gpt-5.4 (oversize requests 400 with "prompt tokens (N)
exceed model maximum (1050000)", and a 1,030,590-token request completes
on every one of them), while the AWS model cards still quote 272K for
gpt-5.5 and gpt-5.4. mode must stay "responses": Mantle's native
/v1/chat/completions rejects function tools unless reasoning_effort is
"none", so chat traffic has to keep bridging to the Responses API
(see the responses_api_bridge tests above).
@ -1781,3 +1782,18 @@ class TestGpt56MantleRegistryEntries:
assert entry["mode"] == "responses"
assert entry["use_openai_responses_path"] is True
assert entry["supported_endpoints"] == ["/v1/chat/completions", "/v1/responses"]
@pytest.mark.parametrize("map_name", ("root", "bundled_backup"))
@pytest.mark.parametrize(
"key",
(
"bedrock_mantle/openai.gpt-5.5",
"bedrock_mantle/openai.gpt-5.4",
),
)
def test_gpt_55_and_54_entries_match_mantle_enforced_limits(self, map_name, key):
entry = _repo_cost_map(map_name)[key]
assert entry["max_input_tokens"] == 1050000
assert entry["max_output_tokens"] == 128000
assert entry["mode"] == "responses"
assert entry["use_openai_responses_path"] is True

View file

@ -3,9 +3,9 @@
The envelope is the single client-held bearer carrying both a litellm identity and the
encrypted upstream grant, with zero server-side storage. These tests pin the security
contract: an envelope opens only under the exact keys that minted it, tampering with any
signed byte is detected, expiry is enforced against the injected clock (capped by the
module TTL ceiling), oversized envelopes are rejected rather than truncated, and no
error value, model repr, or raised exception ever contains the inner access token.
signed byte is detected, expiry is enforced against the injected clock and provider
lifetime, oversized envelopes are rejected rather than truncated, and no error value,
model repr, or raised exception ever contains the inner access token.
"""
import base64
@ -30,6 +30,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import
DecryptFailed,
EnvelopeIdentity,
EnvelopeKeys,
EnvelopeLifetimeUnrepresentable,
EnvelopeMintError,
EnvelopeTooLarge,
Expired,
MalformedPayload,
@ -159,6 +161,20 @@ def test_claim_layout_and_no_plaintext_token_in_envelope():
assert _REFRESH_TOKEN not in json.dumps(claims)
def test_unrepresentable_access_lifetime_is_a_typed_mint_error():
grant = UpstreamTokenGrant(
access_token=SecretStr(_ACCESS_TOKEN),
token_type="Bearer",
expires_in=10**30,
)
result = mint_envelope(_IDENTITY, grant, _KEYS, _NOW)
assert isinstance(result, EnvelopeLifetimeUnrepresentable)
assert result.tag == "envelope_lifetime_unrepresentable"
assert result.expires_in == 10**30
def _refresh_credential() -> RefreshCredential:
return RefreshCredential(refresh_token=SecretStr(_REFRESH_TOKEN), scope="read:tools", expires_in=None)
@ -243,11 +259,11 @@ def test_refresh_envelope_never_leaks_the_refresh_token_in_plaintext():
"expires_in, expected_ttl",
[
(600, 600),
(MAX_ENVELOPE_TTL_SECONDS + 82800, MAX_ENVELOPE_TTL_SECONDS),
(MAX_ENVELOPE_TTL_SECONDS + 82800, MAX_ENVELOPE_TTL_SECONDS + 82800),
(None, MAX_ENVELOPE_TTL_SECONDS),
],
)
def test_exp_is_min_of_upstream_expires_in_and_cap(expires_in, expected_ttl):
def test_exp_matches_upstream_lifetime_or_uses_missing_lifetime_fallback(expires_in: int | None, expected_ttl: int):
grant = UpstreamTokenGrant(access_token=SecretStr(_ACCESS_TOKEN), token_type="Bearer", expires_in=expires_in)
sealed = mint_envelope(_IDENTITY, grant, _KEYS, _NOW)
assert isinstance(sealed, SealedEnvelope)
@ -261,13 +277,13 @@ def test_expiry_honored_against_injected_clock():
assert isinstance(open_envelope(token, _KEYS, _NOW + timedelta(seconds=601)), Expired)
def test_ttl_cap_enforced_on_open_even_when_upstream_token_lives_longer():
def test_upstream_token_lifetime_is_enforced_on_open():
grant = UpstreamTokenGrant(access_token=SecretStr(_ACCESS_TOKEN), token_type="Bearer", expires_in=86400)
token = _sealed_token(grant)
just_before_cap = _NOW + timedelta(seconds=MAX_ENVELOPE_TTL_SECONDS - 1)
at_cap = _NOW + timedelta(seconds=MAX_ENVELOPE_TTL_SECONDS)
assert isinstance(open_envelope(token, _KEYS, just_before_cap), OpenedEnvelope)
assert isinstance(open_envelope(token, _KEYS, at_cap), Expired)
just_before_expiry = _NOW + timedelta(seconds=86399)
at_expiry = _NOW + timedelta(seconds=86400)
assert isinstance(open_envelope(token, _KEYS, just_before_expiry), OpenedEnvelope)
assert isinstance(open_envelope(token, _KEYS, at_expiry), Expired)
def test_tampering_any_payload_or_signature_byte_is_bad_signature():
@ -420,7 +436,7 @@ def test_decryptable_blob_that_is_not_a_grant_is_malformed_payload():
assert isinstance(open_envelope(forged, _KEYS, _NOW), MalformedPayload)
def _mint_with_token_len(n: int) -> SealedEnvelope | EnvelopeTooLarge:
def _mint_with_token_len(n: int) -> SealedEnvelope | EnvelopeMintError:
grant = UpstreamTokenGrant(access_token=SecretStr("a" * n), token_type="Bearer")
return mint_envelope(_IDENTITY, grant, _KEYS, _NOW)

View file

@ -5539,6 +5539,30 @@ async def test_bridge_envelope_too_large_upstream_token_is_502():
assert json.loads(response.body)["error"] == "server_error"
@pytest.mark.asyncio
async def test_bridge_envelope_unrepresentable_upstream_lifetime_is_502():
from litellm.types.mcp import MCPAuth
server = _bridge_server(auth_type=MCPAuth.oauth_delegate)
upstream = {
"access_token": "UPSTREAM-SECRET-TOKEN",
"token_type": "Bearer",
"expires_in": 10**30,
}
response = await _exchange_for_bridge_server(
server,
upstream,
key_hash="hashed-litellm-key-77",
)
assert response.status_code == 502
assert json.loads(response.body) == {
"error": "server_error",
"error_description": "the upstream token response reports an unrepresentable lifetime",
}
@pytest.mark.asyncio
async def test_bridge_access_envelope_never_carries_upstream_refresh_token():
"""The upstream refresh token is never sealed into the ACCESS envelope, the bearer forwarded upstream

View file

@ -1948,14 +1948,14 @@ class FakePodLockManager:
if self.redis_cache is not None:
self.redis_cache.async_get_cache = AsyncMock(return_value="another-pod" if held_by_other else None)
self._acquired = acquired
self.acquire_calls: List[Dict[str, Any]] = []
self.acquire_calls: List[Dict[str, str | int | None]] = []
self.release_calls: List[str] = []
@staticmethod
def get_redis_lock_key(cronjob_id: str) -> str:
return f"cronjob_lock:{cronjob_id}"
async def acquire_lock(self, cronjob_id: str, ttl: Any = None) -> bool:
async def acquire_lock(self, cronjob_id: str, ttl: int | None = None) -> bool:
self.acquire_calls.append({"cronjob_id": cronjob_id, "ttl": ttl})
return self._acquired

View file

@ -6,11 +6,14 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import respx
from fastapi import FastAPI
from fastapi.testclient import TestClient
from prisma.errors import ClientNotConnectedError, HTTPClientClosedError, PrismaError
import litellm
import litellm.proxy.health_endpoints._health_endpoints as _health_endpoints_module
from litellm.litellm_core_utils.health_check_helpers import TEST_IMAGE_BASE64
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -2674,3 +2677,38 @@ class TestNoRedisWarning:
):
details = await _health_endpoints_module._get_health_readiness_details()
assert details["show_no_redis_warning"] is False
def test_test_model_connection_accepts_image_edit_mode(monkeypatch):
"""
Regression: /health/test_connection rejected mode=image_edit with a 422
before image_edit was added to its mode Literal, breaking the UI Test
Connection button for image edit deployments.
"""
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
app = FastAPI()
app.include_router(_health_endpoints_module.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
client = TestClient(app)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: the endpoint reads the proxy-global DB client and 500s when it is None; it has no injection seam
respx.mock(assert_all_called=True) as respx_mock,
):
respx_mock.post(host="api.openai.com", path="/v1/images/edits").respond(
json={"created": 1700000000, "data": [{"b64_json": TEST_IMAGE_BASE64}]}
)
response = client.post(
"/health/test_connection",
json={
"mode": "image_edit",
"litellm_params": {"model": "openai/gpt-image-2", "api_key": "sk-test"},
},
)
assert response.status_code == 200, response.text
assert response.json()["status"] == "success"

View file

@ -780,3 +780,132 @@ class TestBlockRequestsForModelsWithoutPricing:
assert response.status_code == 500
assert "error" in response.json()["detail"]
AN_ALIAS = "onprem/alias"
AN_UNDERLYING_MODEL = "vendor/model"
A_MAPPED_MODEL = "openai/mapped-only-model"
INPUT_TOKENS = 1000
OUTPUT_TOKENS = 500
def _router_pricing(**pricing: float) -> MagicMock:
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
{
"model_name": AN_ALIAS,
"litellm_params": {
"model": AN_UNDERLYING_MODEL,
"custom_llm_provider": "openai",
**pricing,
},
"model_info": {},
}
]
return mock_router
async def _estimate(mock_router: MagicMock | None, model: str = AN_ALIAS, **overrides: int):
from litellm.proxy._types import CostEstimateRequest
from litellm.proxy.management_endpoints.cost_tracking_settings import estimate_cost
request = CostEstimateRequest(
model=model,
input_tokens=INPUT_TOKENS,
output_tokens=OUTPUT_TOKENS,
**overrides,
)
with patch( # test-quality-ok: proxy_server module global is the endpoint's only injection point
"litellm.proxy.proxy_server.llm_router", mock_router
):
return await estimate_cost(request=request, user_api_key_dict=MagicMock())
class TestEstimateCostPartiallyPricedDeployments:
@pytest.mark.asyncio
async def test_a_deployment_that_prices_only_input_bills_output_at_zero(self):
response = await _estimate(_router_pricing(input_cost_per_token=0.000001))
assert response.input_cost_per_token == pytest.approx(0.000001)
assert response.output_cost_per_token == 0.0
assert response.cost_per_request == pytest.approx(0.001)
@pytest.mark.asyncio
async def test_a_deployment_that_prices_only_output_bills_input_at_zero(self):
response = await _estimate(_router_pricing(output_cost_per_token=0.000002))
assert response.input_cost_per_token == 0.0
assert response.output_cost_per_token == pytest.approx(0.000002)
assert response.cost_per_request == pytest.approx(0.001)
@pytest.mark.asyncio
async def test_a_model_priced_only_by_the_cost_map_reports_that_price_and_provider(self, monkeypatch):
monkeypatch.setitem(
litellm.model_cost,
A_MAPPED_MODEL,
{
"input_cost_per_token": 0.000005,
"output_cost_per_token": 0.000006,
"litellm_provider": "openai",
"mode": "chat",
},
)
response = await _estimate(None, model=A_MAPPED_MODEL)
assert response.input_cost_per_token == pytest.approx(0.000005)
assert response.output_cost_per_token == pytest.approx(0.000006)
assert response.provider == "openai"
class TestEstimateCostPeriodTotals:
@pytest.mark.asyncio
async def test_zero_requests_a_day_reports_no_daily_cost_rather_than_zero(self):
response = await _estimate(
_router_pricing(input_cost_per_token=0.000001, output_cost_per_token=0.000002),
num_requests_per_day=0,
)
assert response.daily_cost is None
assert response.daily_input_cost is None
assert response.daily_output_cost is None
@pytest.mark.asyncio
async def test_daily_totals_scale_every_component_by_the_request_count(self):
response = await _estimate(
_router_pricing(input_cost_per_token=0.000001, output_cost_per_token=0.000002),
num_requests_per_day=100,
)
assert response.input_cost_per_request == pytest.approx(0.001)
assert response.output_cost_per_request == pytest.approx(0.001)
assert response.daily_input_cost == pytest.approx(0.1)
assert response.daily_output_cost == pytest.approx(0.1)
assert response.daily_cost == pytest.approx(0.2)
@pytest.mark.asyncio
async def test_a_month_and_a_day_are_totalled_from_their_own_request_counts(self):
response = await _estimate(
_router_pricing(input_cost_per_token=0.000001, output_cost_per_token=0.000002),
num_requests_per_day=100,
num_requests_per_month=3000,
)
assert response.daily_cost == pytest.approx(0.2)
assert response.monthly_cost == pytest.approx(6.0)
assert response.monthly_input_cost == pytest.approx(3.0)
assert response.monthly_output_cost == pytest.approx(3.0)
@pytest.mark.asyncio
async def test_a_configured_margin_is_totalled_per_period_like_the_other_components(self, monkeypatch):
monkeypatch.setattr(litellm, "cost_margin_config", {"openai": 0.10})
response = await _estimate(
_router_pricing(input_cost_per_token=0.000001, output_cost_per_token=0.000002),
num_requests_per_day=100,
)
assert response.margin_cost_per_request == pytest.approx(0.0002)
assert response.cost_per_request == pytest.approx(0.0022)
assert response.daily_margin_cost == pytest.approx(0.02)
assert response.daily_cost == pytest.approx(0.22)

View file

@ -33,6 +33,15 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
TeamMappings,
)
_SSO_PROVIDER_ENV_VARS = (
"DISABLE_ADMIN_UI",
"MICROSOFT_CLIENT_ID",
"GOOGLE_CLIENT_ID",
"GENERIC_CLIENT_ID",
"SAML_IDP_METADATA_URL",
"SAML_IDP_METADATA_XML",
)
def _wire_team_create_tx(prisma_client):
"""`/team/new` inserts the team and mirrors it onto the access groups in one transaction,
@ -2796,10 +2805,15 @@ class TestCLIKeyRegenerationFlow:
mock_request.base_url = "https://proxy.example.com/"
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {"poll_secret_hash": "h"}
env_without_sso_providers = {
name: value
for name, value in os.environ.items()
if name not in _SSO_PROVIDER_ENV_VARS
}
async def drive(enabled: bool):
with (
patch.dict(os.environ, {}, clear=True),
patch.dict(os.environ, env_without_sso_providers, clear=True),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
@ -2825,15 +2839,13 @@ class TestCLIKeyRegenerationFlow:
return_value=None,
) as mock_get_cli_state,
):
try:
await google_login(
request=mock_request,
source="litellm-cli",
key="cli-validsessionkey123456",
user_code="WXYZ-2345",
)
except Exception:
pass
await google_login(
request=mock_request,
source="litellm-cli",
key="cli-validsessionkey123456",
user_code="WXYZ-2345",
)
assert mock_get_cli_state.called
return mock_get_cli_state.call_args.kwargs["user_code"]
assert await drive(enabled=True) == "WXYZ-2345"

View file

@ -1,3 +1,4 @@
import base64
import contextlib
import json
import os
@ -16,6 +17,7 @@ from starlette.datastructures import FormData
import litellm
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
BaseOpenAIPassThroughHandler,
RouteChecks,
@ -37,6 +39,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
vllm_proxy_route,
)
from litellm.proxy._types import LitellmUserRoles, SpecialHeaders, UserAPIKeyAuth
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
@ -592,7 +595,7 @@ class TestVertexAIPassThroughHandler:
"method": "POST",
"path": endpoint,
"headers": [
(b"authorization", b"Bearer test-creds"),
(b"authorization", b"Bearer sk-test-creds"),
],
}
)
@ -617,7 +620,7 @@ class TestVertexAIPassThroughHandler:
):
mock_ensure_token.return_value = ("test-auth-header", test_project)
mock_get_token.return_value = (test_token, "")
mock_auth.return_value = MagicMock()
mock_auth.return_value = UserAPIKeyAuth(api_key="sk-test-creds")
with pytest.raises(HTTPException) as exc_info:
await vertex_proxy_route(
@ -3342,14 +3345,14 @@ class TestVertexRawPredictStreamingClassification:
),
mock.patch(f"{module}.create_pass_through_route", side_effect=fake_create_pass_through_route),
mock.patch(f"{module}.get_litellm_virtual_key", return_value="Bearer test-key"),
mock.patch(f"{module}.user_api_key_auth", new=AsyncMock(return_value={"api_key": "test-key"})),
mock.patch(f"{module}.user_api_key_auth", new=AsyncMock(return_value=UserAPIKeyAuth(api_key="test-key"))),
mock.patch(f"{module}.get_vertex_pass_through_handler", return_value=mock_handler),
):
await vertex_proxy_route(
endpoint=endpoint,
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(token="test-key"),
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
)
assert captured, "create_pass_through_route was never called"
@ -3447,6 +3450,13 @@ def test_is_passthrough_request_streaming_tolerates_non_object_bodies(request_bo
assert is_passthrough_request_streaming(request_body) is expected
def _unsigned_jwt(claims: Mapping[str, str]) -> str:
def segment(payload: Mapping[str, str]) -> str:
return base64.urlsafe_b64encode(json.dumps(dict(payload)).encode()).rstrip(b"=").decode()
return ".".join((segment({"alg": "RS256", "typ": "JWT"}), segment(claims), "c2lnbmF0dXJl"))
class TestVertexCredentiallessPassthroughVirtualKeyLeak:
"""Regression coverage for LIT-5997.
@ -3466,6 +3476,12 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak:
genuine bring-your-own Google credential that must still pass through. The
by-value strip also covers a virtual key sent in the operator-configured
``general_settings.litellm_key_header_name``, whatever that header is named.
The by-value strip keys off what actually authenticated the caller (the
master key, or the LiteLLM key whose hash ``user_api_key_auth`` resolved as
``api_key``), never off header precedence: a custom auth or JWT that
authenticated the caller without consuming ``Authorization`` leaves the
caller's own Google token there, and it must keep flowing.
"""
VKEY = "sk-litellm-victim-key"
@ -3475,8 +3491,14 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak:
)
async def _run(
self, monkeypatch, headers: list[tuple[bytes, bytes]]
self,
monkeypatch,
headers: list[tuple[bytes, bytes]],
authenticated: UserAPIKeyAuth | None = None,
master_key: str | None = "sk-master-1234",
) -> tuple[HTTPException | None, dict | None]:
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", master_key)
caller: Final = authenticated if authenticated is not None else UserAPIKeyAuth(api_key=self.VKEY)
from litellm.proxy.pass_through_endpoints.passthrough_endpoint_router import (
PassthroughEndpointRouter,
)
@ -3509,7 +3531,7 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak:
raised: HTTPException | None = None
with (
mock.patch(f"{module}.create_pass_through_route", side_effect=fake_create_pass_through_route),
mock.patch(f"{module}.user_api_key_auth", new=AsyncMock(return_value=UserAPIKeyAuth(token="hashed"))),
mock.patch(f"{module}.user_api_key_auth", new=AsyncMock(return_value=caller)),
mock.patch(f"{module}.get_vertex_pass_through_handler", return_value=mock_handler),
):
try:
@ -3517,7 +3539,7 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak:
endpoint=self.ENDPOINT,
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(token="hashed"),
user_api_key_dict=caller,
)
except HTTPException as exc:
raised = exc
@ -3767,6 +3789,135 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak:
assert forwarded is None, "a virtual key in the mapped-route litellm_user_api_key header must be dropped, not forwarded"
assert raised is not None and raised.status_code == 401
GOOGLE_OAUTH_TOKEN = "ya29.byo-google-oauth-token"
LITELLM_JWT_CLAIMS = MappingProxyType({"sub": "jwt-subject", "iss": "https://idp.example.com"})
LITELLM_JWT = _unsigned_jwt(LITELLM_JWT_CLAIMS)
GOOGLE_SERVICE_ACCOUNT_JWT = _unsigned_jwt(
{
"sub": "vertex-caller@my-proj.iam.gserviceaccount.com",
"iss": "vertex-caller@my-proj.iam.gserviceaccount.com",
"aud": "https://aiplatform.googleapis.com/",
}
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("master_key", "authenticated"),
[
pytest.param(
"sk-master-1234",
UserAPIKeyAuth(api_key="best-api-key-ever", user_role=LitellmUserRoles.PROXY_ADMIN),
id="custom-auth-returning-its-own-identifier",
),
pytest.param(
"sk-master-1234",
UserAPIKeyAuth(api_key=None, user_id="jwt-subject", jwt_claims=dict(LITELLM_JWT_CLAIMS)),
id="jwt-auth",
),
pytest.param(None, UserAPIKeyAuth(api_key=GOOGLE_OAUTH_TOKEN), id="no-master-key-echoes-raw-header"),
],
)
async def test_google_token_in_authorization_is_forwarded_when_auth_did_not_consume_it(
self, monkeypatch, master_key: str | None, authenticated: UserAPIKeyAuth
):
raised, forwarded = await self._run(
monkeypatch,
[
(b"authorization", f"Bearer {self.GOOGLE_OAUTH_TOKEN}".encode()),
(b"content-type", b"application/json"),
],
authenticated=authenticated,
master_key=master_key,
)
assert raised is None, f"the caller's own Google token must not be mistaken for a LiteLLM key: {raised}"
assert forwarded is not None
assert forwarded.get("authorization") == f"Bearer {self.GOOGLE_OAUTH_TOKEN}"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("credential", "authenticated"),
[
pytest.param("modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential"),
pytest.param(
LITELLM_JWT,
UserAPIKeyAuth(api_key=LITELLM_JWT, user_id="jwt-subject"),
id="custom-auth-echoing-jwt",
),
pytest.param(
LITELLM_JWT,
UserAPIKeyAuth(api_key=None, user_id="jwt-subject", jwt_claims=dict(LITELLM_JWT_CLAIMS)),
id="jwt-auth",
),
pytest.param(
LITELLM_JWT,
UserAPIKeyAuth(
api_key=None,
user_id="jwt-subject",
jwt_claims={
**LITELLM_JWT_CLAIMS,
JWTHandler.LITELLM_JWT_ISSUER_CLAIM: "https://idp.example.com",
JWTHandler.LITELLM_USER_ID_CLAIM: "jwt-subject",
},
),
id="multi-issuer-jwt-auth-normalized-claims",
),
],
)
async def test_non_sk_litellm_credential_that_authenticated_is_rejected_not_forwarded(
self, monkeypatch, credential: str, authenticated: UserAPIKeyAuth
):
raised, forwarded = await self._run(
monkeypatch,
[(b"authorization", f"Bearer {credential}".encode()), (b"content-type", b"application/json")],
authenticated=authenticated,
)
assert forwarded is None, "the credential that authenticated the caller must never reach the upstream forwarder"
assert raised is not None and raised.status_code == 401
@pytest.mark.asyncio
async def test_jwt_authenticated_caller_keeps_a_different_byo_google_jwt(self, monkeypatch):
raised, forwarded = await self._run(
monkeypatch,
[
(b"x-litellm-api-key", self.LITELLM_JWT.encode()),
(b"authorization", f"Bearer {self.GOOGLE_SERVICE_ACCOUNT_JWT}".encode()),
(b"content-type", b"application/json"),
],
authenticated=UserAPIKeyAuth(api_key=None, user_id="jwt-subject", jwt_claims=dict(self.LITELLM_JWT_CLAIMS)),
)
assert raised is None, f"a Google JWT that is not the one that authenticated must keep flowing: {raised}"
assert forwarded is not None
assert forwarded.get("authorization") == f"Bearer {self.GOOGLE_SERVICE_ACCOUNT_JWT}"
assert "x-litellm-api-key" not in forwarded
@pytest.mark.asyncio
async def test_master_key_in_authorization_alone_is_rejected(self, monkeypatch):
raised, forwarded = await self._run(
monkeypatch,
[(b"authorization", b"Bearer sk-master-1234"), (b"content-type", b"application/json")],
authenticated=UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert forwarded is None, "the master key must never reach the upstream forwarder"
assert raised is not None and raised.status_code == 401
@pytest.mark.asyncio
async def test_master_key_is_stripped_and_byo_x_goog_api_key_forwards(self, monkeypatch):
raised, forwarded = await self._run(
monkeypatch,
[
(b"authorization", b"Bearer sk-master-1234"),
(b"x-goog-api-key", b"AIza-real-google-api-key"),
(b"content-type", b"application/json"),
],
authenticated=UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert raised is None
assert forwarded is not None
assert forwarded.get("x-goog-api-key") == "AIza-real-google-api-key"
assert "authorization" not in forwarded
assert "sk-master-1234" not in " ".join(f"{name}:{value}" for name, value in forwarded.items())
class TestGetAzureAISearchIndexFromEndpoint:
"""The operable index is only the segment right after ``indexes``.

View file

@ -6,7 +6,6 @@ non-streaming pass-through responses. Addresses issue #20270.
"""
import json
import sys
from contextlib import ExitStack
from unittest.mock import AsyncMock, MagicMock, patch
@ -66,21 +65,6 @@ def _make_mock_request():
return mock_request
def _ensure_proxy_server_mock():
"""Insert a mock proxy_server module if the real one can't import."""
key = "litellm.proxy.proxy_server"
if key not in sys.modules:
mock_mod = MagicMock()
mock_mod.proxy_logging_obj = MagicMock()
sys.modules[key] = mock_mod
import litellm.proxy
if not hasattr(litellm.proxy, "proxy_server"):
litellm.proxy.proxy_server = sys.modules[key]
_ensure_proxy_server_mock()
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
pass_through_request,
)

View file

@ -2,6 +2,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_base_vertex_proxy_route,
)
@ -323,6 +324,7 @@ async def test_vertex_passthrough_forwards_anthropic_beta_header():
vertex_location="us-central1",
base_target_url="https://us-central1-aiplatform.googleapis.com",
get_vertex_pass_through_handler=mock_handler,
user_api_key_dict=UserAPIKeyAuth(api_key="sk-litellm-secret-key"),
)
# Verify that allowlisted headers are preserved
@ -417,6 +419,7 @@ async def test_vertex_passthrough_does_not_forward_litellm_auth_token():
vertex_location="us-central1",
base_target_url="https://us-central1-aiplatform.googleapis.com",
get_vertex_pass_through_handler=mock_handler,
user_api_key_dict=UserAPIKeyAuth(api_key="sk-litellm-secret-key"),
)
# The ONLY Authorization header should be the Vertex token

View file

@ -1,4 +1,7 @@
import json
from pathlib import Path
import pytest
@ -14,7 +17,13 @@ from litellm.cost_calculator import (
response_cost_calculator,
)
from litellm.types.llms.openai import OpenAIRealtimeStreamList
from litellm.types.utils import ModelInfo, ModelResponse, PromptTokensDetailsWrapper, Usage
from litellm.types.utils import (
CacheCreationTokenDetails,
ModelInfo,
ModelResponse,
PromptTokensDetailsWrapper,
Usage,
)
from litellm.utils import TranscriptionResponse
@ -2856,6 +2865,40 @@ def test_anthropic_fast_multiplier_only_on_models_with_fast_mode(_local_model_co
assert entry["provider_specific_entry"].get("fast") == expected_fast
@pytest.mark.parametrize(
"model",
["claude-sonnet-4-6", "claude-mythos-5", "claude-mythos-preview"],
)
def test_anthropic_us_data_residency_uplift_on_claude_4_6_and_later_models(
_local_model_cost_map, monkeypatch, model
):
"""
Anthropic bills every Claude 4.6+ model served with ``inference_geo="us"`` at
1.1x, and echoes that geo back in the response usage, so each of these real
cost-map entries has to carry the ``us`` multiplier or US-pinned traffic is
under-reported by 10%.
"""
from litellm.llms.anthropic.cost_calculation import (
cost_per_token as anthropic_cost_per_token,
)
from litellm.types.utils import Usage
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
def make_usage() -> "Usage":
return Usage(prompt_tokens=1_000, completion_tokens=100, total_tokens=1_100)
base_prompt_cost, base_completion_cost = anthropic_cost_per_token(model=model, usage=make_usage())
geo_usage = make_usage()
geo_usage.inference_geo = "us"
geo_prompt_cost, geo_completion_cost = anthropic_cost_per_token(model=model, usage=geo_usage)
assert base_prompt_cost > 0
assert geo_prompt_cost == pytest.approx(base_prompt_cost * 1.1)
assert geo_completion_cost == pytest.approx(base_completion_cost * 1.1)
def test_gemini_cache_tokens_details_no_negative_values():
"""
Test for Issue #18750: Negative text_tokens with Gemini caching
@ -3797,3 +3840,51 @@ def test_completion_cost_prices_anthropic_shaped_cache_read_tokens(_local_model_
)
assert cost == pytest.approx(3 * 4e-6 + 4014 * 4e-7 + 5 * 2e-5, rel=1e-9)
@pytest.mark.parametrize(
("model", "expected_1hr_rate"),
[("claude-3-haiku-20240307", 5e-07), ("claude-3-opus-20240229", 3e-05)],
)
def test_claude_3_one_hour_cache_writes_bill_at_double_input(
_local_model_cost_map, model: str, expected_1hr_rate: float
):
"""Regression: both models carried the Sonnet 1h cache-write rate (6e-06) instead of
2x their own input price, overbilling haiku 12x and underbilling opus 5x."""
usage = Usage(
prompt_tokens=1000,
completion_tokens=0,
total_tokens=1000,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=0,
cache_creation_tokens=1000,
cache_creation_token_details=CacheCreationTokenDetails(
ephemeral_5m_input_tokens=0, ephemeral_1h_input_tokens=1000
),
),
)
prompt_cost, _ = cost_per_token(model=model, usage_object=usage, custom_llm_provider="anthropic")
assert prompt_cost == pytest.approx(1000 * expected_1hr_rate, rel=1e-9)
def test_every_one_hour_cache_write_rate_is_double_its_input_rate():
"""Guard against pasting one model's 1h cache-write price onto another: every provider
LiteLLM tracks (Anthropic, Bedrock, Vertex, Azure) publishes the 1h write at 2x input."""
cost_map = json.loads(
(Path(__file__).parents[2] / "model_prices_and_context_window.json").read_text()
)
one_hour_prefix = "cache_creation_input_token_cost_above_1hr"
deviations = {
(name, key): (entry["input_cost_per_token" + key[len(one_hour_prefix) :]], entry[key])
for name, entry in cost_map.items()
if isinstance(entry, dict)
for key in entry
if key.startswith(one_hour_prefix)
and entry[key] != pytest.approx(2 * entry["input_cost_per_token" + key[len(one_hour_prefix) :]], rel=1e-9)
}
assert deviations == {}

View file

@ -1,13 +1,21 @@
import json
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import redis
import redis.asyncio as async_redis
from redis.credentials import CredentialProvider
import litellm
from litellm._redis import (
_async_auth_kwargs,
_get_redis_client_logic,
_get_redis_cluster_kwargs,
_get_redis_env_kwarg_mapping,
_get_redis_kwargs,
_get_redis_url_kwargs,
_pretty_print_redis_config,
get_redis_async_client,
get_redis_client,
get_redis_connection_pool,
@ -18,9 +26,69 @@ from litellm._redis_credential_provider import (
GCPIAMCredentialProvider,
_token_cache,
)
from litellm.caching.redis_cache import RedisCache
from litellm.caching.redis_cluster_cache import RedisClusterCache
from litellm.constants import REDIS_CLUSTER_HEALTH_CHECK_INTERVAL
class _StubCredentialProvider(CredentialProvider):
def __init__(self, token: str = "stub-token") -> None:
self._token = token
def get_credentials(self):
return (self._token,)
async def get_credentials_async(self):
return (self._token,)
class _HostileCredentialProvider(CredentialProvider):
def __init__(self, secret: str) -> None:
self._payload = secret
def get_credentials(self):
return (self._payload,)
async def get_credentials_async(self):
return (self._payload,)
def __repr__(self):
raise AssertionError("provider repr must never be invoked")
def __str__(self):
raise AssertionError("provider str must never be invoked")
def __reduce__(self):
raise AssertionError("provider must never be serialized")
def __getstate__(self):
raise AssertionError("provider state must never be inspected")
def _gcp_marker_callback() -> MagicMock:
callback = MagicMock()
callback._gcp_service_account = "projects/-/serviceAccounts/sa@project.iam.gserviceaccount.com"
return callback
@pytest.fixture
def clean_redis_environment(monkeypatch):
for var in (
"REDIS_URL",
"REDIS_CLUSTER_NODES",
"REDIS_SENTINEL_NODES",
*_get_redis_env_kwarg_mapping(),
):
monkeypatch.delenv(var, raising=False)
@pytest.fixture
def clear_llm_client_cache():
litellm.in_memory_llm_clients_cache.flush_cache()
yield
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.fixture(autouse=True)
def clear_gcp_iam_token_cache():
"""Reset the module-level GCP IAM token cache between tests."""
@ -29,6 +97,364 @@ def clear_gcp_iam_token_cache():
_token_cache.clear()
def test_redis_allowlists_include_credential_provider():
assert "credential_provider" in _get_redis_kwargs()
assert "credential_provider" in _get_redis_url_kwargs()
assert "credential_provider" in _get_redis_cluster_kwargs()
def test_credential_provider_is_not_environment_derived():
mapping = _get_redis_env_kwarg_mapping()
assert "REDIS_CREDENTIAL_PROVIDER" not in mapping
assert "credential_provider" not in mapping.values()
def test_sync_direct_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
client = get_redis_client(host="redis-host", port=6379, credential_provider=provider)
assert client.connection_pool.connection_kwargs["credential_provider"] is provider
def test_sync_direct_provider_supersedes_static_credentials(clean_redis_environment):
provider = _StubCredentialProvider()
client = get_redis_client(
host="redis-host",
port=6379,
username="redis-user",
password="redis-password",
credential_provider=provider,
)
connection = client.connection_pool.make_connection()
assert connection.credential_provider is provider
assert connection.username is None
assert connection.password is None
def test_sync_direct_provider_supersedes_environment_credentials(clean_redis_environment, monkeypatch):
provider = _StubCredentialProvider()
monkeypatch.setenv("REDIS_USERNAME", "redis-user")
monkeypatch.setenv("REDIS_PASSWORD", "redis-password")
client = get_redis_client(host="redis-host", port=6379, credential_provider=provider)
connection = client.connection_pool.make_connection()
assert connection.credential_provider is provider
assert connection.username is None
assert connection.password is None
def test_sync_url_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
client = get_redis_client(url="redis://redis-host:6379", credential_provider=provider)
assert client.connection_pool.connection_kwargs["credential_provider"] is provider
def test_async_direct_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
client = get_redis_async_client(host="redis-host", port=6379, credential_provider=provider)
assert client.connection_pool.connection_kwargs["credential_provider"] is provider
def test_async_url_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
client = get_redis_async_client(url="redis://redis-host:6379", credential_provider=provider)
assert client.connection_pool.connection_kwargs["credential_provider"] is provider
def test_sync_url_credentials_do_not_replace_explicit_provider(clean_redis_environment):
provider = _StubCredentialProvider()
client = get_redis_client(
url="redis://url-user:url-pass@redis-host:6379",
credential_provider=provider,
)
connection = client.connection_pool.make_connection()
assert connection.credential_provider is provider
assert connection.username is None
assert connection.password is None
def test_async_url_credentials_do_not_replace_explicit_provider(clean_redis_environment):
provider = _StubCredentialProvider()
client = get_redis_async_client(
url="redis://url-user:url-pass@redis-host:6379",
credential_provider=provider,
)
connection = client.connection_pool.make_connection()
assert connection.credential_provider is provider
assert connection.username is None
assert connection.password is None
def test_async_host_port_pool_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
pool = get_redis_connection_pool(host="redis-host", port=6379, credential_provider=provider)
assert pool is not None
assert pool.connection_kwargs["credential_provider"] is provider
def test_async_url_pool_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
pool = get_redis_connection_pool(url="redis://redis-host:6379", credential_provider=provider)
assert pool is not None
assert pool.connection_kwargs["credential_provider"] is provider
def test_async_url_pool_strips_userinfo_for_the_provider(clean_redis_environment):
provider = _StubCredentialProvider()
pool = get_redis_connection_pool(url="rediss://url-user:url-pass@redis-host:6379/3", credential_provider=provider)
connection = pool.make_connection()
assert connection.credential_provider is provider
assert connection.username is None
assert connection.password is None
assert connection.db == 3
def test_sync_cluster_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
startup_nodes = [{"host": "cluster-node", "port": 6379}]
with patch("redis.RedisCluster", autospec=True) as mock_cluster_cls:
get_redis_client(startup_nodes=startup_nodes, credential_provider=provider, password="redis-secret")
cluster_kwargs = mock_cluster_cls.call_args.kwargs
assert cluster_kwargs["credential_provider"] is provider
assert "password" not in cluster_kwargs
assert [(node.host, node.port) for node in cluster_kwargs["startup_nodes"]] == [("cluster-node", 6379)]
def test_async_cluster_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
startup_nodes = [{"host": "cluster-node", "port": 6379}]
client = get_redis_async_client(startup_nodes=startup_nodes, credential_provider=provider)
assert client.connection_kwargs["credential_provider"] is provider
assert client.connection_kwargs["socket_keepalive"] is True
assert client.connection_kwargs["health_check_interval"] == REDIS_CLUSTER_HEALTH_CHECK_INTERVAL
def test_explicit_provider_skips_automatic_auth_and_callback(clean_redis_environment, monkeypatch):
provider = _StubCredentialProvider()
monkeypatch.setenv("REDIS_GCP_SERVICE_ACCOUNT", "service-account@example.com")
monkeypatch.setenv("REDIS_AZURE_AD_TOKEN", "true")
with (
patch( # test-quality-ok: an auto-auth callback built here is popped again by the provider branch, so the builders are the only place the wasted work is visible
"litellm._redis.create_gcp_iam_redis_connect_func"
) as mock_gcp,
patch( # test-quality-ok: same as above, and reaching this one also builds an Azure credential the caller never asked for
"litellm._redis.create_azure_ad_redis_connect_func"
) as mock_azure,
):
redis_kwargs = _get_redis_client_logic(
host="redis-host",
port=6379,
credential_provider=provider,
redis_connect_func=_gcp_marker_callback(),
)
mock_gcp.assert_not_called()
mock_azure.assert_not_called()
assert redis_kwargs["credential_provider"] is provider
assert "redis_connect_func" not in redis_kwargs
@pytest.mark.parametrize(
"overrides",
[
{"gcp_ssl_ca_certs": "/tmp/ca.pem"},
{"gcp_service_account": "sa@example.com", "gcp_ssl_ca_certs": "/tmp/ca.pem"},
],
ids=["certs-without-service-account", "both-alongside-a-provider"],
)
def test_gcp_kwargs_never_survive_client_logic(clean_redis_environment, overrides):
redis_kwargs = _get_redis_client_logic(
host="redis-host",
port=6379,
credential_provider=_StubCredentialProvider() if "gcp_service_account" in overrides else None,
**overrides,
)
assert "gcp_service_account" not in redis_kwargs
assert "gcp_ssl_ca_certs" not in redis_kwargs
def test_provider_keeps_the_rest_of_the_url_intact(clean_redis_environment):
provider = _StubCredentialProvider()
redis_kwargs = _get_redis_client_logic(
url="rediss://url-user:url-pass@redis-host:6379/3?protocol=3",
credential_provider=provider,
)
assert redis_kwargs["url"] == "rediss://redis-host:6379/3?protocol=3"
def test_provider_free_url_is_left_untouched(clean_redis_environment):
url = "redis://url-user:url-pass@redis-host:6379/3"
redis_kwargs = _get_redis_client_logic(url=url)
assert redis_kwargs["url"] == url
def test_async_auth_kwargs_supersedes_credentials_an_explicit_provider_replaces():
provider = _StubCredentialProvider()
auth_kwargs = _async_auth_kwargs(
{
"host": "redis-host",
"port": 6379,
"credential_provider": provider,
"redis_connect_func": _gcp_marker_callback(),
"username": "url-user",
"password": "url-pass",
}
)
assert auth_kwargs["credential_provider"] is provider
assert auth_kwargs["host"] == "redis-host"
assert auth_kwargs["port"] == 6379
assert "redis_connect_func" not in auth_kwargs
assert "username" not in auth_kwargs
assert "password" not in auth_kwargs
def test_async_auth_kwargs_leaves_provider_free_kwargs_alone():
redis_kwargs = {"host": "redis-host", "port": 6379, "username": "url-user", "password": "url-pass"}
assert _async_auth_kwargs(redis_kwargs) == redis_kwargs
@pytest.mark.asyncio
async def test_redis_cache_test_connection_uses_shared_factory(clean_redis_environment):
provider = _StubCredentialProvider()
with (
patch("redis.Redis", autospec=True),
patch("redis.asyncio.BlockingConnectionPool", autospec=True),
patch("redis.asyncio.Redis", autospec=True) as mock_async_redis,
):
mock_async_redis.return_value.ping = AsyncMock(return_value=True)
mock_async_redis.return_value.aclose = AsyncMock()
cache = RedisCache(host="redis-host", port=6379, credential_provider=provider, password="redis-secret")
result = await cache.test_connection()
client_kwargs = mock_async_redis.call_args.kwargs
assert result["status"] == "success"
assert client_kwargs["credential_provider"] is provider
assert "password" not in client_kwargs
@pytest.mark.asyncio
async def test_redis_cluster_cache_test_connection_uses_shared_factory(clean_redis_environment):
provider = _StubCredentialProvider()
recorder = MagicMock()
class _StubAsyncCluster:
def __init__(self, **kwargs):
recorder(**kwargs)
async def ping(self):
return True
async def aclose(self):
return None
with (
patch("redis.RedisCluster", autospec=True),
patch("redis.asyncio.cluster.RedisCluster", _StubAsyncCluster),
):
cache = RedisClusterCache(startup_nodes=[{"host": "redis-host", "port": 6379}], credential_provider=provider)
result = await cache.test_connection()
cluster_kwargs = recorder.call_args.kwargs
assert result["status"] == "success"
assert cluster_kwargs["credential_provider"] is provider
def test_redis_cache_key_does_not_inspect_provider(clear_llm_client_cache):
provider = _HostileCredentialProvider("synthetic-secret")
second_provider = _StubCredentialProvider("another-token")
with (
patch("redis.Redis", autospec=True),
patch("redis.asyncio.BlockingConnectionPool", autospec=True),
):
cache = RedisCache(host="redis-host", port=6379, credential_provider=provider)
second_cache = RedisCache(host="redis-host", port=6379, credential_provider=second_provider)
first_key = cache._get_async_client_cache_key()
assert first_key == cache._get_async_client_cache_key()
assert first_key != second_cache._get_async_client_cache_key()
def test_pretty_print_never_expands_credential_provider(capsys):
secret = "aaaa-UNIQUE-SENTINEL-bbbb"
with patch( # test-quality-ok: enable the debug-only printer without changing process-wide logger state
"litellm._redis.verbose_logger.isEnabledFor", return_value=True
):
_pretty_print_redis_config(
redis_kwargs={
"host": "redis-host",
"port": 6379,
"credential_provider": _HostileCredentialProvider(secret),
}
)
output = capsys.readouterr().out
assert secret not in output
assert "UNIQUE" not in output
assert "_payload" not in output
assert "credential_provider" in output
def test_redis_cache_key_does_not_serialize_connect_func():
def connect(connection):
return None
cache = RedisCache.__new__(RedisCache)
cache.redis_kwargs = {"host": "redis-host", "port": 6379, "redis_connect_func": connect}
first_key = cache._get_async_client_cache_key()
assert first_key == cache._get_async_client_cache_key()
def test_redis_cache_key_keys_opaque_kwargs_by_identity():
class _Opaque:
pass
first = RedisCache.__new__(RedisCache)
first.redis_kwargs = {"host": "redis-host", "retry": _Opaque()}
second = RedisCache.__new__(RedisCache)
second.redis_kwargs = {"host": "redis-host", "retry": _Opaque()}
assert first._get_async_client_cache_key() == first._get_async_client_cache_key()
assert first._get_async_client_cache_key() != second._get_async_client_cache_key()
def test_get_redis_url_from_environment_single_url(monkeypatch):
"""Test when REDIS_URL is directly provided"""
# Set the environment variable
@ -500,6 +926,27 @@ def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_
)
@patch("redis.Sentinel")
def test_sync_sentinel_keeps_provider_off_monitors_and_on_master(mock_sentinel_cls):
provider = _StubCredentialProvider()
mock_sentinel = MagicMock()
mock_sentinel_cls.return_value = mock_sentinel
get_redis_client(
sentinel_nodes=[("sentinel-1", 26379)],
sentinel_password="sentinel-secret",
service_name="mymaster",
password="redis-secret",
credential_provider=provider,
)
sentinel_kwargs = mock_sentinel_cls.call_args.kwargs["sentinel_kwargs"]
assert sentinel_kwargs["password"] == "sentinel-secret"
assert "credential_provider" not in sentinel_kwargs
assert mock_sentinel.master_for.call_args.kwargs["credential_provider"] is provider
assert "password" not in mock_sentinel.master_for.call_args.kwargs
@patch("litellm._redis.async_redis.Sentinel")
def test_async_sentinel_uses_sentinel_password_and_master_password(
mock_sentinel_cls,

View file

@ -367,6 +367,45 @@ async def test_router_order_fallback_with_wildcard_model_group():
assert response._hidden_params["model_id"] == "2"
@pytest.mark.asyncio
async def test_router_order_fallback_with_hidden_model_group_alias():
router = Router(
model_list=[
{
"model_name": "canonical-model",
"litellm_params": {
"model": "gpt-4o",
"api_key": "bad",
"mock_response": Exception("fail order 1"),
"order": 1,
},
"model_info": {"id": "1"},
},
{
"model_name": "canonical-model",
"litellm_params": {
"model": "gpt-4o",
"api_key": "good",
"mock_response": "success from order 2",
"order": 2,
},
"model_info": {"id": "2"},
},
],
model_group_alias={"hidden-alias": {"model": "canonical-model", "hidden": True}},
num_retries=0,
)
assert "hidden-alias" not in {deployment["model_name"] for deployment in router.get_model_list() or []}
response = await router.acompletion(
model="hidden-alias",
messages=[{"role": "user", "content": "hi"}],
)
assert response._hidden_params["model_id"] == "2"
def test_check_non_standard_fallback_format():
from litellm.router_utils.fallback_event_handlers import (
_check_non_standard_fallback_format,

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22749
"limit": 22733
},
"LIT002": {
"limit": 26866
"limit": 26864
},
"LIT003": {
"limit": 269
@ -27,7 +27,7 @@
"limit": 0
},
"LIT010": {
"limit": 16655
"limit": 16621
},
"LIT011": {
"limit": 5585

View file

@ -1465,9 +1465,6 @@
"src/components/add_model/conditional_public_model_name.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"local/no-complex-jsx-arrow": {
"count": 1
}
},
"src/components/add_model/handle_add_auto_router_submit.tsx": {

View file

@ -10,7 +10,6 @@ import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { Info, TriangleAlert } from "lucide-react";
import React, { useEffect, useState } from "react";
import NewBadge from "@/components/common_components/NewBadge";
import { useBaseUrl } from "@/components/constants";
import { toast } from "@/lib/toast";
import { addAllowedIP, deleteAllowedIP, getAllowedIPs, getSSOSettings } from "@/components/networking";
@ -378,12 +377,7 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
},
{
key: "ui-settings",
label: (
<span className="flex items-center gap-1.5">
UI Settings
<NewBadge />
</span>
),
label: "UI Settings",
children: (
<div className="flex flex-col gap-4">
<UISettings />

View file

@ -15,7 +15,6 @@ import {
AlertDialogHeader,
AlertDialogTitle,
} from "@/components/ui/alert-dialog";
import NewBadge from "@/components/common_components/NewBadge";
import React, { useEffect, useState, useMemo, useCallback } from "react";
import { useQuery } from "@tanstack/react-query";
import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers";
@ -538,8 +537,8 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
</TabsTrigger>
)}
{isAdminRole(userRole) && (
<TabsTrigger value="submitted" className="flex-none gap-2 rounded-none px-4 py-2">
Submitted MCPs <NewBadge />
<TabsTrigger value="submitted" className="flex-none rounded-none px-4 py-2">
Submitted MCPs
</TabsTrigger>
)}
</TabsList>

View file

@ -6,6 +6,7 @@ export const TEST_MODES = [
{ value: "audio_speech", label: "Audio Speech - /audio/speech" },
{ value: "audio_transcription", label: "Audio Transcription - /audio/transcriptions" },
{ value: "image_generation", label: "Image Generation - /images/generations" },
{ value: "image_edit", label: "Image Edit - /images/edits" },
{ value: "video_generation", label: "Video Generation - /videos" },
{ value: "rerank", label: "Rerank - /rerank" },
{ value: "realtime", label: "Realtime - /realtime" },

View file

@ -1,4 +1,5 @@
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import React, { useEffect, useRef } from "react";
import { useFormContext, useWatch } from "react-hook-form";
import { describe, expect, it } from "vitest";
@ -69,4 +70,24 @@ describe("ConditionalPublicModelName", () => {
expect(screen.getByText("my-custom-model")).toBeInTheDocument();
expect(screen.queryByDisplayValue("custom")).not.toBeInTheDocument();
});
it("keeps the public name input focused across keystrokes", async () => {
const user = userEvent.setup();
render(
<MountedFormHost
defaultValues={{
model: ["gpt-4"],
model_mappings: [{ public_name: "gpt-4", litellm_model: "gpt-4" }],
}}
>
<ConditionalPublicModelName />
</MountedFormHost>,
);
const input = screen.getByDisplayValue("gpt-4");
await user.type(input, "-prod");
expect(input).toHaveValue("gpt-4-prod");
expect(input).toHaveFocus();
});
});

View file

@ -36,6 +36,82 @@ const modelMappingsRule = {
const tooltipCodeClassName = "rounded-sm bg-background/20 px-1 py-0.5 font-mono text-xs";
const ANTHROPIC_1M_HEADERS = JSON.stringify({ extra_headers: { "anthropic-beta": "context-1m-2025-08-07" } }, null, 2);
const publicNameTooltipContent = (
<div className="flex flex-col gap-2 text-left font-normal">
<div>The name you specify in your API calls to LiteLLM Proxy</div>
<div>
<strong>Example:</strong> If you name your public model <code className={tooltipCodeClassName}>example-name</code>
, and choose <code className={tooltipCodeClassName}>openai/qwen-plus-latest</code> as the LiteLLM model
</div>
<div>
<strong>Usage:</strong> You make an API call to the LiteLLM proxy with{" "}
<code className={tooltipCodeClassName}>model = &quot;example-name&quot;</code>
</div>
<div>
<strong>Result:</strong> LiteLLM sends <code className={tooltipCodeClassName}>qwen-plus-latest</code> to the
provider
</div>
</div>
);
const PublicNameInput: React.FC<{ readonly index: number; readonly value: string }> = ({ index, value }) => {
const form = useFormContext<MountedFormValues>();
const selectedProvider = useWatch({ control: form.control, name: "custom_llm_provider" });
const handleChange = (event: React.ChangeEvent<HTMLInputElement>) => {
const typed = event.target.value;
const litellmParams = form.getValues("litellm_extra_params") as string | undefined;
const wantsAnthropic1m =
selectedProvider === Providers.Anthropic && typed.endsWith("-1m") && (litellmParams ?? "").trim() === "";
if (wantsAnthropic1m) {
form.setValue("litellm_extra_params", ANTHROPIC_1M_HEADERS);
}
const publicName = wantsAnthropic1m ? typed.slice(0, -"-1m".length) : typed;
const current = (form.getValues("model_mappings") as ModelMapping[]) ?? [];
form.setValue(
"model_mappings",
current.map((mapping, mappingIndex) =>
mappingIndex === index ? { ...mapping, public_name: publicName } : mapping,
),
);
};
return <Input value={value} onChange={handleChange} />;
};
/**
* Module-level so the header and cell renderers keep a stable identity: React treats a renderer
* declared inside the component as a new element type on every render and remounts the input,
* which drops focus after each keystroke.
*/
const columns: ColumnDef<ModelMapping>[] = [
{
id: "public_name",
accessorKey: "public_name",
header: () => (
<span className="flex items-center">
Public Model Name
<SimpleTooltip content={publicNameTooltipContent} width="500px" />
</span>
),
cell: ({ row }) => <PublicNameInput index={row.index} value={row.original.public_name} />,
},
{
id: "litellm_model",
accessorKey: "litellm_model",
header: () => (
<span className="flex items-center">
LiteLLM Model Name
<SimpleTooltip content={<div>The model name LiteLLM will send to the LLM API</div>} width="360px" />
</span>
),
},
];
const ConditionalPublicModelName: React.FC = () => {
const form = useFormContext<MountedFormValues>();
@ -124,85 +200,6 @@ const ConditionalPublicModelName: React.FC = () => {
if (!showPublicModelName) return null;
const publicNameTooltipContent = (
<div className="flex flex-col gap-2 text-left font-normal">
<div>The name you specify in your API calls to LiteLLM Proxy</div>
<div>
<strong>Example:</strong> If you name your public model{" "}
<code className={tooltipCodeClassName}>example-name</code>, and choose{" "}
<code className={tooltipCodeClassName}>openai/qwen-plus-latest</code> as the LiteLLM model
</div>
<div>
<strong>Usage:</strong> You make an API call to the LiteLLM proxy with{" "}
<code className={tooltipCodeClassName}>model = &quot;example-name&quot;</code>
</div>
<div>
<strong>Result:</strong> LiteLLM sends <code className={tooltipCodeClassName}>qwen-plus-latest</code> to the
provider
</div>
</div>
);
const liteLLMModelTooltipContent = <div>The model name LiteLLM will send to the LLM API</div>;
const columns: ColumnDef<ModelMapping>[] = [
{
id: "public_name",
accessorKey: "public_name",
header: () => (
<span className="flex items-center">
Public Model Name
<SimpleTooltip content={publicNameTooltipContent} width="500px" />
</span>
),
cell: ({ row }) => {
return (
<Input
value={row.original.public_name}
onChange={(e) => {
const newValue = e.target.value;
const newMappings = [...((form.getValues("model_mappings") as ModelMapping[]) ?? [])];
// Check conditions for Anthropic -1m suffix handling
const isAnthropic = selectedProvider === Providers.Anthropic;
const endsWith1m = newValue.endsWith("-1m");
const litellmParams = form.getValues("litellm_extra_params") as string | undefined;
const isLitellmParamsEmpty = !litellmParams || litellmParams.trim() === "";
let finalPublicName = newValue;
if (isAnthropic && endsWith1m && isLitellmParamsEmpty) {
// Set litellm params with extra_headers
const litellmParamsValue = JSON.stringify(
{ extra_headers: { "anthropic-beta": "context-1m-2025-08-07" } },
null,
2,
);
form.setValue("litellm_extra_params", litellmParamsValue);
// Remove -1m suffix from public_name
finalPublicName = newValue.slice(0, -3); // Remove "-1m" (3 characters)
}
newMappings[row.index].public_name = finalPublicName;
form.setValue("model_mappings", newMappings);
}}
/>
);
},
},
{
id: "litellm_model",
accessorKey: "litellm_model",
header: () => (
<span className="flex items-center">
LiteLLM Model Name
<SimpleTooltip content={liteLLMModelTooltipContent} width="360px" />
</span>
),
},
];
return (
<MountedFormField
name="model_mappings"

View file

@ -74,7 +74,6 @@ import {
rolesWithWriteAccess,
} from "../utils/roles";
import BetaBadge from "./BetaBadge";
import NewBadge from "./common_components/NewBadge";
import SidebarAccountMenu from "./SidebarAccountMenu/SidebarAccountMenu";
import SidebarUsageCard from "./SidebarUsageCard";
import { MIGRATED_PAGES, migratedHref, legacyPageHref } from "@/utils/migratedPages";
@ -320,11 +319,7 @@ const menuGroups: MenuGroup[] = [
{
key: "settings",
page: "settings",
label: (
<span className="flex items-center gap-2">
Settings <NewBadge />
</span>
),
label: "Settings",
icon: <SettingsIcon {...ICON} />,
roles: all_admin_roles,
children: [
@ -345,14 +340,7 @@ const menuGroups: MenuGroup[] = [
{
key: "admin-panel",
page: "admin-panel",
label: (
<span className="flex items-center gap-2">
Admin Settings{" "}
<NewBadge dot>
<span />
</NewBadge>
</span>
),
label: "Admin Settings",
icon: <SettingsIcon {...ICON} />,
roles: all_admin_roles,
},

View file

@ -22703,7 +22703,7 @@ export interface components {
* Mode
* @description The mode to test the model with. If not provided, auto-detected from model capabilities.
*/
mode?: ("chat" | "completion" | "embedding" | "audio_speech" | "audio_transcription" | "image_generation" | "video_generation" | "batch" | "rerank" | "realtime" | "responses" | "ocr") | null;
mode?: ("chat" | "completion" | "embedding" | "audio_speech" | "audio_transcription" | "image_generation" | "image_edit" | "video_generation" | "batch" | "rerank" | "realtime" | "responses" | "ocr") | null;
/**
* Model Info
* @description Model info for the health check
@ -36262,7 +36262,9 @@ export interface components {
} | null;
/** Model Max Budget Usage */
model_max_budget_usage?: {
[key: string]: unknown;
[key: string]: {
[key: string]: unknown;
};
} | null;
/**
* Models