mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
ece187ea24
80 changed files with 2348 additions and 621 deletions
5
.github/mutmut-coverage.rc
vendored
Normal file
5
.github/mutmut-coverage.rc
vendored
Normal 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
|
||||
10
.github/workflows/mutation-test.yml
vendored
10
.github/workflows/mutation-test.yml
vendored
|
|
@ -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
2
.gitignore
vendored
|
|
@ -3,6 +3,8 @@
|
|||
tests/e2e/.fixtures/
|
||||
.venv-typecheck
|
||||
.venv_policy_test
|
||||
.venv-mutmut
|
||||
mutants/
|
||||
.env
|
||||
.claude
|
||||
CLAUDE.local.md
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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']}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1796,6 +1796,7 @@ async def test_model_connection(
|
|||
"audio_speech",
|
||||
"audio_transcription",
|
||||
"image_generation",
|
||||
"image_edit",
|
||||
"video_generation",
|
||||
"batch",
|
||||
"rerank",
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == {}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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 />
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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" },
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 = "example-name"</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 = "example-name"</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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
|
|
|
|||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue