[Fix] /v1/messages — plumb client-supplied timeout to httpx call (#26754)

Squash-merged by litellm-agent from kimsehwan96's PR.
This commit is contained in:
hayden 2026-05-12 09:57:17 +09:00 committed by GitHub
parent e182a5e0ba
commit 0e674042cd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
102 changed files with 8970 additions and 623 deletions

131
.github/workflows/mutation-test.yml vendored Normal file
View file

@ -0,0 +1,131 @@
name: "Mutation Test (manual)"
# Manually-triggered mutation testing. Runs mutmut against the scope
# configured in [tool.mutmut] in pyproject.toml (currently the
# litellm/proxy/management_endpoints/ folder). Intended cadence is roughly
# weekly — clicked from the Actions tab when someone wants a fresh report.
#
# Uploads a structured `mutation-report.md` (Meta ACH-style: original +
# mutated function with `# MUTANT START`/`# MUTANT END` delimiters + the
# existing tests + a task instruction) as a workflow artifact. Failures
# do not block anything because nothing depends on this workflow.
on:
workflow_dispatch:
permissions:
contents: read
concurrency:
group: mutation-test-${{ github.ref }}
cancel-in-progress: true
jobs:
mutation:
name: Run mutmut
runs-on: ubuntu-latest
# Whole-folder mutation against ~15 files / ~7.5k LOC can take hours.
# 350 minutes is just under the GitHub-hosted job cap of 360 minutes.
timeout-minutes: 350
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
with:
version: "0.10.9"
- name: Cache uv dependencies
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Install dependencies
run: |
uv sync --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
# mutmut 3.x runs tests inside a `mutants/` sandbox where it injects
# mutation trampolines. uv installs the project as editable by default,
# which puts the original source dir on sys.path via a .pth file and
# shadows the sandbox copy — so tests would never exercise the mutated
# code. Reinstalling non-editable removes the .pth shadow.
- name: Reinstall litellm non-editable (so mutants/ is not shadowed)
run: |
uv pip uninstall litellm
uv pip install . --no-deps
# pytest-retry's pytest_configure hook crashes with
# `INTERNALERROR: no option named 'filtered_exceptions'` when invoked
# via mutmut's in-process pytest.main() call. The entry-point name
# doesn't normalize cleanly with `-p no:<name>`, so just remove the
# package outright. Reruns are wrong for mutation testing anyway —
# rerunning a "failed" mutant test would mask which mutants are killed.
- name: Remove pytest plugins that conflict with mutmut
run: |
uv pip uninstall pytest-retry || true
- name: Run mutmut
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
run: |
set -o pipefail
mkdir -p mutants
uv run --no-sync --with mutmut==3.5.0 mutmut run 2>&1 | tee mutmut-run.log
# Generate the structured report. The script embeds the enclosing
# function source for each survivor (via Python AST) and includes the
# existing test files, so an LLM agent has enough context to write
# killing tests without further file lookups. Modeled on Meta's ACH
# prompt template (arXiv 2501.12862).
- name: Generate detailed mutation report
if: always()
run: |
set +e
uv run --no-sync --with mutmut==3.5.0 mutmut export-cicd-stats > /dev/null 2>&1
uv run --no-sync --with mutmut==3.5.0 mutmut results > mutmut-results.txt 2>&1
uv run --no-sync python scripts/mutation_report.py
# The full report can be very long for big test files; the run-page
# summary cuts off at 1 MB. Append the head of the report (summary
# + survivor list) and link out to the artifact for the full body.
{
head -c 900000 mutation-report.md
echo ""
echo ""
echo "_Full report (with embedded function bodies and test files) is in the workflow artifact._"
} >> "$GITHUB_STEP_SUMMARY"
- name: Upload mutmut artifacts
if: always()
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: mutmut-${{ github.run_id }}-${{ github.run_attempt }}
path: |
mutation-report.md
mutmut-results.txt
mutmut-run.log
mutants/mutmut-stats.json
mutants/mutmut-cicd-stats.json
mutants/litellm/proxy/management_endpoints/**/*.py
if-no-files-found: warn
retention-days: 14

View file

@ -1426,6 +1426,12 @@ if TYPE_CHECKING:
)
from .llms.datarobot.chat.transformation import DataRobotConfig as DataRobotConfig
from .llms.anthropic.chat.transformation import AnthropicConfig as AnthropicConfig
from .llms.bedrock.claude_platform.transformation import (
BedrockClaudePlatformConfig as BedrockClaudePlatformConfig,
)
from .llms.bedrock.claude_platform.messages_transformation import (
BedrockClaudePlatformMessagesConfig as BedrockClaudePlatformMessagesConfig,
)
from .llms.anthropic.completion.transformation import (
AnthropicTextConfig as AnthropicTextConfig,
)

View file

@ -131,6 +131,7 @@ LLM_CONFIG_NAMES = (
"OpenrouterConfig",
"DataRobotConfig",
"AnthropicConfig",
"BedrockClaudePlatformConfig",
"AnthropicTextConfig",
"GroqSTTConfig",
"TritonConfig",
@ -170,6 +171,7 @@ LLM_CONFIG_NAMES = (
"SagemakerNovaConfig",
"CohereChatConfig",
"AnthropicMessagesConfig",
"BedrockClaudePlatformMessagesConfig",
"AmazonAnthropicClaudeMessagesConfig",
"AmazonMantleMessagesConfig",
"TogetherAIConfig",
@ -610,6 +612,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
"OpenrouterConfig": (".llms.openrouter.chat.transformation", "OpenrouterConfig"),
"DataRobotConfig": (".llms.datarobot.chat.transformation", "DataRobotConfig"),
"AnthropicConfig": (".llms.anthropic.chat.transformation", "AnthropicConfig"),
"BedrockClaudePlatformConfig": (
".llms.bedrock.claude_platform.transformation",
"BedrockClaudePlatformConfig",
),
"AnthropicTextConfig": (
".llms.anthropic.completion.transformation",
"AnthropicTextConfig",
@ -712,6 +718,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
".llms.anthropic.experimental_pass_through.messages.transformation",
"AnthropicMessagesConfig",
),
"BedrockClaudePlatformMessagesConfig": (
".llms.bedrock.claude_platform.messages_transformation",
"BedrockClaudePlatformMessagesConfig",
),
"AmazonAnthropicClaudeMessagesConfig": (
".llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation",
"AmazonAnthropicClaudeMessagesConfig",

View file

@ -19,6 +19,7 @@ import redis.asyncio as async_redis # type: ignore
from litellm import get_secret, get_secret_str
from litellm._redis_credential_provider import (
AzureADCredentialProvider,
GCPIAMCredentialProvider,
_generate_gcp_iam_access_token,
)
@ -27,6 +28,8 @@ from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from ._logging import verbose_logger
AZURE_REDIS_SCOPE = "https://redis.azure.com/.default"
def _get_redis_kwargs():
arg_spec = inspect.getfullargspec(redis.Redis)
@ -43,6 +46,10 @@ def _get_redis_kwargs():
"redis_connect_func",
"gcp_service_account",
"gcp_ssl_ca_certs",
"azure_redis_ad_token",
"azure_client_id",
"azure_tenant_id",
"azure_client_secret",
]
available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args
@ -89,6 +96,10 @@ def _get_redis_cluster_kwargs(client=None):
) # Needed for sync clusters and IAM detection
available_args.append("gcp_service_account")
available_args.append("gcp_ssl_ca_certs")
available_args.append("azure_redis_ad_token")
available_args.append("azure_client_id")
available_args.append("azure_tenant_id")
available_args.append("azure_client_secret")
available_args.append("max_connections")
return available_args
@ -155,6 +166,125 @@ def create_gcp_iam_redis_connect_func(
return iam_connect
def _build_azure_credential(
azure_client_id: Optional[str] = None,
azure_tenant_id: Optional[str] = None,
azure_client_secret: Optional[str] = None,
):
"""
Build a long-lived Azure credential object.
Azure SDK credentials cache tokens internally and handle expiry/refresh
transparently, so this should be called once and the result reused.
"""
try:
from azure.identity import (
ClientSecretCredential,
DefaultAzureCredential,
ManagedIdentityCredential,
)
except ImportError:
raise ImportError(
"azure-identity is required for Azure AD Redis authentication. "
"Install it with: pip install azure-identity"
)
_client_id = azure_client_id or os.environ.get("AZURE_CLIENT_ID")
_tenant_id = azure_tenant_id or os.environ.get("AZURE_TENANT_ID")
_client_secret = azure_client_secret or os.environ.get("AZURE_CLIENT_SECRET")
if _client_id and _tenant_id and _client_secret:
return ClientSecretCredential(
client_id=_client_id,
tenant_id=_tenant_id,
client_secret=_client_secret,
)
elif _client_id:
return ManagedIdentityCredential(client_id=_client_id)
else:
return DefaultAzureCredential()
def _generate_azure_ad_redis_token(
azure_client_id: Optional[str] = None,
azure_tenant_id: Optional[str] = None,
azure_client_secret: Optional[str] = None,
) -> str:
"""
One-shot helper that builds a credential and fetches a single Azure AD
access token for Redis. Each call rebuilds the credential and performs a
network round-trip, so it should not be used in steady-state Redis flows
the sync (``create_azure_ad_redis_connect_func``) and async paths
(``AzureADCredentialProvider``) keep the credential alive across
connections so the Azure SDK's internal cache + silent refresh apply.
"""
credential = _build_azure_credential(
azure_client_id=azure_client_id,
azure_tenant_id=azure_tenant_id,
azure_client_secret=azure_client_secret,
)
token = credential.get_token(AZURE_REDIS_SCOPE)
return token.token
def create_azure_ad_redis_connect_func(
azure_client_id: Optional[str] = None,
azure_tenant_id: Optional[str] = None,
azure_client_secret: Optional[str] = None,
) -> Callable:
"""
Creates a custom Redis connection function for Azure AD authentication.
Used for sync Redis clients. The credential is created once (captured by the
closure) and reused across connections the Azure SDK handles token caching
and silent renewal internally. Only ``get_token`` is called per connection.
"""
credential = _build_azure_credential(
azure_client_id=azure_client_id,
azure_tenant_id=azure_tenant_id,
azure_client_secret=azure_client_secret,
)
def ad_connect(self):
"""Initialize the connection and authenticate using Azure AD"""
from redis.exceptions import (
AuthenticationError,
AuthenticationWrongNumberOfArgsError,
)
from redis.utils import str_if_bytes
self._parser.on_connect(self)
access_token = credential.get_token(AZURE_REDIS_SCOPE).token
# Only include username when explicitly set — sending AUTH "" <token>
# is invalid for most ACL-configured Azure Redis instances.
username = os.environ.get("REDIS_USERNAME", "")
if username:
auth_args = (username, access_token)
else:
auth_args = (access_token,)
self.send_command("AUTH", *auth_args, check_health=False)
try:
auth_response = self.read_response()
except AuthenticationWrongNumberOfArgsError:
# Fallback: try with just the token (Redis < 6 / no ACL)
self.send_command("AUTH", access_token, check_health=False)
auth_response = self.read_response()
if str_if_bytes(auth_response) != "OK":
raise AuthenticationError("Azure AD authentication failed for Redis")
# Attach the live credential object so async paths can wrap it in
# AzureADCredentialProvider for refresh-aware token retrieval. The raw
# client_id/tenant_id/secret are intentionally NOT exposed here — the
# credential closure already holds them.
ad_connect._azure_credential = credential # type: ignore[attr-defined]
return ad_connect
def get_redis_url_from_environment():
if "REDIS_URL" in os.environ:
return os.environ["REDIS_URL"]
@ -179,7 +309,7 @@ def get_redis_url_from_environment():
return f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}"
def _get_redis_client_logic(**env_overrides):
def _get_redis_client_logic(**env_overrides): # noqa: PLR0915
"""
Common functionality across sync + async redis client implementations
"""
@ -253,6 +383,52 @@ def _get_redis_client_logic(**env_overrides):
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 = redis_kwargs.get("azure_redis_ad_token") or get_secret(
"REDIS_AZURE_AD_TOKEN"
)
_azure_ad_enabled = (
_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 None:
_azure_client_id = redis_kwargs.get("azure_client_id") or get_secret_str(
"AZURE_CLIENT_ID"
)
_azure_tenant_id = redis_kwargs.get("azure_tenant_id") or get_secret_str(
"AZURE_TENANT_ID"
)
_azure_client_secret = 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 # type: ignore[attr-defined]
# Always remove Azure-specific kwargs that shouldn't be passed to Redis client
redis_kwargs.pop("azure_redis_ad_token", None)
redis_kwargs.pop("azure_client_id", None)
redis_kwargs.pop("azure_tenant_id", None)
redis_kwargs.pop("azure_client_secret", None)
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
@ -373,7 +549,7 @@ def get_redis_client(**env_overrides):
return redis.Redis(**redis_kwargs)
def get_redis_async_client(
def get_redis_async_client( # noqa: PLR0915
connection_pool: Optional[async_redis.BlockingConnectionPool] = None,
**env_overrides,
) -> Union[async_redis.Redis, async_redis.RedisCluster]:
@ -398,6 +574,14 @@ def get_redis_async_client(
cluster_kwargs["credential_provider"] = GCPIAMCredentialProvider(
redis_connect_func._gcp_service_account
)
# Handle Azure AD authentication for async clusters via CredentialProvider
# so the credential's internal cache + silent refresh runs per connection
# (mirrors GCP IAM above; avoids static-token-baked-in-pool expiry).
elif redis_connect_func and hasattr(redis_connect_func, "_azure_credential"):
cluster_kwargs["credential_provider"] = AzureADCredentialProvider(
redis_connect_func._azure_credential,
username=os.environ.get("REDIS_USERNAME") or None,
)
new_startup_nodes: List[ClusterNode] = []
@ -431,6 +615,22 @@ def get_redis_async_client(
# Check for Redis Sentinel
if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs:
return _init_async_redis_sentinel(redis_kwargs)
# Wrap GCP / Azure AD auth in a CredentialProvider for the standard async
# Redis client. The async client doesn't support redis_connect_func, but it
# does honour credential_provider — which is called per connection, so the
# underlying SDK can refresh tokens silently before they expire.
redis_connect_func = redis_kwargs.pop("redis_connect_func", None)
if redis_connect_func and hasattr(redis_connect_func, "_azure_credential"):
redis_kwargs["credential_provider"] = AzureADCredentialProvider(
redis_connect_func._azure_credential,
username=os.environ.get("REDIS_USERNAME") or None,
)
elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(
redis_connect_func._gcp_service_account
)
_pretty_print_redis_config(redis_kwargs=redis_kwargs)
if connection_pool is not None:
@ -464,6 +664,21 @@ def get_redis_connection_pool(
redis_kwargs["max_connections"],
)
return async_redis.BlockingConnectionPool.from_url(**pool_kwargs)
# Wrap GCP / Azure AD auth in a CredentialProvider so pool-managed
# connections re-fetch tokens via the SDK's internal cache + silent refresh
# rather than reusing a single token captured at pool creation.
redis_connect_func = redis_kwargs.pop("redis_connect_func", None)
if redis_connect_func and hasattr(redis_connect_func, "_azure_credential"):
redis_kwargs["credential_provider"] = AzureADCredentialProvider(
redis_connect_func._azure_credential,
username=os.environ.get("REDIS_USERNAME") or None,
)
elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(
redis_connect_func._gcp_service_account
)
connection_class = async_redis.Connection
if "ssl" in redis_kwargs:
connection_class = async_redis.SSLConnection

View file

@ -1,10 +1,13 @@
import asyncio
import threading
import time
from typing import Dict, Tuple
from typing import Any, Dict, Optional, Tuple, Union
from redis.credentials import CredentialProvider # type: ignore[attr-defined]
# Azure AD scope for Redis Cache for Azure.
AZURE_REDIS_SCOPE = "https://redis.azure.com/.default"
# GCP IAM tokens are valid for 1 hour. Cache for 55 minutes to refresh before expiry.
_GCP_IAM_TOKEN_TTL_SECONDS = 3300
@ -101,3 +104,33 @@ class GCPIAMCredentialProvider(CredentialProvider):
_get_cached_gcp_iam_token, self._gcp_service_account
)
return (token,)
class AzureADCredentialProvider(CredentialProvider):
"""
redis.credentials.CredentialProvider implementation that supplies Azure AD
tokens for Redis authentication.
Wraps an azure-identity credential object so the Azure SDK's internal token
cache and silent refresh are honoured on every Redis connection. This avoids
the static-token-baked-in-pool issue where pool-managed connections would
fail authentication after the initial token expired (~1 hour TTL).
"""
def __init__(self, credential: Any, username: Optional[str] = None) -> None:
self._credential = credential
self._username = username
def get_credentials(self) -> Union[Tuple[str], Tuple[str, str]]:
token = self._credential.get_token(AZURE_REDIS_SCOPE).token
if self._username:
return (self._username, token)
return (token,)
async def get_credentials_async(self) -> Union[Tuple[str], Tuple[str, str]]:
token_obj = await asyncio.to_thread(
self._credential.get_token, AZURE_REDIS_SCOPE
)
if self._username:
return (self._username, token_obj.token)
return (token_obj.token,)

View file

@ -617,24 +617,35 @@ def retrieve_batch(
_is_async = kwargs.pop("aretrieve_batch", False) is True
client = kwargs.get("client", None)
# Check if this is an async invoke ARN (different from regular batch ARN)
# Async invoke ARNs have format: arn:aws(-[^:]+)?:bedrock:[a-z0-9-]{1,20}:[0-9]{12}:async-invoke/[a-z0-9]{12}
if (
batch_id.startswith("arn:aws")
and ":bedrock:" in batch_id
and ":async-invoke/" in batch_id
):
# Handle async invoke status check
# Remove aws_region_name from kwargs to avoid duplicate parameter
async_kwargs = kwargs.copy()
async_kwargs.pop("aws_region_name", None)
# Bedrock has two distinct ARN families that need different APIs:
# * async-invoke ARNs (Twelve Labs Marengo embeddings) -> bedrock-runtime data plane
# * model-invocation-job ARNs (CreateModelInvocationJob batch) -> bedrock control plane
# They live on different AWS service endpoints and can't share a handler.
# ARN shapes:
# arn:aws(-[^:]+)?:bedrock:<region>:<account>:async-invoke/<id>
# arn:aws(-[^:]+)?:bedrock:<region>:<account>:model-invocation-job/<id>
if batch_id.startswith("arn:aws") and ":bedrock:" in batch_id:
if ":async-invoke/" in batch_id:
# Remove aws_region_name from kwargs to avoid duplicate parameter
async_kwargs = kwargs.copy()
async_kwargs.pop("aws_region_name", None)
return BedrockBatchesHandler._handle_async_invoke_status(
batch_id=batch_id,
aws_region_name=kwargs.get("aws_region_name", "us-east-1"),
logging_obj=litellm_logging_obj,
**async_kwargs,
)
return BedrockBatchesHandler._handle_async_invoke_status(
batch_id=batch_id,
aws_region_name=kwargs.get("aws_region_name", "us-east-1"),
logging_obj=litellm_logging_obj,
**async_kwargs,
)
if ":model-invocation-job/" in batch_id:
mij_kwargs = kwargs.copy()
mij_kwargs.pop("aws_region_name", None)
return BedrockBatchesHandler._handle_model_invocation_job_status(
batch_id=batch_id,
aws_region_name=kwargs.get("aws_region_name"),
logging_obj=litellm_logging_obj,
**mij_kwargs,
)
# Try to use provider config first (for providers like bedrock)
model: Optional[str] = kwargs.get("model", None)

View file

@ -119,6 +119,20 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
def __init__(self):
pass
def _normalize_tool_choice_for_responses_api(self, tool_choice: Any) -> Any:
"""Chat tool_choice uses function.name; Responses API expects top-level name."""
if not isinstance(tool_choice, dict) or tool_choice.get("type") != "function":
return tool_choice
if isinstance(tool_choice.get("name"), str) and tool_choice.get("name"):
# Return only Responses shape so stray chat ``function`` key is not sent upstream.
return {"type": "function", "name": tool_choice["name"]}
fn = tool_choice.get("function")
if isinstance(fn, dict):
fn_name = fn.get("name")
if isinstance(fn_name, str) and fn_name:
return {"type": "function", "name": fn_name}
return tool_choice
def _handle_raw_dict_response_item(
self, item: Dict[str, Any], index: int
) -> Tuple[Optional[Any], int]:
@ -309,6 +323,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
text_format = self._transform_response_format_to_text_format(value)
if text_format:
responses_api_request["text"] = text_format # type: ignore
elif key == "tool_choice":
responses_api_request["tool_choice"] = ( # type: ignore[assignment]
self._normalize_tool_choice_for_responses_api(value)
)
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
responses_api_request[key] = value # type: ignore
elif key == "previous_response_id":

View file

@ -1771,17 +1771,41 @@ class OpenTelemetry(CustomLogger):
value=safe_dumps(transformed_messages),
)
if kwargs.get("system_instructions"):
transformed_system_instructions = (
self._transform_messages_to_otel_semantic_conventions(
kwargs.get("system_instructions")
# Coalesce the different kwarg names that carry the system
# prompt depending on the call path:
# - "system_instructions" — Vertex AI Gemini chat-completion
# - "instructions" — OpenAI Responses API
# - "system" — Anthropic Messages API
# Use `is not None` rather than truthiness to avoid falsy
# values (e.g. []) falling through to the wrong kwarg.
system_instructions = (
kwargs.get("system_instructions")
if kwargs.get("system_instructions") is not None
else (
kwargs.get("instructions")
if kwargs.get("instructions") is not None
else kwargs.get("system")
)
)
if system_instructions:
if isinstance(system_instructions, str):
# Plain text system prompt — no transformation needed
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_SYSTEM_INSTRUCTIONS.value,
value=system_instructions,
)
else:
transformed_system_instructions = (
self._transform_messages_to_otel_semantic_conventions(
system_instructions
)
)
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_SYSTEM_INSTRUCTIONS.value,
value=safe_dumps(transformed_system_instructions),
)
)
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_SYSTEM_INSTRUCTIONS.value,
value=safe_dumps(transformed_system_instructions),
)
self.safe_set_attribute(
span=span,
@ -1840,6 +1864,57 @@ class OpenTelemetry(CustomLogger):
value=value,
)
elif response_obj.get("output"):
# Responses API: ResponsesAPIResponse has an "output"
# list instead of "choices". Each item with
# type="message" contains a "content" list of
# OutputText objects (type="output_text").
output_items = response_obj.get("output")
output_messages = self._transform_responses_api_output_to_otel(
output_items
)
if output_messages:
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_OUTPUT_MESSAGES.value,
value=safe_dumps(output_messages),
)
# Emit per-tool-call span attributes (parity with
# the choices branch that calls _tool_calls_kv_pair).
# Convert Responses API function_call items to the
# ChatCompletionMessageToolCall format expected by
# _tool_calls_kv_pair.
tool_calls = []
for out_item in output_items:
item_d = self._to_dict(out_item)
if item_d and item_d.get("type") == "function_call":
tool_calls.append(
{
"function": {
"name": item_d.get("name", ""),
"arguments": item_d.get("arguments", ""),
}
}
)
if tool_calls:
kv_pairs = OpenTelemetry._tool_calls_kv_pair(tool_calls) # type: ignore
for key, value in kv_pairs.items():
self.safe_set_attribute(
span=span,
key=key,
value=value,
)
# Extract finish reason from ResponsesAPIResponse.status
status = response_obj.get("status")
if status:
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_RESPONSE_FINISH_REASONS.value,
value=safe_dumps([status]),
)
except Exception as e:
self.handle_callback_failure(
callback_name=self.callback_name or "opentelemetry"
@ -1935,6 +2010,78 @@ class OpenTelemetry(CustomLogger):
transformed.append(transformed_msg)
return transformed
@staticmethod
def _to_dict(obj) -> Optional[dict]:
"""Normalize an object to a plain dict.
Handles three forms that appear in practice:
1. Plain ``dict`` returned as-is.
2. LiteLLM's ``BaseLiteLLMOpenAIResponseObject`` — exposes a
``.get()`` method that delegates to ``__dict__``.
3. Raw Pydantic v2 models from the ``openai`` SDK (e.g.
``ResponseOutputMessage``, ``ResponseOutputText``) these do
**not** have ``.get()`` but do have ``.model_dump()``.
Returns ``None`` for anything else so callers can skip it.
"""
if isinstance(obj, dict):
return obj
if hasattr(obj, "get"):
# BaseLiteLLMOpenAIResponseObject duck-type
return obj # type: ignore[return-value]
if hasattr(obj, "model_dump"):
# Raw Pydantic v2 model (e.g. openai SDK types)
return obj.model_dump() # type: ignore[union-attr]
return None
def _transform_responses_api_output_to_otel(self, output: List) -> List[dict]:
"""
Transform Responses API output items into OTEL GenAI 1.38 format.
The Responses API returns output as a list of items, each with a
``type`` field. Message items (``type="message"``) contain a
``content`` list of ``OutputText`` objects with ``type="output_text"``
and ``text`` fields.
Items may be plain dicts, LiteLLM wrapper objects (with ``.get()``),
or raw Pydantic v2 models from the ``openai`` SDK (with
``.model_dump()``). We normalize each item to a dict via
``_to_dict`` before processing.
This method converts them to the same ``{"role": ..., "parts": [...]}``
format used by ``_transform_choices_to_otel_semantic_conventions``.
"""
transformed = []
for raw_item in output:
item = self._to_dict(raw_item)
if item is None:
continue
if item.get("type") == "message":
role = item.get("role", "assistant")
parts = []
for raw_content in item.get("content", []):
content = self._to_dict(raw_content)
if content is None:
continue
if content.get("type") == "output_text":
text = content.get("text", "")
if text:
parts.append({"type": "text", "content": text})
if parts:
transformed.append({"role": role, "parts": parts})
elif item.get("type") == "function_call":
# Surface tool calls from Responses API output
part: dict = {
"type": "tool_call",
"name": item.get("name", ""),
"arguments": item.get("arguments", ""),
}
if item.get("call_id"):
part["id"] = item["call_id"]
transformed.append({"role": "assistant", "parts": [part]})
return transformed
def set_raw_request_attributes(self, span: Span, kwargs, response_obj):
try:
# Only set provider-specific raw payload attributes on this span.

View file

@ -1809,9 +1809,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
Translate messages to anthropic format.
"""
## VALIDATE REQUEST
"""
Anthropic doesn't support tool calling without `tools=` param specified.
"""
"""Anthropic requires ``tools`` when messages include tool blocks; LiteLLM injects a dummy tool if omitted (no ``modify_params`` needed)."""
from litellm.litellm_core_utils.prompt_templates.factory import (
anthropic_messages_pt,
)
@ -1821,16 +1819,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
and messages is not None
and has_tool_call_blocks(messages)
):
if litellm.modify_params:
optional_params["tools"], _ = self._map_tools(
add_dummy_tool(custom_llm_provider="anthropic")
)
else:
raise litellm.UnsupportedParamsError(
message="Anthropic doesn't support tool calling without `tools=` param specified. Pass `tools=` param OR set `litellm.modify_params = True` // `litellm_settings::modify_params: True` to add dummy tool to the request.",
model="",
llm_provider="anthropic",
)
optional_params["tools"], _ = self._map_tools(
add_dummy_tool(custom_llm_provider="anthropic")
)
# Drop thinking param if thinking is enabled but thinking_blocks are missing
# This prevents the error: "Expected thinking or redacted_thinking, but found tool_use"

View file

@ -454,6 +454,16 @@ def anthropic_messages_handler(
"display": "summarized",
}
# Mirror Router._get_timeout: prefer `stream_timeout` when streaming,
# fall back to `timeout`. Coerce string form to float for httpx.
_resolved_timeout = (
litellm_params.stream_timeout
if stream and litellm_params.stream_timeout is not None
else litellm_params.timeout
)
if isinstance(_resolved_timeout, str):
_resolved_timeout = float(_resolved_timeout)
return base_llm_http_handler.anthropic_messages_handler(
model=model,
messages=messages,
@ -469,5 +479,6 @@ def anthropic_messages_handler(
api_key=api_key,
api_base=api_base,
stream=stream,
timeout=_resolved_timeout,
kwargs=kwargs,
)

View file

@ -1428,7 +1428,13 @@ class BaseAWSLLM:
def _sign_request(
self,
service_name: Literal["bedrock", "sagemaker", "bedrock-agentcore", "s3vectors"],
service_name: Literal[
"bedrock",
"sagemaker",
"bedrock-agentcore",
"s3vectors",
"aws-external-anthropic",
],
headers: dict,
optional_params: dict,
request_data: dict,

View file

@ -1,8 +1,79 @@
from datetime import datetime
from typing import Any, Optional, cast
from openai.types.batch import BatchRequestCounts
from openai.types.batch import Metadata as OpenAIBatchMetadata
from litellm.types.utils import LiteLLMBatch
# AWS Bedrock model-invocation-job statuses → OpenAI Batch statuses.
# Mirrors the mapping used by `BedrockBatchesConfig.transform_create_batch_response`
# so create / retrieve return consistent statuses.
_BEDROCK_MIJ_STATUS_TO_OPENAI = {
"Submitted": "validating",
"Validating": "validating",
"Scheduled": "validating",
"InProgress": "in_progress",
"Stopping": "cancelling",
"Stopped": "cancelled",
"Completed": "completed",
"PartiallyCompleted": "completed",
"Failed": "failed",
"Expired": "expired",
}
def _extract_region_from_bedrock_arn(arn: str) -> Optional[str]:
"""ARN shape: ``arn:aws:bedrock:<region>:<account>:<type>/<id>``"""
try:
parts = arn.split(":")
if len(parts) >= 4 and parts[2] == "bedrock":
return parts[3] or None
except Exception:
pass
return None
def _extract_job_id_from_arn(arn: str) -> Optional[str]:
"""``arn:aws:bedrock:<region>:<acct>:model-invocation-job/<job-id>`` -> ``<job-id>``."""
if ":model-invocation-job/" not in arn:
return None
return arn.rsplit("/", 1)[-1] or None
def _predict_output_file_uri(
output_prefix: str, input_uri: str, job_id: Optional[str]
) -> Optional[str]:
"""
Compute the deterministic per-job result file URI Bedrock writes to.
Bedrock lays results out as::
<output_prefix>/<job-id>/<basename(input_uri)>.out
We compute it client-side so OpenAI-style ``client.files.content(output_file_id)``
works without an extra S3 ``ListObjectsV2`` round-trip. Returns ``None`` if we
don't have enough info; callers should fall back to the bare prefix.
"""
if not output_prefix or not input_uri or not job_id:
return None
if not output_prefix.endswith("/"):
output_prefix = output_prefix + "/"
input_basename = input_uri.rsplit("/", 1)[-1]
if not input_basename:
return None
return f"{output_prefix}{job_id}/{input_basename}.out"
def _to_epoch(value: Any) -> Optional[int]:
if value is None:
return None
if isinstance(value, (int, float)):
return int(value)
if isinstance(value, datetime):
return int(value.timestamp())
return None
class BedrockBatchesHandler:
"""
@ -97,3 +168,173 @@ class BedrockBatchesHandler:
with concurrent.futures.ThreadPoolExecutor() as executor:
future = executor.submit(run_in_thread)
return future.result()
@staticmethod
def _handle_model_invocation_job_status(
batch_id: str,
aws_region_name: Optional[str] = None,
logging_obj=None,
**kwargs,
) -> "LiteLLMBatch":
"""
Handle ``GetModelInvocationJob`` status check for AWS Bedrock bulk batch
inference jobs (the ARN type returned by ``CreateModelInvocationJob``).
``CreateModelInvocationJob`` lives on the Bedrock **control plane**
(``bedrock.<region>.amazonaws.com``), distinct from the data-plane
``bedrock-runtime`` endpoint that serves Twelve Labs async-invoke ARNs.
The two ARN families therefore can't share a handler — see
``litellm/batches/main.py`` for the dispatch.
Args:
batch_id: A ``arn:aws:bedrock:<region>:<acct>:model-invocation-job/<id>``
ARN (or just the trailing job id; both are accepted by
``GetModelInvocationJob``).
aws_region_name: Region for the boto3 ``bedrock`` client. If omitted,
we fall back to parsing the region out of ``batch_id`` itself.
logging_obj: Optional litellm logging object.
**kwargs: Optional AWS credential overrides
(``aws_access_key_id``, ``aws_secret_access_key``,
``aws_session_token``, ``aws_profile_name``,
``aws_role_name``, ``aws_session_name``,
``aws_web_identity_token``, ``aws_sts_endpoint``,
``aws_external_id``). Unknown keys are ignored.
Returns:
``LiteLLMBatch`` shaped like an OpenAI Batch resource. Note that
``request_counts`` is always ``(0, 0, 0)`` because
``GetModelInvocationJob`` does not surface per-record counts;
callers that need accurate counts should parse
``manifest.json.out`` from the output S3 prefix.
"""
try:
import boto3
except ImportError as exc:
raise ImportError(
"Missing boto3 to call bedrock. Run 'pip install boto3'."
) from exc
# Resolve region: explicit > parsed-from-ARN > us-east-1 (boto3 default).
region = (
aws_region_name or _extract_region_from_bedrock_arn(batch_id) or "us-east-1"
)
# Resolve credentials through the same path the rest of the bedrock
# provider uses, so model_list / env / role-assumption configs are
# honored. We instantiate BedrockBatchesConfig (which extends
# BaseAWSLLM) lazily to avoid a circular import at module load.
from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig
creds = BedrockBatchesConfig().get_credentials(
aws_access_key_id=kwargs.get("aws_access_key_id"),
aws_secret_access_key=kwargs.get("aws_secret_access_key"),
aws_session_token=kwargs.get("aws_session_token"),
aws_region_name=region,
aws_session_name=kwargs.get("aws_session_name"),
aws_profile_name=kwargs.get("aws_profile_name"),
aws_role_name=kwargs.get("aws_role_name"),
aws_web_identity_token=kwargs.get("aws_web_identity_token"),
aws_sts_endpoint=kwargs.get("aws_sts_endpoint"),
aws_external_id=kwargs.get("aws_external_id"),
)
client = boto3.client(
"bedrock",
region_name=region,
aws_access_key_id=creds.access_key,
aws_secret_access_key=creds.secret_key,
aws_session_token=creds.token,
)
if logging_obj is not None:
# Use the bare job id in the logged URL so we don't double up the
# `model-invocation-job/` segment when `batch_id` is a full ARN.
# `GetModelInvocationJob` accepts either form, but only the bare id
# produces a sensible-looking URL in logs.
url_path_id = _extract_job_id_from_arn(batch_id) or batch_id
logging_obj.pre_call(
input=batch_id,
api_key="",
additional_args={
"complete_input_dict": {"jobIdentifier": batch_id},
"api_base": (
f"https://bedrock.{region}.amazonaws.com/"
f"model-invocation-job/{url_path_id}"
),
},
)
response = client.get_model_invocation_job(jobIdentifier=batch_id)
if logging_obj is not None:
logging_obj.post_call(
input=batch_id,
api_key="",
original_response=response,
additional_args={"complete_input_dict": {"jobIdentifier": batch_id}},
)
bedrock_status = str(response.get("status", ""))
openai_status = cast(
Any,
_BEDROCK_MIJ_STATUS_TO_OPENAI.get(bedrock_status, "in_progress"),
)
input_uri = (
response.get("inputDataConfig", {})
.get("s3InputDataConfig", {})
.get("s3Uri", "")
)
output_prefix = (
response.get("outputDataConfig", {})
.get("s3OutputDataConfig", {})
.get("s3Uri", "")
)
# Bedrock returns the output *prefix* the user supplied at job creation.
# Actual results land at <prefix>/<job-id>/<basename(input)>.out — we
# surface that single-file URI as `output_file_id` so the OpenAI-style
# download flow works without an extra S3 listing call. We deliberately
# do NOT fall back to the bare prefix when prediction fails: a prefix
# is not a downloadable object, so handing it back as `output_file_id`
# would reproduce the very NoSuchKey bug this handler exists to fix.
# The bare prefix is preserved in metadata for callers that want the
# `manifest.json.out` or want to do their own listing.
job_arn = response.get("jobArn", batch_id)
job_id = _extract_job_id_from_arn(job_arn)
output_file_uri = _predict_output_file_uri(output_prefix, input_uri, job_id)
completed_at = _to_epoch(response.get("endTime"))
# Note: metadata uses "" (not None) for unknown URIs to satisfy the
# OpenAI Batch metadata schema, which is `dict[str, str]`. The
# `output_file_id` field on the LiteLLMBatch itself does carry None
# correctly (see below), so callers should branch on that, not on
# `metadata["output_file_uri"]`.
openai_batch_metadata: OpenAIBatchMetadata = {
"model_arn": response.get("modelId", ""),
"job_arn": job_arn,
"job_name": response.get("jobName", ""),
"failure_message": response.get("message") or "",
"input_s3_uri": input_uri,
"output_s3_uri": output_prefix,
"output_file_uri": output_file_uri or "",
}
return LiteLLMBatch(
id=job_arn,
object="batch",
status=openai_status,
created_at=_to_epoch(response.get("submitTime")) or 0,
in_progress_at=_to_epoch(response.get("lastModifiedTime")),
completed_at=completed_at if openai_status == "completed" else None,
failed_at=completed_at if openai_status == "failed" else None,
cancelled_at=completed_at if openai_status == "cancelled" else None,
expired_at=completed_at if openai_status == "expired" else None,
request_counts=BatchRequestCounts(total=0, completed=0, failed=0),
metadata=openai_batch_metadata,
completion_window="24h",
endpoint="/v1/chat/completions",
input_file_id=input_uri,
output_file_id=output_file_uri if openai_status == "completed" else None,
)

View file

@ -0,0 +1,8 @@
from .transformation import (
BedrockClaudePlatformConfig,
)
from .messages_transformation import (
BedrockClaudePlatformMessagesConfig,
)
__all__ = ["BedrockClaudePlatformConfig", "BedrockClaudePlatformMessagesConfig"]

View file

@ -0,0 +1,107 @@
from typing import Literal, Optional, Tuple
import litellm
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.secret_managers.main import get_secret_str
CLAUDE_PLATFORM_SERVICE_NAME: Literal["aws-external-anthropic"] = (
"aws-external-anthropic"
)
CLAUDE_PLATFORM_BEDROCK_ROUTE = "claude_platform/"
def strip_claude_platform_route(model: str) -> str:
if model.startswith(CLAUDE_PLATFORM_BEDROCK_ROUTE):
return model.replace(CLAUDE_PLATFORM_BEDROCK_ROUTE, "", 1)
return model
class BedrockClaudePlatformMixin(BaseAWSLLM):
@staticmethod
def _get_workspace_id(optional_params: dict, litellm_params: dict) -> Optional[str]:
workspace_id = (
optional_params.get("workspace_id")
or litellm_params.get("workspace_id")
or optional_params.get("aws_workspace_id")
or litellm_params.get("aws_workspace_id")
or optional_params.get("anthropic-workspace-id")
or litellm_params.get("anthropic-workspace-id")
)
if workspace_id is None:
workspace_id = optional_params.get(
"anthropic_workspace_id"
) or litellm_params.get("anthropic_workspace_id")
if workspace_id is not None:
return str(workspace_id)
return get_secret_str("ANTHROPIC_AWS_WORKSPACE_ID") or get_secret_str(
"ANTHROPIC_WORKSPACE_ID"
)
def _get_required_aws_region_name(self, optional_params: dict) -> str:
aws_region_name = (
optional_params.get("aws_region_name")
or get_secret_str("AWS_REGION_NAME")
or get_secret_str("AWS_REGION")
or get_secret_str("AWS_DEFAULT_REGION")
)
if aws_region_name is None:
raise litellm.AuthenticationError(
message=(
"Missing AWS region for Claude Platform on AWS. Pass "
"`aws_region_name` or set a standard AWS region environment value."
),
llm_provider="bedrock",
model="",
)
self._validate_aws_region_name(str(aws_region_name))
return str(aws_region_name)
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
api_base = (
api_base
or litellm.api_base
or get_secret_str("ANTHROPIC_AWS_BASE_URL")
or get_secret_str("ANTHROPIC_AWS_API_BASE")
)
if api_base is None:
aws_region_name = self._get_required_aws_region_name(optional_params)
api_base = (
f"https://{CLAUDE_PLATFORM_SERVICE_NAME}.{aws_region_name}.api.aws"
)
if not api_base.endswith("/v1/messages"):
api_base = f"{api_base.rstrip('/')}/v1/messages"
return api_base
def sign_request(
self,
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
api_key: Optional[str] = None,
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, Optional[bytes]]:
if api_key or get_secret_str("ANTHROPIC_AWS_API_KEY"):
return headers, None
return self._sign_request(
service_name=CLAUDE_PLATFORM_SERVICE_NAME,
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
model=model,
stream=stream,
fake_stream=fake_stream,
)

View file

@ -0,0 +1,71 @@
from typing import Any, Dict, List, Optional, Tuple
import litellm
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
DEFAULT_ANTHROPIC_API_VERSION,
AnthropicMessagesConfig,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.router import GenericLiteLLMParams
from .common_utils import BedrockClaudePlatformMixin, strip_claude_platform_route
class BedrockClaudePlatformMessagesConfig(
BedrockClaudePlatformMixin, AnthropicMessagesConfig
):
def validate_anthropic_messages_environment(
self,
headers: dict,
model: str,
messages: List[Any],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> Tuple[dict, Optional[str]]:
workspace_id = self._get_workspace_id(optional_params, litellm_params)
if workspace_id is None:
raise litellm.AuthenticationError(
message=(
"Missing workspace ID for Claude Platform on AWS. Pass "
"`workspace_id` or configure the provider workspace setting."
),
llm_provider="bedrock",
model=model,
)
resolved_api_key = api_key or get_secret_str("ANTHROPIC_AWS_API_KEY")
headers = {
**headers,
"anthropic-version": headers.get(
"anthropic-version", DEFAULT_ANTHROPIC_API_VERSION
),
"content-type": headers.get("content-type", "application/json"),
"anthropic-workspace-id": workspace_id,
}
if resolved_api_key and "x-api-key" not in headers:
headers["x-api-key"] = resolved_api_key
headers = self._update_headers_with_anthropic_beta(
headers=headers,
optional_params=optional_params,
)
return headers, api_base
def transform_anthropic_messages_request(
self,
model: str,
messages: List[Dict],
anthropic_messages_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Dict:
return super().transform_anthropic_messages_request(
model=strip_claude_platform_route(model),
messages=messages,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)

View file

@ -0,0 +1,94 @@
from typing import Any, Dict, List, Optional
import litellm
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from .common_utils import BedrockClaudePlatformMixin
class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig):
"""
Bedrock Claude Platform uses Anthropic's Messages API with AWS gateway auth.
"""
@property
def custom_llm_provider(self) -> Optional[str]:
return "bedrock"
def validate_environment(
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> Dict:
workspace_id = self._get_workspace_id(optional_params, litellm_params)
if workspace_id is None:
raise litellm.AuthenticationError(
message=(
"Missing workspace ID for Claude Platform on AWS. Pass "
"`workspace_id` or configure the provider workspace setting."
),
llm_provider="bedrock",
model=model,
)
api_key = api_key or get_secret_str("ANTHROPIC_AWS_API_KEY")
anthropic_headers = self.get_anthropic_headers(
api_key=api_key,
auth_token=None,
computer_tool_used=self.is_computer_tool_used(
tools=optional_params.get("tools")
),
prompt_caching_set=self.is_cache_control_set(messages=messages),
pdf_used=self.is_pdf_used(messages=messages),
file_id_used=self.is_file_id_used(messages=messages),
mcp_server_used=self.is_mcp_server_used(
mcp_servers=optional_params.get("mcp_servers")
),
web_search_tool_used=self.is_web_search_tool_used(
tools=optional_params.get("tools")
),
tool_search_used=self.is_tool_search_used(
tools=optional_params.get("tools")
),
programmatic_tool_calling_used=self.is_programmatic_tool_calling_used(
tools=optional_params.get("tools")
),
input_examples_used=self.is_input_examples_used(
tools=optional_params.get("tools")
),
effort_used=self.is_effort_used(
optional_params=optional_params, model=model
),
user_anthropic_beta_headers=self._get_user_anthropic_beta_headers(
anthropic_beta_header=headers.get("anthropic-beta")
),
code_execution_tool_used=self.is_code_execution_tool_used(
tools=optional_params.get("tools")
),
container_with_skills_used=self.is_container_with_skills_used(
optional_params=optional_params
),
)
anthropic_headers["anthropic-workspace-id"] = workspace_id
return {**headers, **anthropic_headers}
def get_model_response_iterator(
self,
streaming_response: Any,
sync_stream: bool,
json_mode: Optional[bool] = False,
) -> Any:
from litellm.llms.anthropic.chat.handler import ModelResponseIterator
return ModelResponseIterator(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=bool(json_mode),
)

View file

@ -692,6 +692,7 @@ class BedrockModelInfo(BaseLLMModelInfo):
) -> Literal[
"converse",
"invoke",
"claude_platform",
"converse_like",
"agent",
"agentcore",
@ -706,6 +707,7 @@ class BedrockModelInfo(BaseLLMModelInfo):
str,
Literal[
"invoke",
"claude_platform",
"converse_like",
"converse",
"agent",
@ -716,6 +718,7 @@ class BedrockModelInfo(BaseLLMModelInfo):
],
] = {
"invoke/": "invoke",
"claude_platform/": "claude_platform",
"converse_like/": "converse_like",
"converse/": "converse",
"agent/": "agent",
@ -753,6 +756,36 @@ class BedrockModelInfo(BaseLLMModelInfo):
"""
return "converse/" in model
@staticmethod
def _explicit_claude_platform_route(model: str) -> bool:
"""
Check if the model is an explicit Claude Platform on AWS route.
"""
return "claude_platform/" in model
@staticmethod
def get_claude_platform_model(model: str) -> str:
"""
Strip the Claude Platform route prefix from a Bedrock model name.
"""
return model.replace("claude_platform/", "", 1)
@staticmethod
def map_claude_platform_auth_params(
passed_params: dict, optional_params: dict
) -> dict:
"""
Map Claude Platform route auth params that are not OpenAI request params.
"""
for key in (
"workspace_id",
"aws_workspace_id",
"anthropic_workspace_id",
):
if key in passed_params:
optional_params[key] = passed_params[key]
return optional_params
@staticmethod
def _explicit_invoke_route(model: str) -> bool:
"""
@ -815,6 +848,12 @@ class BedrockModelInfo(BaseLLMModelInfo):
All other routes should return None since they will go through litellm.completion
"""
#########################################################
# Claude Platform route uses Anthropic Messages API via the AWS gateway.
#########################################################
if BedrockModelInfo._explicit_claude_platform_route(model):
return litellm.BedrockClaudePlatformMessagesConfig()
#########################################################
# Converse routes should go through litellm.completion()
if BedrockModelInfo._explicit_converse_route(model):
@ -860,7 +899,9 @@ def get_bedrock_chat_config(model: str):
base_model = BedrockModelInfo.get_base_model(model)
# Handle explicit routes first
if bedrock_route == "converse" or bedrock_route == "converse_like":
if bedrock_route == "claude_platform":
return litellm.BedrockClaudePlatformConfig()
elif bedrock_route == "converse" or bedrock_route == "converse_like":
return litellm.AmazonConverseConfig()
elif bedrock_route == "openai":
return litellm.AmazonBedrockOpenAIConfig()

View file

@ -408,6 +408,47 @@ class AmazonAnthropicClaudeMessagesConfig(
if self._supports_tool_search_on_bedrock(model):
beta_set.add("tool-search-tool-2025-10-19")
@staticmethod
def _filter_context_management_for_bedrock_invoke(
anthropic_messages_request: Dict,
beta_set: set,
) -> None:
"""
Bedrock InvokeModel accepts ``context_management`` only when it carries
``compact_20260112`` edits paired with the ``compact-2026-01-12``
anthropic-beta header. Other edit types (notably ``clear_thinking_20251015``,
which Claude Code sends on every request) are LiteLLM-internal and would
cause Bedrock to 400 with ``"context_management: Extra inputs are not
permitted"``.
Filter the edits list to the supported subset, add the beta header when
compact edits remain, and drop ``context_management`` entirely when no
supported edits are left so the safety-net allowlist can pass it through.
Ref: https://github.com/BerriAI/litellm/issues/27532
"""
cm = anthropic_messages_request.get("context_management")
if not isinstance(cm, dict):
return
edits = cm.get("edits")
if not isinstance(edits, list):
anthropic_messages_request.pop("context_management", None)
return
compact_edits = [
e
for e in edits
if isinstance(e, dict) and e.get("type") == "compact_20260112"
]
if compact_edits:
beta_set.add("compact-2026-01-12")
anthropic_messages_request["context_management"] = {
**cm,
"edits": compact_edits,
}
else:
anthropic_messages_request.pop("context_management", None)
def _convert_output_format_to_inline_schema(
self,
output_format: Dict,
@ -551,6 +592,11 @@ class AmazonAnthropicClaudeMessagesConfig(
if injected_thinking_for_clear_thinking:
beta_set.add("interleaved-thinking-2025-05-14")
self._filter_context_management_for_bedrock_invoke(
anthropic_messages_request=anthropic_messages_request,
beta_set=beta_set,
)
self._get_tool_search_beta_header_for_bedrock(
model=model,
tool_search_used=tool_search_used,
@ -597,8 +643,9 @@ class AmazonAnthropicClaudeMessagesConfig(
anthropic_messages_request.pop("output_config", None)
# 7. Final safety net: filter top-level fields to the Bedrock Invoke allowlist.
# Catches Anthropic-only extensions (context_management, output_config, speed,
# mcp_servers, ...) and any future additions Claude Code may start sending.
# Catches Anthropic-only extensions (output_config, speed, mcp_servers, ...)
# and any future additions Claude Code may start sending. ``context_management``
# has already been pre-filtered to its Bedrock-supported subset above.
allowed = self.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS
stripped = sorted(k for k in anthropic_messages_request if k not in allowed)
if stripped:

View file

@ -1860,6 +1860,7 @@ class BaseLLMHTTPHandler:
litellm_params: GenericLiteLLMParams,
api_key: Optional[str],
model: str,
timeout: Optional[Union[float, httpx.Timeout]] = None,
) -> httpx.Response:
max_attempts = max(
provider_config.max_retry_on_anthropic_messages_http_error, 1
@ -1874,6 +1875,7 @@ class BaseLLMHTTPHandler:
data=signed_json_body or json.dumps(request_body),
stream=stream or False,
logging_obj=logging_obj,
timeout=timeout,
)
response.raise_for_status()
return response
@ -1928,6 +1930,7 @@ class BaseLLMHTTPHandler:
api_key: Optional[str] = None,
api_base: Optional[str] = None,
stream: Optional[bool] = False,
timeout: Optional[Union[float, httpx.Timeout]] = None,
kwargs: Optional[Dict[str, Any]] = None,
) -> Union[AnthropicMessagesResponse, AsyncIterator]:
from litellm.litellm_core_utils.get_provider_specific_headers import (
@ -2065,6 +2068,7 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params,
api_key=api_key,
model=model,
timeout=timeout,
)
# used for logging + cost tracking
@ -2131,6 +2135,7 @@ class BaseLLMHTTPHandler:
api_key: Optional[str] = None,
api_base: Optional[str] = None,
stream: Optional[bool] = False,
timeout: Optional[Union[float, httpx.Timeout]] = None,
kwargs: Optional[Dict[str, Any]] = None,
) -> Union[
AnthropicMessagesResponse,
@ -2153,6 +2158,7 @@ class BaseLLMHTTPHandler:
api_key=api_key,
api_base=api_base,
stream=stream,
timeout=timeout,
kwargs=kwargs,
)
raise ValueError("anthropic_messages_handler is not implemented for sync calls")

View file

@ -3,7 +3,10 @@
from typing import Optional, Union
import litellm
from litellm.utils import _is_explicitly_disabled_factory, _supports_factory
from litellm.utils import (
_is_explicitly_disabled_factory,
_supports_factory,
)
from .gpt_transformation import OpenAIGPTConfig

View file

@ -156,5 +156,17 @@ class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
text = response_json.get("text") or response_json.get("transcript") or ""
response = TranscriptionResponse(text=text)
# OVHCloud field migration (deadline: 2026-05-11):
# `duration` is replaced by `seconds` in STT responses.
# Prefer `seconds`, fall back to `duration`, normalize to `duration`
# so downstream consumers see a consistent key.
duration = (
response_json["seconds"]
if "seconds" in response_json and response_json["seconds"] is not None
else response_json.get("duration")
)
if duration is not None:
response_json["duration"] = duration
response._hidden_params = response_json
return response

View file

@ -13,6 +13,7 @@ from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.llms.ovhcloud.utils import OVHCloudException
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.openai import AllMessageValues
@ -98,10 +99,16 @@ class OVHCloudChatCompletionStreamingHandler(BaseModelResponseIterator):
new_choices = []
for choice in chunk["choices"]:
if "delta" in choice and "reasoning" in choice["delta"]:
choice["delta"]["reasoning_content"] = choice["delta"].get(
"reasoning"
)
if "delta" in choice:
delta = choice["delta"]
# OVHCloud field migration (deadline: 2026-05-11):
# `reasoning_content` is replaced by `reasoning`.
# Normalise to `reasoning_content` so downstream consumers
# see a consistent key during the transition window.
reasoning_new = delta.get("reasoning")
reasoning_legacy = delta.get("reasoning_content")
if reasoning_new is not None and reasoning_legacy is None:
delta["reasoning_content"] = reasoning_new
new_choices.append(choice)
return ModelResponseStream(

View file

@ -1,3 +1,5 @@
# LiteLLM main module: public completion, embedding, streaming, and moderation entrypoints.
#
# +-----------------------------------------------+
# | |
# | Give Feedback / Get Help |
@ -59,7 +61,13 @@ import litellm
from litellm import client
# Other utils are imported directly to avoid circular imports
from litellm.utils import exception_type, get_litellm_params, get_optional_params
from litellm.utils import (
exception_type,
get_litellm_params,
get_optional_params,
peek_reasoning_summary_aliases,
strip_reasoning_summary_aliases_from_optional_params,
)
# Logging is imported lazily when needed to avoid loading litellm_logging at import time
if TYPE_CHECKING:
@ -946,6 +954,7 @@ def responses_api_bridge_check(
web_search_options: Optional[OpenAIWebSearchOptions] = None,
tools: Optional[List[Any]] = None,
reasoning_effort: Optional[Any] = None,
reasoning_summary: Optional[Any] = None,
) -> Tuple[dict, str]:
model_info: Dict[str, Any] = {}
@ -982,14 +991,23 @@ def responses_api_bridge_check(
mode = "responses"
model_info["mode"] = mode
# OpenAI/Azure gpt-5.4+ chat-completions calls with both tools + reasoning_effort
# must be bridged to Responses API.
# OpenAI/Azure GPT-5 chat-completions that need Responses-only fields (e.g.
# ``reasoningSummary`` in ``extra_body``) must be bridged; Chat Completions rejects
# those keys.
#
# - gpt-5.4+: tools + reasoning_effort (original) or any reasoning-summary alias.
# - Older GPT-5 names (e.g. ``gpt-5``, ``gpt-5.1``): bridge only when a reasoning
# summary alias is present with ``reasoning_effort`` (tools alone stay on chat).
if (
custom_llm_provider in ("openai", "azure")
and OpenAIGPT5Config.is_model_gpt_5_4_plus_model(model)
and tools
and reasoning_effort is not None
and model_info.get("mode") != "responses"
and OpenAIGPT5Config.is_model_gpt_5_model(model)
and not OpenAIGPT5Config.is_model_gpt_5_search_model(model)
and reasoning_effort is not None
and (
reasoning_summary is not None
or (OpenAIGPT5Config.is_model_gpt_5_4_plus_model(model) and tools)
)
):
model_info["mode"] = "responses"
model = model.replace("responses/", "")
@ -1634,8 +1652,10 @@ def completion( # type: ignore # noqa: PLR0915
## RESPONSES API BRIDGE LOGIC ## - check if model has 'mode: responses' in litellm.model_cost map
# Only run the second bridge check if the first one didn't already
# detect responses mode (e.g. via the "responses/" prefix). The second
# check handles cases like gpt-5.4+ with tools+reasoning_effort that
# the first (early) check doesn't cover.
# check handles cases like gpt-5.4+ with tools+reasoning_effort or
# reasoningSummary/reasoning_summary without tools (AI SDK) that the first
# (early) check doesn't cover.
_reasoning_summary_for_bridge = peek_reasoning_summary_aliases(optional_params)
if responses_api_model_info.get("mode") != "responses":
responses_api_model_info, model = responses_api_bridge_check(
model=model,
@ -1643,14 +1663,29 @@ def completion( # type: ignore # noqa: PLR0915
web_search_options=web_search_options,
tools=tools,
reasoning_effort=reasoning_effort,
reasoning_summary=_reasoning_summary_for_bridge,
)
if responses_api_model_info.get("mode") == "responses":
from litellm.completion_extras import responses_api_bridge
optional_params, rs_val = (
strip_reasoning_summary_aliases_from_optional_params(optional_params)
)
if isinstance(reasoning_effort, dict) and "summary" in reasoning_effort:
optional_params = dict(optional_params)
optional_params["reasoning_effort"] = reasoning_effort
elif rs_val is not None:
eff = optional_params.get("reasoning_effort", reasoning_effort)
if isinstance(eff, dict):
optional_params["reasoning_effort"] = {**eff, "summary": rs_val}
elif eff is not None:
optional_params["reasoning_effort"] = {
"effort": eff,
"summary": rs_val,
}
else:
optional_params["reasoning_effort"] = {"summary": rs_val}
return responses_api_bridge.completion(
model=model,
@ -1669,6 +1704,16 @@ def completion( # type: ignore # noqa: PLR0915
encoding=_get_encoding(),
stream=stream,
)
elif (
custom_llm_provider == "openai"
and OpenAIGPT5Config.is_model_gpt_5_model(model)
) or (
custom_llm_provider == "azure"
and litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model)
):
optional_params, _ = strip_reasoning_summary_aliases_from_optional_params(
optional_params
)
if custom_llm_provider == "azure":
# azure configs
@ -3813,7 +3858,33 @@ def completion( # type: ignore # noqa: PLR0915
)
bedrock_route = BedrockModelInfo.get_bedrock_route(model)
if bedrock_route == "converse":
if bedrock_route == "claude_platform":
provider_config = ProviderConfigManager.get_provider_chat_config(
model=model,
provider=LlmProviders.BEDROCK,
)
model = BedrockModelInfo.get_claude_platform_model(model)
response = base_llm_http_handler.completion(
model=model,
stream=stream,
messages=messages,
acompletion=acompletion,
api_base=api_base,
model_response=model_response,
optional_params=optional_params,
litellm_params=litellm_params,
shared_session=shared_session,
custom_llm_provider="bedrock",
timeout=timeout,
headers=headers,
encoding=_get_encoding(),
api_key=api_key,
logging_obj=logging,
client=client,
provider_config=provider_config,
)
return response
elif bedrock_route == "converse":
model = model.replace("converse/", "")
response = bedrock_converse_chat_completion.completion(
model=model,

View file

@ -12,7 +12,8 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
validate_loopback_redirect_uri,
get_request_base_url,
validate_trusted_redirect_uri,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
@ -29,51 +30,6 @@ router = APIRouter(
)
def get_request_base_url(request: Request) -> str:
"""
Get the base URL for the request, considering X-Forwarded-* headers.
X-Forwarded-Proto / X-Forwarded-Host / X-Forwarded-Port are only honoured
when the request comes from a configured trusted proxy
(``use_x_forwarded_for`` enabled AND caller in ``mcp_trusted_proxy_ranges``).
Otherwise the request's literal ``base_url`` is returned, so an
untrusted caller cannot poison OAuth-discovery / redirect_uri values
by injecting headers.
Args:
request: FastAPI Request object
Returns:
The reconstructed base URL (e.g., "https://proxy.example.com")
"""
base_url = str(request.base_url).rstrip("/")
parsed = urlparse(base_url)
if not IPAddressUtils.is_request_from_trusted_proxy(request):
return base_url
x_forwarded_proto = request.headers.get("X-Forwarded-Proto")
x_forwarded_host = request.headers.get("X-Forwarded-Host")
x_forwarded_port = request.headers.get("X-Forwarded-Port")
scheme = x_forwarded_proto if x_forwarded_proto else parsed.scheme
if x_forwarded_host:
# X-Forwarded-Host may already include port (e.g., "example.com:8080")
if ":" in x_forwarded_host and not x_forwarded_host.startswith("["):
netloc = x_forwarded_host
elif x_forwarded_port:
netloc = f"{x_forwarded_host}:{x_forwarded_port}"
else:
netloc = x_forwarded_host
else:
netloc = parsed.netloc
if x_forwarded_port and ":" not in netloc:
netloc = f"{netloc}:{x_forwarded_port}"
return urlunparse((scheme, netloc, parsed.path, "", "", ""))
def encode_state_with_base_url(
base_url: str,
original_state: str,
@ -127,12 +83,14 @@ def decode_state_hash(encrypted_state: str) -> dict:
return state_data
def _get_validated_client_redirect_uri(state_data: Dict[str, Any]) -> str:
"""Return a loopback client redirect URI from OAuth state."""
def _get_validated_client_redirect_uri(
request: Request, state_data: Dict[str, Any]
) -> str:
"""Return a trusted (same-origin or loopback) client redirect URI from OAuth state."""
redirect_uri = state_data.get("client_redirect_uri") or state_data.get("base_url")
if not redirect_uri or not isinstance(redirect_uri, str):
raise HTTPException(status_code=400, detail="Invalid redirect URI")
validate_loopback_redirect_uri(redirect_uri)
validate_trusted_redirect_uri(request, redirect_uri)
return redirect_uri
@ -338,12 +296,12 @@ async def authorize_with_server(
status_code=400, detail="MCP server authorization url is not set"
)
# Loopback-only redirect_uri. The URI is encrypted into the OAuth
# state and decoded on /callback to redirect the user back; a non-
# loopback URI would be an open-redirect + code-theft primitive
# (VERIA-57 root cause B). MCP clients are native apps — loopback is
# the spec-compliant callback pattern.
validate_loopback_redirect_uri(redirect_uri)
# Loopback OR same-origin redirect_uri. The URI is encrypted into the
# OAuth state and decoded on /callback to redirect the user back;
# restricting to trusted origins blocks the open-redirect +
# code-theft primitive (VERIA-57 root cause B). Loopback supports
# native MCP clients; same-origin supports the proxy's own UI callback.
validate_trusted_redirect_uri(request, redirect_uri)
parsed = urlparse(redirect_uri)
base_url = urlunparse(parsed._replace(query=""))
request_base_url = get_request_base_url(request)
@ -660,17 +618,18 @@ async def token_endpoint(
@router.get("/callback")
async def callback(code: str, state: str):
async def callback(request: Request, code: str, state: str):
try:
state_data = decode_state_hash(state)
original_state = state_data["original_state"]
# Re-validate loopback at the sink. /authorize rejects non-loopback
# Re-validate at the sink. /authorize rejects untrusted
# redirect_uri before encoding into state, but encrypted states
# minted before that check was added have no expiry and remain
# valid indefinitely. Validating here blocks the open-redirect +
# code-theft primitive even for pre-fix states.
redirect_uri = _get_validated_client_redirect_uri(state_data)
# valid indefinitely. Validating here (same-origin OR loopback)
# blocks the open-redirect + code-theft primitive even for pre-fix
# states while allowing the UI's same-origin callback to work.
redirect_uri = _get_validated_client_redirect_uri(request, state_data)
params = {"code": code, "state": original_state}
complete_returned_url = _append_query_params(redirect_uri, params)

View file

@ -2,15 +2,63 @@
(BYOK + discoverable / pass-through OAuth proxy)."""
from ipaddress import ip_address
from urllib.parse import urlparse
from urllib.parse import urlparse, urlunparse
from fastapi import HTTPException
from fastapi import HTTPException, Request
from litellm._logging import verbose_logger
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
# RFC 6749 §5.1 / OAuth 2.1 draft-15 §4.1.3: token-endpoint responses
# must not be cached — both success and error bodies may reveal secrets.
TOKEN_NO_CACHE_HEADERS = {"Cache-Control": "no-store", "Pragma": "no-cache"}
def get_request_base_url(request: Request) -> str:
"""
Get the base URL for the request, considering X-Forwarded-* headers.
X-Forwarded-Proto / X-Forwarded-Host / X-Forwarded-Port are only honoured
when the request comes from a configured trusted proxy
(``use_x_forwarded_for`` enabled AND caller in ``mcp_trusted_proxy_ranges``).
Otherwise the request's literal ``base_url`` is returned, so an
untrusted caller cannot poison OAuth-discovery / redirect_uri values
by injecting headers.
Args:
request: FastAPI Request object
Returns:
The reconstructed base URL (e.g., "https://proxy.example.com")
"""
base_url = str(request.base_url).rstrip("/")
parsed = urlparse(base_url)
if not IPAddressUtils.is_request_from_trusted_proxy(request):
return base_url
x_forwarded_proto = request.headers.get("X-Forwarded-Proto")
x_forwarded_host = request.headers.get("X-Forwarded-Host")
x_forwarded_port = request.headers.get("X-Forwarded-Port")
scheme = x_forwarded_proto if x_forwarded_proto else parsed.scheme
if x_forwarded_host:
# X-Forwarded-Host may already include port (e.g., "example.com:8080")
if ":" in x_forwarded_host and not x_forwarded_host.startswith("["):
netloc = x_forwarded_host
elif x_forwarded_port:
netloc = f"{x_forwarded_host}:{x_forwarded_port}"
else:
netloc = x_forwarded_host
else:
netloc = parsed.netloc
if x_forwarded_port and ":" not in netloc:
netloc = f"{netloc}:{x_forwarded_port}"
return urlunparse((scheme, netloc, parsed.path, "", "", ""))
def validate_loopback_redirect_uri(redirect_uri: str) -> None:
"""Require a loopback ``redirect_uri`` (OAuth 2.1 §4.1.2.1 + RFC 8252
§7.3 native-app pattern). MCP clients are native apps that listen on
@ -46,3 +94,60 @@ def validate_loopback_redirect_uri(redirect_uri: str) -> None:
# don't let it bubble up as a 500.
pass
raise HTTPException(status_code=400, detail="invalid_request")
def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None:
"""Accept same-origin (proxy's own origin) OR loopback ``redirect_uri``.
Same-origin is required for the LiteLLM UI's OAuth flow: the UI
redirects to ``<proxy>/ui/mcp/oauth/callback`` which is not loopback
but is on the proxy's own trusted HTTPS origin. An attacker cannot
host content on the proxy's own origin without already owning the
proxy, so the open-redirect / code-theft primitive that motivated
:func:`validate_loopback_redirect_uri` does not apply here.
Loopback continues to be accepted for native MCP clients (per
OAuth 2.1 §4.1.2.1 + RFC 8252 §7.3).
Use this in the discoverable OAuth proxy endpoints that serve both
native clients and the proxy's own UI. BYOK endpoints that only
support native clients should keep
:func:`validate_loopback_redirect_uri`.
"""
try:
parsed = urlparse(redirect_uri)
except ValueError:
raise HTTPException(status_code=400, detail="invalid_request")
if parsed.scheme not in ("http", "https"):
raise HTTPException(status_code=400, detail="invalid_request")
if parsed.fragment:
raise HTTPException(status_code=400, detail="invalid_request")
# Same-origin: scheme + netloc (host[:port]) must match the proxy's
# own base URL at this request (honouring trusted X-Forwarded-*).
try:
proxy_base = urlparse(get_request_base_url(request))
if (
parsed.netloc
and parsed.scheme == proxy_base.scheme
and parsed.netloc.lower() == proxy_base.netloc.lower()
):
return
except Exception as exc:
# If we can't determine the proxy's origin, fall through to
# loopback. Log so the failure is diagnosable in production.
verbose_logger.warning(
"validate_trusted_redirect_uri: could not determine proxy origin, "
"falling back to loopback-only check. error=%s",
exc,
)
host = (parsed.hostname or "").lower()
if host == "localhost":
return
try:
if ip_address(host).is_loopback:
return
except ValueError:
pass
raise HTTPException(status_code=400, detail="invalid_request")

View file

@ -239,6 +239,7 @@ class KeyManagementRoutes(str, enum.Enum):
KEY_BLOCK = "/key/block"
KEY_UNBLOCK = "/key/unblock"
KEY_BULK_UPDATE = "/key/bulk_update"
TEAM_KEY_BULK_UPDATE = "/team/key/bulk_update"
KEY_RESET_SPEND = "/key/{key_id}/reset_spend"
# info and health routes
@ -540,6 +541,7 @@ class LiteLLMRoutes(enum.Enum):
KeyManagementRoutes.KEY_BLOCK.value,
KeyManagementRoutes.KEY_UNBLOCK.value,
KeyManagementRoutes.KEY_BULK_UPDATE.value,
KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value,
KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value,
KeyManagementRoutes.SPEND_LOGS.value,
KeyManagementRoutes.KEY_RESET_SPEND.value,
@ -675,6 +677,10 @@ class LiteLLMRoutes(enum.Enum):
"/global/activity",
"/global/activity/model",
"/global/activity/cache_hits",
# Tag usage endpoints scope internal users to tags produced by
# their own keys in tag_management_endpoints.py.
"/tag/daily/activity",
"/tag/list",
"/v1/models/{model_id}",
"/models/{model_id}",
"/guardrails/list",
@ -687,7 +693,16 @@ class LiteLLMRoutes(enum.Enum):
+ compliance_check_routes
)
internal_user_view_only_routes = spend_tracking_routes
internal_user_view_only_routes = (
spend_tracking_routes
+ compliance_check_routes
+ [
# Tag usage endpoints scope internal viewers to tags produced by
# their own keys in tag_management_endpoints.py.
"/tag/daily/activity",
"/tag/list",
]
)
self_managed_routes = [
"/team/member_add",
@ -4346,10 +4361,16 @@ class JWTRoutingOverride(BaseModel):
A rule matches when all provided selectors match token claims.
If matched, request is routed to the configured auth path.
Wildcard selectors use shell-style patterns (* and ?) and are matched with
case-sensitive semantics; use the same casing your IdP emits in JWT claims.
Space-delimited tokenization applies only to the ``scope`` claim (OAuth/OIDC
scope strings), not to ``iss``, ``aud``, or ``client_id``.
"""
iss: Union[str, List[str]]
client_id: Optional[Union[str, List[str]]] = None
scope: Optional[Union[str, List[str]]] = None
aud: Optional[Union[str, List[str]]] = None
path: Literal["oauth2"] = "oauth2"

View file

@ -224,6 +224,41 @@ class JWTHandler:
return []
def get_all_jwt_team_ids(self, token: dict) -> List[str]:
"""
Return team IDs from both the plural ``team_ids_jwt_field`` and the
singular ``team_id_jwt_field`` claim (string or list of strings), as a
deduplicated list preserving plural-first order.
Membership-reconciliation paths (SSO callback, JWT-bearer sync) need
to consider both claim shapes. Reading only the plural field as
callers historically did silently dropped users whose IdP populates
the singular field, which is what Okta and Auth0 default to when a
user has a single primary team.
This intentionally does NOT consult ``team_id_default``: that fallback
is a property of how the JWT-bearer auth flow resolves a single
request-bound team, not of the token's claims. Callers that want the
default-team behavior should still go through ``get_team_id``.
"""
team_ids: List[str] = list(self.get_team_ids_from_jwt(token))
if self.litellm_jwtauth.team_id_jwt_field is not None:
singular = get_nested_value(
data=token,
key_path=self.litellm_jwtauth.team_id_jwt_field,
default=None,
)
if isinstance(singular, list):
for item in singular:
if item is None:
continue
sid = str(item)
if sid and sid not in team_ids:
team_ids.append(sid)
elif singular and str(singular) not in team_ids:
team_ids.append(str(singular))
return team_ids
def get_end_user_id(
self, token: dict, default_value: Optional[str]
) -> Optional[str]:

View file

@ -50,6 +50,7 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES = frozenset(
KeyManagementRoutes.KEY_BLOCK.value,
KeyManagementRoutes.KEY_UNBLOCK.value,
KeyManagementRoutes.KEY_BULK_UPDATE.value,
KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value,
]
)
@ -671,6 +672,7 @@ class RouteChecks:
"/key/service-account/generate",
"/key/block",
"/key/unblock",
"/team/key/bulk_update",
]
)

View file

@ -8,6 +8,7 @@ Returns a UserAPIKeyAuth object if the API key is valid
"""
import asyncio
import fnmatch
import re
import secrets
from datetime import datetime, timezone
@ -183,22 +184,54 @@ def _get_bearer_token_or_received_api_key(api_key: str) -> str:
def _routing_selector_matches_claim(
selector_value: Optional[Any], claim_value: Optional[Any]
selector_value: Optional[Any],
claim_value: Optional[Any],
*,
split_space_delimited: bool = False,
) -> bool:
if selector_value is None:
return True
selector_list = (
selector_list: List[str] = (
[str(v) for v in selector_value]
if isinstance(selector_value, list)
else [str(selector_value)]
)
if claim_value is None:
return False
if isinstance(claim_value, list):
claim_list = [str(v) for v in claim_value]
return any(v in claim_list for v in selector_list)
elif (
split_space_delimited
and isinstance(claim_value, str)
and " " in claim_value.strip()
):
# OAuth/OIDC often sends scope as a single space-delimited string. Only split
# for the scope selector: iss/aud/client_id must stay exact full-string match
# on unverified claims (see routing override security review). The elif guard
# (`" " in claim_value.strip()`) ensures at least two non-empty tokens survive.
claim_list = [v for v in claim_value.strip().split(" ") if v]
else:
claim_list = [str(claim_value)]
return str(claim_value) in selector_list if claim_value is not None else False
def _selector_matches_claim(selector: str, claim: str) -> bool:
# NOTE: wildcard matching is case-sensitive (fnmatch.fnmatchcase).
if "*" in selector or "?" in selector:
# Without scope splitting, do not let `*` span whitespace: a malformed
# iss like "trusted.example.com evil.com" must not match "trusted.*".
# Scope uses split_space_delimited so each claim token is checked separately.
if not split_space_delimited and any(ch.isspace() for ch in claim):
return False
return fnmatch.fnmatchcase(claim, selector)
return selector == claim
return any(
_selector_matches_claim(selector=s, claim=c)
for s in selector_list
for c in claim_list
)
def _matches_routing_override(
@ -209,6 +242,11 @@ def _matches_routing_override(
and _routing_selector_matches_claim(
override.client_id, token_claims.get("client_id")
)
and _routing_selector_matches_claim(
override.scope,
token_claims.get("scope"),
split_space_delimited=True,
)
and _routing_selector_matches_claim(override.aud, token_claims.get("aud"))
)

View file

@ -2,7 +2,7 @@ import asyncio
import json
import time
from datetime import datetime, timezone
from typing import Any, List, Literal, Optional, Union
from typing import Any, Callable, List, Literal, Optional, Union
import litellm
from litellm._logging import verbose_proxy_logger
@ -83,93 +83,139 @@ class ResetBudgetJob:
"Failed to reset spend counter %s: %s", counter_key, e
)
@staticmethod
async def _invalidate_user_api_key_cache_entry(cache_key: str) -> None:
"""Drop a stale management-cache entry so the next read fetches from DB.
Some entity types (notably tags and end-users) are not handled by
SpendCounterReseed.from_db, so when a spend counter expires the
budget check falls back to ``cached_obj.spend``. If that cached
object lingers in ``user_api_key_cache`` past a budget reset, the
stale ``.spend`` keeps the entity blocked indefinitely. Deleting
the cache entry forces the next auth-time fetch to reload the
zeroed row from Postgres.
"""
try:
from litellm.proxy.proxy_server import user_api_key_cache
await user_api_key_cache.async_delete_cache(key=cache_key)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to invalidate user_api_key_cache entry %s: %s",
cache_key,
e,
)
async def _cascade_reset_spend_for_budget_link(
self,
budgets_to_reset: List[LiteLLM_BudgetTableFull],
table: Any,
counter_key_fn: Callable[[Any], str],
log_subject: str,
extra_where: Optional[dict] = None,
cache_key_fn: Optional[Callable[[Any], str]] = None,
):
"""
Generic cascade: zero spend on rows whose budget_id is in the reset set.
``cache_key_fn`` is optional: when provided, after the DB update each
matching row's entry in ``user_api_key_cache`` is also dropped. This
is required for entities whose spend counter is read with the cached
object's ``.spend`` as fallback (tags, end-users) — otherwise the
stale cached object pins enforcement to the pre-reset spend until
its TTL expires.
"""
budget_ids = [b.budget_id for b in budgets_to_reset if b.budget_id is not None]
if not budget_ids:
return
where: dict = {"budget_id": {"in": budget_ids}}
if extra_where:
where.update(extra_where)
try:
rows = await table.find_many(where=where)
except Exception as e:
rows = []
verbose_proxy_logger.warning(
"Failed to fetch %s for counter invalidation: %s", log_subject, e
)
update_result = await table.update_many(where=where, data={"spend": 0})
for row in rows:
await self._invalidate_spend_counter(counter_key_fn(row))
if cache_key_fn is not None:
await self._invalidate_user_api_key_cache_entry(cache_key_fn(row))
return update_result
async def reset_budget_for_litellm_team_members(
self, budgets_to_reset: List[LiteLLM_BudgetTableFull]
):
"""
Resets the budget for all LiteLLM Team Members if their budget has expired
"""
budget_ids = [
budget.budget_id
for budget in budgets_to_reset
if budget.budget_id is not None
]
try:
memberships = await self.prisma_client.db.litellm_teammembership.find_many(
where={"budget_id": {"in": budget_ids}}
)
except Exception as e:
memberships = []
verbose_proxy_logger.warning(
"Failed to fetch team memberships for counter invalidation: %s", e
)
update_result = await self.prisma_client.db.litellm_teammembership.update_many(
where={"budget_id": {"in": budget_ids}},
data={
"spend": 0,
},
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=self.prisma_client.db.litellm_teammembership,
counter_key_fn=lambda m: f"spend:team_member:{m.user_id}:{m.team_id}",
log_subject="team memberships",
)
for m in memberships:
await self._invalidate_spend_counter(
f"spend:team_member:{m.user_id}:{m.team_id}"
)
return update_result
async def reset_budget_for_keys_linked_to_budgets(
self, budgets_to_reset: List[LiteLLM_BudgetTableFull]
):
"""
Resets the spend for keys linked to budget tiers that are being reset.
This handles keys that have budget_id but no budget_duration set on the key
itself. Keys with budget_id rely on their linked budget tier's reset schedule
rather than having their own budget_duration.
Keys that have their own budget_duration are already handled by
reset_budget_for_litellm_keys() and are excluded here to avoid
double-resetting.
Excludes keys with their own budget_duration; those are reset by
reset_budget_for_litellm_keys() to avoid double-resetting.
"""
budget_ids = [
budget.budget_id
for budget in budgets_to_reset
if budget.budget_id is not None
]
if not budget_ids:
return
where_clause: dict = {
"budget_id": {"in": budget_ids},
"budget_duration": None, # only keys without their own reset schedule
"spend": {"gt": 0}, # only reset keys that have accumulated spend
}
try:
keys = await self.prisma_client.db.litellm_verificationtoken.find_many(
where=where_clause
)
except Exception as e:
keys = []
verbose_proxy_logger.warning(
"Failed to fetch keys for counter invalidation: %s", e
)
update_result = (
await self.prisma_client.db.litellm_verificationtoken.update_many(
where=where_clause,
data={
"spend": 0,
},
)
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=self.prisma_client.db.litellm_verificationtoken,
counter_key_fn=lambda k: f"spend:key:{k.token}",
log_subject="keys",
extra_where={"budget_duration": None, "spend": {"gt": 0}},
)
for k in keys:
await self._invalidate_spend_counter(f"spend:key:{k.token}")
async def reset_budget_for_orgs_linked_to_budgets(
self, budgets_to_reset: List[LiteLLM_BudgetTableFull]
):
"""
Resets the spend for orgs linked to budget tiers that are being reset.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=self.prisma_client.db.litellm_organizationtable,
counter_key_fn=lambda o: f"spend:org:{o.organization_id}",
log_subject="orgs",
extra_where={"spend": {"gt": 0}},
)
return update_result
async def reset_budget_for_tags_linked_to_budgets(
self, budgets_to_reset: List[LiteLLM_BudgetTableFull]
):
"""
Resets the spend for tags linked to budget tiers that are being reset.
Also drops each tag's ``user_api_key_cache`` entry so the next
``_tag_max_budget_check`` reloads the zeroed row from the DB.
``SpendCounterReseed.from_db`` intentionally returns ``None`` for
tags, so the budget check falls back to the cached
``LiteLLM_TagTable.spend`` once the spend counter expires; without
this invalidation, that stale ``.spend`` keeps the tag over-budget
indefinitely.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=self.prisma_client.db.litellm_tagtable,
counter_key_fn=lambda t: f"spend:tag:{t.tag_name}",
log_subject="tags",
extra_where={"spend": {"gt": 0}},
cache_key_fn=lambda t: f"tag:{t.tag_name}",
)
async def reset_budget_for_litellm_budget_table(self):
"""
@ -237,6 +283,14 @@ class ResetBudgetJob:
budgets_to_reset=budgets_to_reset
)
await self.reset_budget_for_orgs_linked_to_budgets(
budgets_to_reset=budgets_to_reset
)
await self.reset_budget_for_tags_linked_to_budgets(
budgets_to_reset=budgets_to_reset
)
if endusers_to_reset is not None and len(endusers_to_reset) > 0:
for enduser in endusers_to_reset:
try:

View file

@ -1131,10 +1131,12 @@ class DBSpendUpdateWriter:
timeout=timedelta(seconds=60)
) as transaction:
async with transaction.batch_() as batcher:
for (
user_id,
response_cost,
) in user_list_transactions.items():
# Sort by ID for consistent lock ordering across pods to prevent deadlocks.
# batch_() issues statements sequentially within the tx, so iteration
# order = lock acquisition order.
for user_id, response_cost in sorted(
user_list_transactions.items()
):
batcher.litellm_usertable.update_many(
where={"user_id": user_id},
data={"spend": {"increment": response_cost}},
@ -1186,10 +1188,10 @@ class DBSpendUpdateWriter:
timeout=timedelta(seconds=60)
) as transaction:
async with transaction.batch_() as batcher:
for (
token,
response_cost,
) in key_list_transactions.items():
# Sort by token for consistent lock ordering across pods to prevent deadlocks.
for token, response_cost in sorted(
key_list_transactions.items()
):
batcher.litellm_verificationtoken.update_many( # 'update_many' prevents error from being raised if no row exists
where={"token": token},
data={
@ -1230,10 +1232,10 @@ class DBSpendUpdateWriter:
timeout=timedelta(seconds=60)
) as transaction:
async with transaction.batch_() as batcher:
for (
team_id,
response_cost,
) in team_list_transactions.items():
# Sort by team_id for consistent lock ordering across pods to prevent deadlocks.
for team_id, response_cost in sorted(
team_list_transactions.items()
):
verbose_proxy_logger.debug(
"Updating spend for team id={} by {}".format(
team_id, response_cost
@ -1288,10 +1290,11 @@ class DBSpendUpdateWriter:
timeout=timedelta(seconds=60)
) as transaction:
async with transaction.batch_() as batcher:
for (
key,
response_cost,
) in team_member_list_transactions.items():
# Sort by composite key for consistent lock ordering across pods to prevent deadlocks.
# Key format "team_id::<v>::user_id::<v>" makes the string sort equivalent to sorting by (team_id, user_id).
for key, response_cost in sorted(
team_member_list_transactions.items()
):
# key is "team_id::<value>::user_id::<value>"
team_id = key.split("::")[1]
user_id = key.split("::")[3]
@ -1348,10 +1351,10 @@ class DBSpendUpdateWriter:
timeout=timedelta(seconds=60)
) as transaction:
async with transaction.batch_() as batcher:
for (
org_id,
response_cost,
) in org_list_transactions.items():
# Sort by org_id for consistent lock ordering across pods to prevent deadlocks.
for org_id, response_cost in sorted(
org_list_transactions.items()
):
batcher.litellm_organizationtable.update_many( # 'update_many' prevents error from being raised if no row exists
where={"organization_id": org_id},
data={"spend": {"increment": response_cost}},
@ -1439,7 +1442,10 @@ class DBSpendUpdateWriter:
timeout=timedelta(seconds=60)
) as transaction:
async with transaction.batch_() as batcher:
for entity_id, response_cost in transactions.items():
# Sort by entity_id for consistent lock ordering across pods to prevent deadlocks.
for entity_id, response_cost in sorted(
transactions.items()
):
verbose_proxy_logger.debug(
f"Updating spend for {entity_name} {where_field}={entity_id} by {response_cost}"
)

View file

@ -242,6 +242,10 @@ async def list_guardrails_v2(
gid = guardrail.get("guardrail_id")
if gid in seen_guardrail_ids:
continue
# Skip stale DB-backed entries — the DB row was deleted (likely by
# another pod) and reconciliation hasn't fired yet on this pod.
if gid is not None and IN_MEMORY_GUARDRAIL_HANDLER.get_source(gid) == "db":
continue
if not is_admin:
g_team_id = guardrail.get("team_id")
if g_team_id is not None and g_team_id not in caller_team_ids:
@ -360,7 +364,7 @@ async def create_guardrail(
try:
IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
guardrail=cast(Guardrail, result)
guardrail=cast(Guardrail, result), source="db"
)
verbose_proxy_logger.info(
f"Immediate sync: Successfully initialized guardrail '{guardrail_name}' (ID: {guardrail_id})"
@ -1017,7 +1021,7 @@ async def approve_guardrail_submission(
}
try:
IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
guardrail=cast(Guardrail, guardrail_dict)
guardrail=cast(Guardrail, guardrail_dict), source="db"
)
verbose_proxy_logger.info(
"Approved guardrail %s (ID: %s) and initialized in memory",
@ -1295,10 +1299,18 @@ async def get_guardrail_info(guardrail_id: str):
guardrail_id=guardrail_id, prisma_client=prisma_client
)
if result is None:
result = IN_MEMORY_GUARDRAIL_HANDLER.get_guardrail_by_id(
in_memory = IN_MEMORY_GUARDRAIL_HANDLER.get_guardrail_by_id(
guardrail_id=guardrail_id
)
guardrail_definition_location = GUARDRAIL_DEFINITION_LOCATION.CONFIG
# Only return config-loaded entries here. A DB-backed entry that's
# missing from the DB is stale (deleted on another pod, awaiting
# reconciliation on this one) and must surface as 404.
if (
in_memory is not None
and IN_MEMORY_GUARDRAIL_HANDLER.get_source(guardrail_id) == "config"
):
result = in_memory
guardrail_definition_location = GUARDRAIL_DEFINITION_LOCATION.CONFIG
if result is None:
raise HTTPException(

View file

@ -3,7 +3,7 @@
import importlib
import os
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional, Type, cast
from typing import Any, Dict, List, Literal, Optional, Set, Type, cast
import litellm
from litellm import Router
@ -403,11 +403,19 @@ class InMemoryGuardrailHandler:
Guardrail id to CustomGuardrail object mapping
"""
self._sources: Dict[str, Literal["db", "config"]] = {}
"""
Guardrail id to provenance marker. "db" entries are reconciled against
the DB on each polling tick; "config" entries are owned by proxy_config.yaml
and never deleted by reconciliation.
"""
def initialize_guardrail(
self,
guardrail: Guardrail,
config_file_path: Optional[str] = None,
llm_router: Optional["Router"] = None,
source: Literal["db", "config"] = "config",
) -> Optional[Guardrail]:
"""
Initialize a guardrail from a dictionary and add it to the litellm callback manager
@ -420,6 +428,10 @@ class InMemoryGuardrailHandler:
verbose_proxy_logger.debug(
"guardrail_id already exists in IN_MEMORY_GUARDRAILS"
)
# Honor the caller's source even on the early-return path so a
# racing polling tick or a hot-reload of config can correct an
# entry's provenance.
self._sources[guardrail_id] = source
return self.IN_MEMORY_GUARDRAILS[guardrail_id]
custom_guardrail_callback: Optional[CustomGuardrail] = None
@ -497,6 +509,7 @@ class InMemoryGuardrailHandler:
# store references to the guardrail in memory
self.IN_MEMORY_GUARDRAILS[guardrail_id] = parsed_guardrail
self.guardrail_id_to_custom_guardrail[guardrail_id] = custom_guardrail_callback
self._sources[guardrail_id] = source
return parsed_guardrail
@ -557,7 +570,10 @@ class InMemoryGuardrailHandler:
return _guardrail_callback
def update_in_memory_guardrail(
self, guardrail_id: str, guardrail: Guardrail
self,
guardrail_id: str,
guardrail: Guardrail,
source: Literal["db", "config"] = "db",
) -> None:
"""
Update a guardrail in memory
@ -566,6 +582,7 @@ class InMemoryGuardrailHandler:
- updates the guardrail params in litellm.callback_manager
"""
self.IN_MEMORY_GUARDRAILS[guardrail_id] = guardrail
self._sources[guardrail_id] = source
custom_guardrail_callback = self.guardrail_id_to_custom_guardrail.get(
guardrail_id
@ -584,6 +601,7 @@ class InMemoryGuardrailHandler:
"""
# Remove from in-memory storage
self.IN_MEMORY_GUARDRAILS.pop(guardrail_id, None)
self._sources.pop(guardrail_id, None)
# Remove the callback from litellm.callbacks
custom_guardrail_callback = self.guardrail_id_to_custom_guardrail.pop(
@ -608,6 +626,34 @@ class InMemoryGuardrailHandler:
"""
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]:
"""
Return the provenance of an in-memory guardrail.
"""
return self._sources.get(guardrail_id)
def reconcile_db_guardrails(self, db_guardrail_ids: Set[str]) -> List[str]:
"""
Drop in-memory entries that originated from the DB but are no longer
present in db_guardrail_ids. Config-loaded guardrails are never touched.
Called by the periodic DB polling tick so that a guardrail deleted
on another pod is eventually purged from this pod's memory + callbacks.
"""
stale_ids = [
guardrail_id
for guardrail_id, source in self._sources.items()
if source == "db" and guardrail_id not in db_guardrail_ids
]
for guardrail_id in stale_ids:
verbose_proxy_logger.info(
"Reconcile: removing stale DB-backed guardrail '%s' from memory "
"(deleted in DB by another pod)",
guardrail_id,
)
self.delete_in_memory_guardrail(guardrail_id)
return stale_ids
def _has_guardrail_params_changed(
self, guardrail_id: str, new_guardrail: Guardrail
) -> bool:
@ -661,7 +707,10 @@ class InMemoryGuardrailHandler:
return len(changed_fields) > 0
def reinitialize_guardrail(
self, guardrail: Guardrail, config_file_path: Optional[str] = None
self,
guardrail: Guardrail,
config_file_path: Optional[str] = None,
source: Literal["db", "config"] = "config",
) -> Optional[Guardrail]:
"""
Force re-initialization of a guardrail even if it exists in memory.
@ -680,7 +729,7 @@ class InMemoryGuardrailHandler:
# Initialize fresh (will add new callback to litellm.callbacks)
return self.initialize_guardrail(
guardrail=guardrail, config_file_path=config_file_path
guardrail=guardrail, config_file_path=config_file_path, source=source
)
def sync_guardrail_from_db(
@ -701,9 +750,15 @@ class InMemoryGuardrailHandler:
f"Guardrail '{guardrail_name}' (ID: {guardrail_id}) params changed, re-initializing..."
)
return self.reinitialize_guardrail(
guardrail=guardrail, config_file_path=config_file_path
guardrail=guardrail,
config_file_path=config_file_path,
source="db",
)
# Params unchanged but the entry is still DB-backed; make sure the
# source marker reflects that even if it was previously set differently
# (e.g. a config entry whose UUID later collided with a DB row).
self._sources[guardrail_id] = "db"
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)

View file

@ -30,6 +30,7 @@ def init_guardrails_v2(
guardrail=cast(Guardrail, guardrail),
config_file_path=config_file_path,
llm_router=llm_router,
source="config",
)
if initialized_guardrail:
guardrail_list.append(initialized_guardrail)

View file

@ -253,27 +253,63 @@ class SharedHealthCheckManager:
# Always release the lock
await self.release_health_check_lock()
else:
# Lock not acquired, wait briefly and try to get cached results
# If Redis is not configured, skip polling — there is no cache
# to wait for.
if self.redis_cache is None:
return await perform_health_check(
model_list=model_list,
details=details,
max_concurrency=max_concurrency,
)
# Lock not acquired — poll for cached results until the lock
# holder finishes or the lock expires, rather than falling back
# to a redundant local health check after only 2 seconds.
verbose_proxy_logger.debug(
"Pod %s waiting for other pod to complete health check", self.pod_id
)
# Wait a bit for the other pod to complete
await asyncio.sleep(2)
poll_interval = 5 # seconds between cache checks
max_wait = self.lock_ttl # wait at most as long as the lock can live
elapsed = 0
# Try to get cached results again
cached_results = await self.get_cached_health_check_results()
if cached_results is not None:
return (
cached_results.get("healthy_endpoints", []),
cached_results.get("unhealthy_endpoints", []),
{},
)
while elapsed < max_wait:
await asyncio.sleep(poll_interval)
elapsed += poll_interval
# Still no cache, fall back to local health check
cached_results = await self.get_cached_health_check_results()
if cached_results is not None:
verbose_proxy_logger.info(
"Pod %s using cached health check results after waiting %ds",
self.pod_id,
elapsed,
)
return (
cached_results.get("healthy_endpoints", []),
cached_results.get("unhealthy_endpoints", []),
{},
)
# Check if the lock is still held — if it was released without
# caching (e.g. the holder crashed), stop waiting early.
try:
lock_key = self.get_health_check_lock_key()
current_owner = await self.redis_cache.async_get_cache(lock_key)
if current_owner is None:
verbose_proxy_logger.debug(
"Pod %s detected lock released without cache, stopping wait",
self.pod_id,
)
break
except Exception:
# Redis hiccup — continue polling rather than crashing out
pass
# Exhausted wait — fall back to local health check
verbose_proxy_logger.warning(
"Pod %s falling back to local health check (no cache available)",
"Pod %s falling back to local health check after waiting %ds (no cache available)",
self.pod_id,
elapsed,
)
return await perform_health_check(

View file

@ -1742,29 +1742,54 @@ async def test_model_connection(
# Look up model configuration from router if model name is provided
# This gets the litellm_params from proxy config (with resolved env vars)
config_litellm_params: dict = {}
if model_name and llm_router is not None:
if llm_router is not None:
# Prefer disambiguation by deployment id (`model_info.id`) when
# the caller supplies it. This is required when multiple
# deployments share a `model_name` (e.g. wildcard `openai/*`
# with multiple `api_base` values for failover): the UI's
# "Test Connection" button targets a specific row, and that
# row's id is the only thing that uniquely identifies which
# deployment to probe. Without this, all duplicates collapse
# onto `deployments[0]`.
request_model_info = model_info or {}
request_model_id = request_model_info.get("id")
try:
# First try to find by proxy model_name (e.g., "gpt-4o")
deployments = llm_router.get_model_list(model_name=model_name)
# If not found, try to find by litellm model name (e.g., "azure/gpt-4o")
if not deployments or len(deployments) == 0:
all_deployments = llm_router.get_model_list(model_name=None)
if all_deployments:
for deployment in all_deployments:
if (
deployment.get("litellm_params", {}).get("model")
== model_name
):
deployments = [deployment]
break
if deployments and len(deployments) > 0:
# Use the first deployment's litellm_params as base config
# These already have resolved environment variables from proxy config
config_litellm_params = dict(
deployments[0].get("litellm_params", {})
deployment_by_id = None
if request_model_id:
deployment_by_id = llm_router.get_deployment(
model_id=request_model_id
)
if deployment_by_id is not None:
config_litellm_params = deployment_by_id.litellm_params.model_dump(
exclude_none=True
)
elif model_name:
# Fall back to model_name lookup for callers (e.g. the
# "Add Model" wizard, or curl) that don't supply an id.
# First try to find by proxy model_name (e.g., "gpt-4o")
deployments = llm_router.get_model_list(model_name=model_name)
# If not found, try to find by litellm model name
# (e.g., "azure/gpt-4o")
if not deployments or len(deployments) == 0:
all_deployments = llm_router.get_model_list(model_name=None)
if all_deployments:
for deployment in all_deployments:
if (
deployment.get("litellm_params", {}).get("model")
== model_name
):
deployments = [deployment]
break
if deployments and len(deployments) > 0:
# Use the first deployment's litellm_params as base
# config. These already have resolved environment
# variables from proxy config.
config_litellm_params = dict(
deployments[0].get("litellm_params", {})
)
except Exception as e:
verbose_proxy_logger.debug(
f"Could not find model {model_name} in router: {e}. "

View file

@ -1,5 +1,6 @@
import asyncio
import copy
import json
import re
import time
from collections import OrderedDict
@ -794,8 +795,17 @@ class LiteLLMProxyRequestSetup:
)
)
for k, v in litellm_logging_metadata_headers.items():
if v is not None:
if v is None:
continue
# httpx requires header values to be str or bytes; coerce numbers/bools
# to str and JSON-encode dict/list (e.g. user_api_key_spend is float,
# user_api_key_auth_metadata is dict). See #27458.
if isinstance(v, (dict, list)):
returned_headers["x-litellm-{}".format(k)] = json.dumps(v)
elif isinstance(v, (str, bytes)):
returned_headers["x-litellm-{}".format(k)] = v
else:
returned_headers["x-litellm-{}".format(k)] = str(v)
return returned_headers
@ -1731,6 +1741,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
data=data,
user_api_key_dict=user_api_key_dict,
pre_alias_model_name=_pre_alias_model,
llm_router=llm_router,
)
## ENFORCED PARAMS CHECK
@ -1864,6 +1875,7 @@ def _apply_credential_overrides_from_model_config(
data: dict,
user_api_key_dict: UserAPIKeyAuth,
pre_alias_model_name: Optional[str] = None,
llm_router: Optional[Router] = None,
) -> None:
"""
Walk the model_config precedence chain in team/project metadata.
@ -1899,10 +1911,19 @@ def _apply_credential_overrides_from_model_config(
if not project_model_config and not team_model_config:
return
# Extract provider hint from model name (e.g. "azure/gpt-4" -> "azure")
# Extract provider hint from model name (e.g. "azure/gpt-4" -> "azure").
# When the user-facing name has no provider prefix, fall back to the
# deployment's litellm_params so multi-provider defaultconfig entries
# don't silently match the first dict key (#27516).
provider: Optional[str] = None
if "/" in model_name:
provider = model_name.split("/", 1)[0]
elif llm_router is not None:
provider = _resolve_provider_from_deployment(
llm_router=llm_router,
model_name=model_name,
pre_alias_model_name=pre_alias_model_name,
)
credential_name = _resolve_credential_from_model_config(
model_name=model_name,
@ -1938,6 +1959,48 @@ def _apply_credential_overrides_from_model_config(
)
def _resolve_provider_from_deployment(
llm_router: Router,
model_name: str,
pre_alias_model_name: Optional[str] = None,
) -> Optional[str]:
"""
Resolve a provider hint from the deployment's litellm_params when the
user-facing model name has no provider prefix.
Tries the post-alias name first (the resolved model group), then the
pre-alias name. Returns None if no deployment is found or the deployment
has no usable provider info.
"""
candidates = [model_name]
if pre_alias_model_name and pre_alias_model_name != model_name:
candidates.append(pre_alias_model_name)
for name in candidates:
try:
deployment = llm_router.get_deployment_by_model_group_name(
model_group_name=name
)
except Exception:
deployment = None
if deployment is None:
continue
litellm_params = getattr(deployment, "litellm_params", None)
if litellm_params is None:
continue
custom_provider = getattr(litellm_params, "custom_llm_provider", None)
if custom_provider:
return custom_provider
deployment_model = getattr(litellm_params, "model", "") or ""
if "/" in deployment_model:
return deployment_model.split("/", 1)[0]
return None
def _resolve_credential_from_model_config(
model_name: str,
project_model_config: Optional[dict],

View file

@ -88,8 +88,8 @@ from litellm.router import Router
from litellm.secret_managers.main import get_secret
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequest,
BulkUpdateKeyRequestItem,
BulkUpdateKeyResponse,
BulkUpdateTeamKeysRequest,
FailedKeyUpdate,
SuccessfulKeyUpdate,
)
@ -1881,7 +1881,7 @@ async def _get_and_validate_existing_key(
async def _process_single_key_update(
key_update_item: BulkUpdateKeyRequestItem,
update_key_request: UpdateKeyRequest,
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: Optional[str],
prisma_client: Optional[PrismaClient],
@ -1889,6 +1889,7 @@ async def _process_single_key_update(
proxy_logging_obj: Any,
llm_router: Optional[Router],
user_custom_key_update: Optional[Callable] = None,
existing_key_row: Optional[LiteLLM_VerificationToken] = None,
) -> Dict[str, Any]:
"""
Process a single key update with all validations and checks.
@ -1897,13 +1898,14 @@ async def _process_single_key_update(
including validation, permission checks, team checks, and database updates.
Args:
key_update_item: The key update request item
update_key_request: Fully-constructed UpdateKeyRequest for the target key
user_api_key_dict: The authenticated user's API key info
litellm_changed_by: Optional header for tracking who made the change
prisma_client: Prisma client instance
user_api_key_cache: User API key cache
proxy_logging_obj: Proxy logging object
llm_router: LLM router instance
existing_key_row: Optional pre-fetched key row to avoid redundant lookups
Returns:
Dict containing the updated key information
@ -1912,13 +1914,14 @@ async def _process_single_key_update(
HTTPException: For various validation and permission errors
"""
# Validate max_budget
_validate_max_budget(key_update_item.max_budget)
_validate_max_budget(update_key_request.max_budget)
# Get and validate existing key
existing_key_row = await _get_and_validate_existing_key(
token=key_update_item.key,
prisma_client=prisma_client,
)
if existing_key_row is None:
existing_key_row = await _get_and_validate_existing_key(
token=update_key_request.key,
prisma_client=prisma_client,
)
# Check team member permissions
if prisma_client is not None:
@ -1930,15 +1933,6 @@ async def _process_single_key_update(
user_api_key_cache=user_api_key_cache,
)
# Create UpdateKeyRequest from BulkUpdateKeyRequestItem
update_key_request = UpdateKeyRequest(
key=key_update_item.key,
budget_id=key_update_item.budget_id,
max_budget=key_update_item.max_budget,
team_id=key_update_item.team_id,
tags=key_update_item.tags,
)
# Custom key update hook
if user_custom_key_update is not None:
if inspect.iscoroutinefunction(user_custom_key_update):
@ -2003,12 +1997,12 @@ async def _process_single_key_update(
detail={"error": "Database not connected"},
)
_data = {**non_default_values, "token": key_update_item.key}
response = await prisma_client.update_data(token=key_update_item.key, data=_data)
_data = {**non_default_values, "token": update_key_request.key}
response = await prisma_client.update_data(token=update_key_request.key, data=_data)
# Delete cache
await _delete_cache_key_object(
hashed_token=_hash_token_if_needed(key_update_item.key),
hashed_token=_hash_token_if_needed(update_key_request.key),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
@ -2598,9 +2592,15 @@ async def bulk_update_keys(
for key_update_item in data.keys:
try:
# Process single key update using reusable function
update_key_request = UpdateKeyRequest(
key=key_update_item.key,
budget_id=key_update_item.budget_id,
max_budget=key_update_item.max_budget,
team_id=key_update_item.team_id,
tags=key_update_item.tags,
)
updated_key_info = await _process_single_key_update(
key_update_item=key_update_item,
update_key_request=update_key_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
prisma_client=prisma_client,
@ -2665,6 +2665,223 @@ async def bulk_update_keys(
)
def _build_failed_team_key_update(
token: str,
exception: Exception,
existing_key_row: Optional[LiteLLM_VerificationToken],
) -> FailedKeyUpdate:
"""Normalize an exception from the per-key update loop into a FailedKeyUpdate."""
if isinstance(exception, HTTPException):
detail = exception.detail
if isinstance(detail, dict):
error_message = detail.get("error", str(exception))
else:
error_message = str(detail)
elif isinstance(exception, ProxyException):
error_message = exception.message
else:
error_message = str(exception)
key_info: Optional[Dict[str, Any]] = None
if existing_key_row is not None:
if hasattr(existing_key_row, "model_dump"):
key_info = existing_key_row.model_dump()
elif hasattr(existing_key_row, "dict"):
key_info = existing_key_row.dict()
if key_info:
key_info.pop("token", None)
return FailedKeyUpdate(key=token, key_info=key_info, failed_reason=error_message)
@router.post(
"/team/key/bulk_update",
tags=["key management"],
dependencies=[Depends(user_api_key_auth)],
response_model=BulkUpdateKeyResponse,
)
@management_endpoint_wrapper
async def bulk_update_team_keys(
data: BulkUpdateTeamKeysRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: Optional[str] = Header(
None,
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
),
):
"""
Apply one update payload to many keys inside a single team.
Pass `team_id` plus either `key_ids` or `all_keys_in_team=True`. The
`update_fields` payload is broadcast to every selected key. Per-key
failures are returned in `failed_updates` rather than aborting the batch.
Callable by proxy admins, or by team admins with `KEY_UPDATE` permission.
"""
from litellm.proxy.proxy_server import (
llm_router,
prisma_client,
proxy_logging_obj,
user_api_key_cache,
user_custom_key_update,
)
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={"error": "Database not connected"},
)
if not data.team_id:
raise HTTPException(
status_code=400,
detail={"error": "team_id is required"},
)
MAX_BATCH_SIZE = 500
if data.key_ids is not None and len(data.key_ids) > MAX_BATCH_SIZE:
raise HTTPException(
status_code=400,
detail={
"error": f"Maximum {MAX_BATCH_SIZE} keys can be updated at once. Found {len(data.key_ids)} key_ids."
},
)
if data.all_keys_in_team:
# "all" excludes blocked/expired — bulk refresh shouldn't revive a key an admin disabled.
# `blocked` is Boolean? with no default; `/key/generate` writes NULL. Prisma's `NOT`
# excludes NULLs, so explicitly OR `false` with `null` to include them.
now = datetime.now(timezone.utc)
existing_keys = await prisma_client.db.litellm_verificationtoken.find_many(
where={
"team_id": data.team_id,
"AND": [
{"OR": [{"blocked": False}, {"blocked": None}]},
{"OR": [{"expires": None}, {"expires": {"gt": now}}]},
],
},
order={"token": "asc"},
take=MAX_BATCH_SIZE + 1,
)
if len(existing_keys) > MAX_BATCH_SIZE:
raise HTTPException(
status_code=400,
detail={
"error": f"Team {data.team_id} has more than {MAX_BATCH_SIZE} keys. Use `key_ids` to update in batches of {MAX_BATCH_SIZE}."
},
)
requested_tokens = [row.token for row in existing_keys]
else:
if data.key_ids is None or len(data.key_ids) == 0:
raise HTTPException(
status_code=400,
detail={
"error": "key_ids must be provided when all_keys_in_team is False"
},
)
# Dedupe by hashed form — duplicates collapse to one update.
requested_tokens = []
hashed_key_ids = []
seen_hashes = set()
for k in data.key_ids:
h = _hash_token_if_needed(k)
if h in seen_hashes:
continue
seen_hashes.add(h)
requested_tokens.append(k)
hashed_key_ids.append(h)
existing_keys = await prisma_client.db.litellm_verificationtoken.find_many(
where={"team_id": data.team_id, "token": {"in": hashed_key_ids}}
)
# Anchor membership check on data.team_id (not existing_keys[0]); empty result must still gate non-admins.
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
auth_anchor = (
existing_keys[0]
if existing_keys
else LiteLLM_VerificationToken(
token="__team_scope_auth_check__",
team_id=data.team_id,
models=[],
)
)
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
user_api_key_dict=user_api_key_dict,
route=KeyManagementRoutes.KEY_UPDATE,
prisma_client=prisma_client,
existing_key_row=auth_anchor,
user_api_key_cache=user_api_key_cache,
)
# Block metadata.allowed_passthrough_routes for non-admins — the runtime
# route checker reads it from key/team metadata to grant passthrough.
_check_passthrough_routes_caller_permission(
data=data.update_fields, user_api_key_dict=user_api_key_dict
)
if not requested_tokens:
raise HTTPException(
status_code=404,
detail={"error": f"No keys found for team {data.team_id}"},
)
existing_by_token = {row.token: row for row in existing_keys}
update_field_dict = data.update_fields.model_dump(exclude_unset=True)
successful_updates: List[SuccessfulKeyUpdate] = []
failed_updates: List[FailedKeyUpdate] = []
for token in requested_tokens:
db_token = _hash_token_if_needed(token)
try:
if db_token not in existing_by_token:
raise HTTPException(
status_code=404,
detail={"error": f"Key not found in team {data.team_id}"},
)
# team_id from validated scope, never user payload — drives _check_team_key_limits.
update_key_request = UpdateKeyRequest(
key=token,
team_id=data.team_id,
**update_field_dict,
)
updated_key_info = await _process_single_key_update(
update_key_request=update_key_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
user_custom_key_update=user_custom_key_update,
existing_key_row=existing_by_token[db_token],
)
successful_updates.append(
SuccessfulKeyUpdate(key=token, key_info=updated_key_info)
)
except Exception as e:
# Log the hashed prefix — `token` may be a raw sk-... and ERROR logs persist.
verbose_proxy_logger.exception(
f"Failed to update key {db_token[:12]}... in team {data.team_id}: {e}"
)
failed_updates.append(
_build_failed_team_key_update(
token=token,
exception=e,
existing_key_row=existing_by_token.get(db_token),
)
)
return BulkUpdateKeyResponse(
total_requested=len(requested_tokens),
successful_updates=successful_updates,
failed_updates=failed_updates,
)
async def validate_key_team_change(
key: LiteLLM_VerificationToken,
team: LiteLLM_TeamTable,

View file

@ -12,12 +12,13 @@ All /tag management endpoints
import asyncio
import json
from typing import TYPE_CHECKING, Dict, List, Optional
from datetime import datetime
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends, HTTPException, Query
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import UserAPIKeyAuth, user_api_key_has_admin_view
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
@ -39,6 +40,72 @@ if TYPE_CHECKING:
router = APIRouter()
async def _get_internal_user_api_keys(
prisma_client,
user_api_key_dict: UserAPIKeyAuth,
) -> List[str]:
user_role = user_api_key_dict.user_role
if user_role is None or not user_role.is_internal_user_role:
return []
user_api_keys = set()
if user_api_key_dict.api_key:
user_api_keys.add(user_api_key_dict.api_key)
user_id = user_api_key_dict.user_id
if user_id is None:
return sorted(user_api_keys)
key_records = await prisma_client.db.litellm_verificationtoken.find_many(
where={"user_id": user_id},
select={"token": True},
)
user_api_keys.update(
key_record.token
for key_record in key_records
if getattr(key_record, "token", None)
)
return sorted(user_api_keys)
async def _get_tag_list_scope(
prisma_client,
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, dict]]:
user_role = user_api_key_dict.user_role
if user_api_key_has_admin_view(user_api_key_dict) or (
user_role is None or not user_role.is_internal_user_role
):
return None
scoped_api_keys = await _get_internal_user_api_keys(
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
)
return {"api_key": {"in": scoped_api_keys}}
async def _get_tag_daily_activity_api_key_filter(
prisma_client,
user_api_key_dict: UserAPIKeyAuth,
requested_api_key: Optional[str],
) -> Optional[Union[str, List[str]]]:
user_role = user_api_key_dict.user_role
if user_api_key_has_admin_view(user_api_key_dict) or (
user_role is None or not user_role.is_internal_user_role
):
return requested_api_key
scoped_api_keys = await _get_internal_user_api_keys(
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
)
if requested_api_key is not None:
return requested_api_key if requested_api_key in scoped_api_keys else []
return scoped_api_keys
async def _get_model_names(prisma_client, model_ids: list) -> Dict[str, str]:
"""Helper function to get model names from model IDs"""
try:
@ -395,6 +462,32 @@ async def info_tag(
raise HTTPException(status_code=500, detail=str(e))
def _validate_tag_list_date_range(
start_date: Optional[str], end_date: Optional[str]
) -> None:
"""Require both dates together, and enforce YYYY-MM-DD format with start <= end."""
if (start_date is None) != (end_date is None):
raise HTTPException(
status_code=400,
detail="start_date and end_date must be provided together",
)
if start_date is None:
return
try:
start = datetime.strptime(start_date, "%Y-%m-%d")
end = datetime.strptime(end_date, "%Y-%m-%d") # type: ignore[arg-type]
except ValueError as e:
raise HTTPException(
status_code=400,
detail=f"Invalid date format, expected YYYY-MM-DD: {e}",
)
if start > end:
raise HTTPException(
status_code=400,
detail="start_date must be on or before end_date",
)
@router.get(
"/tag/list",
tags=["tag management"],
@ -402,6 +495,18 @@ async def info_tag(
)
async def list_tags(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
start_date: Optional[str] = Query(
None,
description=(
"Optional start date (YYYY-MM-DD). When provided together with "
"end_date, dynamic tags are limited to those active in the window. "
"Stored tags are always returned."
),
),
end_date: Optional[str] = Query(
None,
description="Optional end date (YYYY-MM-DD). Must be given with start_date.",
),
):
"""
List all available tags with their budget information.
@ -411,10 +516,44 @@ async def list_tags(
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
_validate_tag_list_date_range(start_date, end_date)
try:
tag_scope = await _get_tag_list_scope(
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
)
## QUERY DYNAMIC TAGS ##
# Use group_by instead of find_many(distinct=["tag"]).
# Prisma's distinct fetches all columns for all rows and deduplicates
# in application code, which is extremely slow on large tables.
# See: https://www.prisma.io/docs/orm/prisma-client/queries/aggregation-grouping-summarizing#distinct-under-the-hood
dynamic_tag_where: Dict[str, Any] = {"tag": {"not": None}}
if tag_scope:
dynamic_tag_where = {**dynamic_tag_where, **tag_scope}
if start_date is not None and end_date is not None:
dynamic_tag_where["date"] = {"gte": start_date, "lte": end_date}
dynamic_tag_rows = await prisma_client.db.litellm_dailytagspend.group_by(
by=["tag"],
where=dynamic_tag_where,
min={"created_at": True},
max={"updated_at": True},
)
used_tag_names = [row["tag"] for row in dynamic_tag_rows if row["tag"]]
if tag_scope is not None and not used_tag_names:
return []
stored_tag_where = (
{"tag_name": {"in": used_tag_names}} if tag_scope is not None else None
)
## QUERY STORED TAGS ##
tag_records = await prisma_client.db.litellm_tagtable.find_many(
include={"litellm_budget_table": True}
where=stored_tag_where,
include={"litellm_budget_table": True},
)
stored_tag_names = set()
@ -448,18 +587,6 @@ async def list_tags(
list_of_tags.append(tag_dict)
## QUERY DYNAMIC TAGS ##
# Use group_by instead of find_many(distinct=["tag"]).
# Prisma's distinct fetches all columns for all rows and deduplicates
# in application code, which is extremely slow on large tables.
# See: https://www.prisma.io/docs/orm/prisma-client/queries/aggregation-grouping-summarizing#distinct-under-the-hood
dynamic_tag_rows = await prisma_client.db.litellm_dailytagspend.group_by(
by=["tag"],
where={"tag": {"not": None}},
min={"created_at": True},
max={"updated_at": True},
)
dynamic_tag_config = [
{
"name": row["tag"],
@ -527,6 +654,7 @@ async def get_tag_daily_activity(
api_key: Optional[str] = None,
page: int = 1,
page_size: int = 10,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get daily activity for specific tags or all tags.
@ -545,8 +673,18 @@ async def get_tag_daily_activity(
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
# Convert comma-separated tags string to list if provided
tag_list = tags.split(",") if tags else None
scoped_api_key_filter = await _get_tag_daily_activity_api_key_filter(
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
requested_api_key=api_key,
)
if scoped_api_key_filter == []:
return SpendAnalyticsPaginatedResponse(results=[])
return await get_daily_activity(
prisma_client=prisma_client,
@ -557,7 +695,7 @@ async def get_tag_daily_activity(
start_date=start_date,
end_date=end_date,
model=model,
api_key=api_key,
api_key=scoped_api_key_filter,
page=page,
page_size=page_size,
# metadata_metrics_func=None because litellm_dailytagspend rows are

View file

@ -740,7 +740,7 @@ def generic_response_convertor(
all_teams = []
if sso_jwt_handler is not None:
team_ids = sso_jwt_handler.get_team_ids_from_jwt(cast(dict, response))
team_ids = sso_jwt_handler.get_all_jwt_team_ids(cast(dict, response))
all_teams.extend(team_ids)
if team_mappings is not None and team_mappings.team_ids_jwt_field is not None:
@ -755,7 +755,7 @@ def generic_response_convertor(
f"Loaded team_ids from DB team_mappings.team_ids_jwt_field='{team_mappings.team_ids_jwt_field}': {team_ids_from_db_mapping}"
)
else:
team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response))
team_ids = jwt_handler.get_all_jwt_team_ids(cast(dict, response))
all_teams.extend(team_ids)
# Determine user role based on role_mappings if available

View file

@ -805,7 +805,8 @@ def run_server( # noqa: PLR0915
)
db_connection_pool_limit = 100
db_connection_timeout = 60
# Starts optional due to config fallback checks; guaranteed non-None before use.
db_connection_timeout: Optional[Union[int, float]] = 60
general_settings = {}
### GET DB TOKEN FOR IAM AUTH ###
@ -914,10 +915,15 @@ def run_server( # noqa: PLR0915
"database_connection_pool_limit",
LiteLLMDatabaseConnectionPool.database_connection_pool_limit.value,
)
db_connection_timeout = general_settings.get(
"database_connection_pool_timeout",
LiteLLMDatabaseConnectionPool.database_connection_pool_timeout.value,
)
db_connection_timeout = general_settings.get("database_connection_timeout")
if db_connection_timeout is None:
db_connection_timeout = general_settings.get(
"database_connection_pool_timeout"
)
if db_connection_timeout is None:
db_connection_timeout = (
LiteLLMDatabaseConnectionPool.database_connection_pool_timeout.value
)
if database_url and database_url.startswith("os.environ/"):
original_dir = os.getcwd()
# set the working directory to where this script is

View file

@ -1061,6 +1061,52 @@ vertex_live_passthrough_vertex_base = VertexBase()
from fastapi.routing import APIWebSocketRoute
def _inject_websocket_stubs_into_openapi_schema(
openapi_schema: dict, websocket_routes: list
) -> dict:
"""
Add a synthetic GET stub for each WebSocket route so it appears in Swagger UI.
Merges into any existing path entry rather than replacing it a WebSocket route
that shares its path with an HTTP route must not erase the HTTP operation. If
a "get" operation is already documented on the path, the WebSocket stub is
skipped to preserve the real GET.
"""
for route in websocket_routes:
base_path = route.path.split("{")[0].rstrip("?")
parameters = []
try:
if hasattr(route, "dependant") and route.dependant is not None:
# Handle both FastAPI <0.120 and >=0.120
query_params = getattr(route.dependant, "query_params", [])
if query_params:
for param in query_params:
parameters.append(
{
"name": param.name,
"in": "query",
"required": param.required,
"schema": {"type": "string"},
}
)
except (AttributeError, TypeError):
pass
path_entry = openapi_schema["paths"].setdefault(base_path, {})
if "get" not in path_entry:
path_entry["get"] = {
"summary": f"WebSocket: {route.name or base_path}",
"description": "WebSocket connection endpoint",
"operationId": f"websocket_{route.name or base_path.replace('/', '_')}",
"parameters": parameters,
"responses": {"101": {"description": "WebSocket Protocol Switched"}},
"tags": ["WebSocket"],
}
return openapi_schema
def get_openapi_schema():
if app.openapi_schema:
return app.openapi_schema
@ -1083,43 +1129,11 @@ def get_openapi_schema():
route for route in app.routes if isinstance(route, APIWebSocketRoute)
]
# Add each WebSocket route to the schema
for route in websocket_routes:
# Get the base path without query parameters
base_path = route.path.split("{")[0].rstrip("?")
# Extract parameters from the route
parameters = []
try:
if hasattr(route, "dependant") and route.dependant is not None:
# Handle both FastAPI <0.120 and >=0.120
query_params = getattr(route.dependant, "query_params", [])
if query_params:
for param in query_params:
parameters.append(
{
"name": param.name,
"in": "query",
"required": param.required,
"schema": {
"type": "string"
}, # You can make this more specific if needed
}
)
except (AttributeError, TypeError):
# If we can't access query_params, continue without them
pass
openapi_schema["paths"][base_path] = {
"get": {
"summary": f"WebSocket: {route.name or base_path}",
"description": "WebSocket connection endpoint",
"operationId": f"websocket_{route.name or base_path.replace('/', '_')}",
"parameters": parameters,
"responses": {"101": {"description": "WebSocket Protocol Switched"}},
"tags": ["WebSocket"],
}
}
# Add a synthetic GET stub for each so they render in Swagger UI,
# without clobbering existing HTTP operations on the same path.
openapi_schema = _inject_websocket_stubs_into_openapi_schema(
openapi_schema, websocket_routes
)
# Add LLM API request schema bodies for documentation
from litellm.proxy.common_utils.custom_openapi_spec import CustomOpenAPISpec
@ -5937,10 +5951,20 @@ class ProxyConfig:
verbose_proxy_logger.debug(
"guardrails from the DB %s", str(guardrails_in_db)
)
db_guardrail_ids: set = set()
for guardrail in guardrails_in_db:
guardrail_id = guardrail.get("guardrail_id")
if guardrail_id:
db_guardrail_ids.add(guardrail_id)
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=cast(Guardrail, guardrail),
)
# Drop in-memory DB-backed entries whose row was deleted on another
# pod. Config-loaded entries are never touched.
IN_MEMORY_GUARDRAIL_HANDLER.reconcile_db_guardrails(
db_guardrail_ids=db_guardrail_ids
)
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - {}".format(

View file

@ -4977,10 +4977,10 @@ class ProxyUpdateSpend:
timeout=timedelta(seconds=60)
) as transaction:
async with transaction.batch_() as batcher:
for (
end_user_id,
response_cost,
) in end_user_list_transactions.items():
# Sort by end_user_id for consistent lock ordering across pods to prevent deadlocks.
for end_user_id, response_cost in sorted(
end_user_list_transactions.items()
):
if litellm.max_end_user_budget is not None:
pass
batcher.litellm_endusertable.upsert(

View file

@ -7076,11 +7076,11 @@ class Router:
_shared_model_info = {
k: v for k, v in _model_info.items() if k not in _custom_pricing_fields
}
litellm.register_model(
model_cost={
_model_name: _shared_model_info,
}
)
_backend_alias_cost = {_model_name: _shared_model_info}
if "responses/" in _model_name:
_stripped_model_name = _model_name.replace("responses/", "")
_backend_alias_cost[_stripped_model_name] = _shared_model_info
litellm.register_model(model_cost=_backend_alias_cost)
## Check if LLM Deployment is allowed for this deployment
if (
@ -7752,6 +7752,12 @@ class Router:
# initialize client
self._add_deployment(deployment=deployment)
_model_info_dict: dict = deployment.model_info.model_dump(exclude_none=True)
for field in CustomPricingLiteLLMParams.model_fields.keys():
field_value = deployment.litellm_params.get(field)
if field_value is not None:
_model_info_dict[field] = field_value
# Register custom pricing in litellm.model_cost.
# Mirrors _create_deployment() logic to ensure dynamically-added deployments
# (e.g., loaded from DB) also have their custom pricing registered.
@ -7759,13 +7765,31 @@ class Router:
# zero-cost models, causing budget checks to block free models.
_model_id = deployment.model_info.id
if _model_id is not None:
_model_info_dict: dict = deployment.model_info.model_dump(exclude_none=True)
for field in CustomPricingLiteLLMParams.model_fields.keys():
field_value = deployment.litellm_params.get(field)
if field_value is not None:
_model_info_dict[field] = field_value
litellm.register_model(model_cost={_model_id: _model_info_dict})
## REGISTER MODEL INFO IN LITELLM MODEL COST MAP
## OLD MODEL REGISTRATION ## Kept to prevent breaking changes
_model_name = deployment.litellm_params.model
if deployment.litellm_params.custom_llm_provider is not None:
_model_name = (
deployment.litellm_params.custom_llm_provider + "/" + _model_name
)
# For the shared backend key, strip custom pricing fields so that
# one deployment's pricing overrides don't pollute another
# deployment sharing the same backend model name.
# Each deployment's full pricing is already stored under its
# unique model_id above (when present).
_custom_pricing_fields = CustomPricingLiteLLMParams.model_fields.keys()
_shared_model_info = {
k: v for k, v in _model_info_dict.items() if k not in _custom_pricing_fields
}
_backend_alias_cost = {_model_name: _shared_model_info}
if "responses/" in _model_name:
_stripped_model_name = _model_name.replace("responses/", "")
_backend_alias_cost[_stripped_model_name] = _shared_model_info
litellm.register_model(model_cost=_backend_alias_cost)
# add to model names
self._add_model_to_list_and_index_map(
model=_deployment, model_id=deployment.model_info.id

View file

@ -1042,3 +1042,10 @@ class BedrockInvokeAnthropicMessagesRequest(TypedDict, total=False):
thinking: dict
metadata: dict
output_config: dict
# `context_management` is allowed for Bedrock InvokeModel only when it
# carries `compact_20260112` edits paired with the `compact-2026-01-12`
# anthropic-beta header. The Invoke transformation filters edits to the
# supported subset and strips the field entirely when nothing remains, so
# other edit types (e.g. `clear_thinking_20251015`) never reach Bedrock.
context_management: dict

View file

@ -1,6 +1,7 @@
from typing import Any, Dict, List, Optional
from datetime import datetime
from typing import Any, Dict, List, Literal, Optional
from pydantic import BaseModel
from pydantic import BaseModel, ConfigDict, model_validator
class BulkUpdateKeyRequestItem(BaseModel):
@ -40,3 +41,78 @@ class BulkUpdateKeyResponse(BaseModel):
total_requested: int
successful_updates: List[SuccessfulKeyUpdate]
failed_updates: List[FailedKeyUpdate]
class KeyUpdateFields(BaseModel):
"""Allowlist of bulk-broadcastable fields for /team/key/bulk_update; `extra="forbid"` blocks RBAC/ownership/scope mutations even by team admins."""
model_config = ConfigDict(extra="forbid", protected_namespaces=())
# Budgets
max_budget: Optional[float] = None
budget_id: Optional[str] = None
budget_duration: Optional[str] = None
budget_limits: Optional[List[Any]] = None
model_max_budget: Optional[Dict[str, Any]] = None
# Rate limits
tpm_limit: Optional[int] = None
rpm_limit: Optional[int] = None
model_tpm_limit: Optional[Dict[str, Any]] = None
model_rpm_limit: Optional[Dict[str, Any]] = None
max_parallel_requests: Optional[int] = None
rpm_limit_type: Optional[
Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"]
] = None
tpm_limit_type: Optional[
Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"]
] = None
# Temporary budget grants (auto-expire). `spend` deliberately omitted — bulk-zeroing it bypasses budget enforcement; admin-only via /key/update.
temp_budget_increase: Optional[float] = None
temp_budget_expiry: Optional[datetime] = None
# Expiry
duration: Optional[str] = None
# Operational metadata
tags: Optional[List[str]] = None
metadata: Optional[Dict[str, Any]] = None
@model_validator(mode="after")
def validate_temp_budget(self) -> "KeyUpdateFields":
if self.temp_budget_increase is not None or self.temp_budget_expiry is not None:
if self.temp_budget_increase is None or self.temp_budget_expiry is None:
raise ValueError(
"temp_budget_increase and temp_budget_expiry must be set together"
)
return self
@model_validator(mode="after")
def require_at_least_one_field(self) -> "KeyUpdateFields":
# Reject empty payload — would iterate every key with no-op writes.
if not self.model_fields_set:
raise ValueError("update_fields must specify at least one field to update.")
return self
class BulkUpdateTeamKeysRequest(BaseModel):
"""Apply one update payload to many keys inside a team; provide either `key_ids` or `all_keys_in_team=True`."""
team_id: str
key_ids: Optional[List[str]] = None
all_keys_in_team: bool = False
update_fields: KeyUpdateFields
@model_validator(mode="after")
def validate_selection(self) -> "BulkUpdateTeamKeysRequest":
has_key_ids = self.key_ids is not None and len(self.key_ids) > 0
if has_key_ids and self.all_keys_in_team:
raise ValueError(
"Provide either `key_ids` or `all_keys_in_team=True`, not both."
)
if not has_key_ids and not self.all_keys_in_team:
raise ValueError(
"Must provide either `key_ids` (non-empty) or `all_keys_in_team=True`."
)
return self

View file

@ -1,3 +1,5 @@
"""Utility helpers for LiteLLM core request handling and provider support."""
# from __future__ import annotations must be the first non-comment statement
from __future__ import annotations
@ -4405,6 +4407,10 @@ def get_optional_params( # noqa: PLR0915
else False
),
)
if bedrock_route == "claude_platform":
optional_params = BedrockModelInfo.map_claude_platform_auth_params(
passed_params=passed_params, optional_params=optional_params
)
elif custom_llm_provider == "cloudflare":
optional_params = litellm.CloudflareChatConfig().map_openai_params(
model=model,
@ -9490,6 +9496,49 @@ def get_non_default_completion_params(kwargs: dict) -> dict:
return non_default_params
def peek_reasoning_summary_aliases(optional_params: dict) -> Optional[Any]:
"""Read AI-SDK-style reasoning summary from optional_params or nested extra_body.
Uses key membership (not ``or`` chains) so falsy values like ``""`` are not skipped.
"""
if "reasoningSummary" in optional_params:
return optional_params["reasoningSummary"]
if "reasoning_summary" in optional_params:
return optional_params["reasoning_summary"]
extra_body = optional_params.get("extra_body")
if isinstance(extra_body, dict):
if "reasoningSummary" in extra_body:
return extra_body["reasoningSummary"]
if "reasoning_summary" in extra_body:
return extra_body["reasoning_summary"]
return None
def strip_reasoning_summary_aliases_from_optional_params(
optional_params: dict,
) -> Tuple[dict, Optional[Any]]:
"""Copy optional_params; remove reasoningSummary aliases from top-level and extra_body."""
op = dict(optional_params)
rs_val = op.pop("reasoningSummary", None)
snake_rs_val = op.pop("reasoning_summary", None)
if rs_val is None:
rs_val = snake_rs_val
eb = op.get("extra_body")
if isinstance(eb, dict):
eb = dict(eb)
eb_rs_val = eb.pop("reasoningSummary", None)
eb_snake_rs_val = eb.pop("reasoning_summary", None)
if rs_val is None:
rs_val = eb_rs_val
if rs_val is None:
rs_val = eb_snake_rs_val
if eb:
op["extra_body"] = eb
else:
op.pop("extra_body", None)
return op, rs_val
def get_non_default_transcription_params(kwargs: dict) -> dict:
from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS

View file

@ -275,6 +275,33 @@ filterwarnings = [
"ignore::DeprecationWarning:pytest_asyncio.plugin",
]
[tool.mutmut]
# Mutation-testing scope. Driven by the manually-triggered workflow at
# .github/workflows/mutation-test.yml. mutmut is not part of the project's
# default install; it is pulled in via `uv run --with mutmut==<version>` in CI.
# `also_copy = ["litellm/"]` is required because mutmut runs in a `mutants/`
# sandbox and the test conftest imports from across the litellm package.
paths_to_mutate = [
"litellm/proxy/management_endpoints/",
]
tests_dir = [
"tests/test_litellm/proxy/management_endpoints/",
]
also_copy = [
"litellm/",
]
# Disable rerun/parallel plugins for mutation runs:
# - pytest-retry triggers an `INTERNALERROR: no option named 'filtered_exceptions'`
# when invoked via mutmut's in-process `pytest.main()` call.
# - rerunning a "failed" test on a mutant would mask which mutants are killed
# vs. survive, so reruns are wrong for mutation testing regardless.
# - xdist is unnecessary inside mutmut (mutmut handles its own parallelism).
pytest_add_cli_args = [
"-p", "no:retry",
"-p", "no:rerunfailures",
"-p", "no:xdist",
]
[tool.coverage.run]
source = ["litellm"]
relative_files = true

423
scripts/mutation_report.py Normal file
View file

@ -0,0 +1,423 @@
#!/usr/bin/env python3
"""Generate an agent-actionable mutation testing report.
Reads the mutmut sandbox state at `mutants/` and produces a single
`mutation-report.md` grouped by function. For each function with surviving
mutants, the report embeds the original function source (via AST), the
unified diff for each surviving mutation (via `mutmut show`), and the
existing test file(s) followed by an ACH-style instruction asking the
reader to write tests that kill the survivors.
Run after `mutmut run` and `mutmut export-cicd-stats`. Expects mutmut to be
invokable as `uv run --no-sync --with mutmut==<version> mutmut <subcommand>`.
"""
from __future__ import annotations
import ast
import json
import re
import subprocess
import sys
import tomllib
from collections import defaultdict
from difflib import SequenceMatcher
from pathlib import Path
from textwrap import dedent
ROOT = Path(__file__).resolve().parent.parent
MUTMUT_INVOCATION = ["uv", "run", "--no-sync", "--with", "mutmut==3.5.0", "mutmut"]
def load_mutmut_config() -> dict:
with open(ROOT / "pyproject.toml", "rb") as f:
return tomllib.load(f)["tool"]["mutmut"]
def get_survivors() -> list[str]:
proc = subprocess.run(
[*MUTMUT_INVOCATION, "results"], capture_output=True, text=True, check=False
)
survivors = []
for line in proc.stdout.splitlines():
m = re.match(r"\s*(\S+):\s*survived\s*$", line)
if m:
survivors.append(m.group(1))
return survivors
def get_mutmut_show(mutant_name: str) -> str:
proc = subprocess.run(
[*MUTMUT_INVOCATION, "show", mutant_name],
capture_output=True,
text=True,
check=False,
)
return proc.stdout.strip() or "(mutmut show produced no output)"
def parse_mutant_name(name: str) -> tuple[str, str, str]:
"""Parse `<dotted.module>.x_<function>__mutmut_<N>` -> (module, function, N).
mutmut prefixes mutated functions with `x_` (single underscore). For a
function named `foo`, mutants are `x_foo__mutmut_N`. For a function named
`_foo` (leading underscore), the mutant becomes `x__foo__mutmut_N` so
the regex matches a single underscore after `x` and captures everything
(including any leading underscores) up to `__mutmut_<N>`.
"""
m = re.match(r"^(.+)\.x_(.+)__mutmut_(\d+)$", name)
if not m:
return name, name, "?"
return m.group(1), m.group(2), m.group(3)
def function_anchor(module_path: str, function_name: str) -> str:
return re.sub(r"[^a-z0-9_-]+", "-", f"{module_path}-{function_name}".lower()).strip(
"-"
)
def module_to_file(module_path: str) -> Path | None:
candidate = ROOT / Path(*module_path.split(".")).with_suffix(".py")
return candidate if candidate.exists() else None
def find_function_in_file(
file_path: Path, function_name: str
) -> tuple[int, int, str, list[int]] | None:
"""Find a top-level or nested function by name; returns the first match.
Returns ``(start_line, end_line, source, all_match_lines)`` or ``None``.
``all_match_lines`` is the start line of every function (any nesting
level) in the file with this name. When ``len(all_match_lines) > 1`` the
file defines the same name in multiple places (e.g., a module-level
helper and a class method) mutmut's mutant identifier does not carry
class context, so we can't determine which definition was mutated.
Callers surface a disambiguation note in that case.
"""
src = file_path.read_text()
tree = ast.parse(src)
matches = [
node
for node in ast.walk(tree)
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
and node.name == function_name
]
if not matches:
return None
first = matches[0]
lines = src.splitlines()
return (
first.lineno,
first.end_lineno,
"\n".join(lines[first.lineno - 1 : first.end_lineno]),
[m.lineno for m in matches],
)
def collect_test_files(tests_dir: list[str]) -> list[Path]:
found: list[Path] = []
for entry in tests_dir:
p = ROOT / entry
if p.is_file():
found.append(p)
elif p.is_dir():
found.extend(sorted(p.rglob("test_*.py")))
return found
def _indent_of(line: str) -> str:
return line[: len(line) - len(line.lstrip())]
def render_meta_style_mutant(
module_path: str, function_name: str, mutant_num: str
) -> str | None:
"""Render the mutated function with `# MUTANT START`/`# MUTANT END` delimiters.
Reads `mutants/<module>.py` (the trampoline file mutmut emits), finds
`x_<func>__mutmut_orig` and `x_<func>__mutmut_<N>`, and renders the
mutated version with the lines that differ from `__mutmut_orig` wrapped
in `# MUTANT START`/`# MUTANT END` comments — the format from Meta's
ACH paper (arXiv 2501.12862, Table 1).
The function header is rewritten to use the original function name so
the agent sees the source as it would appear in the file (rather than
mutmut's internal `x_*__mutmut_<N>` name).
Returns None if the trampoline file or either function cannot be found
(the caller falls back to the unified diff).
"""
trampoline = ROOT / "mutants" / Path(*module_path.split(".")).with_suffix(".py")
if not trampoline.exists():
return None
src = trampoline.read_text()
try:
tree = ast.parse(src)
except SyntaxError:
return None
file_lines = src.splitlines()
orig_def = f"x_{function_name}__mutmut_orig"
mutant_def = f"x_{function_name}__mutmut_{mutant_num}"
orig_node = mutated_node = None
for node in ast.walk(tree):
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
if node.name == orig_def:
orig_node = node
elif node.name == mutant_def:
mutated_node = node
if orig_node is None or mutated_node is None:
return None
orig_lines = file_lines[orig_node.lineno - 1 : orig_node.end_lineno]
mutated_lines = file_lines[mutated_node.lineno - 1 : mutated_node.end_lineno]
if not orig_lines or not mutated_lines:
return None
# Rewrite the def line to use the original (non-trampolined) function name
# so the agent sees the function as it appears in the source file.
orig_lines[0] = orig_lines[0].replace(orig_def, function_name, 1)
mutated_lines[0] = mutated_lines[0].replace(mutant_def, function_name, 1)
matcher = SequenceMatcher(a=orig_lines, b=mutated_lines)
out: list[str] = []
in_diff = False
for op, i1, i2, j1, j2 in matcher.get_opcodes():
if op == "equal":
if in_diff:
# Close the block at the indent of the line just inside it.
indent = _indent_of(out[-1]) if out else ""
out.append(f"{indent}# MUTANT END")
in_diff = False
out.extend(mutated_lines[j1:j2])
else:
if not in_diff:
# Open the block at the indent of the first differing line.
if j1 < len(mutated_lines):
indent = _indent_of(mutated_lines[j1])
elif i1 < len(orig_lines):
indent = _indent_of(orig_lines[i1])
else:
indent = ""
out.append(f"{indent}# MUTANT START")
in_diff = True
if op == "delete":
# Mutation removed lines — surface what was deleted as a
# comment so the agent can see the intent of the change.
for deleted in orig_lines[i1:i2]:
indent = _indent_of(deleted)
out.append(f"{indent}# (deleted by mutation): {deleted.lstrip()}")
else:
# replace / insert: take from mutated_lines
out.extend(mutated_lines[j1:j2])
if in_diff:
indent = _indent_of(out[-1]) if out else ""
out.append(f"{indent}# MUTANT END")
return "\n".join(out)
def render(config: dict, survivors: list[str], stats: dict | None) -> str:
by_function: dict[tuple[str, str], list[tuple[str, str]]] = defaultdict(list)
for survivor in survivors:
module_path, function_name, mutant_num = parse_mutant_name(survivor)
by_function[(module_path, function_name)].append((survivor, mutant_num))
out: list[str] = []
out.append("# Mutation Test Report")
out.append("")
out.append("## Summary")
out.append("")
if stats:
total = stats.get("total", 0) or sum(
stats.get(k, 0)
for k in (
"killed",
"survived",
"no_tests",
"skipped",
"suspicious",
"timeout",
"segfault",
)
)
killed = stats.get("killed", 0)
survived = stats.get("survived", 0)
score = (killed / total * 100) if total else 0.0
out.append(f"- Total mutants: **{total}**")
out.append(f"- Killed: **{killed}**")
out.append(f"- Survived: **{survived}**")
out.append(f"- Mutation score: **{score:.1f}%**")
for k in ("no_tests", "skipped", "suspicious", "timeout", "segfault"):
v = stats.get(k, 0)
if v:
out.append(f"- {k.replace('_', ' ').title()}: {v}")
else:
out.append(f"- Survivors found: **{len(survivors)}**")
out.append("- (mutmut-cicd-stats.json not available — full counts unavailable)")
out.append("")
if not survivors:
out.append("**No surviving mutants — the test suite caught every mutation.**")
out.append("")
return "\n".join(out)
out.append("## Surviving mutants by function")
out.append("")
for (module_path, function_name), items in by_function.items():
anchor = function_anchor(module_path, function_name)
out.append(
f"- [`{function_name}`](#{anchor}) — {len(items)} mutant"
f"{'s' if len(items) != 1 else ''} ({module_path})"
)
out.append("")
for (module_path, function_name), items in by_function.items():
anchor = function_anchor(module_path, function_name)
out.append(f'<a id="{anchor}"></a>')
out.append(f"## `{module_path}.{function_name}`")
out.append("")
out.append(f"**Module:** `{module_path}`")
file_path = module_to_file(module_path)
if file_path is None:
out.append("")
out.append(f"_(could not locate source file for module `{module_path}`)_")
out.append("")
else:
rel = file_path.relative_to(ROOT)
out.append(f"**File:** `{rel}`")
out.append("")
found = find_function_in_file(file_path, function_name)
if found:
start, end, fn_src, all_lines = found
out.append(f"### Original function (lines {start}-{end})")
out.append("")
if len(all_lines) > 1:
line_list = ", ".join(str(line) for line in all_lines)
out.append(
f"> **Note:** {len(all_lines)} functions named "
f"`{function_name}` are defined in this file at lines "
f"{line_list}. Showing the first match. mutmut's "
f"mutant identifier does not carry class context, so "
f"the body below may not correspond to the function "
f"that was actually mutated — verify manually before "
f"writing the killing test."
)
out.append("")
out.append("```python")
out.append(fn_src)
out.append("```")
out.append("")
else:
out.append(f"_(could not locate `{function_name}` in {rel} via AST)_")
out.append("")
out.append(f"### Surviving mutations ({len(items)})")
out.append("")
for i, (mutant_name, mutant_num) in enumerate(items, 1):
out.append(f"#### Mutation {i} of {len(items)} — `{mutant_name}`")
out.append("")
meta_style = render_meta_style_mutant(
module_path, function_name, mutant_num
)
if meta_style is not None:
out.append(
"Mutated function (the bug is delimited by "
"`# MUTANT START` / `# MUTANT END`):"
)
out.append("")
out.append("```python")
out.append(meta_style)
out.append("```")
out.append("")
out.append("<details><summary>Unified diff (`mutmut show`)</summary>")
out.append("")
out.append("```diff")
out.append(get_mutmut_show(mutant_name))
out.append("```")
out.append("")
out.append("</details>")
out.append("")
else:
# Fallback: trampoline file or function lookup failed.
out.append("```diff")
out.append(get_mutmut_show(mutant_name))
out.append("```")
out.append("")
test_files = collect_test_files(config.get("tests_dir", []))
if test_files:
out.append("## Existing tests")
out.append("")
out.append(
"These are the test files that mutmut considered when classifying the "
"mutants above. New tests should be added here, matching existing "
"conventions, fixtures, and naming."
)
out.append("")
for tf in test_files:
rel = tf.relative_to(ROOT)
out.append(f"### `{rel}`")
out.append("")
out.append("```python")
out.append(tf.read_text())
out.append("```")
out.append("")
out.append("## Task")
out.append("")
out.append(
dedent(
"""\
For each surviving mutant listed above, write a new test in the
existing test file (matching its conventions, fixtures, and naming
style) that:
- **Fails** when the mutated version of the function is in place.
- **Passes** when the original (correct) version is in place.
Aim for one test per surviving mutant. If multiple mutants in the
same function can be killed by a single test, that is fine note
which mutant numbers in the test name or docstring.
Do not modify the source file. Only add tests.
"""
).strip()
)
out.append("")
return "\n".join(out)
def main() -> int:
config = load_mutmut_config()
stats_file = ROOT / "mutants" / "mutmut-cicd-stats.json"
stats: dict | None = None
if stats_file.exists():
try:
stats = json.loads(stats_file.read_text())
except json.JSONDecodeError as exc:
print(f"warning: could not parse {stats_file}: {exc}", file=sys.stderr)
survivors = get_survivors()
report = render(config, survivors, stats)
out_path = ROOT / "mutation-report.md"
out_path.write_text(report)
print(
f"Wrote {out_path} ({len(survivors)} survivor"
f"{'s' if len(survivors) != 1 else ''}, {len(report)} chars)"
)
return 0
if __name__ == "__main__":
raise SystemExit(main())

View file

@ -233,6 +233,12 @@ async def test_reset_budget_endusers_partial_failure():
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -400,6 +406,12 @@ async def test_reset_budget_continues_other_categories_on_failure():
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -884,6 +896,12 @@ async def test_service_logger_endusers_success():
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -966,6 +984,12 @@ async def test_service_logger_endusers_failure():
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -1060,6 +1084,10 @@ async def test_reset_budget_for_litellm_team_members_called():
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()

View file

@ -11,7 +11,7 @@ import litellm
def test_completion_openrouter_reasoning_content():
litellm._turn_on_debug()
resp = litellm.completion(
model="openrouter/anthropic/claude-3.7-sonnet",
model="openrouter/anthropic/claude-sonnet-4",
messages=[{"role": "user", "content": "Hello world"}],
reasoning={"effort": "high"},
)

View file

@ -268,51 +268,63 @@ def test_aaparallel_function_call_with_anthropic_thinking(model):
from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
_PARALLEL_TOOL_HISTORY_MESSAGES = [
{
"role": "user",
"content": "What's the weather like in San Francisco, Tokyo, and Paris? - give me 3 responses",
},
Message(
content="Here are the current weather conditions for San Francisco, Tokyo, and Paris:",
role="assistant",
tool_calls=[
ChatCompletionMessageToolCall(
index=1,
function=Function(
arguments='{"location": "San Francisco, CA", "unit": "fahrenheit"}',
name="get_current_weather",
),
id="tooluse_Jj98qn6xQlOP_PiQr-w9iA",
type="function",
)
],
function_call=None,
),
{
"tool_call_id": "tooluse_Jj98qn6xQlOP_PiQr-w9iA",
"role": "tool",
"name": "get_current_weather",
"content": '{"location": "San Francisco", "temperature": "72", "unit": "fahrenheit"}',
},
]
@pytest.mark.parametrize(
"model, provider",
"model, messages, expect_unsupported_params_error",
[
# Bedrock Converse still requires modify_params to inject the dummy tool.
(
"anthropic.claude-3-sonnet-20240229-v1:0",
"bedrock",
_PARALLEL_TOOL_HISTORY_MESSAGES,
True,
),
("claude-haiku-4-5-20251001", "anthropic"),
],
)
@pytest.mark.parametrize(
"messages, expected_error_msg",
[
# Anthropic Messages API: dummy tool is injected without modify_params.
(
"claude-haiku-4-5-20251001",
_PARALLEL_TOOL_HISTORY_MESSAGES,
False,
),
(
"anthropic.claude-3-sonnet-20240229-v1:0",
[
{
"role": "user",
"content": "What's the weather like in San Francisco, Tokyo, and Paris? - give me 3 responses",
},
Message(
content="Here are the current weather conditions for San Francisco, Tokyo, and Paris:",
role="assistant",
tool_calls=[
ChatCompletionMessageToolCall(
index=1,
function=Function(
arguments='{"location": "San Francisco, CA", "unit": "fahrenheit"}',
name="get_current_weather",
),
id="tooluse_Jj98qn6xQlOP_PiQr-w9iA",
type="function",
)
],
function_call=None,
),
{
"tool_call_id": "tooluse_Jj98qn6xQlOP_PiQr-w9iA",
"role": "tool",
"name": "get_current_weather",
"content": '{"location": "San Francisco", "temperature": "72", "unit": "fahrenheit"}',
},
}
],
True,
False,
),
(
"claude-haiku-4-5-20251001",
[
{
"role": "user",
@ -324,25 +336,26 @@ from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
],
)
def test_parallel_function_call_anthropic_error_msg(
model, provider, messages, expected_error_msg
model, messages, expect_unsupported_params_error
):
"""
Anthropic doesn't support tool calling without `tools=` param specified.
Tool history without an explicit ``tools`` param:
Ensure this error is thrown when `tools=` param is not specified. But tool call requests are made.
- Bedrock **Converse** still raises ``UnsupportedParamsError`` unless
``litellm.modify_params`` is enabled (dummy tool is only added there).
- **Anthropic** (and Bedrock Invoke via ``AnthropicConfig.transform_request``)
always get a dummy tool so CLIs work with ``modify_params`` left off.
Reference Issue: https://github.com/BerriAI/litellm/issues/5747, https://github.com/BerriAI/litellm/issues/5388
"""
# Ensure modify_params is False so UnsupportedParamsError is raised
# Ensure modify_params is False so Bedrock Converse path still raises.
# (other tests in this file set it to True and don't reset it)
original_modify_params = litellm.modify_params
litellm.modify_params = False
try:
litellm.set_verbose = True
messages = messages
if expected_error_msg:
if expect_unsupported_params_error:
with pytest.raises(litellm.UnsupportedParamsError) as e:
second_response = litellm.completion(
model=model,

View file

@ -587,12 +587,21 @@ def test_foward_litellm_user_info_to_backend_llm_call():
user_api_key_dict=user_api_key_dict,
)
# All header values must be str/bytes so httpx won't reject them when the
# downstream client builds the request (regression: #27458).
for k, v in data.items():
assert isinstance(v, (str, bytes)), (
f"header {k!r} has non-str value {v!r} ({type(v).__name__}); "
"httpx will raise 'Header value must be str or bytes' when the LLM "
"request is built."
)
expected_data = {
"x-litellm-user_api_key_user_id": "test_user_id",
"x-litellm-user_api_key_org_id": "test_org_id",
"x-litellm-user_api_key_hash": "test_api_key",
"x-litellm-user_api_key_spend": 0.0,
"x-litellm-user_api_key_auth_metadata": {},
"x-litellm-user_api_key_spend": "0.0",
"x-litellm-user_api_key_auth_metadata": "{}",
}
assert json.dumps(data, sort_keys=True) == json.dumps(expected_data, sort_keys=True)

View file

@ -2098,6 +2098,56 @@ def test_map_optional_params_preserves_reasoning_summary():
assert responses_api_request["reasoning"]["summary"] == "detailed"
def test_map_optional_params_tool_choice_chat_nested_to_responses_api():
"""Chat tool_choice must become Responses ToolChoiceFunction (top-level name)."""
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
handler = LiteLLMResponsesTransformationHandler()
responses_api_request = ResponsesAPIOptionalRequestParams()
handler._map_optional_params_to_responses_api_request(
{
"stream": False,
"tool_choice": {
"type": "function",
"function": {"name": "Echo"},
},
},
responses_api_request,
)
assert responses_api_request["tool_choice"] == {
"type": "function",
"name": "Echo",
}
@pytest.mark.parametrize(
("tool_choice", "expected"),
[
("auto", "auto"),
("none", "none"),
(
{"type": "function", "name": "Echo"},
{"type": "function", "name": "Echo"},
),
(
{"type": "function", "name": "foo", "function": {"name": "bar"}},
{"type": "function", "name": "foo"},
),
({"type": "required"}, {"type": "required"}),
],
)
def test_normalize_tool_choice_for_responses_api(tool_choice, expected):
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
handler = LiteLLMResponsesTransformationHandler()
assert handler._normalize_tool_choice_for_responses_api(tool_choice) == expected
def test_convert_chat_completion_file_type_to_input_file():
"""
Test that Chat Completion content with type 'file' is correctly mapped

View file

@ -3159,3 +3159,630 @@ class TestResponseIdFallback(unittest.TestCase):
otel.set_attributes(mock_span, kwargs, response_obj)
mock_span.set_attribute.assert_any_call("litellm.call_id", call_id)
class TestOpenTelemetryResponsesAPI(unittest.TestCase):
"""
Tests for Responses API (/v1/responses) OTel span attributes.
The Responses API uses ``output`` (list of output items) instead of
``choices``, ``instructions`` instead of ``system_instructions``, and
``status`` instead of per-choice ``finish_reason``.
See: https://github.com/BerriAI/litellm/issues/25840
"""
def _base_kwargs(self, **overrides):
"""Return minimal kwargs for set_attributes with Responses API defaults."""
kwargs = {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "What is 2+2?"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "resp_abc123",
"call_type": "responses",
"metadata": {},
},
}
kwargs.update(overrides)
return kwargs
def _responses_api_response_obj(self, text="The answer is 4.", status="completed"):
"""Return a dict mimicking ResponsesAPIResponse with a message output."""
return {
"id": "resp_abc123",
"model": "gpt-4o",
"status": status,
"output": [
{
"type": "message",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": text,
}
],
}
],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 20,
"total_tokens": 30,
},
}
def _get_attr(self, mock_span, attr_name):
"""Extract the value set for a specific attribute name, or None."""
calls = [
call
for call in mock_span.set_attribute.call_args_list
if call[0][0] == attr_name
]
if not calls:
return None
return calls[0][0][1]
# ------------------------------------------------------------------
# gen_ai.output.messages
# ------------------------------------------------------------------
def test_output_messages_populated_for_responses_api(self):
"""gen_ai.output.messages must be set when response has output items."""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = self._base_kwargs()
response_obj = self._responses_api_response_obj(text="The answer is 4.")
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
raw = self._get_attr(mock_span, "gen_ai.output.messages")
self.assertIsNotNone(raw, "gen_ai.output.messages should be set")
parsed = json.loads(raw)
self.assertIsInstance(parsed, list)
self.assertEqual(len(parsed), 1)
self.assertEqual(parsed[0]["role"], "assistant")
self.assertIn("parts", parsed[0])
self.assertEqual(parsed[0]["parts"][0]["type"], "text")
self.assertEqual(parsed[0]["parts"][0]["content"], "The answer is 4.")
def test_output_messages_with_multiple_content_items(self):
"""Multiple output_text items in a single message should all appear as parts."""
otel = OpenTelemetry()
mock_span = MagicMock()
response_obj = {
"id": "resp_multi",
"model": "gpt-4o",
"status": "completed",
"output": [
{
"type": "message",
"role": "assistant",
"content": [
{"type": "output_text", "text": "First paragraph."},
{"type": "output_text", "text": "Second paragraph."},
],
}
],
}
otel.set_attributes(
span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj
)
raw = self._get_attr(mock_span, "gen_ai.output.messages")
parsed = json.loads(raw)
self.assertEqual(len(parsed[0]["parts"]), 2)
self.assertEqual(parsed[0]["parts"][0]["content"], "First paragraph.")
self.assertEqual(parsed[0]["parts"][1]["content"], "Second paragraph.")
def test_output_messages_with_function_call(self):
"""function_call output items should appear as tool_call parts."""
otel = OpenTelemetry()
mock_span = MagicMock()
response_obj = {
"id": "resp_fc",
"model": "gpt-4o",
"status": "completed",
"output": [
{
"type": "function_call",
"name": "get_weather",
"call_id": "call_abc",
"arguments": '{"location": "SF"}',
}
],
}
otel.set_attributes(
span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj
)
raw = self._get_attr(mock_span, "gen_ai.output.messages")
parsed = json.loads(raw)
self.assertEqual(len(parsed), 1)
self.assertEqual(parsed[0]["role"], "assistant")
self.assertEqual(parsed[0]["parts"][0]["type"], "tool_call")
self.assertEqual(parsed[0]["parts"][0]["name"], "get_weather")
self.assertEqual(parsed[0]["parts"][0]["arguments"], '{"location": "SF"}')
self.assertEqual(parsed[0]["parts"][0]["id"], "call_abc")
def test_output_messages_mixed_message_and_function_call(self):
"""Mixed output with both message and function_call items."""
otel = OpenTelemetry()
mock_span = MagicMock()
response_obj = {
"id": "resp_mixed",
"model": "gpt-4o",
"status": "completed",
"output": [
{
"type": "message",
"role": "assistant",
"content": [
{"type": "output_text", "text": "Let me check the weather."},
],
},
{
"type": "function_call",
"name": "get_weather",
"call_id": "call_xyz",
"arguments": "{}",
},
],
}
otel.set_attributes(
span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj
)
raw = self._get_attr(mock_span, "gen_ai.output.messages")
parsed = json.loads(raw)
self.assertEqual(len(parsed), 2)
self.assertEqual(parsed[0]["role"], "assistant")
self.assertEqual(parsed[0]["parts"][0]["content"], "Let me check the weather.")
self.assertEqual(parsed[1]["parts"][0]["type"], "tool_call")
def test_output_messages_empty_text_skipped(self):
"""Output items with empty text should not produce parts."""
otel = OpenTelemetry()
mock_span = MagicMock()
response_obj = {
"id": "resp_empty",
"model": "gpt-4o",
"status": "completed",
"output": [
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": ""}],
}
],
}
otel.set_attributes(
span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj
)
# No output messages should be set since the text is empty
raw = self._get_attr(mock_span, "gen_ai.output.messages")
self.assertIsNone(raw, "Empty output text should not produce gen_ai.output.messages")
def test_choices_still_work(self):
"""Existing choices-based responses must still work (no regression)."""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "test-id",
"call_type": "completion",
"metadata": {},
},
}
response_obj = {
"id": "chatcmpl-123",
"model": "gpt-4",
"choices": [
{
"finish_reason": "stop",
"message": {"role": "assistant", "content": "Hi there!"},
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15},
}
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
raw = self._get_attr(mock_span, "gen_ai.output.messages")
parsed = json.loads(raw)
self.assertEqual(parsed[0]["parts"][0]["content"], "Hi there!")
self.assertEqual(parsed[0]["finish_reason"], "stop")
# ------------------------------------------------------------------
# gen_ai.response.finish_reasons
# ------------------------------------------------------------------
def test_finish_reasons_from_status(self):
"""gen_ai.response.finish_reasons should use ResponsesAPIResponse.status."""
otel = OpenTelemetry()
mock_span = MagicMock()
otel.set_attributes(
span=mock_span,
kwargs=self._base_kwargs(),
response_obj=self._responses_api_response_obj(status="completed"),
)
raw = self._get_attr(mock_span, "gen_ai.response.finish_reasons")
self.assertIsNotNone(raw)
parsed = json.loads(raw)
self.assertEqual(parsed, ["completed"])
def test_finish_reasons_incomplete_status(self):
"""Non-completed status values should still be captured."""
otel = OpenTelemetry()
mock_span = MagicMock()
otel.set_attributes(
span=mock_span,
kwargs=self._base_kwargs(),
response_obj=self._responses_api_response_obj(status="incomplete"),
)
raw = self._get_attr(mock_span, "gen_ai.response.finish_reasons")
parsed = json.loads(raw)
self.assertEqual(parsed, ["incomplete"])
# ------------------------------------------------------------------
# gen_ai.system_instructions
# ------------------------------------------------------------------
def test_system_instructions_from_instructions_kwarg(self):
"""Responses API passes system prompt as kwargs['instructions']."""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = self._base_kwargs(instructions="You are a math tutor.")
response_obj = self._responses_api_response_obj()
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
value = self._get_attr(mock_span, "gen_ai.system_instructions")
self.assertEqual(value, "You are a math tutor.")
def test_system_instructions_from_system_kwarg(self):
"""Anthropic Messages API passes system prompt as kwargs['system']."""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = self._base_kwargs(system="You are a helpful assistant.")
response_obj = self._responses_api_response_obj()
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
value = self._get_attr(mock_span, "gen_ai.system_instructions")
self.assertEqual(value, "You are a helpful assistant.")
def test_system_instructions_from_system_instructions_kwarg(self):
"""Vertex AI Gemini path uses kwargs['system_instructions'] (existing behavior)."""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = self._base_kwargs(
system_instructions=[{"role": "system", "content": "Be concise."}]
)
response_obj = self._responses_api_response_obj()
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
raw = self._get_attr(mock_span, "gen_ai.system_instructions")
self.assertIsNotNone(raw)
parsed = json.loads(raw)
self.assertEqual(parsed[0]["role"], "system")
self.assertIn("parts", parsed[0])
def test_system_instructions_precedence(self):
"""system_instructions takes precedence over instructions and system."""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = self._base_kwargs(
system_instructions="From Gemini",
instructions="From Responses API",
system="From Anthropic",
)
response_obj = self._responses_api_response_obj()
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
# system_instructions (string) should win — it's checked first
value = self._get_attr(mock_span, "gen_ai.system_instructions")
self.assertEqual(value, "From Gemini")
def test_no_system_instructions_when_absent(self):
"""No gen_ai.system_instructions attr when none of the kwargs are set."""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = self._base_kwargs()
response_obj = self._responses_api_response_obj()
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
value = self._get_attr(mock_span, "gen_ai.system_instructions")
self.assertIsNone(value)
class TestTransformResponsesAPIOutput(unittest.TestCase):
"""
Unit tests for _transform_responses_api_output_to_otel.
"""
def test_message_with_output_text(self):
otel = OpenTelemetry()
output = [
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Hello!"}],
}
]
result = otel._transform_responses_api_output_to_otel(output)
self.assertEqual(len(result), 1)
self.assertEqual(result[0]["role"], "assistant")
self.assertEqual(result[0]["parts"], [{"type": "text", "content": "Hello!"}])
def test_function_call_item(self):
otel = OpenTelemetry()
output = [
{
"type": "function_call",
"name": "search",
"call_id": "call_1",
"arguments": '{"q": "test"}',
}
]
result = otel._transform_responses_api_output_to_otel(output)
self.assertEqual(len(result), 1)
self.assertEqual(result[0]["role"], "assistant")
self.assertEqual(result[0]["parts"][0]["type"], "tool_call")
self.assertEqual(result[0]["parts"][0]["name"], "search")
self.assertEqual(result[0]["parts"][0]["id"], "call_1")
def test_function_call_without_call_id(self):
otel = OpenTelemetry()
output = [
{
"type": "function_call",
"name": "search",
"arguments": "{}",
}
]
result = otel._transform_responses_api_output_to_otel(output)
self.assertNotIn("id", result[0]["parts"][0])
def test_unknown_type_ignored(self):
otel = OpenTelemetry()
output = [{"type": "reasoning", "content": "thinking..."}]
result = otel._transform_responses_api_output_to_otel(output)
self.assertEqual(result, [])
def test_non_dict_items_ignored(self):
otel = OpenTelemetry()
output = ["not a dict", 42, None]
result = otel._transform_responses_api_output_to_otel(output)
self.assertEqual(result, [])
def test_empty_output(self):
otel = OpenTelemetry()
result = otel._transform_responses_api_output_to_otel([])
self.assertEqual(result, [])
def test_message_with_empty_text_skipped(self):
otel = OpenTelemetry()
output = [
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": ""}],
}
]
result = otel._transform_responses_api_output_to_otel(output)
self.assertEqual(result, [])
def test_message_default_role(self):
"""Messages without explicit role should default to assistant."""
otel = OpenTelemetry()
output = [
{
"type": "message",
"content": [{"type": "output_text", "text": "Hi"}],
}
]
result = otel._transform_responses_api_output_to_otel(output)
self.assertEqual(result[0]["role"], "assistant")
def test_pydantic_like_objects_accepted(self):
"""Items with .get() but not isinstance(dict) should be accepted."""
class FakeOutputItem:
"""Mimics BaseLiteLLMOpenAIResponseObject duck-typing."""
def __init__(self, data):
self._data = data
def get(self, key, default=None):
return self._data.get(key, default)
class FakeContent:
def __init__(self, data):
self._data = data
def get(self, key, default=None):
return self._data.get(key, default)
otel = OpenTelemetry()
output = [
FakeOutputItem(
{
"type": "message",
"role": "assistant",
"content": [
FakeContent({"type": "output_text", "text": "Pydantic works!"}),
],
}
)
]
result = otel._transform_responses_api_output_to_otel(output)
self.assertEqual(len(result), 1)
self.assertEqual(result[0]["parts"][0]["content"], "Pydantic works!")
class TestSystemInstructionsPrecedence(unittest.TestCase):
"""Tests for the is-not-None precedence in system_instructions coalescing."""
def _get_attr(self, mock_span, attr_name):
calls = [
call
for call in mock_span.set_attribute.call_args_list
if call[0][0] == attr_name
]
if not calls:
return None
return calls[0][0][1]
def _base_kwargs(self, **overrides):
kwargs = {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "Hi"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "test-id",
"call_type": "responses",
"metadata": {},
},
}
kwargs.update(overrides)
return kwargs
def test_empty_list_system_instructions_does_not_fallthrough(self):
"""An empty list for system_instructions should NOT fall through to instructions."""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = self._base_kwargs(
system_instructions=[],
instructions="Should not be used",
)
response_obj = {"id": "r1", "model": "gpt-4o"}
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
# system_instructions is [] (falsy but not None), so it wins.
# Since it's an empty list, no attribute should be set (nothing to transform).
value = self._get_attr(mock_span, "gen_ai.system_instructions")
# The empty list is truthy for `is not None` but produces empty
# transformed output — the attribute should NOT contain "Should not be used".
if value is not None:
self.assertNotIn("Should not be used", str(value))
class TestResponsesAPIToolCallSpanAttributes(unittest.TestCase):
"""Tests for per-tool-call span attributes on Responses API function_call items."""
def _base_kwargs(self):
return {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "What is the weather?"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "resp_tc",
"call_type": "responses",
"metadata": {},
},
}
def test_per_tool_call_attributes_emitted(self):
"""function_call output items should produce per-tool-call span attributes."""
otel = OpenTelemetry()
mock_span = MagicMock()
response_obj = {
"id": "resp_tc",
"model": "gpt-4o",
"status": "completed",
"output": [
{
"type": "function_call",
"name": "get_weather",
"call_id": "call_abc",
"arguments": '{"location": "SF"}',
}
],
}
otel.set_attributes(span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj)
# Verify per-tool-call attributes were set (same format as choices branch)
attr_names = [call[0][0] for call in mock_span.set_attribute.call_args_list]
tool_call_attrs = [a for a in attr_names if "function_call" in a]
self.assertTrue(len(tool_call_attrs) > 0, "Per-tool-call span attributes should be emitted")
# Verify the name attribute specifically
mock_span.set_attribute.assert_any_call(
"gen_ai.completion.0.function_call.name", "get_weather"
)
mock_span.set_attribute.assert_any_call(
"gen_ai.completion.0.function_call.arguments", '{"location": "SF"}'
)
def test_multiple_tool_calls_indexed(self):
"""Multiple function_call items should be indexed correctly."""
otel = OpenTelemetry()
mock_span = MagicMock()
response_obj = {
"id": "resp_tc2",
"model": "gpt-4o",
"status": "completed",
"output": [
{
"type": "function_call",
"name": "get_weather",
"call_id": "call_1",
"arguments": "{}",
},
{
"type": "function_call",
"name": "get_time",
"call_id": "call_2",
"arguments": "{}",
},
],
}
otel.set_attributes(span=mock_span, kwargs=self._base_kwargs(), response_obj=response_obj)
mock_span.set_attribute.assert_any_call(
"gen_ai.completion.0.function_call.name", "get_weather"
)
mock_span.set_attribute.assert_any_call(
"gen_ai.completion.1.function_call.name", "get_time"
)

View file

@ -2135,6 +2135,53 @@ def test_validate_effort_for_model_centralises_per_model_gating(
assert err is None
def test_transform_request_injects_dummy_tool_without_tools_param():
"""
Anthropic rejects messages that contain tool turns when ``tools`` is omitted.
LiteLLM must inject a dummy tool without ``litellm.modify_params``.
"""
config = AnthropicConfig()
prev_modify_params = litellm.modify_params
litellm.modify_params = False
try:
messages = [
{"role": "user", "content": "Hello"},
{
"role": "assistant",
"content": "Calling tool",
"tool_calls": [
{
"id": "toolu_test_dummy",
"type": "function",
"function": {"name": "get_x", "arguments": "{}"},
}
],
},
{
"role": "tool",
"tool_call_id": "toolu_test_dummy",
"content": "{}",
},
]
result = config.transform_request(
model="claude-3-5-haiku-20241022",
messages=messages,
optional_params={"max_tokens": 256},
litellm_params={},
headers={},
)
finally:
litellm.modify_params = prev_modify_params
assert "tools" in result
names = [
t.get("name")
for t in result["tools"]
if isinstance(t, dict) and t.get("name") is not None
]
assert "dummy_tool" in names
def test_transform_request_uses_dynamic_max_tokens():
"""
Test that transform_request uses dynamic max_tokens based on model

View file

@ -499,3 +499,98 @@ class TestThinkingSummaryPreservation:
assert result == {
"reasoning_effort": {"effort": "medium", "summary": "concise"}
}
def test_anthropic_messages_handler_passes_litellm_params_timeout_to_base_handler():
"""The wrapper in this module must read `litellm_params.timeout` (set by
the router from the caller's `timeout` / `stream_timeout`) and forward it
to `BaseLLMHTTPHandler.anthropic_messages_handler` so the proxy honors
per-request timeouts on /v1/messages same as /chat/completions does."""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
anthropic_messages_handler,
base_llm_http_handler,
)
with patch.object(
base_llm_http_handler,
"anthropic_messages_handler",
return_value=MagicMock(),
) as mock_base:
try:
anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "hi"}],
model="anthropic/claude-sonnet-4-20250514",
api_key="test-key",
timeout=0.5,
)
except (ValueError, TypeError, AttributeError):
# downstream initialization may fail under mocks; we only care
# about what was passed to the base handler
pass
mock_base.assert_called_once()
assert mock_base.call_args.kwargs["timeout"] == 0.5
def test_anthropic_messages_handler_coerces_string_timeout_to_float():
"""`GenericLiteLLMParams.timeout` allows `str`, but the httpx-level
handlers expect `float | httpx.Timeout`. The wrapper must coerce string
inputs to float before forwarding (mirrors the pattern in
`litellm/main.py::_sleep_for_timeout`)."""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
anthropic_messages_handler,
base_llm_http_handler,
)
with patch.object(
base_llm_http_handler,
"anthropic_messages_handler",
return_value=MagicMock(),
) as mock_base:
try:
anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "hi"}],
model="anthropic/claude-sonnet-4-20250514",
api_key="test-key",
timeout="0.5",
)
except (ValueError, TypeError, AttributeError):
pass
mock_base.assert_called_once()
assert mock_base.call_args.kwargs["timeout"] == 0.5
assert isinstance(mock_base.call_args.kwargs["timeout"], float)
def test_anthropic_messages_handler_prefers_stream_timeout_for_streaming_calls():
"""When `stream=True` and `litellm_params.stream_timeout` is set, the
wrapper must forward `stream_timeout` (not `timeout`) to the base
handler mirrors `Router._get_timeout` so direct (non-router) callers
get the same stream/non-stream timeout selection as proxy callers."""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
anthropic_messages_handler,
base_llm_http_handler,
)
with patch.object(
base_llm_http_handler,
"anthropic_messages_handler",
return_value=MagicMock(),
) as mock_base:
try:
anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "hi"}],
model="anthropic/claude-sonnet-4-20250514",
api_key="test-key",
stream=True,
timeout=10.0,
stream_timeout=0.5,
)
except (ValueError, TypeError, AttributeError):
pass
mock_base.assert_called_once()
assert mock_base.call_args.kwargs["timeout"] == 0.5

View file

@ -0,0 +1,338 @@
"""Unit tests for ``BedrockBatchesHandler._handle_model_invocation_job_status``.
These cover the upstream support for retrieving Bedrock bulk batch jobs
(``arn:aws:bedrock:<region>:<acct>:model-invocation-job/<id>``) the ARN
type returned by ``CreateModelInvocationJob``. We mock the boto3 client so
the tests don't hit AWS.
"""
from __future__ import annotations
import os
import sys
from datetime import datetime, timezone
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.bedrock.batches.handler import ( # noqa: E402
BedrockBatchesHandler,
_extract_job_id_from_arn,
_extract_region_from_bedrock_arn,
_predict_output_file_uri,
_to_epoch,
)
JOB_ID = "abc1234567"
JOB_ARN = f"arn:aws:bedrock:us-west-2:123456789012:model-invocation-job/{JOB_ID}"
INPUT_URI = "s3://my-bucket/inputs/qwen3-235b-a22b-2507-batch.jsonl"
OUTPUT_PREFIX = "s3://my-bucket/litellm-batch-outputs/litellm-bedrock-files-qwen-uuid/"
SUBMIT_TIME = datetime(2026, 4, 28, 12, 0, 0, tzinfo=timezone.utc)
END_TIME = datetime(2026, 4, 28, 12, 30, 0, tzinfo=timezone.utc)
def _fake_boto3_response(status: str = "Completed", end_time=END_TIME):
return {
"jobArn": JOB_ARN,
"jobName": "litellm-bedrock-files-qwen-uuid",
"modelId": "bedrock/qwen.qwen3-235b-a22b-2507-v1:0",
"status": status,
"submitTime": SUBMIT_TIME,
"lastModifiedTime": end_time,
"endTime": end_time,
"inputDataConfig": {"s3InputDataConfig": {"s3Uri": INPUT_URI}},
"outputDataConfig": {"s3OutputDataConfig": {"s3Uri": OUTPUT_PREFIX}},
}
@pytest.fixture
def patched_boto3():
"""Yield a stub bedrock client whose `get_model_invocation_job` is a MagicMock."""
fake_client = MagicMock()
fake_client.get_model_invocation_job.return_value = _fake_boto3_response()
with (
patch("boto3.client", return_value=fake_client) as boto_client_factory,
patch(
"litellm.llms.bedrock.batches.transformation.BedrockBatchesConfig.get_credentials",
return_value=MagicMock(access_key="AKIA", secret_key="SECRET", token=None),
),
):
yield fake_client, boto_client_factory
def test_extract_region_from_arn():
assert _extract_region_from_bedrock_arn(JOB_ARN) == "us-west-2"
assert _extract_region_from_bedrock_arn("arn:aws:bedrock::123:foo/bar") is None
assert _extract_region_from_bedrock_arn("not-an-arn") is None
def test_extract_region_swallows_unexpected_split_errors():
"""Defensive `except Exception` branch — anything that isn't a plain str
should fall through to ``None`` rather than blow up."""
class WeirdArn:
def split(self, _sep):
raise RuntimeError("boom")
assert _extract_region_from_bedrock_arn(WeirdArn()) is None # type: ignore[arg-type]
def test_predict_output_file_uri_returns_none_for_directory_input_uri():
"""Input URI ending in `/` has an empty basename — we must bail rather
than emit ``<prefix>/<job-id>/.out``."""
assert (
_predict_output_file_uri(OUTPUT_PREFIX, "s3://bucket/inputs/", JOB_ID) is None
)
_DT = datetime(2026, 4, 28, 12, 0, 0, tzinfo=timezone.utc)
@pytest.mark.parametrize(
"value,expected",
[
(None, None),
(1730000000, 1730000000),
(1730000000.5, 1730000000),
(_DT, int(_DT.timestamp())),
("2026-04-28T12:00:00Z", None), # strings aren't supported -> None
],
)
def test_to_epoch_handles_supported_types(value, expected):
assert _to_epoch(value) == expected
def test_extract_job_id_from_arn():
assert _extract_job_id_from_arn(JOB_ARN) == JOB_ID
assert (
_extract_job_id_from_arn("arn:aws:bedrock:us-west-2:1:async-invoke/x") is None
)
def test_predict_output_file_uri_happy_path():
expected = f"{OUTPUT_PREFIX}{JOB_ID}/qwen3-235b-a22b-2507-batch.jsonl.out"
assert _predict_output_file_uri(OUTPUT_PREFIX, INPUT_URI, JOB_ID) == expected
def test_predict_output_file_uri_adds_trailing_slash():
prefix_no_slash = OUTPUT_PREFIX.rstrip("/")
expected = f"{OUTPUT_PREFIX}{JOB_ID}/qwen3-235b-a22b-2507-batch.jsonl.out"
assert _predict_output_file_uri(prefix_no_slash, INPUT_URI, JOB_ID) == expected
@pytest.mark.parametrize(
"missing_arg",
[
("", INPUT_URI, JOB_ID),
(OUTPUT_PREFIX, "", JOB_ID),
(OUTPUT_PREFIX, INPUT_URI, None),
],
)
def test_predict_output_file_uri_returns_none_when_missing_input(missing_arg):
assert _predict_output_file_uri(*missing_arg) is None
def test_handle_model_invocation_job_status_completed(patched_boto3):
fake_client, boto_client_factory = patched_boto3
batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN)
fake_client.get_model_invocation_job.assert_called_once_with(jobIdentifier=JOB_ARN)
# Region should be sniffed from the ARN.
_, kwargs = boto_client_factory.call_args
assert kwargs["region_name"] == "us-west-2"
assert batch.id == JOB_ARN
assert batch.status == "completed"
assert batch.input_file_id == INPUT_URI
expected_out = f"{OUTPUT_PREFIX}{JOB_ID}/qwen3-235b-a22b-2507-batch.jsonl.out"
assert batch.output_file_id == expected_out
assert batch.completed_at == int(END_TIME.timestamp())
assert batch.failed_at is None
assert batch.cancelled_at is None
# Per-record counts aren't reported by GetModelInvocationJob, so we leave
# them zeroed; consumers should parse manifest.json.out for accurate counts.
assert batch.request_counts.total == 0
assert batch.metadata["job_arn"] == JOB_ARN
assert batch.metadata["output_file_uri"] == expected_out
assert batch.metadata["output_s3_uri"] == OUTPUT_PREFIX
@pytest.mark.parametrize(
"bedrock_status,openai_status",
[
("Submitted", "validating"),
("Validating", "validating"),
("Scheduled", "validating"),
("InProgress", "in_progress"),
("Stopping", "cancelling"),
("Stopped", "cancelled"),
("Completed", "completed"),
("PartiallyCompleted", "completed"),
("Failed", "failed"),
("Expired", "expired"),
# Unknown/unmapped Bedrock status falls back to "in_progress" so we
# don't 500 on a future AWS-side enum addition.
("MyBrandNewStatus", "in_progress"),
],
)
def test_status_mapping(patched_boto3, bedrock_status, openai_status):
fake_client, _ = patched_boto3
fake_client.get_model_invocation_job.return_value = _fake_boto3_response(
status=bedrock_status
)
batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN)
assert batch.status == openai_status
# output_file_id is only populated for terminal-completed jobs, so callers
# don't accidentally try to download a non-existent file mid-run.
if openai_status == "completed":
assert batch.output_file_id is not None
else:
assert batch.output_file_id is None
def test_explicit_region_overrides_arn(patched_boto3):
_, boto_client_factory = patched_boto3
BedrockBatchesHandler._handle_model_invocation_job_status(
batch_id=JOB_ARN, aws_region_name="eu-central-1"
)
_, kwargs = boto_client_factory.call_args
assert kwargs["region_name"] == "eu-central-1"
def test_failure_message_propagates(patched_boto3):
fake_client, _ = patched_boto3
failed_response = _fake_boto3_response(status="Failed")
failed_response["message"] = "Input file failed validation"
fake_client.get_model_invocation_job.return_value = failed_response
batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN)
assert batch.status == "failed"
assert batch.failed_at == int(END_TIME.timestamp())
assert batch.metadata["failure_message"] == "Input file failed validation"
def test_completed_with_unpredictable_output_uri_stays_none(patched_boto3):
"""
Regression guard for the original NoSuchKey bug: if Bedrock's response is
missing pieces we need to compute the per-job output file path (here, the
input s3Uri), `output_file_id` must stay `None` rather than fall back to
the bare prefix. Falling back to the prefix is what produced the original
NoSuchKey error this PR fixes.
"""
fake_client, _ = patched_boto3
incomplete_response = _fake_boto3_response(status="Completed")
incomplete_response["inputDataConfig"] = {"s3InputDataConfig": {"s3Uri": ""}}
fake_client.get_model_invocation_job.return_value = incomplete_response
batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN)
assert batch.status == "completed"
# output_file_id MUST be None (not the bare prefix) — that's the whole
# point of this regression test. Callers branch on this field.
assert batch.output_file_id is None
# The metadata field uses "" because OpenAI Batch metadata is dict[str, str];
# callers should branch on `output_file_id` (above) instead.
assert batch.metadata["output_file_uri"] == ""
# The bare prefix is still preserved in metadata so callers can list it.
assert batch.metadata["output_s3_uri"] == OUTPUT_PREFIX
def test_cancelled_status_sets_cancelled_at(patched_boto3):
fake_client, _ = patched_boto3
fake_client.get_model_invocation_job.return_value = _fake_boto3_response(
status="Stopped"
)
batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN)
assert batch.status == "cancelled"
assert batch.cancelled_at == int(END_TIME.timestamp())
assert batch.completed_at is None
assert batch.failed_at is None
assert batch.expired_at is None
def test_expired_status_sets_expired_at(patched_boto3):
fake_client, _ = patched_boto3
fake_client.get_model_invocation_job.return_value = _fake_boto3_response(
status="Expired"
)
batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN)
assert batch.status == "expired"
assert batch.expired_at == int(END_TIME.timestamp())
assert batch.completed_at is None
assert batch.failed_at is None
assert batch.cancelled_at is None
def test_logging_obj_pre_and_post_call_invoked(patched_boto3):
"""`pre_call` / `post_call` get called with sensible payloads when a
`logging_obj` is supplied."""
_, _ = patched_boto3
logging_obj = MagicMock()
BedrockBatchesHandler._handle_model_invocation_job_status(
batch_id=JOB_ARN, logging_obj=logging_obj
)
logging_obj.pre_call.assert_called_once()
logging_obj.post_call.assert_called_once()
pre_kwargs = logging_obj.pre_call.call_args.kwargs
assert pre_kwargs["input"] == JOB_ARN
assert pre_kwargs["additional_args"]["complete_input_dict"] == {
"jobIdentifier": JOB_ARN
}
# Logged URL must use the bare job id, not the full ARN, so it doesn't
# double the `model-invocation-job/` segment or embed colons in the path.
assert pre_kwargs["additional_args"]["api_base"] == (
f"https://bedrock.us-west-2.amazonaws.com/model-invocation-job/{JOB_ID}"
)
post_kwargs = logging_obj.post_call.call_args.kwargs
assert post_kwargs["input"] == JOB_ARN
assert post_kwargs["original_response"]["jobArn"] == JOB_ARN
def test_missing_boto3_raises_helpful_import_error():
"""If boto3 isn't installed we should raise a clear, actionable
ImportError rather than letting a NameError escape."""
real_import = (
__builtins__["__import__"]
if isinstance(__builtins__, dict)
else __builtins__.__import__
)
def fake_import(name, *args, **kwargs):
if name == "boto3":
raise ImportError("No module named 'boto3'")
return real_import(name, *args, **kwargs)
with patch("builtins.__import__", side_effect=fake_import):
with pytest.raises(ImportError, match="pip install boto3"):
BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN)
def test_logging_url_uses_bare_id_when_only_id_passed(patched_boto3):
"""If the caller passes just the trailing job id (also valid for
`GetModelInvocationJob`), the logged URL should use it as-is."""
_, _ = patched_boto3
logging_obj = MagicMock()
BedrockBatchesHandler._handle_model_invocation_job_status(
batch_id=JOB_ID, aws_region_name="us-west-2", logging_obj=logging_obj
)
pre_kwargs = logging_obj.pre_call.call_args.kwargs
assert pre_kwargs["additional_args"]["api_base"] == (
f"https://bedrock.us-west-2.amazonaws.com/model-invocation-job/{JOB_ID}"
)

View file

@ -867,10 +867,12 @@ def test_bedrock_messages_explicit_output_config_wins_over_reasoning_effort():
def test_bedrock_messages_strips_context_management():
"""
Ensure context_management is stripped from the request before sending to
Bedrock Invoke, which doesn't support this Anthropic-specific parameter.
Bedrock Invoke when it carries only LiteLLM-internal edits (e.g.
clear_thinking_20251015, which is consumed via thinking injection).
Claude Code sends context_management on every request; leaving it in the body
causes a 400 "context_management: Extra inputs are not permitted" from Bedrock.
Claude Code sends context_management on every request; leaving such edits
in the body causes a 400 "context_management: Extra inputs are not
permitted" from Bedrock.
"""
from litellm.types.router import GenericLiteLLMParams
@ -897,6 +899,77 @@ def test_bedrock_messages_strips_context_management():
assert result.get("max_tokens") == 4096
def test_bedrock_messages_preserves_compact_context_management_and_adds_beta():
"""
Bedrock InvokeModel supports compaction when paired with the
``compact-2026-01-12`` anthropic-beta header, even though the Converse API
does not. The transformation should:
1. Keep ``context_management`` with compact_20260112 edits in the body
(Bedrock rejects unknown top-level fields, but accepts this one with
the right beta).
2. Auto-inject ``compact-2026-01-12`` into ``anthropic_beta``.
Ref: https://github.com/BerriAI/litellm/issues/27532
"""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hi"}]}]
optional_params = {
"max_tokens": 4096,
"context_management": {
"edits": [{"type": "compact_20260112"}]
},
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-sonnet-4-6-20250929-v1:0",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert result.get("context_management") == {
"edits": [{"type": "compact_20260112"}]
}
assert "compact-2026-01-12" in result.get("anthropic_beta", [])
assert result["max_tokens"] == 4096
def test_bedrock_messages_filters_unsupported_context_management_edits():
"""
Mixed edit lists must drop the LiteLLM-internal ``clear_thinking_20251015``
entries while keeping ``compact_20260112`` and adding the compact beta.
"""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hi"}]}]
optional_params = {
"max_tokens": 4096,
"context_management": {
"edits": [
{"type": "clear_thinking_20251015", "keep": "all"},
{"type": "compact_20260112"},
]
},
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-sonnet-4-6-20250929-v1:0",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert result.get("context_management") == {
"edits": [{"type": "compact_20260112"}]
}
assert "compact-2026-01-12" in result.get("anthropic_beta", [])
def test_bedrock_messages_allowlist_filters_anthropic_only_fields():
"""
Bedrock Invoke rejects any top-level body field it doesn't recognize with

View file

@ -0,0 +1,312 @@
import json
from unittest.mock import patch
import httpx
import pytest
def _anthropic_response(url: str) -> httpx.Response:
return httpx.Response(
status_code=200,
json={
"id": "msg_test",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-6",
"content": [{"type": "text", "text": "ok"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 1, "output_tokens": 1},
},
request=httpx.Request("POST", url),
)
def _capture_request(url: str, headers: dict, data: bytes | str | None) -> dict:
raw_body = data.decode("utf-8") if isinstance(data, bytes) else data or "{}"
return {
"path": httpx.URL(url).path,
"headers": headers,
"body": json.loads(raw_body),
}
def test_claude_platform_builds_default_messages_url_from_region():
from litellm.llms.bedrock.claude_platform.transformation import (
BedrockClaudePlatformConfig,
)
config = BedrockClaudePlatformConfig()
assert (
config.get_complete_url(
api_base=None,
api_key=None,
model="claude-sonnet-4-6",
optional_params={"aws_region_name": "us-west-2"},
litellm_params={},
)
== "https://aws-external-anthropic.us-west-2.api.aws/v1/messages"
)
def test_claude_platform_ignores_standard_anthropic_base_url(monkeypatch):
from litellm.llms.bedrock.claude_platform.transformation import (
BedrockClaudePlatformConfig,
)
monkeypatch.setenv("ANTHROPIC_BASE_URL", "https://api.anthropic.example")
monkeypatch.setenv("ANTHROPIC_API_BASE", "https://api.anthropic-api.example")
config = BedrockClaudePlatformConfig()
assert (
config.get_complete_url(
api_base=None,
api_key=None,
model="claude-sonnet-4-6",
optional_params={"aws_region_name": "us-west-2"},
litellm_params={},
)
== "https://aws-external-anthropic.us-west-2.api.aws/v1/messages"
)
def test_claude_platform_uses_bedrock_subroute():
import litellm
from litellm.llms.bedrock.common_utils import BedrockModelInfo
model, provider, _, _ = litellm.get_llm_provider(
model="bedrock/claude_platform/claude-sonnet-4-6"
)
assert provider == "bedrock"
assert model == "claude_platform/claude-sonnet-4-6"
assert BedrockModelInfo.get_bedrock_route(model) == "claude_platform"
assert BedrockModelInfo.get_claude_platform_model(model) == "claude-sonnet-4-6"
def test_claude_platform_requires_workspace_header():
from litellm import AuthenticationError
from litellm.llms.bedrock.claude_platform.transformation import (
BedrockClaudePlatformConfig,
)
config = BedrockClaudePlatformConfig()
with pytest.raises(AuthenticationError) as exc_info:
config.validate_environment(
api_key="fake-platform-key",
headers={},
model="claude-sonnet-4-6",
messages=[{"role": "user", "content": "hello"}],
optional_params={},
litellm_params={},
)
assert "workspace" in str(exc_info.value).lower()
def test_claude_platform_api_key_auth_sets_workspace_and_key_headers():
from litellm.llms.bedrock.claude_platform.transformation import (
BedrockClaudePlatformConfig,
)
config = BedrockClaudePlatformConfig()
headers = config.validate_environment(
api_key="fake-platform-key",
headers={"anthropic-beta": "skills-2025-10-02"},
model="claude-sonnet-4-6",
messages=[{"role": "user", "content": "hello"}],
optional_params={"workspace_id": "wrkspc_test"},
litellm_params={},
)
assert headers["x-api-key"] == "fake-platform-key"
assert headers["anthropic-workspace-id"] == "wrkspc_test"
assert headers["anthropic-beta"] == "skills-2025-10-02"
def test_claude_platform_does_not_use_standard_anthropic_api_key(monkeypatch):
from litellm.llms.bedrock.claude_platform.transformation import (
BedrockClaudePlatformConfig,
)
monkeypatch.setenv("ANTHROPIC_API_KEY", "standard-anthropic-key")
config = BedrockClaudePlatformConfig()
headers = config.validate_environment(
api_key=None,
headers={},
model="claude-sonnet-4-6",
messages=[{"role": "user", "content": "hello"}],
optional_params={"workspace_id": "wrkspc_test"},
litellm_params={},
)
assert "x-api-key" not in headers
def test_claude_platform_sigv4_signs_transformed_request_body():
from litellm.llms.bedrock.claude_platform.transformation import (
BedrockClaudePlatformConfig,
)
config = BedrockClaudePlatformConfig()
request_body = {
"model": "claude-sonnet-4-6",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 10,
}
with patch.object(
config,
"_sign_request",
return_value=({"Authorization": "signed"}, json.dumps(request_body).encode()),
) as mock_sign_request:
headers, signed_body = config.sign_request(
headers={"anthropic-workspace-id": "wrkspc_test"},
optional_params={"aws_region_name": "us-west-2"},
request_data=request_body,
api_base="https://aws-external-anthropic.us-west-2.api.aws/v1/messages",
api_key=None,
model="claude-sonnet-4-6",
)
assert signed_body == json.dumps(request_body).encode()
assert headers["Authorization"] == "signed"
mock_sign_request.assert_called_once()
assert (
mock_sign_request.call_args.kwargs["service_name"] == "aws-external-anthropic"
)
assert mock_sign_request.call_args.kwargs["request_data"] == request_body
def test_claude_platform_standard_anthropic_api_key_does_not_skip_sigv4(monkeypatch):
from litellm.llms.bedrock.claude_platform.transformation import (
BedrockClaudePlatformConfig,
)
monkeypatch.setenv("ANTHROPIC_API_KEY", "standard-anthropic-key")
config = BedrockClaudePlatformConfig()
request_body = {
"model": "claude-sonnet-4-6",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 10,
}
with patch.object(
config,
"_sign_request",
return_value=({"Authorization": "signed"}, json.dumps(request_body).encode()),
) as mock_sign_request:
headers, signed_body = config.sign_request(
headers={"anthropic-workspace-id": "wrkspc_test"},
optional_params={"aws_region_name": "us-west-2"},
request_data=request_body,
api_base="https://aws-external-anthropic.us-west-2.api.aws/v1/messages",
api_key=None,
model="claude-sonnet-4-6",
)
assert signed_body == json.dumps(request_body).encode()
assert headers["Authorization"] == "signed"
mock_sign_request.assert_called_once()
def test_bedrock_claude_platform_messages_config_round_trips_native_body():
import litellm
from litellm.types.utils import LlmProviders
config = litellm.ProviderConfigManager.get_provider_anthropic_messages_config(
model="claude_platform/claude-sonnet-4-6",
provider=LlmProviders.BEDROCK,
)
assert config is not None
headers, _ = config.validate_anthropic_messages_environment(
api_key="fake-platform-key",
headers={},
model="claude_platform/claude-sonnet-4-6",
messages=[{"role": "user", "content": "hello"}],
optional_params={"max_tokens": 10},
litellm_params={"workspace_id": "wrkspc_test"},
)
request_body = config.transform_anthropic_messages_request(
model="claude_platform/claude-sonnet-4-6",
messages=[{"role": "user", "content": "hello"}],
anthropic_messages_optional_request_params={"max_tokens": 10},
litellm_params={},
headers=headers,
)
assert headers["anthropic-workspace-id"] == "wrkspc_test"
assert headers["x-api-key"] == "fake-platform-key"
assert request_body == {
"model": "claude-sonnet-4-6",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 10,
}
def test_chat_completion_routes_bedrock_claude_platform_to_messages_api():
import litellm
requests = []
def mock_post(self, url, data=None, headers=None, **kwargs):
requests.append(_capture_request(url=url, headers=headers or {}, data=data))
return _anthropic_response(url)
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post", mock_post):
response = litellm.completion(
model="bedrock/claude_platform/claude-sonnet-4-6",
messages=[{"role": "user", "content": "hello"}],
max_tokens=10,
api_base="https://aws-external-anthropic.us-west-2.api.aws",
api_key="fake-platform-key",
workspace_id="wrkspc_test",
)
assert response.choices[0].message.content == "ok"
assert len(requests) == 1
assert requests[0]["path"] == "/v1/messages"
assert requests[0]["headers"]["x-api-key"] == "fake-platform-key"
assert requests[0]["headers"]["anthropic-workspace-id"] == "wrkspc_test"
assert requests[0]["body"]["model"] == "claude-sonnet-4-6"
@pytest.mark.asyncio
async def test_anthropic_messages_routes_bedrock_claude_platform_to_messages_api():
import litellm
requests = []
async def mock_post(self, url, data=None, headers=None, **kwargs):
requests.append(_capture_request(url=url, headers=headers or {}, data=data))
return _anthropic_response(url)
try:
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=mock_post,
):
response = await litellm.anthropic_messages(
model="bedrock/claude_platform/claude-sonnet-4-6",
messages=[{"role": "user", "content": "hello"}],
max_tokens=10,
api_base="https://aws-external-anthropic.us-west-2.api.aws",
api_key="fake-platform-key",
workspace_id="wrkspc_test",
)
finally:
await litellm.close_litellm_async_clients()
assert response["content"][0]["text"] == "ok"
assert len(requests) == 1
assert requests[0]["path"] == "/v1/messages"
assert requests[0]["headers"]["x-api-key"] == "fake-platform-key"
assert requests[0]["headers"]["anthropic-workspace-id"] == "wrkspc_test"
assert requests[0]["body"]["messages"] == [{"role": "user", "content": "hello"}]
assert requests[0]["body"]["max_tokens"] == 10
assert requests[0]["body"]["model"] == "claude-sonnet-4-6"

View file

@ -422,3 +422,140 @@ def test_sync_delete_responses_omits_body_for_azure():
assert captured["url"].endswith(
"/openai/responses/resp_xyz?api-version=2025-03-01-preview"
)
# Per-request `timeout` is plumbed through the /v1/messages chain
# (async handler -> retry helper -> httpx post), matching /chat/completions.
@pytest.mark.parametrize("stream", [False, True])
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_forwards_timeout_to_retry_helper(
stream,
):
"""`async_anthropic_messages_handler` must forward `timeout` to
`_async_post_anthropic_messages_with_http_error_retry` for both
streaming and non-streaming calls."""
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com")
)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude-sonnet-4-20250514", "messages": []}
)
mock_config.get_complete_url = Mock(
return_value="https://api.anthropic.com/v1/messages"
)
mock_config.sign_request = Mock(return_value=({}, b"{}"))
mock_logging_obj = Mock()
mock_logging_obj.update_environment_variables = Mock()
mock_logging_obj.update_from_kwargs = Mock()
mock_logging_obj.pre_call = Mock()
mock_logging_obj.model_call_details = {}
mock_logging_obj.stream = False
with patch.object(
handler,
"_async_post_anthropic_messages_with_http_error_retry",
new_callable=AsyncMock,
) as mock_retry, patch(
"litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers",
return_value=None,
):
try:
await handler.async_anthropic_messages_handler(
model="claude-sonnet-4-20250514",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_provider_config=mock_config,
anthropic_messages_optional_request_params={},
custom_llm_provider="anthropic",
litellm_params=GenericLiteLLMParams(),
logging_obj=mock_logging_obj,
stream=stream,
timeout=0.5,
kwargs={},
)
except Exception:
# Downstream response transformation may fail under mocks; we
# only care about what reached the retry helper.
pass
mock_retry.assert_awaited_once()
assert mock_retry.await_args.kwargs["timeout"] == 0.5
assert mock_retry.await_args.kwargs["stream"] is stream
@pytest.mark.parametrize("stream", [False, True])
@pytest.mark.asyncio
async def test_async_post_anthropic_messages_forwards_timeout_to_httpx_post(stream):
"""`_async_post_anthropic_messages_with_http_error_retry` must forward
`timeout` to `async_httpx_client.post` for both streaming and
non-streaming requests (router resolves `stream_timeout` vs `timeout`
upstream into a single value)."""
handler = BaseLLMHTTPHandler()
mock_response = Mock(status_code=200)
mock_response.raise_for_status = Mock(return_value=None)
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_response)
mock_provider_config = Mock()
mock_provider_config.max_retry_on_anthropic_messages_http_error = 1
await handler._async_post_anthropic_messages_with_http_error_retry(
async_httpx_client=mock_client,
request_url="https://api.anthropic.com/v1/messages",
headers={},
signed_json_body=None,
request_body={"model": "claude-sonnet-4-20250514", "messages": []},
stream=stream,
logging_obj=Mock(),
provider_config=mock_provider_config,
litellm_params=GenericLiteLLMParams(),
api_key=None,
model="claude-sonnet-4-20250514",
timeout=0.5,
)
mock_client.post.assert_awaited_once()
assert mock_client.post.await_args.kwargs["timeout"] == 0.5
assert mock_client.post.await_args.kwargs["stream"] is stream
@pytest.mark.asyncio
async def test_async_post_anthropic_messages_defaults_timeout_to_none_when_unset():
"""When no `timeout` is supplied, the retry helper must call
`async_httpx_client.post(timeout=None)` so AsyncHTTPHandler can fall
back to its own default preserves pre-patch behavior for callers
that do not opt into a per-request timeout."""
handler = BaseLLMHTTPHandler()
mock_response = Mock(status_code=200)
mock_response.raise_for_status = Mock(return_value=None)
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_response)
mock_provider_config = Mock()
mock_provider_config.max_retry_on_anthropic_messages_http_error = 1
await handler._async_post_anthropic_messages_with_http_error_retry(
async_httpx_client=mock_client,
request_url="https://api.anthropic.com/v1/messages",
headers={},
signed_json_body=None,
request_body={"model": "claude-sonnet-4-20250514", "messages": []},
stream=False,
logging_obj=Mock(),
provider_config=mock_provider_config,
litellm_params=GenericLiteLLMParams(),
api_key=None,
model="claude-sonnet-4-20250514",
)
mock_client.post.assert_awaited_once()
assert mock_client.post.await_args.kwargs["timeout"] is None

View file

@ -94,27 +94,33 @@ def test_github_copilot_config_get_openai_compatible_provider_info():
@patch("litellm.llms.github_copilot.authenticator.Authenticator.get_api_key")
@patch("litellm.main.openai_chat_completions.completion")
@patch("litellm.llms.openai.openai.OpenAIChatCompletion.completion")
def test_completion_github_copilot_mock_response(mock_completion, mock_get_api_key):
def test_completion_github_copilot_mock_response(
mock_class_completion, mock_instance_completion, mock_get_api_key, monkeypatch
):
"""Test the completion function with GitHub Copilot provider."""
# Mock the API key return value
# Force chat path through the patched openai_chat_completions instance even if
# a previous test left EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER set in the env.
monkeypatch.delenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", raising=False)
mock_api_key = "gh.test-key-123456789"
mock_get_api_key.return_value = mock_api_key
# Mock completion response
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
mock_response.choices[0].message.content = "Hello, I'm GitHub Copilot!"
mock_completion.return_value = mock_response
# Patch both the class method and the live module-level instance to survive
# conftest module reloads that can swap which class object is in use.
mock_class_completion.return_value = mock_response
mock_instance_completion.return_value = mock_response
# Test non-streaming completion
messages = [
{"role": "system", "content": "You're GitHub Copilot, an AI assistant."},
{"role": "user", "content": "Hello, who are you?"},
]
# Create a properly formatted headers dictionary
headers = {
"editor-version": "Neovim/0.9.0",
"Copilot-Integration-Id": "vscode-chat",
@ -128,19 +134,16 @@ def test_completion_github_copilot_mock_response(mock_completion, mock_get_api_k
assert response is not None
# Verify the get_api_key call was made (can be called multiple times)
assert mock_get_api_key.call_count >= 1
# Verify the completion call was made with the expected params
mock_completion.assert_called_once()
args, kwargs = mock_completion.call_args
# Exactly one of the two patched targets should have been used.
invoked = [m for m in (mock_class_completion, mock_instance_completion) if m.called]
assert len(invoked) == 1
invoked[0].assert_called_once()
_, kwargs = invoked[0].call_args
# Check that the proper authorization header is set
assert "headers" in kwargs
# Check that the model name is correctly formatted
assert (
kwargs.get("model") == "gpt-4"
) # Model name should be without provider prefix
assert kwargs.get("model") == "gpt-4"
assert kwargs.get("messages") == messages

View file

@ -1,10 +1,15 @@
import pytest
import litellm
import litellm.main as litellm_main
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
from litellm.llms.openai.openai import OpenAIConfig
from litellm.utils import _is_explicitly_disabled_factory
from litellm.utils import (
_is_explicitly_disabled_factory,
peek_reasoning_summary_aliases,
strip_reasoning_summary_aliases_from_optional_params,
)
@pytest.fixture()
@ -1007,6 +1012,76 @@ def test_gpt5_search_drops_unsupported_params(config: OpenAIConfig):
assert "tools" not in params
def test_gpt5_chat_strips_reasoning_summary_aliases_after_bridge_check(
monkeypatch: pytest.MonkeyPatch,
):
"""Non-bridged GPT-5 chat calls strip Responses-only reasoning summary aliases."""
captured_kwargs = {}
def fake_openai_completion(**kwargs):
captured_kwargs.update(kwargs)
return {}
monkeypatch.setattr(
litellm_main.openai_chat_completions,
"completion",
fake_openai_completion,
)
litellm.completion(
model="gpt-5",
messages=[{"role": "user", "content": "ok"}],
reasoningSummary="auto",
extra_body={"reasoning_summary": "ignored", "metadata": "ok"},
api_key="fake-key",
)
optional_params = captured_kwargs["optional_params"]
assert "reasoningSummary" not in optional_params
assert "reasoning_summary" not in optional_params
assert optional_params["extra_body"] == {"metadata": "ok"}
def test_reasoning_summary_alias_helpers_preserve_falsy_and_strip_all_aliases():
optional_params = {"reasoningSummary": False, "reasoning_summary": "ignored"}
assert peek_reasoning_summary_aliases(optional_params) is False
stripped, rs_val = strip_reasoning_summary_aliases_from_optional_params(
optional_params
)
assert rs_val is False
assert stripped == {}
optional_params = {
"extra_body": {"reasoningSummary": False, "reasoning_summary": "ignored"}
}
assert peek_reasoning_summary_aliases(optional_params) is False
stripped, rs_val = strip_reasoning_summary_aliases_from_optional_params(
optional_params
)
assert rs_val is False
assert stripped == {}
optional_params = {
"extra_body": {
"reasoningSummary": "auto",
"reasoning_summary": "ignored",
"metadata": "ok",
}
}
assert peek_reasoning_summary_aliases(optional_params) == "auto"
stripped, rs_val = strip_reasoning_summary_aliases_from_optional_params(
optional_params
)
assert rs_val == "auto"
assert stripped == {"extra_body": {"metadata": "ok"}}
# GPT-5 unsupported params audit (validated via direct API calls)
def test_gpt5_rejects_params_unsupported_by_openai(config: OpenAIConfig):
"""Params that OpenAI rejects for all GPT-5 reasoning models."""

View file

@ -54,3 +54,61 @@ def test_ovhcloud_audio_transcription_config_installed():
assert config is not None
assert isinstance(config, BaseAudioTranscriptionConfig)
class TestOVHCloudDurationFieldMigration:
"""Tests for OVHCloud duration -> seconds field migration."""
def test_seconds_field_mapped_to_duration(self):
"""New `seconds` field should be normalized to `duration`."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "Hello world",
"seconds": 3.14,
}
result = config.transform_audio_transcription_response(mock_response)
assert result.text == "Hello world"
assert result._hidden_params["duration"] == 3.14
def test_legacy_duration_field_still_works(self):
"""Legacy `duration` field should still be accepted."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "Hello world",
"duration": 2.71,
}
result = config.transform_audio_transcription_response(mock_response)
assert result.text == "Hello world"
assert result._hidden_params["duration"] == 2.71
def test_seconds_zero_mapped_to_duration(self):
"""seconds=0.0 must not be treated as falsy and lost."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {"text": "silence", "seconds": 0.0}
result = config.transform_audio_transcription_response(mock_response)
assert result._hidden_params["duration"] == 0.0

View file

@ -292,3 +292,78 @@ def test_ovhcloud_with_custom_base_url():
if __name__ == "__main__":
pytest.main([__file__, "-v"])
class TestOVHCloudReasoningFieldMigration:
"""Tests for OVHCloud reasoning_content -> reasoning field migration."""
def test_streaming_new_reasoning_field(self):
"""New `reasoning` field should be mapped to `reasoning_content`."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"role": "assistant",
"reasoning": "Let me think...",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "Let me think..."
def test_streaming_legacy_reasoning_content_unchanged(self):
"""Legacy `reasoning_content` field should pass through untouched."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"role": "assistant",
"reasoning_content": "Already correct field.",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "Already correct field."
def test_streaming_both_fields_legacy_wins(self):
"""When both fields present, existing `reasoning_content` is not overwritten."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"reasoning": "new field",
"reasoning_content": "legacy field",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "legacy field"

View file

@ -1135,3 +1135,69 @@ def test_validate_loopback_redirect_uri_rejects_malformed_cleanly():
with pytest.raises(HTTPException) as exc:
validate_loopback_redirect_uri("http://[not-an-ip]/cb")
assert exc.value.status_code == 400
def _mock_request_with_base_url(base_url: str):
req = MagicMock()
req.base_url = base_url
req.headers = {}
return req
def test_validate_trusted_redirect_uri_accepts_same_origin():
"""UI OAuth flow: redirect_uri on the proxy's own origin is allowed."""
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
req = _mock_request_with_base_url("https://proxy.example.com/")
# Should not raise.
validate_trusted_redirect_uri(
req, "https://proxy.example.com/ui/mcp/oauth/callback"
)
def test_validate_trusted_redirect_uri_accepts_loopback():
"""Native MCP client flow: loopback is still allowed."""
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
req = _mock_request_with_base_url("https://proxy.example.com/")
validate_trusted_redirect_uri(req, "http://127.0.0.1:3000/cb")
validate_trusted_redirect_uri(req, "http://localhost:3000/cb")
def test_validate_trusted_redirect_uri_rejects_external_origin():
"""An attacker-controlled origin must still be rejected."""
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
req = _mock_request_with_base_url("https://proxy.example.com/")
with pytest.raises(HTTPException) as exc:
validate_trusted_redirect_uri(req, "https://attacker.example.com/cb")
assert exc.value.status_code == 400
def test_validate_trusted_redirect_uri_rejects_scheme_mismatch():
"""https→http (or vice versa) on the same host is not same-origin."""
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
req = _mock_request_with_base_url("https://proxy.example.com/")
with pytest.raises(HTTPException) as exc:
validate_trusted_redirect_uri(req, "http://proxy.example.com/ui/callback")
assert exc.value.status_code == 400
def test_validate_trusted_redirect_uri_rejects_fragment():
from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
)
req = _mock_request_with_base_url("https://proxy.example.com/")
with pytest.raises(HTTPException) as exc:
validate_trusted_redirect_uri(req, "https://proxy.example.com/ui/cb#code=1")
assert exc.value.status_code == 400

View file

@ -23,6 +23,20 @@ def mock_mcp_client_ip():
yield
def _mock_callback_request(base_url: str = "http://localhost:3000/"):
"""Return a MagicMock Request for callback/authorize same-origin tests.
The callback handler only uses ``request`` to compute the proxy's own
base URL via ``get_request_base_url`` (which reads ``request.base_url``
and trusted ``X-Forwarded-*`` headers). A simple MagicMock with the
right attributes is sufficient.
"""
req = MagicMock()
req.base_url = base_url
req.headers = {}
return req
@pytest.fixture
def trust_xff():
"""Force ``IPAddressUtils.is_request_from_trusted_proxy`` to True.
@ -1844,6 +1858,7 @@ async def test_oauth_callback_redirects_with_state():
# Call callback endpoint with code and state
response = await callback(
request=_mock_callback_request(),
code="test_authorization_code_12345",
state="encrypted_state_value",
)
@ -1887,6 +1902,7 @@ async def test_oauth_callback_preserves_client_redirect_uri_query():
}
response = await callback(
request=_mock_callback_request(),
code="test_authorization_code_12345",
state="encrypted_state_value",
)
@ -1917,6 +1933,7 @@ async def test_oauth_callback_handles_invalid_state():
# Call callback endpoint with invalid state
response = await callback(
request=_mock_callback_request(),
code="test_code",
state="invalid_encrypted_state",
)
@ -1926,6 +1943,40 @@ async def test_oauth_callback_handles_invalid_state():
assert "Authentication incomplete" in response.body.decode()
@pytest.mark.asyncio
async def test_oauth_callback_accepts_same_origin_ui_redirect():
"""UI OAuth flow: the callback should redirect to the proxy's own UI
origin when the encrypted state carries a same-origin client_redirect_uri."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
callback,
)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
) as mock_decode:
mock_decode.return_value = {
"base_url": "https://proxy.example.com/ui/mcp/oauth/callback",
"original_state": "state-123",
"code_challenge": None,
"code_challenge_method": None,
"client_redirect_uri": "https://proxy.example.com/ui/mcp/oauth/callback",
}
response = await callback(
request=_mock_callback_request(base_url="https://proxy.example.com/"),
code="auth-code-123",
state="encrypted_state",
)
assert response.status_code == 302
assert (
"https://proxy.example.com/ui/mcp/oauth/callback"
in response.headers["location"]
)
assert "code=auth-code-123" in response.headers["location"]
assert "state=state-123" in response.headers["location"]
@pytest.mark.asyncio
async def test_oauth_authorize_includes_scopes_from_server_config():
"""Test that authorize endpoint includes scopes from server configuration."""
@ -2307,7 +2358,11 @@ async def test_callback_revalidates_loopback_on_decoded_base_url():
"client_redirect_uri": "https://attacker.example.com/cb",
}
with pytest.raises(HTTPException) as exc_info:
await callback(code="stolen_code", state="encrypted_stale_state")
await callback(
request=_mock_callback_request(),
code="stolen_code",
state="encrypted_stale_state",
)
assert exc_info.value.status_code == 400
@ -2329,7 +2384,11 @@ async def test_callback_revalidates_loopback_on_decoded_client_redirect_uri():
"client_redirect_uri": "https://attacker.example.com/cb",
}
with pytest.raises(HTTPException) as exc_info:
await callback(code="stolen_code", state="encrypted_stale_state")
await callback(
request=_mock_callback_request(),
code="stolen_code",
state="encrypted_stale_state",
)
assert exc_info.value.status_code == 400
@ -2349,7 +2408,11 @@ async def test_callback_rejects_state_missing_redirect_uri():
"code_challenge_method": None,
}
with pytest.raises(HTTPException) as exc_info:
await callback(code="code", state="encrypted_malformed_state")
await callback(
request=_mock_callback_request(),
code="code",
state="encrypted_malformed_state",
)
assert exc_info.value.status_code == 400

View file

@ -1,5 +1,5 @@
from typing import Optional
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -494,6 +494,80 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes():
mock_cache.async_set_cache.assert_not_called()
def test_get_all_jwt_team_ids_unions_singular_and_plural():
"""get_all_jwt_team_ids must include the singular team_id_jwt_field claim
in addition to the plural team_ids_jwt_field, deduplicated."""
jwt_handler = JWTHandler()
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=MagicMock(),
litellm_jwtauth=LiteLLM_JWTAuth(
team_id_jwt_field="team_id",
team_ids_jwt_field="teams",
),
)
# singular only — Okta/Auth0 default shape
assert jwt_handler.get_all_jwt_team_ids({"team_id": "team-low"}) == ["team-low"]
# plural only — pre-fix shape
assert jwt_handler.get_all_jwt_team_ids({"teams": ["a", "b"]}) == ["a", "b"]
# both populated, no overlap
assert jwt_handler.get_all_jwt_team_ids(
{"team_id": "primary", "teams": ["a", "b"]}
) == ["a", "b", "primary"]
# both populated with overlap — singular dedup'd
assert jwt_handler.get_all_jwt_team_ids({"team_id": "a", "teams": ["a", "b"]}) == [
"a",
"b",
]
# singular field as multi-element list (some IdPs) — merge all, preserve plural-first order
assert jwt_handler.get_all_jwt_team_ids(
{"team_id": ["primary", "secondary"], "teams": ["a"]}
) == ["a", "primary", "secondary"]
# neither populated
assert jwt_handler.get_all_jwt_team_ids({}) == []
def test_get_all_jwt_team_ids_does_not_use_team_id_default():
"""team_id_default is a JWT-bearer-flow auth-builder fallback, not a token
claim. It must NOT leak into get_all_jwt_team_ids otherwise SSO logins
would silently start adding users to the default team for any tenant that
has team_id_default configured."""
jwt_handler = JWTHandler()
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=MagicMock(),
litellm_jwtauth=LiteLLM_JWTAuth(
team_id_jwt_field="team_id",
team_ids_jwt_field="teams",
team_id_default="default-team",
),
)
# team_id claim missing — must not fall back to default-team
assert jwt_handler.get_all_jwt_team_ids({"teams": []}) == []
assert jwt_handler.get_all_jwt_team_ids({}) == []
# only the plural is populated — default still must not be added
assert jwt_handler.get_all_jwt_team_ids({"teams": ["a"]}) == ["a"]
# team_id_jwt_field unset entirely + only default configured: still no default
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=MagicMock(),
litellm_jwtauth=LiteLLM_JWTAuth(
team_ids_jwt_field="teams",
team_id_default="default-team",
),
)
assert jwt_handler.get_all_jwt_team_ids({"teams": []}) == []
@pytest.mark.asyncio
async def test_map_jwt_role_to_litellm_role():
"""Test JWT role mapping to LiteLLM roles with various patterns"""

View file

@ -53,14 +53,20 @@ def test_non_admin_config_update_route_rejected():
assert "Your role=internal_user" in str(exc_info.value)
@pytest.mark.parametrize(
"role",
[
LitellmUserRoles.INTERNAL_USER.value,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
],
)
@pytest.mark.parametrize(
"route",
["/compliance/eu-ai-act", "/compliance/gdpr"],
)
def test_compliance_routes_open_to_internal_user(route):
def test_compliance_routes_open_to_non_admin_roles(role, route):
"""Compliance routes are stateless validators on caller-supplied log data
- non-admin internal_user roles can call them."""
role = LitellmUserRoles.INTERNAL_USER.value
both non-admin internal_user roles can call them."""
user_obj = LiteLLM_UserTable(
user_id="test_user",
user_email="test@example.com",
@ -80,56 +86,6 @@ def test_compliance_routes_open_to_internal_user(route):
)
def test_health_test_connection_route_delegates_internal_user_auth_to_endpoint():
"""Team model test-connection requests are authorized by the endpoint."""
role = LitellmUserRoles.INTERNAL_USER.value
user_obj = LiteLLM_UserTable(
user_id="test_user",
user_email="test@example.com",
user_role=role,
)
valid_token = UserAPIKeyAuth(user_id="test_user", user_role=role)
request = MagicMock(spec=Request)
request.query_params = {}
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=role,
route="/health/test_connection",
request=request,
valid_token=valid_token,
request_data={},
)
@pytest.mark.parametrize(
"route",
["/compliance/eu-ai-act", "/compliance/gdpr"],
)
def test_compliance_routes_blocked_for_internal_user_view_only(route):
"""Deprecated internal_user_viewer role must not gain compliance route access."""
role = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value
user_obj = LiteLLM_UserTable(
user_id="test_user",
user_email="test@example.com",
user_role=role,
)
valid_token = UserAPIKeyAuth(user_id="test_user", user_role=role)
request = MagicMock(spec=Request)
request.query_params = {}
with pytest.raises(Exception) as exc_info:
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=role,
route=route,
request=request,
valid_token=valid_token,
request_data={},
)
assert "Only proxy admin can be used" in str(exc_info.value)
def test_proxy_admin_viewer_config_update_route_rejected():
"""Test that proxy admin viewer users are rejected when trying to call /config/update"""
@ -1933,6 +1889,41 @@ def test_non_admin_non_team_admin_cannot_access_config_update_but_can_attempt_re
assert "Only proxy admin can be used to generate" in str(exc_info.value)
@pytest.mark.parametrize(
"user_role",
[
LitellmUserRoles.INTERNAL_USER.value,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
],
)
@pytest.mark.parametrize("route", ["/tag/list", "/tag/daily/activity"])
def test_internal_users_can_access_scoped_tag_usage_routes(user_role, route):
"""
Internal users can read tag usage endpoints because the endpoint handlers
scope results to the caller's own keys.
"""
user_obj = LiteLLM_UserTable(
user_id="test_user",
user_email="test@example.com",
user_role=user_role,
)
valid_token = UserAPIKeyAuth(
user_id="test_user",
user_role=user_role,
)
request = MagicMock(spec=Request)
request.query_params = {}
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=user_role,
route=route,
request=request,
valid_token=valid_token,
request_data={},
)
@pytest.mark.parametrize(
"user_role",
[

View file

@ -31,8 +31,10 @@ from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import (
_route_requires_auth_despite_public,
_matches_routing_override,
_reserve_budget_after_common_checks,
_route_requires_auth_despite_public,
_routing_selector_matches_claim,
_run_centralized_common_checks,
_run_post_custom_auth_checks,
get_api_key,
@ -594,6 +596,151 @@ def _assert_get_api_key_with_custom_litellm_key_header(
) == (api_key, passed_in_key)
@pytest.mark.parametrize(
"selector_value, claim_value, expected, split_space_delimited",
[
(None, "any-value", True, False),
("issuer.example.com", "issuer.example.com", True, False),
("issuer.example.com", "other-issuer.example.com", False, False),
# iss (and other non-scope claims) must not match via space-split injection
(
"trusted.example.com",
"trusted.example.com attacker.example.com",
False,
False,
),
# Wildcard iss must not match space-containing claim strings (fnmatch * spans spaces)
(
"trusted.*",
"trusted.example.com attacker.example.com",
False,
False,
),
("trusted.*", "trusted.example.com", True, False),
(
["issuer-a.example.com", "issuer-b.example.com"],
"issuer-b.example.com",
True,
False,
),
("*MID_LITELLM", "STREAM_MID_LITELLM", True, False),
("*MID_LITELLM", "REDIS_LITELLM", False, False),
("machine-??", "machine-01", True, False),
("machine-??", "machine-001", False, False),
# Wildcard matching is case-sensitive (fnmatch.fnmatchcase)
("*litellm", "BATCH_LITELLM", False, False),
("*LITELLM", "BATCH_LITELLM", True, False),
("App:LiteLLM", "App:LiteLLM openid", True, True),
("App:*", "App:LiteLLM openid", True, True),
(["openid", "App:LiteLLM"], "openid profile", True, True),
(["service-*", "batch-*"], "batch-123", True, False),
(["service-*", "batch-*"], "other-123", False, False),
("App:LiteLLM", ["openid", "App:LiteLLM"], True, False),
("App:LiteLLM", None, False, False),
],
)
def test_routing_selector_matches_claim_parametrized(
selector_value, claim_value, expected, split_space_delimited
):
assert (
_routing_selector_matches_claim(
selector_value=selector_value,
claim_value=claim_value,
split_space_delimited=split_space_delimited,
)
is expected
)
@pytest.mark.parametrize(
"override, token_claims, expected",
[
# Only iss selector is required and should match.
(
JWTRoutingOverride(iss="oauth-issuer.example.com", path="oauth2"),
{"iss": "oauth-issuer.example.com"},
True,
),
# Scope selector narrows the match.
(
JWTRoutingOverride(
iss="oauth-issuer.example.com",
scope="App:LiteLLM",
path="oauth2",
),
{"iss": "oauth-issuer.example.com", "scope": "App:LiteLLM openid"},
True,
),
# client_id wildcard selector narrows the match.
(
JWTRoutingOverride(
iss="oauth-issuer.example.com",
client_id="*MID_LITELLM",
path="oauth2",
),
{"iss": "oauth-issuer.example.com", "client_id": "BATCH_MID_LITELLM"},
True,
),
(
JWTRoutingOverride(
iss="oauth-issuer.example.com",
client_id="*MID_LITELLM",
path="oauth2",
),
{"iss": "oauth-issuer.example.com", "client_id": "BATCH_PORTAL"},
False,
),
# aud selector still works with list claims.
(
JWTRoutingOverride(
iss="oauth-issuer.example.com",
aud=["api://litellm", "api://fallback"],
path="oauth2",
),
{
"iss": "oauth-issuer.example.com",
"aud": ["api://other", "api://litellm"],
},
True,
),
# All provided selectors are AND-ed.
(
JWTRoutingOverride(
iss="oauth-issuer.example.com",
scope="App:LiteLLM",
client_id="*MID_LITELLM",
path="oauth2",
),
{
"iss": "oauth-issuer.example.com",
"scope": "App:LiteLLM openid",
"client_id": "BATCH_MID_LITELLM",
},
True,
),
(
JWTRoutingOverride(
iss="oauth-issuer.example.com",
scope="App:LiteLLM",
client_id="*MID_LITELLM",
path="oauth2",
),
{
"iss": "oauth-issuer.example.com",
"scope": "App:Other openid",
"client_id": "BATCH_MID_LITELLM",
},
False,
),
],
)
def test_matches_routing_override_parametrized(override, token_claims, expected):
assert (
_matches_routing_override(token_claims=token_claims, override=override)
is expected
)
def test_get_api_key_with_custom_litellm_key_header_bearer_prefix():
token = "sk-" + "1" * 8
header = f"Bearer {token}"
@ -1601,6 +1748,206 @@ class TestJWTOAuth2Coexistence:
mock_jwt_auth.assert_not_called()
assert result.user_id == "machine-client-aud-list"
@pytest.mark.asyncio
async def test_routing_override_matches_scope_claim(self):
"""
Match routing override when scope selector is configured and scope claim matches.
"""
jwt_token = (
"eyJhbGciOiJSUzI1NiJ9."
"eyJpc3MiOiJvYXV0aC1pc3N1ZXIuZXhhbXBsZS5jb20iLCJzY29wZSI6IkFwcDpMaXRlTExNIiwiY2xpZW50X2lkIjoiTUFDSElORV9NSURfTElURUxMTSJ9."
"c2ln"
)
general_settings = {
"enable_oauth2_auth": False,
"enable_jwt_auth": True,
}
mock_oauth2_response = UserAPIKeyAuth(
api_key=jwt_token,
user_id="machine-client-scope-match",
)
mock_request = MagicMock()
mock_request.url.path = "/v1/chat/completions"
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
mock_request.query_params = {}
with (
patch("litellm.proxy.proxy_server.general_settings", general_settings),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch(
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
new_callable=AsyncMock,
return_value=mock_oauth2_response,
) as mock_oauth2,
patch(
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
) as mock_jwt_auth,
):
litellm.proxy.proxy_server.jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=DualCache(),
litellm_jwtauth=LiteLLM_JWTAuth(
routing_overrides=[
JWTRoutingOverride(
iss="oauth-issuer.example.com",
scope="App:LiteLLM",
path="oauth2",
)
]
),
)
result = await user_api_key_auth(
request=mock_request,
api_key=f"Bearer {jwt_token}",
)
mock_oauth2.assert_called_once_with(token=jwt_token)
mock_jwt_auth.assert_not_called()
assert result.user_id == "machine-client-scope-match"
@pytest.mark.asyncio
async def test_routing_override_scope_mismatch_falls_back_to_jwt(self):
"""
If scope selector does not match, continue default JWT flow.
"""
jwt_token = (
"eyJhbGciOiJSUzI1NiJ9."
"eyJpc3MiOiJvYXV0aC1pc3N1ZXIuZXhhbXBsZS5jb20iLCJzY29wZSI6IkFwcDpPdGhlciIsImNsaWVudF9pZCI6IlBPUlRBTF9NSURfTElURUxMTSJ9."
"c2ln"
)
general_settings = {
"enable_oauth2_auth": False,
"enable_jwt_auth": True,
}
mock_jwt_result = {
"is_proxy_admin": True,
"team_object": None,
"user_object": None,
"end_user_object": None,
"org_object": None,
"token": jwt_token,
"team_id": "jwt-team",
"user_id": "jwt-user-scope-mismatch",
"end_user_id": None,
"org_id": None,
"team_membership": None,
"jwt_claims": {"sub": "user1"},
}
mock_request = MagicMock()
mock_request.url.path = "/v1/chat/completions"
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
mock_request.query_params = {}
with (
patch("litellm.proxy.proxy_server.general_settings", general_settings),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch(
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
new_callable=AsyncMock,
) as mock_oauth2,
patch(
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
return_value=mock_jwt_result,
) as mock_jwt_auth,
):
litellm.proxy.proxy_server.jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=DualCache(),
litellm_jwtauth=LiteLLM_JWTAuth(
routing_overrides=[
JWTRoutingOverride(
iss="oauth-issuer.example.com",
scope="App:LiteLLM",
path="oauth2",
)
]
),
)
result = await user_api_key_auth(
request=mock_request,
api_key=f"Bearer {jwt_token}",
)
mock_oauth2.assert_not_called()
mock_jwt_auth.assert_called_once()
assert result.user_id == "jwt-user-scope-mismatch"
@pytest.mark.asyncio
async def test_routing_override_matches_scope_and_client_wildcard_when_scope_claim_is_space_delimited(
self,
):
"""
Integration check: combined scope + wildcard selectors match on OAuth2 path
when scope claim is a space-delimited string.
"""
jwt_token = (
"eyJhbGciOiJSUzI1NiJ9."
"eyJpc3MiOiJvYXV0aC1pc3N1ZXIuZXhhbXBsZS5jb20iLCJzY29wZSI6IkFwcDpMaXRlTExNIG9wZW5pZCIsImNsaWVudF9pZCI6IkJBVENIX01JRF9MSVRFTExNIn0."
"c2ln"
)
general_settings = {
"enable_oauth2_auth": False,
"enable_jwt_auth": True,
}
mock_oauth2_response = UserAPIKeyAuth(
api_key=jwt_token,
user_id="machine-client-space-delimited-scope-match",
)
mock_request = MagicMock()
mock_request.url.path = "/v1/chat/completions"
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
mock_request.query_params = {}
with (
patch("litellm.proxy.proxy_server.general_settings", general_settings),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch(
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
new_callable=AsyncMock,
return_value=mock_oauth2_response,
) as mock_oauth2,
patch(
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
) as mock_jwt_auth,
):
litellm.proxy.proxy_server.jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=DualCache(),
litellm_jwtauth=LiteLLM_JWTAuth(
routing_overrides=[
JWTRoutingOverride(
iss="oauth-issuer.example.com",
scope="App:LiteLLM",
client_id="*MID_LITELLM",
path="oauth2",
)
]
),
)
result = await user_api_key_auth(
request=mock_request,
api_key=f"Bearer {jwt_token}",
)
mock_oauth2.assert_called_once_with(token=jwt_token)
mock_jwt_auth.assert_not_called()
assert result.user_id == "machine-client-space-delimited-scope-match"
@pytest.mark.asyncio
async def test_routing_override_routes_jwt_to_oauth2_when_oauth2_globally_disabled(
self,

View file

@ -39,6 +39,46 @@ class MockLiteLLMVerificationToken:
return {"count": 1}
class MockLiteLLMOrganizationTable:
def __init__(self):
self.update_many_calls: List[Dict[str, Any]] = []
self.find_many_calls: List[Dict[str, Any]] = []
self._find_many_results: List[Any] = []
def set_find_many_results(self, results: List[Any]):
self._find_many_results = results
async def find_many(self, where: Dict[str, Any]) -> List[Any]:
self.find_many_calls.append({"where": where})
return self._find_many_results
async def update_many(
self, where: Dict[str, Any], data: Dict[str, Any]
) -> Dict[str, Any]:
self.update_many_calls.append({"where": where, "data": data})
return {"count": 1}
class MockLiteLLMTagTable:
def __init__(self):
self.update_many_calls: List[Dict[str, Any]] = []
self.find_many_calls: List[Dict[str, Any]] = []
self._find_many_results: List[Any] = []
def set_find_many_results(self, results: List[Any]):
self._find_many_results = results
async def find_many(self, where: Dict[str, Any]) -> List[Any]:
self.find_many_calls.append({"where": where})
return self._find_many_results
async def update_many(
self, where: Dict[str, Any], data: Dict[str, Any]
) -> Dict[str, Any]:
self.update_many_calls.append({"where": where, "data": data})
return {"count": 1}
class MockLiteLLMEndUserTable:
def __init__(self):
self.find_many_calls: List[Dict[str, Any]] = []
@ -57,6 +97,8 @@ class MockDB:
self.litellm_teammembership = MockLiteLLMTeamMembership()
self.litellm_verificationtoken = MockLiteLLMVerificationToken()
self.litellm_endusertable = MockLiteLLMEndUserTable()
self.litellm_organizationtable = MockLiteLLMOrganizationTable()
self.litellm_tagtable = MockLiteLLMTagTable()
class MockPrismaClient:
@ -459,6 +501,100 @@ def test_reset_budget_for_keys_linked_to_budgets_empty(
assert len(calls) == 0
def test_reset_budget_for_orgs_linked_to_budgets(reset_budget_job, mock_prisma_client):
"""
Test that when a budget tier is reset, orgs linked to that budget
(via budget_id) also get their spend reset.
"""
now = datetime.now(timezone.utc)
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 100.0,
"budget_duration": "30d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "30d-org-budget",
"created_at": now - timedelta(days=30),
},
)
asyncio.run(
reset_budget_job.reset_budget_for_orgs_linked_to_budgets(
budgets_to_reset=[test_budget]
)
)
calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls
assert len(calls) == 1
call = calls[0]
assert call["where"]["budget_id"] == {"in": ["30d-org-budget"]}
assert call["where"]["spend"] == {"gt": 0}
assert call["data"]["spend"] == 0
def test_reset_budget_for_orgs_linked_to_budgets_empty(
reset_budget_job, mock_prisma_client
):
"""
Test that when there are no budgets to reset, no update is performed
on the organization table.
"""
asyncio.run(
reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[])
)
calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls
assert len(calls) == 0
def test_reset_budget_for_tags_linked_to_budgets(reset_budget_job, mock_prisma_client):
"""
Test that when a budget tier is reset, tags linked to that budget
(via budget_id) also get their spend reset.
"""
now = datetime.now(timezone.utc)
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 50.0,
"budget_duration": "30d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "30d-tag-budget",
"created_at": now - timedelta(days=30),
},
)
asyncio.run(
reset_budget_job.reset_budget_for_tags_linked_to_budgets(
budgets_to_reset=[test_budget]
)
)
calls = mock_prisma_client.db.litellm_tagtable.update_many_calls
assert len(calls) == 1
call = calls[0]
assert call["where"]["budget_id"] == {"in": ["30d-tag-budget"]}
assert call["where"]["spend"] == {"gt": 0}
assert call["data"]["spend"] == 0
def test_reset_budget_for_tags_linked_to_budgets_empty(
reset_budget_job, mock_prisma_client
):
"""
Test that when there are no budgets to reset, no update is performed
on the tag table.
"""
asyncio.run(
reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[])
)
calls = mock_prisma_client.db.litellm_tagtable.update_many_calls
assert len(calls) == 0
@pytest.mark.parametrize(
"budget_duration, expected_day, expected_month",
[
@ -618,6 +754,75 @@ def test_budget_table_reset_also_resets_linked_keys(
assert calls[0]["data"]["spend"] == 0
def test_budget_table_reset_also_resets_linked_orgs(
reset_budget_job, mock_prisma_client
):
"""
Integration-style test: when reset_budget_for_litellm_budget_table runs,
it should also reset spend for orgs linked to the expiring budget tiers
(in addition to end-users, team members, and keys).
"""
now = datetime.now(timezone.utc)
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 100.0,
"budget_duration": "30d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "30d-org-budget",
"created_at": now - timedelta(days=30),
},
)
mock_prisma_client.data["budget"] = [test_budget]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls
assert len(calls) == 1, (
"Expected reset_budget_for_litellm_budget_table to also reset orgs "
f"linked to expiring budgets, but got {len(calls)} update_many calls"
)
assert calls[0]["where"]["budget_id"] == {"in": ["30d-org-budget"]}
assert calls[0]["data"]["spend"] == 0
def test_budget_table_reset_also_resets_linked_tags(
reset_budget_job, mock_prisma_client
):
"""
Integration-style test: when reset_budget_for_litellm_budget_table runs,
it should also reset spend for tags linked to the expiring budget tiers.
"""
now = datetime.now(timezone.utc)
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 50.0,
"budget_duration": "30d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "30d-tag-budget",
"created_at": now - timedelta(days=30),
},
)
mock_prisma_client.data["budget"] = [test_budget]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
calls = mock_prisma_client.db.litellm_tagtable.update_many_calls
assert len(calls) == 1, (
"Expected reset_budget_for_litellm_budget_table to also reset tags "
f"linked to expiring budgets, but got {len(calls)} update_many calls"
)
assert calls[0]["where"]["budget_id"] == {"in": ["30d-tag-budget"]}
assert calls[0]["data"]["spend"] == 0
def test_reset_budget_resets_endusers_with_null_budget_id(
reset_budget_job, mock_prisma_client
):
@ -1057,16 +1262,26 @@ def test_reset_budget_windows_query_error_does_not_break_team_path(monkeypatch):
def _make_counter_invalidation_job(monkeypatch):
"""Stub spend_counter_cache so we can observe invalidation calls."""
"""Stub spend_counter_cache (and user_api_key_cache) so we can observe
invalidation calls.
Both caches are looked up via ``from litellm.proxy.proxy_server import
<name>`` inside the reset job, so we publish them on a fake module.
"""
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = MagicMock()
spend_counter_cache.redis_cache.async_set_cache = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_delete_cache = AsyncMock()
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
fake_module.user_api_key_cache = user_api_key_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
spend_counter_cache.user_api_key_cache = user_api_key_cache
return spend_counter_cache
@ -1205,3 +1420,136 @@ def test_reset_budget_for_keys_linked_to_budgets_invalidates_redis_counter(monke
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:key:sk-linked", value=0.0, ttl=60
)
def test_reset_budget_for_orgs_linked_to_budgets_invalidates_redis_counter(monkeypatch):
"""Resetting orgs via budget tier must clear each linked org's counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_org = type("Org", (), {"organization_id": "org-acme"})
prisma_client = MagicMock()
prisma_client.db.litellm_organizationtable.find_many = AsyncMock(
return_value=[linked_org]
)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 1}
)
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_orgs_linked_to_budgets([expired_budget]))
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:org:org-acme", value=0.0, ttl=60
)
counter_cache.redis_cache.async_set_cache.assert_any_await(
key="spend:org:org-acme", value=0.0, ttl=60
)
def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monkeypatch):
"""Resetting tags via budget tier must clear each linked tag's counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_tag = type("Tag", (), {"tag_name": "tenant-42"})
prisma_client = MagicMock()
prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag])
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget]))
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:tag:tenant-42", value=0.0, ttl=60
)
counter_cache.redis_cache.async_set_cache.assert_any_await(
key="spend:tag:tenant-42", value=0.0, ttl=60
)
def test_reset_budget_for_tags_linked_to_budgets_invalidates_management_cache(
monkeypatch,
):
"""Regression guard for the bug where tag spend stayed frozen across cycles.
``SpendCounterReseed.from_db`` returns ``None`` for ``spend:tag:*`` keys,
so once the spend counter expires the tag budget check falls back to the
cached ``LiteLLM_TagTable.spend``. If we don't drop the management cache
entry on reset, that cached object lingers (TTL 60s) with the pre-reset
spend, and ``_tag_max_budget_check`` keeps returning HTTP 400 even though
the DB row has been zeroed.
"""
counter_cache = _make_counter_invalidation_job(monkeypatch)
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_tag = type("Tag", (), {"tag_name": "tenant-42"})
prisma_client = MagicMock()
prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag])
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget]))
counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(
key="tag:tenant-42"
)
def test_reset_budget_for_tags_linked_to_budgets_invalidates_each_tag_management_cache(
monkeypatch,
):
"""When multiple tags share the expired budget tier, every one of them
has its ``user_api_key_cache`` entry dropped not just the first."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_tags = [
type("Tag", (), {"tag_name": "tenant-a"}),
type("Tag", (), {"tag_name": "tenant-b"}),
type("Tag", (), {"tag_name": "tenant-c"}),
]
prisma_client = MagicMock()
prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=linked_tags)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 3})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget]))
deleted_keys = {
call.kwargs.get("key")
for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list
}
assert deleted_keys == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"}
def test_reset_budget_for_keys_linked_to_budgets_does_not_touch_management_cache(
monkeypatch,
):
"""Cache invalidation is opt-in: keys / orgs / team-members rely on
``SpendCounterReseed.from_db`` (which DOES handle their counter keys),
so the cache_key_fn hook is intentionally not wired for them. This test
locks in that no-op so a future refactor doesn't accidentally start
clobbering the key cache (which would cost an extra DB round-trip per
reset cycle without fixing anything)."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_key = type("Key", (), {"token": "sk-linked"})
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[linked_key]
)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 1}
)
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget]))
counter_cache.user_api_key_cache.async_delete_cache.assert_not_awaited()

View file

@ -1513,3 +1513,146 @@ async def test_commit_spend_updates_uses_pipeline():
mock_redis_update_buffer.get_all_daily_end_user_spend_update_transactions_from_redis_buffer.assert_not_called()
mock_redis_update_buffer.get_all_daily_agent_spend_update_transactions_from_redis_buffer.assert_not_called()
mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer.assert_not_called()
@pytest.mark.parametrize(
"bucket_name,input_dict,table_attr,method_name,where_key,expected_order",
[
pytest.param(
"user_list_transactions",
{"user_c": 0.1, "user_a": 0.2, "user_b": 0.3},
"litellm_usertable",
"update_many",
"user_id",
["user_a", "user_b", "user_c"],
id="user",
),
pytest.param(
"key_list_transactions",
{"tok_c": 0.1, "tok_a": 0.2, "tok_b": 0.3},
"litellm_verificationtoken",
"update_many",
"token",
["tok_a", "tok_b", "tok_c"],
id="key",
),
pytest.param(
"team_list_transactions",
{"team_c": 0.1, "team_a": 0.2, "team_b": 0.3},
"litellm_teamtable",
"update_many",
"team_id",
["team_a", "team_b", "team_c"],
id="team",
),
pytest.param(
"team_member_list_transactions",
{
"team_id::team_c::user_id::user_x": 0.1,
"team_id::team_a::user_id::user_x": 0.2,
"team_id::team_b::user_id::user_x": 0.3,
},
"litellm_teammembership",
"update_many",
"team_id",
["team_a", "team_b", "team_c"],
id="team_member",
),
pytest.param(
"org_list_transactions",
{"org_c": 0.1, "org_a": 0.2, "org_b": 0.3},
"litellm_organizationtable",
"update_many",
"organization_id",
["org_a", "org_b", "org_c"],
id="org",
),
pytest.param(
"end_user_list_transactions",
{"eu_c": 0.1, "eu_a": 0.2, "eu_b": 0.3},
"litellm_endusertable",
"upsert",
"user_id",
["eu_a", "eu_b", "eu_c"],
id="end_user",
),
pytest.param(
"tag_list_transactions",
{"prod": 0.1, "customer-x": 0.2, "test": 0.3},
"litellm_tagtable",
"update_many",
"tag_name",
["customer-x", "prod", "test"],
id="tag",
),
pytest.param(
"agent_list_transactions",
{"agent_c": 0.1, "agent_a": 0.2, "agent_b": 0.3},
"litellm_agentstable",
"update_many",
"agent_id",
["agent_a", "agent_b", "agent_c"],
id="agent",
),
],
)
@pytest.mark.asyncio
async def test_commit_spend_updates_iterates_in_sorted_order(
bucket_name, input_dict, table_attr, method_name, where_key, expected_order
):
"""
Every spend-bucket code path in _commit_spend_updates_to_db must iterate
in sorted order so concurrent pods acquire row locks in the same order
and avoid PostgreSQL deadlocks. Covers the 5 direct loops (user/key/team/
team_member/org), the end_user path in ProxyUpdateSpend.update_end_user_spend,
and the shared _update_entity_spend_in_db helper (tag, agent).
"""
db_writer = DBSpendUpdateWriter()
captured_where_values = []
def capture(*, where, data):
captured_where_values.append(where[where_key])
mock_batcher = MagicMock()
table_mock = MagicMock()
setattr(table_mock, method_name, MagicMock(side_effect=capture))
setattr(mock_batcher, table_attr, table_mock)
mock_transaction = AsyncMock()
mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction)
mock_transaction.__aexit__ = AsyncMock(return_value=False)
mock_transaction.batch_ = MagicMock(
return_value=AsyncMock(
__aenter__=AsyncMock(return_value=mock_batcher),
__aexit__=AsyncMock(return_value=False),
)
)
mock_prisma_client = MagicMock()
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction)
mock_proxy_logging = MagicMock()
mock_proxy_logging.call_details = {}
buckets = {
"user_list_transactions": {},
"end_user_list_transactions": {},
"key_list_transactions": {},
"team_list_transactions": {},
"team_member_list_transactions": {},
"org_list_transactions": {},
"tag_list_transactions": {},
"agent_list_transactions": {},
}
buckets[bucket_name] = input_dict
await db_writer._commit_spend_updates_to_db(
prisma_client=mock_prisma_client,
n_retry_times=3,
proxy_logging_obj=mock_proxy_logging,
db_spend_update_transactions=buckets,
)
assert captured_where_values == expected_order

View file

@ -106,9 +106,11 @@ def mock_in_memory_handler(mocker):
mock_handler = mocker.Mock(spec=InMemoryGuardrailHandler)
mock_handler.list_in_memory_guardrails.return_value = [MOCK_CONFIG_GUARDRAIL]
mock_handler.get_guardrail_by_id.return_value = MOCK_CONFIG_GUARDRAIL
mock_handler.get_source.return_value = "config"
mock_handler.initialize_guardrail = mocker.Mock()
mock_handler.update_in_memory_guardrail = mocker.Mock()
mock_handler.delete_in_memory_guardrail = mocker.Mock()
mock_handler.reconcile_db_guardrails = mocker.Mock(return_value=[])
return mock_handler
@ -162,6 +164,67 @@ async def test_list_guardrails_v2_with_db_and_config(
assert isinstance(config_guardrail.litellm_params, BaseLitellmParams)
@pytest.mark.asyncio
async def test_list_guardrails_v2_skips_stale_db_backed_in_memory_entries(mocker):
"""
A guardrail that's still in this pod's memory tagged source='db' but is no
longer in the DB result (deleted on another pod, awaiting reconcile) must
NOT surface in the list response pre-fix it leaked as 'config'.
"""
stale_guardrail = {
"guardrail_id": "stale-db-id",
"guardrail_name": "Stale DB Guardrail",
"litellm_params": {"guardrail": "bedrock", "mode": "pre_call"},
"guardrail_info": {},
}
mock_prisma_client = mocker.Mock()
mock_prisma_client.db = mocker.Mock()
mock_prisma_client.db.litellm_guardrailstable = mocker.Mock()
mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[])
mock_in_memory_handler = mocker.Mock()
mock_in_memory_handler.list_in_memory_guardrails.return_value = [stale_guardrail]
mock_in_memory_handler.get_source.return_value = "db"
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)
admin_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
response = await list_guardrails_v2(user_api_key_dict=admin_auth)
assert response.guardrails == []
mock_in_memory_handler.get_source.assert_called_with("stale-db-id")
@pytest.mark.asyncio
async def test_get_guardrail_info_404s_stale_db_backed_entry(
mocker, mock_prisma_client, mock_in_memory_handler
):
"""
Stale DB-backed entry (in-memory but not in DB) must 404 instead of being
returned as if it were a config-loaded guardrail.
"""
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)
mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(
return_value=None
)
# In-memory still has it, but it's tagged as 'db' (stale, awaiting reconcile)
mock_in_memory_handler.get_source.return_value = "db"
with pytest.raises(HTTPException) as exc_info:
await get_guardrail_info("stale-db-id")
assert exc_info.value.status_code == 404
assert "not found" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_list_guardrails_v2_masks_sensitive_data_in_db_guardrails(mocker):
"""Test that sensitive litellm_params are masked for DB guardrails in list response"""
@ -1160,6 +1223,7 @@ async def test_get_guardrail_info_endpoint_config_guardrail(mocker):
# Mock IN_MEMORY_GUARDRAIL_HANDLER at its source to return config guardrail
mock_in_memory_handler = mocker.Mock()
mock_in_memory_handler.get_guardrail_by_id.return_value = MOCK_CONFIG_GUARDRAIL
mock_in_memory_handler.get_source.return_value = "config"
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,

View file

@ -60,3 +60,123 @@ def test_update_in_memory_guardrail():
handler.guardrail_id_to_custom_guardrail["123"].event_hook
is GuardrailEventHooks.pre_call
)
def _make_guardrail(guardrail_id: str, name: str = "g") -> Guardrail:
return Guardrail(
guardrail_id=guardrail_id,
guardrail_name=name,
litellm_params=LitellmParams(guardrail=name, mode="pre_call", default_on=False),
)
def test_reconcile_db_guardrails_drops_stale_db_entries_only():
"""
The reconcile pass must drop in-memory entries marked source='db' that are
missing from the DB result, and never touch source='config' entries.
Models the multi-pod case where another pod deleted a DB-backed guardrail.
"""
handler = InMemoryGuardrailHandler()
# Two DB-backed entries on this pod (synced from earlier polling cycles)
handler.IN_MEMORY_GUARDRAILS["db-keep"] = _make_guardrail("db-keep")
handler.IN_MEMORY_GUARDRAILS["db-stale"] = _make_guardrail("db-stale")
handler._sources["db-keep"] = "db"
handler._sources["db-stale"] = "db"
# One config-loaded entry that must survive reconciliation
handler.IN_MEMORY_GUARDRAILS["cfg"] = _make_guardrail("cfg")
handler._sources["cfg"] = "config"
# The DB now only contains db-keep — db-stale was deleted on another pod.
removed = handler.reconcile_db_guardrails(db_guardrail_ids={"db-keep"})
assert removed == ["db-stale"]
assert "db-stale" not in handler.IN_MEMORY_GUARDRAILS
assert "db-stale" not in handler._sources
assert "db-keep" in handler.IN_MEMORY_GUARDRAILS
assert "cfg" in handler.IN_MEMORY_GUARDRAILS
assert handler._sources["cfg"] == "config"
def test_reconcile_does_not_drop_config_entries_missing_from_db():
"""A config-only guardrail (no DB row) must never be reconciled away."""
handler = InMemoryGuardrailHandler()
handler.IN_MEMORY_GUARDRAILS["cfg-only"] = _make_guardrail("cfg-only")
handler._sources["cfg-only"] = "config"
removed = handler.reconcile_db_guardrails(db_guardrail_ids=set())
assert removed == []
assert "cfg-only" in handler.IN_MEMORY_GUARDRAILS
def test_get_source_returns_marker_set_at_insert():
handler = InMemoryGuardrailHandler()
handler.IN_MEMORY_GUARDRAILS["a"] = _make_guardrail("a")
handler._sources["a"] = "db"
handler.IN_MEMORY_GUARDRAILS["b"] = _make_guardrail("b")
handler._sources["b"] = "config"
assert handler.get_source("a") == "db"
assert handler.get_source("b") == "config"
assert handler.get_source("missing") is None
def test_delete_in_memory_guardrail_clears_source_marker():
handler = InMemoryGuardrailHandler()
handler.IN_MEMORY_GUARDRAILS["a"] = _make_guardrail("a")
handler._sources["a"] = "db"
handler.delete_in_memory_guardrail("a")
assert "a" not in handler.IN_MEMORY_GUARDRAILS
assert "a" not in handler._sources
assert handler.get_source("a") is None
def test_initialize_guardrail_early_return_updates_source_marker():
"""
When initialize_guardrail is called for a guardrail that already exists
in memory, the early-return path must still honor the caller's source.
Otherwise a racing polling tick that placed a DB entry in memory first
would leave a later config-init call wrongly marked as 'db' (or vice
versa), and the entry would be reconciled with the wrong classification.
"""
handler = InMemoryGuardrailHandler()
# Simulate a polling tick already placing the entry as DB-backed.
handler.IN_MEMORY_GUARDRAILS["collide"] = _make_guardrail("collide", name="bedrock")
handler._sources["collide"] = "db"
# Config init re-visits the same id (e.g., hot-reload, or UUID collision).
g = Guardrail(
guardrail_id="collide",
guardrail_name="bedrock",
litellm_params=LitellmParams(
guardrail="bedrock", mode="pre_call", default_on=False
),
)
handler.initialize_guardrail(guardrail=g, source="config")
assert handler.get_source("collide") == "config"
# And the symmetric direction: db sync should override an entry left
# marked as 'config' from a stale init path.
handler.initialize_guardrail(guardrail=g, source="db")
assert handler.get_source("collide") == "db"
def test_sync_guardrail_from_db_marks_source_db_when_unchanged():
"""
sync_guardrail_from_db must enforce source='db' even when params are
unchanged, so a config entry whose UUID happens to collide with a later
DB row gets re-tagged correctly.
"""
handler = InMemoryGuardrailHandler()
g = _make_guardrail("collide")
handler.IN_MEMORY_GUARDRAILS["collide"] = g
handler._sources["collide"] = "config"
handler.sync_guardrail_from_db(g)
assert handler.get_source("collide") == "db"

View file

@ -466,6 +466,236 @@ async def test_test_model_connection_loads_config_from_router():
assert "result" in result
@pytest.mark.asyncio
async def test_test_model_connection_uses_model_info_id_to_disambiguate_duplicate_model_names():
"""
When two deployments share the same `model_name` (e.g. wildcard
`openai/*`) but have different `api_base` values, clicking "Test
Connection" on a specific row in the UI must probe THAT row's
`api_base` not whichever happens to be `deployments[0]`.
The UI passes `model_info.id` to identify the deployment the user
actually clicked on. The backend must use that id to look up the
specific deployment rather than always grabbing the first match.
Regression test for: silent fallback to deployments[0] when
multiple deployments share a wildcard model_name.
"""
from litellm.types.router import Deployment, LiteLLM_Params
mock_request = MagicMock()
mock_user_api_key_dict = MagicMock()
mock_user_api_key_dict.user_id = "test-user"
mock_user_api_key_dict.token = "test-token"
mock_prisma_client = MagicMock()
deployment_a = {
"model_name": "openai/*",
"litellm_params": {
"model": "openai/*",
"api_base": "https://deployment-A-base.invalid/v1",
"api_key": "fake-key-A",
},
"model_info": {"id": "deployment-A-id"},
}
deployment_b = {
"model_name": "openai/*",
"litellm_params": {
"model": "openai/*",
"api_base": "https://deployment-B-base.invalid/v1",
"api_key": "fake-key-B",
},
"model_info": {"id": "deployment-B-id"},
}
mock_router = MagicMock()
mock_router.get_model_list.return_value = [deployment_a, deployment_b]
# Backend uses get_deployment(model_id=...) for O(1) lookup by id.
def _get_deployment_by_id(model_id):
if model_id == "deployment-A-id":
return Deployment(
model_name="openai/*",
litellm_params=LiteLLM_Params(**deployment_a["litellm_params"]),
model_info=deployment_a["model_info"],
)
if model_id == "deployment-B-id":
return Deployment(
model_name="openai/*",
litellm_params=LiteLLM_Params(**deployment_b["litellm_params"]),
model_info=deployment_b["model_info"],
)
return None
mock_router.get_deployment.side_effect = _get_deployment_by_id
mock_can_user_make_model_call = AsyncMock()
mock_health_check_result = {"status": "healthy", "response_time_ms": 50}
mock_ahealth_check = AsyncMock(return_value=mock_health_check_result)
mock_run_with_timeout = AsyncMock(return_value=mock_health_check_result)
def mock_update_params(model_info, litellm_params):
params = litellm_params.copy()
params["messages"] = [{"role": "user", "content": "test"}]
return params
def mock_reject_os_environ(params):
return None
with (
patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
),
patch(
"litellm.proxy.proxy_server.llm_router",
mock_router,
),
patch(
"litellm.proxy.proxy_server.premium_user",
False,
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
mock_can_user_make_model_call,
),
patch(
"litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check",
mock_ahealth_check,
),
patch(
"litellm.proxy.health_endpoints._health_endpoints.run_with_timeout",
mock_run_with_timeout,
),
patch(
"litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check",
mock_update_params,
),
patch(
"litellm.proxy.health_endpoints._health_endpoints._reject_os_environ_references",
mock_reject_os_environ,
),
):
# Click "Test Connection" on deployment B (NOT the first one).
# The UI sends only `model` + `model_info.id` — it does NOT
# send `api_base`/`api_key`, so the backend must resolve them
# from the right deployment.
await health_test_model_connection(
request=mock_request,
mode="chat",
litellm_params={"model": "openai/*"},
model_info={"id": "deployment-B-id"},
user_api_key_dict=mock_user_api_key_dict,
)
# The outbound health check must hit deployment B's api_base.
ahealth_check_call_args = mock_ahealth_check.call_args
assert ahealth_check_call_args is not None
model_params = ahealth_check_call_args.kwargs.get("model_params", {})
assert model_params.get("api_base") == (
"https://deployment-B-base.invalid/v1"
), (
"Expected /health/test_connection to probe deployment B's "
"api_base when model_info.id='deployment-B-id' was provided. "
f"Got: {model_params.get('api_base')!r}. This means the "
"backend silently fell back to deployments[0] (A) instead "
"of disambiguating by model_info.id."
)
assert model_params.get("api_key") == "fake-key-B"
@pytest.mark.asyncio
async def test_test_model_connection_falls_back_to_deployments_zero_without_id():
"""
Backwards-compat: when the request body does NOT include
`model_info.id`, the legacy behavior of using `deployments[0]`
is preserved (single-deployment case, or callers that haven't
been updated to pass an id).
"""
mock_request = MagicMock()
mock_user_api_key_dict = MagicMock()
mock_user_api_key_dict.user_id = "test-user"
mock_user_api_key_dict.token = "test-token"
mock_prisma_client = MagicMock()
deployment_a = {
"model_name": "openai/*",
"litellm_params": {
"model": "openai/*",
"api_base": "https://deployment-A-base.invalid/v1",
"api_key": "fake-key-A",
},
"model_info": {"id": "deployment-A-id"},
}
deployment_b = {
"model_name": "openai/*",
"litellm_params": {
"model": "openai/*",
"api_base": "https://deployment-B-base.invalid/v1",
"api_key": "fake-key-B",
},
"model_info": {"id": "deployment-B-id"},
}
mock_router = MagicMock()
mock_router.get_model_list.return_value = [deployment_a, deployment_b]
mock_can_user_make_model_call = AsyncMock()
mock_health_check_result = {"status": "healthy"}
mock_ahealth_check = AsyncMock(return_value=mock_health_check_result)
mock_run_with_timeout = AsyncMock(return_value=mock_health_check_result)
def mock_update_params(model_info, litellm_params):
params = litellm_params.copy()
params["messages"] = [{"role": "user", "content": "test"}]
return params
def mock_reject_os_environ(params):
return None
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
patch("litellm.proxy.proxy_server.llm_router", mock_router),
patch("litellm.proxy.proxy_server.premium_user", False),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
mock_can_user_make_model_call,
),
patch(
"litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check",
mock_ahealth_check,
),
patch(
"litellm.proxy.health_endpoints._health_endpoints.run_with_timeout",
mock_run_with_timeout,
),
patch(
"litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check",
mock_update_params,
),
patch(
"litellm.proxy.health_endpoints._health_endpoints._reject_os_environ_references",
mock_reject_os_environ,
),
):
await health_test_model_connection(
request=mock_request,
mode="chat",
litellm_params={"model": "openai/*"},
model_info={}, # no id provided
user_api_key_dict=mock_user_api_key_dict,
)
# Without id, deployments[0] (A) should be used (legacy behavior).
model_params = mock_ahealth_check.call_args.kwargs.get("model_params", {})
assert model_params.get("api_base") == "https://deployment-A-base.invalid/v1"
assert model_params.get("api_key") == "fake-key-A"
@pytest.mark.asyncio
async def test_health_services_endpoint_datadog_llm_observability():
"""

View file

@ -5689,7 +5689,7 @@ async def test_process_single_key_update():
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create update request
key_update_item = BulkUpdateKeyRequestItem(
update_key_request = UpdateKeyRequest(
key="test-key-123",
max_budget=100.0,
tags=["production"],
@ -5703,7 +5703,7 @@ async def test_process_single_key_update():
# Call the function
result = await _process_single_key_update(
key_update_item=key_update_item,
update_key_request=update_key_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
prisma_client=mock_prisma_client,
@ -9855,9 +9855,6 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_process_single_key_update,
)
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequestItem,
)
token_hash = "abc123def456"
@ -9900,7 +9897,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash():
new_callable=AsyncMock,
),
):
key_update_item = BulkUpdateKeyRequestItem(
update_key_request = UpdateKeyRequest(
key=token_hash,
max_budget=100.0,
)
@ -9912,7 +9909,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash():
)
await _process_single_key_update(
key_update_item=key_update_item,
update_key_request=update_key_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
prisma_client=mock_prisma_client,
@ -10019,3 +10016,583 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha
call_kwargs = mock_delete_cache.call_args.kwargs
# The token hash should be passed as-is, NOT double-hashed
assert call_kwargs["hashed_token"] == token_hash
# ---------------------------------------------------------------------------
# /team/key/bulk_update tests
# ---------------------------------------------------------------------------
_BULK_PKG = "litellm.proxy.management_endpoints.key_management_endpoints"
def _make_team_key(token: str, team_id: str = "team-abc") -> LiteLLM_VerificationToken:
return LiteLLM_VerificationToken(
token=token,
user_id="user-123",
models=[],
team_id=team_id,
max_budget=None,
)
def _admin() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin"
)
def _internal_user() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-iu", user_id="iu"
)
def _updated(payload):
m = MagicMock()
m.model_dump.return_value = payload
return m
def _setup_team_keys_mocks(
monkeypatch,
*,
find_many=None,
find_unique=None,
update_data=None,
hash_identity=True,
):
"""Set up mocks for bulk_update_team_keys; returns mock_prisma."""
mock_prisma = AsyncMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[] if find_many is None else find_many
)
if find_unique is not None:
mock_prisma.db.litellm_verificationtoken.find_unique = find_unique
if update_data is not None:
mock_prisma.update_data = update_data
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None)
monkeypatch.setattr(
f"{_BULK_PKG}.prepare_key_update_data",
AsyncMock(return_value={"max_budget": 50.0}),
)
monkeypatch.setattr(f"{_BULK_PKG}._delete_cache_key_object", AsyncMock())
monkeypatch.setattr(
f"{_BULK_PKG}.KeyManagementEventHooks.async_key_updated_hook", AsyncMock()
)
monkeypatch.setattr(f"{_BULK_PKG}.get_team_object", AsyncMock(return_value=None))
monkeypatch.setattr(f"{_BULK_PKG}._check_team_key_limits", AsyncMock())
monkeypatch.setattr(
f"{_BULK_PKG}.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
AsyncMock(),
)
if hash_identity:
# Tests use already-hashed tokens; the raw-sk regression opts out.
monkeypatch.setattr(f"{_BULK_PKG}._hash_token_if_needed", lambda token: token)
return mock_prisma
async def _call_as_admin(data):
from litellm.proxy.management_endpoints.key_management_endpoints import (
bulk_update_team_keys,
)
return await bulk_update_team_keys(
data=data, user_api_key_dict=_admin(), litellm_changed_by=None
)
# ---- happy paths ----------------------------------------------------------
@pytest.mark.asyncio
async def test_bulk_update_team_keys_success_with_key_ids(monkeypatch):
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
keys = [_make_team_key("tok-a"), _make_team_key("tok-b")]
find_unique = AsyncMock(side_effect=keys)
mock = _setup_team_keys_mocks(
monkeypatch,
find_many=keys,
find_unique=find_unique,
update_data=AsyncMock(
side_effect=[{"data": _updated({"max_budget": 50.0})}] * 2
),
)
response = await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-abc",
key_ids=["tok-a", "tok-b"],
update_fields=KeyUpdateFields(max_budget=50.0),
)
)
assert len(response.successful_updates) == 2
assert len(response.failed_updates) == 0
where = mock.db.litellm_verificationtoken.find_many.await_args.kwargs["where"]
assert where["team_id"] == "team-abc"
assert where["token"] == {"in": ["tok-a", "tok-b"]}
find_unique.assert_not_called()
@pytest.mark.asyncio
async def test_bulk_update_team_keys_success_all_keys_in_team(monkeypatch):
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
keys = [_make_team_key(f"tok-{i}") for i in range(3)]
find_unique = AsyncMock(side_effect=keys)
mock = _setup_team_keys_mocks(
monkeypatch,
find_many=keys,
find_unique=find_unique,
update_data=AsyncMock(
side_effect=[{"data": _updated({"max_budget": 50.0})}] * 3
),
)
response = await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-abc",
all_keys_in_team=True,
update_fields=KeyUpdateFields(max_budget=50.0),
)
)
assert len(response.successful_updates) == 3
where = mock.db.litellm_verificationtoken.find_many.await_args.kwargs["where"]
# `blocked` is Boolean? with no default → /key/generate writes NULL. Prisma's
# NOT excludes NULLs, so the filter has to OR `false` with `null` explicitly.
blocked_or, expires_or = where["AND"][0]["OR"], where["AND"][1]["OR"]
assert {"blocked": False} in blocked_or and {"blocked": None} in blocked_or
assert {"expires": None} in expires_or
assert any(
"gt" in c.get("expires", {})
for c in expires_or
if isinstance(c.get("expires"), dict)
)
find_unique.assert_not_called()
@pytest.mark.asyncio
async def test_bulk_update_team_keys_key_not_in_team(monkeypatch):
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
in_team = _make_team_key("tok-a")
_setup_team_keys_mocks(
monkeypatch,
find_many=[in_team],
find_unique=AsyncMock(return_value=in_team),
update_data=AsyncMock(return_value={"data": _updated({"max_budget": 50.0})}),
)
response = await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-abc",
key_ids=["tok-a", "tok-foreign"],
update_fields=KeyUpdateFields(max_budget=50.0),
)
)
assert [u.key for u in response.successful_updates] == ["tok-a"]
assert [u.key for u in response.failed_updates] == ["tok-foreign"]
assert "not found in team" in response.failed_updates[0].failed_reason
# ---- error paths ----------------------------------------------------------
@pytest.mark.asyncio
async def test_bulk_update_team_keys_batch_size_cap(monkeypatch):
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
_setup_team_keys_mocks(
monkeypatch,
find_many=[_make_team_key(f"tok-{i}") for i in range(501)],
)
with pytest.raises(HTTPException) as exc:
await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-abc",
all_keys_in_team=True,
update_fields=KeyUpdateFields(max_budget=50.0),
)
)
assert exc.value.status_code == 400
assert "more than 500" in exc.value.detail["error"]
@pytest.mark.asyncio
async def test_bulk_update_team_keys_empty_team_returns_404(monkeypatch):
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
_setup_team_keys_mocks(monkeypatch, find_many=[])
with pytest.raises(HTTPException) as exc:
await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-empty",
all_keys_in_team=True,
update_fields=KeyUpdateFields(max_budget=50.0),
)
)
assert exc.value.status_code == 404
# ---- auth -----------------------------------------------------------------
@pytest.mark.asyncio
async def test_bulk_update_team_keys_team_member_with_permission(monkeypatch):
from litellm.proxy.management_endpoints.key_management_endpoints import (
bulk_update_team_keys,
)
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
key_a = _make_team_key("tok-a")
_setup_team_keys_mocks(
monkeypatch,
find_many=[key_a],
find_unique=AsyncMock(return_value=key_a),
update_data=AsyncMock(return_value={"data": _updated({"max_budget": 50.0})}),
)
auth_check = AsyncMock()
monkeypatch.setattr(
f"{_BULK_PKG}.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
auth_check,
)
response = await bulk_update_team_keys(
data=BulkUpdateTeamKeysRequest(
team_id="team-abc",
all_keys_in_team=True,
update_fields=KeyUpdateFields(max_budget=50.0),
),
user_api_key_dict=_internal_user(),
litellm_changed_by=None,
)
assert len(response.successful_updates) == 1
# Upfront check + per-key check inside _process_single_key_update
assert auth_check.await_count == 2
@pytest.mark.asyncio
async def test_bulk_update_team_keys_team_member_no_permission(monkeypatch):
from litellm.proxy.management_endpoints.key_management_endpoints import (
bulk_update_team_keys,
)
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
mock = _setup_team_keys_mocks(monkeypatch, find_many=[_make_team_key("tok-a")])
monkeypatch.setattr(
f"{_BULK_PKG}.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
AsyncMock(
side_effect=ProxyException(
message="not in team",
type="team_member_permission_error",
param="/key/update",
code=401,
)
),
)
with pytest.raises(ProxyException):
await bulk_update_team_keys(
data=BulkUpdateTeamKeysRequest(
team_id="team-abc",
all_keys_in_team=True,
update_fields=KeyUpdateFields(max_budget=1.0),
),
user_api_key_dict=_internal_user(),
litellm_changed_by=None,
)
mock.update_data.assert_not_called()
# ---- pydantic-layer validation -------------------------------------------
def test_bulk_update_team_keys_request_validation():
"""Allowlist (extra='forbid'), empty-payload rejection, and selection XOR."""
from pydantic import ValidationError
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
forbidden = [
"key",
"key_alias",
"team_id",
"allowed_routes",
"allowed_passthrough_routes",
"permissions",
"object_permission",
"access_group_ids",
"user_id",
"organization_id",
"blocked",
"key_type",
"models",
"config",
"router_settings",
"spend",
]
for f in forbidden:
with pytest.raises(ValidationError, match=f):
KeyUpdateFields(**{f: True})
with pytest.raises(ValidationError, match="at least one"):
KeyUpdateFields()
assert KeyUpdateFields(max_budget=50.0, tags=["x"]).max_budget == 50.0
valid = KeyUpdateFields(max_budget=10)
with pytest.raises(ValidationError):
BulkUpdateTeamKeysRequest(
team_id="t", key_ids=["k"], all_keys_in_team=True, update_fields=valid
)
with pytest.raises(ValidationError):
BulkUpdateTeamKeysRequest(team_id="t", update_fields=valid)
# ---- security regressions ------------------------------------------------
@pytest.mark.asyncio
async def test_bulk_update_team_keys_hashes_raw_sk_key_ids(monkeypatch):
"""Regression: raw sk-... key_ids must be hashed before the find_many lookup."""
from litellm.proxy._types import hash_token
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
raw_sk = "sk-rawkey1234567890"
hashed = hash_token(raw_sk)
row = LiteLLM_VerificationToken(
token=hashed, user_id="u", models=[], team_id="team-abc", max_budget=None
)
mock = _setup_team_keys_mocks(
monkeypatch,
find_many=[row],
find_unique=AsyncMock(return_value=row),
update_data=AsyncMock(return_value={"data": _updated({"max_budget": 50.0})}),
hash_identity=False,
)
response = await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-abc",
key_ids=[raw_sk],
update_fields=KeyUpdateFields(max_budget=50.0),
)
)
where = mock.db.litellm_verificationtoken.find_many.await_args.kwargs["where"]
assert where["token"] == {"in": [hashed]}
# Response reports the user-supplied form, not the hash.
assert response.successful_updates[0].key == raw_sk
@pytest.mark.asyncio
async def test_bulk_update_team_keys_auth_check_runs_when_no_keys_match(monkeypatch):
"""Regression: non-admin with bogus key_ids must still hit the membership gate."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
bulk_update_team_keys,
)
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
mock = _setup_team_keys_mocks(monkeypatch, find_many=[])
auth_check = AsyncMock(
side_effect=ProxyException(
message="not in team",
type="team_member_permission_error",
param="/key/update",
code=401,
)
)
monkeypatch.setattr(
f"{_BULK_PKG}.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
auth_check,
)
with pytest.raises(ProxyException):
await bulk_update_team_keys(
data=BulkUpdateTeamKeysRequest(
team_id="victim-team",
key_ids=["bogus-1", "bogus-2"],
update_fields=KeyUpdateFields(max_budget=1.0),
),
user_api_key_dict=_internal_user(),
litellm_changed_by=None,
)
# Anchored on data.team_id, not existing_keys[0].
assert auth_check.await_args.kwargs["existing_key_row"].team_id == "victim-team"
mock.update_data.assert_not_called()
@pytest.mark.asyncio
async def test_bulk_update_team_keys_does_not_log_raw_sk_token_on_failure(
monkeypatch, caplog
):
"""Regression: per-key failure must not log the raw sk-... (ERROR-level logs persist)."""
import logging
from litellm.proxy._types import hash_token
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
raw_sk = "sk-supersecret1234567890"
row = LiteLLM_VerificationToken(
token=hash_token(raw_sk),
user_id="u",
models=[],
team_id="team-abc",
max_budget=None,
)
_setup_team_keys_mocks(
monkeypatch,
find_many=[row],
update_data=AsyncMock(side_effect=RuntimeError("boom")),
hash_identity=False,
)
with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"):
response = await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-abc",
key_ids=[raw_sk],
update_fields=KeyUpdateFields(max_budget=50.0),
)
)
assert len(response.failed_updates) == 1
log_text = "\n".join(r.getMessage() for r in caplog.records)
assert raw_sk not in log_text
@pytest.mark.asyncio
async def test_bulk_update_team_keys_propagates_team_id_to_per_key_request(monkeypatch):
"""Regression: per-key UpdateKeyRequest carries data.team_id (gates _check_team_key_limits)."""
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
_setup_team_keys_mocks(monkeypatch, find_many=[_make_team_key("tok-a")])
captured = []
async def fake_process(*, update_key_request, **kw):
captured.append(update_key_request)
return {"max_budget": update_key_request.max_budget}
monkeypatch.setattr(f"{_BULK_PKG}._process_single_key_update", fake_process)
await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-abc",
key_ids=["tok-a"],
update_fields=KeyUpdateFields(
tpm_limit=10_000, tpm_limit_type="guaranteed_throughput"
),
)
)
assert captured[0].team_id == "team-abc"
assert captured[0].tpm_limit_type == "guaranteed_throughput"
@pytest.mark.asyncio
async def test_bulk_update_team_keys_dedupes_key_ids(monkeypatch):
"""Duplicate key_ids collapse to a single update (no redundant DB writes, no inflated counts)."""
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
key_a = _make_team_key("tok-a")
update_data = AsyncMock(return_value={"data": _updated({"max_budget": 50.0})})
_setup_team_keys_mocks(
monkeypatch,
find_many=[key_a],
find_unique=AsyncMock(return_value=key_a),
update_data=update_data,
)
response = await _call_as_admin(
BulkUpdateTeamKeysRequest(
team_id="team-abc",
key_ids=["tok-a", "tok-a", "tok-a"],
update_fields=KeyUpdateFields(max_budget=50.0),
)
)
assert response.total_requested == 1
assert len(response.successful_updates) == 1
assert len(response.failed_updates) == 0
update_data.assert_awaited_once()
@pytest.mark.asyncio
async def test_bulk_update_team_keys_blocks_metadata_allowed_passthrough_routes(
monkeypatch,
):
"""Non-admin can't grant passthrough access by smuggling allowed_passthrough_routes through metadata."""
from fastapi import HTTPException
from litellm.proxy.management_endpoints.key_management_endpoints import (
bulk_update_team_keys,
)
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateTeamKeysRequest,
KeyUpdateFields,
)
mock = _setup_team_keys_mocks(monkeypatch, find_many=[_make_team_key("tok-a")])
request = BulkUpdateTeamKeysRequest(
team_id="team-abc",
all_keys_in_team=True,
update_fields=KeyUpdateFields(
metadata={"allowed_passthrough_routes": ["/admin/*"]}
),
)
with pytest.raises(HTTPException) as exc:
await bulk_update_team_keys(
data=request,
user_api_key_dict=_internal_user(),
litellm_changed_by=None,
)
assert exc.value.status_code == 403
assert "allowed_passthrough_routes" in str(exc.value.detail)
mock.update_data.assert_not_called()

View file

@ -380,6 +380,408 @@ async def test_list_tags_no_dynamic_tags():
app.dependency_overrides.clear()
async def test_internal_user_list_tags_only_returns_tags_used_by_their_keys():
"""
Internal users can view tag usage, but the tag list must be scoped to tags
produced by API keys owned by the caller.
"""
from datetime import datetime
from unittest.mock import AsyncMock, Mock
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
api_key="current-owned-key",
user_id="internal-user-123",
user_role=LitellmUserRoles.INTERNAL_USER,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_db = Mock()
mock_prisma.db = mock_db
owned_key_record = Mock()
owned_key_record.token = "owned-key"
mock_db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[owned_key_record]
)
mock_db.litellm_dailytagspend.group_by = AsyncMock(
return_value=[
{
"tag": "stored-owned-tag",
"_min": {"created_at": "2025-02-01T00:00:00Z"},
"_max": {"updated_at": "2025-03-01T00:00:00Z"},
},
{
"tag": "dynamic-owned-tag",
"_min": {"created_at": "2025-02-02T00:00:00Z"},
"_max": {"updated_at": "2025-03-02T00:00:00Z"},
},
]
)
stored_tag = Mock()
stored_tag.tag_name = "stored-owned-tag"
stored_tag.description = "A stored tag used by the caller"
stored_tag.models = ["model-1"]
stored_tag.model_info = {}
stored_tag.spend = 0.0
stored_tag.budget_id = None
stored_tag.created_at = datetime(2025, 1, 1)
stored_tag.updated_at = datetime(2025, 1, 1)
stored_tag.created_by = "admin-user"
stored_tag.litellm_budget_table = None
mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[stored_tag])
response = client.get(
"/tag/list",
headers={"Authorization": "Bearer test-key"},
)
assert response.status_code == 200
assert [tag["name"] for tag in response.json()] == [
"stored-owned-tag",
"dynamic-owned-tag",
]
mock_db.litellm_verificationtoken.find_many.assert_awaited_once_with(
where={"user_id": "internal-user-123"},
select={"token": True},
)
mock_db.litellm_dailytagspend.group_by.assert_awaited_once_with(
by=["tag"],
where={
"tag": {"not": None},
"api_key": {"in": ["current-owned-key", "owned-key"]},
},
min={"created_at": True},
max={"updated_at": True},
)
mock_db.litellm_tagtable.find_many.assert_awaited_once_with(
where={"tag_name": {"in": ["stored-owned-tag", "dynamic-owned-tag"]}},
include={"litellm_budget_table": True},
)
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_list_tags_with_date_range_filters_dynamic_tags():
"""
/tag/list?start_date=...&end_date=... should push the date window into
the dailytagspend group_by WHERE clause so large tables don't get scanned.
"""
from unittest.mock import AsyncMock, Mock
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_db = Mock()
mock_prisma.db = mock_db
mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[])
group_by_mock = AsyncMock(return_value=[])
mock_db.litellm_dailytagspend.group_by = group_by_mock
headers = {"Authorization": "Bearer sk-1234"}
response = client.get(
"/tag/list?start_date=2026-04-01&end_date=2026-04-29",
headers=headers,
)
assert response.status_code == 200
group_by_mock.assert_awaited_once()
where = group_by_mock.await_args.kwargs["where"]
assert where["tag"] == {"not": None}
assert where["date"] == {"gte": "2026-04-01", "lte": "2026-04-29"}
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_internal_user_tag_daily_activity_is_scoped_to_their_keys():
"""
Internal users must not receive proxy-wide tag spend rows when viewing tag
usage daily activity.
"""
from unittest.mock import AsyncMock, Mock
from litellm.proxy.management_endpoints.tag_management_endpoints import (
get_tag_daily_activity,
)
mock_user_auth = UserAPIKeyAuth(
user_id="internal-user-123",
user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
)
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.tag_management_endpoints.get_daily_activity",
new_callable=AsyncMock,
) as mock_get_daily_activity,
):
mock_db = Mock()
mock_prisma.db = mock_db
owned_key_record = Mock()
owned_key_record.token = "owned-key"
mock_db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[owned_key_record]
)
mock_get_daily_activity.return_value = "daily-activity-response"
result = await get_tag_daily_activity(
start_date="2025-01-01",
end_date="2025-01-31",
user_api_key_dict=mock_user_auth,
)
assert result == "daily-activity-response"
mock_get_daily_activity.assert_awaited_once()
assert mock_get_daily_activity.await_args.kwargs["api_key"] == ["owned-key"]
@pytest.mark.asyncio
async def test_internal_user_tag_daily_activity_rejects_unowned_api_key_filter():
"""
If an internal user filters tag usage by an API key they do not own, the
endpoint should return an empty scoped filter instead of exposing that key's
tag spend.
"""
from unittest.mock import AsyncMock, Mock
from litellm.proxy.management_endpoints.tag_management_endpoints import (
get_tag_daily_activity,
)
mock_user_auth = UserAPIKeyAuth(
user_id="internal-user-123",
user_role=LitellmUserRoles.INTERNAL_USER,
)
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.tag_management_endpoints.get_daily_activity",
new_callable=AsyncMock,
) as mock_get_daily_activity,
):
mock_db = Mock()
mock_prisma.db = mock_db
owned_key_record = Mock()
owned_key_record.token = "owned-key"
mock_db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[owned_key_record]
)
result = await get_tag_daily_activity(
start_date="2025-01-01",
end_date="2025-01-31",
api_key="unowned-key",
user_api_key_dict=mock_user_auth,
)
assert result.results == []
assert result.metadata.total_spend == 0
assert result.metadata.total_api_requests == 0
mock_get_daily_activity.assert_not_awaited()
@pytest.mark.asyncio
async def test_internal_user_tag_daily_activity_scopes_to_current_key_without_user_id():
"""
If an internal-user token has no user_id, it should still scope tag usage to
the current request key instead of falling back to proxy-wide tag spend.
"""
from unittest.mock import AsyncMock, Mock
from litellm.proxy.management_endpoints.tag_management_endpoints import (
get_tag_daily_activity,
)
mock_user_auth = UserAPIKeyAuth(
api_key="current-owned-key",
user_id=None,
user_role=LitellmUserRoles.INTERNAL_USER,
)
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.tag_management_endpoints.get_daily_activity",
new_callable=AsyncMock,
) as mock_get_daily_activity,
):
mock_db = Mock()
mock_prisma.db = mock_db
mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_get_daily_activity.return_value = "daily-activity-response"
result = await get_tag_daily_activity(
start_date="2025-01-01",
end_date="2025-01-31",
user_api_key_dict=mock_user_auth,
)
assert result == "daily-activity-response"
mock_db.litellm_verificationtoken.find_many.assert_not_awaited()
mock_get_daily_activity.assert_awaited_once()
assert mock_get_daily_activity.await_args.kwargs["api_key"] == [
"current-owned-key"
]
@pytest.mark.asyncio
async def test_internal_user_tag_daily_activity_without_any_scoped_keys_returns_empty():
"""
If an internal-user token has neither user_id nor api_key, the endpoint must
return an empty response instead of dropping the API key filter.
"""
from unittest.mock import AsyncMock, Mock
from litellm.proxy.management_endpoints.tag_management_endpoints import (
get_tag_daily_activity,
)
mock_user_auth = UserAPIKeyAuth(
user_id=None,
user_role=LitellmUserRoles.INTERNAL_USER,
)
with (
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
patch(
"litellm.proxy.management_endpoints.tag_management_endpoints.get_daily_activity",
new_callable=AsyncMock,
) as mock_get_daily_activity,
):
mock_db = Mock()
mock_prisma.db = mock_db
mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
result = await get_tag_daily_activity(
start_date="2025-01-01",
end_date="2025-01-31",
user_api_key_dict=mock_user_auth,
)
assert result.results == []
assert result.metadata.total_spend == 0
assert result.metadata.total_api_requests == 0
mock_db.litellm_verificationtoken.find_many.assert_not_awaited()
mock_get_daily_activity.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_tag_daily_activity_requires_database_connection():
"""
Tag daily activity should fail with the same explicit DB error used by other
tag endpoints instead of raising an AttributeError during scope resolution.
"""
from litellm.proxy.management_endpoints.tag_management_endpoints import (
get_tag_daily_activity,
)
mock_user_auth = UserAPIKeyAuth(
user_id="internal-user-123",
user_role=LitellmUserRoles.INTERNAL_USER,
)
with patch("litellm.proxy.proxy_server.prisma_client", None):
with pytest.raises(HTTPException) as exc_info:
await get_tag_daily_activity(
start_date="2025-01-01",
end_date="2025-01-31",
user_api_key_dict=mock_user_auth,
)
assert exc_info.value.status_code == 500
assert exc_info.value.detail == "Database not connected"
@pytest.mark.asyncio
async def test_list_tags_without_date_range_omits_date_filter():
"""When no date range is passed, the WHERE clause must not carry a date key."""
from unittest.mock import AsyncMock, Mock
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_db = Mock()
mock_prisma.db = mock_db
mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[])
group_by_mock = AsyncMock(return_value=[])
mock_db.litellm_dailytagspend.group_by = group_by_mock
headers = {"Authorization": "Bearer sk-1234"}
response = client.get("/tag/list", headers=headers)
assert response.status_code == 200
group_by_mock.assert_awaited_once()
where = group_by_mock.await_args.kwargs["where"]
assert "date" not in where
finally:
app.dependency_overrides.clear()
@pytest.mark.parametrize(
"query, expected_detail_fragment",
[
("?start_date=2026-04-01", "must be provided together"),
("?end_date=2026-04-29", "must be provided together"),
("?start_date=2026-04-29&end_date=2026-04-01", "on or before end_date"),
("?start_date=not-a-date&end_date=2026-04-29", "YYYY-MM-DD"),
],
)
@pytest.mark.asyncio
async def test_list_tags_rejects_invalid_date_range(query, expected_detail_fragment):
from unittest.mock import AsyncMock, Mock
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
mock_db = Mock()
mock_prisma.db = mock_db
mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[])
mock_db.litellm_dailytagspend.group_by = AsyncMock(return_value=[])
headers = {"Authorization": "Bearer sk-1234"}
response = client.get(f"/tag/list{query}", headers=headers)
assert response.status_code == 400
assert expected_detail_fragment in response.json()["detail"]
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_get_deployments_by_model_id():
"""

View file

@ -21,6 +21,7 @@ from litellm.proxy.litellm_pre_call_utils import (
_get_enforced_params,
_get_metadata_variable_name,
_resolve_credential_from_model_config,
_resolve_provider_from_deployment,
_update_model_if_key_alias_exists,
add_guardrails_from_policy_engine,
add_litellm_data_to_request,
@ -4043,3 +4044,174 @@ def test_get_guardrail_from_metadata_reads_litellm_metadata_when_no_metadata():
assert result == [
"my-guardrail"
], f"Expected guardrails from litellm_metadata fallback, got: {result}"
# ============================================================================
# Tests for #27516: provider hint resolution from deployment when the
# user-facing model name has no provider prefix.
# ============================================================================
def test_resolve_provider_from_deployment_uses_litellm_params_model():
"""When custom_llm_provider is unset, fall back to the prefix of model."""
router = MagicMock()
deployment = MagicMock()
deployment.litellm_params.model = "bedrock/us.anthropic.claude-sonnet-4-6"
deployment.litellm_params.custom_llm_provider = None
router.get_deployment_by_model_group_name.return_value = deployment
assert (
_resolve_provider_from_deployment(router, "claude-sonnet-4.6") == "bedrock"
)
def test_resolve_provider_from_deployment_prefers_custom_llm_provider():
"""Explicit custom_llm_provider on the deployment wins over model prefix."""
router = MagicMock()
deployment = MagicMock()
deployment.litellm_params.model = "us.anthropic.claude-sonnet-4-6"
deployment.litellm_params.custom_llm_provider = "bedrock"
router.get_deployment_by_model_group_name.return_value = deployment
assert (
_resolve_provider_from_deployment(router, "claude-sonnet-4.6") == "bedrock"
)
def test_resolve_provider_from_deployment_no_match():
"""No deployment for the model group -> None."""
router = MagicMock()
router.get_deployment_by_model_group_name.return_value = None
assert _resolve_provider_from_deployment(router, "unknown-model") is None
def test_resolve_provider_from_deployment_router_raises():
"""Router exceptions must not propagate — fall back to None."""
router = MagicMock()
router.get_deployment_by_model_group_name.side_effect = RuntimeError("boom")
assert _resolve_provider_from_deployment(router, "claude-sonnet-4.6") is None
def test_resolve_provider_from_deployment_falls_back_to_pre_alias():
"""If post-alias lookup fails, the pre-alias name is also tried."""
router = MagicMock()
deployment = MagicMock()
deployment.litellm_params.model = "bedrock/anthropic.claude-sonnet-4-6"
deployment.litellm_params.custom_llm_provider = None
def lookup(model_group_name):
if model_group_name == "pre-alias-name":
return deployment
return None
router.get_deployment_by_model_group_name.side_effect = lookup
result = _resolve_provider_from_deployment(
router, "post-alias-name", pre_alias_model_name="pre-alias-name"
)
assert result == "bedrock"
def test_apply_overrides_multi_provider_default_picks_correct_provider(
setup_test_credentials,
):
"""
Regression for #27516: when defaultconfig has multiple providers and the
request model has no '/' prefix, the deployment's custom_llm_provider must
drive provider matching instead of falling through to dict insertion order.
"""
litellm.credential_list.append(
CredentialItem(
credential_name="bedrock-team-1",
credential_info={},
credential_values={"api_key": "ABSK-bedrock-key-for-team-1"},
)
)
litellm.credential_list.append(
CredentialItem(
credential_name="gemini-team-1",
credential_info={},
credential_values={"api_key": "gemini-key-for-team-1"},
)
)
data = {"model": "claude-sonnet-4.6"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {
# gemini comes first in insertion order — the bug picked it.
"gemini": {"litellm_credentials": "gemini-team-1"},
"bedrock": {"litellm_credentials": "bedrock-team-1"},
}
}
},
)
router = MagicMock()
deployment = MagicMock()
deployment.litellm_params.model = "us.anthropic.claude-sonnet-4-6"
deployment.litellm_params.custom_llm_provider = "bedrock"
router.get_deployment_by_model_group_name.return_value = deployment
_apply_credential_overrides_from_model_config(
data=data,
user_api_key_dict=user_api_key_dict,
llm_router=router,
)
assert data["api_key"] == "ABSK-bedrock-key-for-team-1"
def test_apply_overrides_no_router_keeps_legacy_behaviour(setup_test_credentials):
"""
Without a router, the function still works for the single-provider case
(the historical behaviour). Multi-provider configs with no '/' prefix
keep the legacy first-entry behaviour because there is no way to
disambiguate this preserves backwards compatibility.
"""
data = {"model": "gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {
"azure": {"litellm_credentials": "hotel-azure-eastus"}
}
}
},
)
_apply_credential_overrides_from_model_config(
data=data, user_api_key_dict=user_api_key_dict, llm_router=None
)
assert data["api_base"] == "https://hotel-eastus.openai.azure.com/"
assert data["api_key"] == "key-hotel-eastus"
def test_apply_overrides_provider_prefix_in_model_skips_router_lookup(
setup_test_credentials,
):
"""
When the request model already has a 'provider/...' prefix, the router
lookup must be skipped the explicit prefix is authoritative.
"""
data = {"model": "azure/gpt-4"}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
team_metadata={
"model_config": {
"defaultconfig": {
"azure": {"litellm_credentials": "hotel-azure-eastus"},
"bedrock": {"litellm_credentials": "hotel-rec-azure"},
}
}
},
)
router = MagicMock()
_apply_credential_overrides_from_model_config(
data=data, user_api_key_dict=user_api_key_dict, llm_router=router
)
assert data["api_base"] == "https://hotel-eastus.openai.azure.com/"
assert data["api_key"] == "key-hotel-eastus"
router.get_deployment_by_model_group_name.assert_not_called()

View file

@ -140,3 +140,110 @@ class TestCredentialEndpointsOpenAPISchema:
assert (
"credential_name" in sig.parameters
), "get_credential_by_name must have a credential_name parameter"
class TestWebSocketStubInjection:
"""
Regression test for the v1.82.3 bug where adding a WebSocket route on a path
that already had an HTTP route silently dropped the HTTP operation from the
OpenAPI schema.
Related case: 2026-05-05-madhu-swagger-responses-missing
"""
def _make_fake_ws_route(self, path: str, name: str = "fake_ws"):
"""Minimal stand-in for fastapi.routing.APIWebSocketRoute for the helper's purposes."""
from types import SimpleNamespace
return SimpleNamespace(path=path, name=name, dependant=None)
def test_websocket_stub_does_not_clobber_existing_post(self):
"""
When a WebSocket route shares its path with an existing POST operation,
the POST must survive the WebSocket stub is added alongside, not on top.
"""
from litellm.proxy.proxy_server import (
_inject_websocket_stubs_into_openapi_schema,
)
schema = {
"paths": {
"/v1/responses": {
"post": {"summary": "responses_api", "operationId": "responses_api"}
}
}
}
ws_routes = [self._make_fake_ws_route("/v1/responses", name="responses_ws")]
result = _inject_websocket_stubs_into_openapi_schema(schema, ws_routes)
assert (
"post" in result["paths"]["/v1/responses"]
), "POST operation must be preserved when a WebSocket route shares the path"
assert (
result["paths"]["/v1/responses"]["post"]["operationId"] == "responses_api"
)
assert (
"get" in result["paths"]["/v1/responses"]
), "WebSocket stub should also be added under 'get'"
assert result["paths"]["/v1/responses"]["get"]["tags"] == ["WebSocket"]
def test_websocket_stub_added_when_path_is_new(self):
"""
When a WebSocket route's path is not already in the schema, the stub
creates a fresh entry preserving the original behavior for WebSocket-only
paths.
"""
from litellm.proxy.proxy_server import (
_inject_websocket_stubs_into_openapi_schema,
)
schema = {"paths": {}}
ws_routes = [self._make_fake_ws_route("/ws_only", name="ws_only")]
result = _inject_websocket_stubs_into_openapi_schema(schema, ws_routes)
assert "/ws_only" in result["paths"]
assert "get" in result["paths"]["/ws_only"]
assert result["paths"]["/ws_only"]["get"]["tags"] == ["WebSocket"]
def test_websocket_stub_skipped_when_existing_get(self):
"""
If a real GET is already documented on the path, the WebSocket stub is
skipped a real operation always wins over the synthetic stub. This
closes the same trap for future GET-vs-WebSocket collisions.
"""
from litellm.proxy.proxy_server import (
_inject_websocket_stubs_into_openapi_schema,
)
schema = {
"paths": {
"/health": {
"get": {"summary": "health_check", "operationId": "real_get"}
}
}
}
ws_routes = [self._make_fake_ws_route("/health", name="health_ws")]
result = _inject_websocket_stubs_into_openapi_schema(schema, ws_routes)
assert (
result["paths"]["/health"]["get"]["operationId"] == "real_get"
), "Real GET must take precedence over WebSocket stub"
def test_responses_post_routes_registered_on_router(self):
"""
Sanity check: the three POST routes for the responses API are still wired
on the responses router. Guards against accidental removal at the source.
"""
from litellm.proxy.response_api_endpoints.endpoints import router
post_paths = {
route.path
for route in router.routes
if hasattr(route, "methods")
and "POST" in (route.methods or set())
and route.path in {"/v1/responses", "/responses", "/openai/v1/responses"}
}
assert post_paths == {"/v1/responses", "/responses", "/openai/v1/responses"}

View file

@ -1,6 +1,6 @@
import os
import sys
from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -387,6 +387,102 @@ class TestProxyInitializationHelpers:
), f"exit_code={result.exit_code}, output={result.output}"
mock_uvicorn_run.assert_called_once()
@pytest.mark.parametrize(
"timeout_config,expected_timeout",
[
({"database_connection_timeout": 30}, 30),
({"database_connection_pool_timeout": 45}, 45),
(
{
"database_connection_timeout": 30,
"database_connection_pool_timeout": 45,
},
30,
),
],
)
@patch("subprocess.run")
@patch("atexit.register")
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
@patch(
"litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False
)
def test_db_timeout_settings_are_forwarded_to_pool_timeout(
self,
mock_should_update,
mock_setup_db,
mock_atexit_register,
mock_subprocess_run,
timeout_config,
expected_timeout,
):
from click.testing import CliRunner
from litellm.proxy.proxy_cli import run_server
runner = CliRunner()
mock_subprocess_run.return_value = MagicMock(returncode=0)
mock_proxy_module = MagicMock(
app=MagicMock(),
ProxyConfig=MagicMock(),
KeyManagementSettings=MagicMock(),
save_worker_config=MagicMock(),
)
mock_proxy_module.ProxyConfig.return_value.get_config = AsyncMock(
return_value={
"general_settings": {
"database_url": "postgresql://test:test@localhost:5432/test",
"database_connection_pool_limit": 5,
**timeout_config,
}
}
)
clean_env = {
k: v
for k, v in os.environ.items()
if k not in ("DATABASE_URL", "DIRECT_URL")
}
with (
patch.dict(os.environ, clean_env, clear=True),
patch.dict(
"sys.modules",
{
"proxy_server": mock_proxy_module,
"litellm.proxy.proxy_server": mock_proxy_module,
},
),
patch(
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
) as mock_get_args,
patch(
"litellm.proxy.proxy_cli.append_query_params",
side_effect=lambda url, params: (
f"{url}?connection_limit={params['connection_limit']}&pool_timeout={params['pool_timeout']}"
),
) as mock_append_query_params,
):
mock_get_args.return_value = {
"app": "litellm.proxy.proxy_server:app",
"host": "localhost",
"port": 8000,
}
result = runner.invoke(
run_server,
["--local", "--config", "test-config.yaml", "--skip_server_startup"],
)
assert (
result.exit_code == 0
), f"exit_code={result.exit_code}, output={result.output}"
mock_append_query_params.assert_called()
appended_params = mock_append_query_params.call_args.args[1]
assert appended_params["connection_limit"] == 5
assert appended_params["pool_timeout"] == expected_timeout
@patch("uvicorn.run")
@patch("atexit.register")
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")

View file

@ -322,13 +322,13 @@ class TestSharedHealthCheckManager:
async def test_perform_shared_health_check_lock_failed_then_cache(
self, shared_health_manager, mock_redis_cache
):
"""Test performing shared health check when lock fails but cache becomes available"""
"""Test performing shared health check when lock fails but cache becomes available during polling"""
# First call: no cache, lock fails
# Second call: cache available
# Polling finds cache on first iteration
mock_redis_cache.async_get_cache.side_effect = [
None, # No cache initially
None, # No cache initially (get_cached_health_check_results)
json.dumps(
{ # Cache available after waiting
{ # Cache available on first poll iteration
"healthy_endpoints": [{"model": "cached-model"}],
"unhealthy_endpoints": [],
"healthy_count": 1,
@ -350,18 +350,68 @@ class TestSharedHealthCheckManager:
)
)
# Should wait and then get cached results
mock_sleep.assert_called_once_with(2)
# Should poll once (5s interval) and find cached results
mock_sleep.assert_called_once_with(5)
assert healthy == [{"model": "cached-model"}]
assert unhealthy == []
@pytest.mark.asyncio
async def test_perform_shared_health_check_fallback(
async def test_perform_shared_health_check_fallback(self, mock_redis_cache):
"""Test performing shared health check with fallback to local health check"""
# Use short lock_ttl so the polling loop only runs 2 iterations
manager = SharedHealthCheckManager(
redis_cache=mock_redis_cache,
health_check_ttl=300,
lock_ttl=10,
)
# No cache ever, lock always held by another pod
mock_redis_cache.async_get_cache.side_effect = [
None, # Initial cache check
None, # Iteration 1: cache check
"other_pod", # Iteration 1: lock check (still held)
None, # Iteration 2: cache check
"other_pod", # Iteration 2: lock check (still held)
]
mock_redis_cache.async_set_cache.return_value = False # Lock acquisition fails
model_list = [
{"model_name": "test-model", "litellm_params": {"model": "test-model"}}
]
expected_healthy = [{"model": "test-model", "status": "healthy"}]
expected_unhealthy = []
with (
patch("asyncio.sleep") as mock_sleep,
patch(
"litellm.proxy.health_check_utils.shared_health_check_manager.perform_health_check"
) as mock_perform,
):
mock_perform.return_value = (expected_healthy, expected_unhealthy, {})
healthy, unhealthy, _ = await manager.perform_shared_health_check(
model_list, details=True
)
# Should poll twice (5s * 2 = 10s >= lock_ttl) then fall back
assert mock_sleep.call_count == 2
mock_sleep.assert_called_with(5)
mock_perform.assert_called_once_with(
model_list=model_list, details=True, max_concurrency=None
)
assert healthy == expected_healthy
assert unhealthy == expected_unhealthy
@pytest.mark.asyncio
async def test_perform_shared_health_check_early_exit_orphaned_lock(
self, shared_health_manager, mock_redis_cache
):
"""Test performing shared health check with fallback to local health check"""
# No cache, lock fails, no cache after waiting
mock_redis_cache.async_get_cache.return_value = None
"""Test that polling exits early when the lock disappears without a cache write (crash recovery)"""
mock_redis_cache.async_get_cache.side_effect = [
None, # Initial cache check
None, # Iteration 1: cache check (still no cache)
None, # Iteration 1: lock check -> lock gone (holder crashed)
]
mock_redis_cache.async_set_cache.return_value = False # Lock acquisition fails
model_list = [
@ -384,8 +434,77 @@ class TestSharedHealthCheckManager:
)
)
# Should fall back to local health check
mock_sleep.assert_called_once_with(2)
# Should detect orphaned lock after 1 iteration and fall back immediately
mock_sleep.assert_called_once_with(5)
mock_perform.assert_called_once_with(
model_list=model_list, details=True, max_concurrency=None
)
assert healthy == expected_healthy
assert unhealthy == expected_unhealthy
@pytest.mark.asyncio
async def test_perform_shared_health_check_redis_error_during_polling(
self, shared_health_manager, mock_redis_cache
):
"""Test that a transient Redis error during lock polling doesn't crash the loop"""
cached_data = json.dumps(
{
"healthy_endpoints": [{"model": "cached-model"}],
"unhealthy_endpoints": [],
"healthy_count": 1,
"unhealthy_count": 0,
"timestamp": time.time() - 100,
}
)
mock_redis_cache.async_get_cache.side_effect = [
None, # Initial cache check
None, # Iteration 1: cache check
Exception("Redis connection lost"), # Iteration 1: lock check errors
cached_data, # Iteration 2: cache check -> found!
]
mock_redis_cache.async_set_cache.return_value = False # Lock acquisition fails
model_list = [
{"model_name": "test-model", "litellm_params": {"model": "test-model"}}
]
with patch("asyncio.sleep") as mock_sleep:
healthy, unhealthy, _ = (
await shared_health_manager.perform_shared_health_check(
model_list, details=True
)
)
# Should survive the Redis error and find cache on iteration 2
assert mock_sleep.call_count == 2
assert healthy == [{"model": "cached-model"}]
assert unhealthy == []
@pytest.mark.asyncio
async def test_perform_shared_health_check_no_redis_skips_polling(self):
"""Test that polling is skipped entirely when redis_cache is None"""
manager = SharedHealthCheckManager(redis_cache=None)
model_list = [
{"model_name": "test-model", "litellm_params": {"model": "test-model"}}
]
expected_healthy = [{"model": "test-model", "status": "healthy"}]
expected_unhealthy = []
with (
patch("asyncio.sleep") as mock_sleep,
patch(
"litellm.proxy.health_check_utils.shared_health_check_manager.perform_health_check"
) as mock_perform,
):
mock_perform.return_value = (expected_healthy, expected_unhealthy, {})
healthy, unhealthy, _ = await manager.perform_shared_health_check(
model_list, details=True
)
# Should NOT sleep at all — falls back to local health check immediately
mock_sleep.assert_not_called()
mock_perform.assert_called_once_with(
model_list=model_list, details=True, max_concurrency=None
)

View file

@ -757,6 +757,60 @@ def test_responses_api_bridge_check_gpt_5_4_tools_without_reasoning_stays_chat()
assert model_info.get("mode") != "responses"
def test_responses_api_bridge_check_gpt_5_4_reasoning_summary_without_tools_routes_to_responses():
"""gpt-5.4+ with reasoning_effort + reasoningSummary but no tools should bridge (AI SDK)."""
from litellm.main import responses_api_bridge_check
with patch("litellm.main._get_model_info_helper") as mock_get_model_info:
mock_get_model_info.return_value = {"max_tokens": 128000}
model_info, model = responses_api_bridge_check(
model="gpt-5.4",
custom_llm_provider="openai",
tools=None,
reasoning_effort="medium",
reasoning_summary="auto",
)
assert model == "gpt-5.4"
assert model_info.get("mode") == "responses"
def test_responses_api_bridge_check_gpt_5_reasoning_summary_routes_to_responses():
"""Bare ``gpt-5`` with reasoning_effort + reasoningSummary should bridge (not 5.4+)."""
from litellm.main import responses_api_bridge_check
with patch("litellm.main._get_model_info_helper") as mock_get_model_info:
mock_get_model_info.return_value = {"max_tokens": 128000}
model_info, model = responses_api_bridge_check(
model="gpt-5",
custom_llm_provider="openai",
tools=None,
reasoning_effort="medium",
reasoning_summary="auto",
)
assert model == "gpt-5"
assert model_info.get("mode") == "responses"
def test_responses_api_bridge_check_gpt_5_tools_without_summary_stays_chat():
"""gpt-5 with tools + reasoning_effort but no summary should stay on chat."""
from litellm.main import responses_api_bridge_check
with patch("litellm.main._get_model_info_helper") as mock_get_model_info:
mock_get_model_info.return_value = {"max_tokens": 128000}
model_info, model = responses_api_bridge_check(
model="gpt-5",
custom_llm_provider="openai",
tools=[{"type": "function", "function": {"name": "get_capital"}}],
reasoning_effort="medium",
reasoning_summary=None,
)
assert model == "gpt-5"
assert model_info.get("mode") != "responses"
@patch("litellm.completion_extras.responses_api_bridge.completion")
def test_gpt_5_4_responses_bridge_preserves_reasoning_summary_dict(
mock_responses_completion,
@ -794,6 +848,93 @@ def test_gpt_5_4_responses_bridge_preserves_reasoning_summary_dict(
}
@patch("litellm.completion_extras.responses_api_bridge.completion")
def test_gpt_5_4_responses_bridge_merges_reasoning_summary_kwarg_without_tools(
mock_responses_completion,
):
"""reasoningSummary without tools should route and merge into reasoning_effort dict."""
mock_responses_completion.return_value = MagicMock()
import litellm
litellm.completion(
model="gpt-5.4",
messages=[{"role": "user", "content": "ok"}],
reasoning_effort="medium",
reasoningSummary="auto",
api_key="fake-key",
)
assert mock_responses_completion.called is True
optional_params = mock_responses_completion.call_args.kwargs["optional_params"]
assert optional_params["reasoning_effort"] == {
"effort": "medium",
"summary": "auto",
}
assert "reasoningSummary" not in optional_params
assert "reasoning_summary" not in optional_params
@patch("litellm.completion_extras.responses_api_bridge.completion")
def test_responses_bridge_preserves_reasoning_summary_without_effort(
mock_responses_completion,
):
"""Reasoning summary should survive responses routing even without effort."""
mock_responses_completion.return_value = MagicMock()
import litellm
with patch.object(litellm, "route_all_chat_openai_to_responses", True):
litellm.completion(
model="gpt-4o",
messages=[{"role": "user", "content": "ok"}],
reasoningSummary="auto",
api_key="fake-key",
)
assert mock_responses_completion.called is True
optional_params = mock_responses_completion.call_args.kwargs["optional_params"]
assert optional_params["reasoning_effort"] == {"summary": "auto"}
assert "reasoningSummary" not in optional_params
assert "reasoning_summary" not in optional_params
@patch("litellm.completion_extras.responses_api_bridge.completion")
def test_gpt_5_responses_bridge_tools_and_reasoning_summary(
mock_responses_completion,
):
"""Bare gpt-5 with tools + reasoningSummary should bridge (OpenCode-style)."""
mock_responses_completion.return_value = MagicMock()
import litellm
litellm.completion(
model="gpt-5",
messages=[{"role": "user", "content": "ok"}],
tools=[
{
"type": "function",
"function": {
"name": "apply_patch",
"parameters": {"type": "object", "properties": {}},
},
}
],
tool_choice="auto",
reasoning_effort="medium",
reasoningSummary="auto",
stream=True,
api_key="fake-key",
)
assert mock_responses_completion.called is True
optional_params = mock_responses_completion.call_args.kwargs["optional_params"]
assert optional_params.get("reasoning_effort") == {
"effort": "medium",
"summary": "auto",
}
def test_responses_api_bridge_check_handles_exception():
"""Test that responses_api_bridge_check handles exceptions and still processes responses/ models."""
from litellm.main import responses_api_bridge_check

View file

@ -0,0 +1,13 @@
from pathlib import Path
def test_main_py_starts_with_brief_file_description():
repo_root = Path(__file__).resolve().parents[2]
main_py = repo_root / "litellm" / "main.py"
first_two_lines = main_py.read_text(encoding="utf-8").splitlines()[:2]
assert any(
"LiteLLM main module" in line and "entrypoints" in line
for line in first_two_lines
)

View file

@ -0,0 +1,162 @@
"""Cover the Bedrock-ARN dispatch in ``litellm.batches.main.retrieve_batch``.
The dispatch picks one of two Bedrock handlers depending on the ARN
family in ``batch_id``:
* ``:async-invoke/<id>`` -> ``_handle_async_invoke_status`` (data plane)
* ``:model-invocation-job/<id>`` -> ``_handle_model_invocation_job_status``
(control plane, added in this PR)
Anything else falls through to the generic ``provider_config`` retrieve
flow. We mock the two handlers so the tests don't hit AWS — the focus
here is purely the dispatch logic that lives in ``main.py``.
"""
from __future__ import annotations
import os
import sys
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../.."))
import litellm # noqa: E402
ASYNC_INVOKE_ARN = "arn:aws:bedrock:us-west-2:123456789012:async-invoke/abc123def456"
MIJ_ARN = "arn:aws:bedrock:us-west-2:123456789012:model-invocation-job/abc1234567"
@pytest.fixture
def mock_handlers():
"""Patch both Bedrock retrieve handlers and yield the mocks.
We patch at the import site (litellm.batches.main) rather than the
definition site so the ``BedrockBatchesHandler`` reference inside
``retrieve_batch`` resolves to our mocks.
"""
fake_batch = MagicMock(name="LiteLLMBatch")
with (
patch(
"litellm.batches.main.BedrockBatchesHandler._handle_async_invoke_status",
return_value=fake_batch,
) as async_invoke,
patch(
"litellm.batches.main.BedrockBatchesHandler._handle_model_invocation_job_status",
return_value=fake_batch,
) as mij,
):
yield async_invoke, mij, fake_batch
def test_async_invoke_arn_routes_to_async_invoke_handler(mock_handlers):
"""``:async-invoke/`` ARNs go to the data-plane handler."""
async_invoke, mij, fake_batch = mock_handlers
result = litellm.retrieve_batch(
batch_id=ASYNC_INVOKE_ARN,
custom_llm_provider="bedrock",
aws_region_name="us-west-2",
)
assert result is fake_batch
async_invoke.assert_called_once()
mij.assert_not_called()
call_kwargs = async_invoke.call_args.kwargs
assert call_kwargs["batch_id"] == ASYNC_INVOKE_ARN
assert call_kwargs["aws_region_name"] == "us-west-2"
# Region must be stripped from the forwarded kwargs to avoid TypeError
# (it's already an explicit positional/keyword arg).
assert "aws_region_name" not in {
k
for k in call_kwargs
if k not in {"batch_id", "aws_region_name", "logging_obj"}
}
def test_async_invoke_arn_falls_back_to_default_region_when_unset(mock_handlers):
"""If no ``aws_region_name`` is passed, the data-plane handler defaults
to ``us-east-1`` (preserving prior behavior on this branch)."""
async_invoke, _mij, _ = mock_handlers
litellm.retrieve_batch(
batch_id=ASYNC_INVOKE_ARN,
custom_llm_provider="bedrock",
)
async_invoke.assert_called_once()
assert async_invoke.call_args.kwargs["aws_region_name"] == "us-east-1"
def test_model_invocation_job_arn_routes_to_mij_handler(mock_handlers):
"""``:model-invocation-job/`` ARNs go to the new control-plane handler."""
_async_invoke, mij, fake_batch = mock_handlers
result = litellm.retrieve_batch(
batch_id=MIJ_ARN,
custom_llm_provider="bedrock",
aws_region_name="us-west-2",
)
assert result is fake_batch
mij.assert_called_once()
_async_invoke.assert_not_called()
call_kwargs = mij.call_args.kwargs
assert call_kwargs["batch_id"] == MIJ_ARN
assert call_kwargs["aws_region_name"] == "us-west-2"
def test_model_invocation_job_arn_with_no_region_passes_none(mock_handlers):
"""The MIJ handler is responsible for sniffing region from the ARN
when none is explicitly provided. Dispatch must forward ``None``
rather than substituting a default otherwise per-region jobs in
other AWS regions would silently route to ``us-east-1``."""
_async_invoke, mij, _ = mock_handlers
litellm.retrieve_batch(
batch_id=MIJ_ARN,
custom_llm_provider="bedrock",
)
mij.assert_called_once()
assert mij.call_args.kwargs["aws_region_name"] is None
def test_unrelated_bedrock_arn_falls_through_to_provider_config(mock_handlers):
"""Bedrock ARNs that aren't async-invoke or model-invocation-job
must NOT hit either special handler they should fall through to
the existing generic provider_config path. We don't fully exercise
that path here (it requires a real provider config); we just assert
neither special handler is invoked."""
async_invoke, mij, _ = mock_handlers
# Use a plausible-but-unsupported Bedrock ARN family.
unrelated_arn = "arn:aws:bedrock:us-west-2:123456789012:provisioned-model/xyz"
with pytest.raises(Exception):
# Will raise because no provider_config exists for this path —
# that's fine, we just need to assert neither bedrock handler ran
# before the failure.
litellm.retrieve_batch(
batch_id=unrelated_arn,
custom_llm_provider="bedrock",
)
async_invoke.assert_not_called()
mij.assert_not_called()
def test_non_bedrock_id_skips_bedrock_dispatch_entirely(mock_handlers):
"""Plain (non-ARN) batch ids must not even enter the Bedrock dispatch
block they belong to other providers' retrieve flows."""
async_invoke, mij, _ = mock_handlers
with pytest.raises(Exception):
litellm.retrieve_batch(
batch_id="batch_abc123",
custom_llm_provider="openai",
)
async_invoke.assert_not_called()
mij.assert_not_called()

View file

@ -18,6 +18,7 @@ sys.path.insert(
import litellm
from litellm import Router
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
def test_should_not_pollute_shared_key_with_zero_cost_pricing():
@ -266,3 +267,59 @@ def test_should_preserve_builtin_pricing_regardless_of_deployment_order():
f"Order should not matter. Expected {builtin_output_cost}, "
f"got {info_std_2['output_cost_per_token']}"
)
def test_responses_prefix_stripped_alias_registered_for_model_list():
"""
Register ``litellm.model_cost`` under the backend key with ``responses/`` and
under the stripped key (``responses_api_bridge_check`` removes that segment).
"""
uid = "responses-strip-alias-test-a1b2c3d4"
Router(
model_list=[
{
"model_name": "azure-responses-strip-test",
"litellm_params": {
"model": "responses/gpt-strip-test-a1b2c3d4",
"custom_llm_provider": "azure",
"api_key": "fake-key-strip",
},
"model_info": {
"id": uid,
"supports_native_streaming": True,
},
}
],
)
assert "azure/responses/gpt-strip-test-a1b2c3d4" in litellm.model_cost
assert "azure/gpt-strip-test-a1b2c3d4" in litellm.model_cost
assert (
litellm.model_cost["azure/gpt-strip-test-a1b2c3d4"].get(
"supports_native_streaming"
)
is True
)
def test_responses_prefix_stripped_alias_registered_for_add_deployment():
"""Dynamic ``add_deployment`` must mirror ``_create_deployment`` registration."""
uid = "add-dep-responses-strip-e5f6a7b8"
router = Router(model_list=[])
deployment = Deployment(
model_name="dyn-responses-strip",
litellm_params=LiteLLM_Params(
model="responses/gpt-add-strip-e5f6a7b8",
custom_llm_provider="azure",
api_key="fake-key-add",
),
model_info=ModelInfo(id=uid, supports_native_streaming=True),
)
router.add_deployment(deployment=deployment)
assert "azure/responses/gpt-add-strip-e5f6a7b8" in litellm.model_cost
assert "azure/gpt-add-strip-e5f6a7b8" in litellm.model_cost
assert (
litellm.model_cost["azure/gpt-add-strip-e5f6a7b8"].get(
"supports_native_streaming"
)
is True
)

View file

@ -2817,6 +2817,128 @@ def test_generate_gcp_iam_access_token_import_error():
assert "pip install google-cloud-iam" in str(exc_info.value)
def test_generate_azure_ad_redis_token():
"""Test _generate_azure_ad_redis_token with mocked Azure credential."""
from unittest.mock import Mock, patch
expected_token = "azure-access-token-12345"
mock_token = Mock()
mock_token.token = expected_token
mock_credential = Mock()
mock_credential.get_token.return_value = mock_token
mock_azure_identity = Mock()
mock_azure_identity.DefaultAzureCredential = Mock(return_value=mock_credential)
mock_azure_identity.ClientSecretCredential = Mock()
mock_azure_identity.ManagedIdentityCredential = Mock()
with patch.dict(
"sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()}
):
from litellm._redis import _generate_azure_ad_redis_token
result = _generate_azure_ad_redis_token()
assert result == expected_token
mock_credential.get_token.assert_called_once_with(
"https://redis.azure.com/.default"
)
def test_generate_azure_ad_redis_token_service_principal():
"""Test _generate_azure_ad_redis_token with service principal credentials."""
from unittest.mock import Mock, patch
expected_token = "sp-access-token-67890"
mock_token = Mock()
mock_token.token = expected_token
mock_credential = Mock()
mock_credential.get_token.return_value = mock_token
mock_client_secret_credential = Mock(return_value=mock_credential)
mock_azure_identity = Mock()
mock_azure_identity.DefaultAzureCredential = Mock()
mock_azure_identity.ClientSecretCredential = mock_client_secret_credential
mock_azure_identity.ManagedIdentityCredential = Mock()
with patch.dict(
"sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()}
):
from litellm._redis import _generate_azure_ad_redis_token
result = _generate_azure_ad_redis_token(
azure_client_id="test-client-id",
azure_tenant_id="test-tenant-id",
azure_client_secret="test-secret",
)
assert result == expected_token
mock_client_secret_credential.assert_called_once_with(
client_id="test-client-id",
tenant_id="test-tenant-id",
client_secret="test-secret",
)
def test_generate_azure_ad_redis_token_import_error():
"""Test that _generate_azure_ad_redis_token raises ImportError when azure-identity is missing."""
from unittest.mock import patch
from litellm._redis import _generate_azure_ad_redis_token
with patch.dict("sys.modules", {"azure.identity": None}):
with pytest.raises(ImportError) as exc_info:
_generate_azure_ad_redis_token()
assert "azure-identity is required" in str(exc_info.value)
def test_redis_client_logic_azure_ad_auth():
"""Test that _get_redis_client_logic sets up Azure AD auth when REDIS_AZURE_AD_TOKEN=true.
Mocks ``azure.identity`` via ``sys.modules`` so the test does not require
the real ``azure-identity`` package to be installed in the CI environment.
"""
from unittest.mock import Mock, patch
mock_credential = Mock()
mock_azure_identity = Mock()
mock_azure_identity.DefaultAzureCredential = Mock(return_value=mock_credential)
mock_azure_identity.ClientSecretCredential = Mock(return_value=mock_credential)
mock_azure_identity.ManagedIdentityCredential = Mock(return_value=mock_credential)
with patch.dict(
"sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()}
):
from litellm._redis import _get_redis_client_logic
redis_kwargs = _get_redis_client_logic(
host="myredis.redis.cache.windows.net",
port="6380",
azure_redis_ad_token="true",
ssl=True,
)
assert "redis_connect_func" in redis_kwargs
# Marker for async paths to detect Azure AD auth
assert hasattr(redis_kwargs["redis_connect_func"], "_azure_redis_ad_token")
assert redis_kwargs["redis_connect_func"]._azure_redis_ad_token is True
# Live credential object (not raw secret) is exposed for async paths
assert hasattr(redis_kwargs["redis_connect_func"], "_azure_credential")
# Raw credentials must NOT be exposed on the function
assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_client_secret")
assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_client_id")
assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_tenant_id")
# Azure-specific kwargs should be removed from the dict passed to Redis
assert "azure_redis_ad_token" not in redis_kwargs
assert "azure_client_id" not in redis_kwargs
if __name__ == "__main__":
# Allow running this test file directly for debugging
pytest.main([__file__, "-v"])

View file

@ -0,0 +1,11 @@
import ast
from pathlib import Path
def test_utils_module_has_docstring():
utils_path = Path(__file__).parents[2] / "litellm" / "utils.py"
module = ast.parse(utils_path.read_text())
assert ast.get_docstring(module) == (
"Utility helpers for LiteLLM core request handling and provider support."
)

View file

@ -535,6 +535,21 @@ describe("ModelSelect", () => {
});
});
it("should not render an empty optgroup when includeSpecialOptions is omitted", async () => {
renderWithProviders(<ModelSelect onChange={mockOnChange} context="global" />);
await waitFor(() => {
expect(screen.getByTestId("model-select")).toBeInTheDocument();
});
const optgroups = document.querySelectorAll("optgroup");
// Wildcard Options + Models — no blank leading group
expect(optgroups.length).toBe(2);
optgroups.forEach((g) => {
expect(g.getAttribute("label")).toBeTruthy();
});
});
it("should render maxTagPlaceholder when many items are selected", async () => {
// Create many models to trigger maxTagCount responsive behavior
const manyModels: ProxyModel[] = Array.from({ length: 20 }, (_, i) => ({

View file

@ -141,36 +141,38 @@ export const ModelSelect = (props: ModelSelectProps) => {
onChange={handleChange}
style={style}
options={[
includeSpecialOptions
? {
label: <span>Special Options</span>,
title: "Special Options",
options: [
...(shouldShowAllProxyModels
? [
{
label: <span>All Proxy Models</span>,
value: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
disabled:
value.length > 0 &&
value.some(
(v) => isSpecialOption(v) && v !== MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
),
key: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
},
]
: []),
{
label: <span>No Default Models</span>,
value: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value,
disabled:
value.length > 0 &&
value.some((v) => isSpecialOption(v) && v !== MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value),
key: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value,
},
],
}
: [],
...(includeSpecialOptions
? [
{
label: <span>Special Options</span>,
title: "Special Options",
options: [
...(shouldShowAllProxyModels
? [
{
label: <span>All Proxy Models</span>,
value: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
disabled:
value.length > 0 &&
value.some(
(v) => isSpecialOption(v) && v !== MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
),
key: MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value,
},
]
: []),
{
label: <span>No Default Models</span>,
value: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value,
disabled:
value.length > 0 &&
value.some((v) => isSpecialOption(v) && v !== MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value),
key: MODEL_SELECT_NO_DEFAULT_MODELS_SPECIAL_VALUE.value,
},
],
},
]
: []),
...(wildcard.length > 0
? [
{

View file

@ -57,7 +57,10 @@ vi.mock("./EndpointUsage/EndpointUsage", () => ({
vi.mock("./UsageViewSelect/UsageViewSelect", async () => {
const React = await import("react");
const UsageViewSelect = ({ value, onChange }: any) => {
const UsageViewSelect = ({ value, onChange, canViewTagUsage = false }: any) => {
const tagOption = canViewTagUsage
? React.createElement("option", { value: "tag" }, "Tag Usage")
: null;
return React.createElement(
"select",
{
@ -70,7 +73,7 @@ vi.mock("./UsageViewSelect/UsageViewSelect", async () => {
React.createElement("option", { value: "team" }, "Team Usage"),
React.createElement("option", { value: "organization" }, "Organization Usage"),
React.createElement("option", { value: "customer" }, "Customer Usage"),
React.createElement("option", { value: "tag" }, "Tag Usage"),
tagOption,
React.createElement("option", { value: "agent" }, "Agent Usage"),
React.createElement("option", { value: "user-agent-activity" }, "User Agent Activity"),
);
@ -639,6 +642,29 @@ describe("UsagePage", () => {
});
});
it("should show tag usage selector option for internal users", async () => {
mockUseAuthorized.mockReturnValue({
isLoading: false,
isAuthorized: true,
token: "mock-token",
accessToken: "test-token",
userId: "user-123",
userEmail: "test@example.com",
userRole: "internal_user",
premiumUser: true,
disabledPersonalKeyCreation: false,
showSSOBanner: false,
});
renderWithProviders(<UsagePage {...defaultProps} />);
await waitFor(() => {
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled();
});
expect(screen.getByRole("option", { name: "Tag Usage" })).toBeInTheDocument();
});
it("should show organization usage banner and view for admins", async () => {
renderWithProviders(<UsagePage {...defaultProps} organizations={mockOrganizations} />);

View file

@ -31,7 +31,7 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { useCurrentUser } from "@/app/(dashboard)/hooks/users/useCurrentUser";
import { useInfiniteUsers } from "@/app/(dashboard)/hooks/users/useUsers";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import { all_admin_roles } from "../../../utils/roles";
import { all_admin_roles, internalUserRoles } from "../../../utils/roles";
import { ActivityMetrics, processActivityData } from "../../activity_metrics";
import CloudZeroExportModal from "../../cloudzero_export_modal";
import EntityUsageExportModal from "../../EntityUsageExport";
@ -84,6 +84,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
console.log(`currentUser: ${JSON.stringify(currentUser)}`);
console.log(`currentUser max budget: ${currentUser?.max_budget}`);
const isAdmin = all_admin_roles.includes(userRole || "");
const canViewTagUsage = isAdmin || internalUserRoles.includes(userRole || "");
// Debounced search for user selector
const [userSearchInput, setUserSearchInput] = useState("");
@ -145,23 +146,6 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
const [topKeysLimit, setTopKeysLimit] = useState<number>(5);
const [topModelsLimit, setTopModelsLimit] = useState<number>(5);
const [showTokenBreakdown, setShowTokenBreakdown] = useState(false);
const getAllTags = async () => {
if (!accessToken) {
return;
}
const tags = await tagListCall(accessToken);
setAllTags(
Object.values(tags).map((tag: Tag) => ({
label: tag.name,
value: tag.name,
})),
);
};
useEffect(() => {
getAllTags();
}, [accessToken]);
// Sync selectedUserId when auth state settles (isAdmin/userID may be null on initial render)
useEffect(() => {
if (!isAdmin && userID) {
@ -175,6 +159,30 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
const startTime = useMemo(() => (dateValue.from ? new Date(dateValue.from) : null), [dateValue.from]);
const endTime = useMemo(() => (dateValue.to ? new Date(dateValue.to) : null), [dateValue.to]);
useEffect(() => {
if (!accessToken) return;
let cancelled = false;
(async () => {
try {
const tags = await tagListCall(accessToken, startTime, endTime);
if (cancelled) return;
setAllTags(
Object.values(tags).map((tag: Tag) => ({
label: tag.name,
value: tag.name,
})),
);
} catch (e) {
if (!cancelled) {
console.error("Failed to fetch tag list", e);
}
}
})();
return () => {
cancelled = true;
};
}, [accessToken, startTime, endTime]);
// Try aggregated endpoint first, fall back to paginated on failure
const aggregatedFetchIdRef = useRef(0);
useEffect(() => {
@ -437,7 +445,12 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
<div className="flex items-end justify-between gap-6 mb-6">
<div className="flex-1">
<div className="flex items-end justify-between gap-6 mb-4 w-full">
<UsageViewSelect value={usageView} onChange={(value) => setUsageView(value)} isAdmin={isAdmin} />
<UsageViewSelect
value={usageView}
onChange={(value) => setUsageView(value)}
isAdmin={isAdmin}
canViewTagUsage={canViewTagUsage}
/>
<AdvancedDatePicker value={dateValue} onValueChange={handleDateChange} />
</div>
{paginatedResult.isFetchingMore && (

View file

@ -110,4 +110,16 @@ describe("UsageViewSelect", () => {
expect(mockOnChange).toHaveBeenCalledWith("team");
});
it("should show Tag Usage for non-admin users with tag usage permission", () => {
render(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} canViewTagUsage={true} />);
expect(screen.getByRole("option", { name: "Tag Usage" })).toBeInTheDocument();
});
it("should hide Tag Usage for non-admin users without tag usage permission", () => {
render(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} />);
expect(screen.queryByRole("option", { name: "Tag Usage" })).not.toBeInTheDocument();
});
});

View file

@ -16,6 +16,7 @@ export interface UsageViewSelectProps {
value: UsageOption;
onChange: (value: UsageOption) => void;
isAdmin: boolean;
canViewTagUsage?: boolean;
title?: string;
description?: string;
"data-id"?: string;
@ -106,12 +107,16 @@ export const UsageViewSelect: React.FC<UsageViewSelectProps> = ({
value,
onChange,
isAdmin,
canViewTagUsage = false,
title = "Usage View",
description = "Select the usage data you want to view",
"data-id": dataId,
}) => {
const getFilteredOptions = () => {
return OPTIONS.filter((option) => {
if (option.value === "tag" && canViewTagUsage) {
return true;
}
if (option.adminOnly && !isAdmin) {
return false;
}

View file

@ -250,6 +250,33 @@ describe("ModelInfoView", () => {
});
});
it("should pass model_info.id to disambiguate duplicate model_name deployments", async () => {
// Regression test: when two deployments share `model_name` (e.g.
// wildcard `openai/*` with different `api_base` values), the UI
// must forward the clicked row's `model_info.id` to the backend.
// Otherwise /health/test_connection silently probes deployments[0]
// instead of the deployment the user actually selected.
const user = userEvent.setup();
render(<ModelInfoView {...DEFAULT_ADMIN_PROPS} />, { wrapper });
await waitFor(() => {
expect(screen.getByText("Model Settings")).toBeInTheDocument();
});
const testButton = screen.getByRole("button", { name: /test connection/i });
await user.click(testButton);
await waitFor(() => {
expect(mockTestConnectionRequest).toHaveBeenCalled();
});
const callArgs = mockTestConnectionRequest.mock.calls[0];
// Signature: (accessToken, litellm_params, model_info, mode)
const modelInfoArg = callArgs[2] as Record<string, unknown>;
expect(modelInfoArg).toBeDefined();
expect(modelInfoArg.id).toBe("123");
});
it("should display error notification when connection test fails", async () => {
const user = userEvent.setup();
mockTestConnectionRequest.mockRejectedValue(new Error("Connection failed"));

View file

@ -379,6 +379,12 @@ export default function ModelInfoView({
model: localModelData.litellm_model_name,
},
{
// `id` is required to disambiguate when multiple deployments
// share the same model_name (e.g. wildcard `openai/*` with two
// different `api_base` values for failover). Without it the
// backend silently falls back to deployments[0] and probes
// the wrong endpoint.
id: localModelData.model_info?.id,
mode: localModelData.model_info?.mode,
},
localModelData.model_info?.mode,

View file

@ -7288,10 +7288,29 @@ export const tagInfoCall = async (accessToken: string, tagNames: string[]): Prom
}
};
export const tagListCall = async (accessToken: string): Promise<TagListResponse> => {
const formatYmd = (value: Date): string => {
const year = value.getFullYear();
const month = String(value.getMonth() + 1).padStart(2, "0");
const day = String(value.getDate()).padStart(2, "0");
return `${year}-${month}-${day}`;
};
export const tagListCall = async (
accessToken: string,
startTime?: Date | null,
endTime?: Date | null,
): Promise<TagListResponse> => {
try {
let url = proxyBaseUrl ? `${proxyBaseUrl}/tag/list` : `/tag/list`;
if (startTime && endTime) {
const params = new URLSearchParams({
start_date: formatYmd(startTime),
end_date: formatYmd(endTime),
});
url = `${url}?${params.toString()}`;
}
const response = await fetch(url, {
method: "GET",
headers: {

View file

@ -158,8 +158,8 @@ describe("KeyEditView", () => {
const { getByText } = renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -176,8 +176,8 @@ describe("KeyEditView", () => {
const { getByText } = renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -194,8 +194,8 @@ describe("KeyEditView", () => {
const { getByLabelText } = renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -219,7 +219,7 @@ describe("KeyEditView", () => {
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={onCancelMock}
onSubmit={async () => { }}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -241,8 +241,8 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -259,8 +259,8 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -277,8 +277,8 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -295,8 +295,8 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -314,7 +314,7 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
@ -344,8 +344,8 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={keyDataWithManagementRoutes}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -367,8 +367,8 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={keyDataWithInfoRoutes}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={""}
userID={""}
userRole={""}
@ -385,8 +385,8 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onSubmit={async () => { }}
onCancel={() => {}}
onSubmit={async () => {}}
accessToken={"test-token"}
userID={""}
userRole={""}
@ -404,7 +404,7 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
@ -434,10 +434,14 @@ describe("KeyEditView", () => {
it("should handle empty allowed routes string on submit", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
const keyDataWithRoutes = {
...MOCK_KEY_DATA,
allowed_routes: ["llm_api_routes"],
};
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
keyData={keyDataWithRoutes}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
@ -463,6 +467,101 @@ describe("KeyEditView", () => {
});
});
it("should omit allowed_routes from submit when value is unchanged", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
const aiApisKeyData = {
...MOCK_KEY_DATA,
allowed_routes: ["llm_api_routes"],
};
renderWithProviders(
<KeyEditView
keyData={aiApisKeyData}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
userRole={"admin"}
premiumUser={false}
/>,
);
await waitFor(() => {
expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument();
});
const submitButton = screen.getByRole("button", { name: /save changes/i });
await userEvent.click(submitButton);
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalled();
const callArgs = onSubmitMock.mock.calls[0][0];
expect("allowed_routes" in callArgs).toBe(false);
});
});
it("should omit allowed_routes from submit when keyData.allowed_routes is null and form is untouched", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
const keyDataNullRoutes = {
...MOCK_KEY_DATA,
allowed_routes: null as unknown as string[],
};
renderWithProviders(
<KeyEditView
keyData={keyDataNullRoutes}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
userRole={"admin"}
premiumUser={false}
/>,
);
await waitFor(() => {
expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument();
});
const submitButton = screen.getByRole("button", { name: /save changes/i });
await userEvent.click(submitButton);
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalled();
const callArgs = onSubmitMock.mock.calls[0][0];
expect("allowed_routes" in callArgs).toBe(false);
});
});
it("should omit allowed_routes from submit when server returned routes in a different order", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
const keyDataReordered = {
...MOCK_KEY_DATA,
allowed_routes: ["beta_routes", "alpha_routes"],
};
renderWithProviders(
<KeyEditView
keyData={keyDataReordered}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
userRole={"admin"}
premiumUser={false}
/>,
);
await waitFor(() => {
expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument();
});
const submitButton = screen.getByRole("button", { name: /save changes/i });
await userEvent.click(submitButton);
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalled();
const callArgs = onSubmitMock.mock.calls[0][0];
expect("allowed_routes" in callArgs).toBe(false);
});
});
it("should pass access_group_ids to onSubmit when saving key with access groups", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
@ -554,7 +653,7 @@ describe("KeyEditView", () => {
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => { }}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
@ -576,10 +675,13 @@ describe("KeyEditView", () => {
});
// Wait for the cancel button to actually be disabled (state update may take a moment)
await waitFor(() => {
const cancelButton = screen.getByRole("button", { name: /cancel/i });
expect(cancelButton).toBeDisabled();
}, { timeout: 3000 });
await waitFor(
() => {
const cancelButton = screen.getByRole("button", { name: /cancel/i });
expect(cancelButton).toBeDisabled();
},
{ timeout: 3000 },
);
// Clean up: resolve the promise to allow the form to complete
if (resolveSubmit) {

View file

@ -78,7 +78,6 @@ const getKeyTypeFromRoutes = (allowedRoutes: string[] | null | undefined): strin
return "default";
};
export function KeyEditView({
keyData,
onCancel,
@ -106,7 +105,7 @@ export function KeyEditView({
const [neverExpire, setNeverExpire] = useState<boolean>(!keyData.expires);
const [isKeySaving, setIsKeySaving] = useState(false);
const [budgetLimits, setBudgetLimits] = useState<BudgetWindowEntry[]>(
Array.isArray(keyData.budget_limits) ? keyData.budget_limits : []
Array.isArray(keyData.budget_limits) ? keyData.budget_limits : [],
);
const { data: organizations, isLoading: isOrganizationsLoading } = useOrganizations();
const { data: projects } = useProjects();
@ -116,9 +115,7 @@ export function KeyEditView({
const projectDisplay = (() => {
if (!keyData.project_id) return null;
const project = projects?.find((p) => p.project_id === keyData.project_id);
return project?.project_alias
? `${project.project_alias} (${keyData.project_id})`
: keyData.project_id;
return project?.project_alias ? `${project.project_alias} (${keyData.project_id})` : keyData.project_id;
})();
useEffect(() => {
@ -198,9 +195,10 @@ export function KeyEditView({
access_group_ids: keyData.access_group_ids || [],
auto_rotate: keyData.auto_rotate || false,
...(keyData.rotation_interval && { rotation_interval: keyData.rotation_interval }),
allowed_routes: Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0
? keyData.allowed_routes.join(", ")
: "",
allowed_routes:
Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0
? keyData.allowed_routes.join(", ")
: "",
};
useEffect(() => {
@ -226,9 +224,10 @@ export function KeyEditView({
access_group_ids: keyData.access_group_ids || [],
auto_rotate: keyData.auto_rotate || false,
...(keyData.rotation_interval && { rotation_interval: keyData.rotation_interval }),
allowed_routes: Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0
? keyData.allowed_routes.join(", ")
: "",
allowed_routes:
Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0
? keyData.allowed_routes.join(", ")
: "",
});
}, [keyData, form]);
@ -275,12 +274,25 @@ export function KeyEditView({
}
// If it's already an array (shouldn't happen, but handle it), keep as is
// Backend rejects non-empty allowed_routes from non-admins, so re-sending
// an unchanged value 403s a team admin. Set compare tolerates reorder.
const originalRoutesSet = new Set<string>(Array.isArray(keyData.allowed_routes) ? keyData.allowed_routes : []);
const submittedRoutesSet = new Set<string>(Array.isArray(values.allowed_routes) ? values.allowed_routes : []);
const allowedRoutesUnchanged =
originalRoutesSet.size === submittedRoutesSet.size &&
[...submittedRoutesSet].every((r) => originalRoutesSet.has(r));
if (allowedRoutesUnchanged) {
delete values.allowed_routes;
}
if (neverExpire) {
values.duration = null;
}
// Include multi-window budget limits (filter out incomplete entries)
const validWindows = budgetLimits.filter((w) => w.budget_duration && w.max_budget !== null && w.max_budget !== undefined);
const validWindows = budgetLimits.filter(
(w) => w.budget_duration && w.max_budget !== null && w.max_budget !== undefined,
);
values.budget_limits = validWindows.length > 0 ? validWindows : undefined;
await onSubmit(values);
@ -305,9 +317,13 @@ export function KeyEditView({
{({ getFieldValue, setFieldValue }) => {
const allowedRoutesValue = getFieldValue("allowed_routes") || "";
// Convert string to array for checking
const allowedRoutes = typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== ""
? allowedRoutesValue.split(",").map((r: string) => r.trim()).filter((r: string) => r.length > 0)
: [];
const allowedRoutes =
typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== ""
? allowedRoutesValue
.split(",")
.map((r: string) => r.trim())
.filter((r: string) => r.length > 0)
: [];
const isDisabled = allowedRoutes.includes("management_routes") || allowedRoutes.includes("info_routes");
const models = getFieldValue("models") || [];
@ -348,9 +364,13 @@ export function KeyEditView({
{({ getFieldValue, setFieldValue }) => {
const allowedRoutesValue = getFieldValue("allowed_routes") || "";
// Convert string to array for getKeyTypeFromRoutes
const allowedRoutes = typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== ""
? allowedRoutesValue.split(",").map((r: string) => r.trim()).filter((r: string) => r.length > 0)
: [];
const allowedRoutes =
typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== ""
? allowedRoutesValue
.split(",")
.map((r: string) => r.trim())
.filter((r: string) => r.length > 0)
: [];
const keyTypeValue = getKeyTypeFromRoutes(allowedRoutes);
return (
@ -415,9 +435,7 @@ export function KeyEditView({
}
name="allowed_routes"
>
<Input
placeholder="Enter allowed routes (comma-separated). Special values: llm_api_routes, management_routes. Examples: llm_api_routes, /chat/completions, /keys/*. Leave empty to allow all routes"
/>
<Input placeholder="Enter allowed routes (comma-separated). Special values: llm_api_routes, management_routes. Examples: llm_api_routes, /chat/completions, /keys/*. Leave empty to allow all routes" />
</Form.Item>
<Form.Item label="Max Budget (USD)" name="max_budget">
@ -442,10 +460,7 @@ export function KeyEditView({
</span>
}
>
<BudgetWindowsEditor
value={budgetLimits}
onChange={setBudgetLimits}
/>
<BudgetWindowsEditor value={budgetLimits} onChange={setBudgetLimits} />
</Form.Item>
<Form.Item label="TPM Limit" name="tpm_limit">
@ -579,7 +594,7 @@ export function KeyEditView({
!premiumUser
? "Premium feature - Upgrade to set allowed pass through routes by key"
: Array.isArray(keyData.metadata?.allowed_passthrough_routes) &&
keyData.metadata.allowed_passthrough_routes.length > 0
keyData.metadata.allowed_passthrough_routes.length > 0
? `Current: ${keyData.metadata.allowed_passthrough_routes.join(", ")}`
: "Select or enter allowed pass through routes"
}
@ -690,14 +705,13 @@ export function KeyEditView({
return team.team_alias?.toLowerCase().includes(input.toLowerCase()) ?? false;
}}
>
{(selectedOrganizationId
? teams?.filter((t) => t.organization_id === selectedOrganizationId)
: teams
)?.map((team) => (
<Select.Option key={team.team_id} value={team.team_id}>
{`${team.team_alias} (${team.team_id})`}
</Select.Option>
))}
{(selectedOrganizationId ? teams?.filter((t) => t.organization_id === selectedOrganizationId) : teams)?.map(
(team) => (
<Select.Option key={team.team_id} value={team.team_id}>
{`${team.team_alias} (${team.team_id})`}
</Select.Option>
),
)}
</Select>
</Form.Item>
{enableProjectsUI && hasProject && (

View file

@ -396,7 +396,7 @@ export default function KeyInfoView({
};
return (
<div className="w-full h-screen p-4">
<div className="w-full h-full overflow-y-auto p-4">
<KeyInfoHeader
data={{
keyName: currentKeyData.key_alias || "Virtual Key",
@ -614,7 +614,7 @@ export default function KeyInfoView({
{/* Settings Panel */}
<TabPanel>
<Card className="overflow-y-auto max-h-[65vh]">
<Card>
<div className="flex justify-between items-center mb-4">
<Title>Key Settings</Title>
{!isEditing && canModifyKey && (

View file

@ -19,6 +19,7 @@ export const MCP_CALL_TYPES = ["call_mcp_tool", "list_mcp_tools"];
export const AGENT_CALL_TYPES = ["asend_message"];
export const QUICK_SELECT_OPTIONS: { label: string; value: number; unit: string }[] = [
{ label: "Last Minute", value: 1, unit: "minutes" },
{ label: "Last 15 Minutes", value: 15, unit: "minutes" },
{ label: "Last Hour", value: 1, unit: "hours" },
{ label: "Last 4 Hours", value: 4, unit: "hours" },

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