Merge remote-tracking branch 'origin/litellm_cost_shard_provider_wires' into litellm_cost_shard_proxy_behaviour

This commit is contained in:
kerry 2026-09-21 21:15:49 +00:00
commit 2760ea2e6c
22 changed files with 1251 additions and 37 deletions

View file

@ -11,6 +11,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03",
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": "fast-mode-2026-02-01",
"files-api-2025-04-14": "files-api-2025-04-14",
@ -44,6 +45,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": null,
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": "files-api-2025-04-14",
@ -76,6 +78,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": null,
"dangerous-tool-use-2026-09-03": null,
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
@ -109,6 +112,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03",
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
@ -143,6 +147,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03",
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
@ -177,6 +182,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03",
"effort-2025-11-24": null,
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
@ -210,6 +216,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": null,
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": "fast-mode-2026-02-01",
"files-api-2025-04-14": "files-api-2025-04-14",

View file

@ -80,6 +80,8 @@ class _AsyncRedisCommands(Protocol):
def ttl(self, name: str) -> Awaitable[int]: ...
def expire(self, name: str, time: int) -> Awaitable[bool]: ...
def rpush(self, name: str, *values: str | bytes | float) -> Awaitable[int]: ...
def lpop(self, name: str, count: int | None = None) -> Awaitable[object]: ...
@ -1948,6 +1950,14 @@ class RedisCache(BaseCache):
_record_swallowed_redis_failure(self._circuit_breaker, e)
return None
@_redis_circuit_breaker_guard
async def async_refresh_ttl(self, key: str, ttl: int | None = None) -> bool:
"""EXPIRE an existing key without touching its value. False when the key is absent."""
_used_ttl: Final = self.get_ttl(ttl=ttl)
if _used_ttl is None:
return False
return await self._async_commands().expire(self.check_and_fix_namespace(key=key), _used_ttl)
@_redis_circuit_breaker_guard
async def async_rpush(
self,

View file

@ -533,6 +533,9 @@ class AmazonAnthropicClaudeMessagesConfig(
if anthropic_model_info.is_eager_input_streaming_used(tools):
beta_set.add(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER)
if anthropic_messages_optional_request_params.get("safeguards") is not None:
beta_set.add(ANTHROPIC_BETA_HEADER_VALUES.DANGEROUS_TOOL_USE_2026_09_03.value)
self._filter_context_management_for_bedrock_invoke(
anthropic_messages_request=anthropic_messages_request,
beta_set=beta_set,

View file

@ -108,6 +108,9 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert
if anthropic_model_info.is_tool_search_used(tools):
beta_values.add(get_tool_search_beta_header("vertex_ai"))
if optional_params.get("safeguards") is not None:
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.DANGEROUS_TOOL_USE_2026_09_03.value)
if beta_values:
headers["anthropic-beta"] = ",".join(beta_values)

View file

@ -43011,21 +43011,21 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
"input_cost_per_token": 9.1089e-07,
"input_cost_per_token": 9.0741e-07,
"input_cost_per_token_cache_hit": 4.4e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 1.82178e-06,
"output_cost_per_token": 1.81482e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"cache_read_input_token_cost": 7.59075e-08,
"cache_read_input_token_cost": 7.56175e-08,
"supports_audio_input": false,
"supports_pdf_input": false,
"supports_vision": false,

View file

@ -28,6 +28,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.ptu_pricing import (
CUSTOM_PRICING_FIELDS,
PTU_EMPTIED_PRICING_FIELDS,
@ -94,6 +95,7 @@ from litellm.proxy.spend_tracking.ptu_feature_flag import (
is_ptu_cost_attribution_enabled,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.credentials_repository import CredentialsRepository
from litellm.repositories.model_repository import ModelRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import ModelTableRepository
@ -145,7 +147,7 @@ if TYPE_CHECKING:
from prisma import types as prisma_types
router: Final = APIRouter()
CLEARABLE_LITELLM_PARAMS: Final = frozenset({"cache_control_injection_points"})
CLEARABLE_LITELLM_PARAMS: Final = frozenset({"cache_control_injection_points", "litellm_credential_name"})
NULL_CLEARABLE_LITELLM_PARAMS: Final = frozenset((*SPECIAL_MODEL_INFO_PARAMS, *CLEARABLE_LITELLM_PARAMS))
@ -332,6 +334,36 @@ def _raise_on_strategy_router_write_violation(
)
async def _raise_on_invalid_credential_name(
litellm_params: updateLiteLLMParams | None, prisma_client: PrismaClient
) -> None:
if litellm_params is None or "litellm_credential_name" not in litellm_params.model_fields_set:
return
credential_name: Final = litellm_params.litellm_credential_name
if credential_name is None:
return
if credential_name == "":
raise ProxyException(
message="litellm_credential_name cannot be an empty string. Send null to detach the stored credential or omit the field to leave it unchanged.",
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="litellm_credential_name",
)
if CredentialAccessor.find_credential(credential_name) is not None:
return
stored_credential: Final = await CredentialsRepository(WriterPinnedClient(prisma_client.db)).find_by_name(
credential_name
)
if stored_credential is not None:
return
raise ProxyException(
message=f"Credential '{credential_name}' not found. Create it via /credentials before attaching it to a model.",
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="litellm_credential_name",
)
AUTO_ROUTER_CAPABILITY_SLOT_LOCK_KEY: Final = 5_872_301
_CAPABILITY_LOCK_SQL: Final = "SELECT 1 AS locked FROM pg_advisory_xact_lock($1)"
_STORED_LITELLM_PARAMS_SQL: Final = (
@ -1110,7 +1142,9 @@ async def patch_model(
litellm_params=patch_data.litellm_params,
user_api_key_dict=user_api_key_dict,
existing_litellm_params=db_model.litellm_params,
null_detaches=True,
)
await _raise_on_invalid_credential_name(patch_data.litellm_params, prisma_client)
ModelManagementAuthChecks.can_user_set_aws_session_tags(
litellm_params=patch_data.litellm_params,
@ -1920,22 +1954,33 @@ class ModelManagementAuthChecks:
litellm_params: GenericLiteLLMParams | None,
user_api_key_dict: UserAPIKeyAuth,
existing_litellm_params: GenericLiteLLMParams | None = None,
*,
null_detaches: bool = False,
) -> Literal[True]:
if litellm_params is None or litellm_params.litellm_credential_name is None:
if litellm_params is None:
return True
if existing_litellm_params is not None and existing_litellm_params.litellm_credential_name is not None:
existing_credential_name: Final = decrypt_value_helper(
if "litellm_credential_name" not in litellm_params.model_fields_set:
return True
if litellm_params.litellm_credential_name is None and not null_detaches:
return True
existing_credential_name: Final = (
decrypt_value_helper(
value=existing_litellm_params.litellm_credential_name,
key="litellm_credential_name",
exception_type="debug",
return_original_value=True,
)
if litellm_params.litellm_credential_name == existing_credential_name:
return True
if existing_litellm_params is not None and existing_litellm_params.litellm_credential_name is not None
else None
)
requested_credential_name: Final = litellm_params.litellm_credential_name
if requested_credential_name == existing_credential_name:
return True
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
return True
action: Final = "detach" if requested_credential_name is None else "attach"
raise ProxyException(
message=f"Only a proxy admin can attach a stored credential (litellm_credential_name) to a model. Your role={user_api_key_dict.user_role}.",
message=f"Only a proxy admin can {action} a stored credential (litellm_credential_name) on a model. Your role={user_api_key_dict.user_role}.",
type=ProxyErrorTypes.auth_error.value,
code=status.HTTP_403_FORBIDDEN,
param="litellm_credential_name",

View file

@ -3549,6 +3549,16 @@ async def increment_spend_counter(counter_key: str, increment: float):
return await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
async def refresh_spend_counter_ttl(counter_key: str) -> bool:
if spend_counter_cache.redis_cache is None:
return False
try:
return await spend_counter_cache.redis_cache.async_refresh_ttl(key=counter_key)
except Exception as e:
verbose_proxy_logger.debug("spend counter TTL refresh skipped for %s: %s", counter_key, e)
return False
async def _increment_spend_counter_cache(counter_key: str, increment: float):
if spend_counter_cache.redis_cache is not None:
try:

View file

@ -3,6 +3,7 @@ from __future__ import annotations
import asyncio
import json
import math
import time
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@ -105,6 +106,48 @@ def get_reserved_counter_keys(budget_reservation: dict | None) -> set:
}
_lease_renewals: Final[set[asyncio.Task[None]]] = set() # mutable-ok: asyncio only weak-refs pending tasks
def _start_reservation_lease_renewal(budget_reservation: Mapping[str, object], counter_keys: frozenset[str]) -> None:
"""A reservation lives inside spend counter keys that expire on their Redis TTL. Renew the TTL
while the request is in flight so a request longer than the TTL does not drop its
reservation and admit concurrent requests against the DB floor on any worker."""
from litellm.proxy.proxy_server import spend_counter_cache
if spend_counter_cache.redis_cache is None or not counter_keys:
return
task: Final = asyncio.create_task(
_renew_reservation_lease(
budget_reservation=budget_reservation,
counter_keys=counter_keys,
interval=spend_counter_cache.redis_cache.default_ttl / 2,
request_task=asyncio.current_task(),
)
)
_lease_renewals.add(task)
task.add_done_callback(_lease_renewals.discard)
async def _renew_reservation_lease(
budget_reservation: Mapping[str, object],
counter_keys: frozenset[str],
interval: float,
request_task: asyncio.Task[object] | None,
) -> None:
"""Stops on finalization or once the request task that took the reservation is gone, so a
disconnect path that skipped reconciliation falls back to the plain counter TTL."""
from litellm.proxy.proxy_server import refresh_spend_counter_ttl
deadline: Final = time.monotonic() + litellm.request_timeout
while time.monotonic() < deadline:
await asyncio.sleep(interval)
if budget_reservation.get("finalized") is True or (request_task is not None and request_task.done()):
return
for counter_key in counter_keys:
await refresh_spend_counter_ttl(counter_key=counter_key)
def _key_reservation_should_release_for_throttle(counter_key: str, valid_token: UserAPIKeyAuth | None) -> bool:
"""
Whether an over-budget key's own ``max_budget`` reservation should be
@ -319,13 +362,18 @@ async def reserve_budget_for_request(
llm_router=llm_router,
input_token_counts=input_token_counts,
)
return {
budget_reservation: Final = {
"reserved_cost": reservation_cost,
"entries": applied_entries,
"finalized": False,
"input_cost": min(float(input_cost or 0.0), reservation_cost),
"input_tokens": max(input_token_counts.values(), default=None),
}
_start_reservation_lease_renewal(
budget_reservation=budget_reservation,
counter_keys=frozenset(get_reserved_counter_keys(budget_reservation=budget_reservation)),
)
return budget_reservation
async def reconcile_budget_reservation(

View file

@ -751,6 +751,7 @@ class ANTHROPIC_BETA_HEADER_VALUES(str, Enum):
FAST_MODE_2026_02_01 = "fast-mode-2026-02-01"
ADVISOR_TOOL_2026_03_01 = "advisor-tool-2026-03-01"
PER_TURN_CONTROL_2026_07_01 = "per-turn-control-2026-07-01"
DANGEROUS_TOOL_USE_2026_09_03 = "dangerous-tool-use-2026-09-03"
# Tool search beta header constant (for Anthropic direct API and Microsoft Foundry)

View file

@ -1238,6 +1238,7 @@ class BedrockInvokeAnthropicMessagesRequest(TypedDict, total=False):
thinking: dict
metadata: dict
output_config: dict
safeguards: list
# `context_management` is allowed for Bedrock InvokeModel only when it
# carries `compact_20260112` edits paired with the `compact-2026-01-12`

View file

@ -43011,21 +43011,21 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
"input_cost_per_token": 9.1089e-07,
"input_cost_per_token": 9.0741e-07,
"input_cost_per_token_cache_hit": 4.4e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 1.82178e-06,
"output_cost_per_token": 1.81482e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"cache_read_input_token_cost": 7.59075e-08,
"cache_read_input_token_cost": 7.56175e-08,
"supports_audio_input": false,
"supports_pdf_input": false,
"supports_vision": false,

View file

@ -2,7 +2,7 @@ import asyncio
import json
import os
import uuid
from typing import Any, Dict, List
from typing import Any, Dict, Final, List
import httpx
import pytest
@ -1584,3 +1584,104 @@ async def test_anthropic_messages_streaming_forwards_safeguards_and_keeps_safegu
assert captured["body"]["safeguards"] == safeguards
assert events[0]["message"]["safeguard_results"] == safeguard_results
assert [e for e in events if e["type"] == "message_delta"][0]["delta"]["safeguard_results"] == safeguard_results
def _claude_code_auto_mode_request() -> tuple[list[dict[str, object]], list[dict[str, object]]]:
"""Shapes are what Claude Code 2.1.278 sends and Bedrock Invoke / Vertex rawPredict return, captured 2026-09-21."""
safeguards = [{"type": "dangerous_tool_use", "classifier_context": {"v": 1, "permission_mode": "auto"}}]
tool_verdicts = {"toolu_01": {"type": "evaluated", "outcome": "not_flagged"}}
safeguard_results = [{"type": "dangerous_tool_use", "status": {"type": "available", "tool_uses": tool_verdicts}}]
return safeguards, safeguard_results
def _upstream_answering_with(safeguard_results: list[dict[str, object]], captured: dict[str, object]) -> AsyncHTTPHandler:
def upstream_records_the_request(request: httpx.Request) -> httpx.Response:
captured["body"] = json.loads(request.content)
captured["anthropic-beta"] = request.headers.get("anthropic-beta")
return httpx.Response(
200,
json={
"id": "msg_1",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
"content": [{"type": "text", "text": "ok"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 1, "output_tokens": 1},
"safeguard_results": safeguard_results,
},
request=request,
)
upstream = AsyncHTTPHandler()
upstream.client = httpx.AsyncClient(transport=httpx.MockTransport(upstream_records_the_request))
return upstream
_CLIENT_BETA_HEADERS: Final = (
pytest.param({"anthropic-beta": "dangerous-tool-use-2026-09-03,interleaved-thinking-2025-05-14"}, id="client_sends_beta"),
pytest.param({"anthropic-beta": "interleaved-thinking-2025-05-14"}, id="client_omits_beta"),
pytest.param({}, id="client_sends_no_beta_header"),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("client_headers", _CLIENT_BETA_HEADERS)
async def test_anthropic_messages_forwards_safeguards_and_dangerous_tool_use_beta_to_bedrock_invoke(
local_beta_headers_config, client_headers
):
"""Bedrock Invoke takes betas in the body's `anthropic_beta` and 400s on `safeguards` without the beta, so the beta rides along with the field."""
from litellm.llms.anthropic.experimental_pass_through.messages import handler
safeguards, safeguard_results = _claude_code_auto_mode_request()
captured: dict[str, object] = {}
response = await handler.anthropic_messages(
max_tokens=16,
messages=[{"role": "user", "content": "hi"}],
model="bedrock/us.anthropic.claude-sonnet-5",
custom_llm_provider="bedrock",
aws_access_key_id="test-access-key",
aws_secret_access_key="test-secret-key",
aws_region_name="us-east-1",
client=_upstream_answering_with(safeguard_results, captured),
safeguards=safeguards,
extra_headers=client_headers,
)
assert captured["body"]["safeguards"] == safeguards
assert captured["body"]["anthropic_beta"] == ["dangerous-tool-use-2026-09-03"]
assert response["safeguard_results"] == safeguard_results
@pytest.mark.asyncio
@pytest.mark.parametrize("client_headers", _CLIENT_BETA_HEADERS)
async def test_anthropic_messages_forwards_safeguards_and_dangerous_tool_use_beta_to_vertex(
local_beta_headers_config, client_headers
):
"""Vertex rawPredict takes the beta as the `anthropic-beta` header and 400s on `safeguards` without it, so the beta rides along with the field."""
from litellm.llms.anthropic.experimental_pass_through.messages import handler
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
safeguards, safeguard_results = _claude_code_auto_mode_request()
captured: dict[str, object] = {}
with patch.object(VertexBase, "_ensure_access_token", return_value=("test-token", "test-project")):
response = await handler.anthropic_messages(
max_tokens=16,
messages=[{"role": "user", "content": "hi"}],
model="vertex_ai/claude-sonnet-5",
custom_llm_provider="vertex_ai",
vertex_project="test-project",
vertex_location="global",
vertex_credentials="{}",
client=_upstream_answering_with(safeguard_results, captured),
safeguards=safeguards,
extra_headers=client_headers,
)
assert captured["body"]["safeguards"] == safeguards
assert "anthropic_beta" not in captured["body"]
assert captured["anthropic-beta"].split(",").count("dangerous-tool-use-2026-09-03") == 1
assert response["safeguard_results"] == safeguard_results

View file

@ -1651,6 +1651,92 @@ def test_bedrock_messages_allowlist_filters_anthropic_only_fields():
assert set(result).issubset(cfg.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS)
@pytest.mark.parametrize(
"client_beta_header",
["dangerous-tool-use-2026-09-03,interleaved-thinking-2025-05-14", "interleaved-thinking-2025-05-14"],
ids=["client_sends_beta", "client_omits_beta"],
)
def test_bedrock_messages_forwards_safeguards_with_dangerous_tool_use_beta(local_beta_headers_config, client_beta_header):
"""
Claude Code's server-side auto-mode classifier sends `safeguards` alongside the
dangerous-tool-use-2026-09-03 beta. Bedrock Invoke accepts the pair, answers
"safeguards: Extra inputs are not permitted" for the field alone, and returns
`safeguard_results: []` for the beta alone, so the field reaches it unchanged
and the beta rides along whether or not the client sent it, as every other
body-driven beta does here.
"""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
safeguards = [{"type": "dangerous_tool_use", "classifier_context": {"v": 1, "permission_mode": "auto"}}]
result = cfg.transform_anthropic_messages_request(
model="us.anthropic.claude-sonnet-5",
messages=[{"role": "user", "content": [{"type": "text", "text": "Hello"}]}],
anthropic_messages_optional_request_params={"max_tokens": 64, "safeguards": safeguards},
litellm_params=GenericLiteLLMParams(),
headers={"anthropic-beta": client_beta_header},
)
assert result["safeguards"] == safeguards
assert result["anthropic_beta"].count("dangerous-tool-use-2026-09-03") == 1
def test_bedrock_messages_does_not_add_dangerous_tool_use_beta_without_safeguards(local_beta_headers_config):
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
result = cfg.transform_anthropic_messages_request(
model="us.anthropic.claude-sonnet-5",
messages=[{"role": "user", "content": [{"type": "text", "text": "Hello"}]}],
anthropic_messages_optional_request_params={"max_tokens": 64},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "safeguards" not in result
assert "dangerous-tool-use-2026-09-03" not in result.get("anthropic_beta", [])
def test_bedrock_messages_stream_decoder_keeps_safeguard_results():
"""Bedrock streams the classifier verdicts on message_start and on the final message_delta, exactly as api.anthropic.com does."""
decoder = AmazonAnthropicClaudeMessagesStreamDecoder(model="us.anthropic.claude-sonnet-5")
tool_verdicts = {"toolu_01": {"type": "evaluated", "outcome": "not_flagged"}}
safeguard_results = [{"type": "dangerous_tool_use", "status": {"type": "available", "tool_uses": tool_verdicts}}]
message_start = decoder._chunk_parser(
{
"type": "message_start",
"message": {
"id": "msg_01",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 3, "output_tokens": 0},
"safeguard_results": safeguard_results,
},
}
)
assert isinstance(message_start, dict)
assert message_start["message"]["safeguard_results"] == safeguard_results
message_delta = decoder._chunk_parser(
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None, "safeguard_results": safeguard_results},
"usage": {"output_tokens": 1},
"amazon-bedrock-invocationMetrics": {"inputTokenCount": 3, "outputTokenCount": 1},
}
)
assert isinstance(message_delta, dict)
assert message_delta["delta"]["safeguard_results"] == safeguard_results
def test_bedrock_messages_filters_user_provided_unsupported_beta_header():
"""
In proxy deployments the client (e.g. Claude Code) doesn't know the backend

View file

@ -439,6 +439,19 @@ class TestBetaHeadersOnTheWire:
assert _sent_betas(route) == ["context-1m-2025-08-07", "context-management-2025-06-27"]
assert _sent_body(route)["context_management"] == {"edits": [{"type": "clear_tool_uses_20250919"}]}
@pytest.mark.asyncio
@respx.mock
async def test_safeguards_reach_mantle_with_the_dangerous_tool_use_beta(self):
"""Mantle answers 400 "safeguards: Extra inputs are not permitted" when the field
arrives without dangerous-tool-use-2026-09-03 (probed 2026-09-21), so the beta
has to ride along even when the client never sent the header."""
safeguards = [{"type": "dangerous_tool_use", "classifier_context": {"v": 1, "permission_mode": "auto"}}]
route = await self._send(safeguards=safeguards)
assert _sent_betas(route) == ["dangerous-tool-use-2026-09-03"]
assert _sent_body(route)["safeguards"] == safeguards
@pytest.mark.asyncio
@respx.mock
async def test_betas_and_version_never_travel_in_the_body(self):

View file

@ -3,6 +3,8 @@ import json
import os
from unittest.mock import MagicMock, patch
import pytest
from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.experimental_pass_through.transformation import (
VertexAIPartnerModelsAnthropicMessagesConfig,
)
@ -67,6 +69,63 @@ def test_web_search_header_added_for_messages_endpoint():
)
@pytest.mark.parametrize(
"client_headers",
[{"anthropic-beta": "dangerous-tool-use-2026-09-03"}, {}],
ids=["client_sends_beta", "client_omits_beta"],
)
def test_safeguards_add_dangerous_tool_use_beta_header(client_headers):
"""Vertex rejects `safeguards` without the dangerous-tool-use beta, so the beta rides along with the field the way the web search and context management betas do."""
config = VertexAIPartnerModelsAnthropicMessagesConfig()
litellm_params = {
"vertex_ai_project": "test-project",
"vertex_ai_location": "global",
"vertex_credentials": "{}",
}
optional_params = {
"safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1, "permission_mode": "auto"}}]
}
with (
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"),
):
updated_headers, _ = config.validate_anthropic_messages_environment(
headers=client_headers,
model="claude-sonnet-5",
messages=[],
optional_params=optional_params,
litellm_params=litellm_params,
api_base=None,
)
assert updated_headers["anthropic-beta"].split(",").count("dangerous-tool-use-2026-09-03") == 1
def test_no_safeguards_leaves_dangerous_tool_use_beta_header_out():
config = VertexAIPartnerModelsAnthropicMessagesConfig()
litellm_params = {
"vertex_ai_project": "test-project",
"vertex_ai_location": "global",
"vertex_credentials": "{}",
}
with (
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"),
):
updated_headers, _ = config.validate_anthropic_messages_environment(
headers={},
model="claude-sonnet-5",
messages=[],
optional_params={"max_tokens": 64},
litellm_params=litellm_params,
api_base=None,
)
assert "dangerous-tool-use-2026-09-03" not in updated_headers.get("anthropic-beta", "")
def test_web_search_header_not_added_without_tool():
"""Test that beta header is NOT added when web search tool is not present"""
config = VertexAIPartnerModelsAnthropicMessagesConfig()

View file

@ -523,6 +523,92 @@ def test_wildcard_credential_hydration_preserves_missing_credential_name(
}
def test_hydrate_credential_name_none_leaves_params_untouched(monkeypatch):
import litellm
from litellm.proxy.auth.model_checks import _hydrate_litellm_credential_name
from litellm.types.router import LiteLLM_Params
from litellm.types.utils import CredentialItem
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="shared-credential",
credential_info={},
credential_values={"api_key": "sk-shared"},
)
],
)
params = LiteLLM_Params(model="openai/gpt-4o", litellm_credential_name=None)
result = _hydrate_litellm_credential_name(params)
assert result is not None
assert result.api_key is None
assert result.litellm_credential_name is None
def test_hydrate_replaced_credential_uses_new_credential_values(monkeypatch):
import litellm
from litellm.proxy.auth.model_checks import _hydrate_litellm_credential_name
from litellm.types.router import LiteLLM_Params
from litellm.types.utils import CredentialItem
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="shared-credential",
credential_info={},
credential_values={"api_key": "sk-shared"},
),
CredentialItem(
credential_name="other-credential",
credential_info={},
credential_values={"api_key": "sk-other"},
),
],
)
params = LiteLLM_Params(model="openai/gpt-4o", litellm_credential_name="other-credential")
result = _hydrate_litellm_credential_name(params)
assert result is not None
assert result.api_key == "sk-other"
assert result.litellm_credential_name is None
def test_hydrate_inline_api_key_wins_over_stored_credential(monkeypatch):
import litellm
from litellm.proxy.auth.model_checks import _hydrate_litellm_credential_name
from litellm.types.router import LiteLLM_Params
from litellm.types.utils import CredentialItem
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="shared-credential",
credential_info={},
credential_values={"api_key": "sk-shared"},
)
],
)
params = LiteLLM_Params(
model="openai/gpt-4o",
api_key="sk-inline",
litellm_credential_name="shared-credential",
)
result = _hydrate_litellm_credential_name(params)
assert result is not None
assert result.api_key == "sk-inline"
@pytest.mark.asyncio
async def test_get_available_models_for_user_expands_query_team_wildcard(
monkeypatch,

View file

@ -10,6 +10,7 @@ import pytest
from fastapi.testclient import TestClient
from litellm._uuid import uuid
from litellm.models.credentials import CredentialItem
from litellm.proxy._types import (
LiteLLM_ModelTable,
@ -308,6 +309,62 @@ class TestModelManagementAuthChecks:
)
assert result is True
def test_can_user_attach_credential_non_admin_explicit_null_clear_fails(self):
from litellm.proxy._types import ProxyException
from litellm.types.router import updateLiteLLMParams as litellm_params
with pytest.raises(ProxyException) as exc_info:
ModelManagementAuthChecks.can_user_attach_credential(
litellm_params=litellm_params(litellm_credential_name=None),
user_api_key_dict=self.team_admin_user,
existing_litellm_params=LiteLLM_Params(
model="test_model", litellm_credential_name="shared-credential"
),
null_detaches=True,
)
assert exc_info.value.code == "403"
assert exc_info.value.param == "litellm_credential_name"
def test_can_user_attach_credential_admin_explicit_null_clear_succeeds(self):
from litellm.types.router import updateLiteLLMParams as litellm_params
result = ModelManagementAuthChecks.can_user_attach_credential(
litellm_params=litellm_params(litellm_credential_name=None),
user_api_key_dict=self.admin_user,
existing_litellm_params=LiteLLM_Params(
model="test_model", litellm_credential_name="shared-credential"
),
null_detaches=True,
)
assert result is True
def test_can_user_attach_credential_null_without_existing_allows_any_role(self):
from litellm.types.router import updateLiteLLMParams as litellm_params
result = ModelManagementAuthChecks.can_user_attach_credential(
litellm_params=litellm_params(litellm_credential_name=None),
user_api_key_dict=self.team_admin_user,
existing_litellm_params=LiteLLM_Params(model="test_model"),
null_detaches=True,
)
assert result is True
def test_can_user_attach_credential_null_is_noop_when_null_does_not_detach(self):
from litellm.types.router import updateLiteLLMParams as litellm_params
result = ModelManagementAuthChecks.can_user_attach_credential(
litellm_params=litellm_params(litellm_credential_name=None),
user_api_key_dict=self.team_admin_user,
existing_litellm_params=LiteLLM_Params(
model="test_model", litellm_credential_name="shared-credential"
),
)
assert result is True
def test_can_user_attach_credential_unchanged_encrypted_existing_allows_any_role(self, monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234")
encrypted_name = encrypt_value_helper(value="shared-credential")
@ -1249,6 +1306,60 @@ class TestUpdateModel:
mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once()
mock_clear_cache.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_update_model_legacy_null_credential_name_is_not_a_detach_for_non_admin(self):
from litellm.proxy.management_endpoints.model_management_endpoints import update_model
model_id = "legacy-null-credential"
existing = Deployment(
model_name="legacy-model",
litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini", litellm_credential_name="shared-credential"),
model_info={"id": model_id},
)
existing_row = MagicMock()
existing_row.litellm_params = existing.litellm_params.model_dump()
existing_row.model_dump.return_value = existing.model_dump()
updated_row = MagicMock()
updated_row.model_dump_json.return_value = "{}"
mock_prisma = MagicMock()
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row)
mock_router = MagicMock()
mock_router.get_model_ids.return_value = [model_id]
team_admin = UserAPIKeyAuth(user_id="team-admin", user_role=LitellmUserRoles.INTERNAL_USER)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.llm_router", mock_router),
patch("litellm.proxy.proxy_server.store_model_in_db", True),
patch("litellm.proxy.proxy_server.premium_user", True),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
side_effect=lambda value: value,
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.clear_cache",
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
),
):
await update_model(
model_params=updateDeployment(
litellm_params=updateLiteLLMParams(
model="openai/gpt-4o-mini", litellm_credential_name=None
),
model_info=ModelInfo(id=model_id),
),
user_api_key_dict=team_admin,
)
mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once()
persisted = json.loads(mock_prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"]["litellm_params"])
assert persisted["litellm_credential_name"] == "shared-credential"
class TestUpdatePublicModelGroups:
"""Test that update_public_model_groups correctly sets litellm.public_model_groups
@ -4000,6 +4111,401 @@ class TestUpdateDBModelClearCacheControlInjectionPoints:
assert params["tpm"] == 10
class TestUpdateDBModelClearCredentialName:
def test_explicit_null_removes_stored_credential_name(self):
from litellm.proxy.management_endpoints.model_management_endpoints import update_db_model
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
api_key="sk-real",
tpm=100,
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1", team_id="team-keep", access_groups=["prod"]),
)
update_patch: Final = updateDeployment(
litellm_params=updateLiteLLMParams(litellm_credential_name=None)
)
with patch("litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", side_effect=lambda value, **kwargs: value):
result: Final = update_db_model(db_model=db_model, updated_patch=update_patch)
params: Final = json.loads(result["litellm_params"])
info: Final = json.loads(result["model_info"])
assert "litellm_credential_name" not in params
assert params["model"] == "openai/gpt-4o"
assert params["api_base"] == "https://api.openai.com/v1"
assert params["api_key"] == "sk-real"
assert params["tpm"] == 100
assert info["team_id"] == "team-keep"
assert info["access_groups"] == ["prod"]
def test_omitted_credential_name_keeps_stored_association(self):
from litellm.proxy.management_endpoints.model_management_endpoints import update_db_model
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
api_key="sk-real",
tpm=100,
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1", team_id="team-keep", access_groups=["prod"]),
)
update_patch: Final = updateDeployment(litellm_params=updateLiteLLMParams(tpm=10))
with patch("litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", side_effect=lambda value, **kwargs: value):
result: Final = update_db_model(db_model=db_model, updated_patch=update_patch)
params: Final = json.loads(result["litellm_params"])
assert params["litellm_credential_name"] == "shared-credential"
assert params["tpm"] == 10
def test_null_clear_on_model_without_credential_is_noop(self):
from litellm.proxy.management_endpoints.model_management_endpoints import update_db_model
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(model="openai/gpt-4o", api_base="https://api.openai.com/v1"),
model_info=ModelInfo(id="dep-cred-1", team_id="team-keep", access_groups=["prod"]),
)
update_patch: Final = updateDeployment(
litellm_params=updateLiteLLMParams(litellm_credential_name=None)
)
with patch("litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", side_effect=lambda value, **kwargs: value):
result: Final = update_db_model(db_model=db_model, updated_patch=update_patch)
params: Final = json.loads(result["litellm_params"])
assert "litellm_credential_name" not in params
assert params["model"] == "openai/gpt-4o"
assert params["api_base"] == "https://api.openai.com/v1"
def test_null_credential_clear_alongside_pricing_clear(self):
from litellm.proxy.management_endpoints.model_management_endpoints import update_db_model
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
input_cost_per_token=0.000001,
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1", input_cost_per_token=0.000001),
)
update_patch: Final = updateDeployment(
litellm_params=updateLiteLLMParams(
litellm_credential_name=None,
input_cost_per_token=None,
)
)
with patch("litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", side_effect=lambda value, **kwargs: value):
result: Final = update_db_model(db_model=db_model, updated_patch=update_patch)
params: Final = json.loads(result["litellm_params"])
info: Final = json.loads(result["model_info"])
assert "litellm_credential_name" not in params
assert "input_cost_per_token" not in params
assert "input_cost_per_token" not in info
def test_replace_credential_name_keeps_other_params(self):
from litellm.proxy.management_endpoints.model_management_endpoints import update_db_model
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
api_key="sk-real",
tpm=100,
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1", team_id="team-keep", access_groups=["prod"]),
)
update_patch: Final = updateDeployment(
litellm_params=updateLiteLLMParams(litellm_credential_name="other-credential")
)
with patch("litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", side_effect=lambda value, **kwargs: value):
result: Final = update_db_model(db_model=db_model, updated_patch=update_patch)
params: Final = json.loads(result["litellm_params"])
assert params["litellm_credential_name"] == "other-credential"
assert params["api_base"] == "https://api.openai.com/v1"
assert params["api_key"] == "sk-real"
assert params["tpm"] == 100
class TestPatchModelCredentialName:
@staticmethod
async def _patch_model(
monkeypatch,
db_model: Deployment,
user_api_key_dict: UserAPIKeyAuth,
credential_name: str | None,
db_credential: CredentialItem | None = None,
credentials_repository: MagicMock | None = None,
) -> list[dict[str, object]]:
import litellm
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model, update_db_model
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="shared-credential",
credential_info={},
credential_values={"api_key": "sk-shared"},
),
CredentialItem(
credential_name="other-credential",
credential_info={},
credential_values={"api_key": "sk-other"},
),
],
)
credentials_repository = credentials_repository or MagicMock()
credentials_repository.find_by_name = AsyncMock(return_value=db_credential)
persisted: Final[list[dict[str, object]]] = []
async def persist_model(**kwargs):
row: Final = update_db_model(db_model=kwargs["db_model"], updated_patch=kwargs["patch_data"])
persisted.append(row)
updated_row: Final = MagicMock()
updated_row.model_dump_json.return_value = "{}"
return updated_row
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch("litellm.proxy.proxy_server.store_model_in_db", True),
patch("litellm.proxy.proxy_server.premium_user", True),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.CredentialsRepository",
return_value=credentials_repository,
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.get_db_model",
new=AsyncMock(return_value=db_model),
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints._update_team_model_in_db",
new=AsyncMock(side_effect=persist_model),
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.clear_cache",
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
side_effect=lambda value, **kwargs: value,
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.raise_if_reload_degraded_serving"
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.create_object_audit_log",
new=AsyncMock(),
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.live_model_ids_snapshot",
return_value=frozenset(),
),
):
await patch_model(
model_id="dep-cred-1",
patch_data=updateDeployment(
litellm_params=updateLiteLLMParams(litellm_credential_name=credential_name)
),
user_api_key_dict=user_api_key_dict,
)
return persisted
@pytest.mark.asyncio
async def test_patch_model_rejects_empty_string_credential_name(self, monkeypatch):
from litellm.proxy._types import ProxyException
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1"),
)
with pytest.raises(ProxyException) as exc_info:
await self._patch_model(
monkeypatch,
db_model,
self._admin_user(),
"",
)
assert exc_info.value.code == "400"
assert exc_info.value.param == "litellm_credential_name"
assert "empty" in exc_info.value.message.lower()
@staticmethod
def _admin_user() -> UserAPIKeyAuth:
return UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
@staticmethod
def _team_admin_user() -> UserAPIKeyAuth:
return UserAPIKeyAuth(user_id="team-admin", user_role=LitellmUserRoles.INTERNAL_USER, team_id="team-keep")
@pytest.mark.asyncio
async def test_patch_model_rejects_unknown_credential_name(self, monkeypatch):
from litellm.proxy._types import ProxyException
credentials_repository = MagicMock()
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1"),
)
with pytest.raises(ProxyException) as exc_info:
await self._patch_model(
monkeypatch,
db_model,
self._admin_user(),
"ghost-credential",
credentials_repository=credentials_repository,
)
assert exc_info.value.code == "400"
assert "not found" in exc_info.value.message.lower()
credentials_repository.find_by_name.assert_awaited_once_with("ghost-credential")
@pytest.mark.asyncio
async def test_patch_model_accepts_credential_known_only_in_db(self, monkeypatch):
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1"),
)
persisted: Final = await self._patch_model(
monkeypatch,
db_model,
self._admin_user(),
"db-only-credential",
db_credential=CredentialItem(
credential_name="db-only-credential",
credential_info={},
credential_values={"api_key": "sk-db"},
),
)
params: Final = json.loads(persisted[0]["litellm_params"])
assert params["litellm_credential_name"] == "db-only-credential"
@pytest.mark.asyncio
async def test_patch_model_replaces_credential_name_and_preserves_other_params(self, monkeypatch):
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1"),
)
persisted: Final = await self._patch_model(monkeypatch, db_model, self._admin_user(), "other-credential")
params: Final = json.loads(persisted[0]["litellm_params"])
assert params["litellm_credential_name"] == "other-credential"
assert params["api_base"] == "https://api.openai.com/v1"
@pytest.mark.asyncio
async def test_patch_model_admin_null_clear_persists_without_credential(self, monkeypatch):
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1"),
)
persisted: Final = await self._patch_model(monkeypatch, db_model, self._admin_user(), None)
params: Final = json.loads(persisted[0]["litellm_params"])
assert "litellm_credential_name" not in params
assert params["api_base"] == "https://api.openai.com/v1"
@pytest.mark.asyncio
async def test_patch_model_rejects_non_admin_explicit_null_clear(self, monkeypatch):
from litellm.proxy._types import ProxyException
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1"),
)
with pytest.raises(ProxyException) as exc_info:
await self._patch_model(monkeypatch, db_model, self._team_admin_user(), None)
assert exc_info.value.code == "403"
assert exc_info.value.param == "litellm_credential_name"
@pytest.mark.asyncio
async def test_patch_model_clear_then_reattach_round_trip(self, monkeypatch):
db_model: Final = Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o",
api_base="https://api.openai.com/v1",
litellm_credential_name="shared-credential",
),
model_info=ModelInfo(id="dep-cred-1"),
)
cleared: Final = await self._patch_model(monkeypatch, db_model, self._admin_user(), None)
cleared_model: Final = Deployment.model_validate(
{
"model_name": db_model.model_name,
"litellm_params": json.loads(cleared[0]["litellm_params"]),
"model_info": json.loads(cleared[0]["model_info"]),
}
)
reattached: Final = await self._patch_model(
monkeypatch,
cleared_model,
self._admin_user(),
"shared-credential",
)
params: Final = json.loads(reattached[0]["litellm_params"])
assert params["litellm_credential_name"] == "shared-credential"
class TestGetModelInfoWithIdBlocked:
"""`ProxyConfig.get_model_info_with_id` must propagate the DB-level `blocked`
column into the in-memory `model_info` dict so the router filter can read it."""

