Merge branch 'litellm_internal_staging' into shin_agent_oss_staging_05_10_2026

Co-authored-by: Cursor <cursoragent@cursor.com>

# Conflicts:
#	litellm/proxy/common_utils/reset_budget_job.py
#	tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
This commit is contained in:
Sameer Kankute 2026-05-11 12:15:35 +05:30
commit 518f9dc7a2
No known key found for this signature in database
65 changed files with 5143 additions and 489 deletions

Binary file not shown.

After

Width:  |  Height:  |  Size: 60 KiB

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

@ -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

@ -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

@ -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 |

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,
@ -687,7 +689,7 @@ 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
self_managed_routes = [
"/team/member_add",
@ -4346,10 +4348,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

@ -83,27 +83,27 @@ class ResetBudgetJob:
"Failed to reset spend counter %s: %s", counter_key, e
)
async def _invalidate_source_cache(self, source_cache_key: str) -> None:
"""Drop a cached entity so the next budget check re-reads spend from DB.
@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.
Without this, a row whose spend was just zeroed in DB can still be
served from user_api_key_cache (e.g. ``tag:<name>``) with the
pre-reset spend value, which is consulted as the cold-start /
DB-unavailable fallback in get_current_spend(). Run this AFTER the
DB write commits, mirroring _invalidate_spend_counter.
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.
"""
if self.proxy_logging_obj is None:
return
user_api_key_cache = self.proxy_logging_obj.call_details.get(
"user_api_key_cache"
)
if user_api_key_cache is None:
return
try:
await user_api_key_cache.async_delete_cache(key=source_cache_key)
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 source cache key %s: %s", source_cache_key, e
"Failed to invalidate user_api_key_cache entry %s: %s",
cache_key,
e,
)
async def _cascade_reset_spend_for_budget_link(
@ -113,14 +113,17 @@ class ResetBudgetJob:
counter_key_fn: Callable[[Any], str],
log_subject: str,
extra_where: Optional[dict] = None,
source_cache_key_fn: Optional[Callable[[Any], str]] = None,
cache_key_fn: Optional[Callable[[Any], str]] = None,
):
"""
Generic cascade: zero spend on rows whose budget_id is in the reset set.
When ``source_cache_key_fn`` is supplied, the corresponding entry in
user_api_key_cache is also evicted so the next request re-reads the
zeroed spend from DB rather than the stale cached object.
``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:
@ -142,8 +145,8 @@ class ResetBudgetJob:
for row in rows:
await self._invalidate_spend_counter(counter_key_fn(row))
if source_cache_key_fn is not None:
await self._invalidate_source_cache(source_cache_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
@ -197,13 +200,13 @@ class ResetBudgetJob:
"""
Resets the spend for tags linked to budget tiers that are being reset.
Also evicts the cached LiteLLM_TagTable object at ``tag:<name>`` in
user_api_key_cache. The auth-time tag budget check (see
``_tag_max_budget_check`` in litellm/proxy/auth/auth_checks.py) reads
``tag_object.spend`` as the DB-unavailable fallback in
``get_current_spend``; if that cached object survives a reset it can
keep blocking otherwise-unblocked tenants under cold-start /
Redis-down scenarios.
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,
@ -211,7 +214,7 @@ class ResetBudgetJob:
counter_key_fn=lambda t: f"spend:tag:{t.tag_name}",
log_subject="tags",
extra_where={"spend": {"gt": 0}},
source_cache_key_fn=lambda t: f"tag:{t.tag_name}",
cache_key_fn=lambda t: f"tag:{t.tag_name}",
)
async def reset_budget_for_litellm_budget_table(self):

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,9 +12,10 @@ 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
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
@ -395,6 +396,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 +429,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,6 +450,8 @@ 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:
## QUERY STORED TAGS ##
tag_records = await prisma_client.db.litellm_tagtable.find_many(
@ -453,9 +494,13 @@ async def list_tags(
# 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 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={"tag": {"not": None}},
where=dynamic_tag_where,
min={"created_at": True},
max={"updated_at": True},
)

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

@ -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

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

@ -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

@ -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

@ -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

@ -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"""

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

@ -1262,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
@ -1460,12 +1470,19 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monke
)
def test_reset_budget_for_tags_linked_to_budgets_invalidates_source_cache(monkeypatch):
"""Resetting tags must also evict the cached LiteLLM_TagTable object so
the auth-time fallback (``tag_object.spend``) does not keep blocking a
tenant after spend has been zeroed in the DB.
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.
"""
_make_counter_invalidation_job(monkeypatch)
counter_cache = _make_counter_invalidation_job(monkeypatch)
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_tag = type("Tag", (), {"tag_name": "tenant-42"})
@ -1474,26 +1491,78 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_source_cache(monkey
prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag])
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1})
user_api_key_cache = MagicMock()
user_api_key_cache.async_delete_cache = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.call_details = {"user_api_key_cache": user_api_key_cache}
job = ResetBudgetJob(
proxy_logging_obj=proxy_logging_obj, prisma_client=prisma_client
)
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget]))
user_api_key_cache.async_delete_cache.assert_any_await(key="tag:tenant-42")
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_no_user_api_key_cache(monkeypatch):
"""When user_api_key_cache is not wired up (e.g. early-boot or tests),
the cascade must still complete without raising the spend counter
invalidation is the load-bearing path.
"""
_make_counter_invalidation_job(monkeypatch)
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()
def test_reset_budget_for_tags_linked_to_budgets_management_cache_delete_failure_still_resets(
monkeypatch,
):
"""If ``async_delete_cache`` raises, the DB cascade must still complete."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
counter_cache.user_api_key_cache.async_delete_cache = AsyncMock(
side_effect=RuntimeError("cache unavailable")
)
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_tag = type("Tag", (), {"tag_name": "tenant-42"})
@ -1502,14 +1571,7 @@ def test_reset_budget_for_tags_linked_to_budgets_no_user_api_key_cache(monkeypat
prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[linked_tag])
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 1})
proxy_logging_obj = MagicMock()
proxy_logging_obj.call_details = {}
job = ResetBudgetJob(
proxy_logging_obj=proxy_logging_obj, prisma_client=prisma_client
)
# Should not raise even though user_api_key_cache is missing.
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget]))
# The DB write must still have happened.
prisma_client.db.litellm_tagtable.update_many.assert_awaited_once()

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,117 @@ async def test_list_tags_no_dynamic_tags():
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_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

@ -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

@ -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

@ -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

@ -145,23 +145,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 +158,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(() => {

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 && (