View file

@ -1,6 +1,7 @@
import asyncio
import threading
from collections.abc import Mapping
import time
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
@ -10,6 +11,7 @@ from fastapi import HTTPException
import litellm
from litellm.caching.dual_cache import DualCache
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
AgenticAnthropicStreamingIterator,
@ -2348,35 +2350,139 @@ async def test_release_non_numeric_counter_reseeds_from_db(spend_counter_state):
class _ExpiringRedisCache:
def __init__(self) -> None:
"""In-memory stand-in for RedisCache with real wall-clock key expiry."""
def __init__(self, default_ttl: float = 60.0, fail_first_refresh: bool = False) -> None:
self.default_ttl = default_ttl
self.store: dict[str, float] = {}
self.expires_at: dict[str, float] = {}
self.refresh_attempts = 0
self.refresh_count = 0
self.fail_first_refresh = fail_first_refresh
def _evict_expired(self, key: str) -> None:
if self.expires_at.get(key, float("inf")) <= time.monotonic():
self.store.pop(key, None)
self.expires_at.pop(key, None)
async def async_get_cache(self, key: str, *args: object, **kwargs: object) -> float | None:
self._evict_expired(key)
return self.store.get(key)
async def async_increment(self, key: str, value: float, **kwargs: object) -> float:
self._evict_expired(key)
self.store[key] = self.store.get(key, 0.0) + float(value)
self.expires_at[key] = time.monotonic() + self.default_ttl
return self.store[key]
async def async_set_max(self, key: str, value: float, **kwargs: object) -> float:
self._evict_expired(key)
self.store[key] = max(self.store.get(key, float("-inf")), float(value))
self.expires_at[key] = time.monotonic() + self.default_ttl
return self.store[key]
async def async_set_cache(self, key: str, value: float, *args: object, **kwargs: object) -> bool:
self.store[key] = float(value)
self.expires_at[key] = time.monotonic() + self.default_ttl
return True
async def async_delete_cache(self, key: str, *args: object, **kwargs: object) -> None:
self.store.pop(key, None)
self.expires_at.pop(key, None)
async def async_increment_pipeline(self, increment_list, **kwargs):
results = []
for op in increment_list:
results.append(await self.async_increment(op["key"], op["increment_value"]))
return results
async def async_refresh_ttl(self, key: str, ttl: int | None = None) -> bool:
self.refresh_attempts += 1
if self.fail_first_refresh and self.refresh_attempts == 1:
raise ConnectionError("Redis circuit breaker is open")
self._evict_expired(key)
if key not in self.store:
return False
self.refresh_count += 1
self.expires_at[key] = time.monotonic() + (ttl if ttl is not None else self.default_ttl)
return True
def get_ttl(self, **kwargs) -> None:
return None
async def async_increment_pipeline(
self, increment_list: Sequence[RedisPipelineIncrementOperation], **kwargs: object
) -> list[float]:
return [await self.async_increment(op["key"], op["increment_value"]) for op in increment_list]
def get_ttl(self, **kwargs: object) -> int | None:
return int(self.default_ttl)
@pytest.mark.asyncio
async def test_reservation_survives_redis_counter_ttl_while_request_in_flight(
spend_counter_state,
):
"""A request that runs longer than the counter TTL must keep its reservation in Redis
(so a concurrent request on any worker still sees it), and renewal must stop once the
reservation is reconciled so an idle counter still expires on its own."""
counter_cache, key_cache = spend_counter_state
redis_cache = _ExpiringRedisCache(default_ttl=0.2)
counter_cache.redis_cache = redis_cache
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(token="key-lease", spend=0.0, max_budget=1.0)
counter_key = "spend:key:key-lease"
reservation = await _reserve(valid_token, 0.6, key_cache, proxy_logging_obj)
assert reservation is not None
await asyncio.sleep(0.5)
assert await redis_cache.async_get_cache(key=counter_key) == pytest.approx(0.6)
concurrent = await _reserve(valid_token, 0.6, key_cache, proxy_logging_obj)
assert concurrent is not None
assert concurrent["reserved_cost"] == pytest.approx(0.4)
await release_budget_reservation(reservation)
await release_budget_reservation(concurrent)
await asyncio.sleep(0.15)
refreshes_after_release = redis_cache.refresh_count
await asyncio.sleep(0.35)
assert redis_cache.refresh_count == refreshes_after_release
assert await redis_cache.async_get_cache(key=counter_key) is None
@pytest.mark.asyncio
async def test_reservation_lease_keeps_renewing_after_transient_redis_failure(
spend_counter_state,
):
"""One failed EXPIRE (Redis blip, open circuit breaker) must not end renewal for the
rest of the request."""
counter_cache, key_cache = spend_counter_state
redis_cache = _ExpiringRedisCache(default_ttl=0.2, fail_first_refresh=True)
counter_cache.redis_cache = redis_cache
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(token="key-lease-blip", spend=0.0, max_budget=1.0)
reservation = await _reserve(valid_token, 0.6, key_cache, proxy_logging_obj)
assert reservation is not None
await asyncio.sleep(0.5)
assert redis_cache.refresh_attempts >= 3
await release_budget_reservation(reservation)
@pytest.mark.asyncio
async def test_reservation_lease_stops_when_request_task_ends_without_reconciling(
spend_counter_state,
):
"""A request whose task ends without reconciling (client disconnect path that skips the
cost callbacks) must not keep renewing: the counter falls back to its plain TTL instead of
pinning the reservation until the request timeout."""
counter_cache, key_cache = spend_counter_state
redis_cache = _ExpiringRedisCache(default_ttl=0.2)
counter_cache.redis_cache = redis_cache
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(token="key-lease-orphan", spend=0.0, max_budget=1.0)
counter_key = "spend:key:key-lease-orphan"
reservation = await asyncio.create_task(_reserve(valid_token, 0.6, key_cache, proxy_logging_obj))
assert reservation is not None
assert reservation["finalized"] is False
await asyncio.sleep(0.5)
assert redis_cache.refresh_count == 0
assert await redis_cache.async_get_cache(key=counter_key) is None
class _TeamMembershipFloorDb:

View file

@ -443,6 +443,20 @@ class TestAnthropicBetaHeadersFiltering:
assert filtered == ["thinking-binding-controls-2026-08-01"]
@pytest.mark.parametrize("provider", ["anthropic", "bedrock", "bedrock_mantle", "vertex_ai"])
def test_dangerous_tool_use_forwarded(self, provider):
"""Claude Code's server-side auto-mode classifier sends `safeguards` together with
dangerous-tool-use-2026-09-03. Bedrock Invoke, Bedrock Mantle, and Vertex rawPredict
all answer "safeguards: Extra inputs are not permitted" when the body field arrives
without the beta (probed 2026-09-21), so dropping the header turned every auto-mode
turn into a 400 on Vertex and silently disabled the classifier on Bedrock."""
filtered = filter_and_transform_beta_headers(
beta_headers=["dangerous-tool-use-2026-09-03"],
provider=provider,
)
assert filtered == ["dangerous-tool-use-2026-09-03"]
def test_null_value_headers_filtered(self):
"""Test that headers with null values are always filtered out."""
for provider in [

View file

@ -102,7 +102,7 @@ export interface ModelEditFormValues {
vector_store_ids?: string[];
tags?: string[];
health_check_model?: string | null;
litellm_credential_name?: string;
litellm_credential_name?: string | null;
litellm_extra_params?: string;
model_info?: string;
team_id?: string;
@ -139,7 +139,7 @@ const modelEditShape = {
vector_store_ids: z.array(z.string()).optional(),
tags: z.array(z.string()).optional(),
health_check_model: z.string().nullish(),
litellm_credential_name: textish,
litellm_credential_name: z.string().nullish(),
litellm_extra_params: textish,
model_info: textish,
team_id: textish,
@ -254,7 +254,7 @@ export const toModelEditFormValues = (localModelData: any, isWildcardModel: bool
tags: Array.isArray(localModelData.litellm_params?.tags) ? localModelData.litellm_params.tags : [],
// antd never mounted this field for a non-wildcard model, so the key must be absent, not null.
...(isWildcardModel ? { health_check_model: localModelData.model_info?.health_check_model } : {}),
litellm_credential_name: localModelData.litellm_params?.litellm_credential_name || "",
litellm_credential_name: localModelData.litellm_params?.litellm_credential_name ?? null,
litellm_extra_params: JSON.stringify(
Object.fromEntries(
Object.entries(localModelData.litellm_params || {}).filter(
@ -635,8 +635,8 @@ const ModelInfoEditForm: React.FC<ModelInfoEditFormProps> = ({
{isEditing ? (
<FormField control={form.control} name="litellm_credential_name">
{({ id, value, onChange, onBlur }) => {
const items = [
{ value: "", label: "None" },
const items: { value: string | null; label: string }[] = [
{ value: null, label: "None" },
...credentialsList.map((credential) => ({
value: credential.credential_name,
label: credential.credential_name,
@ -645,15 +645,15 @@ const ModelInfoEditForm: React.FC<ModelInfoEditFormProps> = ({
return (
<Select
items={items}
value={(value as string) ?? ""}
onValueChange={(selected: string | null) => onChange(selected ?? "")}
value={(value as string | null) ?? null}
onValueChange={(selected: string | null) => onChange(selected)}
>
<SelectTrigger id={id} className="w-full" onBlur={onBlur}>
<SelectValue placeholder="Select or search for existing credentials" />
</SelectTrigger>
<SelectContent>
{items.map((item) => (
<SelectItem key={item.value} value={item.value}>
<SelectItem key={item.value ?? "none"} value={item.value}>
{item.label}
</SelectItem>
))}

View file

@ -655,7 +655,7 @@ describe("ModelInfoView", () => {
});
const updatePayload = mockModelPatchUpdateCall.mock.calls[0][1];
expect(updatePayload.litellm_params.litellm_credential_name).toBe("selected-credential");
expect(updatePayload.litellm_params).not.toHaveProperty("litellm_credential_name");
expect(updatePayload.litellm_params.litellm_credential_name).not.toBe("from-json");
});
@ -1545,6 +1545,18 @@ describe("ModelInfoView", () => {
await screen.findByRole("combobox", { expanded: true });
};
const openCredentialSelect = async (user: ReturnType<typeof userEvent.setup>, triggerText?: string) => {
const trigger = screen
.getAllByRole("combobox")
.filter((element) => element.getAttribute("data-slot") === "select-trigger")
.find((element) => triggerText === undefined || element.textContent?.includes(triggerText));
if (trigger === undefined) {
throw new Error(`Could not find credential selector${triggerText ? ` with ${triggerText}` : ""}`);
}
await user.click(trigger);
await screen.findByRole("combobox", { expanded: true });
};
const save = async (user: ReturnType<typeof userEvent.setup>) => {
await user.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => expect(mockModelPatchUpdateCall).toHaveBeenCalled());
@ -1566,7 +1578,6 @@ describe("ModelInfoView", () => {
model: "gpt-4",
api_base: "https://api.openai.com/v1",
custom_llm_provider: "openai",
litellm_credential_name: "selected-credential",
tags: [],
guardrails: [],
},
@ -1805,6 +1816,104 @@ describe("ModelInfoView", () => {
expect(payload.litellm_params.litellm_credential_name).toBe("other-credential");
});
it("sends explicit null when None is picked for a model with a stored credential", async () => {
const user = userEvent.setup();
await enterEditMode(user);
await openCredentialSelect(user, "selected-credential");
await user.click(await screen.findByRole("option", { name: "None" }));
const payload = await save(user);
expect(payload.litellm_params.litellm_credential_name).toBeNull();
expect("litellm_credential_name" in payload.litellm_params).toBe(true);
});
it("omits the credential when it is cleared and then restored before saving", async () => {
const user = userEvent.setup();
await enterEditMode(user);
await openCredentialSelect(user, "selected-credential");
await user.click(await screen.findByRole("option", { name: "None" }));
await openCredentialSelect(user);
await user.click(await screen.findByRole("option", { name: "selected-credential" }));
const payload = await save(user);
expect(payload.litellm_params).not.toHaveProperty("litellm_credential_name");
});
it("omits the credential when None is picked for a model that never had one", async () => {
const { litellm_credential_name: _storedCredential, ...litellmParamsWithoutCredential } =
defaultModelData.litellm_params;
const modelWithoutCredential = {
...defaultModelData,
litellm_params: litellmParamsWithoutCredential,
};
mockUseModelsInfo.mockReturnValue({ data: { data: [modelWithoutCredential] }, isLoading: false, error: null });
mockModelInfoV1Call.mockResolvedValue({ data: [modelWithoutCredential] });
const user = userEvent.setup();
await enterEditMode(user);
await openCredentialSelect(user);
await user.click(await screen.findByRole("option", { name: "None" }));
const payload = await save(user);
expect(payload.litellm_params).not.toHaveProperty("litellm_credential_name");
expect(payload.litellm_params.litellm_credential_name).not.toBe("");
});
it("restores the stored credential in the selector after cancel", async () => {
const user = userEvent.setup();
await enterEditMode(user);
await openCredentialSelect(user, "selected-credential");
await user.click(await screen.findByRole("option", { name: "None" }));
await user.click(screen.getByRole("button", { name: /cancel/i }));
await user.click(await screen.findByRole("button", { name: /edit settings/i }));
const credentialTrigger: HTMLElement = screen
.getAllByRole("combobox")
.filter((element) => element.getAttribute("data-slot") === "select-trigger")
.at(0) as HTMLElement;
expect(credentialTrigger).toHaveTextContent("selected-credential");
});
it("shows Manual in read mode after saving None", async () => {
const user = userEvent.setup();
await enterEditMode(user);
await openCredentialSelect(user, "selected-credential");
await user.click(await screen.findByRole("option", { name: "None" }));
await save(user);
expect(await screen.findByText("Manual")).toBeInTheDocument();
expect(screen.queryByText("selected-credential")).not.toBeInTheDocument();
});
it("keeps showing the stored credential in read mode after an untouched save", async () => {
const user = userEvent.setup();
await enterEditMode(user);
await save(user);
expect(await screen.findByText("selected-credential")).toBeInTheDocument();
});
it("keeps the form open and surfaces the error when the backend rejects the detach", async () => {
mockModelPatchUpdateCall.mockRejectedValueOnce(new Error("403 Only a proxy admin can detach"));
const user = userEvent.setup();
await enterEditMode(user);
await openCredentialSelect(user, "selected-credential");
await user.click(await screen.findByRole("option", { name: "None" }));
await user.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => expect(mockToast.fromError).toHaveBeenCalled());
expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument();
expect(mockToast.success).not.toHaveBeenCalledWith("Model settings updated successfully");
});
it("sends the vector stores picked in the knowledge base selector", async () => {
const user = userEvent.setup();
await enterEditMode(user);

View file

@ -340,8 +340,10 @@ export default function ModelInfoView({
}
}
if (values.litellm_credential_name) {
updatedLitellmParams.litellm_credential_name = values.litellm_credential_name;
const storedCredentialName: string | null = localModelData?.litellm_params?.litellm_credential_name ?? null;
const selectedCredentialName: string | null = values.litellm_credential_name ?? null;
if (selectedCredentialName !== storedCredentialName) {
updatedLitellmParams.litellm_credential_name = selectedCredentialName;
} else {
delete updatedLitellmParams.litellm_credential_name;
}
@ -397,6 +399,7 @@ export default function ModelInfoView({
// without this strip a masked value would be re-encrypted over the real secret.
// Credential rotation has its own dedicated path (UpdateModelCredentialsModal).
const safeLitellmParams = stripMaskedSecrets(updatedLitellmParams);
const { litellm_credential_name: _sentCredential, ...localLitellmParams } = safeLitellmParams;
const updateData = {
model_name: values.model_name,
@ -410,7 +413,10 @@ export default function ModelInfoView({
...localModelData,
model_name: values.model_name,
litellm_model_name: values.litellm_model_name,
litellm_params: safeLitellmParams,
litellm_params:
selectedCredentialName === null
? localLitellmParams
: { ...localLitellmParams, litellm_credential_name: selectedCredentialName },
model_info: updatedModelInfo,
};