diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json
index 57ca267e504..26e4e06a796 100644
--- a/basedpyright-code-budget.json
+++ b/basedpyright-code-budget.json
@@ -105,7 +105,7 @@
"limit": 109
},
"reportUnknownMemberType": {
- "limit": 38271
+ "limit": 38269
},
"reportUnknownParameterType": {
"limit": 19584
diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
index 486904d0abe..c7e1b94a2ef 100644
--- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
+++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
@@ -1801,7 +1801,16 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Remove conflicting keys from data to avoid duplicate keyword arguments
filtered_data = {k: v for k, v in data.items() if k not in ("model", "file_id")}
for model_id, model_file_id in specific_model_file_id_mapping.items():
- delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **filtered_data) # type: ignore
+ credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
+ delete_data = {
+ **{k: v for k, v in filtered_data.items() if k != "_litellm_internal_model_credentials"},
+ **(
+ {"_litellm_internal_model_credentials": MappingProxyType(dict(credentials))}
+ if credentials is not None
+ else {}
+ ),
+ }
+ delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
stored_file_object = await self.delete_unified_file_id(file_id, litellm_parent_otel_span)
@@ -1812,7 +1821,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
prom_logger.record_managed_file_deleted(result="success")
if stored_file_object:
- return stored_file_object
+ return OpenAIFileObject.model_validate(stored_file_object).model_copy(update={"id": file_id})
elif delete_response:
delete_response.id = file_id
return delete_response
diff --git a/litellm/__init__.py b/litellm/__init__.py
index ede8a73453d..5a461801b62 100644
--- a/litellm/__init__.py
+++ b/litellm/__init__.py
@@ -546,7 +546,7 @@ _key_management_system: Optional["KeyManagementSystem"] = None
#### PII MASKING ####
output_parse_pii: bool = False
#############################################
-from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
+from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map, mark_litellm_import_complete
model_cost = get_model_cost_map(url=model_cost_map_url)
cost_discount_config: Dict[str, float] = {} # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount
@@ -2405,3 +2405,5 @@ def __getattr__(name: str) -> Any:
# ALL_LITELLM_RESPONSE_TYPES is lazy-loaded via __getattr__ to avoid loading utils at import time
+
+mark_litellm_import_complete()
diff --git a/litellm/_internal_context.py b/litellm/_internal_context.py
index f856fe0f2b3..8132008731f 100644
--- a/litellm/_internal_context.py
+++ b/litellm/_internal_context.py
@@ -6,9 +6,33 @@ be settable from user input. Context variables are scoped to the current
asyncio task and cannot be injected via HTTP request bodies.
"""
+from collections.abc import Generator
+from contextlib import contextmanager
from contextvars import ContextVar
+from datetime import datetime, timezone
from typing import Final
# When True, suppresses async logging and billing for internal sub-calls
# (e.g., emulated file-search steps that make nested LLM calls).
is_internal_call: Final[ContextVar[bool]] = ContextVar("is_internal_call", default=False)
+
+# One request prices its totals, its per-token-type lines and the rates it reports on
+# separate code paths. Each reads the clock for off-peak pricing, so without a pinned
+# moment they can land on either side of a window boundary and disagree with each other.
+_billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", default=None)
+
+
+@contextmanager
+def pinned_billing_time(moment: datetime) -> Generator[None]:
+ """Price every rate lookup inside this block at ``moment`` rather than at each one's own clock read."""
+ token: Final = _billing_time.set(moment)
+ try:
+ yield
+ finally:
+ _billing_time.reset(token)
+
+
+def current_billing_time() -> datetime:
+ """The pinned billing moment, or now in UTC outside a pinned block."""
+ pinned: Final = _billing_time.get()
+ return pinned if pinned is not None else datetime.now(timezone.utc)
diff --git a/litellm/constants.py b/litellm/constants.py
index 6f2384c8c6c..f9389d22dea 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -143,6 +143,7 @@ DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD: Final = float(
os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD", 0.3)
)
MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH: Final = int(os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150))
+MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH: Final = 2048
DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS: Final = 2000
@@ -197,6 +198,7 @@ LITELLM_UI_ALLOW_HEADERS: Final = [
"x-litellm-adaptive-router-model",
"x-litellm-applied-guardrails",
"x-litellm-guardrail-scan-id",
+ "x-litellm-guardrail-scan-metadata",
"x-litellm-cache-key",
]
@@ -333,6 +335,7 @@ DEFAULT_SSL_CIPHERS: Final = os.getenv(
########### v2 Architecture constants for managing writing updates to the database ###########
REDIS_UPDATE_BUFFER_KEY: Final = "litellm_spend_update_buffer"
+REDIS_GATEWAY_REQUESTS_BUFFER_KEY: Final = "litellm_gateway_requests_buffer"
REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_spend_update_buffer"
REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_team_spend_update_buffer"
REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_org_spend_update_buffer"
diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py
index d4c6c87efc8..814eaaf76f7 100644
--- a/litellm/cost_calculator.py
+++ b/litellm/cost_calculator.py
@@ -25,6 +25,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
TranscriptionUsageObjectTransformation,
)
from litellm.litellm_core_utils.llm_cost_calc.utils import (
+ BilledTokenRates,
CostCalculatorUtils,
_generic_cost_per_character,
_get_regional_uplift_multiplier,
@@ -1125,6 +1126,7 @@ def _store_cost_breakdown_in_logging_obj(
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
+ billed_token_rates: BilledTokenRates | None = None,
) -> None:
"""
Helper function to store cost breakdown in the logging object.
@@ -1169,6 +1171,7 @@ def _store_cost_breakdown_in_logging_obj(
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
+ billed_token_rates=billed_token_rates,
)
except Exception as breakdown_error:
@@ -1737,6 +1740,7 @@ def completion_cost(
_reasoning_cost: float | None = None
_cache_read_cost: float | None = None
_cache_creation_cost: float | None = None
+ _billed_token_rates: BilledTokenRates | None = None
if cost_per_token_usage_object is not None and model:
_breakdown_provider: str | None = (
custom_llm_provider if isinstance(custom_llm_provider, str) else None
@@ -1748,10 +1752,12 @@ def completion_cost(
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
+ custom_cost_per_token=custom_cost_per_token,
)
_reasoning_cost = _token_type_breakdown.reasoning_cost
_cache_read_cost = _token_type_breakdown.cache_read_cost
_cache_creation_cost = _token_type_breakdown.cache_creation_cost
+ _billed_token_rates = _token_type_breakdown.rates
_store_cost_breakdown_in_logging_obj(
litellm_logging_obj=litellm_logging_obj,
prompt_tokens_cost_usd_dollar=prompt_tokens_cost_usd_dollar,
@@ -1771,6 +1777,7 @@ def completion_cost(
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
+ billed_token_rates=_billed_token_rates,
)
return _final_cost
diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py
index 3503468c735..8c7f4557992 100644
--- a/litellm/experimental_mcp_client/client.py
+++ b/litellm/experimental_mcp_client/client.py
@@ -18,6 +18,7 @@ from mcp import ClientSession, McpError, ReadResourceResult, Resource, StdioServ
from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client
from mcp.shared.message import SessionMessage
+from mcp.shared.session import RequestResponder
from typing_extensions import Unpack
_TransportStreams: TypeAlias = tuple[
@@ -56,10 +57,13 @@ def missing_streamable_http_client_error() -> ImportError:
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import (
+ ClientResult,
GetPromptRequestParams,
GetPromptResult,
Prompt,
ResourceTemplate,
+ ServerNotification,
+ ServerRequest,
TextContent,
)
from mcp.types import Tool as MCPTool
@@ -146,8 +150,8 @@ _SDK_READ_TIMEOUT_CODE: Final = int(httpx.codes.REQUEST_TIMEOUT)
otherwise carries JSON-RPC error codes."""
-def _as_read_timeout(exc: BaseException) -> TimeoutError | None:
- """The session read timeout elapsing, re-expressed as a ``TimeoutError``, or ``None``.
+def as_mcp_read_timeout(exc: BaseException) -> TimeoutError | None:
+ """Normalize an MCP SDK read timeout for client and gateway diagnostics, or return ``None``.
The SDK reports its own elapsed read timeout as ``McpError`` carrying an HTTP status code in a
field that otherwise holds JSON-RPC error codes, and it relays an upstream's JSON-RPC error
@@ -442,6 +446,18 @@ class MCPClient:
in_flight_error: BaseException | None = None
try:
read_stream, write_stream = transport[0], transport[1]
+ stream_error: Final[asyncio.Future[Exception]] = asyncio.get_running_loop().create_future()
+
+ async def receive_message(
+ message: RequestResponder[ServerRequest, ClientResult] | ServerNotification | Exception,
+ ) -> None:
+ if not isinstance(message, (ValueError, httpx.RequestError, OSError)):
+ return
+ if not stream_error.done():
+ stream_error.set_result(message)
+ # The SDK closes pending requests when its message handler raises.
+ raise RuntimeError("MCP response stream failed")
+
# Build session kwargs with optional callbacks
session_kwargs: Final[dict[str, Any]] = {}
if self._sampling_callback is not None:
@@ -456,6 +472,7 @@ class MCPClient:
read_stream,
write_stream,
read_timeout_seconds=timedelta(seconds=self.timeout),
+ message_handler=receive_message,
**session_kwargs,
)
session: Final = await session_ctx.__aenter__()
@@ -467,6 +484,10 @@ class MCPClient:
if isinstance(ins, str) and ins.strip():
self._last_initialize_instructions = ins.strip()
return await operation(session)
+ except McpError:
+ if stream_error.done():
+ raise stream_error.result()
+ raise
finally:
try:
await session_ctx.__aexit__(None, None, None)
@@ -501,11 +522,10 @@ class MCPClient:
transport_ctx, http_client = self._create_transport_context()
return await self._execute_session_operation(transport_ctx, operation)
except Exception as e:
- read_timeout: Final = _as_read_timeout(e)
+ read_timeout: Final = as_mcp_read_timeout(e)
if read_timeout is not None:
verbose_logger.warning(
- "MCP client timed out after %ss waiting for %s to answer; the server accepted the "
- "request and ended its response stream without a JSON-RPC reply",
+ "MCP client timed out after %ss waiting for a valid MCP response from %s",
self.timeout,
self.server_url or "stdio",
)
diff --git a/litellm/files/main.py b/litellm/files/main.py
index 19da77b7364..218518eb3cd 100644
--- a/litellm/files/main.py
+++ b/litellm/files/main.py
@@ -31,7 +31,7 @@ FileCreateProvider = Literal[
FileRetrieveProvider = Literal[
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic"
]
-FileDeleteProvider = Literal["openai", "azure", "gemini", "litellm_proxy", "manus", "anthropic"]
+FileDeleteProvider = Literal["openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic"]
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic"]
import litellm
from litellm import get_secret_str
diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py
index 37d6a7e793d..77bf4820a1a 100644
--- a/litellm/integrations/custom_guardrail.py
+++ b/litellm/integrations/custom_guardrail.py
@@ -850,20 +850,24 @@ class CustomGuardrail(CustomLogger):
if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True:
return None
- # CHECK IF GUARDRAIL REJECTS THE REQUEST
target: Final = self._deployment_hook_target()
- hook_request_data: Final = {**request_data, "guardrail_to_apply": self} if target is not self else request_data
- result: Final = await target.async_post_call_success_hook(
- user_api_key_dict=UserAPIKeyAuth(
- user_id=request_data.get("user_api_key_user_id"),
- team_id=request_data.get("user_api_key_team_id"),
- end_user_id=request_data.get("user_api_key_end_user_id"),
- api_key=request_data.get("user_api_key_hash"),
- request_route=request_data.get("user_api_key_request_route"),
- ),
- data=hook_request_data,
- response=response,
- )
+ try:
+ if target is not self:
+ request_data["guardrail_to_apply"] = self # rebind-ok: dispatch consumes this key
+ result: Final = await target.async_post_call_success_hook(
+ user_api_key_dict=UserAPIKeyAuth(
+ user_id=request_data.get("user_api_key_user_id"),
+ team_id=request_data.get("user_api_key_team_id"),
+ end_user_id=request_data.get("user_api_key_end_user_id"),
+ api_key=request_data.get("user_api_key_hash"),
+ request_route=request_data.get("user_api_key_request_route"),
+ ),
+ data=request_data,
+ response=response,
+ )
+ finally:
+ if target is not self:
+ request_data.pop("guardrail_to_apply", None)
if not self._is_valid_response_type(result):
return None
diff --git a/litellm/litellm_core_utils/get_model_cost_map.py b/litellm/litellm_core_utils/get_model_cost_map.py
index cdc4810ff04..f81ddbfee2e 100644
--- a/litellm/litellm_core_utils/get_model_cost_map.py
+++ b/litellm/litellm_core_utils/get_model_cost_map.py
@@ -13,6 +13,7 @@ import hashlib
import json
import os
import random
+import threading
import time
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, replace
@@ -176,6 +177,11 @@ class GetModelCostMap:
RETRYABLE_FETCH_STATUS_CODES: Final = frozenset({429, 500, 502, 503, 504})
MODEL_COST_MAP_FETCH_MAX_ATTEMPTS: Final = 3
MODEL_COST_MAP_FETCH_MAX_WAIT_SECONDS: Final = 30.0
+_litellm_import_complete = threading.Event()
+
+
+def mark_litellm_import_complete() -> None:
+ _litellm_import_complete.set()
@dataclass(frozen=True, slots=True)
@@ -314,12 +320,13 @@ async def _fetch_remote_model_cost_map_with_retry(
def _fetch_remote_model_cost_map_with_retry_sync(
url: str,
timeout: int,
- max_attempts: int,
+ attempts: range,
sleep: Callable[[float], None],
rng: random.Random,
client: _SyncGetClient,
) -> ModelCostMapReloadResult:
- for attempt in range(1, max_attempts + 1):
+ max_attempts: Final = attempts.stop - 1
+ for attempt in attempts:
outcome = _attempt_fetch_sync(client=client, url=url, timeout=timeout)
if not isinstance(outcome, _FetchAttemptRetryable):
return outcome
@@ -520,6 +527,68 @@ def _finalize_loaded_model_cost_map(loaded: ModelCostMapReloaded) -> ModelCostMa
return replace(loaded, model_cost_map=_finalize_model_cost_map(loaded.model_cost_map))
+def adopt_model_cost_map(
+ new_model_cost_map: dict, # mutable-ok: public API preserves the mutable cost-map contract
+) -> int:
+ import litellm
+ from litellm import utils
+
+ litellm.model_cost = new_model_cost_map
+ utils._invalidate_model_cost_lowercase_map() # pyright: ignore[reportPrivateUsage] # required cache invalidation
+ litellm.add_known_models(model_cost_map=new_model_cost_map)
+ fetched_model_count: Final = len(new_model_cost_map) if new_model_cost_map else 0
+ utils.reapply_runtime_model_cost_registrations()
+ return fetched_model_count
+
+
+def _retry_remote_fetch_in_background(
+ url: str,
+ timeout: int,
+ max_attempts: int,
+ sleep: Callable[[float], None],
+ rng: random.Random,
+ client: _SyncGetClient,
+ first_outcome: _FetchAttemptRetryable,
+) -> None:
+ try:
+ first_wait: Final = _next_retry_wait(outcome=first_outcome, attempt=1, max_attempts=max_attempts, rng=rng)
+ if isinstance(first_wait, ModelCostMapReloadUnavailable):
+ return
+ sleep(first_wait)
+ result: Final = _fetch_remote_model_cost_map_with_retry_sync(
+ url=url,
+ timeout=timeout,
+ attempts=range(2, max_attempts + 1),
+ sleep=sleep,
+ rng=rng,
+ client=client,
+ )
+ if isinstance(result, ModelCostMapReloadUnavailable):
+ verbose_logger.warning(
+ "LiteLLM: Failed to fetch remote model cost map from %s after %d attempts; keeping local backup",
+ url,
+ max_attempts,
+ )
+ return
+ _litellm_import_complete.wait()
+ if not GetModelCostMap.validate_model_cost_map(
+ fetched_map=result.model_cost_map,
+ backup_model_count=GetModelCostMap._get_backup_model_count(), # pyright: ignore[reportPrivateUsage] # integrity cache
+ ):
+ verbose_logger.warning(
+ "LiteLLM: Fetched model cost map failed integrity check. Using local backup instead. url=%s",
+ url,
+ )
+ return
+ finalized: Final = _finalize_loaded_model_cost_map(result).model_cost_map
+ _cost_map_source_info.source = "remote"
+ _cost_map_source_info.fallback_reason = None
+ _cost_map_source_info.loaded_at = datetime.now(timezone.utc)
+ adopt_model_cost_map(finalized)
+ except Exception as e:
+ verbose_logger.warning("LiteLLM: Background model cost map retry failed: %s", e)
+
+
def get_model_cost_map(
url: str,
timeout: int = 5,
@@ -532,9 +601,7 @@ def get_model_cost_map(
Public entry point — returns the model cost map dict.
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set, uses the local backup only.
- 2. Otherwise fetches from ``url``, retrying transient HTTP errors
- (429/5xx/transport) with Retry-After-aware backoff, validates
- integrity, and falls back to the local backup on any failure.
+ 2. Otherwise fetches from ``url``, retrying transient errors in a background thread.
Only the backup model count is cached (a single int) for validation.
The full backup dict is only parsed when it must be *returned* as a
@@ -553,24 +620,34 @@ def get_model_cost_map(
_cost_map_source_info.url = url
_cost_map_source_info.is_env_forced = False
- result: Final = _fetch_remote_model_cost_map_with_retry_sync(
- url=url,
- timeout=timeout,
- max_attempts=max_attempts,
- sleep=sleep,
- rng=rng if rng is not None else random.Random(),
- client=client if client is not None else httpx,
- )
- if isinstance(result, ModelCostMapReloadUnavailable):
+ fetch_client: Final = client if client is not None else httpx
+ fetch_rng: Final = rng if rng is not None else random.Random()
+ outcome: Final = _attempt_fetch_sync(client=fetch_client, url=url, timeout=timeout)
+ if isinstance(outcome, _FetchAttemptRetryable) and max_attempts > 1:
+ threading.Thread(
+ target=_retry_remote_fetch_in_background,
+ kwargs={ # mutable-ok: threading requires a mutable keyword-arguments mapping
+ "url": url,
+ "timeout": timeout,
+ "max_attempts": max_attempts,
+ "sleep": sleep,
+ "rng": fetch_rng,
+ "client": fetch_client,
+ "first_outcome": outcome,
+ },
+ name="litellm-model-cost-map-retry",
+ daemon=True,
+ ).start()
+ if not isinstance(outcome, ModelCostMapReloaded):
verbose_logger.warning(
"LiteLLM: Failed to fetch remote model cost map from %s: %s. Falling back to local backup.",
url,
- result.reason,
+ outcome.reason,
)
_cost_map_source_info.source = "local"
- _cost_map_source_info.fallback_reason = f"Remote fetch failed: {result.reason}"
+ _cost_map_source_info.fallback_reason = f"Remote fetch failed: {outcome.reason}"
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision()).model_cost_map
- content: Final = result.model_cost_map
+ content: Final = outcome.model_cost_map
# Validate using cached count (cheap int comparison, no file I/O)
if not GetModelCostMap.validate_model_cost_map(
@@ -587,4 +664,4 @@ def get_model_cost_map(
_cost_map_source_info.source = "remote"
_cost_map_source_info.fallback_reason = None
- return _finalize_loaded_model_cost_map(result).model_cost_map
+ return _finalize_loaded_model_cost_map(outcome).model_cost_map
diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py
index b0d6db20b31..e6b2bb164ef 100644
--- a/litellm/litellm_core_utils/litellm_logging.py
+++ b/litellm/litellm_core_utils/litellm_logging.py
@@ -203,6 +203,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
+ from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
@@ -590,6 +591,7 @@ class Logging(LiteLLMLoggingBaseClass):
# Initialize cost breakdown field
self.cost_breakdown: CostBreakdown | None = None
+ self.billed_token_rates: BilledTokenRates | None = None
# Init Caching related details
self.caching_details: CachingDetails | None = None
@@ -1587,6 +1589,7 @@ class Logging(LiteLLMLoggingBaseClass):
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
+ billed_token_rates: "BilledTokenRates | None" = None,
) -> None:
"""
Helper method to store cost breakdown in the logging object.
@@ -1606,8 +1609,10 @@ class Logging(LiteLLMLoggingBaseClass):
service_tier: Tier the costs above were priced on, already resolved
data_residency: Region uplift the costs above were priced on, already resolved
vertex_location: Vertex AI location the costs above were priced on, already resolved
+ billed_token_rates: Per-token rates the costs above were billed at, already resolved
"""
+ self.billed_token_rates = billed_token_rates
self.cost_breakdown = CostBreakdown(
input_cost=input_cost,
output_cost=output_cost,
diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py
index c05d4c29a5e..5675c59733d 100644
--- a/litellm/litellm_core_utils/llm_cost_calc/utils.py
+++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py
@@ -10,6 +10,7 @@ from typing import Any, Final, Literal, TypedDict, cast
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
import litellm
+from litellm._internal_context import current_billing_time
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import (
select_tier_for_input,
@@ -19,6 +20,7 @@ from litellm.types.utils import (
CacheCreationTokenDetails,
CallTypes,
CompletionTokensDetailsWrapper,
+ CostPerToken,
DataResidency,
ImageResponse,
ModelInfo,
@@ -305,7 +307,7 @@ def _is_within_off_peak_window(off_peak_hours_utc: str | Sequence[str], current_
than being localised, so callers must pass datetime.now(timezone.utc), never datetime.now(),
or every window shifts by the host's offset.
"""
- reference: Final = current_time if current_time is not None else datetime.now(timezone.utc)
+ reference: Final = current_time if current_time is not None else current_billing_time()
now: Final = (reference.astimezone(timezone.utc) if reference.tzinfo is not None else reference).time()
windows: Final = (off_peak_hours_utc,) if isinstance(off_peak_hours_utc, str) else off_peak_hours_utc
for window in windows:
@@ -392,7 +394,7 @@ def _is_off_peak(off_peak: Mapping[str, object], current_time: datetime | None =
rules: the flat hours_utc windows, which apply every day, or any entry in windows, whose
hours apply only on its weekdays.
"""
- reference: Final = current_time if current_time is not None else datetime.now(timezone.utc)
+ reference: Final = current_time if current_time is not None else current_billing_time()
reference_utc: Final = (
reference.astimezone(timezone.utc) if reference.tzinfo is not None else reference.replace(tzinfo=timezone.utc)
)
@@ -1195,7 +1197,7 @@ def generic_cost_per_token(
usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens, 0
)
- billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
+ billing_time: Final = current_time if current_time is not None else current_billing_time()
(
prompt_base_cost,
completion_base_cost,
@@ -1309,42 +1311,90 @@ def _coerce_token_count(value: object) -> int:
return value if isinstance(value, int) and value > 0 else 0
+@dataclass(frozen=True, slots=True)
+class BilledTokenRates:
+ """Per-token rates one request's usage bills at, after token tiers, off-peak windows and the
+ regional multipliers the totals apply, so each cost line equals its token count times its rate."""
+
+ input_cost_per_token: float
+ output_cost_per_token: float
+ cache_read_input_token_cost: float
+ cache_creation_input_token_cost: float
+ cache_creation_input_token_cost_above_1hr: float
+ output_cost_per_reasoning_token: float
+
+ def scaled(self, multiplier: float) -> "BilledTokenRates":
+ if multiplier == 1.0:
+ return self
+ return BilledTokenRates(
+ input_cost_per_token=self.input_cost_per_token * multiplier,
+ output_cost_per_token=self.output_cost_per_token * multiplier,
+ cache_read_input_token_cost=self.cache_read_input_token_cost * multiplier,
+ cache_creation_input_token_cost=self.cache_creation_input_token_cost * multiplier,
+ cache_creation_input_token_cost_above_1hr=self.cache_creation_input_token_cost_above_1hr * multiplier,
+ output_cost_per_reasoning_token=self.output_cost_per_reasoning_token * multiplier,
+ )
+
+
@dataclass(frozen=True, slots=True)
class TokenTypeCostBreakdown:
reasoning_cost: float
cache_read_cost: float
cache_creation_cost: float
+ rates: BilledTokenRates | None = None
+ """Rates these lines were billed at, so a caller reporting both cannot resolve them a second,
+ differently-argued way. None when the model's pricing could not be resolved."""
-def get_token_type_cost_breakdown(
- model: str,
- custom_llm_provider: str | None,
+def _reasoning_token_count(usage: Usage) -> int:
+ parsed: Final = (
+ parse_completion_tokens_details(usage)["reasoning_tokens"] if usage.completion_tokens_details is not None else 0
+ )
+ return parsed or _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
+
+
+def _cache_token_counts(usage: Usage) -> tuple[int, int, CacheCreationTokenDetails | None]:
+ """(cache read tokens, cache creation tokens, cache creation details): read from prompt_tokens_details
+ first, then the private top-level counters the Usage constructor mirrors cache tokens onto for
+ providers/callers that bypass the details."""
+ parsed: Final = parse_prompt_tokens_details(usage) if usage.prompt_tokens_details is not None else None
+ parsed_read: Final = parsed["cache_hit_tokens"] if parsed is not None else 0
+ parsed_creation: Final = parsed["cache_creation_tokens"] if parsed is not None else 0
+ return (
+ parsed_read or _coerce_token_count(getattr(usage, "_cache_read_input_tokens", 0)),
+ parsed_creation or _coerce_token_count(getattr(usage, "_cache_creation_input_tokens", 0)),
+ parsed["cache_creation_token_details"] if parsed is not None else None,
+ )
+
+
+def _custom_pricing_rates(custom_cost_per_token: CostPerToken) -> BilledTokenRates:
+ """Flat custom pricing has no tiers, uplifts or reasoning rate: cache tokens bill at the configured
+ cache rates (else the input rate) and reasoning at the output rate, as _cost_per_token_custom_pricing_helper does."""
+ input_rate: Final = custom_cost_per_token["input_cost_per_token"]
+ output_rate: Final = custom_cost_per_token["output_cost_per_token"]
+ cache_creation_rate: Final = custom_cost_per_token.get("cache_creation_input_token_cost", input_rate)
+ return BilledTokenRates(
+ input_cost_per_token=input_rate,
+ output_cost_per_token=output_rate,
+ cache_read_input_token_cost=custom_cost_per_token.get("cache_read_input_token_cost", input_rate),
+ cache_creation_input_token_cost=cache_creation_rate,
+ cache_creation_input_token_cost_above_1hr=cache_creation_rate,
+ output_cost_per_reasoning_token=output_rate,
+ )
+
+
+def _cost_map_billed_rates(
+ model_info: ModelInfo,
usage: Usage,
- service_tier: str | None = None,
- data_residency: str | None = None,
- vertex_location: str | None = None,
- current_time: datetime | None = None,
-) -> TokenTypeCostBreakdown:
- """
- Provider-agnostic cost of reasoning and cache tokens, derived from the usage
- object and model pricing alone.
-
- This works for every provider, including Perplexity/Cerebras/Dashscope whose
- cost calculators bypass ``generic_cost_per_token``, because cache tokens always
- land on ``prompt_tokens_details`` (via the Usage constructor and provider
- transformations) and reasoning tokens on ``completion_tokens_details``. It reuses
- the same rate-resolution primitives as the total-cost path so the breakdown can
- never drift from the totals. Returns zeros (never raises) when the model or its
- pricing cannot be resolved.
- """
- try:
- model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
- except Exception:
- return TokenTypeCostBreakdown(0.0, 0.0, 0.0)
-
- billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
+ custom_llm_provider: str | None,
+ service_tier: str | None,
+ data_residency: str | None,
+ vertex_location: str | None,
+ current_time: datetime | None,
+) -> BilledTokenRates:
+ billing_time: Final = current_time if current_time is not None else current_billing_time()
(
- _prompt_base_cost,
+ prompt_base_cost,
completion_base_cost,
cache_creation_cost_rate,
cache_creation_cost_above_1hr_rate,
@@ -1356,13 +1406,6 @@ def get_token_type_cost_breakdown(
current_time=billing_time,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
)
-
- reasoning_tokens = (
- parse_completion_tokens_details(usage)["reasoning_tokens"] if usage.completion_tokens_details is not None else 0
- )
- if not reasoning_tokens:
- reasoning_tokens = _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
-
reasoning_rate: Final = _resolve_billed_reasoning_rate(
model_info=model_info,
usage=usage,
@@ -1370,57 +1413,103 @@ def get_token_type_cost_breakdown(
completion_base_cost=completion_base_cost,
current_time=billing_time,
)
- reasoning_cost = float(reasoning_tokens) * reasoning_rate
+ multiplier: Final = (
+ _get_regional_uplift_multiplier(model_info, data_residency)
+ * get_vertex_regional_endpoint_uplift(model_info, vertex_location)
+ * get_provider_specific_geo_multiplier(model_info=model_info, usage=usage)
+ )
+ return BilledTokenRates(
+ input_cost_per_token=prompt_base_cost,
+ output_cost_per_token=completion_base_cost,
+ cache_read_input_token_cost=cache_read_cost_rate,
+ cache_creation_input_token_cost=cache_creation_cost_rate,
+ cache_creation_input_token_cost_above_1hr=cache_creation_cost_above_1hr_rate,
+ output_cost_per_reasoning_token=reasoning_rate,
+ ).scaled(multiplier)
- cache_read_tokens = 0
- cache_creation_tokens = 0
- cache_creation_token_details: CacheCreationTokenDetails | None = None
- if usage.prompt_tokens_details is not None:
- prompt_tokens_details: Final = parse_prompt_tokens_details(usage)
- cache_read_tokens = prompt_tokens_details["cache_hit_tokens"]
- cache_creation_tokens = prompt_tokens_details["cache_creation_tokens"]
- cache_creation_token_details = prompt_tokens_details["cache_creation_token_details"]
- # Fall back to the private top-level counters the Usage constructor mirrors cache
- # tokens onto, so providers/callers that bypass prompt_tokens_details are covered.
- if not cache_read_tokens:
- cache_read_tokens = _coerce_token_count(getattr(usage, "_cache_read_input_tokens", 0))
- if not cache_creation_tokens:
- cache_creation_tokens = _coerce_token_count(getattr(usage, "_cache_creation_input_tokens", 0))
- cache_read_cost = float(cache_read_tokens) * cache_read_cost_rate
- cache_creation_cost = calculate_cache_writing_cost(
- cache_creation_tokens=cache_creation_tokens,
- cache_creation_token_details=cache_creation_token_details,
- cache_creation_cost_above_1hr=cache_creation_cost_above_1hr_rate,
- cache_creation_cost=cache_creation_cost_rate,
+def get_billed_token_rates(
+ model: str,
+ custom_llm_provider: str | None,
+ usage: Usage,
+ service_tier: str | None = None,
+ data_residency: str | None = None,
+ vertex_location: str | None = None,
+ current_time: datetime | None = None,
+ custom_cost_per_token: CostPerToken | None = None,
+) -> BilledTokenRates | None:
+ """Rates the cost calculator bills ``usage`` at, resolved exactly as the totals and the token-type
+ breakdown resolve them. None when the model's pricing cannot be resolved."""
+ if custom_cost_per_token is not None:
+ return _custom_pricing_rates(custom_cost_per_token)
+ try:
+ model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
+ except Exception:
+ return None
+ return _cost_map_billed_rates(
+ model_info=model_info,
+ usage=usage,
+ custom_llm_provider=custom_llm_provider,
+ service_tier=service_tier,
+ data_residency=data_residency,
+ vertex_location=vertex_location,
+ current_time=current_time,
)
- # Apply the same flat regional-processing uplift the totals get, so per-type
- # costs stay reconciled with input_cost/output_cost for regionalized OpenAI hosts.
- uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency)
- if uplift != 1.0:
- reasoning_cost *= uplift
- cache_read_cost *= uplift
- cache_creation_cost *= uplift
- vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
- if vertex_uplift != 1.0:
- reasoning_cost *= vertex_uplift
- cache_read_cost *= vertex_uplift
- cache_creation_cost *= vertex_uplift
+def get_token_type_cost_breakdown(
+ model: str,
+ custom_llm_provider: str | None,
+ usage: Usage,
+ service_tier: str | None = None,
+ data_residency: str | None = None,
+ vertex_location: str | None = None,
+ current_time: datetime | None = None,
+ custom_cost_per_token: CostPerToken | None = None,
+) -> TokenTypeCostBreakdown:
+ """
+ Provider-agnostic cost of reasoning and cache tokens, derived from the usage
+ object and model pricing alone.
- # Mirror the provider-specific geo uplift (e.g. Anthropic us: 1.1) the totals
- # apply, so cache and reasoning line items stay reconciled with them.
- geo_multiplier: Final = get_provider_specific_geo_multiplier(model_info=model_info, usage=usage)
- if geo_multiplier != 1.0:
- reasoning_cost *= geo_multiplier
- cache_read_cost *= geo_multiplier
- cache_creation_cost *= geo_multiplier
+ This works for every provider, including Perplexity/Cerebras/Dashscope whose
+ cost calculators bypass ``generic_cost_per_token``, because cache tokens always
+ land on ``prompt_tokens_details`` (via the Usage constructor and provider
+ transformations) and reasoning tokens on ``completion_tokens_details``. It reuses
+ the same rate resolution as the total-cost path (``get_billed_token_rates``) so the
+ breakdown can never drift from the totals. A deployment billed by
+ ``custom_cost_per_token`` is priced from those flat rates instead of the cost map and,
+ like its totals, bills cache writes flat rather than by their 5m/1h split.
+ Returns zeros (never raises) when the model or its pricing cannot be resolved.
+ """
+ rates: Final = get_billed_token_rates(
+ model=model,
+ custom_llm_provider=custom_llm_provider,
+ usage=usage,
+ service_tier=service_tier,
+ data_residency=data_residency,
+ vertex_location=vertex_location,
+ current_time=current_time,
+ custom_cost_per_token=custom_cost_per_token,
+ )
+ if rates is None:
+ return TokenTypeCostBreakdown(0.0, 0.0, 0.0)
+ cache_read_tokens, cache_creation_tokens, cache_creation_token_details = _cache_token_counts(usage)
+ cache_creation_cost: Final = (
+ float(cache_creation_tokens) * rates.cache_creation_input_token_cost
+ if custom_cost_per_token is not None
+ else calculate_cache_writing_cost(
+ cache_creation_tokens=cache_creation_tokens,
+ cache_creation_token_details=cache_creation_token_details,
+ cache_creation_cost_above_1hr=rates.cache_creation_input_token_cost_above_1hr,
+ cache_creation_cost=rates.cache_creation_input_token_cost,
+ )
+ )
return TokenTypeCostBreakdown(
- reasoning_cost=reasoning_cost,
- cache_read_cost=cache_read_cost,
+ reasoning_cost=float(_reasoning_token_count(usage)) * rates.output_cost_per_reasoning_token,
+ cache_read_cost=float(cache_read_tokens) * rates.cache_read_input_token_cost,
cache_creation_cost=cache_creation_cost,
+ rates=rates,
)
diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py
index 33b27943ad8..9875ac2b9c3 100644
--- a/litellm/llms/bedrock/files/transformation.py
+++ b/litellm/llms/bedrock/files/transformation.py
@@ -7,13 +7,13 @@ from contextlib import suppress
from functools import cache
from itertools import chain
from types import MappingProxyType
-from typing import Any, Final, TypeAlias, TypedDict
+from typing import Any, Final, Literal, TypeAlias, TypedDict
from urllib.parse import unquote
import httpx
from httpx import Headers, Response
from openai.types.file_deleted import FileDeleted
-from pydantic import BaseModel, ConfigDict, TypeAdapter
+from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
@@ -60,11 +60,12 @@ from litellm.utils import get_llm_provider
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError, merge_bedrock_aws_request_params, resolve_s3_encryption_key_id
-# litellm_params key used to hand the SigV4-signed GET headers from
-# `transform_file_content_request` to `validate_environment` (the only hook
-# the shared file-content HTTP handler exposes for setting request headers).
-# Same pattern as the `upload_url` handoff in `transform_create_file_request`.
-S3_SIGNED_GET_HEADERS_PARAM: Final = "_s3_signed_get_headers"
+S3_SIGNED_REQUEST_HEADERS_PARAM: Final = "_s3_signed_request_headers"
+
+
+class _S3DeleteContext(BaseModel):
+ file_id: str = Field(min_length=1)
+
# litellm_params key carrying the size of the body uploaded to S3, handed from
# `transform_create_file_request` to `transform_create_file_response`.
@@ -291,7 +292,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
) -> dict:
result: Final[dict[str, object]] = {}
result.update(headers)
- signed_headers: Final = litellm_params.pop(S3_SIGNED_GET_HEADERS_PARAM, None)
+ signed_headers: Final = litellm_params.pop(S3_SIGNED_REQUEST_HEADERS_PARAM, None)
if isinstance(signed_headers, Mapping):
result.update(signed_headers) # any-ok: untyped handoff headers
# otherwise no extra headers - AWS credentials are handled by BaseAWSLLM
@@ -1187,18 +1188,27 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
def transform_delete_file_request(
self,
file_id: str,
- optional_params: dict,
- litellm_params: dict,
- ) -> tuple[str, dict]:
- raise NotImplementedError("BedrockFilesConfig does not support file deletion")
+ optional_params: Mapping[str, object],
+ litellm_params: MutableMapping[str, object],
+ ) -> tuple[str, dict[str, str]]:
+ return self._transform_s3_file_request(
+ file_id=file_id, method="DELETE", optional_params=optional_params, litellm_params=litellm_params
+ )
def transform_delete_file_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
- litellm_params: dict,
+ litellm_params: Mapping[str, object],
) -> FileDeleted:
- raise NotImplementedError("BedrockFilesConfig does not support file deletion")
+ if raw_response.status_code != 204:
+ raise BedrockError(
+ status_code=raw_response.status_code if raw_response.status_code >= 400 else 502,
+ message=raw_response.text or f"S3 file deletion returned HTTP {raw_response.status_code}",
+ headers=raw_response.headers,
+ )
+ context: Final = _S3DeleteContext.model_validate(logging_obj.model_call_details.get("additional_args"))
+ return FileDeleted(id=context.file_id, deleted=True, object="file")
def transform_list_files_request(
self,
@@ -1233,6 +1243,18 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
if not file_id:
raise ValueError("file_id is required for Bedrock file content retrieval")
+ return self._transform_s3_file_request(
+ file_id=file_id, method="GET", optional_params=optional_params, litellm_params=litellm_params
+ )
+
+ def _transform_s3_file_request(
+ self,
+ *,
+ file_id: str,
+ method: Literal["GET", "DELETE"],
+ optional_params: Mapping[str, object],
+ litellm_params: MutableMapping[str, object],
+ ) -> tuple[str, dict[str, str]]:
s3_uri: Final = extract_s3_uri_from_file_id(file_id)
bucket_name, object_key = _validate_file_id_against_configured_buckets(
s3_uri=s3_uri,
@@ -1240,40 +1262,32 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(litellm_params),
)
- # The shared file-content handler passes optional_params={}, so AWS
- # credentials/region arrive via litellm_params here (unlike the upload
- # path). s3_region_name wins over aws_region_name, same priority as
- # get_complete_file_url above.
- merged_params: Final[dict[str, object]] = {}
- merged_params.update(litellm_params)
- merged_params.update(optional_params)
- request_params: Final = _BedrockS3RequestParams.model_validate(merged_params)
+ request_params: Final = _BedrockS3RequestParams.model_validate({**litellm_params, **optional_params})
region_preference: Final = request_params.s3_region_name or request_params.aws_region_name
region_params: Final[dict[str, str | None]] = {"aws_region_name": region_preference}
aws_region_name: Final = self._get_aws_region_name(optional_params=region_params, model="")
- s3_endpoint_url = (
+ s3_endpoint_url: Final = (
request_params.s3_endpoint_url or f"https://s3.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}"
).rstrip("/")
url: Final = f"{s3_endpoint_url}/{bucket_name}/{encode_s3_object_key_for_url(object_key)}"
- litellm_params[S3_SIGNED_GET_HEADERS_PARAM] = self._sign_s3_get_request(
+ litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM] = self._sign_s3_request_without_body(
api_base=url,
aws_region_name=aws_region_name,
request_params=request_params,
+ method=method,
)
return url, {}
- def _sign_s3_get_request(
+ def _sign_s3_request_without_body(
self,
api_base: str,
aws_region_name: str,
request_params: _BedrockS3RequestParams,
+ method: Literal["GET", "DELETE"] = "GET",
) -> dict[str, str]:
- """
- SigV4-sign an S3 GetObject request, mirroring `_sign_s3_request` (PUT).
- """
try:
import hashlib
@@ -1297,7 +1311,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
empty_body_hash: Final = hashlib.sha256(b"").hexdigest()
aws_request: Final = AWSRequest( # any-ok: botocore AWSRequest is untyped
- method="GET",
+ method=method,
url=api_base,
headers={"x-amz-content-sha256": empty_body_hash},
)
diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py
index 329dddbdf05..5b74caee3a2 100644
--- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py
+++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py
@@ -3,12 +3,15 @@ import importlib
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
+from traceback import walk_tb
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
+from uuid import uuid4
import anyio
import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
+from pydantic import ValidationError
from starlette.datastructures import Headers
from litellm._logging import verbose_logger
@@ -30,6 +33,8 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
list_fault_http_status,
outcome_wire_value,
)
+from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree
+from litellm.proxy._experimental.mcp_server.oauth_utils import _redact_mcp_resource_url
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
acting_user_auth,
build_effective_auth_contexts,
@@ -78,11 +83,39 @@ _MCP_GUARDRAIL_REJECTIONS: Final = (
def _connection_error_message(exc: BaseException, url: str | None, timeout_seconds: float) -> str:
+ reference: Final = uuid4().hex
+ verbose_logger.error(
+ "MCP connection test failed (reference=%s): %s",
+ reference,
+ tuple(
+ (
+ type(cause).__name__,
+ tuple(
+ (frame.f_code.co_filename, lineno, frame.f_code.co_name)
+ for frame, lineno in walk_tb(cause.__traceback__)
+ ),
+ )
+ for cause in iter_exception_tree(exc)
+ ),
+ )
+ return next(
+ (
+ message
+ for cause in iter_exception_tree(exc)
+ if (message := _known_connection_error_message(cause, url, timeout_seconds)) is not None
+ ),
+ "An unexpected error occurred while testing the MCP connection. "
+ f"Retry; if it persists, share reference {reference} with your gateway administrator.",
+ )
+
+
+def _known_connection_error_message(exc: BaseException, url: str | None, timeout_seconds: float) -> str | None:
if isinstance(exc, MCPServerURLCredentialsError):
return str(exc.detail)
if isinstance(exc, TimeoutError):
return (
- f"Failed to connect to MCP server: no response from {url or 'the server'} "
+ "Failed to connect to MCP server: no valid MCP response received from "
+ f"{_redact_mcp_resource_url(url) or 'the server'} "
f"within {timeout_seconds:.0f}s. Check that the LiteLLM proxy can reach this URL "
"from its network (DNS, egress rules, firewalls) and that the server answers MCP requests."
)
@@ -99,13 +132,45 @@ def _connection_error_message(exc: BaseException, url: str | None, timeout_secon
return "Failed to connect to MCP server: the connection timed out."
if isinstance(exc, httpx.HTTPStatusError):
return f"Failed to connect to MCP server: it returned HTTP {exc.response.status_code}."
- return "Failed to connect to MCP server. Check proxy logs for details."
+ if isinstance(exc, (httpx.NetworkError, httpx.RemoteProtocolError, ConnectionError)):
+ return (
+ "Failed to connect to MCP server: the connection was interrupted. "
+ "Check the server and network connection, then retry."
+ )
+ if isinstance(exc, ValueError) and str(exc).startswith("Unexpected content type:"):
+ return (
+ "Failed to connect to MCP server: the endpoint returned an unsupported content type. "
+ "Check that the URL is an MCP endpoint, not a web page, and matches the selected transport."
+ )
+ if isinstance(exc, ValidationError) and exc.title in ("JSONRPCMessage", "InitializeResult", "ListToolsResult"):
+ return (
+ "Failed to connect to MCP server: the endpoint returned invalid JSON or an invalid MCP response. "
+ "Check the MCP endpoint URL and the server's protocol implementation."
+ )
+ if MCP_AVAILABLE and isinstance(exc, McpError):
+ if exc.error.code == -32000 and exc.error.message == "Connection closed":
+ return (
+ "Failed to connect to MCP server: the connection was closed before the request completed. "
+ "Check that the server stays running and returns a complete MCP response, then retry."
+ )
+ if exc.error.code == 32600 and exc.error.message == "Session terminated":
+ return (
+ "Failed to connect to MCP server: the MCP session was terminated. "
+ "Check that the URL points to an MCP endpoint and matches the selected transport, "
+ "then retry to start a new session."
+ )
+ return (
+ f"Failed to connect to MCP server: the MCP request failed (JSON-RPC code {exc.error.code}). "
+ "Check that the endpoint supports MCP initialization and tool listing, and check the upstream server logs."
+ )
+ return None
if MCP_AVAILABLE:
+ from mcp.shared.exceptions import McpError
from mcp.types import Tool as MCPTool
- from litellm.experimental_mcp_client.client import MCPClient
+ from litellm.experimental_mcp_client.client import MCPClient, as_mcp_read_timeout
from litellm.llms.litellm_proxy.skills.skill_search import (
DEFAULT_SKILL_SEARCH_TOP_K,
)
@@ -1342,11 +1407,18 @@ if MCP_AVAILABLE:
except (KeyboardInterrupt, SystemExit, asyncio.CancelledError):
raise
except BaseException as e:
- verbose_logger.error("Error in MCP operation: %s", e, exc_info=True)
+ effective_timeout: Final = (
+ min(request.timeout if request.timeout is not None else MCP_CLIENT_TIMEOUT, timeout_seconds)
+ if any(
+ isinstance(cause, McpError) and as_mcp_read_timeout(cause) is not None
+ for cause in iter_exception_tree(e)
+ )
+ else timeout_seconds
+ )
return {
"status": "error",
"error": True,
- "message": _connection_error_message(e, request.url, timeout_seconds),
+ "message": _connection_error_message(e, request.url, effective_timeout),
}
async def _preview_openapi_tools(spec_path: str) -> dict:
diff --git a/litellm/proxy/_experimental/mcp_server/toolset_db.py b/litellm/proxy/_experimental/mcp_server/toolset_db.py
index ecaaf35e817..48bad178927 100644
--- a/litellm/proxy/_experimental/mcp_server/toolset_db.py
+++ b/litellm/proxy/_experimental/mcp_server/toolset_db.py
@@ -132,10 +132,18 @@ async def update_mcp_toolset(
data: UpdateMCPToolsetRequest,
touched_by: str,
) -> MCPToolset | None:
- data_dict: Final = data.model_dump(exclude_none=True, exclude={"toolset_id"})
- if "tools" in data_dict:
- data_dict["tools"] = json.dumps(data_dict["tools"])
- data_dict["updated_by"] = touched_by
+ """A partial update: absent keeps, null clears. A toolset always has a name and a
+ tool list, so a null ``toolset_name`` or ``tools`` is a no-op rather than a clear;
+ emptying the tool selection is an explicit ``[]``, which cannot be mistaken for a
+ caller that left the field out."""
+ data_dict: Final = dict( # mutable-ok: Prisma requires a plain dict for JSON query serialization
+ (
+ (field, json.dumps(value) if field == "tools" else value)
+ for field, value in data.model_dump(exclude_unset=True).items()
+ if field != "toolset_id" and (field not in ("toolset_name", "tools") or value is not None)
+ ),
+ updated_by=touched_by,
+ )
try:
row: Final = await _toolset_table(prisma_client).update(
where={"toolset_id": data.toolset_id},
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index 19fd19e5c8a..e1f14c3fac9 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -3,6 +3,7 @@ import json
import os
from collections.abc import Callable, Mapping
from datetime import datetime
+from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple
import httpx
@@ -1297,6 +1298,13 @@ class UpdateKeyRequest(KeyRequestBase):
rotation_interval: str | None = None
organization_id: str | None = None
+ @model_validator(mode="before")
+ @classmethod
+ def drop_blank_team_id(cls, values: object) -> object:
+ if isinstance(values, Mapping) and values.get("team_id") == "":
+ return MappingProxyType({k: v for k, v in values.items() if k != "team_id"})
+ return values
+
@field_validator("organization_id", mode="before")
@classmethod
def treat_cleared_organization_id_as_unset(cls, v: object) -> object:
@@ -5151,9 +5159,26 @@ class CostEstimateRequest(LiteLLMPydanticObjectBase):
model: str = Field(description="Model name (from /model_group/info)")
input_tokens: int = Field(description="Expected input tokens per request", ge=0)
output_tokens: int = Field(description="Expected output tokens per request", ge=0)
+ cache_read_input_tokens: int = Field(
+ default=0, description="Input tokens read from the prompt cache; counted within input_tokens", ge=0
+ )
+ cache_creation_input_tokens: int = Field(
+ default=0, description="Input tokens written to the prompt cache; counted within input_tokens", ge=0
+ )
+ reasoning_tokens: int = Field(
+ default=0, description="Reasoning tokens the model emits; counted within output_tokens", ge=0
+ )
num_requests_per_day: int | None = Field(default=None, description="Number of requests per day", ge=0)
num_requests_per_month: int | None = Field(default=None, description="Number of requests per month", ge=0)
+ @model_validator(mode="after")
+ def validate_token_subsets(self) -> "CostEstimateRequest":
+ if self.cache_read_input_tokens + self.cache_creation_input_tokens > self.input_tokens:
+ raise ValueError("cache_read_input_tokens plus cache_creation_input_tokens cannot exceed input_tokens")
+ if self.reasoning_tokens > self.output_tokens:
+ raise ValueError("reasoning_tokens cannot exceed output_tokens")
+ return self
+
class CostEstimateResponse(LiteLLMPydanticObjectBase):
"""Response body for /cost/estimate endpoint."""
@@ -5161,6 +5186,9 @@ class CostEstimateResponse(LiteLLMPydanticObjectBase):
model: str
input_tokens: int
output_tokens: int
+ cache_read_input_tokens: int = 0
+ cache_creation_input_tokens: int = 0
+ reasoning_tokens: int = 0
num_requests_per_day: int | None = None
num_requests_per_month: int | None = None
# Per-request costs
@@ -5168,17 +5196,33 @@ class CostEstimateResponse(LiteLLMPydanticObjectBase):
input_cost_per_request: float = Field(description="Input token cost per request (before margin)")
output_cost_per_request: float = Field(description="Output token cost per request (before margin)")
margin_cost_per_request: float = Field(default=0.0, description="Margin/fee added per request")
+ cache_read_cost_per_request: float = Field(default=0.0, description="Cache-read share of input_cost_per_request")
+ cache_creation_cost_per_request: float = Field(
+ default=0.0, description="Cache-write share of input_cost_per_request"
+ )
+ reasoning_cost_per_request: float = Field(default=0.0, description="Reasoning share of output_cost_per_request")
# Daily costs (if num_requests_per_day provided)
daily_cost: float | None = Field(default=None, description="Total daily cost (includes margin)")
daily_input_cost: float | None = Field(default=None, description="Daily input token cost")
daily_output_cost: float | None = Field(default=None, description="Daily output token cost")
daily_margin_cost: float | None = Field(default=None, description="Daily margin/fee")
+ daily_cache_read_cost: float | None = Field(default=None, description="Cache-read share of daily_input_cost")
+ daily_cache_creation_cost: float | None = Field(default=None, description="Cache-write share of daily_input_cost")
+ daily_reasoning_cost: float | None = Field(default=None, description="Reasoning share of daily_output_cost")
# Monthly costs (if num_requests_per_month provided)
monthly_cost: float | None = Field(default=None, description="Total monthly cost (includes margin)")
monthly_input_cost: float | None = Field(default=None, description="Monthly input token cost")
monthly_output_cost: float | None = Field(default=None, description="Monthly output token cost")
monthly_margin_cost: float | None = Field(default=None, description="Monthly margin/fee")
- # Pricing info
- input_cost_per_token: float | None = None
- output_cost_per_token: float | None = None
+ monthly_cache_read_cost: float | None = Field(default=None, description="Cache-read share of monthly_input_cost")
+ monthly_cache_creation_cost: float | None = Field(
+ default=None, description="Cache-write share of monthly_input_cost"
+ )
+ monthly_reasoning_cost: float | None = Field(default=None, description="Reasoning share of monthly_output_cost")
+ # Pricing info: the rates this request's usage bills at, after token tiers and regional multipliers
+ input_cost_per_token: float | None = Field(default=None, description="Rate billed per input token")
+ output_cost_per_token: float | None = Field(default=None, description="Rate billed per output token")
+ cache_read_input_token_cost: float | None = Field(default=None, description="Rate billed per cache-read token")
+ cache_creation_input_token_cost: float | None = Field(default=None, description="Rate billed per cache-write token")
+ output_cost_per_reasoning_token: float | None = Field(default=None, description="Rate billed per reasoning token")
provider: str | None = None
diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py
index dc693317de0..a71a1993064 100644
--- a/litellm/proxy/auth/auth_checks.py
+++ b/litellm/proxy/auth/auth_checks.py
@@ -475,6 +475,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None
_NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({})
+_TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True})
def _has_ptu_flat_cost(model: str, llm_router: "Router") -> bool:
@@ -2858,7 +2859,9 @@ class TeamNotFoundError(HTTPException):
async def _get_team_db_check(
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
) -> "_PrismaTeamRow | None":
- response = await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id})
+ response = await _team_table(TeamRepository(prisma_client)).find_unique(
+ where={"team_id": team_id}, include=_TEAM_GRANT_RELATIONS
+ )
if response is None and team_id_upsert:
from litellm.proxy.management_endpoints.team_endpoints import new_team
@@ -3158,7 +3161,9 @@ async def get_team_object_by_alias(
# Query database by team_alias
try:
- teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(where={"team_alias": team_alias})
+ teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(
+ where={"team_alias": team_alias}, include=_TEAM_GRANT_RELATIONS
+ )
if not teams:
raise HTTPException(
diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py
index 0795cee7409..69091ee8344 100644
--- a/litellm/proxy/auth/handle_jwt.py
+++ b/litellm/proxy/auth/handle_jwt.py
@@ -53,6 +53,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.auth.route_checks import RouteChecks
+from litellm.proxy.auth.team_grants import team_model_aliases
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
get_management_object_ttl,
@@ -1595,7 +1596,7 @@ class JWTAuthManager:
model=requested_model,
team_object=team_object,
llm_router=llm_router,
- team_model_aliases=None,
+ team_model_aliases=team_model_aliases(team_object),
)
):
is_allowed = allowed_routes_check(
@@ -2132,7 +2133,7 @@ class JWTAuthManager:
model=requested_model,
team_object=team_object,
llm_router=llm_router,
- team_model_aliases=None,
+ team_model_aliases=team_model_aliases(team_object),
)
except ProxyException:
continue
diff --git a/litellm/proxy/auth/team_grants.py b/litellm/proxy/auth/team_grants.py
new file mode 100644
index 00000000000..1196011dcdd
--- /dev/null
+++ b/litellm/proxy/auth/team_grants.py
@@ -0,0 +1,122 @@
+"""Project a team row (plus the caller's membership in it) onto the ``team_*`` fields of ``UserAPIKeyAuth``.
+
+The virtual-key path gets these fields for free from the combined-view SQL join. Every other auth path
+starts from a ``LiteLLM_TeamTable`` object instead and has to copy them over by hand, which is how JWT
+callers kept losing grants (aliases, permissions, limits) one field at a time. Build the badge through
+``team_grants`` and the two paths cannot drift.
+"""
+
+from collections.abc import Mapping, Sequence
+from types import MappingProxyType
+from typing import Annotated, Final
+
+from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError
+from pydantic.main import IncEx
+from typing_extensions import ReadOnly, TypedDict
+
+from litellm.proxy._types import (
+ LiteLLM_ObjectPermissionTable,
+ LiteLLM_TeamMembership,
+ LiteLLM_TeamTable,
+ Member,
+)
+
+_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
+_JSON_COLUMNS: Final[Mapping[str, IncEx | bool]] = MappingProxyType(
+ {"metadata": True, "litellm_model_table": MappingProxyType({"model_aliases": True})}
+)
+
+
+def _decode_model_aliases(value: object) -> object:
+ """``LiteLLM_ModelTable.model_aliases`` is typed ``str | dict``; writers hand Prisma ``json.dumps(...)``, so take both."""
+ if not isinstance(value, str):
+ return value
+ try:
+ return _MODEL_ALIASES_ADAPTER.validate_json(value)
+ except ValidationError:
+ return None
+
+
+class TeamModelAliasTable(BaseModel):
+ model_config = ConfigDict(protected_namespaces=())
+
+ model_aliases: Annotated[Mapping[str, str] | None, BeforeValidator(_decode_model_aliases)] = None
+
+
+class _TeamJsonColumns(BaseModel):
+ """The two loosely typed columns on ``LiteLLM_TeamTable``, re-read with the shape the badge needs."""
+
+ metadata: Mapping[str, object] | None = None
+ litellm_model_table: TeamModelAliasTable | None = None
+
+
+class TeamGrants(TypedDict, total=False):
+ """Keyword arguments for ``UserAPIKeyAuth``. Empty when the caller has no team, so the model's own defaults apply."""
+
+ team_alias: ReadOnly[str | None]
+ team_tpm_limit: ReadOnly[int | None]
+ team_rpm_limit: ReadOnly[int | None]
+ team_max_budget: ReadOnly[float | None]
+ team_soft_budget: ReadOnly[float | None]
+ team_spend: ReadOnly[float | None]
+ team_models: ReadOnly[Sequence[str]]
+ team_blocked: ReadOnly[bool]
+ team_metadata: ReadOnly[Mapping[str, object] | None]
+ team_model_aliases: ReadOnly[Mapping[str, str] | None]
+ team_object_permission_id: ReadOnly[str | None]
+ team_object_permission: ReadOnly[LiteLLM_ObjectPermissionTable | None]
+ team_member: ReadOnly[Member | None]
+ team_member_spend: ReadOnly[float | None]
+ team_member_tpm_limit: ReadOnly[int | None]
+ team_member_rpm_limit: ReadOnly[int | None]
+
+
+def _json_columns(team_object: LiteLLM_TeamTable) -> _TeamJsonColumns:
+ try:
+ return _TeamJsonColumns.model_validate(team_object.model_dump(include=_JSON_COLUMNS))
+ except ValidationError:
+ return _TeamJsonColumns()
+
+
+def team_model_aliases(team_object: LiteLLM_TeamTable | None) -> Mapping[str, str] | None:
+ if team_object is None:
+ return None
+ alias_table: Final = _json_columns(team_object).litellm_model_table
+ return alias_table.model_aliases if alias_table is not None else None
+
+
+def team_grants(
+ team_object: LiteLLM_TeamTable | None,
+ team_membership: LiteLLM_TeamMembership | None,
+ user_id: str | None,
+) -> TeamGrants:
+ if team_object is None:
+ return TeamGrants()
+ json_columns: Final = _json_columns(team_object)
+ return TeamGrants(
+ team_alias=team_object.team_alias,
+ team_tpm_limit=team_object.tpm_limit,
+ team_rpm_limit=team_object.rpm_limit,
+ team_max_budget=team_object.max_budget,
+ team_soft_budget=team_object.soft_budget,
+ team_spend=team_object.spend,
+ team_models=tuple(team_object.models),
+ team_blocked=team_object.blocked,
+ team_metadata=json_columns.metadata,
+ team_model_aliases=(
+ json_columns.litellm_model_table.model_aliases if json_columns.litellm_model_table is not None else None
+ ),
+ team_object_permission_id=team_object.object_permission_id,
+ team_object_permission=team_object.object_permission,
+ team_member=next(
+ (m for m in team_object.members_with_roles if user_id is not None and m.user_id == user_id),
+ None,
+ ),
+ team_member_spend=team_membership.spend if team_membership is not None else None,
+ team_member_tpm_limit=(
+ team_membership.safe_get_team_member_tpm_limit() if team_membership is not None else None
+ ),
+ team_member_rpm_limit=(
+ team_membership.safe_get_team_member_rpm_limit() if team_membership is not None else None
+ ),
+ )
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index 1ab1811fef8..2a1d0709212 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -82,6 +82,7 @@ from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
from litellm.proxy.auth.resolvers import CredentialRef, Principal
from litellm.proxy.auth.resolvers.store import IdentityStore
from litellm.proxy.auth.route_checks import RouteChecks
+from litellm.proxy.auth.team_grants import team_grants
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
from litellm.proxy.common_utils.http_parsing_utils import (
@@ -1480,24 +1481,16 @@ async def _user_api_key_auth_builder(
user_id=user_id,
user_email=user_email,
team_id=team_id,
- team_alias=(team_object.team_alias if team_object is not None else None),
- team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
- team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
- team_models=(team_object.models if team_object is not None else []),
- team_metadata=(team_object.metadata if team_object is not None else None),
org_id=org_id,
end_user_id=end_user_id,
parent_otel_span=parent_otel_span,
jwt_claims=jwt_claims,
+ **team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
valid_token = UserAPIKeyAuth(
api_key=None,
team_id=team_id,
- team_alias=(team_object.team_alias if team_object is not None else None),
- team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
- team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
- team_models=(team_object.models if team_object is not None else []),
user_role=(
LitellmUserRoles(user_object.user_role)
if user_object is not None and user_object.user_role is not None
@@ -1511,17 +1504,8 @@ async def _user_api_key_auth_builder(
user_tpm_limit=(user_object.tpm_limit if user_object is not None else None),
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
- team_member_rpm_limit=(
- team_membership.safe_get_team_member_rpm_limit() if team_membership is not None else None
- ),
- team_member_tpm_limit=(
- team_membership.safe_get_team_member_tpm_limit() if team_membership is not None else None
- ),
- team_metadata=(team_object.metadata if team_object is not None else None),
jwt_claims=jwt_claims,
- )
- valid_token.team_object_permission = (
- team_object.object_permission if team_object is not None else None
+ **team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py
index 770963a1f24..561a53409f4 100644
--- a/litellm/proxy/common_utils/callback_utils.py
+++ b/litellm/proxy/common_utils/callback_utils.py
@@ -1,10 +1,12 @@
import copy
+import json
import os
from collections.abc import Callable, Iterable, Mapping
from dataclasses import dataclass
+from itertools import accumulate
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias
-from typing_extensions import assert_never
+from typing_extensions import ReadOnly, TypedDict, assert_never
import litellm
from litellm import get_secret
@@ -12,6 +14,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import (
CLIENT_OUTPUT_CEILING_METADATA_KEY,
CONSUMED_REQUEST_TAGS_METADATA_KEY,
+ MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH,
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
ROUTING_REQUEST_TAGS_METADATA_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
@@ -28,6 +31,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
from litellm.proxy.types_utils.utils import get_instance_fn
+from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
StandardLoggingGuardrailInformation,
StandardLoggingPayload,
@@ -52,6 +56,15 @@ reset_color_code: Final = "\033[0m"
TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY: Final = "_pillar_response_headers_trusted"
GUARDRAIL_SCAN_IDS_METADATA_KEY: Final = "guardrail_scan_ids"
+GUARDRAIL_SCAN_METADATA_METADATA_KEY: Final = "guardrail_scan_metadata"
+
+
+class GuardrailScanMetadata(TypedDict):
+ guardrail: ReadOnly[str | None]
+ stage: ReadOnly[str]
+ provider: ReadOnly[str]
+ scan_id: ReadOnly[str]
+
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
@@ -450,6 +463,16 @@ def get_remaining_tokens_and_requests_from_request_data(data: dict) -> dict[str,
return headers
+def _serialize_scan_metadata_header(entries: Iterable[object], *, max_length: int) -> str | None:
+ """Compact JSON list of scan metadata entries, dropping trailing entries so the header fits in max_length."""
+ encoded: Final = tuple(json.dumps(entry, separators=(",", ":")) for entry in entries)
+ lengths: Final = tuple(accumulate(len(item) + 1 for item in encoded))
+ kept: Final = sum(1 for length in lengths if length + 1 <= max_length)
+ if kept == 0:
+ return None
+ return f"[{','.join(encoded[:kept])}]"
+
+
def get_logging_caching_headers(request_data: dict) -> dict | None:
_metadata: Final[dict] = {}
metadata_bucket: Final = request_data.get("metadata")
@@ -468,6 +491,15 @@ def get_logging_caching_headers(request_data: dict) -> dict | None:
if scan_ids:
headers["x-litellm-guardrail-scan-id"] = ",".join(scan_ids)
+ scan_metadata: Final = _metadata.get(GUARDRAIL_SCAN_METADATA_METADATA_KEY)
+ scan_metadata_header: Final = (
+ _serialize_scan_metadata_header(scan_metadata, max_length=MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH)
+ if isinstance(scan_metadata, (list, tuple))
+ else None
+ )
+ if scan_metadata_header:
+ headers["x-litellm-guardrail-scan-metadata"] = scan_metadata_header
+
if "applied_policies" in _metadata:
headers["x-litellm-applied-policies"] = ",".join(_metadata["applied_policies"])
@@ -501,6 +533,7 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset(
"applied_policies",
"applied_guardrails",
GUARDRAIL_SCAN_IDS_METADATA_KEY,
+ GUARDRAIL_SCAN_METADATA_METADATA_KEY,
"policy_sources",
"guardrails",
"guardrail_config",
@@ -565,21 +598,40 @@ def add_guardrail_to_applied_guardrails_header(request_data: dict, guardrail_nam
_metadata["applied_guardrails"] = [guardrail_name]
-def add_guardrail_scan_id(request_data: dict, scan_id: str | None) -> None:
+def add_guardrail_scan_id(
+ request_data: dict[str, object],
+ scan_id: str | None,
+ *,
+ guardrail_name: str | None,
+ provider: str,
+ stage: GuardrailEventHooks,
+) -> None:
"""
- Record a provider scan id so it can be surfaced to the caller.
+ Record a provider scan id, keyed to the guardrail execution that produced it, so it can be surfaced to the caller.
Guardrails only return scan details to the client when they block, so allowed requests carry no
- audit trail. Ids recorded here become the x-litellm-guardrail-scan-id response header.
+ audit trail. Ids recorded here become the x-litellm-guardrail-scan-id response header, and the
+ (guardrail, stage, provider, scan_id) entries become the x-litellm-guardrail-scan-metadata header.
"""
if not scan_id:
return
_, _metadata = get_or_create_metadata_bucket(request_data)
existing: Final = _metadata.get(GUARDRAIL_SCAN_IDS_METADATA_KEY)
- scan_ids: Final = tuple(existing) if isinstance(existing, (list, tuple)) else ()
+ scan_ids: Final[tuple[object, ...]] = tuple(existing) if isinstance(existing, (list, tuple)) else ()
if scan_id not in scan_ids:
_metadata[GUARDRAIL_SCAN_IDS_METADATA_KEY] = (*scan_ids, scan_id)
+ entry: Final[GuardrailScanMetadata] = {
+ "guardrail": guardrail_name,
+ "stage": stage.value,
+ "provider": provider,
+ "scan_id": scan_id,
+ }
+ existing_entries: Final = _metadata.get(GUARDRAIL_SCAN_METADATA_METADATA_KEY)
+ entries: Final[tuple[object, ...]] = tuple(existing_entries) if isinstance(existing_entries, (list, tuple)) else ()
+ if entry not in entries:
+ _metadata[GUARDRAIL_SCAN_METADATA_METADATA_KEY] = (*entries, entry)
+
def add_policy_to_applied_policies_header(request_data: dict, policy_name: str | None):
"""
diff --git a/litellm/proxy/db/db_url_settings.py b/litellm/proxy/db/db_url_settings.py
index 4a0231ad9df..d1e4b3e92b8 100644
--- a/litellm/proxy/db/db_url_settings.py
+++ b/litellm/proxy/db/db_url_settings.py
@@ -34,12 +34,20 @@ writer's connection params (pool size, timeouts, pgbouncer mode) for the
ones the reader URL does not pin itself.
"""
+import _ssl
+import hashlib
import os
+import socket
+import ssl
+import struct
+import sys
+import tempfile
import urllib.parse
-from collections.abc import Mapping
+from collections.abc import Callable, Mapping, Sequence
from functools import partial
+from pathlib import Path
from types import MappingProxyType
-from typing import Annotated, Final, cast
+from typing import Annotated, Final, Protocol, TypeAlias, cast
from pydantic import AliasChoices, BeforeValidator, Field
from pydantic_settings import BaseSettings, SettingsConfigDict
@@ -126,21 +134,100 @@ def add_missing_query_params(url: str, params: Mapping[str, str | int | float])
LIBPQ_VERIFY_SSLMODES: Final[frozenset[str]] = frozenset({"verify-ca", "verify-full"})
+PEM_CERT_HEADER: Final = b"-----BEGIN CERTIFICATE-----"
+PG_SSL_REQUEST: Final = struct.pack("!ii", 8, 80877103)
+TLS_PROBE_TIMEOUT_SECONDS: Final = 10.0
+
+RootCertResolver: TypeAlias = Callable[[str, str, int], str] # mutable-ok: Callable parameter syntax
-def translate_libpq_ssl_params(url: str) -> str:
+class _VerifiedChainSource(Protocol):
+ def get_verified_chain(self) -> Sequence[_ssl.Certificate] | None: ...
+
+
+def _verified_chain_der(tls: ssl.SSLSocket) -> tuple[bytes, ...]:
+ if sys.version_info >= (3, 13):
+ return tuple(tls.get_verified_chain())
+ legacy: Final = cast( # cast-ok: the stub omits _sslobj, the C object has get_verified_chain since 3.10
+ "_VerifiedChainSource | None",
+ tls._sslobj, # pyright: ignore[reportAttributeAccessIssue, reportUnknownMemberType] # public API only from 3.13
+ )
+ chain: Final = () if legacy is None else legacy.get_verified_chain() or ()
+ return tuple(cert.public_bytes(_ssl.ENCODING_DER) for cert in chain)
+
+
+def _server_trust_anchor(cafile: str, host: str, port: int) -> bytes | None:
+ try:
+ context: Final = ssl.create_default_context(cafile=cafile)
+ with socket.create_connection((host, port), timeout=TLS_PROBE_TIMEOUT_SECONDS) as raw:
+ raw.sendall(PG_SSL_REQUEST)
+ if raw.recv(1) != b"S":
+ return None
+ with context.wrap_socket(raw, server_hostname=host) as tls:
+ chain: Final = _verified_chain_der(tls)
+ except (OSError, ValueError):
+ return None
+ return chain[-1] if chain else None
+
+
+def pin_bundle_root(cert_path: str, host: str, port: int) -> str:
+ """Reduce a multi-root CA bundle to the one root that verifies ``host``.
+
+ Prisma's ``sslcert`` loads a single PEM certificate (native-tls
+ ``Certificate::from_pem``), so pointing it at a bundle such as the AWS RDS
+ global bundle trusts only the first of its 108 regional roots and the
+ handshake fails with "unable to get local issuer certificate" for every
+ other region. A single-certificate file is returned as is. For a bundle,
+ one verifying handshake (chain and hostname, whole bundle as trust store)
+ identifies the trust anchor the server actually chains to, which is
+ written to a single-certificate file for Prisma. If the probe fails the
+ bundle path is returned unchanged, so Prisma fails closed exactly as
+ before rather than trusting anything the bundle would not.
+ """
+ try:
+ if Path(cert_path).read_bytes().count(PEM_CERT_HEADER) < 2:
+ return cert_path
+ except OSError:
+ return cert_path
+ root: Final = _server_trust_anchor(cert_path, host, port)
+ if root is None:
+ return cert_path
+ pinned: Final = Path(tempfile.gettempdir()) / f"litellm-sslcert-{hashlib.sha256(root).hexdigest()[:16]}.pem"
+ return str(pinned) if _replace_file(pinned, ssl.DER_cert_to_PEM_cert(root)) else cert_path
+
+
+def _replace_file(target: Path, content: str) -> bool:
+ """Write ``content`` to a private temp file and rename it over ``target``, so
+ readers never see a partial file and a symlink planted at ``target`` is
+ replaced rather than followed."""
+ try:
+ fd, staged = tempfile.mkstemp(dir=target.parent, prefix=f"{target.name}.")
+ except OSError:
+ return False
+ try:
+ with os.fdopen(fd, "w") as handle:
+ handle.write(content)
+ os.replace(staged, target)
+ except OSError:
+ Path(staged).unlink(missing_ok=True)
+ return False
+ return True
+
+
+def translate_libpq_ssl_params(url: str, resolve_root_cert: RootCertResolver = pin_bundle_root) -> str:
"""Rewrite libpq's certificate-verification params into Prisma's dialect.
Prisma's engine only knows ``sslmode=disable|prefer|require``, ``sslcert``
- (the CA bundle) and ``sslaccept=strict``. It silently discards
+ (a single CA certificate) and ``sslaccept=strict``. It silently discards
``sslrootcert`` and downgrades ``sslmode=verify-ca`` / ``verify-full`` to
``prefer``, so a URL copied from libpq / RDS docs connects over TLS with no
certificate check at all. ``verify-ca`` and ``verify-full`` both become
``require`` (Prisma has no CA-only mode), ``sslrootcert`` becomes
- ``sslcert``, and either one turns on ``sslaccept=strict`` (chain and
- hostname), matching libpq where a root cert makes ``require`` verify.
- Prisma params the operator pinned themselves win; anything else is left
- untouched.
+ ``sslcert`` (run through ``resolve_root_cert``, which pins a multi-root
+ bundle down to the server's root), and either one turns on
+ ``sslaccept=strict`` (chain and hostname), matching libpq where a root
+ cert makes ``require`` verify. Prisma params the operator pinned
+ themselves win; anything else is left untouched.
"""
parsed: Final = urllib.parse.urlsplit(url)
pairs: Final = tuple(urllib.parse.parse_qsl(parsed.query, keep_blank_values=True))
@@ -154,7 +241,9 @@ def translate_libpq_ssl_params(url: str) -> str:
if key != "sslrootcert"
)
root_cert: Final = tuple(
- ("sslcert", value) for key, value in pairs if key == "sslrootcert" and "sslcert" not in keys
+ ("sslcert", resolve_root_cert(value, parsed.hostname or "", parsed.port or int(DEFAULT_POSTGRES_PORT)))
+ for key, value in pairs
+ if key == "sslrootcert" and "sslcert" not in keys
)
strict: Final = () if "sslaccept" in keys else (("sslaccept", "strict"),)
query: Final = urllib.parse.urlencode(translated + root_cert + strict)
diff --git a/litellm/proxy/db/gateway_request_tracking.py b/litellm/proxy/db/gateway_request_tracking.py
index bebd74e877c..c9ace68db33 100644
--- a/litellm/proxy/db/gateway_request_tracking.py
+++ b/litellm/proxy/db/gateway_request_tracking.py
@@ -10,13 +10,28 @@ strings rather than passing the raw path through. Nothing a caller sends can
add a key, so the fold and the table it commits to are bounded by (days x
routes) however much traffic arrives, and the response path carries no
unbounded queue that would block once full.
+
+A flush commits its whole snapshot as one multi-row ``INSERT ... ON CONFLICT DO
+UPDATE`` rather than one upsert per key, so a worker costs the primary one
+statement per interval however many routes it served. With
+``use_redis_transaction_buffer`` on, workers instead push their snapshot to a
+Redis list and one lock-holding pod folds every entry and writes the table, so
+the deployment as a whole costs the primary one statement per interval.
"""
-from dataclasses import asdict
+import json
+from collections.abc import AsyncIterator, Iterable
from datetime import datetime, timezone
-from typing import TYPE_CHECKING, Final
+from itertools import chain
+from types import MappingProxyType
+from typing import TYPE_CHECKING, Final, TypeAlias
+
+from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
+from litellm.caching import RedisCache
+from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_GATEWAY_REQUESTS_BUFFER_KEY
+from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.middleware.billable_request_metrics_middleware import BillableCategory
from litellm.types.proxy.gateway_requests import (
GatewayRequestCounts,
@@ -28,6 +43,15 @@ if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
_EMPTY: Final = GatewayRequestCounts(successful_requests=0, failed_requests=0)
+_TABLE: Final = '"LiteLLM_DailyGatewayRequests"'
+_COLUMNS_PER_ROW: Final = 5
+_UTC_NOW: Final = "(NOW() AT TIME ZONE 'UTC')"
+GATEWAY_REQUESTS_JOB_NAME: Final = "update_gateway_requests_job"
+
+_BufferedRows: TypeAlias = tuple[tuple[str, str, str, int, int], ...]
+_BUFFERED_ROWS: Final = TypeAdapter(_BufferedRows)
+_BUFFERED_ENTRIES: Final = TypeAdapter(tuple[str | bytes, ...])
+_NO_COUNTS: Final[GatewayRequestSnapshot] = MappingProxyType({})
def _utc_date() -> str:
@@ -59,20 +83,54 @@ class GatewayRequestAccumulator:
route) however long the database is unreachable.
This buys at-least-once, not exactly-once, and the cost is worth stating.
- The batch commits inside its context manager's ``__aexit__``, so a failure
- raised after the transaction committed (a connection dropped while reading
- the acknowledgement) restores counts that are already persisted, and the
- next flush increments them a second time. Exactly-once would need a dedup
- key the upserts could ignore on replay. For a traffic-volume metric a rare
+ The statement commits on the server before its acknowledgement is read, so
+ a failure raised after the commit (a connection dropped while reading the
+ acknowledgement) restores counts that are already persisted, and the next
+ flush increments them a second time. Exactly-once would need a dedup key
+ the upsert could ignore on replay. For a traffic-volume metric a rare
overcount on a dropped acknowledgement beats losing a whole interval to
every database blip, so the trade is deliberate.
"""
- for key, counts in snapshot.items():
- existing = self._counts.get(key, _EMPTY)
- self._counts[key] = GatewayRequestCounts(
- successful_requests=existing.successful_requests + counts.successful_requests,
- failed_requests=existing.failed_requests + counts.failed_requests,
- )
+ self._counts = dict(fold_counts(chain(self._counts.items(), snapshot.items()))) # mutable-ok: fold replaced
+
+
+def fold_counts(items: Iterable[tuple[GatewayRequestKey, GatewayRequestCounts]]) -> GatewayRequestSnapshot:
+ """Sum counts key-wise; the result stays bounded by (date x category x route)."""
+ folded: Final[dict[GatewayRequestKey, GatewayRequestCounts]] = {} # mutable-ok: local fold returned once
+ for key, counts in items:
+ existing = folded.get(key, _EMPTY)
+ folded[key] = GatewayRequestCounts(
+ successful_requests=existing.successful_requests + counts.successful_requests,
+ failed_requests=existing.failed_requests + counts.failed_requests,
+ )
+ return folded
+
+
+def build_gateway_requests_upsert(snapshot: GatewayRequestSnapshot) -> tuple[str, tuple[str | int, ...]]:
+ """
+ One ``INSERT ... ON CONFLICT DO UPDATE`` that increments every (date, category,
+ route) in the snapshot. Rows are ordered by the conflict key so concurrent
+ writers lock rows in the same order and cannot deadlock.
+ """
+ ordered: Final = sorted(snapshot.items(), key=lambda item: (item[0].date, item[0].category, item[0].route))
+ rows: Final = ", ".join(
+ f"(${base + 1}::text, ${base + 2}::text, ${base + 3}::text, ${base + 4}::bigint, ${base + 5}::bigint, {_UTC_NOW})"
+ for base in range(0, len(ordered) * _COLUMNS_PER_ROW, _COLUMNS_PER_ROW)
+ )
+ sql: Final = (
+ f'INSERT INTO {_TABLE} ("date", "category", "route", "successful_requests", "failed_requests", "updated_at")\n'
+ f"VALUES {rows}\n"
+ 'ON CONFLICT ("date", "category", "route") DO UPDATE SET\n'
+ f' "successful_requests" = {_TABLE}."successful_requests" + EXCLUDED."successful_requests",\n'
+ f' "failed_requests" = {_TABLE}."failed_requests" + EXCLUDED."failed_requests",\n'
+ f' "updated_at" = {_UTC_NOW}'
+ )
+ params: Final[tuple[str | int, ...]] = tuple(
+ value
+ for key, counts in ordered
+ for value in (key.date, key.category, key.route, counts.successful_requests, counts.failed_requests)
+ )
+ return sql, params
async def commit_gateway_requests_to_db(
@@ -80,50 +138,130 @@ async def commit_gateway_requests_to_db(
prisma_client: "PrismaClient",
snapshot: GatewayRequestSnapshot,
) -> None:
- """Upsert one incrementing row per (date, category, route)."""
+ """Increment every (date, category, route) in the snapshot with a single statement."""
if not snapshot:
return
- ordered: Final = sorted(snapshot.items(), key=lambda item: (item[0].date, item[0].category, item[0].route))
+ sql, params = build_gateway_requests_upsert(snapshot)
+ await prisma_client.db.execute_raw(sql, *params) # pyright: ignore[reportAny] # untyped prisma client
- # pyright: ignore[reportAny] on both lines -- prisma's generated client is untyped,
- # so .db and every table action off it resolve to Any at this boundary. The dict
- # literals below are the shape prisma's generated inputs require.
- async with prisma_client.db.batch_() as batcher: # pyright: ignore[reportAny] # untyped prisma client
- for key, counts in ordered:
- columns = asdict(key)
- batcher.litellm_dailygatewayrequests.upsert( # pyright: ignore[reportAny] # untyped prisma client
- where={"date_category_route": columns}, # mutable-ok: prisma input is dict-shaped
- data={ # mutable-ok: prisma input is dict-shaped
- "create": { # mutable-ok: prisma input is dict-shaped
- **columns,
- "successful_requests": counts.successful_requests,
- "failed_requests": counts.failed_requests,
- },
- "update": { # mutable-ok: prisma input is dict-shaped
- "successful_requests": {"increment": counts.successful_requests}, # mutable-ok: as above
- "failed_requests": {"increment": counts.failed_requests}, # mutable-ok: as above
- },
- },
+ verbose_proxy_logger.debug(
+ "Gateway request tracking - committed %d aggregated rows in one statement", len(snapshot)
+ )
+
+
+class GatewayRequestRedisBuffer:
+ """
+ Folds every worker's snapshot through one Redis list so a single pod per
+ interval writes the table, mirroring the spend writer's transaction buffer.
+
+ Each entry is one worker's snapshot as JSON rows; the lock holder pops them,
+ sums them, and commits one statement. A commit failure pushes the summed
+ rows back so the next holder retries, keeping the at-least-once guarantee.
+ If that push fails too, the rows go back to the holder's own accumulator so
+ they ride along with its next flush instead of vanishing with the pop.
+ """
+
+ def __init__(self, *, redis_cache: RedisCache, pod_lock_manager: PodLockManager) -> None:
+ self._redis_cache: Final = redis_cache
+ self._pod_lock_manager: Final = pod_lock_manager
+
+ async def push(self, snapshot: GatewayRequestSnapshot) -> None:
+ if not snapshot:
+ return
+ rows: Final[_BufferedRows] = tuple(
+ (key.date, key.category, key.route, counts.successful_requests, counts.failed_requests)
+ for key, counts in snapshot.items()
+ )
+ await self._redis_cache.async_rpush(key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, values=(json.dumps(rows),))
+
+ async def _pop_batch(self) -> tuple[str | bytes, ...]:
+ popped: Final[object] = await self._redis_cache.async_lpop( # pyright: ignore[reportAny] # redis returns Any
+ key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT
+ )
+ if not popped:
+ return ()
+ return _BUFFERED_ENTRIES.validate_python(popped if isinstance(popped, list) else (popped,))
+
+ async def _pop_all(self) -> AsyncIterator[str | bytes]:
+ while True:
+ batch = await self._pop_batch()
+ for entry in batch:
+ yield entry
+ if len(batch) < MAX_REDIS_BUFFER_DEQUEUE_COUNT:
+ return
+
+ async def pop(self) -> GatewayRequestSnapshot:
+ entries: Final = tuple([entry async for entry in self._pop_all()])
+ return fold_counts(
+ (
+ GatewayRequestKey(date=date, category=category, route=route),
+ GatewayRequestCounts(successful_requests=succeeded, failed_requests=failed),
)
+ for entry in entries
+ for date, category, route, succeeded, failed in _BUFFERED_ROWS.validate_json(entry)
+ )
- verbose_proxy_logger.debug("Gateway request tracking - committed %d aggregated rows", len(ordered))
+ async def commit_if_leader(self, prisma_client: "PrismaClient") -> GatewayRequestSnapshot:
+ """
+ Drain the list and write it as one statement, but only on the pod holding the job lock.
+
+ The lock is a lease, never released: the holder re-enters it on every flush and
+ keeps committing alone until the TTL lapses, so the primary sees one statement
+ per flush interval deployment-wide instead of one per worker.
+
+ Returns the popped rows that could be neither committed nor re-queued, for the
+ caller to keep in memory. Empty on success.
+ """
+ if not await self._pod_lock_manager.acquire_lock(cronjob_id=GATEWAY_REQUESTS_JOB_NAME):
+ return _NO_COUNTS
+ buffered: Final = await self.pop()
+ try:
+ await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=buffered)
+ except Exception: # noqa: BLE001 -- a failed commit must not stop the scheduler
+ verbose_proxy_logger.warning(
+ "Gateway request tracking - failed to commit %d buffered rows, re-queuing to Redis for the next flush",
+ len(buffered),
+ exc_info=True,
+ )
+ return await self._requeue(buffered)
+ return _NO_COUNTS
+
+ async def _requeue(self, snapshot: GatewayRequestSnapshot) -> GatewayRequestSnapshot:
+ try:
+ await self.push(snapshot)
+ except Exception: # noqa: BLE001 -- the rows go back to the caller's accumulator instead
+ verbose_proxy_logger.warning(
+ "Gateway request tracking - Redis re-queue failed, keeping %d rows in memory for the next flush",
+ len(snapshot),
+ exc_info=True,
+ )
+ return snapshot
+ return _NO_COUNTS
async def flush_gateway_requests(
prisma_client: "PrismaClient",
accumulator: GatewayRequestAccumulator,
+ redis_buffer: GatewayRequestRedisBuffer | None = None,
) -> None:
"""
Scheduler entrypoint. Never raises: a metering failure must not kill the job.
+ With ``redis_buffer`` the snapshot goes to Redis and only the lease holder
+ writes to Postgres. Shutdown passes no buffer so a departing worker writes its
+ own counts directly instead of parking them behind a lease it may not hold.
+
``CancelledError`` is deliberately not caught, so a flush cancelled during
shutdown drops its snapshot rather than restoring counts onto an accumulator
the process is about to discard.
"""
snapshot: Final = accumulator.drain()
try:
- await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=snapshot)
+ if redis_buffer is None:
+ await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=snapshot)
+ else:
+ await redis_buffer.push(snapshot)
except Exception: # noqa: BLE001 -- a failed flush must not stop the scheduler
accumulator.restore(snapshot)
verbose_proxy_logger.warning(
@@ -131,3 +269,13 @@ async def flush_gateway_requests(
len(snapshot),
exc_info=True,
)
+ return
+ if redis_buffer is None:
+ return
+ try:
+ accumulator.restore(await redis_buffer.commit_if_leader(prisma_client))
+ except Exception: # noqa: BLE001 -- entries still in Redis are drained by the next flush
+ verbose_proxy_logger.warning(
+ "Gateway request tracking - leader drain failed, buffered rows stay in Redis for the next flush",
+ exc_info=True,
+ )
diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py
index 10683550f85..c22d35509c1 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py
@@ -17,7 +17,8 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
-from litellm.types.guardrails import GuardrailEventHooks
+from litellm.proxy.common_utils.callback_utils import add_guardrail_scan_id
+from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailIntegrations
from litellm.types.utils import (
GenericGuardrailAPIInputs,
GuardrailStatus,
@@ -218,6 +219,13 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
metadata: Final = request_data.get("metadata") or {}
request_data["metadata"] = metadata
metadata["_openai_moderation_response"] = moderation_response.model_dump()
+ add_guardrail_scan_id(
+ request_data=request_data,
+ scan_id=moderation_response.id,
+ guardrail_name=self.guardrail_name,
+ provider=SupportedGuardrailIntegrations.OPENAI_MODERATION.value,
+ stage=GuardrailEventHooks.post_call if input_type == "response" else GuardrailEventHooks.pre_call,
+ )
# Check if content is flagged and raise exception if needed
self._check_moderation_result(moderation_response)
diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py
index b73d3adb99e..3bc0dfabefc 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py
@@ -721,10 +721,18 @@ class PanwPrismaAirsHandler(CustomGuardrail):
}
}
- def _record_scan_id(self, request_data: dict[str, object], scan_result: Mapping[str, object]) -> None:
+ def _record_scan_id(
+ self, request_data: dict[str, object], scan_result: Mapping[str, object], stage: GuardrailEventHooks
+ ) -> None:
"""Surface the AIRS scan id on the response, so allowed calls are auditable too."""
scan_id: Final = scan_result.get("scan_id")
- add_guardrail_scan_id(request_data=request_data, scan_id=str(scan_id) if scan_id else None)
+ add_guardrail_scan_id(
+ request_data=request_data,
+ scan_id=str(scan_id) if scan_id else None,
+ guardrail_name=self.guardrail_name,
+ provider=self._PROVIDER_NAME,
+ stage=stage,
+ )
def _handle_api_error_with_logging(
self,
@@ -948,7 +956,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
event_type=GuardrailEventHooks.post_call,
)
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
- self._record_scan_id(request_data, scan_result)
+ self._record_scan_id(request_data, scan_result, GuardrailEventHooks.post_call)
def _check_and_mark_scanned(self, data: dict, scan_type: str) -> bool:
"""
@@ -1078,7 +1086,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.pre_call,
)
- self._record_scan_id(data, scan_result)
+ self._record_scan_id(data, scan_result, GuardrailEventHooks.pre_call)
action: Final = scan_result.get("action", "block")
category: Final = scan_result.get("category", "unknown")
@@ -1199,7 +1207,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.post_call,
)
- self._record_scan_id(data, scan_result)
+ self._record_scan_id(data, scan_result, GuardrailEventHooks.post_call)
action: Final = scan_result.get("action", "block")
category: Final = scan_result.get("category", "unknown")
@@ -1401,7 +1409,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.post_call,
)
- self._record_scan_id(request_data, scan_result)
+ self._record_scan_id(request_data, scan_result, GuardrailEventHooks.post_call)
# Add guardrail to applied guardrails header for observability
add_guardrail_to_applied_guardrails_header(
@@ -1475,7 +1483,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
)
continue
- self._record_scan_id(request_data, scan_result)
+ self._record_scan_id(
+ request_data,
+ scan_result,
+ GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call,
+ )
action = scan_result.get("action", "block")
masked_args = self._masked_tool_call_arguments(
@@ -1829,7 +1841,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
new_texts.append(text)
continue
- self._record_scan_id(request_data, scan_result)
+ self._record_scan_id(
+ request_data,
+ scan_result,
+ GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call,
+ )
action = scan_result.get("action", "block")
masked_text = self._get_masked_text(scan_result, is_response=is_response)
@@ -1901,7 +1917,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
)
# If we reach here, fallback_on_error="allow"
else:
- self._record_scan_id(request_data, mcp_scan_result)
+ self._record_scan_id(request_data, mcp_scan_result, GuardrailEventHooks.pre_call)
action = mcp_scan_result.get("action", "block")
masked_text = self._get_masked_text(mcp_scan_result, is_response=False)
if action == "allow":
diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py
index 924f84be5f4..3250ae5cca9 100644
--- a/litellm/proxy/litellm_pre_call_utils.py
+++ b/litellm/proxy/litellm_pre_call_utils.py
@@ -235,6 +235,7 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
"applied_policies",
"policy_sources",
"guardrail_scan_ids",
+ "guardrail_scan_metadata",
"routing_decision",
GATEWAY_INJECTED_CACHE_METADATA_KEY,
"pillar_response_headers",
@@ -291,6 +292,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS: Final = (
"applied_policies",
"policy_sources",
"guardrail_scan_ids",
+ "guardrail_scan_metadata",
"routing_decision",
GATEWAY_INJECTED_CACHE_METADATA_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
diff --git a/litellm/proxy/management_endpoints/cost_tracking_settings.py b/litellm/proxy/management_endpoints/cost_tracking_settings.py
index 204051c3715..493f75008c3 100644
--- a/litellm/proxy/management_endpoints/cost_tracking_settings.py
+++ b/litellm/proxy/management_endpoints/cost_tracking_settings.py
@@ -18,6 +18,7 @@ from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
import litellm
+from litellm._internal_context import current_billing_time, pinned_billing_time
from litellm._logging import verbose_proxy_logger
from litellm.cost_calculator import completion_cost
from litellm.proxy._types import (
@@ -27,7 +28,15 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
-from litellm.types.utils import CostPerToken, LlmProvidersSet, ModelInfo
+from litellm.types.utils import (
+ CostBreakdown,
+ CostPerToken,
+ LlmProvidersSet,
+ ModelInfo,
+ ModelResponse,
+ PromptTokensDetailsWrapper,
+ Usage,
+)
router: Final = APIRouter()
@@ -46,13 +55,15 @@ def _configured_price(key: str, sources: tuple[Mapping[str, object], ...]) -> fl
def _extract_custom_pricing(
- litellm_params: Mapping[str, object], model_info: Mapping[str, object]
+ litellm_params: Mapping[str, object], model_info: Mapping[str, object], builtin: ModelInfo | None
) -> CostPerToken | None:
"""
Pull per-token pricing configured on a deployment so on-prem / self-hosted
models (absent from the public cost map) still estimate a real cost.
Pricing may live on ``litellm_params`` or ``model_info``; ``litellm_params``
- wins, matching the router's cost-map registration precedence.
+ wins, matching the router's cost-map registration precedence. Cache rates the
+ deployment leaves unset come from the backend model's built-in entry, then its
+ own input rate, again matching what the router registers for live billing.
"""
sources: Final = (litellm_params, model_info)
input_price: Final = _configured_price("input_cost_per_token", sources)
@@ -61,15 +72,21 @@ def _extract_custom_pricing(
if input_price is None and output_price is None:
return None
+ input_rate: Final = input_price or 0.0
+ cache_sources: Final = sources if builtin is None else (*sources, builtin)
+ cache_read_price: Final = _configured_price("cache_read_input_token_cost", cache_sources)
+ cache_creation_price: Final = _configured_price("cache_creation_input_token_cost", cache_sources)
return CostPerToken(
- input_cost_per_token=input_price or 0.0,
+ input_cost_per_token=input_rate,
output_cost_per_token=output_price or 0.0,
+ cache_read_input_token_cost=input_rate if cache_read_price is None else cache_read_price,
+ cache_creation_input_token_cost=input_rate if cache_creation_price is None else cache_creation_price,
)
-def _lookup_model_info(model: str) -> ModelInfo | None:
+def _lookup_model_info(model: str, custom_llm_provider: str | None = None) -> ModelInfo | None:
try:
- return litellm.get_model_info(model=model)
+ return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception:
return None
@@ -98,17 +115,14 @@ def _resolve_model_for_cost_lookup(model: str) -> ResolvedCostModel:
model_info: Final = first_deployment.get("model_info", {})
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
provider: Final = str(custom_llm_provider) if custom_llm_provider is not None else None
- custom_cost_per_token: Final = _extract_custom_pricing(litellm_params, model_info)
-
- # Check base_model first (needed for Azure custom deployment names)
+ # base_model wins (needed for Azure custom deployment names)
base_model: Final = model_info.get("base_model") or litellm_params.get("base_model")
- if base_model:
- verbose_proxy_logger.debug("Resolved model '%s' to base_model '%s' from router", model, base_model)
- return ResolvedCostModel(str(base_model), provider, custom_cost_per_token)
-
- resolved_model: Final = litellm_params.get("model")
+ resolved_model: Final = base_model or litellm_params.get("model")
if resolved_model:
verbose_proxy_logger.debug("Resolved model '%s' to '%s' from router", model, resolved_model)
+ custom_cost_per_token: Final = _extract_custom_pricing(
+ litellm_params, model_info, _lookup_model_info(str(resolved_model), provider)
+ )
return ResolvedCostModel(str(resolved_model), provider, custom_cost_per_token)
except Exception as e:
verbose_proxy_logger.debug("Could not resolve model '%s' from router: %s", model, e)
@@ -117,19 +131,59 @@ def _resolve_model_for_cost_lookup(model: str) -> ResolvedCostModel:
return ResolvedCostModel(model, None, None)
-def _calculate_period_costs(num_requests, cost_per_request, input_cost, output_cost, margin_cost):
- """
- Calculate costs for a given number of requests.
+@dataclass(frozen=True, slots=True)
+class CostLines:
+ """Cost of one request split the way the spend logs split it: the cache lines are
+ shares of input_cost and the reasoning line is a share of output_cost."""
- Returns tuple of (total_cost, input_cost, output_cost, margin_cost) or all None if num_requests is None/0.
- """
- if not num_requests:
- return None, None, None, None
- return (
- cost_per_request * num_requests,
- input_cost * num_requests,
- output_cost * num_requests,
- margin_cost * num_requests,
+ total_cost: float
+ input_cost: float
+ output_cost: float
+ margin_cost: float
+ cache_read_cost: float
+ cache_creation_cost: float
+ reasoning_cost: float
+
+ def times(self, num_requests: int | None) -> "CostLines | None":
+ if not num_requests:
+ return None
+ return CostLines(
+ total_cost=self.total_cost * num_requests,
+ input_cost=self.input_cost * num_requests,
+ output_cost=self.output_cost * num_requests,
+ margin_cost=self.margin_cost * num_requests,
+ cache_read_cost=self.cache_read_cost * num_requests,
+ cache_creation_cost=self.cache_creation_cost * num_requests,
+ reasoning_cost=self.reasoning_cost * num_requests,
+ )
+
+
+def _cost_lines(cost_per_request: float, cost_breakdown: CostBreakdown | None) -> CostLines:
+ breakdown: Final = cost_breakdown if cost_breakdown is not None else CostBreakdown()
+ return CostLines(
+ total_cost=cost_per_request,
+ input_cost=breakdown.get("input_cost", 0.0),
+ output_cost=breakdown.get("output_cost", 0.0),
+ margin_cost=breakdown.get("margin_total_amount", 0.0),
+ cache_read_cost=breakdown.get("cache_read_cost", 0.0),
+ cache_creation_cost=breakdown.get("cache_creation_cost", 0.0),
+ reasoning_cost=breakdown.get("reasoning_cost", 0.0),
+ )
+
+
+def _usage_for_estimate(request: CostEstimateRequest) -> Usage:
+ cache_tokens: Final = request.cache_read_input_tokens + request.cache_creation_input_tokens
+ return Usage(
+ prompt_tokens=request.input_tokens,
+ completion_tokens=request.output_tokens,
+ total_tokens=request.input_tokens + request.output_tokens,
+ reasoning_tokens=request.reasoning_tokens,
+ prompt_tokens_details=PromptTokensDetailsWrapper(
+ cached_tokens=request.cache_read_input_tokens,
+ cache_creation_tokens=request.cache_creation_input_tokens,
+ )
+ if cache_tokens
+ else None,
)
@@ -530,11 +584,14 @@ async def estimate_cost(
- model: Model name (e.g., "gpt-4", "claude-3-opus")
- input_tokens: Expected input tokens per request
- output_tokens: Expected output tokens per request
+ - cache_read_input_tokens: Cache-read tokens per request, counted within input_tokens (optional)
+ - cache_creation_input_tokens: Cache-write tokens per request, counted within input_tokens (optional)
+ - reasoning_tokens: Reasoning tokens per request, counted within output_tokens (optional)
- num_requests_per_day: Number of requests per day (optional)
- num_requests_per_month: Number of requests per month (optional)
Returns cost breakdown including:
- - Per-request costs (input, output, margin)
+ - Per-request costs (input, output, margin, plus the cache-read, cache-write and reasoning shares)
- Daily costs (if num_requests_per_day provided)
- Monthly costs (if num_requests_per_month provided)
@@ -543,14 +600,15 @@ async def estimate_cost(
{
"model": "gpt-4",
"input_tokens": 1000,
+ "cache_read_input_tokens": 800,
"output_tokens": 500,
+ "reasoning_tokens": 200,
"num_requests_per_day": 100,
"num_requests_per_month": 3000
}
```
"""
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
- from litellm.types.utils import ModelResponse, Usage
# Resolve model name (handles router aliases like 'e-model-router' -> 'azure_ai/gpt-4')
resolved: Final = _resolve_model_for_cost_lookup(request.model)
@@ -559,15 +617,8 @@ async def estimate_cost(
verbose_proxy_logger.debug("Cost estimate: request.model='%s' resolved to '%s'", request.model, resolved_model)
- # Create a mock response with usage for completion_cost
- mock_response: Final = ModelResponse(
- model=resolved_model,
- usage=Usage(
- prompt_tokens=request.input_tokens,
- completion_tokens=request.output_tokens,
- total_tokens=request.input_tokens + request.output_tokens,
- ),
- )
+ usage: Final = _usage_for_estimate(request)
+ mock_response: Final = ModelResponse(model=resolved_model, usage=usage)
# Create a logging object to capture cost breakdown
litellm_logging_obj: Final = LiteLLMLoggingObj(
@@ -580,92 +631,73 @@ async def estimate_cost(
function_id="cost-estimate",
)
- # Use completion_cost which handles all the logic including margins/discounts
- try:
- cost_per_request: Final = completion_cost(
- completion_response=mock_response,
- model=resolved_model,
- custom_llm_provider=resolved_provider,
- custom_cost_per_token=resolved.custom_cost_per_token,
- litellm_logging_obj=litellm_logging_obj,
- )
- except Exception as e:
- raise HTTPException(
- status_code=404,
- detail={
- "error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e}"
- },
- )
+ # Pinning one moment keeps an off-peak window that opens mid-quote from pricing the totals on
+ # one side of it and the reported rates on the other.
+ with pinned_billing_time(current_billing_time()):
+ # Use completion_cost which handles all the logic including margins/discounts
+ try:
+ cost_per_request: Final = completion_cost(
+ completion_response=mock_response,
+ model=resolved_model,
+ custom_llm_provider=resolved_provider,
+ custom_cost_per_token=resolved.custom_cost_per_token,
+ litellm_logging_obj=litellm_logging_obj,
+ )
+ except Exception as e:
+ raise HTTPException(
+ status_code=404,
+ detail={
+ "error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e}"
+ },
+ )
- # Get cost breakdown from the logging object
- cost_breakdown: Final = litellm_logging_obj.cost_breakdown
+ # The rates come back from the pricing call itself rather than a second lookup, so they are the
+ # ones the cost lines above billed at even when completion_cost infers a provider this endpoint
+ # never resolved (an unrouted "xai/grok-4" prices on xai's inclusive tier thresholds; a lookup
+ # here without that provider would report the sub-200k rate for a line billed above it).
+ rates: Final = litellm_logging_obj.billed_token_rates
+ per_request: Final = _cost_lines(cost_per_request, litellm_logging_obj.cost_breakdown)
+ daily: Final = per_request.times(request.num_requests_per_day)
+ monthly: Final = per_request.times(request.num_requests_per_month)
- input_cost: Final = cost_breakdown.get("input_cost", 0.0) if cost_breakdown else 0.0
- output_cost: Final = cost_breakdown.get("output_cost", 0.0) if cost_breakdown else 0.0
- margin_cost: Final = cost_breakdown.get("margin_total_amount", 0.0) if cost_breakdown else 0.0
-
- model_info: Final = _lookup_model_info(resolved_model)
- mapped_input_price: Final = model_info.get("input_cost_per_token") if model_info is not None else None
- mapped_output_price: Final = model_info.get("output_cost_per_token") if model_info is not None else None
+ model_info: Final = _lookup_model_info(resolved_model, resolved_provider)
mapped_provider: Final = model_info.get("litellm_provider") if model_info is not None else None
-
- input_cost_per_token: Final = (
- resolved.custom_cost_per_token["input_cost_per_token"]
- if resolved.custom_cost_per_token is not None
- else mapped_input_price
- )
- output_cost_per_token: Final = (
- resolved.custom_cost_per_token["output_cost_per_token"]
- if resolved.custom_cost_per_token is not None
- else mapped_output_price
- )
custom_llm_provider: Final = mapped_provider if mapped_provider is not None else resolved_provider
- # Calculate daily and monthly costs
- (
- daily_cost,
- daily_input_cost,
- daily_output_cost,
- daily_margin_cost,
- ) = _calculate_period_costs(
- num_requests=request.num_requests_per_day,
- cost_per_request=cost_per_request,
- input_cost=input_cost,
- output_cost=output_cost,
- margin_cost=margin_cost,
- )
- (
- monthly_cost,
- monthly_input_cost,
- monthly_output_cost,
- monthly_margin_cost,
- ) = _calculate_period_costs(
- num_requests=request.num_requests_per_month,
- cost_per_request=cost_per_request,
- input_cost=input_cost,
- output_cost=output_cost,
- margin_cost=margin_cost,
- )
-
return CostEstimateResponse(
model=request.model,
input_tokens=request.input_tokens,
output_tokens=request.output_tokens,
+ cache_read_input_tokens=request.cache_read_input_tokens,
+ cache_creation_input_tokens=request.cache_creation_input_tokens,
+ reasoning_tokens=request.reasoning_tokens,
num_requests_per_day=request.num_requests_per_day,
num_requests_per_month=request.num_requests_per_month,
- cost_per_request=cost_per_request,
- input_cost_per_request=input_cost,
- output_cost_per_request=output_cost,
- margin_cost_per_request=margin_cost,
- daily_cost=daily_cost,
- daily_input_cost=daily_input_cost,
- daily_output_cost=daily_output_cost,
- daily_margin_cost=daily_margin_cost,
- monthly_cost=monthly_cost,
- monthly_input_cost=monthly_input_cost,
- monthly_output_cost=monthly_output_cost,
- monthly_margin_cost=monthly_margin_cost,
- input_cost_per_token=input_cost_per_token,
- output_cost_per_token=output_cost_per_token,
+ cost_per_request=per_request.total_cost,
+ input_cost_per_request=per_request.input_cost,
+ output_cost_per_request=per_request.output_cost,
+ margin_cost_per_request=per_request.margin_cost,
+ cache_read_cost_per_request=per_request.cache_read_cost,
+ cache_creation_cost_per_request=per_request.cache_creation_cost,
+ reasoning_cost_per_request=per_request.reasoning_cost,
+ daily_cost=daily.total_cost if daily is not None else None,
+ daily_input_cost=daily.input_cost if daily is not None else None,
+ daily_output_cost=daily.output_cost if daily is not None else None,
+ daily_margin_cost=daily.margin_cost if daily is not None else None,
+ daily_cache_read_cost=daily.cache_read_cost if daily is not None else None,
+ daily_cache_creation_cost=daily.cache_creation_cost if daily is not None else None,
+ daily_reasoning_cost=daily.reasoning_cost if daily is not None else None,
+ monthly_cost=monthly.total_cost if monthly is not None else None,
+ monthly_input_cost=monthly.input_cost if monthly is not None else None,
+ monthly_output_cost=monthly.output_cost if monthly is not None else None,
+ monthly_margin_cost=monthly.margin_cost if monthly is not None else None,
+ monthly_cache_read_cost=monthly.cache_read_cost if monthly is not None else None,
+ monthly_cache_creation_cost=monthly.cache_creation_cost if monthly is not None else None,
+ monthly_reasoning_cost=monthly.reasoning_cost if monthly is not None else None,
+ input_cost_per_token=rates.input_cost_per_token if rates is not None else None,
+ output_cost_per_token=rates.output_cost_per_token if rates is not None else None,
+ cache_read_input_token_cost=rates.cache_read_input_token_cost if rates is not None else None,
+ cache_creation_input_token_cost=rates.cache_creation_input_token_cost if rates is not None else None,
+ output_cost_per_reasoning_token=rates.output_cost_per_reasoning_token if rates is not None else None,
provider=custom_llm_provider,
)
diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
index d5c3427f29a..2aa7fdc7393 100644
--- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
@@ -2673,6 +2673,8 @@ if MCP_AVAILABLE:
"""
Updates the MCP Server in the db.
+ Partial update: a field left out of the payload keeps its stored value, and a field sent as null is cleared.
+
Parameters:
- payload: UpdateMCPServerRequest - Required. The updated mcp server data.
```
@@ -3098,6 +3100,8 @@ if MCP_AVAILABLE:
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: str | None = Header(None),
):
+ """Partial update: a field left out keeps its stored value, and a field sent as null is cleared, except
+ ``toolset_name`` and ``tools``, which a toolset always has; empty the tool selection with an explicit []."""
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
raise HTTPException(
diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py
index 2ec68a10f65..c050368b3fe 100644
--- a/litellm/proxy/management_endpoints/team_endpoints.py
+++ b/litellm/proxy/management_endpoints/team_endpoints.py
@@ -5601,7 +5601,7 @@ async def team_model_add(
updated_team: Final = await _team_db(prisma_client).update(
where={"team_id": data.team_id},
data={"updated_at": datetime.now(timezone.utc)},
- include={"object_permission": True},
+ include={"litellm_model_table": True, "object_permission": True},
)
if updated_team is None:
raise HTTPException(
@@ -5688,7 +5688,7 @@ async def team_model_delete(
updated_team: Final = await _team_db(prisma_client).update(
where={"team_id": data.team_id},
data={"models": updated_models},
- include={"object_permission": True},
+ include={"litellm_model_table": True, "object_permission": True},
)
if updated_team is None:
raise HTTPException(
diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py
index 3e6434a5afd..c60888e298f 100644
--- a/litellm/proxy/management_endpoints/ui_sso.py
+++ b/litellm/proxy/management_endpoints/ui_sso.py
@@ -22,7 +22,6 @@ from html import escape
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
- Annotated,
Any,
Final,
Literal,
@@ -42,7 +41,7 @@ if TYPE_CHECKING:
import jwt
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
from fastapi.responses import RedirectResponse
-from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError
+from pydantic import BaseModel, TypeAdapter, ValidationError
import litellm
from litellm._logging import verbose_proxy_logger
@@ -95,6 +94,7 @@ from litellm.proxy.auth.auth_utils import (
)
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
+from litellm.proxy.auth.team_grants import TeamModelAliasTable
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.admin_ui_utils import (
admin_ui_disabled,
@@ -209,31 +209,14 @@ def _team_detail_db(repo: TeamRepository) -> "TableActions[_TeamDetailRow]":
return repo.table
-_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
-def _decode_model_aliases(value: object) -> object:
- """``/team/new`` stores team model aliases as a JSON-encoded string in the Json column."""
- if not isinstance(value, str):
- return value
- try:
- return _MODEL_ALIASES_ADAPTER.validate_json(value)
- except ValidationError:
- return None
-
-
-class _TeamModelAliasTable(BaseModel):
- model_config = ConfigDict(protected_namespaces=())
-
- model_aliases: Annotated[Mapping[str, str] | None, BeforeValidator(_decode_model_aliases)] = None
-
-
class _TeamRowGrants(BaseModel):
team_id: str
team_alias: str | None = None
models: tuple[str, ...] = ()
- litellm_model_table: _TeamModelAliasTable | None = None
+ litellm_model_table: TeamModelAliasTable | None = None
class CliSsoTeamDetail(BaseModel):
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 7fcf89904c2..8f631ccde52 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -17,6 +17,7 @@ import time
import traceback
import warnings
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Collection, Mapping, MutableMapping, Sequence
+from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import MappingProxyType, UnionType
from typing import (
@@ -131,6 +132,7 @@ from litellm.router_utils.auto_router_tuning_baseline import (
snapshot_tuning_baselines,
tuning_limit_violation,
)
+from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
@@ -138,11 +140,7 @@ from litellm.types.utils import (
TextCompletionResponse,
TokenCountResponse,
)
-from litellm.utils import (
- _invalidate_model_cost_lowercase_map,
- load_credentials_from_list,
- reapply_runtime_model_cost_registrations,
-)
+from litellm.utils import load_credentials_from_list
if TYPE_CHECKING:
from aiohttp import ClientSession
@@ -426,6 +424,7 @@ from litellm.proxy.db.exception_handler import (
)
from litellm.proxy.db.gateway_request_tracking import (
GatewayRequestAccumulator,
+ GatewayRequestRedisBuffer,
flush_gateway_requests,
)
from litellm.proxy.db.proxy_worker_heartbeat import (
@@ -2357,6 +2356,17 @@ open_telemetry_logger: OpenTelemetry | None = None
gateway_request_accumulator: Final = GatewayRequestAccumulator()
### INITIALIZE GLOBAL LOGGING OBJECT ###
proxy_logging_obj: ProxyLogging = ProxyLogging(user_api_key_cache=user_api_key_cache, premium_user=premium_user)
+
+
+def _gateway_request_redis_buffer() -> GatewayRequestRedisBuffer | None:
+ """Shares the spend writer's transaction-buffer Redis and pod lock when use_redis_transaction_buffer is on."""
+ writer: Final = proxy_logging_obj.db_spend_update_writer
+ redis_cache: Final = writer.redis_update_buffer.redis_cache
+ if redis_cache is None or not writer.redis_update_buffer._should_commit_spend_updates_to_redis():
+ return None
+ return GatewayRequestRedisBuffer(redis_cache=redis_cache, pod_lock_manager=writer.pod_lock_manager)
+
+
### REDIS QUEUE ###
async_result: Final = None
celery_app_conn: Final = None
@@ -2707,6 +2717,12 @@ async def _read_spend_counter_estimate(counter_key: str, fallback_spend: float)
return fallback_spend, False
+@dataclass(frozen=True, slots=True)
+class _PendingSpendIncrement:
+ counter_key: str
+ increment: float
+
+
async def increment_spend_counters(
token: str | None,
team_id: str | None,
@@ -2741,7 +2757,7 @@ async def increment_spend_counters(
cost: Final[float] = response_cost
- async def _key_scope(key_token: str) -> None:
+ async def _key_scope(key_token: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
# key_token arrives pre-hashed from metadata["user_api_key"] (auth flow
# hashes raw "sk-..." keys before they reach the callback). The
# startswith("sk-") check is a safety net matching update_cache —
@@ -2752,30 +2768,29 @@ async def increment_spend_counters(
hash_token(token=key_token) if isinstance(key_token, str) and key_token.startswith("sk-") else key_token
)
key_counter_key: Final = f"spend:key:{hashed_token}"
- if key_counter_key not in reserved_counter_keys:
- await _init_and_increment_spend_counter(
- counter_key=key_counter_key,
- source_cache_key=hashed_token,
- increment=cost,
+ key_pending: Final[tuple[_PendingSpendIncrement, ...]] = (
+ ()
+ if key_counter_key in reserved_counter_keys
+ else (
+ await _prepare_spend_counter_increment(
+ counter_key=key_counter_key,
+ source_cache_key=hashed_token,
+ increment=cost,
+ ),
)
-
- key_obj: Final[object] = await user_api_key_cache.async_get_cache(key=hashed_token)
- if key_obj is None:
- return
- key_budget_limits = getattr(key_obj, "budget_limits", None) or (
- key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
)
- if isinstance(key_budget_limits, str):
- key_budget_limits = json.loads(key_budget_limits)
- if not isinstance(key_budget_limits, list):
- return
- for window in key_budget_limits:
- duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
- key_window_reset_at = window.get("reset_at") if isinstance(window, dict) else window.reset_at
- key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
+
+ async def _key_window_increment(window: object) -> _PendingSpendIncrement | None:
+ duration = (
+ window["budget_duration"] if isinstance(window, dict) else getattr(window, "budget_duration", None)
+ )
+ key_window_reset_at = (
+ window.get("reset_at") if isinstance(window, dict) else getattr(window, "reset_at", None)
+ )
+ key_window_counter: Final = f"spend:key:{hashed_token}:window:{duration}"
key_window_start = get_budget_window_start(window)
- if key_window_counter not in reserved_counter_keys:
- await _init_and_increment_window_spend_counter(
+ pending_window: Final = (
+ await _prepare_window_spend_counter_increment(
counter_key=key_window_counter,
entity_type="Key",
entity_id=hashed_token,
@@ -2783,6 +2798,9 @@ async def increment_spend_counters(
window_start=key_window_start,
increment=cost,
)
+ if key_window_counter not in reserved_counter_keys
+ else None
+ )
await _enqueue_window_spend_row_update(
entity_type=Litellm_EntityType.KEY,
entity_id=hashed_token,
@@ -2792,33 +2810,48 @@ async def increment_spend_counters(
increment=cost,
request_started_at=request_started_at,
)
+ return pending_window
- async def _team_scope(scope_team_id: str) -> None:
- team_counter_key: Final = f"spend:team:{scope_team_id}"
- if team_counter_key not in reserved_counter_keys:
- await _init_and_increment_spend_counter(
- counter_key=team_counter_key,
- source_cache_key=f"team_id:{scope_team_id}",
- increment=cost,
- )
-
- team_obj: Final[object] = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
- if team_obj is None:
- return
- team_budget_limits = getattr(team_obj, "budget_limits", None) or (
- team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
+ key_obj: Final[object] = await user_api_key_cache.async_get_cache(key=hashed_token)
+ if key_obj is None:
+ return key_pending
+ key_budget_limits = getattr(key_obj, "budget_limits", None) or (
+ key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
)
- if isinstance(team_budget_limits, str):
- team_budget_limits = json.loads(team_budget_limits)
- if not isinstance(team_budget_limits, list):
- return
- for window in team_budget_limits:
- duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
- team_window_reset_at = window.get("reset_at") if isinstance(window, dict) else window.reset_at
- team_window_counter = f"spend:team:{scope_team_id}:window:{duration}"
+ if isinstance(key_budget_limits, str):
+ key_budget_limits = json.loads(key_budget_limits)
+ if not isinstance(key_budget_limits, list):
+ return key_pending
+ window_pending: Final = await asyncio.gather(
+ *(_key_window_increment(window) for window in key_budget_limits), return_exceptions=True
+ )
+ return key_pending + tuple(item for item in window_pending if item is not None)
+
+ async def _team_scope(scope_team_id: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
+ team_counter_key: Final = f"spend:team:{scope_team_id}"
+ team_pending: Final[tuple[_PendingSpendIncrement, ...]] = (
+ ()
+ if team_counter_key in reserved_counter_keys
+ else (
+ await _prepare_spend_counter_increment(
+ counter_key=team_counter_key,
+ source_cache_key=f"team_id:{scope_team_id}",
+ increment=cost,
+ ),
+ )
+ )
+
+ async def _team_window_increment(window: object) -> _PendingSpendIncrement | None:
+ duration = (
+ window["budget_duration"] if isinstance(window, dict) else getattr(window, "budget_duration", None)
+ )
+ team_window_reset_at = (
+ window.get("reset_at") if isinstance(window, dict) else getattr(window, "reset_at", None)
+ )
+ team_window_counter: Final = f"spend:team:{scope_team_id}:window:{duration}"
team_window_start = get_budget_window_start(window)
- if team_window_counter not in reserved_counter_keys:
- await _init_and_increment_window_spend_counter(
+ pending_window: Final = (
+ await _prepare_window_spend_counter_increment(
counter_key=team_window_counter,
entity_type="Team",
entity_id=scope_team_id,
@@ -2826,6 +2859,9 @@ async def increment_spend_counters(
window_start=team_window_start,
increment=cost,
)
+ if team_window_counter not in reserved_counter_keys
+ else None
+ )
await _enqueue_window_spend_row_update(
entity_type=Litellm_EntityType.TEAM,
entity_id=scope_team_id,
@@ -2835,25 +2871,47 @@ async def increment_spend_counters(
increment=cost,
request_started_at=request_started_at,
)
+ return pending_window
- async def _team_member_scope(scope_user_id: str, scope_team_id: str) -> None:
+ team_obj: Final[object] = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
+ if team_obj is None:
+ return team_pending
+ team_budget_limits = getattr(team_obj, "budget_limits", None) or (
+ team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
+ )
+ if isinstance(team_budget_limits, str):
+ team_budget_limits = json.loads(team_budget_limits)
+ if not isinstance(team_budget_limits, list):
+ return team_pending
+ window_pending: Final = await asyncio.gather(
+ *(_team_window_increment(window) for window in team_budget_limits), return_exceptions=True
+ )
+ return team_pending + tuple(item for item in window_pending if item is not None)
+
+ async def _team_member_scope(
+ scope_user_id: str, scope_team_id: str
+ ) -> tuple[_PendingSpendIncrement | BaseException, ...]:
team_member_counter_key: Final = f"spend:team_member:{scope_user_id}:{scope_team_id}"
if team_member_counter_key in reserved_counter_keys:
- return
- await _init_and_increment_spend_counter(
- counter_key=team_member_counter_key,
- source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
- increment=cost,
+ return ()
+ return (
+ await _prepare_spend_counter_increment(
+ counter_key=team_member_counter_key,
+ source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
+ increment=cost,
+ ),
)
- async def _user_scope(scope_user_id: str) -> None:
+ async def _user_scope(scope_user_id: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
user_counter_key: Final = f"spend:user:{scope_user_id}"
if user_counter_key in reserved_counter_keys:
- return
- await _init_and_increment_spend_counter(
- counter_key=user_counter_key,
- source_cache_key=scope_user_id,
- increment=cost,
+ return ()
+ return (
+ await _prepare_spend_counter_increment(
+ counter_key=user_counter_key,
+ source_cache_key=scope_user_id,
+ increment=cost,
+ ),
)
scope_coros: Final = tuple(
@@ -2863,7 +2921,7 @@ async def increment_spend_counters(
_team_scope(team_id) if team_id is not None else None,
_team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None,
_user_scope(user_id) if user_id is not None else None,
- _increment_end_user_and_tag_spend_counters(
+ _prepare_end_user_and_tag_spend_increments(
end_user_id=end_user_id,
tags=tags,
response_cost=cost,
@@ -2871,14 +2929,14 @@ async def increment_spend_counters(
)
if end_user_id is not None or tags is not None
else None,
- _increment_model_access_group_spend_counters(
+ _prepare_model_access_group_spend_increments(
model_access_groups=model_access_groups,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
)
if model_access_groups
else None,
- _increment_org_spend_counter(
+ _prepare_org_spend_increment(
org_id=org_id,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
@@ -2893,7 +2951,20 @@ async def increment_spend_counters(
# as orphaned tasks that race the caller's reservation-counter invalidation;
# all scopes settle, then the first error propagates as before.
scope_results: Final = await asyncio.gather(*scope_coros, return_exceptions=True)
- scope_errors: Final = [r for r in scope_results if isinstance(r, BaseException)]
+ scope_errors: Final = tuple(
+ item
+ for scope in scope_results
+ for item in (scope if isinstance(scope, tuple) else (scope,))
+ if isinstance(item, BaseException)
+ )
+ pending: Final = tuple(
+ item
+ for scope in scope_results
+ if not isinstance(scope, BaseException)
+ for item in scope
+ if not isinstance(item, BaseException)
+ )
+ await _apply_spend_counter_increments(pending=pending)
if scope_errors:
raise scope_errors[0]
@@ -2936,41 +3007,49 @@ async def _reconcile_budget_reservation_for_counter_update(
return reserved_counter_keys
-async def _increment_end_user_and_tag_spend_counters(
+async def _prepare_end_user_and_tag_spend_increments(
end_user_id: str | None,
tags: list[str] | None,
response_cost: float,
reserved_counter_keys: set[str],
-) -> None:
- if end_user_id is not None:
- await _init_and_increment_unreserved_spend_counter(
- counter_key=f"spend:end_user:{end_user_id}",
- source_cache_key=end_user_cache_key(end_user_id),
- increment=response_cost,
- reserved_counter_keys=reserved_counter_keys,
- )
-
- if tags is None:
- return
-
- seen_tags: Final[set[str]] = set()
- for tag_name in tags:
- if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags:
- continue
- seen_tags.add(tag_name)
- await _init_and_increment_unreserved_spend_counter(
- counter_key=f"spend:tag:{tag_name}",
- source_cache_key=tag_cache_key(tag_name),
- increment=response_cost,
- reserved_counter_keys=reserved_counter_keys,
- )
+) -> tuple[_PendingSpendIncrement | BaseException, ...]:
+ unique_tags: Final = (
+ tuple(dict.fromkeys(tag for tag in tags if tag and isinstance(tag, str))) if tags is not None else ()
+ )
+ results: Final = await asyncio.gather(
+ *(
+ coro
+ for coro in (
+ _prepare_unreserved_spend_counter_increment(
+ counter_key=f"spend:end_user:{end_user_id}",
+ source_cache_key=end_user_cache_key(end_user_id),
+ increment=response_cost,
+ reserved_counter_keys=reserved_counter_keys,
+ )
+ if end_user_id is not None
+ else None,
+ *(
+ _prepare_unreserved_spend_counter_increment(
+ counter_key=f"spend:tag:{tag_name}",
+ source_cache_key=tag_cache_key(tag_name),
+ increment=response_cost,
+ reserved_counter_keys=reserved_counter_keys,
+ )
+ for tag_name in unique_tags
+ ),
+ )
+ if coro is not None
+ ),
+ return_exceptions=True,
+ )
+ return tuple(item for item in results if item is not None)
-async def _increment_model_access_group_spend_counters(
+async def _prepare_model_access_group_spend_increments(
model_access_groups: Sequence[object],
response_cost: float,
reserved_counter_keys: set[str],
-) -> None:
+) -> tuple[_PendingSpendIncrement | BaseException, ...]:
"""Charge the model access groups that authorized this request.
Without this the counter auth reads is written only by the reservation path, so
@@ -2984,55 +3063,63 @@ async def _increment_model_access_group_spend_counters(
unique_groups: Final = tuple(
dict.fromkeys(group for group in model_access_groups if group and isinstance(group, str))
)
- for group in unique_groups:
- await _init_and_increment_unreserved_spend_counter(
- counter_key=model_access_group_spend_counter_key(group),
- source_cache_key=model_access_group_cache_key(group),
- increment=response_cost,
- reserved_counter_keys=reserved_counter_keys,
- )
+ results: Final = await asyncio.gather(
+ *(
+ _prepare_unreserved_spend_counter_increment(
+ counter_key=model_access_group_spend_counter_key(group),
+ source_cache_key=model_access_group_cache_key(group),
+ increment=response_cost,
+ reserved_counter_keys=reserved_counter_keys,
+ )
+ for group in unique_groups
+ ),
+ return_exceptions=True,
+ )
+ return tuple(item for item in results if item is not None)
-async def _increment_org_spend_counter(
+async def _prepare_org_spend_increment(
org_id: str | None,
response_cost: float,
reserved_counter_keys: set[str],
-) -> None:
+) -> tuple[_PendingSpendIncrement, ...]:
if org_id is None:
- return
+ return ()
- await _init_and_increment_unreserved_spend_counter(
+ pending: Final = await _prepare_unreserved_spend_counter_increment(
counter_key=f"spend:org:{org_id}",
source_cache_key=[f"org_id:{org_id}:with_budget", f"org_id:{org_id}"],
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
+ return (pending,) if pending is not None else ()
-async def _init_and_increment_unreserved_spend_counter(
+async def _prepare_unreserved_spend_counter_increment(
counter_key: str,
source_cache_key: str | list[str],
increment: float,
reserved_counter_keys: set[str],
-) -> None:
+) -> _PendingSpendIncrement | None:
if counter_key in reserved_counter_keys:
- return
+ return None
- await _init_and_increment_spend_counter(
+ return await _prepare_spend_counter_increment(
counter_key=counter_key,
source_cache_key=source_cache_key,
increment=increment,
)
-async def _init_and_increment_spend_counter(
+async def _prepare_spend_counter_increment(
counter_key: str,
source_cache_key: str | list[str],
increment: float,
-):
+) -> _PendingSpendIncrement:
"""
Initialize counter from the authoritative DB spend value if not yet
- set, then atomically increment in both in-memory and Redis.
+ set, then return the pending increment for the caller to apply in one
+ pipelined Redis call.
On first access per pod:
1. Check spend_counter_cache (in-memory -> Redis via DualCache)
@@ -3044,13 +3131,13 @@ async def _init_and_increment_spend_counter(
the counter as absent and seed it. Using increment means the worst case
is over-counting (conservative, blocks slightly early) rather than
under-counting (would allow overspend).
- 4. Increment atomically (both in-memory + Redis)
+ 4. Increment is returned for the caller to apply via pipeline
"""
await _ensure_spend_counter_initialized(
counter_key=counter_key,
source_cache_key=source_cache_key,
)
- await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
+ return _PendingSpendIncrement(counter_key=counter_key, increment=increment)
async def _enqueue_window_spend_row_update(
@@ -3102,20 +3189,20 @@ async def _enqueue_window_spend_row_update(
)
-async def _init_and_increment_window_spend_counter(
+async def _prepare_window_spend_counter_increment(
counter_key: str,
entity_type: str,
entity_id: str,
window_duration: str | None,
window_start: datetime | None,
increment: float,
-):
+) -> _PendingSpendIncrement | None:
if window_start is None:
verbose_proxy_logger.warning(
"Skipping spend counter increment for invalid budget window %s",
counter_key,
)
- return
+ return None
initialized: Final = await _ensure_window_spend_counter_initialized(
counter_key=counter_key,
@@ -3125,8 +3212,8 @@ async def _init_and_increment_window_spend_counter(
window_start=window_start,
)
if initialized is False:
- return
- await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
+ return None
+ return _PendingSpendIncrement(counter_key=counter_key, increment=increment)
async def _ensure_spend_counter_initialized(
@@ -3259,6 +3346,32 @@ async def _invalidate_spend_counter(counter_key: str):
)
+async def _apply_spend_counter_increments(pending: Sequence[_PendingSpendIncrement]) -> None:
+ if not pending:
+ return
+ redis_cache: Final = spend_counter_cache.redis_cache
+ if redis_cache is None:
+ for item in pending:
+ await spend_counter_cache.async_increment_cache(
+ key=item.counter_key,
+ value=item.increment,
+ refresh_ttl=True,
+ )
+ return
+ ttl: Final = redis_cache.get_ttl()
+ increment_list: Final = [ # mutable-ok: async_increment_pipeline signature requires list[RedisPipelineIncrementOperation]
+ RedisPipelineIncrementOperation(key=item.counter_key, increment_value=item.increment, ttl=ttl)
+ for item in pending
+ ]
+ try:
+ results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list)
+ except Exception:
+ await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending))
+ raise
+ for item, current_value in zip(pending, results or ()):
+ spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value)
+
+
async def update_cache(
token: str | None,
user_id: str | None,
@@ -4436,20 +4549,9 @@ def resolve_classifier_plugin(
def _swap_in_model_cost_map(new_model_cost_map: dict) -> int:
- """Adopt a freshly fetched cost map into this process's litellm state, return the model count"""
- litellm.model_cost = new_model_cost_map
- # Invalidate case-insensitive lookup map since model_cost was replaced
- _invalidate_model_cost_lowercase_map()
- # Repopulate provider model sets (e.g. litellm.anthropic_models) so that
- # wildcard patterns like "anthropic/*" include any newly added models.
- litellm.add_known_models(model_cost_map=new_model_cost_map)
- # Counted before the re-apply below, which writes into this same dict, so the
- # number reported describes the fetched price data alone.
- fetched_model_count: Final = len(new_model_cost_map) if new_model_cost_map else 0
- # The swap discards everything registered at runtime (deployment model_info,
- # register_model overrides), so put it back on top of the fresh catalog.
- reapply_runtime_model_cost_registrations()
- return fetched_model_count
+ from litellm.litellm_core_utils.get_model_cost_map import adopt_model_cost_map
+
+ return adopt_model_cost_map(new_model_cost_map)
def should_load_db_object(object_type: str | SupportedDBObjectType) -> bool:
@@ -9543,7 +9645,7 @@ class ProxyStartupEvent:
flush_gateway_requests,
"interval",
seconds=batch_writing_interval,
- args=(prisma_client, gateway_request_accumulator),
+ args=(prisma_client, gateway_request_accumulator, _gateway_request_redis_buffer()),
id="update_gateway_requests_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py
index 540d492beec..599e978df6a 100644
--- a/litellm/responses/utils.py
+++ b/litellm/responses/utils.py
@@ -543,6 +543,49 @@ class ResponsesAPIRequestUtils:
return request_input
+ @staticmethod
+ def strip_encrypted_reasoning_from_input(request_input: object) -> None:
+ """Drop reasoning items the routed deployment cannot decrypt, keeping their readable summary.
+
+ Mutates ``request_input`` in place: the router's fallback snapshot shares this
+ list object, so a rebound list would replay the stripped items on the fallback hop.
+ """
+ if not isinstance(request_input, list):
+ return
+ items: Final = cast(list[object], request_input) # cast-ok: untyped client json
+ stripped: Final = tuple(ResponsesAPIRequestUtils._without_encrypted_reasoning(item) for item in items)
+ items[:] = (item for item in stripped if item is not None) # rebind-ok: list shared with fallback snapshot
+
+ @staticmethod
+ def _without_encrypted_reasoning(item: object) -> object | None:
+ if not isinstance(item, dict):
+ return item
+ reasoning: Final = cast(Mapping[str, object], item) # cast-ok: untyped client json
+ if reasoning.get("type") != "reasoning" or not reasoning.get("encrypted_content"):
+ return reasoning
+ readable: Final = any(
+ ResponsesAPIRequestUtils._has_readable_text(reasoning.get(key)) for key in ("summary", "content")
+ )
+ if not readable:
+ return None
+ kept: Final[dict[str, object]] = { # mutable-ok: request item rebuilt without the undecryptable keys
+ key: value for key, value in reasoning.items() if key not in ("encrypted_content", "id")
+ }
+ return kept
+
+ @staticmethod
+ def _has_readable_text(value: object) -> bool:
+ """A reasoning item's ``summary``/``content`` carries readable text: a non-empty string, or a
+ list holding at least one block with a non-empty ``text`` field (summary_text / output_text)."""
+ if isinstance(value, str):
+ return bool(value.strip())
+ if isinstance(value, list):
+ return any(
+ isinstance(block, dict) and bool(cast(Mapping[str, object], block).get("text")) # cast-ok: untyped json
+ for block in value
+ )
+ return False
+
@staticmethod
def _build_responses_api_response_id(
custom_llm_provider: str | None,
diff --git a/litellm/router.py b/litellm/router.py
index 934a4ac86a9..68ace283949 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -1317,6 +1317,43 @@ class Router:
if isinstance(litellm.input_callback, list):
litellm.input_callback = [c for c in litellm.input_callback if id(c) not in selector_ids]
+ def _apply_updated_routing_strategy_args(self) -> None:
+ """
+ Re-link the default group's selector to the current `routing_strategy_args`.
+
+ Selectors freeze their `RoutingArgs` at construction, so a runtime args
+ update would otherwise keep serving the boot-time values until restart.
+ Latency/usage state survives the rebuild: it lives in the shared router
+ cache, not on the selector.
+ """
+ strategy: Final = self._normalize_strategy(self.routing_strategy)
+ if strategy == "lar1":
+ from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy
+
+ apply_lar1_routing_strategy(self, self.routing_strategy_args)
+ return
+
+ attr: Final = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy or "")
+ current: Final = getattr(self, attr, None) if attr is not None else None
+ if attr is None or current is None:
+ return
+
+ try:
+ rebuilt: Final = self._build_strategy_selector(
+ strategy=strategy or "",
+ routing_strategy_args=self.routing_strategy_args,
+ )
+ except (TypeError, ValidationError):
+ verbose_router_logger.exception(
+ "Invalid routing_strategy_args %s for '%s'; keeping the previous ones",
+ self.routing_strategy_args,
+ strategy,
+ )
+ return
+
+ self._unregister_router_selectors((current,))
+ setattr(self, attr, rebuilt)
+
def routing_strategy_init(self, routing_strategy: RoutingStrategy | str, routing_strategy_args: dict):
verbose_router_logger.info("Routing strategy: %s", routing_strategy)
self._validate_routing_strategy(routing_strategy)
@@ -11130,6 +11167,40 @@ class Router:
return ids
+ def get_candidate_model_ids_for_route(self, model: str, team_id: str | None = None) -> frozenset[str]:
+ """
+ Deployment ids that could serve ``model`` for ``team_id``, unioned across the paths
+ the router resolves a route through: ``model_group_alias``, a routing group, the
+ ``model_name`` and team indexes, and wildcard pattern routes. Read-only and
+ side-effect-free, unlike ``_common_checks_available_deployment`` which also applies
+ fallbacks and can raise. Lets a pre-call check tell a genuine cross-group route from
+ same-group unavailability without re-deriving that precedence at the call site, and
+ without leaking deployment ids into request kwargs bound for the provider.
+ """
+ resolved: Final = self._get_model_from_alias(model=model) or model
+ routing_group_members: Final = self._get_routing_group_deployments(model=resolved, team_id=team_id)
+ if routing_group_members is not None:
+ return self._deployment_ids(routing_group_members)
+ if resolved in self.model_names:
+ return self._deployment_ids(self._get_all_deployments(model_name=resolved, team_id=team_id))
+ team_router: Final = self.team_pattern_routers.get(team_id) if team_id is not None else None
+ return self._deployment_ids(
+ (
+ *self._get_all_deployments(model_name=resolved, team_id=team_id),
+ *(self.pattern_router.route(resolved) or ()),
+ *((team_router.route(resolved) or ()) if team_router is not None else ()),
+ )
+ )
+
+ @staticmethod
+ def _deployment_ids(deployments: Sequence[Mapping[str, object]]) -> frozenset[str]:
+ return frozenset(
+ str(model_info["id"])
+ for deployment in deployments
+ for model_info in (deployment.get("model_info"),)
+ if isinstance(model_info, Mapping) and model_info.get("id") is not None
+ )
+
def has_model_id(self, candidate_id: str) -> bool:
"""
O(1) membership check for a deployment ID without allocating large lists.
@@ -11847,7 +11918,7 @@ class Router:
_existing_router_settings: Final = self.get_settings()
rebuild_routing_groups = False
- relink_lar1_from_args = False
+ routing_args_updated = False
for var in kwargs:
if var in RUNTIME_UPDATABLE_ROUTER_SETTINGS:
if var in _int_settings:
@@ -11886,15 +11957,13 @@ class Router:
)
rebuild_routing_groups = True
elif var == "routing_strategy_args":
- relink_lar1_from_args = True
+ routing_args_updated = True
setattr(self, var, value)
else:
verbose_router_logger.debug("Setting %s is not allowed", var)
- if relink_lar1_from_args and self._normalize_strategy(self.routing_strategy) == "lar1":
- from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy
-
- apply_lar1_routing_strategy(self, self.routing_strategy_args)
+ if routing_args_updated:
+ self._apply_updated_routing_strategy_args()
if rebuild_routing_groups:
self._init_routing_groups(self._routing_groups_input)
diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py
index 805d4ff9080..e902192811c 100644
--- a/litellm/router_strategy/lowest_latency.py
+++ b/litellm/router_strategy/lowest_latency.py
@@ -3,8 +3,11 @@
import random
from collections.abc import Sequence
from datetime import datetime, timedelta
+from math import ceil
from typing import TYPE_CHECKING, Any, Final
+from pydantic import Field
+
import litellm
from litellm import ModelResponse, token_counter, verbose_logger
from litellm.caching.caching import DualCache
@@ -24,6 +27,7 @@ class RoutingArgs(LiteLLMPydanticObjectBase):
ttl: float = 1 * 60 * 60 # 1 hour
lowest_latency_buffer: float = 0
max_latency_list_size: int = 10
+ ttft_percentile: float | None = Field(default=None, gt=0, le=1)
def _average_latency(samples: Sequence[float]) -> float:
@@ -32,6 +36,12 @@ def _average_latency(samples: Sequence[float]) -> float:
return sum(samples) / len(samples)
+def _percentile_latency(samples: Sequence[float], percentile: float) -> float:
+ values: Final = sorted(samples)
+ index: Final = ceil(len(values) * percentile) - 1
+ return values[index]
+
+
def _ttft_seconds(elapsed: timedelta | float) -> float:
if isinstance(elapsed, timedelta):
return elapsed.total_seconds()
@@ -427,14 +437,17 @@ class LowestLatencyLoggingHandler(CustomLogger):
item_rpm = item_map.get(precise_minute, {}).get("rpm", 0)
item_tpm = item_map.get(precise_minute, {}).get("tpm", 0)
- # get average latency or average ttft (depending on streaming/non-streaming)
use_ttft = (
request_kwargs is not None
and request_kwargs.get("stream", None) is not None
and request_kwargs["stream"] is True
and len(item_ttft_latency) > 0
)
- average_latency = _average_latency(item_ttft_latency if use_ttft else item_latency)
+ selected_latency = (
+ _percentile_latency(item_ttft_latency, self.routing_args.ttft_percentile)
+ if use_ttft and self.routing_args.ttft_percentile is not None
+ else _average_latency(item_ttft_latency if use_ttft else item_latency)
+ )
# -------------- #
# Debugging Logic
@@ -443,7 +456,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
# this helps a user to debug why the router picked a specfic deployment #
_deployment_api_base = _deployment.get("litellm_params", {}).get("api_base", "")
if _deployment_api_base is not None:
- _latency_per_deployment[_deployment_api_base] = average_latency
+ _latency_per_deployment[_deployment_api_base] = selected_latency
# -------------- #
# End of Debugging Logic
# -------------- #
@@ -453,7 +466,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
): # if user passed in tpm / rpm in the model_list
continue
else:
- potential_deployments.append((_deployment, average_latency))
+ potential_deployments.append((_deployment, selected_latency))
if len(potential_deployments) == 0:
return None
diff --git a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py
index b623e31ce06..d10153881d9 100644
--- a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py
+++ b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py
@@ -37,13 +37,13 @@ Safe to enable globally:
"""
import time
+from collections.abc import Mapping
from typing import TYPE_CHECKING, Final, Optional, Protocol, cast
import httpx
from litellm._logging import verbose_router_logger
from litellm.exceptions import (
- BadRequestError,
RateLimitError,
ServiceUnavailableError,
)
@@ -158,6 +158,23 @@ class EncryptedContentAffinityCheck(CustomLogger):
return deployment
return None
+ @staticmethod
+ def _request_team_id(request_kwargs: Mapping[str, object]) -> str | None:
+ containers: Final = (request_kwargs.get("metadata"), request_kwargs.get("litellm_metadata"))
+ team_ids: Final = (c.get("user_api_key_team_id") for c in containers if isinstance(c, Mapping))
+ return next((tid for tid in team_ids if isinstance(tid, str)), None)
+
+ def _routed_group_candidate_model_ids(self, request_kwargs: Mapping[str, object], model: str) -> frozenset[str]:
+ """
+ Deployment ids that could serve this turn's routed ``model``, as the router
+ resolves a route (model_group_alias / routing group / model_name / team /
+ pattern). Delegates to the router so the full precedence is not re-derived here
+ and no deployment ids are written into request kwargs bound for the provider.
+ """
+ if self.router is None:
+ return frozenset()
+ return self.router.get_candidate_model_ids_for_route(model=model, team_id=self._request_team_id(request_kwargs))
+
@staticmethod
def _encryption_boundary_key(
litellm_params: object,
@@ -225,10 +242,14 @@ class EncryptedContentAffinityCheck(CustomLogger):
"""
If the request ``input`` contains litellm-encoded item IDs, decode the
embedded ``model_id`` and pin the request to that deployment. Raises
- ``RateLimitError`` / ``ServiceUnavailableError`` / ``BadRequestError``
- when the originating deployment is unavailable and no encryption-boundary
- peer exists, rather than dispatching a doomed request to a non-peer
- deployment. The 429/503 split mirrors the originating cooldown's status:
+ ``RateLimitError`` / ``ServiceUnavailableError`` when the originating
+ deployment is a member of the routed model group but currently unavailable
+ and no encryption-boundary peer exists, rather than dispatching a doomed
+ request to a non-peer deployment. When the origin is not a member of the
+ routed group (an auto-router tier change, a model switch with no peer, a
+ removed deployment, or an unknown/forged marker), the encrypted reasoning is
+ stripped and the request dispatches with its readable history instead. The
+ 429/503 split mirrors the originating cooldown's status:
a 429-induced cooldown surfaces as 429 (with ``Retry-After`` set to the
remaining cooldown window) so OpenAI-compatible clients back off and
retry after the deployment is eligible again.
@@ -285,12 +306,34 @@ class EncryptedContentAffinityCheck(CustomLogger):
request_kwargs["_encrypted_content_affinity_pinned"] = True
return boundary_matches
- # Dispatching to a non-peer would guarantee an upstream
- # `invalid_encrypted_content` 400, so fail fast with a clearer error.
+ # The origin cannot serve this turn's routed group and no peer shares the boundary, so its
+ # encrypted reasoning can never decrypt here. Strip it, keep the readable history, and dispatch
+ # to the routed group instead of failing. Membership is tested by deployment id against the set
+ # the router actually resolved for this route, not by model-group name, so an alias, a
+ # provider-qualified spelling, a team-public name, or a pattern route of the same group is not
+ # mistaken for a tier change. An unknown origin (a removed deployment, or a forged marker) is
+ # treated the same as a cross-group one, which also denies an authenticated caller a
+ # deployment-id existence oracle: a real cross-group id and a nonexistent id both strip and
+ # dispatch rather than returning distinguishable responses. Only a genuine same-group member
+ # that is currently unavailable falls through to the fail-fast, preserving the cooldown contract.
+ routed_group_model_ids: Final = (
+ self._routed_group_candidate_model_ids(request_kwargs, model) if originating is not None else frozenset()
+ )
+ if str(model_id) not in routed_group_model_ids:
+ verbose_router_logger.debug(
+ "EncryptedContentAffinityCheck: model_id=%s is not a candidate for the routed group %s; "
+ "forwarding without its encrypted reasoning",
+ model_id,
+ model,
+ )
+ ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
+ return typed_healthy_deployments
+
+ # The origin is a member of the routed group but currently unavailable (cooled down); fail fast
+ # rather than dispatching to a non-peer, which would guarantee an upstream 400.
raise await self._unavailable_origin_error(
model=model,
model_id=model_id,
- originating=originating,
parent_otel_span=parent_otel_span,
)
@@ -298,25 +341,11 @@ class EncryptedContentAffinityCheck(CustomLogger):
self,
model: str,
model_id: str,
- originating: Deployment | None,
parent_otel_span: Span | None,
) -> Exception:
# Public error messages intentionally omit the originating ``model_id`` so
# an authenticated caller forging encrypted-content markers cannot use the
# error surface to enumerate which deployment IDs exist on this router.
- if originating is None:
- return BadRequestError(
- message=(
- "The deployment that produced this encrypted_content is no "
- "longer configured on this router, and no deployment on the "
- "same encryption boundary is available. Re-issue the request "
- "without the stale encrypted_content items, or restore the "
- "originating deployment."
- ),
- model=model,
- llm_provider="",
- )
-
cooldown: Final = await self._get_origin_cooldown(model_id=model_id, parent_otel_span=parent_otel_span)
if cooldown is not None and str(cooldown.get("status_code")) == "429":
diff --git a/litellm/utils.py b/litellm/utils.py
index a3142cb54fe..f56341a6dec 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -3079,7 +3079,7 @@ def register_model(
# Convert stringified numbers to appropriate numeric types
loaded_model_cost = model_cost
elif isinstance(model_cost, str):
- loaded_model_cost = litellm.get_model_cost_map(url=model_cost)
+ loaded_model_cost = litellm.get_model_cost_map(url=model_cost, max_attempts=1)
if persist_across_reloads:
_registrations: Final[Mapping[str, Mapping[str, object]]] = loaded_model_cost
diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py
index 0af29f069c6..a11f015743b 100644
--- a/tests/code_coverage_tests/router_code_coverage.py
+++ b/tests/code_coverage_tests/router_code_coverage.py
@@ -87,6 +87,7 @@ ignored_function_names = [
"_delete_claude_code_session_router_binding", # Tested through Redis cleanup failure in test_router.py
"_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py
"_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py
+ "_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name)
]
diff --git a/tests/e2e/batches/COVERAGE.md b/tests/e2e/batches/COVERAGE.md
index 8a7b68511ec..919c39f21a2 100644
--- a/tests/e2e/batches/COVERAGE.md
+++ b/tests/e2e/batches/COVERAGE.md
@@ -120,6 +120,36 @@ create traverse gateway -> gateway -> OpenAI (LIT-5347, PR #36240). The pin:
nested managed ids round-trip retrieve. This self-chaining only needs the proxy to
reach its own `PROXY_BASE_URL`, which holds both locally and on the e2e stage.
+## Cleanup
+
+Batch teardown cancels active batches before deleting their input files and keys.
+Raw file IDs from both `model_param` and `provider_fallback` uploads use the upload
+provider when deleted. Model-encoded and managed file IDs route themselves
+
+File deletion and batch cancellation check their responses and retry transient
+failures up to three times. Teardown attempts every registered cleanup before
+reporting failures as test errors. Already deleted files and batches that are
+terminal are safe to clean up again. Managed batch cancellation polls for up to eleven minutes
+before input deletion: the ten-minute provider window plus a propagation margin.
+Accepted cancellation may still report validating or in_progress while the provider
+updates its state. Raw and model-encoded batches are polled until cancelling or
+terminal before input deletion. OpenAI and Azure lifecycle cleanup also deletes
+output and error files returned by terminal batches. Bedrock deletion uses a signed S3 DELETE
+restricted to the configured storage buckets and managed file prefixes. The low-RPM
+test submits with its restricted key and cleans up with the test administrator key
+
+Managed deletion forwards the deployment's trusted bucket configuration and returns
+the requested managed file ID even when stored output metadata carries a provider ID
+
+Azure input uploads request `expires_after` anchored to `created_at` with
+`seconds=1209600`, and the lifecycle tests check the returned expiry. This is a
+fallback for interrupted runs: immediate deletion remains the normal cleanup.
+Azure's minimum supported native expiry is 14 days, so a three-day expiry cannot
+be requested through its Files API
+
+The Azure entry in `files_settings` must use `api_version: 2025-04-01-preview`
+for raw uploads to honor expiry, matching the batch deployment's API version
+
## Terminal state + cost write-back (cross-run marker baton)
The 24h completion window rules out submit-and-wait inside one run, so
diff --git a/tests/e2e/batches/batch_cleanup.py b/tests/e2e/batches/batch_cleanup.py
new file mode 100644
index 00000000000..9284882ad82
--- /dev/null
+++ b/tests/e2e/batches/batch_cleanup.py
@@ -0,0 +1,140 @@
+from builtins import ExceptionGroup
+from collections.abc import Callable
+from itertools import count
+from time import monotonic, sleep
+from typing import Final, Protocol
+
+from batch_client import BatchObject, FileDeleteResponse
+from capabilities import is_managed_id
+from e2e_http import NetworkError, RateLimitedError, Result, Success, UnknownApiError
+from pydantic import BaseModel
+
+CLEANUP_DELAYS: Final = (1.0, 2.0, 4.0)
+BATCH_TERMINAL_STATUSES: Final = frozenset({"completed", "failed", "expired", "cancelled"})
+BATCH_PENDING_STATUSES: Final = frozenset({"validating", "in_progress", "finalizing", "cancelling"})
+BATCH_CANCEL_TIMEOUT_SECONDS: Final = 660.0
+BATCH_CANCEL_POLL_SECONDS: Final = 10.0
+
+
+class BatchCleanupClient(Protocol):
+ def delete_file(self, file_id: str, *, key: str, provider: str | None = None) -> Result[FileDeleteResponse]: ...
+
+ def retrieve_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]: ...
+
+ def cancel_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]: ...
+
+
+def cleanup_result[R: BaseModel](
+ action: Callable[[], Result[R]], *, wait: Callable[[float], None] = sleep
+) -> Result[R]:
+ for delay, result in ((delay, action()) for delay in CLEANUP_DELAYS):
+ match result:
+ case NetworkError() | RateLimitedError():
+ wait(delay)
+ case UnknownApiError(status_code=code) if code in {408, 429, 500, 502, 503, 504}:
+ wait(delay)
+ case _:
+ return result
+ return action()
+
+
+def _require_cleanup_success[R: BaseModel](result: Result[R], operation: str) -> R:
+ match result:
+ case Success(data=data):
+ return data
+ case UnknownApiError(status_code=code):
+ raise AssertionError(f"{operation} failed: HTTP {code}")
+ case _:
+ raise AssertionError(f"{operation} failed: {result.kind}")
+
+
+def cleanup_file(client: BatchCleanupClient, file_id: str, *, key: str, provider: str | None = None) -> None:
+ result: Final = cleanup_result(lambda: client.delete_file(file_id, key=key, provider=provider))
+ if isinstance(result, UnknownApiError) and result.status_code == 404:
+ return
+ deleted: Final = _require_cleanup_success(result, f"Delete file {file_id}")
+ assert deleted.deleted is True or (
+ deleted.deleted is None and is_managed_id(file_id) and deleted.id == file_id and deleted.object == "file"
+ ), f"Delete file {file_id} did not confirm deletion"
+
+
+def cleanup_batch(
+ client: BatchCleanupClient,
+ batch_id: str,
+ *,
+ key: str,
+ provider: str | None = None,
+ delete_output_files: bool = False,
+ wait: Callable[[float], None] = sleep,
+ clock: Callable[[], float] = monotonic,
+) -> None:
+ needs_terminal_state: Final = is_managed_id(batch_id)
+ fetched: Final = _require_cleanup_success(
+ cleanup_result(lambda: client.retrieve_batch(batch_id, key=key, provider=provider)),
+ f"Retrieve batch {batch_id} for cleanup",
+ )
+ if fetched.status in BATCH_TERMINAL_STATUSES:
+ if delete_output_files:
+ _cleanup_batch_outputs(client, fetched, key=key, provider=provider)
+ return
+ if fetched.status == "cancelling" and not needs_terminal_state:
+ return
+ result: Final = (
+ Success(status_code=200, data=fetched)
+ if fetched.status == "cancelling"
+ else cleanup_result(lambda: client.cancel_batch(batch_id, key=key, provider=provider))
+ )
+ conflicted: Final = isinstance(result, UnknownApiError) and result.status_code in {400, 409}
+ if not conflicted:
+ cancelled: Final = _require_cleanup_success(result, f"Cancel batch {batch_id}")
+ assert cancelled.status in BATCH_TERMINAL_STATUSES | BATCH_PENDING_STATUSES, (
+ f"Cancel batch {batch_id} left status {cancelled.status}"
+ )
+ if cancelled.status in BATCH_TERMINAL_STATUSES:
+ if delete_output_files:
+ _cleanup_batch_outputs(client, cancelled, key=key, provider=provider)
+ return
+ if cancelled.status == "cancelling" and not needs_terminal_state:
+ return
+ deadline: Final = clock() + BATCH_CANCEL_TIMEOUT_SECONDS
+ for current in (
+ _require_cleanup_success(
+ cleanup_result(lambda: client.retrieve_batch(batch_id, key=key, provider=provider)),
+ f"Retrieve batch {batch_id} after cancellation",
+ )
+ for _ in count()
+ ):
+ if current.status in BATCH_TERMINAL_STATUSES:
+ if delete_output_files:
+ _cleanup_batch_outputs(client, current, key=key, provider=provider)
+ return
+ assert current.status in ({"cancelling"} if conflicted else BATCH_PENDING_STATUSES), (
+ f"Cancel batch {batch_id} left status {current.status}"
+ )
+ if current.status == "cancelling" and not needs_terminal_state:
+ return
+ assert clock() < deadline, (
+ f"Batch {batch_id} cancellation did not finish within {BATCH_CANCEL_TIMEOUT_SECONDS}s"
+ )
+ wait(BATCH_CANCEL_POLL_SECONDS)
+
+
+def _cleanup_batch_outputs(client: BatchCleanupClient, batch: BatchObject, *, key: str, provider: str | None) -> None:
+ errors: Final = tuple(
+ error
+ for file_id in dict.fromkeys((batch.output_file_id, batch.error_file_id))
+ if file_id is not None and file_id != batch.input_file_id
+ if (error := _output_cleanup_error(client, file_id, key=key, provider=provider)) is not None
+ )
+ if errors:
+ raise ExceptionGroup(f"Batch {batch.id} output cleanup failed", errors)
+
+
+def _output_cleanup_error(
+ client: BatchCleanupClient, file_id: str, *, key: str, provider: str | None
+) -> Exception | None:
+ try:
+ cleanup_file(client, file_id, key=key, provider=provider)
+ except Exception as error:
+ return error
+ return None
diff --git a/tests/e2e/batches/batch_client.py b/tests/e2e/batches/batch_client.py
index 31e49f22450..c9c77e1f12e 100644
--- a/tests/e2e/batches/batch_client.py
+++ b/tests/e2e/batches/batch_client.py
@@ -13,8 +13,9 @@ co-located here because only this suite uses them.
from __future__ import annotations
from dataclasses import dataclass
+from typing import Final, Literal
-from pydantic import BaseModel
+from pydantic import BaseModel, Field
from proxy_client import ProxyClient
from e2e_http import (
@@ -27,6 +28,18 @@ from e2e_http import (
from models import LiteLLMParamsBody
UPLOAD_FILENAME = "batch_input.jsonl"
+AZURE_FILE_EXPIRY_SECONDS: Final = 14 * 24 * 60 * 60
+
+
+class ExpiringFileUploadForm(FileUploadForm):
+ expires_after_anchor: Literal["created_at"] = Field(default="created_at", alias="expires_after[anchor]")
+ expires_after_seconds: int = Field(default=AZURE_FILE_EXPIRY_SECONDS, alias="expires_after[seconds]")
+
+
+def batch_upload_form(provider: str, *, target_model_names: str | None = None) -> FileUploadForm:
+ if provider == "azure":
+ return ExpiringFileUploadForm(target_model_names=target_model_names)
+ return FileUploadForm(target_model_names=target_model_names)
class FileObject(BaseModel):
@@ -37,6 +50,7 @@ class FileObject(BaseModel):
bytes: int | None = None
status: str | None = None
created_at: int | None = None
+ expires_at: int | None = None
class FileList(BaseModel):
@@ -85,7 +99,7 @@ class BatchList(BaseModel):
class FileDeleteResponse(BaseModel):
id: str
object: str | None = None
- deleted: bool
+ deleted: bool | None = None
class BatchCreateBody(BaseModel):
diff --git a/tests/e2e/batches/capabilities.py b/tests/e2e/batches/capabilities.py
index 1bcea0a61ee..17749c2fb87 100644
--- a/tests/e2e/batches/capabilities.py
+++ b/tests/e2e/batches/capabilities.py
@@ -108,6 +108,10 @@ class Capability:
def id(self) -> str:
return f"{self.provider}-{self.scenario}"
+ @property
+ def file_provider(self) -> str | None:
+ return self.provider if self.scenario in {"model_param", "provider_fallback"} else None
+
@property
def jsonl_model(self) -> str:
# Always the provider deployment name. Unified routes via
diff --git a/tests/e2e/batches/conftest.py b/tests/e2e/batches/conftest.py
index 3b133fab680..91a365b6b92 100644
--- a/tests/e2e/batches/conftest.py
+++ b/tests/e2e/batches/conftest.py
@@ -13,7 +13,7 @@ the proxy config.
from __future__ import annotations
import os
-from typing import Iterator
+from typing import Final, Iterator
import pytest
@@ -21,6 +21,7 @@ from batch_client import BatchClient, build_client
from capabilities import PROVIDERS
from e2e_config import MANAGED_FILES_OPT_IN_ENV
from e2e_http import NoBody
+from lifecycle import ResourceManager
from proxy_client import ProxyClient
@@ -52,6 +53,13 @@ def client(proxy: ProxyClient) -> BatchClient:
return build_client(proxy)
+@pytest.fixture
+def resources(client: BatchClient) -> Iterator[ResourceManager]:
+ manager: Final = ResourceManager(client=client.proxy, strict_cleanup=True)
+ yield manager
+ manager.teardown()
+
+
@pytest.fixture(scope="session")
def batch_deployments(client: BatchClient) -> Iterator[None]:
probe = client.proxy.probe("/health/liveliness", params=NoBody())
diff --git a/tests/e2e/batches/test_batch_cleanup.py b/tests/e2e/batches/test_batch_cleanup.py
new file mode 100644
index 00000000000..d0038139dcf
--- /dev/null
+++ b/tests/e2e/batches/test_batch_cleanup.py
@@ -0,0 +1,313 @@
+from builtins import ExceptionGroup
+from collections.abc import Callable
+from typing import Final
+from unittest.mock import Mock, call
+
+import pytest
+from batch_cleanup import BATCH_CANCEL_TIMEOUT_SECONDS, CLEANUP_DELAYS, cleanup_batch, cleanup_file, cleanup_result
+from batch_client import AZURE_FILE_EXPIRY_SECONDS, BatchObject, FileDeleteResponse, batch_upload_form
+from capabilities import CAPABILITIES, Capability
+from e2e_http import NetworkError, RateLimitedError, Result, Success, UnknownApiError
+from lifecycle import ResourceManager
+from models import KeyGenerateBody
+
+MANAGED_FILE_ID: Final = "bGl0ZWxsbV9wcm94eTtmaWxlLTE="
+MANAGED_BATCH_ID: Final = "bGl0ZWxsbV9wcm94eTtiYXRjaC0x"
+
+
+class ExpectedCalls[T]:
+ def __init__(self, values: tuple[T, ...]) -> None:
+ self.values: Final = values
+ self.recorder: Final = Mock()
+
+ def __call__(self, value: T) -> None:
+ self.recorder(value)
+
+ def assert_done(self) -> None:
+ assert tuple(self.recorder.call_args_list) == tuple(call(value) for value in self.values)
+
+
+class CleanupClient:
+ def __init__(
+ self,
+ *,
+ calls: ExpectedCalls[str],
+ files: tuple[Result[FileDeleteResponse], ...] = (),
+ batches: tuple[Result[BatchObject], ...] = (),
+ cancellations: tuple[Result[BatchObject], ...] = (),
+ ) -> None:
+ self.calls: Final = calls
+ self.file_response: Final[Callable[[], Result[FileDeleteResponse]]] = Mock(side_effect=files)
+ self.batch_response: Final[Callable[[], Result[BatchObject]]] = Mock(side_effect=batches)
+ self.cancel_response: Final[Callable[[], Result[BatchObject]]] = Mock(side_effect=cancellations)
+
+ def delete_file(self, file_id: str, *, key: str, provider: str | None = None) -> Result[FileDeleteResponse]:
+ self.calls(f"delete {provider} {file_id}")
+ return self.file_response()
+
+ def retrieve_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]:
+ self.calls(f"retrieve {provider} {batch_id}")
+ return self.batch_response()
+
+ def cancel_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]:
+ self.calls(f"cancel {provider} {batch_id}")
+ return self.cancel_response()
+
+ def generate_key(self, body: KeyGenerateBody) -> str:
+ return "test-key"
+
+ def delete_key(self, key: str) -> None:
+ self.calls(f"delete key {key}")
+
+ def delete_customers(self, user_ids: list[str]) -> None:
+ self.calls(f"delete customers {user_ids}")
+
+
+def batch(status: str) -> Success[BatchObject]:
+ return Success(status_code=200, data=BatchObject(id="batch-1", status=status))
+
+
+def deleted_file(*, deleted: bool = True) -> Success[FileDeleteResponse]:
+ return Success(status_code=200, data=FileDeleteResponse(id="file-1", deleted=deleted))
+
+
+class TestFileCleanup:
+ def test_managed_delete_accepts_the_deleted_file_object(self) -> None:
+ response: Final = Success(
+ status_code=200, data=FileDeleteResponse.model_validate({"id": MANAGED_FILE_ID, "object": "file"})
+ )
+ client: Final = CleanupClient(calls=ExpectedCalls((f"delete None {MANAGED_FILE_ID}",)), files=(response,))
+ cleanup_file(client, MANAGED_FILE_ID, key="test-key")
+ client.calls.assert_done()
+
+ @pytest.mark.parametrize("file_id", ["file-1", MANAGED_FILE_ID])
+ def test_a_success_status_without_a_deletion_confirmation_is_rejected(self, file_id: str) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls((f"delete None {file_id}",)),
+ files=(Success(status_code=200, data=FileDeleteResponse(id=file_id)),),
+ )
+ with pytest.raises(AssertionError, match="did not confirm deletion"):
+ cleanup_file(client, file_id, key="test-key")
+ client.calls.assert_done()
+
+ @pytest.mark.parametrize("cap", CAPABILITIES, ids=[cap.id for cap in CAPABILITIES])
+ def test_deletes_raw_files_through_the_upload_provider(self, cap: Capability) -> None:
+ expected_provider: Final = cap.provider if cap.scenario in {"model_param", "provider_fallback"} else None
+ client: Final = CleanupClient(
+ calls=ExpectedCalls((f"delete {expected_provider} file-1",)), files=(deleted_file(),)
+ )
+ cleanup_file(client, "file-1", key="test-key", provider=cap.file_provider)
+ client.calls.assert_done()
+
+ def test_failed_delete_is_reported_after_remaining_resources_are_cleaned(self) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(("delete azure file-1", "delete key test-key")),
+ files=(UnknownApiError(status_code=403, body="secret response"),),
+ )
+ manager: Final = ResourceManager(client=client, strict_cleanup=True)
+ key: Final = manager.key()
+ manager.defer(lambda: cleanup_file(client, "file-1", key=key, provider="azure"))
+ with pytest.raises(ExceptionGroup) as caught:
+ manager.teardown()
+ client.calls.assert_done()
+ assert len(caught.value.exceptions) == 1
+ assert str(caught.value.exceptions[0]) == "Delete file file-1 failed: HTTP 403"
+
+ def test_success_response_must_confirm_deletion(self) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(("delete None file-1",)), files=(deleted_file(deleted=False),)
+ )
+ with pytest.raises(AssertionError, match="did not confirm deletion"):
+ cleanup_file(client, "file-1", key="test-key")
+ client.calls.assert_done()
+
+ def test_cleanup_is_idempotent_when_file_is_already_deleted(self) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(("delete azure file-1",)),
+ files=(UnknownApiError(status_code=404, body="missing"),),
+ )
+ cleanup_file(client, "file-1", key="test-key", provider="azure")
+ client.calls.assert_done()
+
+ def test_default_resource_cleanup_keeps_existing_best_effort_behavior(self) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(("delete None file-1", "delete key test-key")),
+ files=(UnknownApiError(status_code=403, body="forbidden"),),
+ )
+ manager: Final = ResourceManager(client=client)
+ key: Final = manager.key()
+ manager.defer(lambda: cleanup_file(client, "file-1", key=key))
+ manager.teardown()
+ client.calls.assert_done()
+
+
+class TestCleanupRetries:
+ @pytest.mark.parametrize(
+ "failure",
+ [NetworkError(message="offline"), RateLimitedError(), UnknownApiError(status_code=503, body="unavailable")],
+ )
+ def test_transient_error_retries_and_returns_success(self, failure: Result[FileDeleteResponse]) -> None:
+ responses: Final = (failure, deleted_file())
+ outcomes: Final = Mock(side_effect=responses)
+ delays: Final = ExpectedCalls((1.0,))
+ result: Final[Result[FileDeleteResponse]] = cleanup_result(outcomes, wait=delays)
+ assert isinstance(result, Success) and result.data.deleted
+ delays.assert_done()
+
+ def test_persistent_error_has_bounded_retries(self) -> None:
+ failure: Final = UnknownApiError(status_code=503, body="unavailable")
+ outcomes: Final = Mock(return_value=failure)
+ delays: Final = ExpectedCalls(CLEANUP_DELAYS)
+ result: Final[Result[FileDeleteResponse]] = cleanup_result(outcomes, wait=delays)
+ assert result is failure
+ delays.assert_done()
+ assert outcomes.call_count == len(CLEANUP_DELAYS) + 1
+
+ def test_permanent_error_is_not_retried(self) -> None:
+ failure: Final = UnknownApiError(status_code=403, body="forbidden")
+ responses: Final = (failure, deleted_file())
+ outcomes: Final = Mock(side_effect=responses)
+ delays: Final = ExpectedCalls[float](())
+ assert cleanup_result(outcomes, wait=delays) is failure
+ delays.assert_done()
+ assert outcomes.call_count == 1
+
+
+class TestBatchCancellation:
+ def test_cancelling_batch_is_polled_until_terminal_without_cancelling_again(self) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls((f"retrieve None {MANAGED_BATCH_ID}",) * 3),
+ batches=(batch("cancelling"), batch("cancelling"), batch("cancelled")),
+ )
+ delays: Final = ExpectedCalls((10.0,))
+ cleanup_batch(client, MANAGED_BATCH_ID, key="test-key", wait=delays)
+ client.calls.assert_done()
+ delays.assert_done()
+
+ def test_cancellation_timeout_is_reported_but_file_and_key_cleanup_still_run(self) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(
+ (
+ f"retrieve None {MANAGED_BATCH_ID}",
+ f"retrieve None {MANAGED_BATCH_ID}",
+ "delete None file-1",
+ "delete key test-key",
+ )
+ ),
+ batches=(batch("cancelling"), batch("cancelling")),
+ files=(deleted_file(),),
+ )
+ times: Final = (0.0, BATCH_CANCEL_TIMEOUT_SECONDS)
+ ticks: Final[Callable[[], float]] = Mock(side_effect=times)
+ manager: Final = ResourceManager(client=client, strict_cleanup=True)
+ key: Final = manager.key()
+ manager.defer(lambda: cleanup_file(client, "file-1", key=key))
+ manager.defer(lambda: cleanup_batch(client, MANAGED_BATCH_ID, key=key, clock=ticks))
+ with pytest.raises(ExceptionGroup) as caught:
+ manager.teardown()
+ assert "cancellation did not finish" in str(caught.value.exceptions[0])
+ client.calls.assert_done()
+
+ @pytest.mark.parametrize("status", ["completed", "failed", "expired", "cancelled"])
+ def test_inactive_batch_needs_no_cancellation(self, status: str) -> None:
+ client: Final = CleanupClient(calls=ExpectedCalls(("retrieve None batch-1",)), batches=(batch(status),))
+ cleanup_batch(client, "batch-1", key="test-key")
+ client.calls.assert_done()
+
+ def test_active_batch_is_cancelled_through_its_provider(self) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(("retrieve azure batch-1", "cancel azure batch-1")),
+ batches=(batch("in_progress"), batch("cancelled")),
+ cancellations=(batch("cancelling"),),
+ )
+ cleanup_batch(client, "batch-1", key="test-key", provider="azure")
+ client.calls.assert_done()
+
+ @pytest.mark.parametrize("batch_id", ["batch-1", MANAGED_BATCH_ID])
+ @pytest.mark.parametrize("pending_status", ["validating", "in_progress"])
+ def test_accepted_cancellation_waits_through_stale_provider_status(
+ self, batch_id: str, pending_status: str
+ ) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(
+ (
+ f"retrieve vertex_ai {batch_id}",
+ f"cancel vertex_ai {batch_id}",
+ f"retrieve vertex_ai {batch_id}",
+ f"retrieve vertex_ai {batch_id}",
+ f"retrieve vertex_ai {batch_id}",
+ "delete vertex_ai file-1",
+ "delete key test-key",
+ )
+ ),
+ batches=(batch("validating"), batch(pending_status), batch(pending_status), batch("cancelled")),
+ cancellations=(batch(pending_status),),
+ files=(deleted_file(),),
+ )
+ delays: Final = ExpectedCalls((10.0, 10.0))
+ manager: Final = ResourceManager(client=client, strict_cleanup=True)
+ key: Final = manager.key()
+ manager.defer(lambda: cleanup_file(client, "file-1", key=key, provider="vertex_ai"))
+ manager.defer(lambda: cleanup_batch(client, batch_id, key=key, provider="vertex_ai", wait=delays))
+ manager.teardown()
+ client.calls.assert_done()
+ delays.assert_done()
+
+ @pytest.mark.parametrize("output_delete_fails", [False, True])
+ def test_batch_that_completed_before_cleanup_deletes_output_and_error_files(
+ self, output_delete_fails: bool
+ ) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(("retrieve openai batch-1", "delete openai file-output", "delete openai file-error")),
+ batches=(
+ Success(
+ status_code=200,
+ data=BatchObject(
+ id="batch-1",
+ status="completed",
+ input_file_id="file-input",
+ output_file_id="file-output",
+ error_file_id="file-error",
+ ),
+ ),
+ ),
+ files=(
+ UnknownApiError(status_code=403, body="forbidden") if output_delete_fails else deleted_file(),
+ deleted_file(),
+ ),
+ )
+ if output_delete_fails:
+ with pytest.raises(ExceptionGroup, match="output cleanup failed"):
+ cleanup_batch(client, "batch-1", key="test-key", provider="openai", delete_output_files=True)
+ else:
+ cleanup_batch(client, "batch-1", key="test-key", provider="openai", delete_output_files=True)
+ client.calls.assert_done()
+
+ @pytest.mark.parametrize("status", ["completed", "in_progress"])
+ def test_cancellation_conflict_is_accepted_only_when_batch_became_inactive(self, status: str) -> None:
+ client: Final = CleanupClient(
+ calls=ExpectedCalls(("retrieve None batch-1", "cancel None batch-1", "retrieve None batch-1")),
+ batches=(batch("in_progress"), batch(status)),
+ cancellations=(UnknownApiError(status_code=409, body="conflict"),),
+ )
+ if status == "completed":
+ cleanup_batch(client, "batch-1", key="test-key")
+ else:
+ with pytest.raises(AssertionError, match="Cancel batch batch-1 left status in_progress"):
+ cleanup_batch(client, "batch-1", key="test-key")
+ client.calls.assert_done()
+
+
+class TestAzureFileExpiry:
+ def test_azure_form_serializes_native_expiry_for_the_proxy(self) -> None:
+ form: Final = batch_upload_form("azure", target_model_names="azure-test")
+ assert form.model_dump(by_alias=True, exclude_none=True) == {
+ "purpose": "batch",
+ "target_model_names": "azure-test",
+ "expires_after[anchor]": "created_at",
+ "expires_after[seconds]": AZURE_FILE_EXPIRY_SECONDS,
+ }
+
+ @pytest.mark.parametrize("provider", ["openai", "vertex_ai", "bedrock"])
+ def test_other_providers_keep_their_existing_upload_fields(self, provider: str) -> None:
+ assert batch_upload_form(provider).model_dump(by_alias=True, exclude_none=True) == {"purpose": "batch"}
diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py
index ed7cf656d01..c4b699190b8 100644
--- a/tests/e2e/batches/test_batches_e2e.py
+++ b/tests/e2e/batches/test_batches_e2e.py
@@ -21,14 +21,16 @@ import os
import re
import time
from datetime import datetime, timedelta, timezone
-from typing import Callable
import pytest
from pydantic import BaseModel
-from e2e_config import PROXY_BASE_URL, unique_marker
+from e2e_config import MASTER_KEY, PROXY_BASE_URL, unique_marker
+from batch_cleanup import cleanup_batch, cleanup_file
from batch_client import (
+ AZURE_FILE_EXPIRY_SECONDS,
+ batch_upload_form,
UPLOAD_FILENAME,
BatchClient,
BatchCreateBody,
@@ -155,19 +157,19 @@ def upload_for_scenario(
if cap.scenario == "encoded":
return client.upload_file(
content=content,
- form=FileUploadForm(purpose="batch"),
+ form=batch_upload_form(cap.provider),
model=cap.model,
key=key,
)
if cap.scenario == "unified":
return client.upload_file(
content=content,
- form=FileUploadForm(purpose="batch", target_model_names=cap.model),
+ form=batch_upload_form(cap.provider, target_model_names=cap.model),
key=key,
)
return client.upload_file(
content=content,
- form=FileUploadForm(purpose="batch"),
+ form=batch_upload_form(cap.provider),
key=key,
provider=cap.provider,
)
@@ -188,20 +190,11 @@ def create_for_scenario(
def op_provider(cap: Capability) -> str | None:
- """provider_fallback ids are raw, so retrieve/cancel/list/delete need the provider
+ """provider_fallback batch ids are raw, so retrieve/cancel/list need the provider
hint; the other scenarios encode it into the id and route automatically."""
return cap.provider if cap.scenario == "provider_fallback" else None
-def quietly(action: Callable[[], object]) -> Callable[[], None]:
- """Adapt a value-returning call into a best-effort cleanup the teardown can run."""
-
- def run() -> None:
- action()
-
- return run
-
-
def assert_file_object(file: FileObject, *, provider: str) -> None:
assert file.object == "file", f"file.object={file.object!r}"
assert file.purpose == "batch", f"file.purpose={file.purpose!r}"
@@ -209,6 +202,10 @@ def assert_file_object(file: FileObject, *, provider: str) -> None:
if provider != "bedrock":
assert file.bytes > 0, f"file.bytes={file.bytes!r}"
assert file.status, "file.status missing"
+ if provider == "azure":
+ assert file.expires_at is not None, "Azure batch input has no automatic expiry"
+ assert file.created_at is not None
+ assert file.expires_at - file.created_at == AZURE_FILE_EXPIRY_SECONDS
assert (
file.created_at is not None and file.created_at > 0
), "file.created_at missing"
@@ -249,7 +246,7 @@ def test_batch_lifecycle(
file = unwrap(upload_for_scenario(client, cap, render_jsonl(cap.jsonl_model), key))
resources.defer(
- quietly(lambda: client.delete_file(file.id, key=key, provider=provider))
+ lambda: cleanup_file(client, file.id, key=key, provider=cap.file_provider)
)
assert_file_object(file, provider=cap.provider)
assert matches_id_shape(
@@ -260,7 +257,9 @@ def test_batch_lifecycle(
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(
- quietly(lambda: client.cancel_batch(batch.id, key=key, provider=provider))
+ lambda: cleanup_batch(
+ client, batch.id, key=key, provider=provider, delete_output_files=cap.provider in {"openai", "azure"}
+ )
)
assert batch.id, f"create returned no batch id (body={created.body[:200]})"
@@ -339,7 +338,7 @@ def test_batch_key_model_access_denied(
denied_upload = client.upload_file(
content=render_jsonl(AZURE_BATCH_MODEL),
- form=FileUploadForm(purpose="batch"),
+ form=batch_upload_form("azure"),
model=AZURE_BATCH_MODEL,
key=key,
)
@@ -356,7 +355,7 @@ def test_batch_key_model_access_denied(
)
).id
resources.defer(
- quietly(lambda: client.delete_file(raw_file, key=key, provider="openai"))
+ lambda: cleanup_file(client, raw_file, key=key, provider="openai")
)
denied_create = client.create_batch(
@@ -383,6 +382,7 @@ def test_file_upload_and_delete_outputs(
key=key,
)
)
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="openai")
deleted = unwrap(client.delete_file(file.id, key=key))
@@ -458,12 +458,12 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
- resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
+ resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
_ = client.proxy.poll_logs_for_key(key, min_rows=1)
@@ -517,7 +517,7 @@ class TestBatchFileContent:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert file.id
downloaded = client.proxy.transport.download(
@@ -559,11 +559,11 @@ class TestBatchFileContent:
file = unwrap(
client.upload_file(
content=payload,
- form=FileUploadForm(purpose="batch", target_model_names=provider.model),
+ form=batch_upload_form(provider.name, target_model_names=provider.model),
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider=provider.name)
assert is_managed_id(file.id), (
f"{provider.name}: unified upload must return a managed file id, got {file.id!r}"
@@ -626,7 +626,7 @@ class TestOpenAIFiles:
)
)
resources.defer(
- quietly(lambda: client.delete_file(file.id, key=key, provider="openai"))
+ lambda: cleanup_file(client, file.id, key=key, provider="openai")
)
listed = unwrap(client.list_files(key=key))
@@ -690,7 +690,7 @@ class TestOpenAIFiles:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
fetched = unwrap(client.retrieve_file(file.id, key=key))
assert fetched.id == file.id, "retrieve must echo the uploaded file id"
@@ -760,7 +760,7 @@ class TestBatchRateLimitErrorMapping:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
@@ -803,7 +803,7 @@ class TestBatchEnqueuedTokenLimit:
"""
def _upload_batch_file(
- self, client: BatchClient, resources: ResourceManager, key: str
+ self, client: BatchClient, resources: ResourceManager, key: str, *, cleanup_key: str | None = None
) -> FileObject:
file = unwrap(
client.upload_file(
@@ -813,7 +813,7 @@ class TestBatchEnqueuedTokenLimit:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=cleanup_key or key))
return file
def _generate_enqueued_key(
@@ -850,7 +850,7 @@ class TestBatchEnqueuedTokenLimit:
marker="rpm",
rpm_limit=BATCH_RL_RPM_LIMIT,
)
- file = self._upload_batch_file(client, resources, key)
+ file = self._upload_batch_file(client, resources, key, cleanup_key=MASTER_KEY)
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
@@ -861,7 +861,7 @@ class TestBatchEnqueuedTokenLimit:
)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
- resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
+ resources.defer(lambda: cleanup_batch(client, batch.id, key=MASTER_KEY, delete_output_files=True))
@pytest.mark.covers(
"quota_management.ratelimit.batch_enqueued_tokens.blocks_when_exhausted",
@@ -904,7 +904,7 @@ class TestBatchEnqueuedTokenLimit:
first = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(first)
first_batch = BatchObject.model_validate_json(first.body)
- resources.defer(quietly(lambda: client.cancel_batch(first_batch.id, key=key)))
+ resources.defer(lambda: cleanup_batch(client, first_batch.id, key=key))
blocked = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
assert blocked.status_code == 429, (
@@ -928,7 +928,7 @@ class TestBatchEnqueuedTokenLimit:
)
require_successful_call(retried)
retry_batch = BatchObject.model_validate_json(retried.body)
- resources.defer(quietly(lambda: client.cancel_batch(retry_batch.id, key=key)))
+ resources.defer(lambda: cleanup_batch(client, retry_batch.id, key=key))
ASSUME_ROLE_RAW_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
@@ -984,13 +984,13 @@ class TestBedrockBatchAssumeRole:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="bedrock")
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
- resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
+ resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert batch.id, f"assume-role create returned no batch id: {created.body[:200]}"
assert is_managed_id(batch.id), (
@@ -1044,7 +1044,7 @@ class TestGeminiFiles:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="gemini")
assert file.id, "gemini file upload returned no id"
@@ -1099,13 +1099,13 @@ class TestHostedVllmBatch:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="hosted_vllm")
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
- resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
+ resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert batch.id, f"hosted_vllm create returned no batch id: {created.body[:200]}"
assert batch.status in CREATED_BATCH_STATUSES, (
@@ -1192,7 +1192,7 @@ class TestBatchFailurePaths:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
@@ -1243,12 +1243,12 @@ class TestBatchFailurePaths:
file = unwrap(
client.upload_file(
content=render_jsonl(AZURE_BATCH_RAW_MODEL),
- form=FileUploadForm(purpose="batch"),
+ form=batch_upload_form("azure"),
model=AZURE_BATCH_MODEL,
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert decoded_model_from_id(file.id) == AZURE_BATCH_MODEL, (
f"upload did not encode the azure deployment into the file id: {file.id!r}"
)
@@ -1258,7 +1258,7 @@ class TestBatchFailurePaths:
)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
- resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
+ resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert decoded_model_from_id(batch.id) == AZURE_BATCH_MODEL, (
"create with a foreign encoded file id must route by the file's embedded model, "
@@ -1307,7 +1307,7 @@ class TestBatchSecondHop:
key=key,
)
)
- resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
+ resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert is_managed_id(file.id), (
f"second-hop unified upload must return a managed file id, got {file.id!r}"
)
@@ -1315,7 +1315,7 @@ class TestBatchSecondHop:
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
- resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
+ resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert is_managed_id(batch.id), (
f"second-hop create must return a managed batch id, got {batch.id!r}"
diff --git a/tests/e2e/batches/test_managed_files_enforcement_e2e.py b/tests/e2e/batches/test_managed_files_enforcement_e2e.py
index 7ad0b16adc3..4f703cf0fdc 100644
--- a/tests/e2e/batches/test_managed_files_enforcement_e2e.py
+++ b/tests/e2e/batches/test_managed_files_enforcement_e2e.py
@@ -21,6 +21,7 @@ from typing import Iterator
import pytest
from batch_client import BatchClient, FileObject
+from batch_cleanup import cleanup_file
from capabilities import batch_model_name, is_managed_id, openai_batch_params
from e2e_config import unique_marker
from e2e_http import FileUploadForm, Result, UnknownApiError, unwrap
@@ -108,7 +109,7 @@ def test_cross_user_managed_id_denied_owner_allowed(
key=owner_key,
)
)
- resources.defer(lambda: client.delete_file(uploaded.id, key=owner_key))
+ resources.defer(lambda: cleanup_file(client, uploaded.id, key=owner_key))
assert is_managed_id(uploaded.id), f"expected a managed unified file id, got {uploaded.id}"
denied = client.retrieve_file(uploaded.id, key=other_key)
diff --git a/tests/e2e/coverage_registry/mcp.yaml b/tests/e2e/coverage_registry/mcp.yaml
index ab644118a47..f853e9ff8d6 100644
--- a/tests/e2e/coverage_registry/mcp.yaml
+++ b/tests/e2e/coverage_registry/mcp.yaml
@@ -119,3 +119,11 @@
assertions: [succeeds]
source: "server.py:1089"
rationale: Smoke; rarely used; same auth model as tools
+- id: mcp.list_tools.api_key.toolset_scoped
+ module: mcp
+ tier: P0
+ operation: list_tools
+ auth_family: api_key
+ assertions: [toolset_scoped]
+ source: "user_api_key_auth_mcp.py:2137"
+ rationale: "A key granted a toolset lists exactly the toolset's tools: the rest of the server's catalog stays hidden and every stored name resolves"
diff --git a/tests/e2e/coverage_registry/mgmt.yaml b/tests/e2e/coverage_registry/mgmt.yaml
index 860d96a50b4..c8d7037d2fd 100644
--- a/tests/e2e/coverage_registry/mgmt.yaml
+++ b/tests/e2e/coverage_registry/mgmt.yaml
@@ -76,3 +76,13 @@
- {id: mgmt.credential_migration.check.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:4252", rationale: "Encryption migration (smoke)"}
- {id: mgmt.credential.new.serves_request, module: mgmt, tier: P1, surface: api, assertions: [serves_request], source: "credential_endpoints/endpoints.py:42", rationale: "Stored credential resolves into a deployment and serves a live /messages request"}
- {id: mgmt.model.test_connection.happy_path, module: mgmt, tier: P0, surface: api, assertions: [happy_path], source: "_health_endpoints.py:1785", rationale: "Test Connection for a responses-mode Bedrock Mantle deployment reaches the live provider and reports success; this exact shape 500ed on an acompletion partial before v1.91.0", fail_before_fix: proven}
+- {id: mgmt.mcp_server.new.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:1577", rationale: "Every field of an admin-created MCP server reads back verbatim, by id and in the list, on every replica"}
+- {id: mgmt.mcp_server.list.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:1112", rationale: "The MCP page grid lists a created server with the same field values its detail view reports"}
+- {id: mgmt.mcp_server.update.preserves_unrelated_fields, module: mgmt, tier: P0, surface: api, assertions: [preserves_unrelated_fields], source: "mcp_management_endpoints.py:2665", rationale: "A dashboard edit of one field leaves the others intact and is visible on every replica after one save; edits that took several saves to stick were a customer defect"}
+- {id: mgmt.mcp_server.update.clear_persists, module: mgmt, tier: P0, surface: api, assertions: [clear_persists], source: "mcp_management_endpoints.py:2665", rationale: "An explicit null clears the stored field (absent keeps, null clears)"}
+- {id: mgmt.mcp_server.delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:2139", rationale: "A deleted server is gone by id and from the list on every replica"}
+- {id: mgmt.mcp_toolset.new.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:3009", rationale: "Toolset tools read back under the exact server_id and tool_name written; a toolset stored under one name and read under another granted nothing"}
+- {id: mgmt.mcp_toolset.update.preserves_unrelated_fields, module: mgmt, tier: P0, surface: api, assertions: [preserves_unrelated_fields], source: "mcp_management_endpoints.py:3098", rationale: "Editing the description leaves the tools and name intact"}
+- {id: mgmt.mcp_toolset.update.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:3098", rationale: "Narrowing the tools to one entry reads back exactly that entry"}
+- {id: mgmt.mcp_toolset.update.clear_persists, module: mgmt, tier: P0, surface: api, assertions: [clear_persists], source: "mcp_management_endpoints.py:3098", fail_before_fix: proven, rationale: "An explicit null clears the stored description; the update used to drop null and keep the old value"}
+- {id: mgmt.mcp_toolset.delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:3149", rationale: "A deleted toolset is gone by id and from the list on every replica"}
diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py
index 9d5f1658e91..415c72bbb3c 100644
--- a/tests/e2e/e2e_http.py
+++ b/tests/e2e/e2e_http.py
@@ -49,6 +49,11 @@ class AnthropicHeaders(AuthHeaders):
anthropic_version: str = Field(default="2023-06-01", alias="anthropic-version")
+class PartialBody(BaseModel):
+ """A body for a partial-update route (absent = keep, null = clear): a field left
+ unset is omitted from the wire, and a field set to None is sent as JSON null."""
+
+
class NoBody(BaseModel):
"""Empty body/query for routes that take none."""
@@ -252,6 +257,13 @@ def assert_auth_denied(result: StreamingResponse, context: str) -> None:
f"{context}: expected 401/403, got {result.status_code}: {result.body[:300]}"
)
+
+def wire_body(json: BaseModel) -> dict[str, object]:
+ if isinstance(json, PartialBody):
+ return json.model_dump(by_alias=True, exclude_unset=True)
+ return json.model_dump(by_alias=True, exclude_none=True)
+
+
def _headers(headers: BaseModel) -> dict[str, str]:
dumped: dict[str, object] = headers.model_dump(by_alias=True, exclude_none=True)
return {key: str(value) for key, value in dumped.items()}
@@ -307,9 +319,26 @@ def request_with_retry[T: RetryableResponse](
return issue()
-def _classify[R: BaseModel](
- resp: requests.Response, response_type: type[R]
-) -> Result[R]:
+class ClassifiableResponse(Protocol):
+ """What classifying an outcome reads off a response. requests.Response satisfies
+ it, and so does a fake, so the classification rules are testable on their own."""
+
+ @property
+ def status_code(self) -> int: ...
+
+ @property
+ def ok(self) -> bool: ...
+
+ @property
+ def text(self) -> str: ...
+
+ @property
+ def content(self) -> bytes: ...
+
+ def json(self) -> object: ...
+
+
+def classify[R: BaseModel](resp: ClassifiableResponse, response_type: type[R]) -> Result[R]:
if resp.status_code == 401:
return UnauthorizedError(body=resp.text)
if resp.status_code == 429:
@@ -317,7 +346,8 @@ def _classify[R: BaseModel](
if not resp.ok:
return UnknownApiError(status_code=resp.status_code, body=resp.text)
try:
- return Success(status_code=resp.status_code, data=response_type.model_validate(resp.json()))
+ payload: Final[object] = resp.json() if resp.content else {}
+ return Success(status_code=resp.status_code, data=response_type.model_validate(payload))
except Exception as exc: # noqa: BLE001 - any parse/validation failure is a value
return ValidationError(message=str(exc))
@@ -335,13 +365,13 @@ def post[R: BaseModel](
lambda: requests.post(
str(url),
headers=_headers(headers),
- json=json.model_dump(by_alias=True, exclude_none=True),
+ json=wire_body(json),
timeout=timeout,
)
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
- return _classify(resp, response_type)
+ return classify(resp, response_type)
def get[R: BaseModel](
@@ -363,7 +393,7 @@ def get[R: BaseModel](
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
- return _classify(resp, response_type)
+ return classify(resp, response_type)
def get_external[R: BaseModel](
@@ -383,7 +413,7 @@ def get_external[R: BaseModel](
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
- return _classify(resp, response_type)
+ return classify(resp, response_type)
def delete[R: BaseModel](
@@ -400,14 +430,14 @@ def delete[R: BaseModel](
lambda: requests.delete(
str(url),
headers=_headers(headers),
- json=json.model_dump(by_alias=True, exclude_none=True),
+ json=wire_body(json),
params=_params(params),
timeout=timeout,
)
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
- return _classify(resp, response_type)
+ return classify(resp, response_type)
def patch[R: BaseModel](
@@ -423,13 +453,13 @@ def patch[R: BaseModel](
lambda: requests.patch(
str(url),
headers=_headers(headers),
- json=json.model_dump(by_alias=True, exclude_none=True),
+ json=wire_body(json),
timeout=timeout,
)
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
- return _classify(resp, response_type)
+ return classify(resp, response_type)
def put[R: BaseModel](
@@ -445,13 +475,13 @@ def put[R: BaseModel](
lambda: requests.put(
str(url),
headers=_headers(headers),
- json=json.model_dump(by_alias=True, exclude_none=True),
+ json=wire_body(json),
timeout=timeout,
)
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
- return _classify(resp, response_type)
+ return classify(resp, response_type)
def probe(
@@ -555,7 +585,7 @@ def send(
str(url),
headers=_headers(headers),
params=_params(params),
- json=json.model_dump(by_alias=True, exclude_none=True),
+ json=wire_body(json),
stream=stream,
timeout=timeout,
)
@@ -605,7 +635,7 @@ def upload[R: BaseModel](
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
- return _classify(resp, response_type)
+ return classify(resp, response_type)
def stream_binary(
@@ -623,7 +653,7 @@ def stream_binary(
resp = requests.post(
str(url),
headers=_headers(headers),
- json=json.model_dump(by_alias=True, exclude_none=True),
+ json=wire_body(json),
stream=True,
timeout=timeout,
)
diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py
index 1f55a0f9a56..ed112a79b9b 100644
--- a/tests/e2e/guardrails/guardrails_client.py
+++ b/tests/e2e/guardrails/guardrails_client.py
@@ -7,7 +7,7 @@ from __future__ import annotations
import time
from collections.abc import Callable
from dataclasses import dataclass
-from typing import Literal
+from typing import Final, Literal
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, settle_propagation, unique_marker
from e2e_http import NoBody, Result, StreamingResponse, Success, unwrap
@@ -405,6 +405,29 @@ def build_client(proxy: ProxyClient) -> GuardrailsClient:
return GuardrailsClient(proxy=proxy)
+def poll_until_guardrail_applied(
+ call: Callable[[], StreamingResponse],
+ guardrail_name: str,
+ *,
+ timeout: float = POLL_TIMEOUT,
+ interval: float = POLL_INTERVAL,
+ now: Callable[[], float] = time.monotonic,
+ sleep: Callable[[float], None] = time.sleep,
+) -> StreamingResponse:
+ deadline: Final = now() + timeout
+ if not (result := call()).ok:
+ return result
+ while (
+ guardrail_name
+ not in (name.strip() for name in result.headers.get("x-litellm-applied-guardrails", "").split(","))
+ and (remaining := deadline - now()) > 0
+ ):
+ sleep(min(interval, remaining))
+ if now() >= deadline or not (result := call()).ok:
+ break
+ return result
+
+
def poll_until_blocked[R: BaseModel](call: Callable[[], Result[R]]) -> Result[R]:
"""Retry a call that a guardrail should reject until it is, returning the last result.
diff --git a/tests/e2e/guardrails/test_guardrails_client.py b/tests/e2e/guardrails/test_guardrails_client.py
new file mode 100644
index 00000000000..423c2ede599
--- /dev/null
+++ b/tests/e2e/guardrails/test_guardrails_client.py
@@ -0,0 +1,66 @@
+from dataclasses import dataclass
+from itertools import chain, repeat
+from typing import Final
+
+import pytest
+
+from e2e_http import StreamingResponse
+from guardrails_client import poll_until_guardrail_applied
+
+
+@dataclass
+class Clock:
+ elapsed: float = 0.0
+
+ def now(self) -> float:
+ return self.elapsed
+
+ def sleep(self, seconds: float) -> None:
+ self.elapsed += seconds
+
+
+def _response(applied: str, status: int = 200) -> StreamingResponse:
+ return StreamingResponse(status_code=status, body="{}", headers={"x-litellm-applied-guardrails": applied})
+
+
+def test_waits_for_requested_guardrail_after_an_unrelated_global_guardrail() -> None:
+ clock: Final = Clock()
+ expected: Final = _response("global-filter, tool-permission")
+ responses: Final = iter((_response("global-filter"), expected))
+
+ result: Final = poll_until_guardrail_applied(
+ lambda: next(responses), "tool-permission", timeout=5, interval=2, now=clock.now, sleep=clock.sleep
+ )
+
+ assert result is expected
+ assert clock.elapsed == 2
+
+
+@pytest.mark.parametrize("applied", ("", "global-filter", "tool-permission-sibling"))
+def test_missing_exact_guardrail_returns_failure_evidence_at_deadline(applied: str) -> None:
+ clock: Final = Clock()
+ missing: Final = _response(applied)
+ responses: Final = iter((missing, missing, missing))
+
+ result: Final = poll_until_guardrail_applied(
+ lambda: next(responses), "tool-permission", timeout=5, interval=2, now=clock.now, sleep=clock.sleep
+ )
+
+ assert result is missing
+ assert clock.elapsed == 5
+ with pytest.raises(StopIteration):
+ next(responses)
+
+
+@pytest.mark.parametrize("status", (400, 401, 429, 500))
+def test_http_failure_is_not_hidden_by_a_later_success(status: int) -> None:
+ clock: Final = Clock()
+ failed: Final = _response("", status)
+ responses: Final = iter(chain((failed,), repeat(_response("tool-permission"))))
+
+ result: Final = poll_until_guardrail_applied(
+ lambda: next(responses), "tool-permission", timeout=5, interval=2, now=clock.now, sleep=clock.sleep
+ )
+
+ assert result is failed
+ assert clock.elapsed == 0
diff --git a/tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py b/tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py
index 9ef3650625c..8d1047e53c7 100644
--- a/tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py
+++ b/tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py
@@ -30,6 +30,7 @@ from guardrails_client import (
ToolPermissionParamsBody,
ToolPermissionRuleBody,
poll_until_blocked,
+ poll_until_guardrail_applied,
)
from lifecycle import ResourceManager
from models import ChatResponse, ChatTool, ChatToolFunction
@@ -84,8 +85,8 @@ def _register_tool_permission(client: GuardrailsClient, resources: ResourceManag
resources.defer(lambda: client.delete_guardrail(guardrail_id))
-def _applied_guardrails(outcome: StreamingResponse) -> str:
- return outcome.headers.get("x-litellm-applied-guardrails", "")
+def _applied_guardrails(outcome: StreamingResponse) -> tuple[str, ...]:
+ return tuple(name.strip() for name in outcome.headers.get("x-litellm-applied-guardrails", "").split(","))
def _tool_call_names(response: ChatResponse) -> tuple[str, ...]:
@@ -144,14 +145,17 @@ class TestToolPermissionPreCall:
name = f"e2e-toolperm-allow-{unique_marker()}"
_register_tool_permission(client, resources, name=name)
- outcome = client.chat_raw(
- scoped_key,
- MODEL,
- TOOL_PROMPT,
- guardrails=[name],
- max_tokens=128,
- tools=[ALLOWED_TOOL],
- tool_choice="required",
+ outcome = poll_until_guardrail_applied(
+ lambda: client.chat_raw(
+ scoped_key,
+ MODEL,
+ TOOL_PROMPT,
+ guardrails=[name],
+ max_tokens=128,
+ tools=[ALLOWED_TOOL],
+ tool_choice="required",
+ ),
+ name,
)
assert outcome.ok, f"the permitted tool must be served, got {outcome.status_code}: {outcome.body[:400]}"
diff --git a/tests/e2e/lifecycle.py b/tests/e2e/lifecycle.py
index c9a67ebdb8c..eb9704d4dcb 100644
--- a/tests/e2e/lifecycle.py
+++ b/tests/e2e/lifecycle.py
@@ -8,8 +8,9 @@ ResourceManager; the test registers a cleanup for every resource it creates, and
the fixture's teardown releases them all even when the test body raises.
"""
+from builtins import ExceptionGroup
from dataclasses import dataclass, field
-from typing import Callable, List, Protocol, runtime_checkable
+from typing import Callable, Final, List, Protocol, runtime_checkable
from proxy_client import ProxyClient
from models import KeyGenerateBody
@@ -52,6 +53,7 @@ class ResourceManager:
"""
client: ResourceClient
+ strict_cleanup: bool = False
_cleanups: List[Callable[[], object]] = field(
default_factory=list
) # mutable-ok: append-only teardown registry
@@ -82,8 +84,17 @@ class ResourceManager:
return customer_id
def teardown(self) -> None:
- for cleanup in reversed(self._cleanups):
- try:
- cleanup()
- except Exception:
- pass # best-effort: a failed cleanup must not block the rest
+ failures: Final = tuple(
+ failure for cleanup in reversed(self._cleanups)
+ if (failure := _run_cleanup(cleanup)) is not None
+ )
+ if failures and self.strict_cleanup:
+ raise ExceptionGroup("Resource cleanup failed", failures)
+
+
+def _run_cleanup(cleanup: Callable[[], object]) -> Exception | None:
+ try:
+ cleanup()
+ except Exception as exc:
+ return exc
+ return None
diff --git a/tests/e2e/llm_translation/test_responses_bridge_streaming_e2e.py b/tests/e2e/llm_translation/test_responses_bridge_streaming_e2e.py
index 9a45743a0cd..75817340876 100644
--- a/tests/e2e/llm_translation/test_responses_bridge_streaming_e2e.py
+++ b/tests/e2e/llm_translation/test_responses_bridge_streaming_e2e.py
@@ -16,8 +16,10 @@ into a chat completion chunk. Two customer-visible contracts only hold on that p
from __future__ import annotations
+from typing import Final, Literal
+
import pytest
-from pydantic import BaseModel
+from pydantic import BaseModel, Field
from e2e_config import unique_marker
from e2e_http import StreamingResponse
@@ -51,7 +53,8 @@ class _BridgeChoice(BaseModel):
class _BridgeChunk(BaseModel):
id: str
- choices: list[_BridgeChoice] = []
+ object: Literal["chat.completion.chunk"]
+ choices: list[_BridgeChoice] = Field(default_factory=list)
class _WeatherArgs(BaseModel):
@@ -103,16 +106,19 @@ class TestResponsesBridgeChatCompletionsStreaming:
resources.key(),
ChatBody(
model=bridged_model,
- messages=[ChatMessage(role="user", content=f"Count from 1 to 5, one number per line. {unique_marker()}")],
+ messages=[
+ ChatMessage(role="user", content=f"Count from 1 to 5, one number per line. {unique_marker()}")
+ ],
max_tokens=64,
stream=True,
),
)
- chunks = _bridge_chunks(result)
- ids = {chunk.id for chunk in chunks}
+ chunks: Final = _bridge_chunks(result)
+ assert len(chunks) > 1, "the shared-id contract needs more than one streamed chunk"
+ ids: Final = frozenset(chunk.id for chunk in chunks)
assert len(ids) == 1, f"bridged stream used {len(ids)} different chunk ids: {sorted(ids)[:5]}"
- assert ids.pop().startswith("chatcmpl-"), f"bridged chunk id is not chat-completion shaped: {chunks[0].id}"
+ assert chunks[0].id.strip(), "bridged stream emitted an empty chunk id"
@pytest.mark.covers(
"llm.chat_completions.openai.basic.stream.bridge_streams_sse",
@@ -134,9 +140,9 @@ class TestResponsesBridgeChatCompletionsStreaming:
chunks = _bridge_chunks(result)
content = "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices)
assert content.strip(), f"bridged stream completed with no content deltas: {result.stream_events[:3]}"
- assert any(
- choice.finish_reason for chunk in chunks for choice in chunk.choices
- ), f"bridged stream never emitted a finish_reason: {result.stream_events[-3:]}"
+ assert any(choice.finish_reason for chunk in chunks for choice in chunk.choices), (
+ f"bridged stream never emitted a finish_reason: {result.stream_events[-3:]}"
+ )
assert result.stream_done, f"bridged stream did not terminate with [DONE]: {result.stream_events[-2:]}"
@pytest.mark.covers(
diff --git a/tests/e2e/logging/datadog_reader.py b/tests/e2e/logging/datadog_reader.py
index d0f478185c2..368c20cb6aa 100644
--- a/tests/e2e/logging/datadog_reader.py
+++ b/tests/e2e/logging/datadog_reader.py
@@ -12,8 +12,12 @@ empty result. External reads go through ``e2e_http``.
from __future__ import annotations
+import math
+import random
import time
-from dataclasses import dataclass
+from collections.abc import Callable, Mapping
+from dataclasses import dataclass, field
+from typing import Final
import pytest
from pydantic import BaseModel, ConfigDict, Field
@@ -27,17 +31,33 @@ from e2e_config import (
DD_SITE,
POLL_TIMEOUT,
)
-from e2e_http import URL, Headers, RateLimitedError, Success, post
+from e2e_http import URL, Headers, StreamingResponse, send
-#: How many rate-limited responses in a row one search tolerates before the
-#: hard fail; each retry sleeps a full search interval, so this rides out a
-#: burst from a concurrent consumer of the org-wide search budget.
-_RATE_LIMIT_RETRIES = 5
+type SearchCall = Callable[[str, float], StreamingResponse]
+
+
+def _seconds(value: str | None) -> float | None:
+ if value is None:
+ return None
+ try:
+ seconds: Final = float(value)
+ except ValueError:
+ return None
+ return seconds if math.isfinite(seconds) and seconds >= 0 else None
+
+
+def _rate_limit_delay(headers: Mapping[str, str]) -> float:
+ delays: Final = tuple(
+ delay
+ for name in ("x-ratelimit-reset", "retry-after")
+ if (delay := _seconds(headers.get(name))) is not None
+ )
+ return max(1.0, max(delays, default=DD_SEARCH_INTERVAL))
class _DdAuthHeaders(Headers):
- api_key: str = Field(serialization_alias="DD-API-KEY")
- app_key: str = Field(serialization_alias="DD-APPLICATION-KEY")
+ api_key: str = Field(serialization_alias="DD-API-KEY", repr=False)
+ app_key: str = Field(serialization_alias="DD-APPLICATION-KEY", repr=False)
class _SearchFilter(BaseModel):
@@ -88,8 +108,12 @@ class _SearchResponse(BaseModel):
@dataclass(frozen=True, slots=True)
class DdLogsReader:
site: str
- api_key: str
- app_key: str
+ api_key: str = field(repr=False)
+ app_key: str = field(repr=False)
+ search: SearchCall | None = field(default=None, repr=False)
+ now: Callable[[], float] = field(default=time.monotonic, repr=False)
+ sleep: Callable[[float], None] = field(default=time.sleep, repr=False)
+ jitter: Callable[[], float] = field(default=random.random, repr=False)
def events_for_marker(self, marker: str) -> list[DdLogEvent]:
"""Every ingested event whose attributes carry the marker. DataDog
@@ -108,25 +132,28 @@ class DdLogsReader:
a single event. A 429 backs off and retries - the search budget is
org-wide, so another consumer can empty it under us - while any other
failure stays a hard fail."""
- for _ in range(_RATE_LIMIT_RETRIES):
- result = post(
- URL(f"https://api.{self.site}/api/v2/logs/events/search"),
- headers=_DdAuthHeaders(api_key=self.api_key, app_key=self.app_key),
- json=_SearchRequest(filter=_SearchFilter(query=query)),
- response_type=_SearchResponse,
- timeout=30.0,
- )
- match result:
- case Success(data=page):
- return [event.attributes for event in page.data]
- case RateLimitedError(retry_after_seconds=retry_after):
- time.sleep(retry_after if retry_after else DD_SEARCH_INTERVAL)
- case failure:
- pytest.fail(f"DataDog Logs Search API at api.{self.site} failed: {failure}")
+ return self._events_for_query(query, self.now() + POLL_TIMEOUT)
+
+ def _events_for_query(self, query: str, deadline: float) -> list[DdLogEvent]:
+ search: Final = self.search or self._search_page
+ while (remaining := deadline - self.now()) > 0:
+ if (result := search(query, min(30.0, remaining))).ok:
+ return [event.attributes for event in _SearchResponse.model_validate_json(result.body).data]
+ if result.status_code != 429:
+ pytest.fail(f"DataDog Logs Search API at api.{self.site} failed with HTTP {result.status_code}")
+ if (delay := min(_rate_limit_delay(result.headers) + self.jitter(), deadline - self.now())) > 0:
+ self.sleep(delay)
pytest.fail(
- f"DataDog Logs Search API at api.{self.site} still rate-limited after "
- f"{_RATE_LIMIT_RETRIES} retries {DD_SEARCH_INTERVAL}s apart - the org-wide "
- "logs_public_search_api budget (2 requests per 10s) is exhausted by another consumer"
+ f"DataDog Logs Search API at api.{self.site} remained rate-limited for {POLL_TIMEOUT}s; "
+ "the org-wide logs_public_search_api budget is exhausted"
+ )
+
+ def _search_page(self, query: str, timeout: float) -> StreamingResponse:
+ return send(
+ URL(f"https://api.{self.site}/api/v2/logs/events/search"),
+ headers=_DdAuthHeaders(api_key=self.api_key, app_key=self.app_key),
+ json=_SearchRequest(filter=_SearchFilter(query=query)),
+ timeout=timeout,
)
def poll_events_for_marker(self, marker: str) -> list[DdLogEvent]:
@@ -140,33 +167,42 @@ class DdLogsReader:
hide from the exactly-one assertion - real-DataDog jitter can surface
one call's two events tens of seconds apart. Searches pace at
DD_SEARCH_INTERVAL, not POLL_INTERVAL, to respect the search API's
- request budget. At the deadline the last result is returned as-is."""
- deadline = time.monotonic() + POLL_TIMEOUT
- while time.monotonic() < deadline:
- events = self.events_for_query(query)
+ request budget. Discovery, quota retries, and duplicate detection share
+ one POLL_TIMEOUT deadline; an incomplete settle window fails closed."""
+ deadline: Final = self.now() + POLL_TIMEOUT
+ while (remaining := deadline - self.now()) > 0:
+ events = self._events_for_query(query, deadline)
if events:
- return self._settled_events_for_query(query, events)
- time.sleep(DD_SEARCH_INTERVAL)
- return self.events_for_query(query)
+ return self._settled_events_for_query(query, events, deadline)
+ if (remaining := deadline - self.now()) > 0:
+ self.sleep(min(DD_SEARCH_INTERVAL, remaining))
+ return []
- def _settled_events_for_query(self, query: str, events: list[DdLogEvent]) -> list[DdLogEvent]:
+ def _settled_events_for_query(self, query: str, events: list[DdLogEvent], deadline: float) -> list[DdLogEvent]:
"""Re-read at every search interval until the settle window closes; a
duplicate ends the watch early because more waiting cannot clear it.
Keep the last non-empty result: a transient empty search (index lag)
must not erase events already confirmed earlier in the settle window.
+ A successful final search must reach the full settle window before the
+ shared read-back deadline; otherwise duplicate detection is incomplete.
"""
- settle_deadline = time.monotonic() + DD_SETTLE_SECONDS
+ settle_deadline: Final = self.now() + DD_SETTLE_SECONDS
last_nonempty = events
- while time.monotonic() < settle_deadline:
- time.sleep(DD_SEARCH_INTERVAL)
- latest = self.events_for_query(query)
- if not latest:
- continue
+ if len(events) > 1:
+ return events
+ while (remaining := deadline - self.now()) > 0:
+ self.sleep(min(DD_SEARCH_INTERVAL, remaining))
+ if self.now() >= deadline:
+ break
+ latest = self._events_for_query(query, deadline)
if len(latest) > 1:
return latest
- last_nonempty = latest
- return last_nonempty
+ if latest:
+ last_nonempty = latest
+ if self.now() >= settle_deadline:
+ return last_nonempty
+ pytest.fail(f"DataDog log delivery could not complete its duplicate-detection window within {POLL_TIMEOUT}s")
def build_dd_logs_reader() -> DdLogsReader:
diff --git a/tests/e2e/logging/test_datadog_reader.py b/tests/e2e/logging/test_datadog_reader.py
new file mode 100644
index 00000000000..910a1cefd42
--- /dev/null
+++ b/tests/e2e/logging/test_datadog_reader.py
@@ -0,0 +1,223 @@
+import json
+from collections.abc import Iterator, Sequence
+from dataclasses import dataclass
+from typing import Final
+
+import pytest
+
+from datadog_reader import DdLogsReader
+from datadog_reader import _DdAuthHeaders # pyright: ignore[reportPrivateUsage] # verifies private auth-header serialization
+from e2e_config import DD_SEARCH_INTERVAL, POLL_TIMEOUT
+from e2e_http import StreamingResponse
+
+
+def test_failure_diagnostics_hide_credentials_without_changing_auth_headers() -> None:
+ api_key: Final = "test-datadog-api-secret"
+ app_key: Final = "test-datadog-app-secret"
+ reader: Final = DdLogsReader(site="datadoghq.com", api_key=api_key, app_key=app_key)
+ headers: Final = _DdAuthHeaders(api_key=api_key, app_key=app_key)
+
+ for value in (reader, headers):
+ assert api_key not in repr(value)
+ assert app_key not in repr(value)
+
+ assert headers.model_dump(by_alias=True) == {
+ "DD-API-KEY": api_key,
+ "DD-APPLICATION-KEY": app_key,
+ }
+
+
+@dataclass
+class Clock:
+ elapsed: float = 0.0
+
+ def now(self) -> float:
+ return self.elapsed
+
+ def sleep(self, seconds: float) -> None:
+ self.elapsed += seconds
+
+
+@dataclass
+class Search:
+ responses: Iterator[StreamingResponse]
+ calls: tuple[tuple[str, float], ...] = ()
+
+ def __call__(self, query: str, timeout: float) -> StreamingResponse:
+ self.calls += ((query, timeout),)
+ return next(self.responses)
+
+
+def _page(*event_ids: str) -> StreamingResponse:
+ return StreamingResponse(
+ status_code=200,
+ body=json.dumps({"data": [{"attributes": {"attributes": {"id": event_id}}} for event_id in event_ids]}),
+ )
+
+
+def _reader(responses: Sequence[StreamingResponse], clock: Clock) -> tuple[DdLogsReader, Search]:
+ search: Final = Search(iter(responses))
+ return DdLogsReader(
+ site="us5.datadoghq.com",
+ api_key="test-api-secret",
+ app_key="test-app-secret",
+ search=search,
+ now=clock.now,
+ sleep=clock.sleep,
+ jitter=lambda: 0.25,
+ ), search
+
+
+def test_429_honors_server_reset_and_preserves_duplicate_events() -> None:
+ clock: Final = Clock()
+ reader, search = _reader(
+ (StreamingResponse(status_code=429, body="", headers={"x-ratelimit-reset": "6"}), _page("first", "duplicate")),
+ clock,
+ )
+
+ events: Final = reader.events_for_query("test-marker")
+
+ assert tuple(event.attributes["id"] for event in events) == ("first", "duplicate")
+ assert clock.elapsed == 6.25
+ assert search.calls == (("test-marker", 30.0), ("test-marker", 30.0))
+
+
+@pytest.mark.parametrize("reset", ("", "invalid", "nan", "inf", "-1"))
+def test_invalid_reset_uses_search_interval(reset: str) -> None:
+ clock: Final = Clock()
+ reader, _ = _reader(
+ (StreamingResponse(status_code=429, body="", headers={"x-ratelimit-reset": reset}), _page()), clock
+ )
+
+ assert reader.events_for_query("test-marker") == []
+ assert clock.elapsed == DD_SEARCH_INTERVAL + 0.25
+
+
+def test_zero_reset_cannot_create_a_busy_retry_loop() -> None:
+ clock: Final = Clock()
+ reader, _ = _reader(
+ (StreamingResponse(status_code=429, body="", headers={"x-ratelimit-reset": "0"}), _page()), clock
+ )
+
+ assert reader.events_for_query("test-marker") == []
+ assert clock.elapsed == 1.25
+
+
+def test_retry_after_is_not_shortened_by_an_earlier_reset() -> None:
+ clock: Final = Clock()
+ reader, _ = _reader(
+ (StreamingResponse(status_code=429, body="", headers={"x-ratelimit-reset": "2", "retry-after": "8"}), _page()),
+ clock,
+ )
+
+ assert reader.events_for_query("test-marker") == []
+ assert clock.elapsed == 8.25
+
+
+def test_rate_limit_wait_stops_at_deadline_without_issuing_another_request() -> None:
+ clock: Final = Clock()
+ reader, search = _reader(
+ (StreamingResponse(status_code=429, body="", headers={"x-ratelimit-reset": str(POLL_TIMEOUT * 10)}),), clock
+ )
+
+ with pytest.raises(pytest.fail.Exception, match="remained rate-limited"):
+ reader.events_for_query("test-marker")
+
+ assert clock.elapsed == POLL_TIMEOUT
+ assert search.calls == (("test-marker", 30.0),)
+
+
+def test_late_retry_cannot_receive_a_fresh_request_timeout() -> None:
+ clock: Final = Clock()
+ reader, search = _reader(
+ (StreamingResponse(status_code=429, body="", headers={"x-ratelimit-reset": str(POLL_TIMEOUT - 5)}), _page()),
+ clock,
+ )
+
+ assert reader.events_for_query("test-marker") == []
+ assert search.calls == (("test-marker", 30.0), ("test-marker", 4.75))
+
+
+@pytest.mark.parametrize("status", (-1, 401, 403, 500))
+def test_non_quota_failures_are_not_retried_or_treated_as_empty_results(status: int) -> None:
+ clock: Final = Clock()
+ reader, search = _reader((StreamingResponse(status_code=status, body=""), _page()), clock)
+
+ with pytest.raises(pytest.fail.Exception, match=f"failed with HTTP {status}"):
+ reader.events_for_query("test-marker")
+
+ assert search.calls == (("test-marker", 30.0),)
+ assert clock.elapsed == 0
+
+
+def test_polling_quota_retries_share_the_original_deadline() -> None:
+ clock: Final = Clock()
+ reader, search = _reader(
+ (_page(), StreamingResponse(status_code=429, body="", headers={"x-ratelimit-reset": str(POLL_TIMEOUT)})),
+ clock,
+ )
+
+ with pytest.raises(pytest.fail.Exception, match="remained rate-limited"):
+ reader.poll_events_for_query("test-marker")
+
+ assert clock.elapsed == POLL_TIMEOUT
+ assert len(search.calls) == 2
+
+
+def test_empty_polling_does_not_start_a_final_search_after_its_deadline() -> None:
+ clock: Final = Clock()
+ attempts: Final = int(POLL_TIMEOUT / DD_SEARCH_INTERVAL)
+ reader, search = _reader((_page(),) * attempts, clock)
+
+ assert reader.poll_events_for_query("test-marker") == []
+ assert clock.elapsed == POLL_TIMEOUT
+ assert len(search.calls) == attempts
+
+
+def test_settlement_quota_retries_keep_the_remaining_readback_budget() -> None:
+ clock: Final = Clock()
+ empty_reads: Final = int(POLL_TIMEOUT / DD_SEARCH_INTERVAL) - 2
+ reader, search = _reader(
+ (_page(),) * empty_reads
+ + (_page("first"), StreamingResponse(status_code=429, body="", headers={"x-ratelimit-reset": str(POLL_TIMEOUT)})),
+ clock,
+ )
+
+ with pytest.raises(pytest.fail.Exception, match="remained rate-limited"):
+ reader.poll_events_for_query("test-marker")
+
+ assert clock.elapsed == POLL_TIMEOUT
+ assert search.calls[-1] == ("test-marker", DD_SEARCH_INTERVAL)
+ assert len(search.calls) == empty_reads + 2
+
+
+def test_settlement_detects_a_duplicate_on_the_final_search() -> None:
+ clock: Final = Clock()
+ reader, _ = _reader((_page("first"), _page("first"), _page(), _page("first", "duplicate")), clock)
+
+ events: Final = reader.poll_events_for_query("test-marker")
+
+ assert tuple(event.attributes["id"] for event in events) == ("first", "duplicate")
+ assert clock.elapsed == 30
+
+
+def test_settlement_keeps_confirmed_events_through_empty_searches() -> None:
+ clock: Final = Clock()
+ reader, _ = _reader((_page("first"), _page(), _page(), _page()), clock)
+
+ events: Final = reader.poll_events_for_query("test-marker")
+
+ assert tuple(event.attributes["id"] for event in events) == ("first",)
+ assert clock.elapsed == 30
+
+
+def test_late_delivery_cannot_pass_without_a_complete_settle_window() -> None:
+ clock: Final = Clock()
+ empty_reads: Final = int(POLL_TIMEOUT / DD_SEARCH_INTERVAL) - 2
+ reader, search = _reader((_page(),) * empty_reads + (_page("first"), _page("first")), clock)
+
+ with pytest.raises(pytest.fail.Exception, match="duplicate-detection window"):
+ reader.poll_events_for_query("test-marker")
+
+ assert clock.elapsed == POLL_TIMEOUT
+ assert len(search.calls) == empty_reads + 2
diff --git a/tests/e2e/management/management_client.py b/tests/e2e/management/management_client.py
index 2b897f5f07f..1ef0d89a8f9 100644
--- a/tests/e2e/management/management_client.py
+++ b/tests/e2e/management/management_client.py
@@ -43,6 +43,9 @@ from models import (
KeyResetSpendBody,
KeyResetSpendResponse,
KeyUpdateBody,
+ McpServerCreateBody,
+ McpServerRow,
+ McpServerUpdateBody,
ModelDeleteBody,
OrgDeleteBody,
OrgInfoParams,
@@ -537,6 +540,38 @@ class ManagementClient:
).root
)
+ def create_mcp_server(self, body: McpServerCreateBody) -> McpServerRow:
+ return unwrap(
+ self.proxy.transport.post(
+ "/v1/mcp/server",
+ headers=self.proxy.transport.master,
+ json=body,
+ response_type=McpServerRow,
+ )
+ )
+
+ def update_mcp_server(self, body: McpServerUpdateBody) -> McpServerRow:
+ """PUT /v1/mcp/server, the call behind the dashboard's Save Changes: a partial
+ update where a field left unset keeps its stored value and None clears it."""
+ return unwrap(
+ self.proxy.transport.put(
+ "/v1/mcp/server",
+ headers=self.proxy.transport.master,
+ json=body,
+ response_type=McpServerRow,
+ )
+ )
+
+ def delete_mcp_server(self, server_id: str) -> Result[NoBody]:
+ """DELETE /v1/mcp/server/{server_id}. Returns the outcome so the act phase can
+ unwrap it while a deferred teardown can ignore an already-deleted server."""
+ return self.proxy.transport.delete(
+ f"/v1/mcp/server/{server_id}",
+ headers=self.proxy.transport.master,
+ json=NoBody(),
+ response_type=NoBody,
+ )
+
def chat_status(self, key: str, model: str, content: str) -> StreamingResponse:
return self.proxy.transport.send(
"/chat/completions",
diff --git a/tests/e2e/management/test_mcp_lifecycle_e2e.py b/tests/e2e/management/test_mcp_lifecycle_e2e.py
new file mode 100644
index 00000000000..9257d697647
--- /dev/null
+++ b/tests/e2e/management/test_mcp_lifecycle_e2e.py
@@ -0,0 +1,294 @@
+"""Live e2e: the MCP server and toolset management routes' lifecycle contract.
+
+Two customer defects sit on these routes, and each step here is the read-back that
+would have caught one of them: a dashboard edit that took several saves to stick
+because the read landed on a replica the write had not reached, and a toolset whose
+tools were stored under one name and read back under another, so it granted
+nothing. Every read-back therefore polls every replica that serves the route
+(ProxyClient.read_back_everywhere) and asserts the exact values written, and both
+update routes are held to the same partial-update contract: a field left out of the
+payload keeps its stored value, a field sent as null is cleared. The server URL is
+unreachable on purpose; only persistence is under test, never a tool call.
+"""
+
+from __future__ import annotations
+
+from collections.abc import Callable, Mapping
+from typing import Final
+
+import pytest
+from e2e_config import unique_marker
+from e2e_http import unwrap
+from lifecycle import ResourceManager
+from management_client import ManagementClient
+from models import (
+ McpInfo,
+ McpServerCreateBody,
+ McpServerListResponse,
+ McpServerRow,
+ McpServerUpdateBody,
+ ToolsetCreateBody,
+ ToolsetListResponse,
+ ToolsetRow,
+ ToolsetTool,
+ ToolsetUpdateBody,
+)
+
+pytestmark = pytest.mark.e2e
+
+UNREACHABLE_URL: Final = "https://e2e-fake-mcp.test.local/mcp"
+
+
+def _create_server(client: ManagementClient, resources: ResourceManager) -> tuple[McpServerCreateBody, str]:
+ name: Final = f"e2e_mcp_lifecycle_{unique_marker()}"
+ body: Final = McpServerCreateBody(
+ server_name=name,
+ alias=name,
+ url=UNREACHABLE_URL,
+ transport="http",
+ description="e2e lifecycle server",
+ mcp_info=McpInfo(
+ server_name=f"{name} (display)",
+ description="shown on the MCP page",
+ logo_url="https://e2e.test.local/logo.png",
+ ),
+ )
+ server_id: Final = client.create_mcp_server(body).server_id
+ resources.defer(lambda: client.delete_mcp_server(server_id))
+ return body, server_id
+
+
+def _assert_server_matches(row: McpServerRow, written: McpServerCreateBody, *, where: str) -> None:
+ stored: Final = (row.server_name, row.alias, row.url, row.transport, row.description, row.mcp_info)
+ expected: Final = (
+ written.server_name,
+ written.alias,
+ written.url,
+ written.transport,
+ written.description,
+ written.mcp_info,
+ )
+ assert stored == expected, f"{where}: stored {stored}, expected {expected}"
+
+
+def _server_everywhere(
+ client: ManagementClient, server_id: str, *, settled: Callable[[McpServerRow], bool]
+) -> Mapping[str, McpServerRow]:
+ return client.proxy.read_body_back_everywhere(f"/v1/mcp/server/{server_id}", McpServerRow, settled=settled)
+
+
+def _listed_server_everywhere(client: ManagementClient, server_id: str) -> Mapping[str, McpServerRow]:
+ listings: Final = client.proxy.read_body_back_everywhere(
+ "/v1/mcp/server",
+ McpServerListResponse,
+ settled=lambda rows: any(row.server_id == server_id for row in rows.root),
+ )
+ return {replica: next(row for row in rows.root if row.server_id == server_id) for replica, rows in listings.items()}
+
+
+class TestMcpServerLifecycle:
+ @pytest.mark.covers("mgmt.mcp_server.new.persists")
+ def test_create_persists_every_field_on_every_replica(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ body, server_id = _create_server(client, resources)
+
+ by_id: Final = _server_everywhere(client, server_id, settled=lambda row: row.server_id == server_id)
+ for replica, row in by_id.items():
+ _assert_server_matches(row, body, where=f"GET /v1/mcp/server/{server_id} on {replica}")
+
+ @pytest.mark.skip(
+ reason=(
+ "product gap: GET /v1/mcp/server builds each row from the in-memory registry, whose "
+ "_build_mcp_server_table sets description from mcp_info['description'], so the list "
+ "reports the mcp_info description while GET /v1/mcp/server/{server_id} reports the "
+ "stored description column. A server created with both set to different text reads "
+ "back with two different descriptions depending on the route"
+ )
+ )
+ @pytest.mark.covers("mgmt.mcp_server.list.persists")
+ def test_created_server_is_listed_with_every_field(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ body, server_id = _create_server(client, resources)
+
+ for replica, row in _listed_server_everywhere(client, server_id).items():
+ _assert_server_matches(row, body, where=f"GET /v1/mcp/server on {replica}")
+
+ @pytest.mark.covers("mgmt.mcp_server.update.preserves_unrelated_fields")
+ def test_updating_only_the_alias_keeps_every_other_field_on_every_replica(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ body, server_id = _create_server(client, resources)
+ renamed: Final = f"{body.alias}_renamed"
+
+ _ = client.update_mcp_server(McpServerUpdateBody(server_id=server_id, alias=renamed))
+
+ after_one_put: Final = _server_everywhere(client, server_id, settled=lambda row: row.alias == renamed)
+ for replica, row in after_one_put.items():
+ _assert_server_matches(
+ row,
+ body.model_copy(update={"alias": renamed}),
+ where=f"GET /v1/mcp/server/{server_id} on {replica} after one PUT of alias",
+ )
+
+ @pytest.mark.covers("mgmt.mcp_server.update.clear_persists")
+ def test_clearing_the_description_with_null_reads_back_null(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ body, server_id = _create_server(client, resources)
+
+ _ = client.update_mcp_server(McpServerUpdateBody(server_id=server_id, description=None))
+
+ cleared: Final = _server_everywhere(client, server_id, settled=lambda row: row.description is None)
+ for replica, row in cleared.items():
+ _assert_server_matches(
+ row,
+ body.model_copy(update={"description": None}),
+ where=f"GET /v1/mcp/server/{server_id} on {replica} after PUT description=null",
+ )
+
+ @pytest.mark.covers("mgmt.mcp_server.delete.persists")
+ def test_delete_removes_the_server_from_every_replica(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ _, server_id = _create_server(client, resources)
+
+ _ = unwrap(client.delete_mcp_server(server_id))
+
+ gone: Final = client.proxy.gone_everywhere(f"/v1/mcp/server/{server_id}")
+ assert set(gone.values()) == {404}, f"a deleted server must 404 on every replica; got {dict(gone)}"
+ listings: Final = client.proxy.read_body_back_everywhere(
+ "/v1/mcp/server",
+ McpServerListResponse,
+ settled=lambda rows: all(row.server_id != server_id for row in rows.root),
+ )
+ for replica, rows in listings.items():
+ assert all(row.server_id != server_id for row in rows.root), (
+ f"GET /v1/mcp/server on {replica} still lists the deleted server {server_id}"
+ )
+
+
+def _create_toolset(
+ client: ManagementClient, resources: ResourceManager, server_id: str
+) -> tuple[ToolsetCreateBody, str]:
+ body: Final = ToolsetCreateBody(
+ toolset_name=f"e2e_toolset_{unique_marker()}",
+ description="e2e lifecycle toolset",
+ tools=[
+ ToolsetTool(server_id=server_id, tool_name="search_datadog_logs"),
+ ToolsetTool(server_id=server_id, tool_name="get_datadog_metric"),
+ ],
+ )
+ toolset_id: Final = client.proxy.create_toolset(body).toolset_id
+ resources.defer(lambda: client.proxy.delete_toolset(toolset_id))
+ return body, toolset_id
+
+
+def _assert_toolset_matches(row: ToolsetRow, written: ToolsetCreateBody, *, where: str) -> None:
+ stored: Final = (row.toolset_name, row.description, row.tools)
+ expected: Final = (written.toolset_name, written.description, written.tools)
+ assert stored == expected, f"{where}: stored {stored}, expected {expected}"
+
+
+def _toolset_everywhere(
+ client: ManagementClient, toolset_id: str, *, settled: Callable[[ToolsetRow], bool]
+) -> Mapping[str, ToolsetRow]:
+ return client.proxy.read_body_back_everywhere(f"/v1/mcp/toolset/{toolset_id}", ToolsetRow, settled=settled)
+
+
+class TestMcpToolsetLifecycle:
+ @pytest.mark.covers("mgmt.mcp_toolset.new.persists")
+ def test_create_persists_both_tools_under_the_exact_names_written(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ _, server_id = _create_server(client, resources)
+ body, toolset_id = _create_toolset(client, resources, server_id)
+
+ by_id: Final = _toolset_everywhere(client, toolset_id, settled=lambda row: row.toolset_id == toolset_id)
+ for replica, row in by_id.items():
+ _assert_toolset_matches(row, body, where=f"GET /v1/mcp/toolset/{toolset_id} on {replica}")
+ listings: Final = client.proxy.read_body_back_everywhere(
+ "/v1/mcp/toolset",
+ ToolsetListResponse,
+ settled=lambda rows: any(row.toolset_id == toolset_id for row in rows.root),
+ )
+ for replica, rows in listings.items():
+ _assert_toolset_matches(
+ next(row for row in rows.root if row.toolset_id == toolset_id),
+ body,
+ where=f"GET /v1/mcp/toolset on {replica}",
+ )
+
+ @pytest.mark.covers("mgmt.mcp_toolset.update.preserves_unrelated_fields")
+ def test_updating_only_the_description_keeps_the_tools_and_name(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ _, server_id = _create_server(client, resources)
+ body, toolset_id = _create_toolset(client, resources, server_id)
+
+ _ = client.proxy.update_toolset(ToolsetUpdateBody(toolset_id=toolset_id, description="edited"))
+
+ edited: Final = _toolset_everywhere(client, toolset_id, settled=lambda row: row.description == "edited")
+ for replica, row in edited.items():
+ _assert_toolset_matches(
+ row,
+ body.model_copy(update={"description": "edited"}),
+ where=f"GET /v1/mcp/toolset/{toolset_id} on {replica} after PUT of description",
+ )
+
+ @pytest.mark.covers("mgmt.mcp_toolset.update.persists")
+ def test_updating_the_tools_to_one_entry_reads_back_exactly_that_entry(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ _, server_id = _create_server(client, resources)
+ body, toolset_id = _create_toolset(client, resources, server_id)
+ kept: Final = body.tools[:1]
+
+ _ = client.proxy.update_toolset(ToolsetUpdateBody(toolset_id=toolset_id, tools=kept))
+
+ narrowed: Final = _toolset_everywhere(client, toolset_id, settled=lambda row: row.tools == kept)
+ for replica, row in narrowed.items():
+ _assert_toolset_matches(
+ row,
+ body.model_copy(update={"tools": kept}),
+ where=f"GET /v1/mcp/toolset/{toolset_id} on {replica} after PUT of one tool",
+ )
+
+ @pytest.mark.covers("mgmt.mcp_toolset.update.clear_persists")
+ def test_clearing_the_description_with_null_reads_back_null(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ _, server_id = _create_server(client, resources)
+ body, toolset_id = _create_toolset(client, resources, server_id)
+
+ _ = client.proxy.update_toolset(ToolsetUpdateBody(toolset_id=toolset_id, description=None))
+
+ cleared: Final = _toolset_everywhere(client, toolset_id, settled=lambda row: row.description is None)
+ for replica, row in cleared.items():
+ _assert_toolset_matches(
+ row,
+ body.model_copy(update={"description": None}),
+ where=f"GET /v1/mcp/toolset/{toolset_id} on {replica} after PUT description=null",
+ )
+
+ @pytest.mark.covers("mgmt.mcp_toolset.delete.persists")
+ def test_delete_removes_the_toolset_from_every_replica(
+ self, client: ManagementClient, resources: ResourceManager
+ ) -> None:
+ _, server_id = _create_server(client, resources)
+ _, toolset_id = _create_toolset(client, resources, server_id)
+
+ _ = unwrap(client.proxy.delete_toolset(toolset_id))
+
+ gone: Final = client.proxy.gone_everywhere(f"/v1/mcp/toolset/{toolset_id}")
+ assert set(gone.values()) == {404}, f"a deleted toolset must 404 on every replica; got {dict(gone)}"
+ listings: Final = client.proxy.read_body_back_everywhere(
+ "/v1/mcp/toolset",
+ ToolsetListResponse,
+ settled=lambda rows: all(row.toolset_id != toolset_id for row in rows.root),
+ )
+ for replica, rows in listings.items():
+ assert all(row.toolset_id != toolset_id for row in rows.root), (
+ f"GET /v1/mcp/toolset on {replica} still lists the deleted toolset {toolset_id}"
+ )
diff --git a/tests/e2e/mcp/datadog_mcp.py b/tests/e2e/mcp/datadog_mcp.py
index d1ea53a0b3b..352b4446cfd 100644
--- a/tests/e2e/mcp/datadog_mcp.py
+++ b/tests/e2e/mcp/datadog_mcp.py
@@ -3,6 +3,7 @@
from __future__ import annotations
import os
+from collections.abc import Sequence
from e2e_config import datadog_mcp_url, unique_marker
from lifecycle import ResourceManager
@@ -35,7 +36,11 @@ def register_datadog_mcp(
resources: ResourceManager,
*,
mcp_access_groups: list[str] | None = None,
+ allowed_tools: Sequence[str] | None = (SEARCH_LOGS_TOOL,),
) -> str:
+ """Register the core Datadog toolset with its credentials from the env. By default
+ the server exposes only `search_datadog_logs`; pass `allowed_tools=None` to expose
+ every tool the core toolset serves."""
assert_dd_mcp_creds()
name = f"e2e_dd_mcp_{unique_marker()}"
server_id = client.register_server(
@@ -47,7 +52,7 @@ def register_datadog_mcp(
"DD-API-KEY": _dd_api_key(),
"DD-APPLICATION-KEY": _dd_app_key(),
},
- allowed_tools=[SEARCH_LOGS_TOOL],
+ allowed_tools=None if allowed_tools is None else list(allowed_tools),
mcp_access_groups=mcp_access_groups,
)
resources.defer(lambda: client.delete_server(server_id))
diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py
index 73453478e5a..210fc7a1e98 100644
--- a/tests/e2e/mcp/mcp_client.py
+++ b/tests/e2e/mcp/mcp_client.py
@@ -16,11 +16,11 @@ import time
from collections.abc import Mapping
from dataclasses import dataclass
-from pydantic import BaseModel, ConfigDict, Field, RootModel
+from pydantic import BaseModel, ConfigDict, Field
from e2e_config import settle_propagation
from e2e_http import Headers, NoBody, Result, Success, UnknownApiError, unwrap
-from models import KeyGenerateBody, ObjectPermission
+from models import KeyGenerateBody, McpServerListResponse, McpServerRow, ObjectPermission
from proxy_client import ProxyClient
McpToolArg = str | int | float | bool | list[str] | dict[str, str]
@@ -46,16 +46,6 @@ class McpServerNewResponse(BaseModel):
server_id: str
-class McpServerRow(BaseModel):
- server_id: str
- alias: str | None = None
- url: str | None = None
-
-
-class McpServersListResponse(RootModel[list[McpServerRow]]):
- pass
-
-
class McpToolMcpInfo(BaseModel):
server_id: str | None = None
alias: str | None = None
@@ -193,7 +183,7 @@ class McpClient:
"/v1/mcp/server",
headers=self.proxy.transport.master,
params=NoBody(),
- response_type=McpServersListResponse,
+ response_type=McpServerListResponse,
)
).root
@@ -224,11 +214,16 @@ class McpClient:
user_id: str,
mcp_servers: list[str] | None,
mcp_access_groups: list[str] | None = None,
+ mcp_toolsets: list[str] | None = None,
models: list[str] | None = None,
) -> str:
object_permission = (
- ObjectPermission(mcp_servers=mcp_servers, mcp_access_groups=mcp_access_groups)
- if mcp_servers is not None or mcp_access_groups is not None
+ ObjectPermission(
+ mcp_servers=mcp_servers,
+ mcp_access_groups=mcp_access_groups,
+ mcp_toolsets=mcp_toolsets,
+ )
+ if mcp_servers is not None or mcp_access_groups is not None or mcp_toolsets is not None
else None
)
return self.proxy.generate_key(
@@ -272,6 +267,20 @@ class McpClient:
)
time.sleep(self.proxy.poll_interval)
+ def await_tools(self, key: str, server_id: str, *, expected: frozenset[str]) -> frozenset[str]:
+ """Poll tools/list until `server_id`'s tools as `key` sees them are exactly
+ `expected`, and return the last listing either way, so the caller's equality
+ assertion names the difference. Fails at poll_timeout only when the read
+ itself never succeeded."""
+ deadline = time.monotonic() + self.proxy.poll_timeout
+ while True:
+ result = self.list_tools(key)
+ if isinstance(result, Success) and result.data.tool_names_for_server(server_id) == expected:
+ return expected
+ if time.monotonic() >= deadline:
+ return unwrap(result).tool_names_for_server(server_id)
+ time.sleep(self.proxy.poll_interval)
+
def await_call_tool(
self,
key: str,
diff --git a/tests/e2e/mcp/test_mcp_toolset_enforcement_e2e.py b/tests/e2e/mcp/test_mcp_toolset_enforcement_e2e.py
new file mode 100644
index 00000000000..6b901145eb1
--- /dev/null
+++ b/tests/e2e/mcp/test_mcp_toolset_enforcement_e2e.py
@@ -0,0 +1,95 @@
+"""Live e2e: a key granted a toolset lists exactly the toolset's tools.
+
+An admin registers the real Datadog remote MCP server with its whole core toolset
+exposed, discovers two of its tool names through a key granted the server outright,
+and curates a toolset naming exactly those two. A second key is granted the server
+plus that toolset, and its tools/list must come back as exactly those two names: no
+more, so the rest of the server's catalog stays hidden behind the toolset, and no
+fewer, so a tool stored under one name and read under another (which granted
+nothing) fails here first. Requires DD_API_KEY + DD_APP_KEY (the suite's real MCP
+upstream).
+"""
+
+from __future__ import annotations
+
+from typing import Final
+
+import pytest
+from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp
+from e2e_config import unique_marker
+from e2e_http import unwrap
+from lifecycle import ResourceManager
+from mcp_client import McpClient
+from models import ToolsetCreateBody, ToolsetTool
+
+pytestmark = pytest.mark.e2e
+
+
+def _key(
+ client: McpClient,
+ resources: ResourceManager,
+ label: str,
+ *,
+ server_id: str,
+ toolset_id: str | None = None,
+) -> str:
+ key: Final = client.generate_key(
+ user_id=f"e2e-mcp-{label}-{unique_marker()}",
+ mcp_servers=[server_id],
+ mcp_toolsets=None if toolset_id is None else [toolset_id],
+ )
+ resources.defer(lambda: client.proxy.delete_key(key))
+ return key
+
+
+def _wire_prefix(wire_name: str, tool_name: str, catalog: frozenset[str]) -> str:
+ """The prefix tools/list puts in front of one server's tool names, measured off a
+ tool whose own name is known rather than guessed from the alias. A toolset grants
+ by the tool's own name, never the wire name, and the prefix is whatever the proxy
+ is configured to build (the alias, or a short server id), so measuring it is the
+ only way to cross between the two."""
+ assert wire_name.endswith(tool_name), f"tools/list served {wire_name!r}, expected it to end with {tool_name!r}"
+ prefix: Final = wire_name[: len(wire_name) - len(tool_name)]
+ unprefixed: Final = frozenset(name for name in catalog if not name.startswith(prefix))
+ assert not unprefixed, (
+ f"every tool of one server shares the wire prefix {prefix!r}, so {sorted(unprefixed)} "
+ f"cannot be reduced to the names a toolset grants by"
+ )
+ return prefix
+
+
+class TestMcpToolsetEnforcement:
+ @pytest.mark.covers("mcp.list_tools.api_key.toolset_scoped")
+ def test_key_granted_a_toolset_lists_exactly_its_tools(self, client: McpClient, resources: ResourceManager) -> None:
+ server_id: Final = register_datadog_mcp(client, resources, allowed_tools=None)
+ client.await_registered(server_id)
+
+ catalog_key: Final = _key(client, resources, "catalog", server_id=server_id)
+ known_wire: Final = client.await_tool(catalog_key, server_id, SEARCH_LOGS_TOOL)
+ catalog: Final = unwrap(client.list_tools(catalog_key)).tool_names_for_server(server_id)
+ assert len(catalog) > 2, (
+ f"the Datadog core toolset must serve more tools than the toolset names, or the "
+ f"restriction has nothing to hide; got {sorted(catalog)}"
+ )
+ prefix: Final = _wire_prefix(known_wire, SEARCH_LOGS_TOOL, catalog)
+ chosen_wire: Final = frozenset(sorted(catalog)[:2])
+ chosen: Final = frozenset(name.removeprefix(prefix) for name in chosen_wire)
+
+ toolset: Final = client.proxy.create_toolset(
+ ToolsetCreateBody(
+ toolset_name=f"e2e_toolset_{unique_marker()}",
+ description="two Datadog tools",
+ tools=[ToolsetTool(server_id=server_id, tool_name=name) for name in sorted(chosen)],
+ )
+ )
+ resources.defer(lambda: client.proxy.delete_toolset(toolset.toolset_id))
+ assert frozenset(tool.tool_name for tool in toolset.tools) == chosen, (
+ f"toolset stored {toolset.tools}, expected the two names {sorted(chosen)} verbatim"
+ )
+
+ scoped_key: Final = _key(client, resources, "toolset", server_id=server_id, toolset_id=toolset.toolset_id)
+ listed: Final = client.await_tools(scoped_key, server_id, expected=chosen_wire)
+ assert listed == chosen_wire, (
+ f"a key granted the toolset must list exactly its two tools; "
+ f"got {sorted(listed)}, expected {sorted(chosen_wire)}"
+ )
diff --git a/tests/e2e/models.py b/tests/e2e/models.py
index 9f2654e0eec..62810e6cfd9 100644
--- a/tests/e2e/models.py
+++ b/tests/e2e/models.py
@@ -10,6 +10,7 @@ from collections.abc import Sequence
from datetime import datetime
from typing import Final, Literal
+from e2e_http import PartialBody
from pydantic import AliasChoices, BaseModel, ConfigDict, Field, RootModel, model_serializer, model_validator
# ---------- keys ----------
@@ -55,6 +56,7 @@ class KeyMetadata(BaseModel):
class ObjectPermission(BaseModel):
mcp_servers: list[str] | None = None
mcp_access_groups: list[str] | None = None
+ mcp_toolsets: list[str] | None = None
class KeyGenerateBody(BaseModel):
@@ -77,7 +79,7 @@ class KeyGenerateBody(BaseModel):
allowed_passthrough_routes: list[str] | None = None
metadata: KeyMetadata | None = None
object_permission: ObjectPermission | None = None
- router_settings: "RouterSettingsOverride | None" = None
+ router_settings: RouterSettingsOverride | None = None
class KeyGenerateResponse(BaseModel):
@@ -516,6 +518,15 @@ class CountTokensResponse(BaseModel):
# ---------- mcp servers ----------
+class McpInfo(BaseModel):
+ """The `mcp_info` display block stored on an MCP server; only the fields the
+ lifecycle test writes and reads back."""
+
+ server_name: str | None = None
+ description: str | None = None
+ logo_url: str | None = None
+
+
class McpServerCreateBody(BaseModel):
"""POST /v1/mcp/server. For a gateway-managed OAuth server, `auth_type` is
`oauth2` and `oauth2_flow` is `authorization_code`; the upstream endpoints
@@ -530,6 +541,18 @@ class McpServerCreateBody(BaseModel):
oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None
authorization_url: str | None = None
token_url: str | None = None
+ server_name: str | None = None
+ description: str | None = None
+ mcp_info: McpInfo | None = None
+
+
+class McpServerUpdateBody(PartialBody):
+ """PUT /v1/mcp/server: a field left unset keeps its stored value, a field set
+ to None is cleared."""
+
+ server_id: str
+ alias: str | None = None
+ description: str | None = None
class McpServerInfo(BaseModel):
@@ -543,6 +566,54 @@ class McpServerInfo(BaseModel):
allow_all_keys: bool | None = None
+class McpServerRow(McpServerInfo):
+ """A stored MCP server as the create, get, and list routes return it: the
+ fields the lifecycle test asserts survive the round trip."""
+
+ server_name: str | None = None
+ transport: str | None = None
+ description: str | None = None
+ mcp_info: McpInfo | None = None
+
+
+class McpServerListResponse(RootModel[list[McpServerRow]]):
+ """GET /v1/mcp/server answers with a bare array of servers."""
+
+
+class ToolsetTool(BaseModel):
+ server_id: str
+ tool_name: str
+
+
+class ToolsetCreateBody(BaseModel):
+ toolset_name: str
+ description: str | None = None
+ tools: list[ToolsetTool]
+
+
+class ToolsetUpdateBody(PartialBody):
+ """PUT /v1/mcp/toolset: a field left unset keeps its stored value, a field set
+ to None is cleared."""
+
+ toolset_id: str
+ description: str | None = None
+ tools: list[ToolsetTool] | None = None
+
+
+class ToolsetRow(BaseModel):
+ """A stored toolset as POST /v1/mcp/toolset, GET /v1/mcp/toolset/{toolset_id},
+ and each row of GET /v1/mcp/toolset return it."""
+
+ toolset_id: str
+ toolset_name: str
+ description: str | None = None
+ tools: list[ToolsetTool] = Field(default_factory=list)
+
+
+class ToolsetListResponse(RootModel[list[ToolsetRow]]):
+ """GET /v1/mcp/toolset answers with a bare array of toolsets."""
+
+
class EmbedBody(BaseModel):
model: str
input: str
diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py
index 520cbfde5a9..1bac5116a9d 100644
--- a/tests/e2e/proxy_client.py
+++ b/tests/e2e/proxy_client.py
@@ -12,6 +12,7 @@ import time
import warnings
from collections.abc import Callable, Mapping
from dataclasses import dataclass
+from functools import reduce
from datetime import datetime
from types import MappingProxyType
from typing import Final
@@ -26,6 +27,7 @@ from e2e_http import (
Result,
StreamingResponse,
Success,
+ UnknownApiError,
is_ok,
unwrap,
)
@@ -70,6 +72,9 @@ from models import (
SpendLogsPage,
SpendLogsPageParams,
SpendLogsParams,
+ ToolsetCreateBody,
+ ToolsetRow,
+ ToolsetUpdateBody,
)
from e2e_config import (
CONTROL_PLANE_BASE_URL,
@@ -82,7 +87,7 @@ from e2e_config import (
SLOW_PROVIDER_TIMEOUT_SECONDS,
settle_propagation,
)
-from transport import HttpTransport, SplitTransport, Transport
+from transport import HttpTransport, SplitTransport, Transport, is_control_plane_path
RowsPredicate = Callable[[list[SpendLogRow]], bool]
@@ -235,6 +240,99 @@ def servable_timeout_message(
)
+type ReplicaRead[T] = Callable[[float], T]
+
+
+@dataclass(frozen=True, slots=True)
+class EverywhereConverged[T]:
+ """Every replica answered with something `settled` accepts, keyed by replica."""
+
+ answers: Mapping[str, T]
+
+
+@dataclass(frozen=True, slots=True)
+class NeverConvergedOn[T]:
+ """`replica` ran out its budget without an answer `settled` accepts; `last` is
+ its final answer, so the failure can say what that replica still serves."""
+
+ replica: str
+ last: T
+
+
+def _last_answer[T](
+ read: ReplicaRead[T],
+ *,
+ settled: Callable[[T], bool],
+ timeout: float,
+ interval: float,
+ request_timeout: float,
+ now: Callable[[], float],
+ sleep: Callable[[float], None],
+) -> T:
+ """Poll `read` until `settled` accepts its answer or `timeout` runs out, and
+ return the last answer either way. Each read's request timeout is clamped to
+ the budget left, and the final poll runs even when less than an interval
+ remains, so a deadline never skips the read that would have settled."""
+ deadline: Final = now() + timeout
+ answer = read(min(request_timeout, timeout))
+ while not settled(answer):
+ remaining = deadline - now()
+ if remaining <= 0:
+ return answer
+ sleep(min(interval, remaining))
+ answer = read(min(request_timeout, remaining))
+ return answer
+
+
+def await_everywhere[T](
+ reads: Mapping[str, ReplicaRead[T]],
+ *,
+ settled: Callable[[T], bool],
+ timeout: float,
+ interval: float,
+ request_timeout: float,
+ now: Callable[[], float],
+ sleep: Callable[[float], None],
+) -> EverywhereConverged[T] | NeverConvergedOn[T]:
+ """`_last_answer` against every replica in turn, each with the full budget, so a
+ write counts as visible only once the last replica reflects it, and stop at the
+ first replica that never converges. Clock and sleep are injected."""
+ def read_replica(
+ outcome: EverywhereConverged[T] | NeverConvergedOn[T],
+ item: tuple[str, ReplicaRead[T]],
+ ) -> EverywhereConverged[T] | NeverConvergedOn[T]:
+ if isinstance(outcome, NeverConvergedOn):
+ return outcome
+ replica, read = item
+ answer: Final = _last_answer(
+ read,
+ settled=settled,
+ timeout=timeout,
+ interval=interval,
+ request_timeout=request_timeout,
+ now=now,
+ sleep=sleep,
+ )
+ if not settled(answer):
+ return NeverConvergedOn(replica=replica, last=answer)
+ return EverywhereConverged(answers=MappingProxyType({**outcome.answers, replica: answer}))
+
+ initial: Final[EverywhereConverged[T] | NeverConvergedOn[T]] = EverywhereConverged(answers=MappingProxyType({}))
+ return reduce(read_replica, reads.items(), initial)
+
+
+def _is_not_found[R: BaseModel](result: Result[R]) -> bool:
+ return isinstance(result, UnknownApiError) and result.status_code == 404
+
+
+def _status_of[R: BaseModel](result: Result[R]) -> int:
+ match result:
+ case Success(status_code=status_code) | UnknownApiError(status_code=status_code):
+ return status_code
+ case _:
+ return -1
+
+
type Poller[T] = Callable[[], T]
@@ -321,6 +419,7 @@ def converge_timeout_message(*, what: str, replica: str, timeout: float, last_re
class ProxyClient:
transport: Transport
replicas: Mapping[str, Transport]
+ control_replicas: Mapping[str, Transport]
poll_timeout: float = 120.0
poll_interval: float = 5.0
model_servable_timeout: float = MODEL_SERVABLE_TIMEOUT
@@ -569,6 +668,112 @@ class ProxyClient:
if not is_ok(result):
warnings.warn(f"delete_model({model_id!r}) failed: {result}", stacklevel=2)
+ # ---- replica read-back ----------------------------------------------
+
+ def replicas_for(self, path: str) -> Mapping[str, Transport]:
+ """The replicas that serve `path`: every data-plane replica for an LLM route,
+ and for a management route the control-plane replicas, since the data-plane
+ replicas trim management routes and answer them 404. A monolith serves both
+ from every replica, so a management read-back polls all of them; a split
+ deployment exposes one control-plane address (there is one backend process
+ behind it on the stack these suites run against), so it polls that. A
+ control plane fronting several backends would need its own replica list to
+ prove each one converged, the way PROXY_REPLICA_URLS does for the gateways.
+ Never empty: a read-back against no replica would assert nothing and pass."""
+ replicas: Final = self.control_replicas if is_control_plane_path(path) else self.replicas
+ assert replicas, f"no replica is configured to serve {path}, so a read-back there would prove nothing"
+ return replicas
+
+ def read_body_back_everywhere[R: BaseModel](
+ self, path: str, response_type: type[R], *, settled: Callable[[R], bool]
+ ) -> Mapping[str, R]:
+ """GET `path` on every replica that serves it, polling each to poll_timeout
+ until `settled` accepts its body, and fail naming the first replica that
+ never converged. Returns each replica's settled body, keyed by replica, so
+ the caller can assert the rest of it."""
+ outcome: Final = await_everywhere(
+ {url: self._reader(transport, path, response_type) for url, transport in self.replicas_for(path).items()},
+ settled=lambda result: isinstance(result, Success) and settled(result.data),
+ timeout=self.poll_timeout,
+ interval=self.poll_interval,
+ request_timeout=REQUEST_TIMEOUT,
+ now=time.monotonic,
+ sleep=time.sleep,
+ )
+ match outcome:
+ case EverywhereConverged(answers=answers):
+ return MappingProxyType({url: unwrap(result) for url, result in answers.items()})
+ case NeverConvergedOn(replica=replica, last=last):
+ raise AssertionError(
+ f"GET {path} on {replica} never converged within {self.poll_timeout}s of the write; "
+ f"last read: {last}"
+ )
+
+ def gone_everywhere(self, path: str) -> Mapping[str, int]:
+ """Poll GET `path` on every replica that serves it until each stops serving
+ it, and fail naming the first replica that still does at poll_timeout.
+ Returns each replica's final status, so the caller asserts the 404 itself."""
+ outcome: Final = await_everywhere(
+ {url: self._reader(transport, path, NoBody) for url, transport in self.replicas_for(path).items()},
+ settled=_is_not_found,
+ timeout=self.poll_timeout,
+ interval=self.poll_interval,
+ request_timeout=REQUEST_TIMEOUT,
+ now=time.monotonic,
+ sleep=time.sleep,
+ )
+ match outcome:
+ case EverywhereConverged(answers=answers):
+ return MappingProxyType({url: _status_of(result) for url, result in answers.items()})
+ case NeverConvergedOn(replica=replica, last=last):
+ raise AssertionError(
+ f"GET {path} on {replica} still answers {self.poll_timeout}s after the delete; last read: {last}"
+ )
+
+ @staticmethod
+ def _reader[R: BaseModel](transport: Transport, path: str, response_type: type[R]) -> ReplicaRead[Result[R]]:
+ return lambda request_timeout: transport.get(
+ path,
+ headers=transport.master,
+ params=NoBody(),
+ response_type=response_type,
+ timeout=request_timeout,
+ )
+
+ # ---- mcp toolsets ---------------------------------------------------
+
+ def create_toolset(self, body: ToolsetCreateBody) -> ToolsetRow:
+ return unwrap(
+ self.transport.post(
+ "/v1/mcp/toolset",
+ headers=self.transport.master,
+ json=body,
+ response_type=ToolsetRow,
+ )
+ )
+
+ def update_toolset(self, body: ToolsetUpdateBody) -> ToolsetRow:
+ """PUT /v1/mcp/toolset: a partial update where a field left unset keeps its
+ stored value and None clears it."""
+ return unwrap(
+ self.transport.put(
+ "/v1/mcp/toolset",
+ headers=self.transport.master,
+ json=body,
+ response_type=ToolsetRow,
+ )
+ )
+
+ def delete_toolset(self, toolset_id: str) -> Result[NoBody]:
+ """DELETE /v1/mcp/toolset/{toolset_id}. Returns the outcome so the act phase
+ can unwrap it while a deferred teardown can ignore an already-deleted row."""
+ return self.transport.delete(
+ f"/v1/mcp/toolset/{toolset_id}",
+ headers=self.transport.master,
+ json=NoBody(),
+ response_type=NoBody,
+ )
+
def create_credential(self, body: CredentialCreateBody) -> None:
unwrap(
self.transport.post(
@@ -736,7 +941,10 @@ def build_proxy_client(
base URLs are the same for a monolithic proxy, so routing is then a no-op.
``replica_urls`` (PROXY_REPLICA_URLS) names every data-plane replica the model
barrier polls directly; it is the data-plane URL itself unless the stack
- exports each gateway's own address.
+ exports each gateway's own address. Management read-backs poll those same
+ replicas when the two planes share a base URL (a monolith, where every replica
+ serves every route) and the control plane alone when they differ (a split
+ deployment, where the data-plane replicas do not serve management routes).
The endpoints are injectable for callers that resolve the proxy some other
way than ``e2e_config``'s env names (see ``claude_code/_env.py``); they must
@@ -764,9 +972,13 @@ def build_proxy_client(
for url in replica_urls
}
)
+ control_replicas: Final = (
+ replicas if control_plane_base_url == base_url else MappingProxyType({control_plane_base_url: split.control})
+ )
return ProxyClient(
transport=split,
replicas=replicas,
+ control_replicas=control_replicas,
poll_timeout=POLL_TIMEOUT,
poll_interval=POLL_INTERVAL,
)
diff --git a/tests/e2e/router/test_auto_router_regressions_e2e.py b/tests/e2e/router/test_auto_router_regressions_e2e.py
index 188db2a8eb5..374badcf5fc 100644
--- a/tests/e2e/router/test_auto_router_regressions_e2e.py
+++ b/tests/e2e/router/test_auto_router_regressions_e2e.py
@@ -41,6 +41,7 @@ which stores either the registered alias or the provider-prefixed form.
import json
import os
from collections.abc import Iterator
+from contextlib import ExitStack
from dataclasses import dataclass
from typing import Final
@@ -120,19 +121,10 @@ class ResponsesApiResponse(BaseModel):
@dataclass(frozen=True, slots=True)
-class TagSplitDeployments:
- """Scenario A mirrors the customer-shaped config from GitHub issue #36619:
- plain deployment registered first, tier deployment and marker both tagged.
- Scenario B flips both axes for GitHub issue #36621: marker registered first
- and its tier deployment left untagged, so routing depends neither on
- registration order nor on tier deployments carrying tags."""
-
- tag_a: str
- shared_a: str
- tier_a: str
- tag_b: str
- shared_b: str
- tier_b: str
+class TagSplitDeployment:
+ tag: str
+ shared: str
+ tier: str
@dataclass(frozen=True, slots=True)
@@ -173,9 +165,7 @@ def _uniform_tier_config(tier_model: str) -> dict[str, object]:
}
-def _key_for(
- proxy: ProxyClient, resources: ResourceManager, models: list[str], tag_filtering: bool = False
-) -> str:
+def _key_for(proxy: ProxyClient, resources: ResourceManager, models: list[str], tag_filtering: bool = False) -> str:
key: Final = proxy.generate_key(
KeyGenerateBody(
models=models,
@@ -211,46 +201,61 @@ def _assert_served_only_by(rows: list[SpendLogRow], allowed: frozenset[str], con
)
-@pytest.fixture(scope="module")
-def split(proxy: ProxyClient) -> Iterator[TagSplitDeployments]:
+@pytest.fixture(scope="class")
+def router_stack() -> Iterator[ExitStack]:
+ with ExitStack() as stack:
+ yield stack
+
+
+def _register_models(
+ proxy: ProxyClient, stack: ExitStack, registrations: tuple[tuple[str, LiteLLMParamsBody], ...]
+) -> None:
+ for name, params in registrations:
+ stack.callback(proxy.delete_model, proxy.create_model(name, params))
+
+
+def _tag_split(proxy: ProxyClient, stack: ExitStack, *, marker_first: bool) -> TagSplitDeployment:
marker: Final = unique_marker()
- deployments: Final = TagSplitDeployments(
- tag_a=f"e2e-split-a-{marker}",
- shared_a=f"e2e-autoroute-a-{marker}",
- tier_a=f"e2e-tier-a-{marker}",
- tag_b=f"e2e-split-b-{marker}",
- shared_b=f"e2e-autoroute-b-{marker}",
- tier_b=f"e2e-tier-b-{marker}",
+ named: Final = TagSplitDeployment(
+ tag=f"e2e-split-{marker}",
+ shared=f"e2e-autoroute-{marker}",
+ tier=f"e2e-tier-{marker}",
)
anthropic_key: Final = _provider_key("ANTHROPIC_API_KEY")
- marker_params_a: Final = LiteLLMParamsBody(
- model="auto_router/complexity_router",
- complexity_router_config=_uniform_tier_config(deployments.tier_a),
- tags=[deployments.tag_a],
+ marker_registration: Final = (
+ named.shared,
+ LiteLLMParamsBody(
+ model="auto_router/complexity_router",
+ complexity_router_config=_uniform_tier_config(named.tier),
+ tags=[named.tag],
+ ),
)
- marker_params_b: Final = LiteLLMParamsBody(
- model="auto_router/complexity_router",
- complexity_router_config=_uniform_tier_config(deployments.tier_b),
- tags=[deployments.tag_b],
+ tier_registration: Final = (
+ named.tier,
+ LiteLLMParamsBody(model=CHEAP_MODEL, api_key=anthropic_key, tags=None if marker_first else [named.tag]),
)
- registrations: Final[tuple[tuple[str, LiteLLMParamsBody], ...]] = (
- (deployments.shared_a, LiteLLMParamsBody(model=PLAIN_MODEL, api_key=anthropic_key)),
- (deployments.tier_a, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=anthropic_key, tags=[deployments.tag_a])),
- (deployments.shared_a, marker_params_a),
- (deployments.shared_b, marker_params_b),
- (deployments.tier_b, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=anthropic_key)),
- (deployments.shared_b, LiteLLMParamsBody(model=PLAIN_MODEL, api_key=anthropic_key)),
+ plain_registration: Final = (named.shared, LiteLLMParamsBody(model=PLAIN_MODEL, api_key=anthropic_key))
+ registrations: Final = (
+ (marker_registration, tier_registration, plain_registration)
+ if marker_first
+ else (plain_registration, tier_registration, marker_registration)
)
- created: Final = tuple(proxy.create_model(name, params) for name, params in registrations)
- try:
- yield deployments
- finally:
- for model_id in created:
- proxy.delete_model(model_id)
+ _register_models(proxy, stack, registrations)
+ return named
-@pytest.fixture(scope="module")
-def zero_priced_alias(proxy: ProxyClient) -> Iterator[ZeroPricedAlias]:
+@pytest.fixture(scope="class")
+def plain_first_split(proxy: ProxyClient, router_stack: ExitStack) -> TagSplitDeployment:
+ return _tag_split(proxy, router_stack, marker_first=False)
+
+
+@pytest.fixture(scope="class")
+def marker_first_split(proxy: ProxyClient, router_stack: ExitStack) -> TagSplitDeployment:
+ return _tag_split(proxy, router_stack, marker_first=True)
+
+
+@pytest.fixture(scope="class")
+def zero_priced_alias(proxy: ProxyClient, router_stack: ExitStack) -> ZeroPricedAlias:
marker: Final = unique_marker()
named: Final = ZeroPricedAlias(alias=f"e2e-priced-alias-{marker}", tier=f"e2e-priced-tier-{marker}")
alias_params: Final = LiteLLMParamsBody(
@@ -263,16 +268,12 @@ def zero_priced_alias(proxy: ProxyClient) -> Iterator[ZeroPricedAlias]:
(named.tier, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=_provider_key("ANTHROPIC_API_KEY"))),
(named.alias, alias_params),
)
- created: Final = tuple(proxy.create_model(name, params) for name, params in registrations)
- try:
- yield named
- finally:
- for model_id in created:
- proxy.delete_model(model_id)
+ _register_models(proxy, router_stack, registrations)
+ return named
-@pytest.fixture(scope="module")
-def heuristic_split(proxy: ProxyClient) -> Iterator[HeuristicSplit]:
+@pytest.fixture(scope="class")
+def heuristic_split(proxy: ProxyClient, router_stack: ExitStack) -> HeuristicSplit:
marker: Final = unique_marker()
named: Final = HeuristicSplit(
alias=f"e2e-heuristic-router-{marker}",
@@ -289,16 +290,12 @@ def heuristic_split(proxy: ProxyClient) -> Iterator[HeuristicSplit]:
(named.strong, LiteLLMParamsBody(model=STRONG_MODEL, api_key=_provider_key("OPENAI_API_KEY"))),
(named.alias, LiteLLMParamsBody(model="auto_router/complexity_router", complexity_router_config=config)),
)
- created: Final = tuple(proxy.create_model(name, params) for name, params in registrations)
- try:
- yield named
- finally:
- for model_id in created:
- proxy.delete_model(model_id)
+ _register_models(proxy, router_stack, registrations)
+ return named
-@pytest.fixture(scope="module")
-def semantic_auto_router(proxy: ProxyClient) -> Iterator[SemanticAutoRouter]:
+@pytest.fixture(scope="class")
+def semantic_auto_router(proxy: ProxyClient, router_stack: ExitStack) -> SemanticAutoRouter:
marker: Final = unique_marker()
named: Final = SemanticAutoRouter(
marker=f"e2e-semantic-router-{marker}",
@@ -321,16 +318,12 @@ def semantic_auto_router(proxy: ProxyClient) -> Iterator[SemanticAutoRouter]:
(named.fallback, LiteLLMParamsBody(model=PLAIN_MODEL, api_key=_provider_key("ANTHROPIC_API_KEY"))),
(named.marker, marker_params),
)
- created: Final = tuple(proxy.create_model(name, params) for name, params in registrations)
- try:
- yield named
- finally:
- for model_id in created:
- proxy.delete_model(model_id)
+ _register_models(proxy, router_stack, registrations)
+ return named
-@pytest.fixture(scope="module")
-def credentialed_alias(proxy: ProxyClient) -> Iterator[CredentialedAlias]:
+@pytest.fixture(scope="class")
+def credentialed_alias(proxy: ProxyClient, router_stack: ExitStack) -> CredentialedAlias:
marker: Final = unique_marker()
named: Final = CredentialedAlias(alias=f"e2e-cred-alias-{marker}", tier=f"e2e-cred-tier-{marker}")
alias_params: Final = LiteLLMParamsBody(
@@ -342,104 +335,110 @@ def credentialed_alias(proxy: ProxyClient) -> Iterator[CredentialedAlias]:
(named.tier, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=_provider_key("ANTHROPIC_API_KEY"))),
(named.alias, alias_params),
)
- created: Final = tuple(proxy.create_model(name, params) for name, params in registrations)
- try:
- yield named
- finally:
- for model_id in created:
- proxy.delete_model(model_id)
+ _register_models(proxy, router_stack, registrations)
+ return named
class TestTagSplitRouting:
@pytest.mark.covers("reliability.routing.tagged_marker.request_tag_selects_marker")
def test_body_tagged_chat_routes_through_the_marker_to_its_tier(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment
) -> None:
"""Pins GitHub issue #36619: with tag filtering on, a chat request whose
body metadata tags match the tagged marker under a shared model name is
answered by the marker's tier deployment, not by the plain deployment
that was registered under the name first."""
- key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True)
- chat: Final = unwrap(proxy.chat(key, _hello_chat_body(split.shared_a, tags=[split.tag_a])))
+ key: Final = _key_for(proxy, resources, [plain_first_split.shared, plain_first_split.tier], tag_filtering=True)
+ chat: Final = unwrap(proxy.chat(key, _hello_chat_body(plain_first_split.shared, tags=[plain_first_split.tag])))
assert chat.choices, "tagged chat through the shared name returned no choices"
rows: Final = proxy.poll_logs_for_key(key, min_rows=1)
- _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_a}, "body-tagged chat on the shared name")
+ _assert_served_only_by(rows, CHEAP_SERVED | {plain_first_split.tier}, "body-tagged chat on the shared name")
@pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment")
def test_untagged_chat_is_always_served_by_the_plain_deployment(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment
) -> None:
"""Pins GitHub issue #36620: untagged chat requests to the shared name
succeed on every call and are all served by the plain deployment; the
tagged marker never captures them, so no intermittent auto-router
errors and no tier hijacking."""
- key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True)
+ key: Final = _key_for(proxy, resources, [plain_first_split.shared, plain_first_split.tier], tag_filtering=True)
for _ in range(5):
- chat = unwrap(proxy.chat(key, _hello_chat_body(split.shared_a)))
+ chat = unwrap(proxy.chat(key, _hello_chat_body(plain_first_split.shared)))
assert chat.choices, "untagged chat through the shared name returned no choices"
rows: Final = proxy.poll_logs_for_key(key, min_rows=5)
- _assert_served_only_by(rows, PLAIN_SERVED | {split.shared_a}, "untagged chat on the shared name")
+ _assert_served_only_by(rows, PLAIN_SERVED | {plain_first_split.shared}, "untagged chat on the shared name")
@pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment")
def test_untagged_messages_is_served_by_the_plain_deployment(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment
) -> None:
"""Pins GitHub issue #36620 on the /v1/messages surface: an untagged
Anthropic-native request to the shared name is served by the plain
deployment, not captured by the tagged marker."""
- key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True)
- answer: Final = unwrap(proxy.messages(key, _hello_messages_body(split.shared_a)))
+ key: Final = _key_for(proxy, resources, [plain_first_split.shared, plain_first_split.tier], tag_filtering=True)
+ answer: Final = unwrap(proxy.messages(key, _hello_messages_body(plain_first_split.shared)))
assert answer.content or answer.choices, "untagged /v1/messages returned neither content nor choices"
rows: Final = proxy.poll_logs_for_key(key, min_rows=1)
- _assert_served_only_by(rows, PLAIN_SERVED | {split.shared_a}, "untagged /v1/messages on the shared name")
+ _assert_served_only_by(
+ rows, PLAIN_SERVED | {plain_first_split.shared}, "untagged /v1/messages on the shared name"
+ )
class TestUntaggedTierDeployments:
@pytest.mark.covers("reliability.routing.tagged_marker.header_tag_selects_marker")
def test_header_tagged_messages_routes_through_the_marker_to_an_untagged_tier(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, marker_first_split: TagSplitDeployment
) -> None:
"""Pins GitHub issue #36621: a /v1/messages request tagged only via the
x-litellm-tags header selects the tagged marker, and the rewrite still
lands on the tier deployment even though that deployment carries no
tags, because the marker consumed the routing tags."""
- key: Final = _key_for(proxy, resources, [split.shared_b, split.tier_b], tag_filtering=True)
- headers: Final = TaggedAnthropicHeaders(authorization=f"Bearer {key}", x_litellm_tags=split.tag_b)
+ key: Final = _key_for(
+ proxy, resources, [marker_first_split.shared, marker_first_split.tier], tag_filtering=True
+ )
+ headers: Final = TaggedAnthropicHeaders(authorization=f"Bearer {key}", x_litellm_tags=marker_first_split.tag)
answer: Final = unwrap(
proxy.transport.post(
"/v1/messages",
headers=headers,
- json=_hello_messages_body(split.shared_b),
+ json=_hello_messages_body(marker_first_split.shared),
response_type=AnthropicMessagesResponse,
)
)
assert answer.content or answer.choices, "header-tagged /v1/messages returned neither content nor choices"
rows: Final = proxy.poll_logs_for_key(key, min_rows=1)
- _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_b}, "header-tagged /v1/messages on the shared name")
+ _assert_served_only_by(
+ rows, CHEAP_SERVED | {marker_first_split.tier}, "header-tagged /v1/messages on the shared name"
+ )
@pytest.mark.covers("reliability.routing.tagged_marker.untagged_tier_deployments_still_served")
def test_body_tagged_chat_reaches_the_untagged_tier_after_marker_rewrite(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, marker_first_split: TagSplitDeployment
) -> None:
"""Pins the tag-consumption half of GitHub issue #36621: after the
tagged marker rewrites the request to its tier model, the consumed
routing tags no longer constrain deployment selection, so the untagged
tier deployment serves the request instead of a strict-tag denial."""
- key: Final = _key_for(proxy, resources, [split.shared_b, split.tier_b], tag_filtering=True)
- chat: Final = unwrap(proxy.chat(key, _hello_chat_body(split.shared_b, tags=[split.tag_b])))
+ key: Final = _key_for(
+ proxy, resources, [marker_first_split.shared, marker_first_split.tier], tag_filtering=True
+ )
+ chat: Final = unwrap(
+ proxy.chat(key, _hello_chat_body(marker_first_split.shared, tags=[marker_first_split.tag]))
+ )
assert chat.choices, "body-tagged chat through the marker-first shared name returned no choices"
rows: Final = proxy.poll_logs_for_key(key, min_rows=1)
- _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_b}, "body-tagged chat with untagged tier")
+ _assert_served_only_by(rows, CHEAP_SERVED | {marker_first_split.tier}, "body-tagged chat with untagged tier")
@pytest.mark.covers("reliability.routing.tagged_marker.tag_semantics_stay_strict")
def test_tagged_call_straight_at_an_untagged_deployment_stays_denied(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, marker_first_split: TagSplitDeployment
) -> None:
"""The tag-consumption fix must not loosen strict tag semantics: a
tagged request aimed directly at an untagged deployment (no marker
involved) is still rejected with the 401 tags-configuration error."""
- key: Final = _key_for(proxy, resources, [split.tier_b], tag_filtering=True)
- result: Final = proxy.chat(key, _hello_chat_body(split.tier_b, tags=[split.tag_b]))
+ key: Final = _key_for(proxy, resources, [marker_first_split.tier], tag_filtering=True)
+ result: Final = proxy.chat(key, _hello_chat_body(marker_first_split.tier, tags=[marker_first_split.tag]))
assert isinstance(result, UnauthorizedError), (
f"expected the tagged direct call to an untagged deployment to be denied with 401, got {result}"
)
@@ -451,37 +450,39 @@ class TestUntaggedTierDeployments:
class TestResponsesApiTagRouting:
@pytest.mark.covers("reliability.routing.tagged_marker.responses_input_routes_through_marker")
def test_header_tagged_responses_with_string_input_routes_to_the_tier(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment
) -> None:
"""Pins the /v1/responses surface of the tag split (GitHub issues
#36620/#36621): a /v1/responses request with string input, tagged via
the x-litellm-tags header, succeeds and routes through the tagged
marker to its tier."""
- key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True)
- headers: Final = TaggedAuthHeaders(authorization=f"Bearer {key}", x_litellm_tags=split.tag_a)
+ key: Final = _key_for(proxy, resources, [plain_first_split.shared, plain_first_split.tier], tag_filtering=True)
+ headers: Final = TaggedAuthHeaders(authorization=f"Bearer {key}", x_litellm_tags=plain_first_split.tag)
body: Final = ResponsesBody(
- model=split.shared_a, input=f"say hello {unique_marker()}", max_output_tokens=64
+ model=plain_first_split.shared, input=f"say hello {unique_marker()}", max_output_tokens=64
)
answer: Final = unwrap(
proxy.transport.post("/v1/responses", headers=headers, json=body, response_type=ResponsesApiResponse)
)
assert answer.id, "header-tagged /v1/responses returned no response id"
rows: Final = proxy.poll_logs_for_key(key, min_rows=1)
- _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_a}, "header-tagged /v1/responses string input")
+ _assert_served_only_by(
+ rows, CHEAP_SERVED | {plain_first_split.tier}, "header-tagged /v1/responses string input"
+ )
@pytest.mark.covers("reliability.routing.tagged_marker.responses_input_routes_through_marker")
def test_body_tagged_responses_with_list_input_routes_to_the_tier(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment
) -> None:
"""Pins the body-tag and list-input combination of the same split:
/v1/responses with litellm_metadata.tags and structured input items
routes through the tagged marker to its tier."""
- key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True)
+ key: Final = _key_for(proxy, resources, [plain_first_split.shared, plain_first_split.tier], tag_filtering=True)
body: Final = ResponsesBody(
- model=split.shared_a,
+ model=plain_first_split.shared,
input=[ResponsesInputItem(role="user", content=f"say hello {unique_marker()}")],
max_output_tokens=64,
- litellm_metadata=ResponsesTagMetadata(tags=[split.tag_a]),
+ litellm_metadata=ResponsesTagMetadata(tags=[plain_first_split.tag]),
)
answer: Final = unwrap(
proxy.transport.post(
@@ -493,18 +494,18 @@ class TestResponsesApiTagRouting:
)
assert answer.id, "body-tagged /v1/responses returned no response id"
rows: Final = proxy.poll_logs_for_key(key, min_rows=1)
- _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_a}, "body-tagged /v1/responses list input")
+ _assert_served_only_by(rows, CHEAP_SERVED | {plain_first_split.tier}, "body-tagged /v1/responses list input")
@pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment")
def test_untagged_responses_is_served_by_the_plain_deployment(
- self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments
+ self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment
) -> None:
"""Pins the untagged half of the /v1/responses tag split: an untagged
request to the shared name is served by the plain deployment, matching
the chat and messages surfaces."""
- key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True)
+ key: Final = _key_for(proxy, resources, [plain_first_split.shared, plain_first_split.tier], tag_filtering=True)
body: Final = ResponsesBody(
- model=split.shared_a, input=f"say hello {unique_marker()}", max_output_tokens=64
+ model=plain_first_split.shared, input=f"say hello {unique_marker()}", max_output_tokens=64
)
answer: Final = unwrap(
proxy.transport.post(
@@ -516,7 +517,9 @@ class TestResponsesApiTagRouting:
)
assert answer.id, "untagged /v1/responses returned no response id"
rows: Final = proxy.poll_logs_for_key(key, min_rows=1)
- _assert_served_only_by(rows, PLAIN_SERVED | {split.shared_a}, "untagged /v1/responses on the shared name")
+ _assert_served_only_by(
+ rows, PLAIN_SERVED | {plain_first_split.shared}, "untagged /v1/responses on the shared name"
+ )
class TestStrategyAliasPricing:
@@ -551,9 +554,7 @@ class TestComplexityHeuristicScope:
while the accompanying ~2KB agent system prompt is packed with enough
reasoning and complexity keywords that scoring the combined text lands
in REASONING; only ask-only scoring keeps this on the cheap tier."""
- key: Final = _key_for(
- proxy, resources, [heuristic_split.alias, heuristic_split.cheap, heuristic_split.strong]
- )
+ key: Final = _key_for(proxy, resources, [heuristic_split.alias, heuristic_split.cheap, heuristic_split.strong])
body: Final = ChatBody(
model=heuristic_split.alias,
messages=[
diff --git a/tests/e2e/test_e2e_http.py b/tests/e2e/test_e2e_http.py
index 66841725d1d..81cd6c8d3d1 100644
--- a/tests/e2e/test_e2e_http.py
+++ b/tests/e2e/test_e2e_http.py
@@ -13,13 +13,24 @@ monkeypatches anything.
from __future__ import annotations
from collections.abc import Callable, Iterator, Mapping, Sequence
-from dataclasses import dataclass, field
+from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
import pytest
-
-from e2e_http import RETRY_ATTEMPTS, TRANSIENT_STATUSES, request_with_retry, streaming_outcome
+from e2e_http import (
+ RETRY_ATTEMPTS,
+ TRANSIENT_STATUSES,
+ NoBody,
+ PartialBody,
+ Success,
+ ValidationError,
+ classify,
+ request_with_retry,
+ streaming_outcome,
+ wire_body,
+)
+from pydantic import BaseModel, TypeAdapter
@dataclass
@@ -33,10 +44,10 @@ class FakeResponse:
@dataclass
class SleepRecorder:
- delays: list[float] = field(default_factory=list)
+ delays: tuple[float, ...] = ()
def __call__(self, seconds: float) -> None:
- self.delays.append(seconds)
+ self.delays += (seconds,)
def _issue_from(responses: Sequence[FakeResponse]) -> Callable[[], FakeResponse]:
@@ -55,7 +66,7 @@ class TestTransientRetryPolicy:
sleep = SleepRecorder()
result = request_with_retry(_issue_from(responses), sleep=sleep)
assert result is responses[0]
- assert sleep.delays == []
+ assert sleep.delays == ()
assert responses[0].close_calls == 0
def test_429_is_never_retried(self) -> None:
@@ -63,7 +74,7 @@ class TestTransientRetryPolicy:
sleep = SleepRecorder()
result = request_with_retry(_issue_from(responses), sleep=sleep)
assert result is responses[0]
- assert sleep.delays == []
+ assert sleep.delays == ()
assert responses[0].close_calls == 0
def test_overloaded_529_retries_with_backoff_then_returns_the_success(self) -> None:
@@ -71,7 +82,7 @@ class TestTransientRetryPolicy:
sleep = SleepRecorder()
result = request_with_retry(_issue_from(responses), sleep=sleep)
assert result is responses[1]
- assert sleep.delays == [0.5]
+ assert sleep.delays == (0.5,)
assert responses[0].close_calls == 1
assert responses[1].close_calls == 0
@@ -80,7 +91,7 @@ class TestTransientRetryPolicy:
sleep = SleepRecorder()
result = request_with_retry(_issue_from(responses), sleep=sleep)
assert result is responses[RETRY_ATTEMPTS - 1]
- assert sleep.delays == [0.5, 1.0]
+ assert sleep.delays == (0.5, 1.0)
assert [r.close_calls for r in responses] == [1, 1, 0, 0]
@@ -134,3 +145,65 @@ class TestStreamEventArrivals:
assert result.stream_events == []
assert result.stream_event_arrivals == []
assert result.body == "bad request"
+
+
+class _ServerUpdate(PartialBody):
+ server_id: str
+ alias: str | None = None
+ description: str | None = None
+
+
+class _ServerCreate(BaseModel):
+ alias: str
+ description: str | None = None
+
+
+class TestWireBody:
+ """A partial-update body must put exactly the caller's choice on the wire: an
+ omitted field stays off it so the route keeps the stored value, and an explicit
+ None goes out as JSON null so the route clears it. Plain bodies keep dropping
+ None, which is what every create route expects."""
+
+ def test_partial_body_omits_unset_fields_and_sends_explicit_none_as_null(self) -> None:
+ assert wire_body(_ServerUpdate(server_id="s1", description=None)) == {"server_id": "s1", "description": None}
+ assert wire_body(_ServerUpdate(server_id="s1", alias="renamed")) == {"server_id": "s1", "alias": "renamed"}
+
+ def test_plain_body_drops_none_fields(self) -> None:
+ assert wire_body(_ServerCreate(alias="a", description=None)) == {"alias": "a"}
+
+
+_JSON: Final[TypeAdapter[object]] = TypeAdapter(object)
+
+
+@dataclass
+class FakeJsonResponse:
+ """The `classify` view of a response: a status, the raw body bytes, and the
+ parse that would raise on an empty one."""
+
+ status_code: int
+ content: bytes
+
+ @property
+ def ok(self) -> bool:
+ return self.status_code < 400
+
+ @property
+ def text(self) -> str:
+ return self.content.decode()
+
+ def json(self) -> object:
+ return _JSON.validate_json(self.content)
+
+
+class TestClassifyEmptyBody:
+ """A delete that answers 202 with no body is a success, not a parse failure:
+ the MCP server and toolset delete routes both answer that way, and reading it
+ as a failure would hide a delete that did not happen behind one that did."""
+
+ def test_empty_2xx_body_is_a_success(self) -> None:
+ result: Final = classify(FakeJsonResponse(status_code=202, content=b""), NoBody)
+ assert isinstance(result, Success) and result.status_code == 202
+
+ def test_body_that_is_not_json_is_still_a_validation_failure(self) -> None:
+ result: Final = classify(FakeJsonResponse(status_code=200, content=b"
"), NoBody)
+ assert isinstance(result, ValidationError)
diff --git a/tests/e2e/test_proxy_client.py b/tests/e2e/test_proxy_client.py
index 2caac58333f..3b84a47e3cc 100644
--- a/tests/e2e/test_proxy_client.py
+++ b/tests/e2e/test_proxy_client.py
@@ -15,28 +15,35 @@ from collections.abc import Iterable, Mapping
from dataclasses import dataclass
from itertools import chain, repeat
from types import MappingProxyType
-from typing import Final
+from typing import Final, cast
import pytest
-
from e2e_config import parse_replica_urls
from e2e_http import Result, Success
from models import KeyInfo, KeyInfoResponse, ModelListEntry, ModelsListResponse
from proxy_client import (
- Poller,
ConvergeOutcome,
Converged,
+ EverywhereConverged,
ModelsPoller,
+ NeverConvergedOn,
NotConverged,
NotServableOn,
+ Poller,
+ ProxyClient,
+ ReplicaRead,
Servable,
await_converged_everywhere,
+ await_everywhere,
await_servable_everywhere,
- first_lagging_replica,
+ build_proxy_client,
converge_timeout_message,
+ first_lagging_replica,
)
+from transport import Transport
MODEL: Final = "gpt-under-test"
+_NO_TRANSPORTS: Final = cast(Transport, None)
TIMEOUT: Final = 10.0
INTERVAL: Final = 2.0
RPM_BEFORE_UPDATE: Final = 100
@@ -187,3 +194,83 @@ class TestParseReplicaUrls:
def test_falls_back_to_the_data_plane_address_when_unset(self) -> None:
assert parse_replica_urls("", "http://lb") == ("http://lb",)
+
+
+def _answers(answers: Iterable[str]) -> ReplicaRead[str]:
+ it: Final = iter(answers)
+ return lambda _timeout: next(it)
+
+
+def _await_everywhere(reads: Mapping[str, ReplicaRead[str]]) -> EverywhereConverged[str] | NeverConvergedOn[str]:
+ clock: Final = FakeClock()
+ return await_everywhere(
+ reads,
+ settled=lambda answer: answer == "renamed",
+ timeout=TIMEOUT,
+ interval=INTERVAL,
+ request_timeout=5.0,
+ now=clock.now,
+ sleep=clock.sleep,
+ )
+
+
+class TestAwaitEverywhere:
+ def test_waits_for_the_lagging_replica_and_returns_every_settled_answer(self) -> None:
+ reads: Final = {
+ "gateway-1": _answers(repeat("renamed")),
+ "gateway-2": _answers(chain(repeat("stale", 2), repeat("renamed"))),
+ }
+ outcome: Final = _await_everywhere(reads)
+ assert isinstance(outcome, EverywhereConverged)
+ assert dict(outcome.answers) == {"gateway-1": "renamed", "gateway-2": "renamed"}
+
+ def test_names_the_replica_that_never_converges_with_what_it_last_served(self) -> None:
+ reads: Final = {
+ "gateway-1": _answers(repeat("renamed")),
+ "gateway-2": _answers(repeat("stale")),
+ }
+ assert _await_everywhere(reads) == NeverConvergedOn(replica="gateway-2", last="stale")
+
+ def test_polls_until_the_deadline_before_giving_up(self) -> None:
+ lagging: Final = chain(repeat("stale", int(TIMEOUT / INTERVAL)), repeat("renamed"))
+ outcome: Final = _await_everywhere({"gateway-1": _answers(lagging)})
+ assert isinstance(outcome, EverywhereConverged), outcome
+
+
+class TestReplicasFor:
+ def test_split_deployment_reads_management_routes_back_from_the_control_plane(self) -> None:
+ client: Final = build_proxy_client(
+ base_url="http://lb",
+ control_plane_base_url="http://backend",
+ replica_urls=("http://gateway-1", "http://gateway-2"),
+ )
+ assert set(client.replicas_for("/key/info")) == {"http://backend"}
+ assert set(client.replicas_for("/v1/models")) == {"http://gateway-1", "http://gateway-2"}
+
+ def test_monolith_reads_management_routes_back_from_every_replica(self) -> None:
+ client: Final = build_proxy_client(
+ base_url="http://lb",
+ control_plane_base_url="http://lb",
+ replica_urls=("http://pod-1", "http://pod-2"),
+ )
+ assert set(client.replicas_for("/key/info")) == {"http://pod-1", "http://pod-2"}
+
+ def test_mcp_admin_routes_read_back_from_every_data_plane_replica(self) -> None:
+ """/v1/mcp/* is a lazily mounted feature, so a data-plane replica serves it
+ too and answers from its own in-memory registry. Routing it to the control
+ plane would leave every replica but that one unproven, and would move the
+ tools/list barrier in mcp_client off the plane that serves tools/list."""
+ client: Final = build_proxy_client(
+ base_url="http://lb",
+ control_plane_base_url="http://backend",
+ replica_urls=("http://gateway-1", "http://gateway-2"),
+ )
+ assert set(client.replicas_for("/v1/mcp/server/abc")) == {"http://gateway-1", "http://gateway-2"}
+ assert set(client.replicas_for("/v1/mcp/toolset/abc")) == {"http://gateway-1", "http://gateway-2"}
+
+ def test_a_route_no_replica_serves_is_refused_rather_than_read_back_vacuously(self) -> None:
+ """A read-back over zero replicas would satisfy every predicate and assert
+ nothing, so asking for one fails instead of passing silently."""
+ client: Final = ProxyClient(transport=_NO_TRANSPORTS, replicas={}, control_replicas={})
+ with pytest.raises(AssertionError, match="no replica is configured"):
+ _ = client.replicas_for("/v1/models")
diff --git a/tests/e2e/ui/helpers/navigation.ts b/tests/e2e/ui/helpers/navigation.ts
index 4a7c4e7baa9..e0e7b4da396 100644
--- a/tests/e2e/ui/helpers/navigation.ts
+++ b/tests/e2e/ui/helpers/navigation.ts
@@ -73,3 +73,13 @@ export async function clickTeamId(page: PlaywrightPage, teamId: string): Promise
await cell.click();
await expect(page.getByText("Back to Teams")).toBeVisible({ timeout: 10_000 });
}
+
+export async function openKeyDetail(page: PlaywrightPage, alias: string): Promise {
+ await page.getByPlaceholder("Search by key alias or ID").fill(alias);
+ const row = page.getByRole("row").filter({ hasText: alias });
+ await expect(row, `key row "${alias}" never appeared on the Virtual Keys page`).toBeVisible({ timeout: 15_000 });
+ await row.getByRole("button", { name: alias }).click();
+ await expect(page.getByText("Back to Keys"), `key detail for "${alias}" never opened`).toBeVisible({
+ timeout: 15_000,
+ });
+}
diff --git a/tests/e2e/ui/helpers/traffic.ts b/tests/e2e/ui/helpers/traffic.ts
index 7f8417cdffb..cb68747b364 100644
--- a/tests/e2e/ui/helpers/traffic.ts
+++ b/tests/e2e/ui/helpers/traffic.ts
@@ -1,4 +1,4 @@
-import { APIRequestContext, expect } from "@playwright/test";
+import { APIRequestContext, APIResponse, expect } from "@playwright/test";
/** Model names served by fixtures/config.yml, both backed by the mock LLM server. */
export const CHAT_MODEL_A = "fake-openai-gpt-4";
@@ -15,6 +15,9 @@ export const masterKey = (): string => process.env.LITELLM_MASTER_KEY || "sk-123
export const rootPath = (): string => process.env.SERVER_ROOT_PATH ?? "";
+/** Date.now() alone collides: `--repeat-each` starts its copies inside the same millisecond. */
+export const uniqueSuffix = (): string => `${Date.now()}-${Math.random().toString(36).slice(2, 8)}`;
+
interface ChatOptions {
model: string;
prompt: string;
@@ -25,9 +28,8 @@ interface ChatOptions {
traceId?: string;
}
-/** POST /v1/chat/completions and return the completion id (the Logs Request ID). */
-export async function sendChatCompletion(request: APIRequestContext, opts: ChatOptions): Promise {
- const res = await request.post(`${rootPath()}/v1/chat/completions`, {
+const postChatCompletion = (request: APIRequestContext, opts: ChatOptions): Promise =>
+ request.post(`${rootPath()}/v1/chat/completions`, {
headers: {
Authorization: `Bearer ${opts.apiKey ?? masterKey()}`,
"Content-Type": "application/json",
@@ -39,12 +41,26 @@ export async function sendChatCompletion(request: APIRequestContext, opts: ChatO
...(opts.traceId ? { litellm_trace_id: opts.traceId } : {}),
},
});
+
+/** POST /v1/chat/completions and return the completion id (the Logs Request ID). */
+export async function sendChatCompletion(request: APIRequestContext, opts: ChatOptions): Promise {
+ const res = await postChatCompletion(request, opts);
expect(res.ok(), `chat completion for ${opts.model} failed (${res.status()}): ${await res.text()}`).toBe(true);
const body = await res.json();
expect(body.choices?.[0]?.message?.content).toContain(MOCK_RESPONSE_TEXT);
return body.id as string;
}
+export interface ChatAttempt {
+ status: number;
+ body: string;
+}
+
+export async function attemptChatCompletion(request: APIRequestContext, opts: ChatOptions): Promise {
+ const res = await postChatCompletion(request, opts);
+ return { status: res.status(), body: await res.text() };
+}
+
/** `key` is the sk- value to authenticate with; `token` is its hash, which spend aggregates are keyed by. */
export async function createVirtualKey(
request: APIRequestContext,
@@ -66,6 +82,33 @@ export async function createVirtualKey(
};
}
+export interface KeyInfo {
+ key_alias: string | null;
+ max_budget: number | null;
+ budget_duration: string | null;
+ budget_reset_at: string | null;
+ blocked: boolean | null;
+ models: string[];
+ team_id: string | null;
+}
+
+export async function readKeyInfo(request: APIRequestContext, token: string): Promise {
+ const res = await request.get(`${rootPath()}/key/info?key=${encodeURIComponent(token)}`, {
+ headers: { Authorization: `Bearer ${masterKey()}` },
+ });
+ expect(res.ok(), `GET /key/info for ${token} failed (${res.status()}): ${await res.text()}`).toBe(true);
+ const body = await res.json();
+ return body.info as KeyInfo;
+}
+
+export async function deleteVirtualKey(request: APIRequestContext, token: string): Promise {
+ const res = await request.post(`${rootPath()}/key/delete`, {
+ headers: { Authorization: `Bearer ${masterKey()}`, "Content-Type": "application/json" },
+ data: { keys: [token] },
+ });
+ expect(res.ok(), `key delete for ${token} failed (${res.status()}): ${await res.text()}`).toBe(true);
+}
+
/** Spend logs are flushed on a timer, so an assertion straight after a completion races the writer. */
export async function waitForSpendLog(
request: APIRequestContext,
diff --git a/tests/e2e/ui/tests/internal-user/internalUserKeyScope.spec.ts b/tests/e2e/ui/tests/internal-user/internalUserKeyScope.spec.ts
new file mode 100644
index 00000000000..f923841257a
--- /dev/null
+++ b/tests/e2e/ui/tests/internal-user/internalUserKeyScope.spec.ts
@@ -0,0 +1,208 @@
+import { test, expect, type APIRequestContext } from "@playwright/test";
+import { Page } from "../../fixtures/pages";
+import {
+ dismissFeedbackPopup,
+ navigateToPage,
+ openKeyDetail,
+} from "../../helpers/navigation";
+import {
+ CHAT_MODEL_A,
+ CHAT_MODEL_B,
+ MOCK_RESPONSE_TEXT,
+ attemptChatCompletion,
+ createVirtualKey,
+ deleteVirtualKey,
+ masterKey,
+ readKeyInfo,
+ rootPath,
+ uniqueSuffix,
+} from "../../helpers/traffic";
+
+const MEMBER_PASSWORD = "E2e-Team-Member-Pass-1!";
+
+interface CreatedTeam {
+ readonly team_id: string;
+}
+
+function assertCreatedTeam(body: unknown): asserts body is CreatedTeam {
+ expect(body, "/team/new returned no team_id").toMatchObject({
+ team_id: expect.any(String),
+ });
+}
+
+async function postAsMaster(
+ request: APIRequestContext,
+ path: string,
+ data: Record,
+): Promise {
+ const res = await request.post(`${rootPath()}${path}`, {
+ headers: {
+ Authorization: `Bearer ${masterKey()}`,
+ "Content-Type": "application/json",
+ },
+ data,
+ });
+ expect(
+ res.ok(),
+ `POST ${path} failed (${res.status()}): ${await res.text()}`,
+ ).toBe(true);
+ return res.json();
+}
+
+test.describe("Internal User - own team key model scope", () => {
+ test.use({ storageState: { cookies: [], origins: [] } });
+
+ test("a team member narrows their own key's models and the proxy enforces it", async ({
+ page,
+ request,
+ }) => {
+ const suffix = uniqueSuffix();
+ const email = `team-member-${suffix}@test.local`;
+ const userId = `e2e-key-scope-user-${suffix}`;
+ const alias = `e2e-key-scope-${suffix}`;
+
+ const team = await postAsMaster(request, "/team/new", {
+ team_alias: `E2E Key Scope ${suffix}`,
+ models: [CHAT_MODEL_A, CHAT_MODEL_B],
+ team_member_permissions: ["/key/generate", "/key/update", "/key/info"],
+ });
+ assertCreatedTeam(team);
+ const teamId = team.team_id;
+
+ try {
+ await postAsMaster(request, "/user/new", {
+ user_id: userId,
+ user_email: email,
+ user_role: "internal_user",
+ auto_create_key: false,
+ });
+ await postAsMaster(request, "/user/update", {
+ user_id: userId,
+ password: MEMBER_PASSWORD,
+ });
+ await postAsMaster(request, "/team/member_add", {
+ team_id: teamId,
+ member: { role: "user", user_id: userId },
+ });
+
+ const created = await createVirtualKey(request, {
+ key_alias: alias,
+ team_id: teamId,
+ user_id: userId,
+ models: [],
+ });
+
+ try {
+ await page.goto("/ui/login");
+ await page.getByPlaceholder("Enter your username").fill(email);
+ await page
+ .getByPlaceholder("Enter your password")
+ .fill(MEMBER_PASSWORD);
+ await page.getByRole("button", { name: "Login", exact: true }).click();
+ await expect(
+ page.locator("a", { hasText: "Virtual Keys" }),
+ `${email} never reached the dashboard`,
+ ).toBeVisible({ timeout: 30_000 });
+ await dismissFeedbackPopup(page);
+
+ await navigateToPage(page, Page.ApiKeys);
+ await openKeyDetail(page, alias);
+
+ await page.getByRole("tab", { name: "Settings" }).click();
+ await page.getByRole("button", { name: "Edit Settings" }).click();
+
+ await page.getByRole("combobox", { name: "Select models" }).click();
+ await expect(
+ page.getByRole("option", { name: CHAT_MODEL_A, exact: true }),
+ `the Models dropdown does not offer ${CHAT_MODEL_A} to a team member`,
+ ).toBeVisible({ timeout: 15_000 });
+ await expect(
+ page.getByRole("option", { name: CHAT_MODEL_B, exact: true }),
+ `the Models dropdown does not offer ${CHAT_MODEL_B} to a team member`,
+ ).toBeVisible();
+
+ await page
+ .getByRole("option", { name: CHAT_MODEL_A, exact: true })
+ .click();
+ await page.keyboard.press("Escape");
+
+ const updated = page.waitForResponse(
+ (res) =>
+ res.url().includes("/key/update") &&
+ res.request().method() === "POST",
+ );
+ await page.getByRole("button", { name: "Save Changes" }).click();
+ const updateStatus = (await updated).status();
+ expect(
+ updateStatus,
+ "a team member's own-key edit was refused",
+ ).toBeGreaterThanOrEqual(200);
+ expect(
+ updateStatus,
+ "a team member's own-key edit was refused",
+ ).toBeLessThan(300);
+ await expect(
+ page.getByText("Key updated successfully").first(),
+ ).toBeVisible({ timeout: 15_000 });
+
+ await expect
+ .poll(
+ async () => (await readKeyInfo(request, created.token)).models,
+ {
+ message: `the narrowed model scope never reached /key/info for ${alias}`,
+ timeout: 20_000,
+ },
+ )
+ .toEqual([CHAT_MODEL_A]);
+
+ await expect
+ .poll(
+ async () =>
+ await attemptChatCompletion(request, {
+ model: CHAT_MODEL_B,
+ prompt: `out of scope ${suffix}`,
+ apiKey: created.key,
+ }),
+ {
+ message: `${CHAT_MODEL_B} was still served after the key was narrowed to ${CHAT_MODEL_A}`,
+ timeout: 30_000,
+ },
+ )
+ .toMatchObject({
+ status: 403,
+ body: expect.stringContaining(CHAT_MODEL_B),
+ });
+
+ const inScope = await attemptChatCompletion(request, {
+ model: CHAT_MODEL_A,
+ prompt: `in scope ${suffix}`,
+ apiKey: created.key,
+ });
+ expect(
+ inScope,
+ `${CHAT_MODEL_A} is no longer served by the narrowed key`,
+ ).toMatchObject({
+ status: 200,
+ body: expect.stringContaining(MOCK_RESPONSE_TEXT),
+ });
+ } finally {
+ await deleteVirtualKey(request, created.token);
+ }
+ } finally {
+ await request.post(`${rootPath()}/user/delete`, {
+ headers: {
+ Authorization: `Bearer ${masterKey()}`,
+ "Content-Type": "application/json",
+ },
+ data: { user_ids: [userId] },
+ });
+ await request.post(`${rootPath()}/team/delete`, {
+ headers: {
+ Authorization: `Bearer ${masterKey()}`,
+ "Content-Type": "application/json",
+ },
+ data: { team_ids: [teamId] },
+ });
+ }
+ });
+});
diff --git a/tests/e2e/ui/tests/internal-user/modelsByTeam.spec.ts b/tests/e2e/ui/tests/internal-user/modelsByTeam.spec.ts
new file mode 100644
index 00000000000..5e2c80b5845
--- /dev/null
+++ b/tests/e2e/ui/tests/internal-user/modelsByTeam.spec.ts
@@ -0,0 +1,196 @@
+import {
+ test as base,
+ expect,
+ type Locator,
+ type Page as PlaywrightPage,
+} from "@playwright/test";
+import {
+ E2E_TEAM_CRUD_ALIAS,
+ E2E_TEAM_ORG_ALIAS,
+ INTERNAL_USER_STORAGE_PATH,
+} from "../../constants";
+import { Page } from "../../fixtures/pages";
+import { navigateToPage } from "../../helpers/navigation";
+import { readBack } from "../../helpers/roundTrip";
+import { CHAT_MODEL_A, CHAT_MODEL_B, masterKey } from "../../helpers/traffic";
+
+const MOCK_LLM_BASE = `http://127.0.0.1:${process.env.MOCK_LLM_PORT ?? "8090"}/v1`;
+const CURRENT_TEAM_VIEW = "Current Team Models";
+const ALL_MODELS_VIEW = "All Available Models";
+const PERSONAL_TEAM = "Personal";
+
+const teamSelector = (page: PlaywrightPage): Locator =>
+ page.getByRole("combobox", { name: "Current team", exact: true });
+const viewSelector = (page: PlaywrightPage): Locator =>
+ page.getByRole("combobox", { name: "View", exact: true });
+
+async function chooseOption(
+ page: PlaywrightPage,
+ selector: Locator,
+ optionName: string,
+): Promise {
+ await selector.click();
+ const option = page.getByRole("option", { name: optionName, exact: true });
+ await expect(option, `option ${optionName} is offered`).toBeVisible({
+ timeout: 10_000,
+ });
+ await option.click();
+ await expect(
+ selector,
+ `${optionName} is the selection the control now reports`,
+ ).toContainText(optionName, {
+ timeout: 10_000,
+ });
+}
+
+async function deleteDeployment(
+ page: PlaywrightPage,
+ id: string,
+): Promise {
+ const post = () =>
+ page.request.post("/model/delete", {
+ headers: { Authorization: `Bearer ${masterKey()}` },
+ data: { id },
+ });
+ const deleted = await post().catch(() => post());
+ expect(
+ deleted.ok(),
+ `cleanup: /model/delete ${id} returned ${deleted.status()}`,
+ ).toBe(true);
+}
+
+function modelRow(page: PlaywrightPage, modelName: string): Locator {
+ return page.getByRole("row").filter({ hasText: modelName });
+}
+
+async function isRegistered(
+ page: PlaywrightPage,
+ modelName: string,
+): Promise {
+ const body = await readBack<{ data: { model_name?: string }[] }>(
+ page,
+ "/v2/model/info",
+ );
+ return body.data.some((row) => row.model_name === modelName);
+}
+
+const uniqueSuffix = (): string =>
+ `${Date.now()}-${Math.random().toString(36).slice(2, 8)}`;
+
+const test = base.extend<{ ungrantedModelName: string }>({
+ ungrantedModelName: async ({ page }, use) => {
+ const ungrantedModelName = `e2e-ungranted-${uniqueSuffix()}`;
+ const created = await page.request.post("/model/new", {
+ headers: { Authorization: `Bearer ${masterKey()}` },
+ data: {
+ model_name: ungrantedModelName,
+ litellm_params: {
+ model: `openai/${ungrantedModelName}`,
+ api_base: MOCK_LLM_BASE,
+ api_key: "fake-key",
+ },
+ model_info: {},
+ },
+ });
+ expect(
+ created.ok(),
+ `/model/new failed: ${created.status()} ${await created.text()}`,
+ ).toBe(true);
+ const ungrantedModelId = (await created.json()).model_info?.id;
+ expect(ungrantedModelId, "model id from /model/new").toBeTruthy();
+
+ try {
+ await expect
+ .poll(async () => await isRegistered(page, ungrantedModelName), {
+ message: `deployment ${ungrantedModelName} never appeared in /v2/model/info after create`,
+ timeout: 60_000,
+ })
+ .toBe(true);
+ await use(ungrantedModelName);
+ } finally {
+ await deleteDeployment(page, ungrantedModelId);
+ }
+ },
+});
+
+test.describe("Models and Endpoints for an internal user", () => {
+ test.use({ storageState: INTERNAL_USER_STORAGE_PATH });
+
+ test("shows an internal user exactly the models of the team they select", async ({
+ page,
+ ungrantedModelName,
+ }) => {
+ await navigateToPage(page, Page.Models);
+
+ await expect(
+ page.getByRole("tab", { name: "Your Models" }),
+ "an internal user lands on their own models tab, not an admin-only view",
+ ).toBeVisible({ timeout: 15_000 });
+ await expect(
+ viewSelector(page),
+ "the models table opens scoped to the selected team",
+ ).toContainText(CURRENT_TEAM_VIEW, { timeout: 15_000 });
+ await expect(
+ modelRow(page, ungrantedModelName),
+ `the personal view lists ${ungrantedModelName}, so it is on the proxy and reachable from this page`,
+ ).toHaveCount(1, { timeout: 30_000 });
+
+ await chooseOption(page, teamSelector(page), E2E_TEAM_CRUD_ALIAS);
+ await expect(
+ modelRow(page, CHAT_MODEL_A),
+ `${E2E_TEAM_CRUD_ALIAS} lists ${CHAT_MODEL_A}`,
+ ).toHaveCount(1, {
+ timeout: 15_000,
+ });
+ await expect(
+ modelRow(page, CHAT_MODEL_B),
+ `${E2E_TEAM_CRUD_ALIAS} lists ${CHAT_MODEL_B}`,
+ ).toHaveCount(1, {
+ timeout: 15_000,
+ });
+ await expect(
+ modelRow(page, ungrantedModelName),
+ `${ungrantedModelName} is on the proxy but not granted to ${E2E_TEAM_CRUD_ALIAS}, so it must not be listed`,
+ ).toHaveCount(0);
+
+ await chooseOption(page, teamSelector(page), E2E_TEAM_ORG_ALIAS);
+ await expect(
+ modelRow(page, CHAT_MODEL_A),
+ `${E2E_TEAM_ORG_ALIAS} lists ${CHAT_MODEL_A}`,
+ ).toHaveCount(1, {
+ timeout: 15_000,
+ });
+ await expect(
+ page.getByTestId("pagination-range"),
+ `${E2E_TEAM_ORG_ALIAS} lists the one model it grants and nothing else`,
+ ).toHaveText("Showing 1-1 of 1", { timeout: 15_000 });
+ await expect(
+ modelRow(page, CHAT_MODEL_B),
+ `${CHAT_MODEL_B} belongs to another team and must not leak into ${E2E_TEAM_ORG_ALIAS}`,
+ ).toHaveCount(0);
+ await expect(
+ modelRow(page, ungrantedModelName),
+ `${ungrantedModelName} is granted to no team and must not leak into ${E2E_TEAM_ORG_ALIAS}`,
+ ).toHaveCount(0);
+
+ await chooseOption(page, viewSelector(page), ALL_MODELS_VIEW);
+ await expect(
+ modelRow(page, CHAT_MODEL_A),
+ `switching to ${ALL_MODELS_VIEW} leaves the table populated rather than blanking it`,
+ ).toHaveCount(1, { timeout: 15_000 });
+
+ await page.reload();
+ await expect(
+ teamSelector(page),
+ "the team selection is not persisted across a reload, so the table returns to the personal view",
+ ).toContainText(PERSONAL_TEAM, { timeout: 15_000 });
+ await expect(
+ viewSelector(page),
+ "the view selection is not persisted across a reload either",
+ ).toContainText(CURRENT_TEAM_VIEW, { timeout: 15_000 });
+ await expect(
+ modelRow(page, ungrantedModelName),
+ "the personal view still renders models after a reload rather than coming back empty",
+ ).toHaveCount(1, { timeout: 30_000 });
+ });
+});
diff --git a/tests/e2e/ui/tests/modelsPage/editLitellmParams.spec.ts b/tests/e2e/ui/tests/modelsPage/editLitellmParams.spec.ts
new file mode 100644
index 00000000000..4480515ae59
--- /dev/null
+++ b/tests/e2e/ui/tests/modelsPage/editLitellmParams.spec.ts
@@ -0,0 +1,252 @@
+import {
+ test as base,
+ expect,
+ type Page as PlaywrightPage,
+} from "@playwright/test";
+import { ADMIN_STORAGE_PATH } from "../../constants";
+import { Page } from "../../fixtures/pages";
+import { navigateToPage } from "../../helpers/navigation";
+import { captureRequestBody, readBack } from "../../helpers/roundTrip";
+import { masterKey, sendChatCompletion } from "../../helpers/traffic";
+
+const MOCK_LLM_BASE = `http://127.0.0.1:${process.env.MOCK_LLM_PORT ?? "8090"}/v1`;
+const CUSTOM_PARAM = "extra_headers";
+const CUSTOM_PARAM_VALUE = { "X-E2E-Edit-Probe": "one" };
+
+type StoredParams = Record;
+
+async function readStoredParams(
+ page: PlaywrightPage,
+ modelId: string,
+): Promise {
+ const body = await readBack<{ data: { litellm_params: StoredParams }[] }>(
+ page,
+ `/model/info?litellm_model_id=${modelId}`,
+ );
+ return body.data[0]?.litellm_params ?? {};
+}
+
+function paramsEditor(page: PlaywrightPage) {
+ return page.getByPlaceholder('"rpm": 100');
+}
+
+async function editParams(
+ page: PlaywrightPage,
+ mutate: (params: StoredParams) => StoredParams,
+): Promise {
+ await page.getByRole("button", { name: "Edit Settings" }).click();
+ const editor = paramsEditor(page);
+ await expect(
+ editor,
+ "the LiteLLM Params editor is reachable on every visit to the edit form",
+ ).toBeVisible({
+ timeout: 15_000,
+ });
+ const shown = JSON.parse(await editor.inputValue()) as StoredParams;
+ await editor.fill(JSON.stringify(mutate(shown), null, 2));
+}
+
+async function deleteDeployment(
+ page: PlaywrightPage,
+ id: string,
+): Promise {
+ const post = () =>
+ page.request.post("/model/delete", {
+ headers: { Authorization: `Bearer ${masterKey()}` },
+ data: { id },
+ });
+ const deleted = await post().catch(() => post());
+ expect(
+ deleted.ok(),
+ `cleanup: /model/delete ${id} returned ${deleted.status()}`,
+ ).toBe(true);
+}
+
+const uniqueSuffix = (): string =>
+ `${Date.now()}-${Math.random().toString(36).slice(2, 8)}`;
+
+const test = base.extend<{
+ deployment: { readonly modelName: string; readonly createdModelId: string };
+}>({
+ deployment: async ({ page, request }, use) => {
+ const modelName = `e2e-edit-params-${uniqueSuffix()}`;
+ const created = await page.request.post("/model/new", {
+ headers: { Authorization: `Bearer ${masterKey()}` },
+ data: {
+ model_name: modelName,
+ litellm_params: {
+ model: `openai/${modelName}`,
+ api_base: MOCK_LLM_BASE,
+ api_key: "fake-key",
+ },
+ model_info: {},
+ },
+ });
+ expect(
+ created.ok(),
+ `/model/new failed: ${created.status()} ${await created.text()}`,
+ ).toBe(true);
+ const createdModelId = (await created.json()).model_info?.id;
+ expect(createdModelId, "model id from /model/new").toBeTruthy();
+
+ try {
+ await expect
+ .poll(
+ async () => {
+ try {
+ await sendChatCompletion(request, {
+ model: modelName,
+ prompt: `warmup ${modelName}`,
+ });
+ return true;
+ } catch {
+ return false;
+ }
+ },
+ {
+ message: `deployment ${modelName} never became routable after /model/new`,
+ timeout: 60_000,
+ },
+ )
+ .toBe(true);
+ await use({ modelName, createdModelId });
+ } finally {
+ await deleteDeployment(page, createdModelId);
+ }
+ },
+});
+
+test.describe("Edit LiteLLM Params on a deployment", () => {
+ test.use({ storageState: ADMIN_STORAGE_PATH });
+
+ test("params added on a deployment can be re-edited, and the deployment keeps serving", async ({
+ page,
+ request,
+ deployment: { modelName, createdModelId },
+ }) => {
+ await navigateToPage(page, Page.Models);
+ const modelIdCell = page.getByTestId(`model-id-${createdModelId}`);
+ await expect(
+ modelIdCell,
+ `the Models table lists ${modelName}`,
+ ).toBeVisible({ timeout: 15_000 });
+ await modelIdCell.click();
+ await expect(page.getByText("Back to Models").first()).toBeVisible({
+ timeout: 15_000,
+ });
+
+ await editParams(page, (params) => ({
+ ...params,
+ temperature: 0.2,
+ [CUSTOM_PARAM]: CUSTOM_PARAM_VALUE,
+ }));
+ const firstSave = await captureRequestBody(
+ page,
+ { method: "PATCH", urlIncludes: `/model/${createdModelId}/update` },
+ async () => {
+ await page.getByRole("button", { name: "Save Changes" }).click();
+ },
+ );
+ expect(
+ firstSave.litellm_params?.temperature,
+ "the added temperature goes on the wire",
+ ).toBe(0.2);
+ expect(
+ firstSave.litellm_params?.[CUSTOM_PARAM],
+ `the added ${CUSTOM_PARAM} goes on the wire`,
+ ).toEqual(CUSTOM_PARAM_VALUE);
+ expect(
+ firstSave.litellm_params?.model,
+ "a params edit does not rewrite the upstream model",
+ ).toBe(`openai/${modelName}`);
+ expect(
+ firstSave.litellm_params?.api_base,
+ "a params edit does not rewrite the api base",
+ ).toBe(MOCK_LLM_BASE);
+ expect(
+ firstSave.litellm_params,
+ "the credential is never re-sent, so a masked placeholder cannot overwrite the stored key",
+ ).not.toHaveProperty("api_key");
+
+ await expect
+ .poll(
+ async () => (await readStoredParams(page, createdModelId)).temperature,
+ {
+ message: "the added temperature never reached the stored deployment",
+ timeout: 20_000,
+ },
+ )
+ .toBe(0.2);
+ const afterFirstSave = await readStoredParams(page, createdModelId);
+ expect(
+ afterFirstSave[CUSTOM_PARAM],
+ `the added ${CUSTOM_PARAM} reached the stored deployment`,
+ ).toEqual(CUSTOM_PARAM_VALUE);
+ expect(
+ afterFirstSave.model,
+ "the stored upstream model survived the edit",
+ ).toBe(`openai/${modelName}`);
+ expect(
+ afterFirstSave.api_base,
+ "the stored api base survived the edit",
+ ).toBe(MOCK_LLM_BASE);
+
+ await editParams(page, (params) => ({
+ ...Object.fromEntries(
+ Object.entries(params).filter(([key]) => key !== CUSTOM_PARAM),
+ ),
+ temperature: 0.7,
+ }));
+ const secondSave = await captureRequestBody(
+ page,
+ { method: "PATCH", urlIncludes: `/model/${createdModelId}/update` },
+ async () => {
+ await page.getByRole("button", { name: "Save Changes" }).click();
+ },
+ );
+ expect(
+ secondSave.litellm_params?.temperature,
+ "a param set by an earlier save can be edited again",
+ ).toBe(0.7);
+ expect(
+ secondSave.litellm_params,
+ `dropping ${CUSTOM_PARAM} from the editor drops it from the request the UI sends`,
+ ).not.toHaveProperty(CUSTOM_PARAM);
+ expect(
+ secondSave.litellm_params?.model,
+ "a second params edit still leaves the upstream model alone",
+ ).toBe(`openai/${modelName}`);
+ expect(
+ secondSave.litellm_params?.api_base,
+ "a second params edit still leaves the api base alone",
+ ).toBe(MOCK_LLM_BASE);
+ expect(
+ secondSave.litellm_params,
+ "the credential is still never re-sent",
+ ).not.toHaveProperty("api_key");
+
+ await expect
+ .poll(
+ async () => (await readStoredParams(page, createdModelId)).temperature,
+ {
+ message:
+ "the re-edited temperature never reached the stored deployment",
+ timeout: 20_000,
+ },
+ )
+ .toBe(0.7);
+
+ await page.reload();
+ await expect(
+ page
+ .getByRole("tabpanel", { name: "Overview" })
+ .getByText('"temperature": 0.7'),
+ "reopening the deployment renders the re-edited value, not the one from the first save",
+ ).toBeVisible({ timeout: 20_000 });
+
+ await sendChatCompletion(request, {
+ model: modelName,
+ prompt: `still serving ${modelName}`,
+ });
+ });
+});
diff --git a/tests/e2e/ui/tests/modelsPage/modelHealthStatus.spec.ts b/tests/e2e/ui/tests/modelsPage/modelHealthStatus.spec.ts
new file mode 100644
index 00000000000..247cce1b85d
--- /dev/null
+++ b/tests/e2e/ui/tests/modelsPage/modelHealthStatus.spec.ts
@@ -0,0 +1,245 @@
+import {
+ test as base,
+ expect,
+ type Locator,
+ type Page as PlaywrightPage,
+} from "@playwright/test";
+import { ADMIN_STORAGE_PATH } from "../../constants";
+import { Page } from "../../fixtures/pages";
+import { navigateToPage } from "../../helpers/navigation";
+import { readBack } from "../../helpers/roundTrip";
+import { masterKey } from "../../helpers/traffic";
+
+const MOCK_LLM_BASE = `http://127.0.0.1:${process.env.MOCK_LLM_PORT ?? "8090"}/v1`;
+const UNREACHABLE_BASE = "http://127.0.0.1:9/v1";
+
+async function isRegistered(
+ page: PlaywrightPage,
+ modelName: string,
+): Promise {
+ const body = await readBack<{ data: { model_name?: string }[] }>(
+ page,
+ "/v2/model/info",
+ );
+ return body.data.some((row) => row.model_name === modelName);
+}
+
+function healthRow(page: PlaywrightPage, modelName: string): Locator {
+ return page.getByRole("row").filter({ hasText: modelName });
+}
+
+function pageOf(label: string): { current: number; total: number } {
+ const [current, total] = label
+ .replace("Page ", "")
+ .split(" of ")
+ .map((part) => Number(part.trim()));
+ return { current, total };
+}
+
+async function locateHealthRow(
+ page: PlaywrightPage,
+ modelName: string,
+): Promise {
+ const pageLabel = page.getByTestId("pagination-page");
+ await expect(
+ pageLabel,
+ "the health table reports which page it is showing",
+ ).toBeVisible({ timeout: 20_000 });
+
+ const deadline = Date.now() + 60_000;
+ while (Date.now() < deadline) {
+ const row = healthRow(page, modelName);
+ const onThisPage = await row
+ .first()
+ .waitFor({ state: "visible", timeout: 3_000 })
+ .then(() => true)
+ .catch(() => false);
+ if (onThisPage) return row;
+
+ const { current, total } = pageOf(await pageLabel.innerText());
+ const goTo = current < total ? current + 1 : 1;
+ if (total === 1) continue;
+ await page
+ .getByRole("button", {
+ name: current < total ? "Go to next page" : "Go to first page",
+ })
+ .click();
+ await expect(pageLabel).toContainText(`Page ${goTo} of`, {
+ timeout: 15_000,
+ });
+ }
+ return healthRow(page, modelName);
+}
+
+async function openHealthTab(page: PlaywrightPage): Promise {
+ await page.getByRole("tab", { name: "Health Status" }).click();
+ await expect(
+ page.getByRole("heading", { name: "Model Health Status" }),
+ ).toBeVisible({ timeout: 15_000 });
+}
+
+async function expectStatus(
+ page: PlaywrightPage,
+ modelName: string,
+ status: string,
+): Promise {
+ const row = await locateHealthRow(page, modelName);
+ await expect(row, `${modelName} has one row in the health table`).toHaveCount(
+ 1,
+ { timeout: 20_000 },
+ );
+ await expect(
+ row.getByText(status, { exact: true }),
+ `the Health Status cell for ${modelName} reads ${status}`,
+ ).toHaveCount(1, { timeout: 60_000 });
+}
+
+async function deleteDeployment(
+ page: PlaywrightPage,
+ id: string,
+): Promise {
+ const post = () =>
+ page.request.post("/model/delete", {
+ headers: { Authorization: `Bearer ${masterKey()}` },
+ data: { id },
+ });
+ const deleted = await post().catch(() => post());
+ expect(
+ deleted.ok(),
+ `cleanup: /model/delete ${id} returned ${deleted.status()}`,
+ ).toBe(true);
+}
+
+const uniqueSuffix = (): string =>
+ `${Date.now()}-${Math.random().toString(36).slice(2, 8)}`;
+
+async function withDeployment(
+ page: PlaywrightPage,
+ prefix: string,
+ apiBase: string,
+ use: (name: string) => Promise,
+): Promise {
+ const name = `${prefix}-${uniqueSuffix()}`;
+ const created = await page.request.post("/model/new", {
+ headers: { Authorization: `Bearer ${masterKey()}` },
+ data: {
+ model_name: name,
+ litellm_params: {
+ model: `openai/${name}`,
+ api_base: apiBase,
+ api_key: "fake-key",
+ },
+ model_info: {},
+ },
+ });
+ expect(
+ created.ok(),
+ `/model/new for ${name} failed: ${created.status()} ${await created.text()}`,
+ ).toBe(true);
+ const id = (await created.json()).model_info?.id;
+ expect(id, `model id from /model/new for ${name}`).toBeTruthy();
+ try {
+ await expect
+ .poll(() => isRegistered(page, name), {
+ message: `deployment ${name} never appeared in /v2/model/info after create`,
+ timeout: 60_000,
+ })
+ .toBe(true);
+ await use(name);
+ } finally {
+ await deleteDeployment(page, id);
+ }
+}
+
+const test = base.extend<{ reachableName: string; unreachableName: string }>({
+ reachableName: async ({ page }, use) => {
+ await withDeployment(page, "e2e-health-up", MOCK_LLM_BASE, use);
+ },
+ unreachableName: async ({ page }, use) => {
+ await withDeployment(page, "e2e-health-down", UNREACHABLE_BASE, use);
+ },
+});
+
+test.describe("Model health status", () => {
+ test.use({ storageState: ADMIN_STORAGE_PATH });
+
+ test("Run Health Check reports a reachable deployment healthy and an unreachable one unhealthy", async ({
+ page,
+ reachableName,
+ unreachableName,
+ }) => {
+ await navigateToPage(page, Page.Models);
+ await openHealthTab(page);
+
+ for (const name of [reachableName, unreachableName]) {
+ const row = await locateHealthRow(page, name);
+ await expect(row, `${name} has one row in the health table`).toHaveCount(
+ 1,
+ { timeout: 20_000 },
+ );
+ await row
+ .getByRole("button", { name: "Run Health Check", exact: true })
+ .click();
+ }
+
+ await expectStatus(page, reachableName, "healthy");
+ await expect(
+ healthRow(page, reachableName).getByText("unhealthy", { exact: true }),
+ "a reachable deployment is never reported unhealthy",
+ ).toHaveCount(0);
+ await expectStatus(page, unreachableName, "unhealthy");
+
+ const successDetail = (
+ await locateHealthRow(page, reachableName)
+ ).getByRole("button", {
+ name: "View response details",
+ });
+ await expect(
+ successDetail,
+ `${reachableName} offers its health check response for inspection`,
+ ).toBeVisible({ timeout: 60_000 });
+ await successDetail.click();
+ const successDialog = page.getByRole("dialog");
+ await expect(
+ successDialog.getByRole("heading", {
+ name: `Health Check Response - ${reachableName}`,
+ }),
+ "the healthy deployment's detail opens its own response dialog",
+ ).toBeVisible({ timeout: 10_000 });
+ await successDialog.getByRole("button", { name: "Close" }).last().click();
+ await expect(successDialog).toBeHidden({ timeout: 10_000 });
+
+ const errorDetail = (
+ await locateHealthRow(page, unreachableName)
+ ).getByRole("button", {
+ name: "View full error details",
+ });
+ await expect(
+ errorDetail,
+ `${unreachableName} offers its health check error for inspection`,
+ ).toBeVisible({ timeout: 60_000 });
+ await errorDetail.click();
+ const errorDialog = page.getByRole("dialog");
+ await expect(
+ errorDialog.getByRole("heading", {
+ name: `Health Check Error - ${unreachableName}`,
+ }),
+ "the unreachable deployment's detail opens its own error dialog",
+ ).toBeVisible({ timeout: 10_000 });
+ await expect(
+ errorDialog,
+ "the error dialog carries the upstream connection failure, not a generic message",
+ ).toContainText(/connection error/i, { timeout: 10_000 });
+ await expect(
+ errorDialog,
+ "the error dialog names the endpoint that could not be reached",
+ ).toContainText(UNREACHABLE_BASE);
+ await errorDialog.getByRole("button", { name: "Close" }).last().click();
+ await expect(errorDialog).toBeHidden({ timeout: 10_000 });
+
+ await page.reload();
+ await openHealthTab(page);
+ await expectStatus(page, reachableName, "healthy");
+ await expectStatus(page, unreachableName, "unhealthy");
+ });
+});
diff --git a/tests/e2e/ui/tests/proxy-admin/keyBlocking.spec.ts b/tests/e2e/ui/tests/proxy-admin/keyBlocking.spec.ts
new file mode 100644
index 00000000000..99a8065a797
--- /dev/null
+++ b/tests/e2e/ui/tests/proxy-admin/keyBlocking.spec.ts
@@ -0,0 +1,112 @@
+import { test as base, expect } from "@playwright/test";
+import { ADMIN_STORAGE_PATH } from "../../constants";
+import { Page } from "../../fixtures/pages";
+import { dismissFeedbackPopup, navigateToPage, openKeyDetail } from "../../helpers/navigation";
+import {
+ CHAT_MODEL_A,
+ MOCK_RESPONSE_TEXT,
+ attemptChatCompletion,
+ createVirtualKey,
+ deleteVirtualKey,
+ readKeyInfo,
+ sendChatCompletion,
+ uniqueSuffix,
+} from "../../helpers/traffic";
+
+interface ScopedKey {
+ alias: string;
+ token: string;
+ apiKey: string;
+}
+
+const test = base.extend<{ scopedKey: ScopedKey }>({
+ scopedKey: async ({ page }, use) => {
+ const alias = `e2e-block-key-${uniqueSuffix()}`;
+ const created = await createVirtualKey(page.request, {
+ key_alias: alias,
+ models: [CHAT_MODEL_A],
+ });
+ await use({ alias, token: created.token, apiKey: created.key });
+ await deleteVirtualKey(page.request, created.token);
+ },
+});
+
+test.describe("Proxy Admin - Key blocking", () => {
+ test.use({ storageState: ADMIN_STORAGE_PATH });
+
+ test("blocking a key stops it serving and unblocking restores it", async ({ page, scopedKey }) => {
+ const { alias, token, apiKey } = scopedKey;
+
+ await sendChatCompletion(page.request, {
+ model: CHAT_MODEL_A,
+ prompt: `pre-block ${alias}`,
+ apiKey,
+ });
+
+ await navigateToPage(page, Page.ApiKeys);
+ await dismissFeedbackPopup(page);
+ await openKeyDetail(page, alias);
+
+ await page.getByRole("button", { name: "More key actions" }).click();
+ await page.getByRole("menuitem", { name: "Block Key" }).click();
+ const blockDialog = page.getByRole("dialog", { name: "Block Key" });
+ await expect(blockDialog, "the Block Key confirmation never opened").toBeVisible({ timeout: 10_000 });
+ await blockDialog.getByRole("button", { name: "Block", exact: true }).click();
+
+ await expect
+ .poll(async () => (await readKeyInfo(page.request, token)).blocked, {
+ message: "the key never came back blocked from /key/info",
+ timeout: 20_000,
+ })
+ .toBe(true);
+
+ await expect
+ .poll(
+ async () =>
+ await attemptChatCompletion(page.request, {
+ model: CHAT_MODEL_A,
+ prompt: "blocked",
+ apiKey,
+ }),
+ {
+ message: "a blocked key was still served by /v1/chat/completions",
+ timeout: 30_000,
+ },
+ )
+ .toMatchObject({ status: 401, body: expect.stringContaining("blocked") });
+
+ await page.reload();
+ await expect(
+ page.getByText("Blocked", { exact: true }),
+ "the reloaded key detail does not show the key as blocked",
+ ).toBeVisible({ timeout: 15_000 });
+
+ await page.getByRole("button", { name: "More key actions" }).click();
+ await page.getByRole("menuitem", { name: "Unblock Key" }).click();
+ const unblockDialog = page.getByRole("dialog", { name: "Unblock Key" });
+ await expect(unblockDialog, "the Unblock Key confirmation never opened").toBeVisible({ timeout: 10_000 });
+ await unblockDialog.getByRole("button", { name: "Unblock", exact: true }).click();
+
+ await expect
+ .poll(async () => (await readKeyInfo(page.request, token)).blocked, {
+ message: "the key never came back unblocked from /key/info",
+ timeout: 20_000,
+ })
+ .toBe(false);
+
+ await expect
+ .poll(
+ async () =>
+ await attemptChatCompletion(page.request, {
+ model: CHAT_MODEL_A,
+ prompt: "unblocked",
+ apiKey,
+ }),
+ {
+ message: "an unblocked key is still refused by /v1/chat/completions",
+ timeout: 30_000,
+ },
+ )
+ .toMatchObject({ status: 200, body: expect.stringContaining(MOCK_RESPONSE_TEXT) });
+ });
+});
diff --git a/tests/e2e/ui/tests/proxy-admin/keyBudgetWindow.spec.ts b/tests/e2e/ui/tests/proxy-admin/keyBudgetWindow.spec.ts
new file mode 100644
index 00000000000..4e4d0a395c3
--- /dev/null
+++ b/tests/e2e/ui/tests/proxy-admin/keyBudgetWindow.spec.ts
@@ -0,0 +1,101 @@
+import { test as base, expect } from "@playwright/test";
+import { ADMIN_STORAGE_PATH, E2E_TEAM_CRUD_ID } from "../../constants";
+import { Page } from "../../fixtures/pages";
+import { dismissFeedbackPopup, navigateToPage, openKeyDetail } from "../../helpers/navigation";
+import { captureRequestBody } from "../../helpers/roundTrip";
+import { CHAT_MODEL_A, createVirtualKey, deleteVirtualKey, readKeyInfo, uniqueSuffix } from "../../helpers/traffic";
+
+interface ScopedKey {
+ alias: string;
+ token: string;
+}
+
+const test = base.extend<{ scopedKey: ScopedKey }>({
+ scopedKey: async ({ page }, use) => {
+ const alias = `e2e-budget-window-${uniqueSuffix()}`;
+ const created = await createVirtualKey(page.request, {
+ key_alias: alias,
+ team_id: E2E_TEAM_CRUD_ID,
+ models: [CHAT_MODEL_A],
+ });
+ await use({ alias, token: created.token });
+ await deleteVirtualKey(page.request, created.token);
+ },
+});
+
+test.describe("Proxy Admin - Key budget window", () => {
+ test.use({ storageState: ADMIN_STORAGE_PATH });
+
+ test("a monthly spend cap survives a reload, and clearing the window keeps the cap", async ({ page, scopedKey }) => {
+ const { alias, token } = scopedKey;
+
+ const before = await readKeyInfo(page.request, token);
+ expect(before.max_budget, "a freshly generated key starts with no budget").toBeNull();
+
+ await navigateToPage(page, Page.ApiKeys);
+ await dismissFeedbackPopup(page);
+ await openKeyDetail(page, alias);
+
+ await page.getByRole("tab", { name: "Settings" }).click();
+ await page.getByRole("button", { name: "Edit Settings" }).click();
+
+ await page.getByRole("spinbutton", { name: "Max Budget (USD)" }).fill("12.5");
+ await page.getByLabel("Reset Budget", { exact: true }).click();
+ await page.getByRole("option", { name: "monthly", exact: true }).click();
+ await page.getByRole("button", { name: "Save Changes" }).click();
+
+ await expect
+ .poll(async () => (await readKeyInfo(page.request, token)).max_budget, {
+ message: "the $12.50 cap never reached /key/info",
+ timeout: 20_000,
+ })
+ .toBe(12.5);
+ await expect
+ .poll(async () => (await readKeyInfo(page.request, token)).budget_duration, {
+ message: "the monthly reset window never reached /key/info",
+ timeout: 20_000,
+ })
+ .toBe("30d");
+
+ const capped = await readKeyInfo(page.request, token);
+ const resetAt = new Date(capped.budget_reset_at ?? "");
+ expect(Number.isNaN(resetAt.getTime()), "a monthly window left the key with no budget_reset_at").toBe(false);
+ expect(resetAt.getTime(), "budget_reset_at was set in the past").toBeGreaterThan(Date.now());
+ expect(resetAt.getUTCDate(), "a monthly window resets on the 1st, a daily one would not").toBe(1);
+
+ await page.reload();
+ await expect(
+ page.getByRole("paragraph").filter({ hasText: "of $12.50" }),
+ "the reloaded key detail does not render the $12.50 cap",
+ ).toBeVisible({ timeout: 15_000 });
+
+ await page.getByRole("tab", { name: "Settings" }).click();
+ await expect(
+ page.getByTestId("budget-reset-value"),
+ "the reloaded key detail does not name the 30d reset window",
+ ).toHaveText(/Every 30d/, { timeout: 15_000 });
+
+ await page.getByRole("button", { name: "Edit Settings" }).click();
+ await page.getByLabel("Reset Budget", { exact: true }).click();
+ await page.getByRole("option", { name: "Never resets", exact: true }).click();
+
+ const cleared = await captureRequestBody(page, { method: "POST", urlIncludes: "/key/update" }, async () => {
+ await page.getByRole("button", { name: "Save Changes" }).click();
+ });
+ expect(cleared).toHaveProperty("budget_duration");
+ expect(cleared.budget_duration, "clearing the window must send budget_duration: null explicitly").toBeNull();
+
+ await expect
+ .poll(async () => (await readKeyInfo(page.request, token)).budget_duration, {
+ message: "the reset window was never cleared on /key/info",
+ timeout: 20_000,
+ })
+ .toBeNull();
+
+ const after = await readKeyInfo(page.request, token);
+ expect(after.budget_reset_at, "clearing the reset window left a stale next-reset timestamp").toBeNull();
+ expect(after.max_budget, "clearing the reset window also wiped the spend cap").toBe(12.5);
+ expect(after.models, "editing the budget left the key's models untouched").toEqual(before.models);
+ expect(after.team_id, "editing the budget left the key's team untouched").toEqual(before.team_id);
+ });
+});
diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
index 091b958d7c3..48fceb50403 100644
--- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
+++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
@@ -1095,6 +1095,110 @@ async def test_afile_content_passes_trusted_model_credentials_to_router():
assert trusted_credentials["s3_bucket_name"] == "my-bucket"
+def _managed_deletion_file_id(provider_file_id):
+ from litellm.types.utils import SpecialEnums
+
+ value = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
+ "application/json", "test-file", "batch-model", provider_file_id, "model-123"
+ )
+ return base64.urlsafe_b64encode(value.encode()).decode().rstrip("=")
+
+
+def _managed_files_with_deletion_row(unified_file_id, provider_file_id, file_object):
+ from litellm.caching import DualCache
+ from litellm.models.managed_files import LiteLLM_ManagedFileTable
+ from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
+
+ row = LiteLLM_ManagedFileTable(
+ unified_file_id=unified_file_id,
+ model_mappings={"model-123": provider_file_id},
+ flat_model_file_ids=[provider_file_id],
+ file_object=file_object,
+ )
+ table = MagicMock(
+ find_first=AsyncMock(return_value=row),
+ delete=AsyncMock(),
+ )
+ return _PROXY_LiteLLMManagedFiles(
+ internal_usage_cache=DualCache(),
+ prisma_client=MagicMock(db=MagicMock(litellm_managedfiletable=table)),
+ ), table
+
+
+@pytest.mark.asyncio
+async def test_afile_delete_bedrock_uses_deployment_bucket_and_signed_s3_delete(monkeypatch):
+ import httpx
+ import respx
+
+ from litellm import Router
+
+ monkeypatch.delenv("AWS_S3_BUCKET_NAME", raising=False)
+ monkeypatch.delenv("AWS_S3_OUTPUT_BUCKET_NAME", raising=False)
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ router = Router(
+ model_list=[
+ {
+ "model_name": "bedrock-batch",
+ "litellm_params": {
+ "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
+ "aws_access_key_id": "AKIAEXAMPLE",
+ "aws_secret_access_key": "secret",
+ "aws_region_name": "us-west-2",
+ "s3_bucket_name": "my-bucket",
+ },
+ "model_info": {"id": "model-123"},
+ }
+ ],
+ num_retries=0,
+ )
+ s3_uri = "s3://my-bucket/litellm-bedrock-files/input.jsonl"
+ unified_file_id = _managed_deletion_file_id(s3_uri)
+ managed_files, table = _managed_files_with_deletion_row(unified_file_id, s3_uri, None)
+ with respx.mock:
+ route = respx.delete(
+ "https://s3.us-west-2.amazonaws.com/my-bucket/litellm-bedrock-files/input.jsonl"
+ ).mock(return_value=httpx.Response(204))
+ response = await managed_files.afile_delete(
+ file_id=unified_file_id,
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ _litellm_internal_model_credentials={"s3_bucket_name": "request-bucket"},
+ )
+
+ assert len(route.calls) == 1
+ assert route.calls[0].request.headers["Authorization"].startswith("AWS4-HMAC-SHA256")
+ assert response.id == unified_file_id
+ assert response.deleted is True
+ table.delete.assert_awaited_once_with(where={"unified_file_id": unified_file_id})
+
+
+@pytest.mark.asyncio
+async def test_afile_delete_returns_managed_id_for_stored_provider_output():
+ from openai.types import FileDeleted
+
+ provider_file_id = "file-error-output"
+ unified_file_id = _managed_deletion_file_id(provider_file_id)
+ stored_file = _make_file_object(provider_file_id)
+ managed_files, table = _managed_files_with_deletion_row(unified_file_id, provider_file_id, stored_file)
+ router = MagicMock(
+ get_deployment_credentials_with_provider=MagicMock(return_value=None),
+ afile_delete=AsyncMock(return_value=FileDeleted(id=provider_file_id, object="file", deleted=True)),
+ )
+ response = await managed_files.afile_delete(
+ file_id=unified_file_id,
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ _litellm_internal_model_credentials={"s3_bucket_name": "request-bucket"},
+ )
+
+ assert response.id == unified_file_id
+ assert response.object == "file"
+ assert response.filename == stored_file.filename
+ assert stored_file.id == provider_file_id
+ router.afile_delete.assert_awaited_once_with(model="model-123", file_id=provider_file_id)
+ table.delete.assert_awaited_once_with(where={"unified_file_id": unified_file_id})
+
+
@pytest.mark.asyncio
async def test_afile_content_bedrock_unified_id_end_to_end(monkeypatch):
"""
diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py
index 4db131da62c..b07e5876e8b 100644
--- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py
+++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py
@@ -1,9 +1,12 @@
import asyncio
import base64
+import json
import os
import sys
+from collections.abc import AsyncIterator
from importlib import metadata
from pathlib import Path
+from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import anyio
@@ -11,15 +14,19 @@ import httpx
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
from mcp import McpError
+from mcp.client.streamable_http import streamable_http_client
+from pydantic import ValidationError
from mcp.shared.message import SessionMessage
from mcp.types import (
LATEST_PROTOCOL_VERSION,
+ CallToolResult,
ErrorData,
Implementation,
InitializeResult,
JSONRPCError,
JSONRPCMessage,
JSONRPCResponse,
+ LoggingMessageNotificationParams,
ServerCapabilities,
)
@@ -29,8 +36,9 @@ import litellm.experimental_mcp_client.client as mcp_client_module
from litellm.experimental_mcp_client.client import (
MCP_STREAMABLE_HTTP_REQUIREMENT,
MCPClient,
- _as_read_timeout,
_first_non_cancelled_cause,
+ _TransportContext,
+ as_mcp_read_timeout,
missing_streamable_http_client_error,
strip_auth_scheme,
)
@@ -859,25 +867,25 @@ def _raise_mcp_error_while_handling_a_timeout(code: int, message: str) -> McpErr
return raised
-def test_as_read_timeout_separates_the_sdk_timeout_from_a_relayed_upstream_error():
+def test_as_mcp_read_timeout_separates_the_sdk_timeout_from_a_relayed_upstream_error():
"""Neither signal alone is enough. The code alone cannot separate the SDK's own timeout from an
upstream JSON-RPC error that happens to use 408, and the context chain alone cannot separate it
from any other relayed error that surfaces while a timeout is being handled, so both must hold.
"""
timeout_code = int(httpx.codes.REQUEST_TIMEOUT)
- translated = _as_read_timeout(_raise_mcp_error_while_handling_a_timeout(timeout_code, "Timed out while waiting"))
+ translated = as_mcp_read_timeout(_raise_mcp_error_while_handling_a_timeout(timeout_code, "Timed out while waiting"))
assert isinstance(translated, TimeoutError)
assert str(translated) == "Timed out while waiting"
relayed_408 = McpError(ErrorData(code=timeout_code, message="upstream said 408"))
- assert _as_read_timeout(relayed_408) is None, "an upstream 408 with no elapsed timeout is not our timeout"
+ assert as_mcp_read_timeout(relayed_408) is None, "an upstream 408 with no elapsed timeout is not our timeout"
relayed_other = _raise_mcp_error_while_handling_a_timeout(-32603, "upstream internal error")
- assert _as_read_timeout(relayed_other) is None, "a non-timeout code is not our timeout, whatever the chain"
+ assert as_mcp_read_timeout(relayed_other) is None, "a non-timeout code is not our timeout, whatever the chain"
- assert _as_read_timeout(McpError(ErrorData(code=-32603, message="boom"))) is None
- assert _as_read_timeout(RuntimeError("not an McpError")) is None
+ assert as_mcp_read_timeout(McpError(ErrorData(code=-32603, message="boom"))) is None
+ assert as_mcp_read_timeout(RuntimeError("not an McpError")) is None
@pytest.mark.asyncio
@@ -1224,14 +1232,14 @@ def test_without_a_configured_slot_the_existing_precedence_is_unchanged():
_REDIRECT_CASES = [
- ("https://upstream.example.com/mcp", "https://upstream.example.com/other"), # same origin
+ ("https://upstream.example.com/mcp", "https://upstream.example.com/other"), # same origin
("https://upstream.example.com/mcp", "https://upstream.example.com:443/other"), # explicit default port
- ("https://upstream.example.com/mcp", "https://attacker.example.com/collect"), # different host
- ("https://upstream.example.com/mcp", "http://upstream.example.com/collect"), # scheme downgrade
- ("https://upstream.example.com/mcp", "https://upstream.example.com:8443/other"), # different port
- ("https://upstream.example.com/mcp", "https://sub.upstream.example.com/x"), # different host
- ("http://upstream.example.com/mcp", "https://upstream.example.com/other"), # http -> https upgrade
- ("http://upstream.example.com/mcp", "http://upstream.example.com/other"), # same origin, plain http
+ ("https://upstream.example.com/mcp", "https://attacker.example.com/collect"), # different host
+ ("https://upstream.example.com/mcp", "http://upstream.example.com/collect"), # scheme downgrade
+ ("https://upstream.example.com/mcp", "https://upstream.example.com:8443/other"), # different port
+ ("https://upstream.example.com/mcp", "https://sub.upstream.example.com/x"), # different host
+ ("http://upstream.example.com/mcp", "https://upstream.example.com/other"), # http -> https upgrade
+ ("http://upstream.example.com/mcp", "http://upstream.example.com/other"), # same origin, plain http
]
@@ -1283,3 +1291,398 @@ def test_a_differently_cased_injected_header_cannot_shadow_the_slot() -> None:
headers = client._get_auth_headers()
assert [v for k, v in headers.items() if k.lower() == "esb-oauth"] == ["Bearer minted-token"]
assert headers["X-Trace"] == "keep"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ ("content_type", "body", "expected_type"),
+ [
+ ("text/html", b"secret-page", ValueError),
+ ("application/json", b"secret-invalid-json", ValidationError),
+ ("application/json", b"", ValidationError),
+ ("application/json", b'{"secret":"invalid-rpc"}', ValidationError),
+ ("application/json", b'{"jsonrpc":"2.0","id":0,"result":{"secret":"invalid-schema"}}', ValidationError),
+ ],
+)
+async def test_invalid_http_response_surfaces_without_waiting_for_timeout(
+ content_type: str, body: bytes, expected_type: type[Exception]
+) -> None:
+ from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
+
+ def respond(request: httpx.Request) -> httpx.Response:
+ return httpx.Response(200, headers={"Content-Type": content_type}, content=body)
+
+ async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
+ client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
+ with pytest.raises(expected_type) as caught:
+ await asyncio.wait_for(
+ client._execute_session_operation(
+ streamable_http_client(client.server_url, http_client=http_client),
+ lambda session: session.list_tools(),
+ ),
+ timeout=3,
+ )
+
+ message: Final = _connection_error_message(caught.value, client.server_url, 30)
+ assert "unsupported content type" in message or "invalid MCP response" in message
+ assert "secret" not in message
+ assert "timed out" not in message
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("status_code", [200, 401, 503])
+async def test_http_response_handler_preserves_success_and_http_errors(status_code: int) -> None:
+ def respond(request: httpx.Request) -> httpx.Response:
+ if request.method == "DELETE":
+ return httpx.Response(200)
+ payload: Final = json.loads(request.content)
+ if "id" not in payload:
+ return httpx.Response(202)
+ result: Final = (
+ {
+ "protocolVersion": LATEST_PROTOCOL_VERSION,
+ "capabilities": {},
+ "serverInfo": {"name": "test", "version": "1"},
+ }
+ if payload["method"] == "initialize"
+ else {"tools": []}
+ )
+ return httpx.Response(status_code, json={"jsonrpc": "2.0", "id": payload["id"], "result": result})
+
+ async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
+ client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
+ operation: Final = client._execute_session_operation(
+ streamable_http_client(client.server_url, http_client=http_client), lambda session: session.list_tools()
+ )
+ if status_code == 200:
+ result: Final = await asyncio.wait_for(operation, timeout=3)
+ assert result.tools == []
+ else:
+ with pytest.raises(httpx.HTTPStatusError) as caught:
+ await asyncio.wait_for(operation, timeout=3)
+ assert caught.value.response.status_code == status_code
+
+
+@pytest.mark.asyncio
+async def test_http_response_handler_preserves_notifications_and_tool_listing() -> None:
+ notification: Final = {
+ "jsonrpc": "2.0",
+ "method": "notifications/message",
+ "params": {"level": "info", "data": "Listing tools"},
+ }
+ logging_callback: Final = AsyncMock()
+
+ def respond(request: httpx.Request) -> httpx.Response:
+ if request.method == "DELETE":
+ return httpx.Response(200)
+ payload: Final = json.loads(request.content)
+ if "id" not in payload:
+ return httpx.Response(202)
+ if payload["method"] == "initialize":
+ return httpx.Response(
+ 200,
+ json={
+ "jsonrpc": "2.0",
+ "id": payload["id"],
+ "result": {
+ "protocolVersion": LATEST_PROTOCOL_VERSION,
+ "capabilities": {"logging": {}, "tools": {}},
+ "serverInfo": {"name": "test", "version": "1"},
+ },
+ },
+ )
+ response: Final = {
+ "jsonrpc": "2.0",
+ "id": payload["id"],
+ "result": {"tools": [{"name": "search", "inputSchema": {"type": "object"}}]},
+ }
+ return httpx.Response(
+ 200,
+ headers={"Content-Type": "text/event-stream"},
+ content="".join(f"event: message\ndata: {json.dumps(message)}\n\n" for message in (notification, response)),
+ )
+
+ async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
+ client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30, logging_callback=logging_callback)
+ result: Final = await asyncio.wait_for(
+ client._execute_session_operation(
+ streamable_http_client(client.server_url, http_client=http_client), lambda session: session.list_tools()
+ ),
+ timeout=3,
+ )
+
+ assert [tool.name for tool in result.tools] == ["search"]
+ logging_callback.assert_awaited_once_with(LoggingMessageNotificationParams(level="info", data="Listing tools"))
+
+
+@pytest.mark.asyncio
+async def test_invalid_tool_list_schema_is_identified_as_an_upstream_response() -> None:
+ from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
+
+ def respond(request: httpx.Request) -> httpx.Response:
+ if request.method == "DELETE":
+ return httpx.Response(200)
+ payload: Final = json.loads(request.content)
+ if "id" not in payload:
+ return httpx.Response(202)
+ result: Final = (
+ {
+ "protocolVersion": LATEST_PROTOCOL_VERSION,
+ "capabilities": {},
+ "serverInfo": {"name": "test", "version": "1"},
+ }
+ if payload["method"] == "initialize"
+ else {"tools": "secret-invalid-tools"}
+ )
+ return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload["id"], "result": result})
+
+ async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
+ client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
+ with pytest.raises(ValidationError) as caught:
+ await asyncio.wait_for(
+ client._execute_session_operation(
+ streamable_http_client(client.server_url, http_client=http_client),
+ lambda session: session.list_tools(),
+ ),
+ timeout=3,
+ )
+
+ message: Final = _connection_error_message(caught.value, client.server_url, 30)
+ assert "invalid MCP response" in message
+ assert "secret" not in message
+
+
+class _DiagnosticSSEStream(httpx.AsyncByteStream):
+ def __init__(self, messages: asyncio.Queue[bytes | Exception | None]) -> None:
+ self.messages = messages
+
+ async def __aiter__(self) -> AsyncIterator[bytes]:
+ yield b"event: endpoint\ndata: /messages\n\n"
+ while True:
+ message: Final = await self.messages.get()
+ if message is None:
+ return
+ if isinstance(message, Exception):
+ raise message
+ yield b"event: message\ndata: " + message + b"\n\n"
+
+
+_DIAGNOSTIC_STDIO_SERVER: Final = """
+import json, sys
+mode, failure_method = sys.argv[1:]
+for line in sys.stdin:
+ request = json.loads(line)
+ if "method" not in request or "id" not in request:
+ continue
+ if request["method"] == failure_method:
+ if mode == "bad-json":
+ print("secret-invalid-json", flush=True)
+ continue
+ if mode == "closed":
+ sys.exit(0)
+ if mode == "silent":
+ print(json.dumps({"jsonrpc": "2.0", "method": "notifications/message", "params": {"level": "info", "data": "Waiting"}}), flush=True)
+ continue
+ if request["method"] == "initialize":
+ result = {"protocolVersion": request["params"]["protocolVersion"], "capabilities": {"tools": {}, "logging": {}}, "serverInfo": {"name": "diagnostic", "version": "1"}}
+ elif request["method"] == "tools/list":
+ print(json.dumps({"jsonrpc": "2.0", "method": "notifications/message", "params": {"level": "info", "data": "Listing tools"}}), flush=True)
+ print(json.dumps({"jsonrpc": "2.0", "id": "unmatched", "result": {}}), flush=True)
+ print(json.dumps({"jsonrpc": "2.0", "id": "server-ping", "method": "ping"}), flush=True)
+ result = {"tools": [{"name": "ping", "inputSchema": {"type": "object"}}]}
+ else:
+ result = {"content": [{"type": "text", "text": "pong"}], "isError": False}
+ print(json.dumps({"jsonrpc": "2.0", "id": request["id"], "result": result}), flush=True)
+"""
+
+
+def _diagnostic_transport(transport: MCPTransport, mode: str, failure_method: str) -> _TransportContext:
+ from mcp import StdioServerParameters
+ from mcp.client.sse import sse_client
+ from mcp.client.stdio import stdio_client
+
+ if transport == MCPTransport.stdio:
+ return stdio_client(
+ StdioServerParameters(
+ command=sys.executable, args=["-u", "-c", _DIAGNOSTIC_STDIO_SERVER, mode, failure_method]
+ )
+ )
+ messages: Final[asyncio.Queue[bytes | Exception | None]] = asyncio.Queue()
+
+ async def respond(request: httpx.Request) -> httpx.Response:
+ if request.method == "GET":
+ return httpx.Response(
+ 200, headers={"Content-Type": "text/event-stream"}, stream=_DiagnosticSSEStream(messages)
+ )
+ payload: Final = json.loads(request.content)
+ if "method" not in payload or "id" not in payload:
+ return httpx.Response(202)
+ if payload["method"] == failure_method and mode != "ok":
+ if mode == "bad-json":
+ await messages.put(b"secret-invalid-json")
+ elif mode == "io-error":
+ await messages.put(httpx.ReadError("secret-read-error"))
+ elif mode == "closed":
+ await messages.put(None)
+ elif mode == "silent":
+ await messages.put(
+ b'{"jsonrpc":"2.0","method":"notifications/message","params":{"level":"info","data":"Waiting"}}'
+ )
+ return httpx.Response(202)
+ if payload["method"] == "tools/list":
+ for message in (
+ {
+ "jsonrpc": "2.0",
+ "method": "notifications/message",
+ "params": {"level": "info", "data": "Listing tools"},
+ },
+ {"jsonrpc": "2.0", "id": "unmatched", "result": {}},
+ {"jsonrpc": "2.0", "id": "server-ping", "method": "ping"},
+ ):
+ await messages.put(json.dumps(message).encode())
+ result: Final = (
+ {
+ "protocolVersion": LATEST_PROTOCOL_VERSION,
+ "capabilities": {"tools": {}, "logging": {}},
+ "serverInfo": {"name": "diagnostic", "version": "1"},
+ }
+ if payload["method"] == "initialize"
+ else {"tools": [{"name": "ping", "inputSchema": {"type": "object"}}]}
+ if payload["method"] == "tools/list"
+ else {"content": [{"type": "text", "text": "pong"}], "isError": False}
+ )
+ await messages.put(json.dumps({"jsonrpc": "2.0", "id": payload["id"], "result": result}).encode())
+ return httpx.Response(202)
+
+ def factory(
+ headers: dict[str, str] | None = None,
+ timeout: httpx.Timeout | None = None,
+ auth: httpx.Auth | None = None,
+ ) -> httpx.AsyncClient:
+ return httpx.AsyncClient(transport=httpx.MockTransport(respond), headers=headers, timeout=timeout, auth=auth)
+
+ return sse_client("https://example.com/sse", httpx_client_factory=factory)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
+@pytest.mark.parametrize("failure_method", ["initialize", "tools/list"])
+async def test_transport_parsing_failure_is_preserved(transport: MCPTransport, failure_method: str) -> None:
+ client: Final = MCPClient(server_url="https://example.com/sse", transport_type=transport, timeout=0.2)
+ with pytest.raises(ValidationError):
+ await asyncio.wait_for(
+ client._execute_session_operation(
+ _diagnostic_transport(transport, "bad-json", failure_method), lambda session: session.list_tools()
+ ),
+ timeout=3,
+ )
+
+
+@pytest.mark.asyncio
+async def test_sse_read_failure_is_preserved() -> None:
+ client: Final = MCPClient(server_url="https://example.com/sse", transport_type=MCPTransport.sse, timeout=0.2)
+ with pytest.raises(httpx.ReadError, match="secret-read-error"):
+ await asyncio.wait_for(
+ client._execute_session_operation(
+ _diagnostic_transport(MCPTransport.sse, "io-error", "tools/list"), lambda session: session.list_tools()
+ ),
+ timeout=3,
+ )
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
+@pytest.mark.parametrize("mode", ["ok", "closed", "silent"])
+async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str) -> None:
+ from mcp import ClientSession
+ from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
+
+ logging_callback: Final = AsyncMock()
+ client: Final = MCPClient(
+ server_url="https://example.com/sse", transport_type=transport, timeout=0.2, logging_callback=logging_callback
+ )
+
+ async def operation(session: ClientSession) -> CallToolResult:
+ tools: Final = await session.list_tools()
+ assert [tool.name for tool in tools.tools] == ["ping"]
+ return await session.call_tool("ping", {})
+
+ pending: Final = client._execute_session_operation(_diagnostic_transport(transport, mode, "tools/list"), operation)
+ if mode == "ok":
+ result: Final = await asyncio.wait_for(pending, timeout=3)
+ assert result.isError is False
+ assert result.content[0].text == "pong"
+ logging_callback.assert_awaited_once_with(LoggingMessageNotificationParams(level="info", data="Listing tools"))
+ else:
+ with pytest.raises(McpError) as caught:
+ await asyncio.wait_for(pending, timeout=3)
+ if mode == "closed":
+ assert "connection was closed" in _connection_error_message(caught.value, client.server_url, 0.2)
+ else:
+ assert isinstance(as_mcp_read_timeout(caught.value), TimeoutError)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
+async def test_transport_cancellation_cleans_up_a_pending_request(transport: MCPTransport) -> None:
+ ready: Final = asyncio.Event()
+
+ async def on_log(message: LoggingMessageNotificationParams) -> None:
+ if message.data == "Waiting":
+ ready.set()
+
+ client: Final = MCPClient(
+ server_url="https://example.com/sse", transport_type=transport, timeout=30, logging_callback=on_log
+ )
+ task: Final = asyncio.create_task(
+ client._execute_session_operation(
+ _diagnostic_transport(transport, "silent", "tools/list"), lambda session: session.list_tools()
+ )
+ )
+ try:
+ await asyncio.wait_for(ready.wait(), timeout=3)
+ finally:
+ task.cancel()
+ with pytest.raises(asyncio.CancelledError):
+ await asyncio.wait_for(task, timeout=3)
+
+
+class _InterruptedHTTPBody(httpx.AsyncByteStream):
+ async def __aiter__(self) -> AsyncIterator[bytes]:
+ yield b'{"jsonrpc":'
+ raise httpx.RemoteProtocolError("secret-incomplete-response")
+
+
+@pytest.mark.asyncio
+async def test_interrupted_http_response_preserves_the_transport_failure() -> None:
+ def respond(request: httpx.Request) -> httpx.Response:
+ return httpx.Response(200, headers={"Content-Type": "application/json"}, stream=_InterruptedHTTPBody())
+
+ async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
+ client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
+ with pytest.raises(httpx.RemoteProtocolError, match="secret-incomplete-response"):
+ await asyncio.wait_for(
+ client._execute_session_operation(
+ streamable_http_client(client.server_url, http_client=http_client),
+ lambda session: session.list_tools(),
+ ),
+ timeout=3,
+ )
+
+
+@pytest.mark.asyncio
+async def test_empty_http_event_stream_uses_the_existing_request_deadline() -> None:
+ def respond(request: httpx.Request) -> httpx.Response:
+ return httpx.Response(200, headers={"Content-Type": "text/event-stream"}, content=b"")
+
+ async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
+ client: Final = MCPClient(server_url="https://example.com/mcp", timeout=0.2)
+ with pytest.raises(McpError) as caught:
+ await asyncio.wait_for(
+ client._execute_session_operation(
+ streamable_http_client(client.server_url, http_client=http_client),
+ lambda session: session.list_tools(),
+ ),
+ timeout=3,
+ )
+ assert isinstance(as_mcp_read_timeout(caught.value), TimeoutError)
diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py
index cd8d609cf71..ddc8439a83a 100644
--- a/tests/test_litellm/integrations/test_custom_guardrail.py
+++ b/tests/test_litellm/integrations/test_custom_guardrail.py
@@ -1842,12 +1842,14 @@ class _ApplyStyleGuardrail(CustomGuardrail):
self.block = block
self.apply_called = False
self.seen_texts = None
+ self.seen_request_data = None
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
from fastapi import HTTPException
self.apply_called = True
self.seen_texts = inputs.get("texts")
+ self.seen_request_data = request_data
if self.block:
raise HTTPException(status_code=400, detail={"error": "Violated moderation policy"})
return inputs
@@ -2646,6 +2648,91 @@ class TestCustomGuardrailPostCallSuccessDeploymentHook:
which starved every later callback in litellm.callbacks (notably the lazily-appended
VectorStorePreCallHook that attaches provider_specific_fields["search_results"])."""
+ @pytest.mark.asyncio
+ async def test_apply_guardrail_retains_request_identity(self) -> None:
+ from litellm.types.guardrails import GuardrailEventHooks
+ from litellm.types.utils import Choices, Message, ModelResponse
+
+ guardrail: Final = _ApplyStyleGuardrail(block=False)
+ guardrail.event_hook = GuardrailEventHooks.post_call
+ request_data: Final = {"guardrails": ["apply-style-guardrail"]}
+ response: Final = ModelResponse(choices=[Choices(message=Message(content="review me"))])
+
+ await guardrail.async_post_call_success_deployment_hook(
+ request_data=request_data, response=response, call_type=CallTypes.acompletion
+ )
+
+ assert guardrail.seen_request_data is request_data
+ assert guardrail.seen_texts == ["review me"]
+ assert "guardrail_to_apply" not in request_data
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("call_type", (None, CallTypes.acompletion))
+ async def test_apply_guardrail_masks_response_and_records_metadata(self, call_type: CallTypes | None) -> None:
+ from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
+ ContentFilterGuardrail,
+ )
+ from litellm.types.guardrails import BlockedWord, ContentFilterAction, GuardrailEventHooks
+ from litellm.types.utils import Choices, Message, ModelResponse
+
+ guardrail: Final = ContentFilterGuardrail(
+ guardrail_name="response-filter",
+ event_hook=GuardrailEventHooks.post_call,
+ blocked_words=[BlockedWord(keyword="secret", action=ContentFilterAction.MASK)],
+ )
+ request_data: Final = {"guardrails": ["response-filter"]}
+ response: Final = ModelResponse(choices=[Choices(message=Message(content="a secret"))])
+
+ result: Final = await guardrail.async_post_call_success_deployment_hook(
+ request_data=request_data, response=response, call_type=call_type
+ )
+
+ assert isinstance(result, ModelResponse)
+ assert result.choices[0].message.content == f"a {guardrail.keyword_redaction_tag}"
+ entries: Final = _guardrail_entries(request_data)
+ assert len(entries) == 1
+ assert entries[0]["guardrail_name"] == "response-filter"
+ assert entries[0]["guardrail_mode"] == "post_call"
+ assert "guardrail_to_apply" not in request_data
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("error_type", (None, RuntimeError, asyncio.CancelledError))
+ async def test_dispatch_cleans_up_request_on_every_exit(self, error_type: type[BaseException] | None) -> None:
+ from contextlib import nullcontext
+
+ from litellm.integrations.custom_logger import CustomLogger
+ from litellm.types.utils import LLMResponseTypes, ModelResponse
+
+ error: Final = error_type("dispatch interrupted") if error_type is not None else None
+
+ class Dispatch(CustomLogger):
+ request_data: dict[str, object] | None = None
+
+ async def async_post_call_success_hook(
+ self, data: dict[str, object], user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
+ ) -> LLMResponseTypes:
+ self.request_data = data
+ if error is not None:
+ raise error
+ return response
+
+ dispatch: Final = Dispatch()
+
+ class Guardrail(_ApplyStyleGuardrail):
+ def _deployment_hook_target(self) -> CustomLogger:
+ return dispatch
+
+ guardrail: Final = Guardrail(block=False)
+ guardrail.event_hook = GuardrailEventHooks.post_call
+ request_data: Final = {"guardrails": ["apply-style-guardrail"]}
+ with pytest.raises(error_type) if error_type is not None else nullcontext():
+ await guardrail.async_post_call_success_deployment_hook(
+ request_data=request_data, response=ModelResponse(), call_type=CallTypes.acompletion
+ )
+
+ assert dispatch.request_data is request_data
+ assert "guardrail_to_apply" not in request_data
+
@pytest.mark.asyncio
async def test_returns_none_when_request_has_no_guardrails(self):
from litellm.types.utils import ModelResponse
@@ -2740,4 +2827,5 @@ class TestCustomGuardrailPostCallSuccessDeploymentHook:
assert result is response
assert response.choices[0].message.content == "filtered response"
- assert request_data == {"guardrails": ["test-guardrail"]}
+ assert "guardrail_to_apply" not in request_data
+ assert len(_guardrail_entries(request_data)) == 1
diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py
index 65a6dd2a4ca..fbb9d178390 100644
--- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py
+++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py
@@ -1,9 +1,11 @@
import json
+from datetime import datetime, timezone
import pytest
from fastapi.testclient import TestClient
import litellm
+from litellm._internal_context import pinned_billing_time
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
StandardBuiltInToolCostTracking,
)
@@ -27,10 +29,10 @@ from litellm.types.utils import (
)
from litellm.litellm_core_utils.llm_cost_calc.utils import (
+ BilledTokenRates,
CostCalculatorUtils,
PromptTokensDetailsResult,
TokenRates,
- TokenTypeCostBreakdown,
_calculate_input_cost,
_get_token_base_cost,
_is_off_peak,
@@ -38,6 +40,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import (
apply_off_peak_pricing,
calculate_cache_writing_cost,
generic_cost_per_token,
+ get_billed_token_rates,
get_token_type_cost_breakdown,
)
from litellm.types.utils import CacheCreationTokenDetails, Usage
@@ -3906,6 +3909,200 @@ def test_token_type_cost_breakdown_reconciles_with_generic_total(_local_model_co
assert text_input_cost + breakdown.cache_read_cost == pytest.approx(prompt_cost)
+def _custom_priced_usage() -> Usage:
+ return Usage(
+ prompt_tokens=1000,
+ completion_tokens=500,
+ total_tokens=1500,
+ prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=800, cache_creation_tokens=100),
+ completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=200),
+ )
+
+
+def test_token_type_cost_breakdown_prices_custom_pricing_from_its_flat_rates():
+ """
+ A custom-priced deployment, usually absent from the cost map, used to get zero cache and
+ reasoning lines while its total already billed cache tokens at the custom cache rates.
+ The lines must come from the same flat rates: a configured cache rate, else the input
+ rate for cache tokens and the output rate for reasoning tokens.
+ """
+ from litellm.types.utils import CostPerToken
+
+ breakdown = get_token_type_cost_breakdown(
+ model="openai/onprem-model",
+ custom_llm_provider="openai",
+ usage=_custom_priced_usage(),
+ custom_cost_per_token=CostPerToken(
+ input_cost_per_token=1e-6, output_cost_per_token=2e-6, cache_read_input_token_cost=1e-7
+ ),
+ )
+
+ assert breakdown.cache_read_cost == pytest.approx(800 * 1e-7)
+ assert breakdown.cache_creation_cost == pytest.approx(100 * 1e-6)
+ assert breakdown.reasoning_cost == pytest.approx(200 * 2e-6)
+
+
+def test_token_type_cost_breakdown_reconciles_with_custom_pricing_totals():
+ from litellm.cost_calculator import cost_per_token
+ from litellm.types.utils import CostPerToken
+
+ usage = _custom_priced_usage()
+ custom_cost_per_token = CostPerToken(
+ input_cost_per_token=1e-6,
+ output_cost_per_token=2e-6,
+ cache_read_input_token_cost=1e-7,
+ cache_creation_input_token_cost=1.25e-6,
+ )
+
+ prompt_cost, completion_cost = cost_per_token(
+ model="openai/onprem-model",
+ custom_llm_provider="openai",
+ prompt_tokens=1000,
+ completion_tokens=500,
+ usage_object=usage,
+ custom_cost_per_token=custom_cost_per_token,
+ )
+ breakdown = get_token_type_cost_breakdown(
+ model="openai/onprem-model",
+ custom_llm_provider="openai",
+ usage=usage,
+ custom_cost_per_token=custom_cost_per_token,
+ )
+
+ assert 100 * 1e-6 + breakdown.cache_read_cost + breakdown.cache_creation_cost == pytest.approx(prompt_cost)
+ assert 300 * 2e-6 + breakdown.reasoning_cost == pytest.approx(completion_cost)
+
+
+def test_billed_token_rates_follow_the_token_tier_the_breakdown_bills_at(monkeypatch):
+ monkeypatch.setitem(
+ litellm.model_cost,
+ "tiered-cache-model",
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "cache_read_input_token_cost": 3e-7,
+ "cache_creation_input_token_cost": 3.75e-6,
+ "input_cost_per_token_above_200k_tokens": 6e-6,
+ "output_cost_per_token_above_200k_tokens": 3e-5,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-7,
+ "cache_creation_input_token_cost_above_200k_tokens": 7.5e-6,
+ "litellm_provider": "openai",
+ "mode": "chat",
+ },
+ )
+ usage = Usage(
+ prompt_tokens=250_000,
+ completion_tokens=1_000,
+ total_tokens=251_000,
+ prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=200_000, cache_creation_tokens=10_000),
+ completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=200),
+ )
+
+ rates = get_billed_token_rates(model="tiered-cache-model", custom_llm_provider="openai", usage=usage)
+ breakdown = get_token_type_cost_breakdown(model="tiered-cache-model", custom_llm_provider="openai", usage=usage)
+
+ assert rates == BilledTokenRates(
+ input_cost_per_token=6e-6,
+ output_cost_per_token=3e-5,
+ cache_read_input_token_cost=6e-7,
+ cache_creation_input_token_cost=7.5e-6,
+ cache_creation_input_token_cost_above_1hr=0.0,
+ output_cost_per_reasoning_token=3e-5,
+ )
+ assert breakdown.cache_read_cost == pytest.approx(200_000 * rates.cache_read_input_token_cost)
+ assert breakdown.cache_creation_cost == pytest.approx(10_000 * rates.cache_creation_input_token_cost)
+ assert breakdown.reasoning_cost == pytest.approx(200 * rates.output_cost_per_reasoning_token)
+
+
+def test_a_pinned_billing_time_prices_the_totals_and_the_reported_rates_at_one_moment(monkeypatch):
+ """Totals and reported rates resolve off-peak pricing on separate paths that each read the
+ clock, so a window opening between the two reads used to leave them describing one request
+ at two different prices. Pinned, both must answer for the pinned moment."""
+ monkeypatch.setitem(
+ litellm.model_cost,
+ "off-peak-model",
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "off_peak_pricing": {
+ "hours_utc": "02:00-03:00",
+ "input_cost_per_token": 1e-6,
+ "output_cost_per_token": 5e-6,
+ },
+ "litellm_provider": "openai",
+ "mode": "chat",
+ },
+ )
+ usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
+
+ with pinned_billing_time(datetime(2026, 1, 1, 2, 30, tzinfo=timezone.utc)):
+ off_peak_prompt_cost, off_peak_completion_cost = generic_cost_per_token(
+ model="off-peak-model", usage=usage, custom_llm_provider="openai"
+ )
+ off_peak_rates = get_billed_token_rates(model="off-peak-model", custom_llm_provider="openai", usage=usage)
+ with pinned_billing_time(datetime(2026, 1, 1, 12, 30, tzinfo=timezone.utc)):
+ peak_prompt_cost, peak_completion_cost = generic_cost_per_token(
+ model="off-peak-model", usage=usage, custom_llm_provider="openai"
+ )
+ peak_rates = get_billed_token_rates(model="off-peak-model", custom_llm_provider="openai", usage=usage)
+
+ assert off_peak_rates.input_cost_per_token == pytest.approx(1e-6)
+ assert peak_rates.input_cost_per_token == pytest.approx(3e-6)
+ assert off_peak_prompt_cost == pytest.approx(1000 * off_peak_rates.input_cost_per_token)
+ assert off_peak_completion_cost == pytest.approx(500 * off_peak_rates.output_cost_per_token)
+ assert peak_prompt_cost == pytest.approx(1000 * peak_rates.input_cost_per_token)
+ assert peak_completion_cost == pytest.approx(500 * peak_rates.output_cost_per_token)
+
+
+def test_the_token_type_breakdown_carries_the_rates_it_billed_at(monkeypatch):
+ """Callers that report both the lines and the rates read the rates off the breakdown rather than
+ resolving them a second time, so the breakdown has to hand back exactly what it billed at."""
+ monkeypatch.setitem(
+ litellm.model_cost,
+ "xai/tiered-model",
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "cache_read_input_token_cost": 3e-7,
+ "input_cost_per_token_above_200k_tokens": 6e-6,
+ "output_cost_per_token_above_200k_tokens": 3e-5,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-7,
+ "litellm_provider": "xai",
+ "mode": "chat",
+ },
+ )
+ usage = Usage(
+ prompt_tokens=200_000,
+ completion_tokens=1_000,
+ total_tokens=201_000,
+ prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=100_000),
+ )
+
+ breakdown = get_token_type_cost_breakdown(model="xai/tiered-model", custom_llm_provider="xai", usage=usage)
+
+ assert breakdown.rates == get_billed_token_rates(
+ model="xai/tiered-model", custom_llm_provider="xai", usage=usage
+ )
+ assert breakdown.rates.cache_read_input_token_cost == pytest.approx(6e-7)
+ assert breakdown.cache_read_cost == pytest.approx(100_000 * breakdown.rates.cache_read_input_token_cost)
+
+
+def test_the_token_type_breakdown_reports_no_rates_for_an_unpriced_model():
+ usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
+
+ breakdown = get_token_type_cost_breakdown(
+ model="no-such-model-anywhere", custom_llm_provider="openai", usage=usage
+ )
+
+ assert breakdown.rates is None
+
+
+def test_billed_token_rates_are_none_for_an_unpriced_model():
+ usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
+
+ assert get_billed_token_rates(model="no-such-model-anywhere", custom_llm_provider="openai", usage=usage) is None
+
+
def test_token_type_cost_breakdown_zero_without_special_tokens(_local_model_cost_map):
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
@@ -3913,9 +4110,7 @@ def test_token_type_cost_breakdown_zero_without_special_tokens(_local_model_cost
model="gpt-4o", custom_llm_provider="openai", usage=usage
)
- assert breakdown == TokenTypeCostBreakdown(
- reasoning_cost=0.0, cache_read_cost=0.0, cache_creation_cost=0.0
- )
+ assert (breakdown.reasoning_cost, breakdown.cache_read_cost, breakdown.cache_creation_cost) == (0.0, 0.0, 0.0)
@pytest.mark.parametrize(
@@ -3987,9 +4182,7 @@ def test_token_type_cost_breakdown_handles_unknown_model_gracefully():
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=5),
),
)
- assert breakdown == TokenTypeCostBreakdown(
- reasoning_cost=0.0, cache_read_cost=0.0, cache_creation_cost=0.0
- )
+ assert (breakdown.reasoning_cost, breakdown.cache_read_cost, breakdown.cache_creation_cost) == (0.0, 0.0, 0.0)
def test_token_type_cost_breakdown_applies_regional_uplift(_local_model_cost_map):
diff --git a/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py b/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py
index 0495440c51c..266c2ca1465 100644
--- a/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py
+++ b/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py
@@ -6,6 +6,7 @@ count actual model entries, not reserved meta keys) and the extraction of the
import json
import os
+import threading
import pytest
@@ -26,9 +27,7 @@ from litellm.litellm_core_utils.get_model_cost_map import (
def _load_root_cost_map() -> dict:
- path = os.path.join(
- os.path.dirname(__file__), "../../../model_prices_and_context_window.json"
- )
+ path = os.path.join(os.path.dirname(__file__), "../../../model_prices_and_context_window.json")
with open(path) as f:
return json.load(f)
@@ -44,9 +43,7 @@ def test_git_blob_id_is_what_git_hash_object_prints():
def _make_models(n: int) -> dict:
- return {
- f"model-{i}": {"litellm_provider": "openai", "mode": "chat"} for i in range(n)
- }
+ return {f"model-{i}": {"litellm_provider": "openai", "mode": "chat"} for i in range(n)}
def test_count_model_entries_excludes_reserved_keys():
@@ -129,9 +126,7 @@ def test_finalize_pops_key_and_installs_rules():
def test_finalize_with_no_block_clears_rules():
previous = list(get_fallback_generalization_rules())
try:
- set_fallback_generalizations(
- [{"name": "stale", "pattern": r"^x", "model_info": {"a": 1}}]
- )
+ set_fallback_generalizations([{"name": "stale", "pattern": r"^x", "model_info": {"a": 1}}])
_finalize_model_cost_map(_make_models(2))
assert match_capability_generalizations("x-1") is None
finally:
@@ -317,9 +312,7 @@ def test_get_model_cost_map_stamps_loaded_at():
from litellm.litellm_core_utils import get_model_cost_map as module
- client, _calls = _mock_client(
- [httpx.Response(200, content=_real_map_bytes())], client_cls=httpx.Client
- )
+ client, _calls = _mock_client([httpx.Response(200, content=_real_map_bytes())], client_cls=httpx.Client)
before = datetime.now(timezone.utc)
module.get_model_cost_map(url="https://example.invalid/cost_map.json", client=client)
@@ -328,6 +321,7 @@ def test_get_model_cost_map_stamps_loaded_at():
assert loaded_at is not None
assert before <= loaded_at <= datetime.now(timezone.utc)
+
# ---------------------------------------------------------------------------
# refetch_model_cost_map: retry/backoff behavior for runtime reloads
# ---------------------------------------------------------------------------
@@ -394,9 +388,7 @@ async def test_refetch_retries_429_honoring_retry_after():
]
)
sleeper = _SleepRecorder()
- result = await refetch_model_cost_map(
- url=_URL, sleep=sleeper, rng=random.Random(0), client=client
- )
+ result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert isinstance(result, ModelCostMapReloaded)
assert len(result.model_cost_map) > 100
assert calls["count"] == 3
@@ -408,9 +400,7 @@ async def test_refetch_gives_up_after_max_attempts_with_exponential_backoff():
"""All 429 without Retry-After: exponential backoff waits, then a failure value."""
client, calls = _mock_client([httpx.Response(429)])
sleeper = _SleepRecorder()
- result = await refetch_model_cost_map(
- url=_URL, sleep=sleeper, rng=random.Random(0), client=client
- )
+ result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert isinstance(result, ModelCostMapReloadUnavailable)
assert "429" in result.reason
assert "after 3 attempts" in result.reason
@@ -430,9 +420,7 @@ async def test_refetch_caps_retry_after_wait():
]
)
sleeper = _SleepRecorder()
- result = await refetch_model_cost_map(
- url=_URL, sleep=sleeper, rng=random.Random(0), client=client
- )
+ result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert isinstance(result, ModelCostMapReloaded)
assert sleeper.waits == [30.0]
@@ -447,9 +435,7 @@ async def test_refetch_retries_transport_errors():
]
)
sleeper = _SleepRecorder()
- result = await refetch_model_cost_map(
- url=_URL, sleep=sleeper, rng=random.Random(0), client=client
- )
+ result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert isinstance(result, ModelCostMapReloaded)
assert calls["count"] == 2
assert len(sleeper.waits) == 1
@@ -460,9 +446,7 @@ async def test_refetch_non_retryable_status_fails_immediately():
"""A 404 is permanent: one attempt, no sleeps, failure value."""
client, calls = _mock_client([httpx.Response(404)])
sleeper = _SleepRecorder()
- result = await refetch_model_cost_map(
- url=_URL, sleep=sleeper, rng=random.Random(0), client=client
- )
+ result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert isinstance(result, ModelCostMapReloadUnavailable)
assert "404" in result.reason
assert calls["count"] == 1
@@ -473,9 +457,7 @@ async def test_refetch_non_retryable_status_fails_immediately():
async def test_refetch_invalid_json_fails_immediately():
client, calls = _mock_client([httpx.Response(200, content=b"not json")])
sleeper = _SleepRecorder()
- result = await refetch_model_cost_map(
- url=_URL, sleep=sleeper, rng=random.Random(0), client=client
- )
+ result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert isinstance(result, ModelCostMapReloadUnavailable)
assert "invalid JSON" in result.reason
assert calls["count"] == 1
@@ -487,9 +469,7 @@ async def test_refetch_shrunk_map_fails_integrity_not_swapped_in():
"""A drastically shrunk upstream file is rejected instead of being adopted."""
tiny = json.dumps(_make_models(60)).encode()
client, _calls = _mock_client([httpx.Response(200, content=tiny)])
- result = await refetch_model_cost_map(
- url=_URL, sleep=_SleepRecorder(), rng=random.Random(0), client=client
- )
+ result = await refetch_model_cost_map(url=_URL, sleep=_SleepRecorder(), rng=random.Random(0), client=client)
assert isinstance(result, ModelCostMapReloadUnavailable)
assert "integrity validation" in result.reason
@@ -586,69 +566,125 @@ from litellm.litellm_core_utils.get_model_cost_map import (
class _SyncSleepRecorder:
"""Injected in place of time.sleep so the boot path's waits are asserted without delay."""
- def __init__(self):
+ def __init__(self, block=False):
self.waits = []
+ self.block = block
+ self.started = threading.Event()
+ self.release = threading.Event()
def __call__(self, seconds: float) -> None:
+ if self.block:
+ self.started.set()
+ self.release.wait(timeout=10)
self.waits.append(seconds)
-def test_boot_load_retries_transient_failures_instead_of_falling_back():
- """A refused connection then a 503 at pod boot used to pin the process to the bundled
- backup for its lifetime; both are transient and must be retried before giving up."""
+def _retry_threads():
+ return [thread for thread in threading.enumerate() if thread.name == "litellm-model-cost-map-retry"]
+
+
+def test_boot_load_success_does_not_start_background_retry():
+ client, calls = _mock_client([httpx.Response(200, content=_real_map_bytes())], client_cls=httpx.Client)
+ sleeper = _SyncSleepRecorder()
+
+ cost_map = get_model_cost_map(
+ url=_URL,
+ sleep=sleeper,
+ rng=random.Random(0),
+ client=client,
+ )
+ assert calls["count"] == 1
+ assert sleeper.waits == []
+ assert _retry_threads() == []
+ assert cost_map.keys() >= _load_root_cost_map().keys() - {"sample_spec", FALLBACK_GENERALIZATIONS_KEY}
+ assert get_model_cost_map_source_info()["source"] == "remote"
+
+
+def test_boot_load_transient_failure_returns_local_then_background_retry_adopts_remote(monkeypatch):
+ import litellm
+ from litellm import utils as litellm_utils
+ from litellm.litellm_core_utils import get_model_cost_map as module
+
+ original_model_cost = litellm.model_cost
+ monkeypatch.setattr(litellm, "model_cost", dict(original_model_cost))
+ for name, provider_models in tuple(vars(litellm).items()):
+ if name.endswith("_models") and isinstance(provider_models, set):
+ monkeypatch.setattr(litellm, name, set(provider_models))
+ monkeypatch.setattr(litellm, "models_by_provider", dict(litellm.models_by_provider))
+ monkeypatch.setattr(
+ litellm_utils,
+ "_runtime_registered_model_cost",
+ dict(litellm_utils._runtime_registered_model_cost),
+ )
+ source_info = module._cost_map_source_info
+ for name in ("source", "url", "is_env_forced", "fallback_reason", "loaded_at", "source_revision", "etag"):
+ monkeypatch.setattr(source_info, name, getattr(source_info, name))
+
+ remote_map = _load_root_cost_map()
+ remote_map["claude-remote-only-test"] = {"litellm_provider": "anthropic", "mode": "chat"}
client, calls = _mock_client(
[
httpx.ConnectError("connection refused"),
- httpx.Response(503),
- httpx.Response(200, content=_real_map_bytes()),
+ httpx.Response(200, content=json.dumps(remote_map).encode()),
],
client_cls=httpx.Client,
)
- sleeper = _SyncSleepRecorder()
+ sleeper = _SyncSleepRecorder(block=True)
+ litellm.register_model({"my-runtime-model": {"litellm_provider": "custom", "max_input_tokens": 4321}})
- cost_map = get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
-
- assert calls["count"] == 3
- assert len(sleeper.waits) == 2
- assert 2.0 <= sleeper.waits[0] < 3.0
- assert 4.0 <= sleeper.waits[1] < 5.0
- source = get_model_cost_map_source_info()
- assert source["source"] == "remote"
- assert source["fallback_reason"] is None
- assert cost_map.keys() >= _load_root_cost_map().keys() - {"sample_spec", FALLBACK_GENERALIZATIONS_KEY}
-
-
-def test_boot_load_honors_retry_after_then_falls_back_after_max_attempts():
- """An outage longer than the retry budget still ends on the bundled backup, and the
- recorded fallback reason says how many attempts were spent so operators can tell."""
- client, calls = _mock_client(
- [httpx.Response(429, headers={"Retry-After": "7"})], client_cls=httpx.Client
+ cost_map = get_model_cost_map(
+ url=_URL,
+ max_attempts=3,
+ sleep=sleeper,
+ rng=random.Random(0),
+ client=client,
)
- sleeper = _SyncSleepRecorder()
- cost_map = get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
-
- assert calls["count"] == 3
- assert sleeper.waits == [7.0, 7.0]
- source = get_model_cost_map_source_info()
- assert source["source"] == "local"
- assert "after 3 attempts" in source["fallback_reason"]
- assert len(cost_map) > 100
+ assert calls["count"] == 1
+ assert sleeper.waits == []
+ assert sleeper.started.wait(timeout=10)
+ threads = _retry_threads()
+ try:
+ assert len(threads) == 1
+ assert "claude-remote-only-test" not in cost_map
+ assert cost_map.keys() == _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map()).keys()
+ source = get_model_cost_map_source_info()
+ assert source["source"] == "local"
+ assert source["fallback_reason"].startswith("Remote fetch failed:")
+ sleeper.release.set()
+ for thread in threads:
+ thread.join(timeout=10)
+ assert all(not thread.is_alive() for thread in threads)
+ assert sleeper.waits and 2.0 <= sleeper.waits[0] < 3.0
+ assert calls["count"] == 2
+ assert "claude-remote-only-test" in litellm.model_cost
+ assert "claude-remote-only-test" in litellm.anthropic_models
+ assert "my-runtime-model" in litellm.model_cost
+ source = get_model_cost_map_source_info()
+ assert source["source"] == "remote"
+ assert source["fallback_reason"] is None
+ finally:
+ sleeper.release.set()
+ for thread in _retry_threads():
+ thread.join(timeout=10)
-def test_boot_load_does_not_retry_permanent_failures():
- """A 404 or a malformed URL cannot heal by waiting: one attempt, no sleeps, backup."""
+def test_boot_load_does_not_retry_non_retryable_failure():
client, calls = _mock_client([httpx.Response(404)], client_cls=httpx.Client)
sleeper = _SyncSleepRecorder()
- get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
+ get_model_cost_map(
+ url=_URL,
+ sleep=sleeper,
+ rng=random.Random(0),
+ client=client,
+ )
assert calls["count"] == 1
assert sleeper.waits == []
- assert get_model_cost_map_source_info()["source"] == "local"
-
- get_model_cost_map(url="not a url", sleep=sleeper, rng=random.Random(0))
- assert sleeper.waits == []
- assert get_model_cost_map_source_info()["source"] == "local"
+ assert _retry_threads() == []
+ source = get_model_cost_map_source_info()
+ assert source["source"] == "local"
+ assert source["fallback_reason"] is not None
def test_boot_load_respects_local_env_override(monkeypatch):
@@ -701,7 +737,9 @@ def test_boot_load_that_fails_the_integrity_check_reports_the_backup_not_the_rej
)
get_model_cost_map(url=_URL, sleep=_SyncSleepRecorder(), rng=random.Random(0), client=remote)
shrunk_body = b'{"gpt-5.4-mini": {"mode": "chat", "input_cost_per_token": 1e-06, "output_cost_per_token": 2e-06}}'
- shrunk, _ = _mock_client([httpx.Response(200, headers={"ETag": 'W/"shrunk"'}, content=shrunk_body)], client_cls=httpx.Client)
+ shrunk, _ = _mock_client(
+ [httpx.Response(200, headers={"ETag": 'W/"shrunk"'}, content=shrunk_body)], client_cls=httpx.Client
+ )
get_model_cost_map(url=_URL, sleep=_SyncSleepRecorder(), rng=random.Random(0), client=shrunk)
diff --git a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py
index 541c0db15d8..3b01a4f2054 100644
--- a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py
+++ b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py
@@ -5,6 +5,8 @@ Test bedrock files transformation functionality
import json
import os
from collections.abc import Mapping
+from contextlib import AsyncExitStack, closing
+from typing import Final
from unittest.mock import MagicMock
from urllib.parse import unquote, urlparse
@@ -1855,6 +1857,104 @@ class TestBedrockBatchNonChatEndpointRecords:
]
+class TestBedrockFileDeletion:
+ S3_URI: Final = "s3://my-bucket/litellm-bedrock-files-model-abc.jsonl"
+ URL: Final = "https://s3.us-west-2.amazonaws.com/my-bucket/litellm-bedrock-files-model-abc.jsonl"
+
+ def test_interleaved_deletions_keep_their_own_file_ids(self, monkeypatch: pytest.MonkeyPatch) -> None:
+ import httpx
+
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ config: Final = BedrockFilesConfig()
+ params: Final = {
+ "aws_access_key_id": "AKIAEXAMPLE",
+ "aws_secret_access_key": "test-secret",
+ "aws_region_name": "us-west-2",
+ }
+ file_ids: Final = (self.S3_URI, "s3://my-bucket/litellm-bedrock-files-model-second.jsonl")
+ for file_id in file_ids:
+ config.transform_delete_file_request(file_id=file_id, optional_params={}, litellm_params=params)
+
+ deleted: Final = tuple(
+ config.transform_delete_file_response(
+ raw_response=httpx.Response(204),
+ logging_obj=MagicMock(model_call_details={"additional_args": {"file_id": file_id}}),
+ litellm_params=params,
+ ).id
+ for file_id in file_ids
+ )
+
+ assert deleted == file_ids
+
+ def test_delete_file_sends_signed_delete_and_returns_matching_id(self, monkeypatch: pytest.MonkeyPatch) -> None:
+ import httpx
+ import respx
+
+ import litellm
+ from litellm.llms.custom_httpx.http_handler import HTTPHandler
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ with respx.mock, closing(HTTPHandler()) as client:
+ route: Final = respx.delete(self.URL).mock(return_value=httpx.Response(204))
+ deleted: Final = litellm.file_delete(
+ file_id=self.S3_URI, custom_llm_provider="bedrock", client=client,
+ aws_access_key_id="AKIAEXAMPLE", aws_secret_access_key="test-secret", aws_region_name="us-west-2",
+ )
+ assert route.call_count == 1
+ request: Final = route.calls[0].request
+ assert request.content == b""
+ signed: Final = AWSRequest(method="DELETE", url=self.URL, headers={
+ "X-Amz-Date": request.headers["X-Amz-Date"],
+ "X-Amz-Content-SHA256": request.headers["X-Amz-Content-SHA256"],
+ })
+ signed.context["timestamp"] = request.headers["X-Amz-Date"]
+ auth: Final = S3SigV4Auth(Credentials("AKIAEXAMPLE", "test-secret"), "s3", "us-west-2")
+ signature: Final = auth.signature(auth.string_to_sign(signed, auth.canonical_request(signed)), signed)
+ assert request.headers["Authorization"].endswith(f"Signature={signature}")
+ assert deleted.id == self.S3_URI and deleted.deleted is True
+
+ @pytest.mark.asyncio
+ async def test_adelete_file_propagates_s3_errors(self, monkeypatch: pytest.MonkeyPatch) -> None:
+ import httpx
+ import respx
+
+ import litellm
+ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
+ async with AsyncExitStack() as stack:
+ client: Final = AsyncHTTPHandler()
+ stack.push_async_callback(client.close)
+ with respx.mock:
+ route: Final = respx.delete(self.URL).mock(
+ return_value=httpx.Response(403, content=b"AccessDenied")
+ )
+ from litellm.llms.bedrock.common_utils import BedrockError
+
+ with pytest.raises(BedrockError, match="AccessDenied"):
+ await litellm.afile_delete(
+ file_id=self.S3_URI, custom_llm_provider="bedrock", client=client,
+ aws_access_key_id="AKIAEXAMPLE", aws_secret_access_key="test-secret", aws_region_name="us-west-2",
+ )
+ assert route.call_count == 1
+
+ @pytest.mark.parametrize("file_id, message", [
+ ("s3://other-bucket/litellm-bedrock-files-model-abc.jsonl", "configured storage bucket"),
+ ("s3://my-bucket/private/data.jsonl", "LiteLLM-managed"),
+ ])
+ def test_delete_rejects_untrusted_objects_before_signing(
+ self, file_id: str, message: str, monkeypatch: pytest.MonkeyPatch
+ ) -> None:
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ with pytest.raises(ValueError, match=message):
+ BedrockFilesConfig().transform_delete_file_request(file_id=file_id, optional_params={}, litellm_params={})
+
+
class TestBedrockFileContentTransformation:
"""SigV4-signed S3 GetObject retrieval of Bedrock batch output files."""
@@ -1873,7 +1973,7 @@ class TestBedrockFileContentTransformation:
import hashlib
from litellm.llms.bedrock.files.transformation import (
- S3_SIGNED_GET_HEADERS_PARAM,
+ S3_SIGNED_REQUEST_HEADERS_PARAM,
BedrockFilesConfig,
)
@@ -1889,7 +1989,7 @@ class TestBedrockFileContentTransformation:
assert url == self.EXPECTED_URL
assert params == {}
- signed_headers = litellm_params[S3_SIGNED_GET_HEADERS_PARAM]
+ signed_headers = litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]
content_hashes = {
value
for name, value in signed_headers.items()
@@ -2139,7 +2239,7 @@ class TestBedrockFileContentTransformation:
def test_s3_region_name_wins_for_content_signing(self, monkeypatch):
"""s3_region_name must override aws_region_name for both the URL and the signature."""
from litellm.llms.bedrock.files.transformation import (
- S3_SIGNED_GET_HEADERS_PARAM,
+ S3_SIGNED_REQUEST_HEADERS_PARAM,
BedrockFilesConfig,
)
@@ -2154,17 +2254,17 @@ class TestBedrockFileContentTransformation:
)
assert url.startswith("https://s3.eu-west-1.amazonaws.com/")
- authorization = litellm_params[S3_SIGNED_GET_HEADERS_PARAM]["Authorization"]
+ authorization = litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]["Authorization"]
assert "/eu-west-1/s3/aws4_request" in authorization
def test_validate_environment_merges_and_pops_signed_get_headers(self):
from litellm.llms.bedrock.files.transformation import (
- S3_SIGNED_GET_HEADERS_PARAM,
+ S3_SIGNED_REQUEST_HEADERS_PARAM,
BedrockFilesConfig,
)
litellm_params = {
- S3_SIGNED_GET_HEADERS_PARAM: {"Authorization": "AWS4-HMAC-SHA256 test"}
+ S3_SIGNED_REQUEST_HEADERS_PARAM: {"Authorization": "AWS4-HMAC-SHA256 test"}
}
headers = BedrockFilesConfig().validate_environment(
@@ -2179,7 +2279,7 @@ class TestBedrockFileContentTransformation:
"x-custom": "kept",
"Authorization": "AWS4-HMAC-SHA256 test",
}
- assert S3_SIGNED_GET_HEADERS_PARAM not in litellm_params
+ assert S3_SIGNED_REQUEST_HEADERS_PARAM not in litellm_params
def test_transform_file_content_response_wraps_binary_content(self):
import httpx
@@ -2379,7 +2479,7 @@ class TestBedrockFilesS3SignatureEncoding:
self, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.llms.bedrock.files.transformation import (
- S3_SIGNED_GET_HEADERS_PARAM,
+ S3_SIGNED_REQUEST_HEADERS_PARAM,
BedrockFilesConfig,
)
@@ -2402,7 +2502,7 @@ class TestBedrockFilesS3SignatureEncoding:
method="GET",
url=url,
body=None,
- headers=litellm_params[S3_SIGNED_GET_HEADERS_PARAM],
+ headers=litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM],
)
@@ -2457,7 +2557,7 @@ def test_sign_s3_request_assumes_role_with_external_id(monkeypatch):
assert "ASIAFILESPUTROLE" in authorization
-def test_sign_s3_get_request_assumes_role_with_external_id(monkeypatch):
+def test_sign_s3_request_without_body_assumes_role_with_external_id(monkeypatch):
"""A trust policy requiring sts:ExternalId must be satisfied when signing the S3 download request."""
import datetime
from unittest.mock import patch
@@ -2504,7 +2604,7 @@ def test_sign_s3_get_request_assumes_role_with_external_id(monkeypatch):
assert request_params.aws_external_id == "external-id-files-get"
with patch.object(boto3, "client", return_value=FakeSTSClient()):
- signed_headers = BedrockFilesConfig()._sign_s3_get_request(
+ signed_headers = BedrockFilesConfig()._sign_s3_request_without_body(
api_base="https://s3.us-east-1.amazonaws.com/safe-bucket/litellm-bedrock-files-model-id-abc.jsonl",
aws_region_name="us-east-1",
request_params=request_params,
diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
index f854d806bdc..a1c28f36e70 100644
--- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
+++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py
@@ -38,6 +38,14 @@ def flush_shared_bedrock_iam_cache():
yield
+@pytest.fixture(autouse=True)
+def _clean_ssl_env(monkeypatch):
+ """get_ssl_verify reads these, so the sts client's verify= would otherwise depend on
+ the ambient environment. The published images set SSL_CERT_FILE."""
+ for env_var in ("SSL_CERT_FILE", "SSL_VERIFY"):
+ monkeypatch.delenv(env_var, raising=False)
+
+
def test_base_aws_llm_instances_share_process_wide_iam_cache():
"""Regression LIT-2662: new instances must reuse iam_cache (Bedrock passthrough is per-request)."""
first = BaseAWSLLM()
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py
index c0d055edb7f..669e094fee4 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py
@@ -1,10 +1,11 @@
"""
-Tests for partial-update semantics of PUT /v1/mcp/server.
+Tests for partial-update semantics of PUT /v1/mcp/server and PUT /v1/mcp/toolset.
A partial update must only write the fields the caller explicitly provided.
Omitting a field must NOT reset it to its Pydantic schema default (e.g.
``transport=sse``, ``mcp_access_groups=[]``, ``allow_all_keys=False``), which
-would silently overwrite the existing DB row.
+would silently overwrite the existing DB row, and a field the caller sent as null
+must be cleared rather than left at its stored value.
"""
import json
@@ -850,3 +851,69 @@ async def test_cf_pair_switch_does_not_clear_dcr_bridge():
data = UpdateMCPServerRequest(server_id="s", auth_type="oauth_delegate")
data_dict = await _run_update_with_existing(data, existing_auth_type="true_passthrough")
assert "dcr_bridge" not in data_dict
+
+
+def _mock_toolset_prisma():
+ """A prisma double whose update answers with a row the reader can expand, so the
+ call under test returns instead of failing inside the row mapper."""
+ updated_row = MagicMock()
+ updated_row.model_dump.return_value = {
+ "toolset_id": "ts-1",
+ "toolset_name": "ops",
+ "description": None,
+ "tools": "[]",
+ }
+ mock_prisma = MagicMock()
+ mock_prisma.db.litellm_mcptoolsettable = AsyncMock()
+ mock_prisma.db.litellm_mcptoolsettable.update = AsyncMock(return_value=updated_row)
+ return mock_prisma
+
+
+async def _run_toolset_update(payload: dict) -> dict:
+ """The columns PUT /v1/mcp/toolset writes for this payload, minus the audit stamp
+ every write carries. The prisma double is injected, so nothing is patched."""
+ from litellm.proxy._experimental.mcp_server.toolset_db import update_mcp_toolset
+ from litellm.types.mcp_server.mcp_toolset import UpdateMCPToolsetRequest
+
+ mock_prisma = _mock_toolset_prisma()
+ await update_mcp_toolset(mock_prisma, UpdateMCPToolsetRequest.model_validate(payload), "test-user")
+ written = dict(mock_prisma.db.litellm_mcptoolsettable.update.call_args[1]["data"])
+ assert written["updated_by"] == "test-user"
+ return {name: value for name, value in written.items() if name != "updated_by"}
+
+
+@pytest.mark.asyncio
+async def test_toolset_partial_update_clears_description_on_explicit_null():
+ """The dump used to drop None, so a null description could never clear the stored
+ one: the toolset kept a description its owner had deleted."""
+ assert await _run_toolset_update({"toolset_id": "ts-1", "description": None}) == {"description": None}
+
+
+@pytest.mark.asyncio
+async def test_toolset_partial_update_omits_the_fields_the_caller_left_out():
+ tools = [{"server_id": "s1", "tool_name": "alpha"}]
+ assert await _run_toolset_update({"toolset_id": "ts-1", "tools": tools}) == {"tools": json.dumps(tools)}
+
+
+@pytest.mark.asyncio
+async def test_toolset_partial_update_ignores_null_tools_rather_than_revoking_them():
+ """A client that sends tools=null means "leave the selection alone", so the grants
+ survive. Clearing them is an explicit [], which cannot be confused with an omitted
+ field; treating null as a clear would silently revoke every tool the toolset grants."""
+ assert await _run_toolset_update({"toolset_id": "ts-1", "tools": None, "description": "kept"}) == {
+ "description": "kept"
+ }
+
+
+@pytest.mark.asyncio
+async def test_toolset_partial_update_empties_the_selection_on_an_explicit_empty_list():
+ assert await _run_toolset_update({"toolset_id": "ts-1", "tools": []}) == {"tools": "[]"}
+
+
+@pytest.mark.asyncio
+async def test_toolset_partial_update_ignores_a_null_name():
+ """A toolset always has a name, so a null toolset_name is a no-op, not a clear
+ that would write a NOT NULL column to null."""
+ assert await _run_toolset_update({"toolset_id": "ts-1", "toolset_name": None, "description": "kept"}) == {
+ "description": "kept"
+ }
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py
index bd692776c82..3487f634251 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py
@@ -3,7 +3,7 @@ import inspect
import json
import sys
from datetime import datetime
-from typing import Any, Dict, Optional
+from typing import Any, Dict, Final, Optional
from unittest.mock import AsyncMock, MagicMock
if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11
@@ -113,7 +113,7 @@ class TestExecuteWithMcpClient:
assert "stack_trace" not in result
@pytest.mark.asyncio
- async def test_timeout_caps_hanging_operation_and_names_url(self, monkeypatch):
+ async def test_timeout_caps_hanging_operation_and_names_origin(self, monkeypatch):
async def fake_create_client(*args, **kwargs):
return object()
@@ -138,7 +138,7 @@ class TestExecuteWithMcpClient:
)
assert result["error"] is True
- assert "https://mcp.example.com/mcp/" in result["message"]
+ assert "https://mcp.example.com" in result["message"]
@pytest.mark.asyncio
async def test_timeout_covers_client_creation(self, monkeypatch):
@@ -166,15 +166,15 @@ class TestExecuteWithMcpClient:
)
assert result["error"] is True
- assert "https://mcp.example.com/mcp/" in result["message"]
+ assert "https://mcp.example.com" in result["message"]
def test_timeout_defaults_to_tool_listing_timeout(self):
default = inspect.signature(rest_endpoints._execute_with_mcp_client).parameters["timeout_seconds"].default
assert default == MCP_TOOL_LISTING_TIMEOUT
- def test_connection_error_message_timeout_names_url_and_budget(self):
+ def test_connection_error_message_timeout_names_origin_and_budget(self):
message = rest_endpoints._connection_error_message(TimeoutError(), "https://api.example.com/mcp/", 30.0)
- assert "https://api.example.com/mcp/" in message
+ assert "https://api.example.com" in message
assert "30s" in message
def test_connection_error_message_hides_arbitrary_http_exception_detail(self):
@@ -592,7 +592,7 @@ class TestExecuteWithMcpClient:
assert result["status"] == "error"
assert result["error"] is True
- assert "Failed to connect to MCP server" in result["message"]
+ assert "reference" in result["message"]
# Error message must not leak raw exception details
assert "cancel scope" not in result["message"]
@@ -3427,10 +3427,191 @@ class TestConnectionErrorMessage:
message = rest_endpoints._connection_error_message(exc, "https://example.com", 30.0)
assert "503" in message
+ @pytest.mark.parametrize(
+ "error_type", [httpx.ReadError, httpx.WriteError, httpx.RemoteProtocolError, ConnectionResetError]
+ )
+ def test_interrupted_connection_message_is_safe(self, error_type: type[Exception]) -> None:
+ message: Final = rest_endpoints._connection_error_message(
+ error_type("secret-transport-detail"), "https://example.com/?token=secret-query", 30
+ )
+ assert "connection was interrupted" in message
+ assert "secret" not in message
+
+ def test_closed_connection_explains_incomplete_request(self) -> None:
+ from mcp import McpError
+ from mcp.types import ErrorData
+
+ message: Final = rest_endpoints._connection_error_message(
+ McpError(ErrorData(code=-32000, message="Connection closed", data="secret-data")), None, 30
+ )
+ assert "connection was closed before the request completed" in message
+ assert "secret" not in message
+
+ def test_timeout_does_not_claim_the_server_sent_nothing(self) -> None:
+ message: Final = rest_endpoints._connection_error_message(TimeoutError(), None, 30)
+ assert "no valid MCP response received" in message
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("sdk_timeout", [True, False])
+ @pytest.mark.parametrize("read_timeout", [0, 1])
+ async def test_timeout_message_uses_the_deadline_that_expired(self, sdk_timeout: bool, read_timeout: int) -> None:
+ from mcp import McpError
+ from mcp.types import ErrorData
+
+ async def operation(client: rest_endpoints.MCPClient) -> dict[str, object]:
+ try:
+ raise TimeoutError("secret-timeout")
+ except TimeoutError as elapsed:
+ if not sdk_timeout:
+ raise
+ try:
+ raise McpError(ErrorData(code=408, message="secret-sdk-timeout")) from elapsed
+ except McpError as sdk_error:
+ raise TimeoutError() from sdk_error
+
+ payload: Final = NewMCPServerRequest(
+ server_name="timeout", url="https://example.com", auth_type=MCPAuth.none, timeout=read_timeout
+ )
+ result: Final = await rest_endpoints._execute_with_mcp_client(payload, operation, timeout_seconds=30)
+ assert (f"within {read_timeout}s" if sdk_timeout else "within 30s") in result["message"]
+ assert "secret" not in result["message"]
+
def test_unknown_error_falls_back_to_generic(self):
message = rest_endpoints._connection_error_message(RuntimeError("weird"), "https://example.com", 30.0)
assert "weird" not in message
- assert "proxy logs" in message.lower()
+ assert "reference" in message.lower()
+
+ def test_sdk_session_terminated_explains_endpoint_and_retry(self) -> None:
+ from mcp.shared.exceptions import McpError
+ from mcp.types import ErrorData
+
+ message: Final = rest_endpoints._connection_error_message(
+ McpError(ErrorData(code=32600, message="Session terminated")), "https://example.com/mcp", 30.0
+ )
+
+ assert "session was terminated" in message
+ assert "MCP endpoint" in message
+ assert "transport" in message
+ assert "retry" in message
+ assert "404" not in message
+
+ @pytest.mark.parametrize("code", [-32700, -32601, -32602, -32603, -32000, 32600, 408])
+ def test_rpc_errors_include_code_without_echoing_upstream_data(self, code: int) -> None:
+ from mcp.shared.exceptions import McpError
+ from mcp.types import ErrorData
+
+ message: Final = rest_endpoints._connection_error_message(
+ McpError(ErrorData(code=code, message="secret-message", data={"token": "secret-data"})),
+ "https://example.com/secret-path?token=secret-query",
+ 30.0,
+ )
+
+ assert f"JSON-RPC code {code}" in message
+ assert "secret" not in message
+ assert "timed out" not in message
+ assert "session was terminated" not in message
+
+ @pytest.mark.parametrize("status_code", [401, 403, 404, 405, 429, 503])
+ def test_wrapped_http_failures_preserve_status(self, status_code: int) -> None:
+ response: Final = httpx.Response(status_code, text="secret-body")
+ upstream: Final = httpx.HTTPStatusError(
+ "secret-exception",
+ request=httpx.Request("POST", "https://example.com/?token=secret-query"),
+ response=response,
+ )
+ wrapped: Final = BaseExceptionGroup(
+ "secret-group", [asyncio.CancelledError(), BaseExceptionGroup("nested", [upstream])]
+ )
+
+ message: Final = rest_endpoints._connection_error_message(wrapped, "https://example.com", 30.0)
+
+ assert f"HTTP {status_code}" in message
+ assert "secret" not in message
+
+ def test_explicit_cause_is_classified_before_incidental_context(self) -> None:
+ wrapped: Final = RuntimeError("secret-wrapper")
+ wrapped.__cause__ = httpx.ConnectError("secret-cause")
+ wrapped.__context__ = TimeoutError("secret-context")
+
+ message: Final = rest_endpoints._connection_error_message(wrapped, "https://example.com", 30.0)
+
+ assert "unreachable" in message
+ assert "secret" not in message
+
+ def test_timeout_url_redacts_credentials_path_query_and_fragment(self) -> None:
+ message: Final = rest_endpoints._connection_error_message(
+ TimeoutError("secret-error"),
+ "https://secret-user:secret-pass@example.com:8443/secret-path?token=secret-query#secret-fragment",
+ 30.0,
+ )
+
+ assert "https://example.com:8443" in message
+ assert "30s" in message
+ assert "secret" not in message
+
+ def test_unknown_failure_reference_matches_safe_diagnostics(self, caplog: pytest.LogCaptureFixture) -> None:
+ import re
+
+ try:
+ raise RuntimeError("secret-exception-body")
+ except RuntimeError as exc:
+ message: Final = rest_endpoints._connection_error_message(
+ exc, "https://secret-user:secret-password@example.com/secret-path?token=secret-query", 30.0
+ )
+
+ reference: Final = re.search(r"reference ([a-f0-9]{32})", message)
+ assert reference is not None
+ diagnostics: Final = tuple(
+ record for record in caplog.records if "MCP connection test failed" in record.message
+ )
+ assert len(diagnostics) == 1
+ assert reference.group(1) in diagnostics[0].message
+ assert "RuntimeError" in diagnostics[0].message
+ assert "test_unknown_failure_reference_matches_safe_diagnostics" in diagnostics[0].message
+ assert diagnostics[0].exc_info is None
+ assert "secret" not in message + diagnostics[0].message
+
+ @pytest.mark.parametrize("exc", [ValueError("secret-config"), HTTPException(500, "secret-detail")])
+ def test_unrelated_errors_are_not_misreported_as_invalid_mcp(self, exc: Exception) -> None:
+ message: Final = rest_endpoints._connection_error_message(exc, "https://example.com", 30.0)
+
+ assert "reference" in message
+ assert "invalid MCP response" not in message
+ assert "secret" not in message
+
+ def test_configuration_validation_error_uses_unknown_fallback(self) -> None:
+ from pydantic import ValidationError
+
+ with pytest.raises(ValidationError) as caught:
+ NewMCPServerRequest.model_validate({"server_name": "example", "transport": "secret-invalid-transport"})
+
+ message: Final = rest_endpoints._connection_error_message(caught.value, "https://example.com", 30.0)
+ assert "reference" in message
+ assert "invalid MCP response" not in message
+ assert "secret" not in message
+
+ @pytest.mark.asyncio
+ async def test_connection_test_preserves_cancellation(self) -> None:
+ async def cancelled_operation(client: rest_endpoints.MCPClient) -> dict[str, object]:
+ raise asyncio.CancelledError
+
+ payload: Final = NewMCPServerRequest(server_name="cancelled", url="https://example.com", auth_type=MCPAuth.none)
+ with pytest.raises(asyncio.CancelledError):
+ await rest_endpoints._execute_with_mcp_client(payload, cancelled_operation)
+
+ @pytest.mark.asyncio
+ async def test_unknown_failure_preserves_response_contract(self) -> None:
+ async def failing_operation(client: rest_endpoints.MCPClient) -> dict[str, object]:
+ raise RuntimeError("secret-operation")
+
+ payload: Final = NewMCPServerRequest(server_name="unknown", url="https://example.com", auth_type=MCPAuth.none)
+ result: Final = await rest_endpoints._execute_with_mcp_client(payload, failing_operation)
+
+ assert result["error"] is True
+ assert result["status"] == "error"
+ assert "reference" in result["message"]
+ assert "secret" not in result["message"]
+ assert "stack_trace" not in result
class TestGetServerAuthHeaderGroupDefault:
diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py
index 2284a05b2e9..5bfef2b6445 100644
--- a/tests/test_litellm/proxy/auth/test_auth_checks.py
+++ b/tests/test_litellm/proxy/auth/test_auth_checks.py
@@ -2374,6 +2374,44 @@ def _mock_prisma_for_team_lookup(find_unique):
return mock_prisma_client
+_TEAM_ALIAS_TABLE_ROW = {"id": 1, "model_aliases": '{"fast": "gpt-4o"}', "created_by": "admin", "updated_by": "admin"}
+
+
+def _prisma_team_row(include):
+ """Mimics Prisma: the `litellm_model_table` relation rides on the row only when the query `include`s it."""
+ columns = {"team_id": "team-aliases", "team_alias": "aliases", "models": ["gpt-4o"]}
+ row = (
+ {**columns, "litellm_model_table": _TEAM_ALIAS_TABLE_ROW}
+ if (include or {}).get("litellm_model_table")
+ else columns
+ )
+ return SimpleNamespace(dict=lambda: row, model_dump=lambda: row)
+
+
+@pytest.mark.asyncio
+async def test_get_team_object_loads_model_aliases_relation():
+ """LIT-5858: the auth path read teams without `include`ing `litellm_model_table`, so every JWT
+ team came back with `model_aliases=None` and alias requests 403'd."""
+ from litellm.proxy.auth.auth_checks import get_team_object
+ from litellm.proxy.auth.team_grants import team_model_aliases
+
+ async def find_unique(where, include=None):
+ return _prisma_team_row(include)
+
+ mock_cache = MagicMock()
+ mock_cache.async_get_cache = AsyncMock(return_value=None)
+ mock_cache.async_set_cache = AsyncMock()
+
+ team = await get_team_object(
+ team_id="team-aliases",
+ prisma_client=_mock_prisma_for_team_lookup(AsyncMock(side_effect=find_unique)),
+ user_api_key_cache=mock_cache,
+ check_db_only=True,
+ )
+
+ assert team_model_aliases(team) == {"fast": "gpt-4o"}
+
+
@pytest.mark.asyncio
async def test_get_team_object_distinguishes_absent_team_from_unreadable_row():
"""A deleted team and a database that would not answer both surface as a 404,
@@ -6195,6 +6233,32 @@ async def test_get_team_object_by_alias_db_fetch_returns_cached_obj():
assert result.models == ["gpt-4"]
+@pytest.mark.asyncio
+async def test_get_team_object_by_alias_loads_model_aliases_relation():
+ """LIT-5858: same regression as `test_get_team_object_loads_model_aliases_relation`, for the
+ `team_alias_jwt_field` lookup."""
+ from litellm.proxy.auth.auth_checks import get_team_object_by_alias
+ from litellm.proxy.auth.team_grants import team_model_aliases
+
+ async def find_many(where, include=None):
+ return [_prisma_team_row(include)]
+
+ mock_prisma_client = MagicMock()
+ mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many)
+
+ mock_cache = MagicMock()
+ mock_cache.async_get_cache = AsyncMock(return_value=None)
+ mock_cache.async_set_cache = AsyncMock()
+
+ team = await get_team_object_by_alias(
+ team_alias="aliases",
+ prisma_client=mock_prisma_client,
+ user_api_key_cache=mock_cache,
+ )
+
+ assert team_model_aliases(team) == {"fast": "gpt-4o"}
+
+
@pytest.mark.asyncio
async def test_get_org_object_by_alias_db_fetch_returns_validated_org():
from litellm.proxy._types import LiteLLM_OrganizationTable
diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py
index 99a0a4c0a8b..94226b5404d 100644
--- a/tests/test_litellm/proxy/auth/test_handle_jwt.py
+++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py
@@ -13,6 +13,7 @@ from litellm.proxy._types import (
DEFAULT_JWKS_STALE_TTL,
JWTLiteLLMRoleMap,
LiteLLM_JWTAuth,
+ LiteLLM_ModelTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
LiteLLM_UserTable,
@@ -1255,6 +1256,57 @@ async def test_find_team_with_model_access_model_group(monkeypatch):
assert team_obj.team_id == "team-1"
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "model_aliases",
+ ['{"fast": "gpt-4o"}', {"fast": "gpt-4o"}],
+ ids=["json-string", "dict"],
+)
+async def test_find_team_with_model_access_resolves_team_model_alias(monkeypatch, model_aliases):
+ """LIT-5858: a JWT team that grants `gpt-4o` under the alias `fast` must resolve a request
+ for `fast`. The JWT path used to pass `team_model_aliases=None`, so every alias request 403'd."""
+ import sys
+ import types
+
+ from litellm.caching import DualCache
+ from litellm.proxy.utils import ProxyLogging
+ from litellm.router import Router
+
+ router = Router(model_list=[{"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}])
+ proxy_server_module = types.ModuleType("proxy_server")
+ proxy_server_module.llm_router = router
+ monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module)
+
+ team = LiteLLM_TeamTable(
+ team_id="team-aliases",
+ models=["gpt-4o"],
+ litellm_model_table=LiteLLM_ModelTable(model_aliases=model_aliases, created_by="admin", updated_by="admin"),
+ )
+
+ async def mock_get_team_object(*args, **kwargs):
+ return team
+
+ monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object)
+
+ jwt_handler = JWTHandler()
+ jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
+ user_api_key_cache = DualCache()
+
+ team_id, team_obj = await JWTAuthManager.find_team_with_model_access(
+ team_ids={"team-aliases"},
+ requested_model="fast",
+ route="/chat/completions",
+ jwt_handler=jwt_handler,
+ prisma_client=None,
+ user_api_key_cache=user_api_key_cache,
+ parent_otel_span=None,
+ proxy_logging_obj=ProxyLogging(user_api_key_cache=user_api_key_cache),
+ )
+
+ assert team_id == "team-aliases"
+ assert team_obj is team
+
+
@pytest.mark.asyncio
async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatch):
"""Regression for #31189: a single-team JWT that grants the requested model
diff --git a/tests/test_litellm/proxy/auth/test_team_grants.py b/tests/test_litellm/proxy/auth/test_team_grants.py
new file mode 100644
index 00000000000..447fc1c93a1
--- /dev/null
+++ b/tests/test_litellm/proxy/auth/test_team_grants.py
@@ -0,0 +1,129 @@
+import pytest
+
+from litellm.proxy._types import (
+ LiteLLM_BudgetTable,
+ LiteLLM_ObjectPermissionTable,
+ LiteLLM_TeamMembership,
+ LiteLLM_TeamTable,
+ LiteLLM_VerificationTokenView,
+ Member,
+ UserAPIKeyAuth,
+)
+from litellm.models.team import LiteLLM_ModelTable
+from litellm.proxy.auth.team_grants import team_grants, team_model_aliases
+
+TEAM_ID = "team-grants"
+USER_ID = "user-in-team"
+ALIASES = {"fast": "gpt-4o-mini", "smart": "gpt-4o"}
+
+
+def _alias_table(model_aliases) -> LiteLLM_ModelTable:
+ return LiteLLM_ModelTable(model_aliases=model_aliases, created_by="admin", updated_by="admin")
+
+
+def _full_team(model_aliases=ALIASES) -> LiteLLM_TeamTable:
+ return LiteLLM_TeamTable(
+ team_id=TEAM_ID,
+ team_alias="grants-team",
+ tpm_limit=1000,
+ rpm_limit=10,
+ max_budget=50.0,
+ soft_budget=25.0,
+ spend=12.5,
+ models=["gpt-4o", "gpt-4o-mini"],
+ blocked=True,
+ metadata={"tier": "gold"},
+ litellm_model_table=_alias_table(model_aliases),
+ object_permission_id="op-1",
+ object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-1", mcp_servers=["mcp-a"]),
+ members_with_roles=[
+ Member(user_id="someone-else", role="user"),
+ Member(user_id=USER_ID, role="admin"),
+ ],
+ )
+
+
+def _membership() -> LiteLLM_TeamMembership:
+ return LiteLLM_TeamMembership(
+ user_id=USER_ID,
+ team_id=TEAM_ID,
+ spend=3.25,
+ litellm_budget_table=LiteLLM_BudgetTable(tpm_limit=500, rpm_limit=5),
+ )
+
+
+def test_team_grants_cover_every_team_field_the_key_path_gets():
+ """Class guard for LIT-5858 and its siblings: every ``team_*`` column the combined-view SQL hands the
+ virtual-key path must come out of the projection too, with the team's actual value, so adding a column
+ to ``LiteLLM_VerificationTokenView`` without teaching ``team_grants`` fails here instead of in prod."""
+ team = _full_team()
+ grants = team_grants(team_object=team, team_membership=_membership(), user_id=USER_ID)
+ token = UserAPIKeyAuth(team_id=TEAM_ID, **grants)
+
+ view_team_fields = {name for name in LiteLLM_VerificationTokenView.model_fields if name.startswith("team_")}
+ assert view_team_fields - {"team_id"} <= set(grants)
+ assert all(grants[name] is not None for name in view_team_fields - {"team_id"})
+
+ assert token.team_alias == "grants-team"
+ assert token.team_tpm_limit == 1000
+ assert token.team_rpm_limit == 10
+ assert token.team_max_budget == 50.0
+ assert token.team_soft_budget == 25.0
+ assert token.team_spend == 12.5
+ assert token.team_models == ["gpt-4o", "gpt-4o-mini"]
+ assert token.team_blocked is True
+ assert token.team_metadata == {"tier": "gold"}
+ assert token.team_model_aliases == ALIASES
+ assert token.team_object_permission_id == "op-1"
+ assert token.team_object_permission is not None
+ assert token.team_object_permission.mcp_servers == ["mcp-a"]
+ assert token.team_member == Member(user_id=USER_ID, role="admin")
+ assert token.team_member_spend == 3.25
+ assert token.team_member_tpm_limit == 500
+ assert token.team_member_rpm_limit == 5
+
+
+def test_team_grants_without_team_leave_token_defaults():
+ token = UserAPIKeyAuth(**team_grants(team_object=None, team_membership=None, user_id=USER_ID))
+ assert token == UserAPIKeyAuth()
+
+
+@pytest.mark.parametrize(
+ "stored_aliases",
+ [ALIASES, '{"fast": "gpt-4o-mini", "smart": "gpt-4o"}'],
+ ids=["json-object", "json-string-as-written-by-team-new"],
+)
+def test_team_model_aliases_decode_both_storage_shapes(stored_aliases):
+ team = _full_team(model_aliases=stored_aliases)
+ assert team_model_aliases(team) == ALIASES
+ assert team_grants(team_object=team, team_membership=None, user_id=None)["team_model_aliases"] == ALIASES
+
+
+@pytest.mark.parametrize("stored_aliases", [None, "not json", '["a", "b"]', {"fast": 3}], ids=str)
+def test_team_model_aliases_treat_unusable_column_as_no_aliases(stored_aliases):
+ team = _full_team(model_aliases=stored_aliases)
+ assert team_model_aliases(team) is None
+ assert team_grants(team_object=team, team_membership=None, user_id=None)["team_model_aliases"] is None
+
+
+def test_team_model_aliases_none_without_relation_loaded():
+ team = _full_team()
+ team.litellm_model_table = None
+ assert team_model_aliases(team) is None
+ assert team_model_aliases(None) is None
+
+
+def test_team_member_is_the_callers_row_only():
+ team = _full_team()
+ assert team_grants(team_object=team, team_membership=None, user_id="someone-else")["team_member"] == Member(
+ user_id="someone-else", role="user"
+ )
+ assert team_grants(team_object=team, team_membership=None, user_id="stranger")["team_member"] is None
+ assert team_grants(team_object=team, team_membership=None, user_id=None)["team_member"] is None
+
+
+def test_membership_limits_absent_without_membership_row():
+ grants = team_grants(team_object=_full_team(), team_membership=None, user_id=USER_ID)
+ assert grants["team_member_spend"] is None
+ assert grants["team_member_tpm_limit"] is None
+ assert grants["team_member_rpm_limit"] is None
diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
index d4dae5ee150..6df7f2118e0 100644
--- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
+++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
@@ -7174,3 +7174,119 @@ async def test_sideband_auth_uses_encrypted_model_for_budget_checks(monkeypatch,
}, AsyncMock(), AsyncMock())
assert await auth_module.user_api_key_auth_websocket(websocket) == "authenticated-with-model"
assert seen == [({"model": "budgeted-voice"}, "Bearer owner")]
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("is_proxy_admin", [False, True], ids=["standard-return", "proxy-admin-return"])
+async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_admin):
+ """LIT-5858: the team-based JWT path hand-built ``UserAPIKeyAuth`` from a short list of team fields, so the
+ team's model aliases (and on the admin return, its object permission) never reached the token and alias
+ requests 403'd. Both returns now go through ``team_grants``; pin the fields that used to be dropped."""
+ import litellm.proxy.proxy_server as _proxy_server_mod
+ from fastapi import Request
+ from starlette.datastructures import URL
+
+ from litellm.models.team import LiteLLM_ModelTable
+ from litellm.proxy._types import (
+ LiteLLM_ObjectPermissionTable,
+ LiteLLM_TeamMembership,
+ LiteLLM_TeamTable,
+ Member,
+ )
+
+ class _AcceptEveryJwt(JWTHandler):
+ def is_jwt(self, token: str) -> bool:
+ return True
+
+ jwt_handler = _AcceptEveryJwt()
+ jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
+
+ team = LiteLLM_TeamTable(
+ team_id="team-jwt-aliases",
+ team_alias="jwt-aliases",
+ models=["gpt-4o"],
+ max_budget=40.0,
+ spend=4.0,
+ blocked=False,
+ metadata={"tier": "gold"},
+ litellm_model_table=LiteLLM_ModelTable(
+ model_aliases='{"fast": "gpt-4o"}', created_by="admin", updated_by="admin"
+ ),
+ object_permission_id="op-jwt",
+ object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-jwt", mcp_servers=["mcp-a"]),
+ members_with_roles=[Member(user_id="jwt-user", role="admin")],
+ )
+ membership = LiteLLM_TeamMembership(user_id="jwt-user", team_id="team-jwt-aliases", spend=1.5)
+ builder_result = {
+ "is_proxy_admin": is_proxy_admin,
+ "team_object": team,
+ "user_object": None,
+ "end_user_object": None,
+ "org_object": None,
+ "token": "jwt",
+ "team_id": "team-jwt-aliases",
+ "user_id": "jwt-user",
+ "user_email": "jwt-user@example.com",
+ "end_user_id": None,
+ "org_id": None,
+ "team_membership": membership,
+ "jwt_claims": {"sub": "jwt-user"},
+ }
+
+ mock_proxy_logging_obj = MagicMock()
+ mock_proxy_logging_obj.internal_usage_cache = MagicMock()
+ mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
+ mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
+ attrs = {
+ "prisma_client": MagicMock(),
+ "user_api_key_cache": DualCache(),
+ "proxy_logging_obj": mock_proxy_logging_obj,
+ "master_key": "sk-master-key",
+ "general_settings": {"enable_jwt_auth": True},
+ "llm_model_list": [],
+ "llm_router": None,
+ "open_telemetry_logger": None,
+ "model_max_budget_limiter": MagicMock(),
+ "user_custom_auth": None,
+ "jwt_handler": jwt_handler,
+ "premium_user": True,
+ "litellm_proxy_admin_name": "admin",
+ }
+ originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
+ try:
+ for k, v in attrs.items():
+ setattr(_proxy_server_mod, k, v)
+ request = Request(scope={"type": "http", "headers": [], "method": "POST"})
+ request._url = URL(url="/chat/completions")
+ with patch( # test-quality-ok: auth_builder is the claim-resolution seam; the regression is how its result is projected onto the token
+ "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
+ new_callable=AsyncMock,
+ return_value=builder_result,
+ ):
+ token = await _user_api_key_auth_builder(
+ request=request,
+ api_key="Bearer header.payload.signature",
+ azure_api_key_header="",
+ anthropic_api_key_header=None,
+ google_ai_studio_api_key_header=None,
+ azure_apim_header=None,
+ request_data={},
+ )
+ finally:
+ for k, v in originals.items():
+ setattr(_proxy_server_mod, k, v)
+
+ assert token.team_id == "team-jwt-aliases"
+ assert token.user_role == (LitellmUserRoles.PROXY_ADMIN if is_proxy_admin else LitellmUserRoles.INTERNAL_USER)
+ assert token.team_model_aliases == {"fast": "gpt-4o"}
+ assert token.team_object_permission is not None
+ assert token.team_object_permission.mcp_servers == ["mcp-a"]
+ assert token.team_object_permission_id == "op-jwt"
+ assert token.team_alias == "jwt-aliases"
+ assert token.team_models == ["gpt-4o"]
+ assert token.team_max_budget == 40.0
+ assert token.team_spend == 4.0
+ assert token.team_metadata == {"tier": "gold"}
+ assert token.team_member == Member(user_id="jwt-user", role="admin")
+ assert token.team_member_spend == 1.5
+ assert token.jwt_claims == {"sub": "jwt-user"}
diff --git a/tests/test_litellm/proxy/common_utils/test_callback_utils.py b/tests/test_litellm/proxy/common_utils/test_callback_utils.py
index 66f77db6da9..ecb2375d495 100644
--- a/tests/test_litellm/proxy/common_utils/test_callback_utils.py
+++ b/tests/test_litellm/proxy/common_utils/test_callback_utils.py
@@ -1,30 +1,33 @@
import copy
+import json
import sys
from types import ModuleType, SimpleNamespace
+from typing import Final
+from unittest.mock import patch
import pytest
-
+import litellm
+from litellm.caching.caching import DualCache
+from litellm.constants import MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH
+from litellm.integrations.custom_logger import CustomLogger
+from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.callback_utils import (
+ _serialize_scan_metadata_header,
add_guardrail_scan_id,
add_policy_to_applied_policies_header,
decrypt_callback_vars,
encrypt_callback_vars,
get_logging_caching_headers,
- initialize_callbacks_on_proxy,
get_remaining_tokens_and_requests_from_request_data,
+ initialize_callbacks_on_proxy,
normalize_callback_names,
+ process_callback,
sanitize_openai_provider_metadata,
strip_callback_config,
)
-import litellm
-from litellm.caching.caching import DualCache
-from litellm.integrations.custom_logger import CustomLogger
-from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
-
-from unittest.mock import patch
-from litellm.proxy.common_utils.callback_utils import process_callback
+from litellm.types.guardrails import GuardrailEventHooks
def test_get_remaining_tokens_and_requests_from_request_data():
@@ -189,20 +192,109 @@ def test_get_logging_caching_headers_merges_metadata_and_litellm_metadata():
assert headers["x-litellm-policy-sources"] == "global-baseline=team_default"
+def _record(
+ request_data: dict[str, object],
+ scan_id: str | None,
+ guardrail_name: str = "airs",
+ provider: str = "panw_prisma_airs",
+ stage: GuardrailEventHooks = GuardrailEventHooks.pre_call,
+) -> None:
+ add_guardrail_scan_id(
+ request_data=request_data, scan_id=scan_id, guardrail_name=guardrail_name, provider=provider, stage=stage
+ )
+
+
def test_add_guardrail_scan_id_dedupes_and_becomes_response_header():
request_data = {"litellm_metadata": {}}
- add_guardrail_scan_id(request_data=request_data, scan_id="scan-1")
- add_guardrail_scan_id(request_data=request_data, scan_id="scan-1")
- add_guardrail_scan_id(request_data=request_data, scan_id="scan-2")
- add_guardrail_scan_id(request_data=request_data, scan_id=None)
+ _record(request_data, "scan-1")
+ _record(request_data, "scan-1")
+ _record(request_data, "scan-2")
+ _record(request_data, None)
assert request_data["litellm_metadata"]["guardrail_scan_ids"] == ("scan-1", "scan-2")
assert get_logging_caching_headers(request_data)["x-litellm-guardrail-scan-id"] == "scan-1,scan-2"
-def test_get_logging_caching_headers_omits_scan_id_header_without_scans():
- assert "x-litellm-guardrail-scan-id" not in get_logging_caching_headers({"litellm_metadata": {}})
+def test_scan_metadata_header_maps_each_id_to_its_guardrail_stage_and_provider():
+ request_data: Final[dict[str, object]] = {"litellm_metadata": {}}
+
+ _record(
+ request_data, "scan-1", guardrail_name="airs", provider="panw_prisma_airs", stage=GuardrailEventHooks.pre_call
+ )
+ _record(
+ request_data, "mod-1", guardrail_name="mod", provider="openai_moderation", stage=GuardrailEventHooks.pre_call
+ )
+ _record(
+ request_data, "scan-2", guardrail_name="airs", provider="panw_prisma_airs", stage=GuardrailEventHooks.post_call
+ )
+ _record(
+ request_data, "scan-2", guardrail_name="airs", provider="panw_prisma_airs", stage=GuardrailEventHooks.post_call
+ )
+ _record(request_data, None, guardrail_name="mod", provider="openai_moderation", stage=GuardrailEventHooks.post_call)
+
+ headers: Final = get_logging_caching_headers(request_data)
+ assert headers is not None
+ assert headers["x-litellm-guardrail-scan-id"] == "scan-1,mod-1,scan-2"
+ assert json.loads(headers["x-litellm-guardrail-scan-metadata"]) == [
+ {"guardrail": "airs", "stage": "pre_call", "provider": "panw_prisma_airs", "scan_id": "scan-1"},
+ {"guardrail": "mod", "stage": "pre_call", "provider": "openai_moderation", "scan_id": "mod-1"},
+ {"guardrail": "airs", "stage": "post_call", "provider": "panw_prisma_airs", "scan_id": "scan-2"},
+ ]
+
+
+def test_scan_metadata_keeps_same_id_reused_across_stages():
+ request_data: Final[dict[str, object]] = {"metadata": {}}
+
+ _record(request_data, "scan-1", stage=GuardrailEventHooks.pre_call)
+ _record(request_data, "scan-1", stage=GuardrailEventHooks.post_call)
+
+ headers: Final = get_logging_caching_headers(request_data)
+ assert headers is not None
+ assert headers["x-litellm-guardrail-scan-id"] == "scan-1"
+ assert [entry["stage"] for entry in json.loads(headers["x-litellm-guardrail-scan-metadata"])] == [
+ "pre_call",
+ "post_call",
+ ]
+
+
+def test_scan_metadata_header_drops_trailing_entries_to_stay_within_length_limit():
+ request_data: Final[dict[str, object]] = {"litellm_metadata": {}}
+ scan_ids: Final = tuple(f"0f9c4b7e-3d2a-4c1b-9e8f-{index:012d}" for index in range(40))
+ for scan_id in scan_ids:
+ _record(request_data, scan_id, stage=GuardrailEventHooks.post_call)
+
+ headers: Final = get_logging_caching_headers(request_data)
+ assert headers is not None
+ assert headers["x-litellm-guardrail-scan-id"] == ",".join(scan_ids)
+ header: Final = headers["x-litellm-guardrail-scan-metadata"]
+ assert len(header) <= MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH
+ kept: Final = json.loads(header)
+ assert 1 < len(kept) < len(scan_ids)
+ assert [entry["scan_id"] for entry in kept] == list(scan_ids[: len(kept)])
+
+
+def test_serialize_scan_metadata_header_keeps_exactly_the_entries_that_fit():
+ entries: Final = ({"scan_id": "a"}, {"scan_id": "b"}, {"scan_id": "c"})
+ two_entries: Final = '[{"scan_id":"a"},{"scan_id":"b"}]'
+
+ assert _serialize_scan_metadata_header(entries, max_length=len(two_entries)) == two_entries
+ assert _serialize_scan_metadata_header(entries, max_length=len(two_entries) - 1) == '[{"scan_id":"a"}]'
+ assert _serialize_scan_metadata_header(entries, max_length=len(two_entries) + 1) == two_entries
+ assert _serialize_scan_metadata_header(entries, max_length=1000) == json.dumps(entries, separators=(",", ":"))
+ assert _serialize_scan_metadata_header(entries, max_length=5) is None
+ assert _serialize_scan_metadata_header((), max_length=1000) is None
+
+
+def test_scan_metadata_is_an_internal_metadata_key():
+ assert sanitize_openai_provider_metadata({"guardrail_scan_metadata": "x", "keep": "y"}) == {"keep": "y"}
+
+
+def test_get_logging_caching_headers_omits_scan_headers_without_scans():
+ headers: Final = get_logging_caching_headers({"litellm_metadata": {}})
+ assert headers is not None
+ assert "x-litellm-guardrail-scan-id" not in headers
+ assert "x-litellm-guardrail-scan-metadata" not in headers
def test_initialize_callbacks_on_proxy_instantiates_compression_interception(
diff --git a/tests/test_litellm/proxy/db/test_db_url_settings.py b/tests/test_litellm/proxy/db/test_db_url_settings.py
index ba342342366..875dca4bee3 100644
--- a/tests/test_litellm/proxy/db/test_db_url_settings.py
+++ b/tests/test_litellm/proxy/db/test_db_url_settings.py
@@ -11,15 +11,30 @@ clobber a pre-existing ``DATABASE_URL_READ_REPLICA``. A pre-existing
``DATABASE_URL`` (password auth) is likewise left untouched.
"""
+import datetime
+import hashlib
import os
+import socket
+import ssl
+import tempfile
+import threading
import urllib.parse
+from collections.abc import Iterator
+from dataclasses import dataclass
+from pathlib import Path
+from typing import Final
from unittest.mock import patch
import pytest
+from cryptography import x509
+from cryptography.hazmat.primitives import hashes, serialization
+from cryptography.hazmat.primitives.asymmetric import ec
from pydantic import ValidationError
from litellm.proxy.db.db_url_settings import (
+ PG_SSL_REQUEST,
DatabaseURLSettings,
+ translate_libpq_ssl_params,
unsupported_db_scheme,
unsupported_db_scheme_message,
)
@@ -381,9 +396,7 @@ def test_writer_password_is_percent_encoded(monkeypatch):
def test_writer_url_not_clobbered_when_already_set(monkeypatch):
"""An operator-pinned DATABASE_URL (e.g. helm's $(VAR) assembly) always
wins over the discrete fields."""
- monkeypatch.setenv(
- "DATABASE_URL", "postgresql://pinned:url@db.example.com:5432/litellm_db"
- )
+ monkeypatch.setenv("DATABASE_URL", "postgresql://pinned:url@db.example.com:5432/litellm_db")
monkeypatch.setenv("DATABASE_HOST", "writer.example.com")
monkeypatch.setenv("DATABASE_USER", "litellm")
monkeypatch.setenv("DATABASE_NAME", "litellm_db")
@@ -515,9 +528,7 @@ def test_apply_to_env_rejects_pinned_sqlite_direct_url(monkeypatch):
def test_apply_to_env_rejects_pinned_non_postgres_reader(monkeypatch):
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
- monkeypatch.setenv(
- "DATABASE_URL_READ_REPLICA", "mysql://u:p@reader.example.com:3306/db"
- )
+ monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "mysql://u:p@reader.example.com:3306/db")
with pytest.raises(RuntimeError, match=r"DATABASE_URL_READ_REPLICA.*mysql"):
_apply()
@@ -542,15 +553,11 @@ def test_reader_inherits_writer_connection_params(monkeypatch):
"DATABASE_URL",
"postgresql://u:p@writer.example.com:5432/db?connection_limit=3&pool_timeout=20&pgbouncer=true",
)
- monkeypatch.setenv(
- "DATABASE_URL_READ_REPLICA", "postgresql://u:p@reader.example.com:5432/db"
- )
+ monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "postgresql://u:p@reader.example.com:5432/db")
_apply()
- query = urllib.parse.parse_qs(
- urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query
- )
+ query = urllib.parse.parse_qs(urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query)
assert query["connection_limit"] == ["3"]
assert query["pool_timeout"] == ["20"]
assert query["pgbouncer"] == ["true"]
@@ -568,9 +575,7 @@ def test_reader_keeps_its_own_pinned_connection_params(monkeypatch):
_apply()
- query = urllib.parse.parse_qs(
- urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query
- )
+ query = urllib.parse.parse_qs(urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query)
assert query["connection_limit"] == ["50"]
assert query["pool_timeout"] == ["20"]
@@ -776,6 +781,170 @@ def test_libpq_verify_full_and_sslrootcert_become_prisma_strict_sslcert(monkeypa
}
+def _issue_cert(
+ subject: str, issuer: x509.Certificate | None, issuer_key: ec.EllipticCurvePrivateKey | None, ca: bool
+) -> tuple[x509.Certificate, ec.EllipticCurvePrivateKey]:
+ key: Final = ec.generate_private_key(ec.SECP256R1())
+ name: Final = x509.Name((x509.NameAttribute(x509.NameOID.COMMON_NAME, subject),))
+ now: Final = datetime.datetime.now(datetime.timezone.utc)
+ builder: Final = (
+ x509.CertificateBuilder()
+ .subject_name(name)
+ .issuer_name(issuer.subject if issuer else name)
+ .public_key(key.public_key())
+ .serial_number(x509.random_serial_number())
+ .not_valid_before(now - datetime.timedelta(minutes=5))
+ .not_valid_after(now + datetime.timedelta(days=1))
+ .add_extension(x509.BasicConstraints(ca=ca, path_length=None), critical=True)
+ .add_extension(x509.SubjectAlternativeName((x509.DNSName("localhost"),)), critical=False)
+ )
+ return builder.sign(issuer_key or key, hashes.SHA256()), key
+
+
+def _pem(cert: x509.Certificate) -> bytes:
+ return cert.public_bytes(serialization.Encoding.PEM)
+
+
+class _TlsPostgresStub:
+ """Answers one libpq ``SSLRequest`` with ``S`` and serves ``leaf + intermediate``."""
+
+ def __init__(self, chain_pem: Path, key_pem: Path) -> None:
+ self.context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
+ self.context.load_cert_chain(str(chain_pem), str(key_pem))
+ self.listener: Final = socket.create_server(("127.0.0.1", 0))
+ self.port: Final[int] = self.listener.getsockname()[1]
+ self.thread: Final = threading.Thread(target=self._serve, daemon=True)
+ self.thread.start()
+
+ def _serve(self) -> None:
+ with self.listener:
+ while True:
+ try:
+ conn: socket.socket = self.listener.accept()[0]
+ except OSError:
+ return
+ with conn:
+ try:
+ if conn.recv(8) == PG_SSL_REQUEST:
+ conn.sendall(b"S")
+ with self.context.wrap_socket(conn, server_side=True) as tls:
+ tls.recv(1)
+ except OSError:
+ continue
+
+
+@dataclass(frozen=True, slots=True)
+class _RdsLikePki:
+ bundle: Path
+ wrong_bundle: Path
+ root: Path
+ port: int
+
+
+@pytest.fixture
+def rds_like_pki(tmp_path: Path) -> Iterator[_RdsLikePki]:
+ """An RDS-shaped trust setup: the server sends leaf + intermediate, the
+ bundle holds only self-signed roots, and the right root is not first."""
+ root, root_key = _issue_cert("Real Root CA", None, None, ca=True)
+ decoys: Final = tuple(_issue_cert(f"Decoy Root CA {i}", None, None, ca=True)[0] for i in range(3))
+ intermediate, intermediate_key = _issue_cert("Intermediate CA", root, root_key, ca=True)
+ leaf, leaf_key = _issue_cert("localhost", intermediate, intermediate_key, ca=False)
+ chain_pem: Final = tmp_path / "server-chain.pem"
+ chain_pem.write_bytes(_pem(leaf) + _pem(intermediate))
+ key_pem: Final = tmp_path / "server.key"
+ key_pem.write_bytes(
+ leaf_key.private_bytes(
+ serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()
+ )
+ )
+ bundle: Final = tmp_path / "global-bundle.pem"
+ bundle.write_bytes(b"".join(_pem(decoy) for decoy in decoys) + _pem(root))
+ wrong_bundle: Final = tmp_path / "wrong-bundle.pem"
+ wrong_bundle.write_bytes(b"".join(_pem(decoy) for decoy in decoys))
+ root_pem: Final = tmp_path / "root.pem"
+ root_pem.write_bytes(_pem(root))
+ stub: Final = _TlsPostgresStub(chain_pem, key_pem)
+ yield _RdsLikePki(bundle=bundle, wrong_bundle=wrong_bundle, root=root_pem, port=stub.port)
+ stub.listener.close()
+
+
+def _params(url: str) -> tuple[tuple[str, str], ...]:
+ return tuple(urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query, keep_blank_values=True))
+
+
+def test_multi_root_bundle_is_pinned_to_the_root_the_server_chains_to(
+ monkeypatch: pytest.MonkeyPatch, rds_like_pki: _RdsLikePki
+):
+ """Prisma's ``sslcert`` loads only the first certificate of the file, so
+ handing it the whole RDS bundle trusts one region's root and fails with
+ "unable to get local issuer certificate" everywhere else. The URL Prisma
+ receives must point at a single-certificate file holding the server's root."""
+ monkeypatch.setenv(
+ "DATABASE_URL",
+ f"postgresql://u:p@localhost:{rds_like_pki.port}/litellm_db?sslmode=verify-full&sslrootcert={rds_like_pki.bundle}",
+ )
+
+ _apply()
+
+ (sslmode, sslcert, sslaccept, _) = _params(os.environ["DATABASE_URL"])
+ assert (sslmode, sslaccept) == (("sslmode", "require"), ("sslaccept", "strict"))
+ assert sslcert[0] == "sslcert" and sslcert[1] != str(rds_like_pki.bundle)
+ assert Path(sslcert[1]).read_bytes() == rds_like_pki.root.read_bytes()
+
+
+def test_pinned_root_replaces_a_planted_symlink_instead_of_writing_through_it(
+ monkeypatch: pytest.MonkeyPatch, rds_like_pki: _RdsLikePki, tmp_path: Path
+):
+ """The pinned file has a predictable name in a shared temp dir, so a symlink
+ planted there must not redirect the write onto its target."""
+ monkeypatch.setattr(tempfile, "tempdir", str(tmp_path))
+ root_der: Final = x509.load_pem_x509_certificate(rds_like_pki.root.read_bytes()).public_bytes(
+ serialization.Encoding.DER
+ )
+ pinned: Final = tmp_path / f"litellm-sslcert-{hashlib.sha256(root_der).hexdigest()[:16]}.pem"
+ victim: Final = tmp_path / "victim.txt"
+ victim.write_text("untouched")
+ pinned.symlink_to(victim)
+ monkeypatch.setenv(
+ "DATABASE_URL",
+ f"postgresql://u:p@localhost:{rds_like_pki.port}/litellm_db?sslmode=verify-full&sslrootcert={rds_like_pki.bundle}",
+ )
+
+ _apply()
+
+ assert ("sslcert", str(pinned)) in _params(os.environ["DATABASE_URL"])
+ assert victim.read_text() == "untouched"
+ assert not pinned.is_symlink() and pinned.read_bytes() == rds_like_pki.root.read_bytes()
+
+
+def test_bundle_without_the_servers_root_is_passed_through_unchanged(
+ monkeypatch: pytest.MonkeyPatch, rds_like_pki: _RdsLikePki
+):
+ """Nothing in the bundle verifies the server, so no root is pinned and
+ Prisma keeps rejecting the connection instead of trusting a root the
+ operator never shipped."""
+ monkeypatch.setenv(
+ "DATABASE_URL",
+ f"postgresql://u:p@localhost:{rds_like_pki.port}/litellm_db"
+ f"?sslmode=verify-full&sslrootcert={rds_like_pki.wrong_bundle}",
+ )
+
+ _apply()
+
+ assert ("sslcert", str(rds_like_pki.wrong_bundle)) in _params(os.environ["DATABASE_URL"])
+
+
+def test_root_cert_resolver_receives_the_urls_host_and_default_port():
+ def resolver(cert_path: str, host: str, port: int) -> str:
+ return f"/pinned/{host}/{port}{cert_path}"
+
+ url: Final = translate_libpq_ssl_params(
+ "postgresql://u:p@db.example.com/litellm_db?sslmode=verify-full&sslrootcert=/certs/bundle.pem", resolver
+ )
+
+ assert ("sslcert", "/pinned/db.example.com/5432/certs/bundle.pem") in _params(url)
+
+
def test_libpq_verify_ca_becomes_prisma_strict(monkeypatch):
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@db.example.com:5432/litellm_db?sslmode=verify-ca")
diff --git a/tests/test_litellm/proxy/db/test_gateway_request_tracking.py b/tests/test_litellm/proxy/db/test_gateway_request_tracking.py
index 93a11a914cb..045261e2d53 100644
--- a/tests/test_litellm/proxy/db/test_gateway_request_tracking.py
+++ b/tests/test_litellm/proxy/db/test_gateway_request_tracking.py
@@ -8,8 +8,11 @@ from datetime import datetime, timezone
import pytest
+from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_GATEWAY_REQUESTS_BUFFER_KEY
from litellm.proxy.db.gateway_request_tracking import (
+ GATEWAY_REQUESTS_JOB_NAME,
GatewayRequestAccumulator,
+ GatewayRequestRedisBuffer,
commit_gateway_requests_to_db,
flush_gateway_requests,
)
@@ -83,40 +86,47 @@ def test_drain_snapshot_is_not_mutated_by_later_records():
# ── commit ────────────────────────────────────────────────────────────────────
-class FakeTable:
- def __init__(self) -> None:
- self.upserts: list[dict] = []
-
- def upsert(self, *, where: dict, data: dict) -> None:
- self.upserts.append({"where": where, "data": data})
-
-
-class FakeBatcher:
- def __init__(self, table: FakeTable) -> None:
- self.litellm_dailygatewayrequests = table
-
- async def __aenter__(self) -> "FakeBatcher":
- return self
-
- async def __aexit__(self, *args: object) -> bool:
- return False
-
-
class FakeDB:
- def __init__(self, table: FakeTable) -> None:
- self._table = table
+ def __init__(self) -> None:
+ self.statements: list[tuple[str, tuple[object, ...]]] = []
- def batch_(self) -> FakeBatcher:
- return FakeBatcher(self._table)
+ async def execute_raw(self, query: str, *args: object) -> int:
+ self.statements.append((query, args))
+ return len(args) // 5
class FakePrismaClient:
def __init__(self) -> None:
- self.table = FakeTable()
- self.db = FakeDB(self.table)
+ self.db = FakeDB()
-def test_commit_upserts_one_incrementing_row_per_key():
+def _rows_written(client: FakePrismaClient) -> list[tuple[object, ...]]:
+ """Every (date, category, route, successful, failed) tuple the database received, in statement order."""
+ return [params[i : i + 5] for _, params in client.db.statements for i in range(0, len(params), 5)]
+
+
+def test_commit_increments_with_a_single_statement_for_the_whole_snapshot():
+ """One statement per flush is the whole point: the previous per-key upsert cost
+ the primary (workers x routes) statements per interval."""
+ client = FakePrismaClient()
+ snapshot = {
+ GatewayRequestKey(date="2026-08-01", category="llm", route=route): (
+ GatewayRequestCounts(successful_requests=7, failed_requests=2)
+ )
+ for route in ("/chat/completions", "/embeddings", "/responses", "/v1/messages", "/mcp")
+ }
+
+ asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot=snapshot))
+
+ assert len(client.db.statements) == 1
+ sql, params = client.db.statements[0]
+ assert sql.count("ON CONFLICT") == 1
+ assert sql.count("(NOW() AT TIME ZONE 'UTC'))") == 5
+ assert len(params) == 25
+
+
+def test_commit_sql_adds_to_the_existing_row_instead_of_replacing_it():
+ """A worker only knows its own share; the SQL must add EXCLUDED onto the stored count."""
client = FakePrismaClient()
snapshot = {
GatewayRequestKey(date="2026-08-01", category="llm", route="/chat/completions"): (
@@ -126,20 +136,36 @@ def test_commit_upserts_one_incrementing_row_per_key():
asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot=snapshot))
- assert len(client.table.upserts) == 1
- written = client.table.upserts[0]
- assert written["where"] == {
- "date_category_route": {
- "date": "2026-08-01",
- "category": "llm",
- "route": "/chat/completions",
- }
+ sql, params = client.db.statements[0]
+ assert 'INSERT INTO "LiteLLM_DailyGatewayRequests"' in sql
+ assert 'ON CONFLICT ("date", "category", "route") DO UPDATE SET' in sql
+ assert (
+ '"successful_requests" = "LiteLLM_DailyGatewayRequests"."successful_requests" + EXCLUDED."successful_requests"'
+ in sql
+ )
+ assert '"failed_requests" = "LiteLLM_DailyGatewayRequests"."failed_requests" + EXCLUDED."failed_requests"' in sql
+ assert params == ("2026-08-01", "llm", "/chat/completions", 7, 2)
+
+
+def test_commit_placeholders_line_up_with_params():
+ """$n positions are generated per row; a drift here silently swaps a route for a count."""
+ client = FakePrismaClient()
+ snapshot = {
+ GatewayRequestKey(date="2026-08-01", category="llm", route="/chat/completions"): (
+ GatewayRequestCounts(successful_requests=1, failed_requests=0)
+ ),
+ GatewayRequestKey(date="2026-08-01", category="mcp", route="/mcp"): (
+ GatewayRequestCounts(successful_requests=0, failed_requests=3)
+ ),
}
- assert written["data"]["update"] == {
- "successful_requests": {"increment": 7},
- "failed_requests": {"increment": 2},
- }
- assert written["data"]["create"]["successful_requests"] == 7
+
+ asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot=snapshot))
+
+ sql, params = client.db.statements[0]
+ assert "($1::text, $2::text, $3::text, $4::bigint, $5::bigint," in sql
+ assert "($6::text, $7::text, $8::text, $9::bigint, $10::bigint," in sql
+ assert "$11" not in sql
+ assert params == ("2026-08-01", "llm", "/chat/completions", 1, 0, "2026-08-01", "mcp", "/mcp", 0, 3)
def test_commit_is_deterministically_ordered():
@@ -154,17 +180,14 @@ def test_commit_is_deterministically_ordered():
asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot=snapshot))
- written_order = [
- (row["where"]["date_category_route"]["date"], row["where"]["date_category_route"]["category"])
- for row in client.table.upserts
- ]
+ written_order = [(row[0], row[1]) for row in _rows_written(client)]
assert written_order == [("2026-08-01", "llm"), ("2026-08-01", "mcp"), ("2026-08-02", "llm")]
def test_commit_skips_the_database_entirely_when_nothing_accumulated():
client = FakePrismaClient()
asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot={}))
- assert client.table.upserts == []
+ assert client.db.statements == []
# ── flush ─────────────────────────────────────────────────────────────────────
@@ -177,12 +200,12 @@ def test_flush_drains_and_commits():
asyncio.run(flush_gateway_requests(client, acc))
- assert len(client.table.upserts) == 1
+ assert len(client.db.statements) == 1
assert acc.drain() == {}
class ExplodingDB:
- def batch_(self):
+ async def execute_raw(self, query: str, *args: object) -> int:
raise RuntimeError("db gone")
@@ -208,10 +231,7 @@ def test_failed_flush_keeps_counts_for_the_next_attempt():
client = FakePrismaClient()
asyncio.run(flush_gateway_requests(client, acc))
- assert client.table.upserts[0]["data"]["update"] == {
- "successful_requests": {"increment": 1},
- "failed_requests": {"increment": 1},
- }
+ assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 1, 1)]
def test_restored_counts_merge_with_requests_recorded_meanwhile():
@@ -223,5 +243,272 @@ def test_restored_counts_merge_with_requests_recorded_meanwhile():
client = FakePrismaClient()
asyncio.run(flush_gateway_requests(client, acc))
- assert len(client.table.upserts) == 1
- assert client.table.upserts[0]["data"]["update"]["successful_requests"] == {"increment": 2}
+ assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
+
+
+class ExplodingDBWithInFlightRequest:
+ """Fails the write after a request has been recorded while it was in flight."""
+
+ def __init__(self, accumulator: GatewayRequestAccumulator) -> None:
+ self.accumulator = accumulator
+
+ async def execute_raw(self, query: str, *args: object) -> int:
+ _record(self.accumulator, 500)
+ raise RuntimeError("db gone")
+
+
+class ExplodingClientWithInFlightRequest:
+ def __init__(self, accumulator: GatewayRequestAccumulator) -> None:
+ self.db = ExplodingDBWithInFlightRequest(accumulator)
+
+
+def test_restore_keeps_requests_recorded_while_the_failed_write_was_in_flight():
+ acc = GatewayRequestAccumulator()
+ _record(acc, 200)
+ asyncio.run(flush_gateway_requests(ExplodingClientWithInFlightRequest(acc), acc))
+
+ client = FakePrismaClient()
+ asyncio.run(flush_gateway_requests(client, acc))
+
+ assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 1, 1)]
+
+
+# ── redis buffer ──────────────────────────────────────────────────────────────
+
+
+class FakeRedis:
+ def __init__(self) -> None:
+ self.lists: dict[str, list[str]] = {}
+
+ async def async_rpush(self, key: str, values: list[str]) -> int:
+ self.lists.setdefault(key, []).extend(values)
+ return len(self.lists[key])
+
+ async def async_lpop(self, key: str, count: int) -> list[str] | None:
+ queue = self.lists.get(key, [])
+ if not queue:
+ return None
+ popped, self.lists[key] = queue[:count], queue[count:]
+ return popped
+
+
+class FakePodLock:
+ def __init__(self, *, leader: bool) -> None:
+ self.leader = leader
+ self.held: list[str] = []
+ self.released: list[str] = []
+
+ async def acquire_lock(self, cronjob_id: str) -> bool:
+ self.held.append(cronjob_id)
+ return self.leader
+
+ async def release_lock(self, cronjob_id: str) -> None:
+ self.released.append(cronjob_id)
+
+
+class FakeLease:
+ """Redis-side view of the job lock: SET NX by pod id, re-entrant for the holder, freed only by release or TTL."""
+
+ def __init__(self) -> None:
+ self.holder: str | None = None
+
+
+class FakeLeasePodLock:
+ def __init__(self, lease: FakeLease, pod_id: str) -> None:
+ self.lease = lease
+ self.pod_id = pod_id
+
+ async def acquire_lock(self, cronjob_id: str) -> bool:
+ if self.lease.holder is None:
+ self.lease.holder = self.pod_id
+ return self.lease.holder == self.pod_id
+
+ async def release_lock(self, cronjob_id: str) -> None:
+ if self.lease.holder == self.pod_id:
+ self.lease.holder = None
+
+
+def _buffer(redis: FakeRedis, *, leader: bool) -> tuple[GatewayRequestRedisBuffer, FakePodLock]:
+ lock = FakePodLock(leader=leader)
+ return GatewayRequestRedisBuffer(redis_cache=redis, pod_lock_manager=lock), lock # pyright: ignore[reportArgumentType] # duck-typed fakes
+
+
+def test_non_leader_workers_push_to_redis_and_never_touch_the_database():
+ redis = FakeRedis()
+ client = FakePrismaClient()
+ for _ in range(3):
+ acc = GatewayRequestAccumulator()
+ _record(acc, 200)
+ buffer, _ = _buffer(redis, leader=False)
+ asyncio.run(flush_gateway_requests(client, acc, buffer))
+
+ assert client.db.statements == []
+ assert len(redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY]) == 3
+
+
+def test_leader_folds_every_workers_snapshot_into_one_statement():
+ """Fifty workers each flushing the same routes must cost the primary one statement, not fifty."""
+ redis = FakeRedis()
+ client = FakePrismaClient()
+ for _ in range(50):
+ acc = GatewayRequestAccumulator()
+ _record(acc, 200)
+ _record(acc, 500, route="/responses")
+ buffer, _ = _buffer(redis, leader=False)
+ asyncio.run(flush_gateway_requests(client, acc, buffer))
+
+ leader_acc = GatewayRequestAccumulator()
+ _record(leader_acc, 200)
+ leader, lock = _buffer(redis, leader=True)
+ asyncio.run(flush_gateway_requests(client, leader_acc, leader))
+
+ assert len(client.db.statements) == 1
+ assert _rows_written(client) == [
+ (_today(), "llm", "/chat/completions", 51, 0),
+ (_today(), "llm", "/responses", 0, 50),
+ ]
+ assert redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
+ assert lock.held == [GATEWAY_REQUESTS_JOB_NAME]
+ assert lock.released == []
+
+
+def test_leader_keeps_the_lease_so_staggered_pods_cost_one_statement_per_interval():
+ """Pods flush on their own clocks; without the lease each one would win the lock in turn and commit alone."""
+ redis = FakeRedis()
+ client = FakePrismaClient()
+ lease = FakeLease()
+ pods = tuple(
+ GatewayRequestRedisBuffer(redis_cache=redis, pod_lock_manager=FakeLeasePodLock(lease, f"pod-{i}")) # pyright: ignore[reportArgumentType] # duck-typed fakes
+ for i in range(4)
+ )
+
+ for _interval in range(3):
+ for pod in pods:
+ acc = GatewayRequestAccumulator()
+ _record(acc, 200)
+ asyncio.run(flush_gateway_requests(client, acc, pod))
+
+ assert lease.holder == "pod-0"
+ assert len(client.db.statements) == 3
+ assert [row[3] for row in _rows_written(client)] == [1, 4, 4]
+ assert len(redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY]) == 3
+
+
+def test_leader_drains_a_backlog_deeper_than_one_capped_pop():
+ """More workers than MAX_REDIS_BUFFER_DEQUEUE_COUNT must not leave a growing tail queued behind the cap."""
+ redis = FakeRedis()
+ client = FakePrismaClient()
+ workers = MAX_REDIS_BUFFER_DEQUEUE_COUNT * 2 + 1
+ for _ in range(workers):
+ acc = GatewayRequestAccumulator()
+ _record(acc, 200)
+ buffer, _ = _buffer(redis, leader=False)
+ asyncio.run(flush_gateway_requests(client, acc, buffer))
+
+ leader, _ = _buffer(redis, leader=True)
+ asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), leader))
+
+ assert len(client.db.statements) == 1
+ assert _rows_written(client) == [(_today(), "llm", "/chat/completions", workers, 0)]
+ assert redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
+
+
+def test_leader_with_nothing_buffered_writes_nothing():
+ redis = FakeRedis()
+ client = FakePrismaClient()
+ leader, lock = _buffer(redis, leader=True)
+
+ asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), leader))
+
+ assert client.db.statements == []
+ assert lock.released == []
+
+
+def test_leader_requeues_to_redis_when_the_database_commit_fails():
+ """Counts popped from Redis are gone from every worker; a failed commit must put them back."""
+ redis = FakeRedis()
+ acc = GatewayRequestAccumulator()
+ _record(acc, 200)
+ _record(acc, 200)
+ leader, lock = _buffer(redis, leader=True)
+
+ asyncio.run(flush_gateway_requests(ExplodingClient(), acc, leader))
+
+ assert len(redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY]) == 1
+ assert lock.released == []
+ assert acc.drain() == {}
+
+ client = FakePrismaClient()
+ retry, _ = _buffer(redis, leader=True)
+ asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), retry))
+ assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
+
+
+class ExplodingRedis(FakeRedis):
+ async def async_rpush(self, key: str, values: list[str]) -> int:
+ raise RuntimeError("redis gone")
+
+
+class UnreadableRedis(FakeRedis):
+ async def async_lpop(self, key: str, count: int) -> list[str] | None:
+ raise RuntimeError("redis gone mid-flush")
+
+
+class UnwritableRedis(FakeRedis):
+ """Pops succeed, pushes fail: a Redis that went read-only between the leader's pop and its re-queue."""
+
+ async def async_rpush(self, key: str, values: list[str]) -> int:
+ raise RuntimeError("redis read-only")
+
+
+def test_leader_keeps_popped_counts_in_memory_when_both_the_database_and_the_requeue_fail():
+ """The pop removed the only copy; if Redis will not take it back the leader itself must carry it."""
+ redis = FakeRedis()
+ worker_acc = GatewayRequestAccumulator()
+ _record(worker_acc, 200)
+ _record(worker_acc, 200)
+ worker, _ = _buffer(redis, leader=False)
+ asyncio.run(flush_gateway_requests(FakePrismaClient(), worker_acc, worker))
+
+ degraded = UnwritableRedis()
+ degraded.lists = redis.lists
+ leader_acc = GatewayRequestAccumulator()
+ leader, _ = _buffer(degraded, leader=True)
+ asyncio.run(flush_gateway_requests(ExplodingClient(), leader_acc, leader))
+ assert degraded.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
+
+ client = FakePrismaClient()
+ retry, _ = _buffer(redis, leader=True)
+ asyncio.run(flush_gateway_requests(client, leader_acc, retry))
+ assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
+
+
+def test_leader_whose_redis_read_fails_leaves_the_pushed_rows_for_the_next_flush():
+ """The scheduler job must not raise, and nothing is popped so nothing needs restoring anywhere."""
+ redis = UnreadableRedis()
+ acc = GatewayRequestAccumulator()
+ _record(acc, 200)
+ client = FakePrismaClient()
+ leader, _ = _buffer(redis, leader=True)
+
+ asyncio.run(flush_gateway_requests(client, acc, leader))
+
+ assert client.db.statements == []
+ assert acc.drain() == {}
+ assert len(redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY]) == 1
+
+
+def test_failed_redis_push_keeps_counts_locally_for_the_next_flush():
+ acc = GatewayRequestAccumulator()
+ _record(acc, 200)
+ _record(acc, 500)
+ buffer, lock = _buffer(ExplodingRedis(), leader=True)
+
+ asyncio.run(flush_gateway_requests(FakePrismaClient(), acc, buffer))
+
+ assert lock.held == []
+ assert acc.drain() == {
+ GatewayRequestKey(date=_today(), category="llm", route="/chat/completions"): (
+ GatewayRequestCounts(successful_requests=1, failed_requests=1)
+ )
+ }
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py
index 2b43720a126..615d06b0f42 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py
@@ -3,14 +3,19 @@
Test OpenAI Moderation Guardrail
"""
+import json
import os
+from typing import Final
from unittest.mock import MagicMock, patch
+import httpx
import pytest
+from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import UserAPIKeyAuth
+from litellm.proxy.common_utils.callback_utils import get_logging_caching_headers
from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import (
OpenAIModerationGuardrail,
)
@@ -989,3 +994,30 @@ async def test_openai_moderation_initialize_guardrail_forwards_streaming_flags()
assert guardrail.streaming_sampling_rate == 2
finally:
litellm.logging_callback_manager._reset_all_callbacks()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(("input_type", "stage"), [("request", "pre_call"), ("response", "post_call")])
+async def test_openai_moderation_records_moderation_id_as_scan_metadata(input_type: str, stage: str):
+ """Each moderation call's id is exposed with the guardrail name, stage and provider that produced it."""
+ payload: Final = {
+ "id": f"modr-{stage}",
+ "model": "omni-moderation-latest",
+ "results": [{"flagged": False, "categories": {}, "category_scores": {}, "category_applied_input_types": {}}],
+ }
+ http_client: Final = AsyncHTTPHandler()
+ http_client.client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _: httpx.Response(200, json=payload)))
+
+ with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
+ guardrail: Final = OpenAIModerationGuardrail(guardrail_name="openai-mod")
+ guardrail.async_handler = http_client
+ request_data: Final[dict[str, object]] = {"metadata": {}}
+
+ await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data=request_data, input_type=input_type)
+
+ headers: Final = get_logging_caching_headers(request_data)
+ assert headers is not None
+ assert headers["x-litellm-guardrail-scan-id"] == f"modr-{stage}"
+ assert json.loads(headers["x-litellm-guardrail-scan-metadata"]) == [
+ {"guardrail": "openai-mod", "stage": stage, "provider": "openai_moderation", "scan_id": f"modr-{stage}"}
+ ]
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py
index 8f29ba66814..3d7c6e06d94 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py
@@ -12,6 +12,7 @@ This test file follows LiteLLM's testing patterns and covers:
import copy
import json
from datetime import datetime
+from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@@ -5785,7 +5786,14 @@ class TestPanwAirsScanIdExposure:
headers = get_logging_caching_headers(data)
assert headers["x-litellm-guardrail-scan-id"] == "scan-abc-123"
- assert "x-litellm-guardrail-scan-metadata" not in headers
+ assert json.loads(headers["x-litellm-guardrail-scan-metadata"]) == [
+ {
+ "guardrail": handler.guardrail_name,
+ "stage": "pre_call",
+ "provider": "panw_prisma_airs",
+ "scan_id": "scan-abc-123",
+ }
+ ]
@pytest.mark.asyncio
async def test_request_and_response_scan_ids_are_both_exposed(self, user_api_key_dict):
@@ -5809,6 +5817,26 @@ class TestPanwAirsScanIdExposure:
headers = get_logging_caching_headers(data)
assert headers["x-litellm-guardrail-scan-id"] == "scan-abc-123,scan-response-456"
+ assert [(e["stage"], e["scan_id"]) for e in json.loads(headers["x-litellm-guardrail-scan-metadata"])] == [
+ ("pre_call", "scan-abc-123"),
+ ("post_call", "scan-response-456"),
+ ]
+
+ @pytest.mark.asyncio
+ async def test_apply_guardrail_response_scan_is_tagged_post_call(self):
+ from litellm.proxy.common_utils.callback_utils import get_logging_caching_headers
+
+ handler: Final = self._handler(self.ALLOW_SCAN_RESULT)
+ request_data: Final[dict[str, object]] = {"litellm_call_id": "test-call-id", "model": "gpt-4", "metadata": {}}
+
+ await handler.apply_guardrail(
+ inputs={"texts": ["Hello world"]}, request_data=request_data, input_type="response"
+ )
+
+ headers: Final = get_logging_caching_headers(request_data)
+ assert headers is not None
+ entries: Final = json.loads(headers["x-litellm-guardrail-scan-metadata"])
+ assert [(e["stage"], e["provider"]) for e in entries] == [("post_call", "panw_prisma_airs")]
@pytest.mark.asyncio
async def test_repeated_scan_id_is_not_duplicated(self, user_api_key_dict):
@@ -5850,6 +5878,8 @@ class TestPanwAirsScanIdExposure:
assert "guardrail_scan_ids" in _UNTRUSTED_METADATA_CONTROL_FIELDS
assert "guardrail_scan_ids" in _UNTRUSTED_ROOT_CONTROL_FIELDS
+ assert "guardrail_scan_metadata" in _UNTRUSTED_METADATA_CONTROL_FIELDS
+ assert "guardrail_scan_metadata" in _UNTRUSTED_ROOT_CONTROL_FIELDS
class TestPanwAirsBlockedErrorDetailPassthrough:
"""Regression tests for the full AIRS scan response on blocks.
diff --git a/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py b/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py
index ec62cc47018..7ece35ceedf 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py
@@ -4,13 +4,17 @@ Tests for cost tracking settings management endpoints.
Tests the GET and PATCH endpoints for managing cost discount configuration.
"""
+from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi.testclient import TestClient
+from pydantic import ValidationError
import litellm
+from litellm._internal_context import pinned_billing_time
+from litellm.proxy._types import CostEstimateRequest
from litellm.proxy.management_endpoints.cost_tracking_settings import router
from litellm.proxy.proxy_server import app
@@ -789,13 +793,13 @@ INPUT_TOKENS = 1000
OUTPUT_TOKENS = 500
-def _router_pricing(**pricing: float) -> MagicMock:
+def _router_pricing(model: str = AN_UNDERLYING_MODEL, **pricing: float) -> MagicMock:
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
{
"model_name": AN_ALIAS,
"litellm_params": {
- "model": AN_UNDERLYING_MODEL,
+ "model": model,
"custom_llm_provider": "openai",
**pricing,
},
@@ -811,9 +815,7 @@ async def _estimate(mock_router: MagicMock | None, model: str = AN_ALIAS, **over
request = CostEstimateRequest(
model=model,
- input_tokens=INPUT_TOKENS,
- output_tokens=OUTPUT_TOKENS,
- **overrides,
+ **{"input_tokens": INPUT_TOKENS, "output_tokens": OUTPUT_TOKENS, **overrides},
)
with patch( # test-quality-ok: proxy_server module global is the endpoint's only injection point
"litellm.proxy.proxy_server.llm_router", mock_router
@@ -909,3 +911,299 @@ class TestEstimateCostPeriodTotals:
assert response.cost_per_request == pytest.approx(0.0022)
assert response.daily_margin_cost == pytest.approx(0.02)
assert response.daily_cost == pytest.approx(0.22)
+
+
+CACHE_READ_TOKENS = 800
+CACHE_CREATION_TOKENS = 100
+REASONING_TOKENS = 200
+TEXT_INPUT_TOKENS = INPUT_TOKENS - CACHE_READ_TOKENS - CACHE_CREATION_TOKENS
+TEXT_OUTPUT_TOKENS = OUTPUT_TOKENS - REASONING_TOKENS
+
+
+async def _estimate_with_cache_and_reasoning(mock_router: MagicMock | None, model: str = AN_ALIAS, **overrides: int):
+ return await _estimate(
+ mock_router,
+ model=model,
+ cache_read_input_tokens=CACHE_READ_TOKENS,
+ cache_creation_input_tokens=CACHE_CREATION_TOKENS,
+ reasoning_tokens=REASONING_TOKENS,
+ **overrides,
+ )
+
+
+class TestEstimateCostCacheAndReasoningTokens:
+ @pytest.mark.asyncio
+ async def test_a_mapped_model_bills_cache_and_reasoning_tokens_at_their_own_rates(self, monkeypatch):
+ monkeypatch.setitem(
+ litellm.model_cost,
+ A_MAPPED_MODEL,
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "cache_read_input_token_cost": 3e-7,
+ "cache_creation_input_token_cost": 3.75e-6,
+ "output_cost_per_reasoning_token": 1e-5,
+ "litellm_provider": "openai",
+ "mode": "chat",
+ },
+ )
+
+ response = await _estimate_with_cache_and_reasoning(None, model=A_MAPPED_MODEL, num_requests_per_day=10)
+
+ assert response.cache_read_cost_per_request == pytest.approx(CACHE_READ_TOKENS * 3e-7)
+ assert response.cache_creation_cost_per_request == pytest.approx(CACHE_CREATION_TOKENS * 3.75e-6)
+ assert response.reasoning_cost_per_request == pytest.approx(REASONING_TOKENS * 1e-5)
+ assert response.input_cost_per_request == pytest.approx(
+ TEXT_INPUT_TOKENS * 3e-6 + CACHE_READ_TOKENS * 3e-7 + CACHE_CREATION_TOKENS * 3.75e-6
+ )
+ assert response.output_cost_per_request == pytest.approx(TEXT_OUTPUT_TOKENS * 15e-6 + REASONING_TOKENS * 1e-5)
+ assert response.cost_per_request == pytest.approx(
+ response.input_cost_per_request + response.output_cost_per_request
+ )
+ assert response.daily_cache_read_cost == pytest.approx(10 * CACHE_READ_TOKENS * 3e-7)
+ assert response.daily_cache_creation_cost == pytest.approx(10 * CACHE_CREATION_TOKENS * 3.75e-6)
+ assert response.daily_reasoning_cost == pytest.approx(10 * REASONING_TOKENS * 1e-5)
+ assert response.monthly_cache_read_cost is None
+ assert response.cache_read_input_token_cost == pytest.approx(3e-7)
+ assert response.cache_creation_input_token_cost == pytest.approx(3.75e-6)
+ assert response.output_cost_per_reasoning_token == pytest.approx(1e-5)
+ assert (
+ response.cache_read_input_tokens,
+ response.cache_creation_input_tokens,
+ response.reasoning_tokens,
+ ) == (CACHE_READ_TOKENS, CACHE_CREATION_TOKENS, REASONING_TOKENS)
+
+ @pytest.mark.asyncio
+ async def test_a_model_without_cache_or_reasoning_prices_estimates_what_the_proxy_bills(self, monkeypatch):
+ """The cost calculator bills cache tokens of a cost-map model without cache prices at zero
+ and its reasoning tokens at the output rate. The estimate reports those effective rates."""
+ monkeypatch.setitem(
+ litellm.model_cost,
+ A_MAPPED_MODEL,
+ {"input_cost_per_token": 5e-6, "output_cost_per_token": 6e-6, "litellm_provider": "openai", "mode": "chat"},
+ )
+
+ response = await _estimate_with_cache_and_reasoning(None, model=A_MAPPED_MODEL)
+
+ assert response.cache_read_cost_per_request == 0.0
+ assert response.cache_creation_cost_per_request == 0.0
+ assert response.reasoning_cost_per_request == pytest.approx(REASONING_TOKENS * 6e-6)
+ assert response.input_cost_per_request == pytest.approx(TEXT_INPUT_TOKENS * 5e-6)
+ assert response.cost_per_request == pytest.approx(TEXT_INPUT_TOKENS * 5e-6 + OUTPUT_TOKENS * 6e-6)
+ assert response.cache_read_input_token_cost == 0.0
+ assert response.cache_creation_input_token_cost == 0.0
+ assert response.output_cost_per_reasoning_token == pytest.approx(6e-6)
+
+ @pytest.mark.asyncio
+ async def test_a_request_without_cache_or_reasoning_tokens_estimates_as_before(self, monkeypatch):
+ monkeypatch.setitem(
+ litellm.model_cost,
+ A_MAPPED_MODEL,
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "cache_read_input_token_cost": 3e-7,
+ "cache_creation_input_token_cost": 3.75e-6,
+ "output_cost_per_reasoning_token": 1e-5,
+ "litellm_provider": "openai",
+ "mode": "chat",
+ },
+ )
+
+ response = await _estimate(None, model=A_MAPPED_MODEL, num_requests_per_day=10)
+
+ assert response.cost_per_request == pytest.approx(INPUT_TOKENS * 3e-6 + OUTPUT_TOKENS * 15e-6)
+ assert response.cache_read_cost_per_request == 0.0
+ assert response.cache_creation_cost_per_request == 0.0
+ assert response.reasoning_cost_per_request == 0.0
+ assert response.daily_cache_read_cost == 0.0
+ assert response.daily_reasoning_cost == 0.0
+
+ @pytest.mark.asyncio
+ async def test_a_custom_priced_deployment_bills_cache_and_reasoning_tokens_from_its_flat_rates(self):
+ response = await _estimate_with_cache_and_reasoning(
+ _router_pricing(input_cost_per_token=1e-6, output_cost_per_token=2e-6, cache_read_input_token_cost=1e-7)
+ )
+
+ assert response.cache_read_cost_per_request == pytest.approx(CACHE_READ_TOKENS * 1e-7)
+ assert response.cache_creation_cost_per_request == pytest.approx(CACHE_CREATION_TOKENS * 1e-6)
+ assert response.reasoning_cost_per_request == pytest.approx(REASONING_TOKENS * 2e-6)
+ assert response.cost_per_request == pytest.approx(
+ TEXT_INPUT_TOKENS * 1e-6 + CACHE_READ_TOKENS * 1e-7 + CACHE_CREATION_TOKENS * 1e-6 + OUTPUT_TOKENS * 2e-6
+ )
+ assert response.cache_read_input_token_cost == pytest.approx(1e-7)
+ assert response.cache_creation_input_token_cost == pytest.approx(1e-6)
+ assert response.output_cost_per_reasoning_token == pytest.approx(2e-6)
+
+ @pytest.mark.asyncio
+ async def test_a_custom_priced_deployment_of_a_mapped_model_inherits_its_built_in_cache_rates(self, monkeypatch):
+ monkeypatch.setitem(
+ litellm.model_cost,
+ A_MAPPED_MODEL,
+ {
+ "input_cost_per_token": 5e-6,
+ "output_cost_per_token": 6e-6,
+ "cache_read_input_token_cost": 5e-7,
+ "cache_creation_input_token_cost": 6.25e-6,
+ "litellm_provider": "openai",
+ "mode": "chat",
+ },
+ )
+
+ response = await _estimate_with_cache_and_reasoning(
+ _router_pricing(model=A_MAPPED_MODEL, input_cost_per_token=1e-6, output_cost_per_token=2e-6)
+ )
+
+ assert response.cache_read_cost_per_request == pytest.approx(CACHE_READ_TOKENS * 5e-7)
+ assert response.cache_creation_cost_per_request == pytest.approx(CACHE_CREATION_TOKENS * 6.25e-6)
+ assert response.input_cost_per_request == pytest.approx(
+ TEXT_INPUT_TOKENS * 1e-6 + CACHE_READ_TOKENS * 5e-7 + CACHE_CREATION_TOKENS * 6.25e-6
+ )
+ assert response.cache_read_input_token_cost == pytest.approx(5e-7)
+ assert response.cache_creation_input_token_cost == pytest.approx(6.25e-6)
+
+ @pytest.mark.asyncio
+ async def test_a_tiered_model_reports_the_rates_its_lines_were_billed_at(self, monkeypatch):
+ """Above a token tier the calculator bills every line at the tier's rate, so the reported
+ rates must be the tier's too: each line equals its token count times the rate next to it."""
+ monkeypatch.setitem(
+ litellm.model_cost,
+ A_MAPPED_MODEL,
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "cache_read_input_token_cost": 3e-7,
+ "cache_creation_input_token_cost": 3.75e-6,
+ "input_cost_per_token_above_200k_tokens": 6e-6,
+ "output_cost_per_token_above_200k_tokens": 3e-5,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-7,
+ "cache_creation_input_token_cost_above_200k_tokens": 7.5e-6,
+ "litellm_provider": "openai",
+ "mode": "chat",
+ },
+ )
+
+ response = await _estimate(
+ None,
+ model=A_MAPPED_MODEL,
+ input_tokens=250_000,
+ cache_read_input_tokens=200_000,
+ cache_creation_input_tokens=10_000,
+ output_tokens=1_000,
+ reasoning_tokens=200,
+ )
+
+ assert response.input_cost_per_token == pytest.approx(6e-6)
+ assert response.output_cost_per_token == pytest.approx(3e-5)
+ assert response.cache_read_input_token_cost == pytest.approx(6e-7)
+ assert response.cache_creation_input_token_cost == pytest.approx(7.5e-6)
+ assert response.output_cost_per_reasoning_token == pytest.approx(3e-5)
+ assert response.cache_read_cost_per_request == pytest.approx(200_000 * response.cache_read_input_token_cost)
+ assert response.cache_creation_cost_per_request == pytest.approx(
+ 10_000 * response.cache_creation_input_token_cost
+ )
+ assert response.reasoning_cost_per_request == pytest.approx(200 * response.output_cost_per_reasoning_token)
+ assert response.input_cost_per_request == pytest.approx(
+ 40_000 * response.input_cost_per_token
+ + response.cache_read_cost_per_request
+ + response.cache_creation_cost_per_request
+ )
+ assert response.output_cost_per_request == pytest.approx(1_000 * response.output_cost_per_token)
+
+ @pytest.mark.asyncio
+ async def test_a_quote_prices_its_totals_and_its_rates_at_the_same_moment(self, monkeypatch):
+ """The totals and the reported rates resolve off-peak pricing on separate paths. A quote
+ taken as a window opens must not bill on one side of it and report rates from the other."""
+ monkeypatch.setitem(
+ litellm.model_cost,
+ A_MAPPED_MODEL,
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "off_peak_pricing": {
+ "hours_utc": "02:00-03:00",
+ "input_cost_per_token": 1e-6,
+ "output_cost_per_token": 5e-6,
+ },
+ "litellm_provider": "openai",
+ "mode": "chat",
+ },
+ )
+
+ with pinned_billing_time(datetime(2026, 1, 1, 2, 30, tzinfo=timezone.utc)):
+ response = await _estimate(None, model=A_MAPPED_MODEL)
+
+ assert response.input_cost_per_token == pytest.approx(1e-6)
+ assert response.output_cost_per_token == pytest.approx(5e-6)
+ assert response.input_cost_per_request == pytest.approx(INPUT_TOKENS * response.input_cost_per_token)
+ assert response.output_cost_per_request == pytest.approx(OUTPUT_TOKENS * response.output_cost_per_token)
+
+
+ @pytest.mark.asyncio
+ async def test_an_unrouted_model_reports_the_rates_of_the_provider_the_calculator_inferred(self, monkeypatch):
+ """The cost calculator infers a provider this endpoint never resolved, and the provider decides
+ whether a tier threshold is inclusive. xai bills a request sitting exactly on the 200k threshold
+ at the tier rate, so the reported rates have to be the tier's rather than the sub-tier base."""
+ an_xai_model = "xai/tiered-model"
+ monkeypatch.setitem(
+ litellm.model_cost,
+ an_xai_model,
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "cache_read_input_token_cost": 3e-7,
+ "input_cost_per_token_above_200k_tokens": 6e-6,
+ "output_cost_per_token_above_200k_tokens": 3e-5,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-7,
+ "litellm_provider": "xai",
+ "mode": "chat",
+ },
+ )
+
+ response = await _estimate(
+ None,
+ model=an_xai_model,
+ input_tokens=200_000,
+ cache_read_input_tokens=100_000,
+ output_tokens=1_000,
+ )
+
+ assert response.input_cost_per_token == pytest.approx(6e-6)
+ assert response.output_cost_per_token == pytest.approx(3e-5)
+ assert response.cache_read_input_token_cost == pytest.approx(6e-7)
+ assert response.cache_read_cost_per_request == pytest.approx(100_000 * response.cache_read_input_token_cost)
+ assert response.input_cost_per_request == pytest.approx(
+ 100_000 * response.input_cost_per_token + response.cache_read_cost_per_request
+ )
+ assert response.output_cost_per_request == pytest.approx(1_000 * response.output_cost_per_token)
+
+
+class TestCostEstimateRequestTokenSubsets:
+ def test_cache_tokens_beyond_the_input_tokens_are_rejected(self):
+ with pytest.raises(ValidationError, match="cannot exceed input_tokens"):
+ CostEstimateRequest(
+ model=AN_ALIAS,
+ input_tokens=INPUT_TOKENS,
+ output_tokens=OUTPUT_TOKENS,
+ cache_read_input_tokens=INPUT_TOKENS,
+ cache_creation_input_tokens=1,
+ )
+
+ def test_reasoning_tokens_beyond_the_output_tokens_are_rejected(self):
+ with pytest.raises(ValidationError, match="cannot exceed output_tokens"):
+ CostEstimateRequest(
+ model=AN_ALIAS,
+ input_tokens=INPUT_TOKENS,
+ output_tokens=OUTPUT_TOKENS,
+ reasoning_tokens=OUTPUT_TOKENS + 1,
+ )
+
+ def test_the_endpoint_answers_422_when_cache_tokens_exceed_input_tokens(self):
+ response = client.post(
+ "/cost/estimate",
+ headers={"Authorization": "Bearer sk-1234"},
+ json={"model": AN_ALIAS, "input_tokens": 1000, "output_tokens": 100, "cache_read_input_tokens": 8000},
+ )
+
+ assert response.status_code == 422
+ assert "cannot exceed input_tokens" in response.text
diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py
index a873a367eab..d2aeba18f7d 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py
@@ -17975,6 +17975,32 @@ def test_key_request_blank_organization_id_is_unset():
assert UpdateKeyRequest(key="sk-1", organization_id="org-1").organization_id == "org-1"
+def test_update_key_request_blank_team_id_is_not_a_team_change():
+ from litellm.proxy._types import UpdateKeyRequest
+ from litellm.proxy.management_endpoints.key_management_endpoints import (
+ is_different_team,
+ )
+
+ blank = UpdateKeyRequest(key="sk-1", team_id="", key_alias="renamed")
+ assert blank.team_id is None
+ assert "team_id" not in blank.model_dump(exclude_unset=True)
+ assert blank.model_dump(exclude_unset=True) == {"key": "sk-1", "key_alias": "renamed"}
+ assert is_different_team(data=blank, existing_key_row=LiteLLM_VerificationToken(token="hashed")) is False
+ assert (
+ is_different_team(data=blank, existing_key_row=LiteLLM_VerificationToken(token="hashed", team_id="team-1"))
+ is False
+ )
+ assert "team_id" in UpdateKeyRequest(key="sk-1", team_id=None).model_dump(exclude_unset=True)
+ assert UpdateKeyRequest(key="sk-1", team_id="team-1").team_id == "team-1"
+ assert (
+ is_different_team(
+ data=UpdateKeyRequest(key="sk-1", team_id="team-1"),
+ existing_key_row=LiteLLM_VerificationToken(token="hashed"),
+ )
+ is True
+ )
+
+
def test_key_generation_check_blank_team_id_uses_personal_permissions(monkeypatch):
"""key_generation_check with team_id="" must take the personal-key path instead
of failing the team lookup with "Unable to find team object" (LIT-3925)."""
diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
index 051e6bed4fd..2f6561046b1 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
@@ -2205,6 +2205,50 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name):
assert update_call_kwargs.get("include", {}).get("object_permission") is True
+@pytest.mark.asyncio
+@pytest.mark.parametrize("endpoint_name", ["team_model_add", "team_model_delete"])
+async def test_team_model_add_delete_keep_model_aliases_in_team_cache(endpoint_name, monkeypatch):
+ """LIT-5858: Prisma only returns `litellm_model_table` when the `update` asks for it, so the refreshed
+ cache entry lost the team's model aliases and JWT alias requests 403'd until the next DB read."""
+ from litellm.proxy._types import TeamModelAddRequest, TeamModelDeleteRequest
+ from litellm.proxy.auth.team_grants import team_model_aliases
+ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
+ from litellm.proxy.management_endpoints.team_endpoints import team_model_add, team_model_delete
+
+ columns = {"team_id": "team-1234", "models": ["gpt-4o", "openai/*"]}
+ alias_table = {"id": 1, "model_aliases": '{"fast": "gpt-4o"}', "created_by": "admin", "updated_by": "admin"}
+
+ async def update(where, data, include=None):
+ row = {**columns, "litellm_model_table": alias_table} if (include or {}).get("litellm_model_table") else columns
+ return SimpleNamespace(team_id="team-1234", model_dump=lambda: row)
+
+ prisma_client = MagicMock()
+ prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(model_dump=lambda: columns))
+ prisma_client.db.litellm_teamtable.update = AsyncMock(side_effect=update)
+ prisma_client.db.execute_raw = AsyncMock(return_value=None)
+ cache = UserApiKeyCache()
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
+ monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
+ monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", None)
+
+ admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin")
+ if endpoint_name == "team_model_add":
+ await team_model_add(
+ data=TeamModelAddRequest(team_id="team-1234", models=["team-byok-1"]),
+ http_request=MagicMock(),
+ user_api_key_dict=admin,
+ )
+ else:
+ await team_model_delete(
+ data=TeamModelDeleteRequest(team_id="team-1234", models=["openai/*"]),
+ http_request=MagicMock(),
+ user_api_key_dict=admin,
+ )
+
+ cached_team = await cache.async_get_cache(key="team_id:team-1234", model_type=LiteLLM_TeamTableCachedObj)
+ assert team_model_aliases(cached_team) == {"fast": "gpt-4o"}
+
+
@pytest.mark.asyncio
@pytest.mark.parametrize(
"endpoint_name",
diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py
index 86e97a334df..c343652efd9 100644
--- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py
+++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py
@@ -4,11 +4,12 @@ Pins covered:
- ``get_current_spend``
- ``increment_spend_counters``
- ``_reconcile_budget_reservation_for_counter_update``
-- ``_increment_end_user_and_tag_spend_counters``
-- ``_increment_org_spend_counter``
-- ``_init_and_increment_unreserved_spend_counter``
-- ``_init_and_increment_spend_counter``
-- ``_init_and_increment_window_spend_counter``
+- ``_prepare_end_user_and_tag_spend_increments``
+- ``_prepare_org_spend_increment``
+- ``_prepare_unreserved_spend_counter_increment``
+- ``_prepare_spend_counter_increment``
+- ``_prepare_window_spend_counter_increment``
+- ``_apply_spend_counter_increments``
- ``_ensure_spend_counter_initialized``
- ``_get_source_cache_base_spend``
- ``_ensure_window_spend_counter_initialized``
@@ -48,9 +49,7 @@ def _make_spend_counter_cache(
cache.in_memory_cache.delete_cache = MagicMock()
if with_redis:
cache.redis_cache = MagicMock()
- cache.redis_cache.async_get_cache = AsyncMock(
- return_value=redis_get_value, side_effect=redis_get_side_effect
- )
+ cache.redis_cache.async_get_cache = AsyncMock(return_value=redis_get_value, side_effect=redis_get_side_effect)
cache.redis_cache.async_increment = AsyncMock(
return_value=redis_increment_value,
side_effect=redis_increment_side_effect,
@@ -58,6 +57,8 @@ def _make_spend_counter_cache(
cache.redis_cache.async_delete_cache = AsyncMock()
cache.redis_cache.async_set_cache = AsyncMock()
cache.redis_cache.async_set_max = AsyncMock()
+ cache.redis_cache.async_increment_pipeline = AsyncMock(return_value=None)
+ cache.redis_cache.get_ttl = MagicMock(return_value=None)
else:
cache.redis_cache = None
cache.async_increment_cache = AsyncMock(return_value=redis_increment_value)
@@ -70,9 +71,7 @@ def _make_spend_counter_cache(
def _make_user_api_key_cache(get_value=None, get_side_effect=None):
cache = MagicMock()
- cache.async_get_cache = AsyncMock(
- return_value=get_value, side_effect=get_side_effect
- )
+ cache.async_get_cache = AsyncMock(return_value=get_value, side_effect=get_side_effect)
cache.async_set_cache_pipeline = AsyncMock()
return cache
@@ -109,9 +108,7 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory(monkeypatch
)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- result = await ps.get_current_spend(
- counter_key="spend:key:abc", fallback_spend=99.0
- )
+ result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=99.0)
assert result == 17.0
@@ -136,9 +133,7 @@ async def test_get_current_spend_floors_stale_low_counter_against_db(monkeypatch
# the stale counter is repaired up to the authoritative DB value via a
# monotonic set-max so other workers read the corrected total, and a
# concurrent increment cannot be clobbered
- fake_cache.redis_cache.async_set_max.assert_awaited_once_with(
- key="spend:key:abc", value=12.0
- )
+ fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key="spend:key:abc", value=12.0)
@pytest.mark.asyncio
@@ -169,9 +164,7 @@ async def test_get_current_spend_no_floor_without_max_budget(monkeypatch):
from_db = AsyncMock(return_value=12.0)
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
- result = await ps.get_current_spend(
- counter_key="spend:key:abc", fallback_spend=12.0
- )
+ result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0)
assert result == 2.0
assert from_db.await_count == 0
@@ -210,12 +203,8 @@ async def test_get_current_spend_floor_caches_db_read(monkeypatch):
from_db = AsyncMock(return_value=12.0)
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
- first = await ps.get_current_spend(
- counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0
- )
- second = await ps.get_current_spend(
- counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0
- )
+ first = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0)
+ second = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0)
assert first == 12.0
assert second == 12.0
@@ -336,9 +325,7 @@ async def test_get_current_spend_floors_window_against_spend_logs(monkeypatch):
assert result == 15.0
assert wfsl.await_count == 1
- fake_cache.redis_cache.async_set_max.assert_awaited_once_with(
- key=counter_key, value=15.0
- )
+ fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key=counter_key, value=15.0)
def _make_window_spend_prisma(row=None, spend_logs_total=0.0):
@@ -379,9 +366,7 @@ async def test_get_current_spend_floors_window_against_maintained_row(monkeypatc
assert result == 15.0
fake_prisma.db.litellm_spendlogs.group_by.assert_not_awaited()
- fake_cache.redis_cache.async_set_max.assert_awaited_once_with(
- key=counter_key, value=15.0
- )
+ fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key=counter_key, value=15.0)
@pytest.mark.asyncio
@@ -393,9 +378,7 @@ async def test_get_current_spend_floors_window_against_logs_when_row_stale(monke
window_start = datetime(2026, 1, 8, tzinfo=timezone.utc)
fake_prisma = _make_window_spend_prisma(
- row=SimpleNamespace(
- window_start=window_start - timedelta(days=7), spend=999.0
- ),
+ row=SimpleNamespace(window_start=window_start - timedelta(days=7), spend=999.0),
spend_logs_total=15.0,
)
fake_cache = _make_spend_counter_cache(redis_get_value=2.0)
@@ -423,21 +406,13 @@ async def test_get_current_spend_fail_closed_rejects_when_unverifiable(monkeypat
rather than admitted on an unverifiable budget."""
from fastapi import HTTPException
- fake_cache = _make_spend_counter_cache(
- redis_get_side_effect=RuntimeError("redis down")
- )
+ fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- monkeypatch.setattr(
- ps, "general_settings", {"fail_closed_budget_enforcement": True}
- )
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
with pytest.raises(HTTPException) as exc:
- await ps.get_current_spend(
- counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
- )
+ await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
assert exc.value.status_code == 503
@@ -445,18 +420,12 @@ async def test_get_current_spend_fail_closed_rejects_when_unverifiable(monkeypat
async def test_get_current_spend_fail_closed_off_admits_when_unverifiable(monkeypatch):
"""Default (flag off): an unverifiable read keeps the existing behavior and
admits using the cached fallback — no new rejection."""
- fake_cache = _make_spend_counter_cache(
- redis_get_side_effect=RuntimeError("redis down")
- )
+ fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "general_settings", {})
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
- result = await ps.get_current_spend(
- counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
- )
+ result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
assert result == 1.0
@@ -466,13 +435,9 @@ async def test_get_current_spend_fail_closed_admits_when_redis_verified(monkeypa
authoritative, so an under-budget request is admitted normally."""
fake_cache = _make_spend_counter_cache(redis_get_value=1.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- monkeypatch.setattr(
- ps, "general_settings", {"fail_closed_budget_enforcement": True}
- )
+ monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
- result = await ps.get_current_spend(
- counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
- )
+ result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
assert result == 1.0
@@ -481,16 +446,10 @@ async def test_get_current_spend_fail_closed_allows_authoritative_fallback(monke
"""End-user/tag callers pass fallback_authoritative=True (their spend is
loaded fresh from the DB in auth), so fail-closed does not reject them even
when the counter path is unreadable."""
- fake_cache = _make_spend_counter_cache(
- redis_get_side_effect=RuntimeError("redis down")
- )
+ fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- monkeypatch.setattr(
- ps, "general_settings", {"fail_closed_budget_enforcement": True}
- )
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
result = await ps.get_current_spend(
counter_key="spend:end_user:e1",
@@ -508,9 +467,7 @@ async def test_get_current_spend_strict_floors_when_fallback_also_stale(monkeypa
re-checks the authoritative DB and enforces against it."""
fake_cache = _make_spend_counter_cache(redis_get_value=0.00001)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- monkeypatch.setattr(
- ps, "general_settings", {"fail_closed_budget_enforcement": True}
- )
+ monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
from_db = AsyncMock(return_value=0.5)
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
@@ -532,9 +489,7 @@ async def test_get_current_spend_strict_floors_when_fallback_also_stale(monkeypa
@pytest.mark.asyncio
async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
- fake_cache = _make_spend_counter_cache(
- redis_get_value=None, redis_increment_value=5.0
- )
+ fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=5.0)
fake_user_cache = _make_user_api_key_cache(get_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
@@ -543,9 +498,7 @@ async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
async def _fake_coalesced(**kwargs):
return None
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced)
- )
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced))
await ps.increment_spend_counters(
token="hashed-tok",
@@ -554,25 +507,36 @@ async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
response_cost=5.0,
)
+ pipeline = fake_cache.redis_cache.async_increment_pipeline
+ pipeline.assert_awaited_once()
+ increment_list = pipeline.await_args.kwargs["increment_list"]
+ assert {op["key"] for op in increment_list} == {
+ "spend:key:hashed-tok",
+ "spend:team:t1",
+ "spend:team_member:u1:t1",
+ "spend:user:u1",
+ }
+ assert all(op["increment_value"] == 5.0 for op in increment_list)
observed = {
"redis_increment_called": fake_cache.redis_cache.async_increment.called,
- "increment_calls": fake_cache.redis_cache.async_increment.call_count,
+ "pipeline_calls": pipeline.await_count,
"user_cache_used": fake_user_cache.async_get_cache.called,
}
assert normalize(observed) == {
- "redis_increment_called": True,
- "increment_calls": 4,
+ "redis_increment_called": False,
+ "pipeline_calls": 1,
"user_cache_used": True,
}
class _ConcurrencyProbe:
- """Stand-in for redis_cache.async_increment that pins concurrency.
+ """Stand-in for redis_cache.async_get_cache that pins concurrency.
- Each call registers itself as in-flight and blocks on ``release`` until the
- test lets it proceed. ``all_arrived`` fires once ``expected`` distinct scope
- increments are simultaneously suspended here, which can only happen if the
- per-scope increments are gathered rather than awaited one after another.
+ Each warm-check read registers itself as in-flight and blocks on ``release``
+ until the test lets it proceed. ``all_arrived`` fires once ``expected``
+ distinct scope warm-checks are simultaneously suspended here, which can only
+ happen if the per-scope prepares are gathered rather than awaited one after
+ another.
"""
def __init__(self, expected_concurrency: int):
@@ -581,36 +545,45 @@ class _ConcurrencyProbe:
self.max_in_flight = 0
self.all_arrived = asyncio.Event()
self.release = asyncio.Event()
- self.values: dict[str, float] = {}
+ self.keys: list[str] = []
- async def async_increment(self, *, key, value, refresh_ttl=True):
+ async def async_get_cache(self, *, key, **kwargs):
self.in_flight += 1
self.max_in_flight = max(self.max_in_flight, self.in_flight)
+ self.keys.append(key)
if self.in_flight >= self.expected:
self.all_arrived.set()
if not self.release.is_set():
await self.release.wait()
self.in_flight -= 1
- self.values[key] = self.values.get(key, 0.0) + value
- return self.values[key]
+ return 1.0
@pytest.mark.asyncio
async def test_increment_spend_counters_runs_scopes_concurrently(monkeypatch):
"""The six independent scopes (key, team, team_member, user, end_user+tags,
- org) must be incremented concurrently. The probe only fires once all six are
- suspended in async_increment at the same time, which is impossible if the
+ org) must prepare their increments concurrently. The probe only fires once
+ all eight warm-check reads (one per counter: 6 scopes + 2 tags) are
+ suspended in async_get_cache at the same time, which is impossible if the
awaits are chained sequentially."""
- probe = _ConcurrencyProbe(expected_concurrency=6)
- fake_cache = _make_spend_counter_cache(redis_get_value=None)
- fake_cache.redis_cache.async_increment = probe.async_increment
+ probe = _ConcurrencyProbe(expected_concurrency=8)
+ fake_cache = _make_spend_counter_cache()
+ fake_cache.redis_cache.async_get_cache = probe.async_get_cache
+ recorded: dict[str, float] = {}
+
+ async def _record_pipeline(increment_list, **_):
+ results = []
+ for op in increment_list:
+ recorded[op["key"]] = recorded.get(op["key"], 0.0) + op["increment_value"]
+ results.append(recorded[op["key"]])
+ return results
+
+ fake_cache.redis_cache.async_increment_pipeline = AsyncMock(side_effect=_record_pipeline)
fake_user_cache = _make_user_api_key_cache(get_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
task = asyncio.create_task(
ps.increment_spend_counters(
@@ -630,16 +603,16 @@ async def test_increment_spend_counters_runs_scopes_concurrently(monkeypatch):
probe.release.set()
await task
pytest.fail(
- "scope increments did not run concurrently; sequential awaits "
- f"detected (peak in-flight was {probe.max_in_flight}, expected 6)"
+ "scope prepares did not run concurrently; sequential awaits "
+ f"detected (peak in-flight was {probe.max_in_flight}, expected 8)"
)
- assert probe.in_flight == 6
- assert probe.max_in_flight == 6
+ assert probe.in_flight == 8
+ assert probe.max_in_flight == 8
probe.release.set()
await task
- assert probe.values == {
+ assert recorded == {
"spend:key:hashed-tok": 5.0,
"spend:team:t1": 5.0,
"spend:team_member:u1:t1": 5.0,
@@ -659,26 +632,25 @@ async def test_increment_spend_counters_skips_reserved_counter_keys(monkeypatch)
import litellm.proxy.spend_tracking.budget_reservation as br
reserved = {"spend:key:hashed-tok", "spend:org:org1"}
- monkeypatch.setattr(
- br, "get_reserved_counter_keys", MagicMock(return_value=set(reserved))
- )
+ monkeypatch.setattr(br, "get_reserved_counter_keys", MagicMock(return_value=set(reserved)))
monkeypatch.setattr(br, "reconcile_budget_reservation", AsyncMock())
recorded: dict[str, float] = {}
- async def _record_increment(*, key, value, refresh_ttl=True):
- recorded[key] = recorded.get(key, 0.0) + value
- return recorded[key]
+ async def _record_pipeline(increment_list, **_):
+ results = []
+ for op in increment_list:
+ recorded[op["key"]] = recorded.get(op["key"], 0.0) + op["increment_value"]
+ results.append(recorded[op["key"]])
+ return results
fake_cache = _make_spend_counter_cache(redis_get_value=None)
- fake_cache.redis_cache.async_increment = _record_increment
+ fake_cache.redis_cache.async_increment_pipeline = AsyncMock(side_effect=_record_pipeline)
fake_user_cache = _make_user_api_key_cache(get_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
reservation = {"finalized": False}
await ps.increment_spend_counters(
@@ -708,27 +680,46 @@ async def test_increment_spend_counters_failing_scope_propagates_after_siblings_
):
"""A failure in one scope must propagate to the caller (so it can invalidate
reserved counters) while every other scope still settles rather than being
- left as an orphaned background task, and the reservation is not finalized."""
- recorded: dict[str, float] = {}
+ left as an orphaned background task, and the reservation is not finalized.
+ The surviving scopes' increments are still applied in the single pipeline:
+ dropping them would under-count spend, the unsafe direction for budget
+ enforcement."""
+ warmed_keys: list[str] = []
- async def _increment(*, key, value, refresh_ttl=True):
+ async def _warm_check(*, key, **kwargs):
+ warmed_keys.append(key)
if key == "spend:team:t1":
- raise RuntimeError("redis increment failed")
- recorded[key] = recorded.get(key, 0.0) + value
- return recorded[key]
+ raise RuntimeError("redis get failed")
+ return 1.0
- fake_cache = _make_spend_counter_cache(redis_get_value=None)
- fake_cache.redis_cache.async_increment = _increment
+ async def _reseed_fails(*, counter_key, **kwargs):
+ if counter_key == "spend:team:t1":
+ raise RuntimeError("reseed failed")
+
+ applied: dict[str, float] = {}
+
+ async def _record_pipeline(increment_list, **_):
+ results = []
+ for op in increment_list:
+ applied[op["key"]] = op["increment_value"]
+ results.append(op["increment_value"])
+ return results
+
+ fake_cache = _make_spend_counter_cache()
+ fake_cache.redis_cache.async_get_cache = AsyncMock(side_effect=_warm_check)
+ fake_cache.redis_cache.async_increment_pipeline = AsyncMock(side_effect=_record_pipeline)
fake_user_cache = _make_user_api_key_cache(get_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
+ ps.SpendCounterReseed,
+ "coalesced",
+ AsyncMock(side_effect=_reseed_fails),
)
reservation = {"finalized": False}
- with pytest.raises(RuntimeError, match="redis increment failed"):
+ with pytest.raises(RuntimeError, match="reseed failed"):
await ps.increment_spend_counters(
token="hashed-tok",
team_id="t1",
@@ -741,7 +732,19 @@ async def test_increment_spend_counters_failing_scope_propagates_after_siblings_
)
assert reservation["finalized"] is False
- assert recorded == {
+ # every sibling scope settled (its warm-check ran) before the error propagated
+ assert set(warmed_keys) == {
+ "spend:key:hashed-tok",
+ "spend:team:t1",
+ "spend:team_member:u1:t1",
+ "spend:user:u1",
+ "spend:end_user:eu1",
+ "spend:tag:a",
+ "spend:org:org1",
+ }
+ # the surviving scopes' increments were still applied, in one pipeline call
+ fake_cache.redis_cache.async_increment_pipeline.assert_awaited_once()
+ assert applied == {
"spend:key:hashed-tok": 5.0,
"spend:team_member:u1:t1": 5.0,
"spend:user:u1": 5.0,
@@ -749,6 +752,7 @@ async def test_increment_spend_counters_failing_scope_propagates_after_siblings_
"spend:tag:a": 5.0,
"spend:org:org1": 5.0,
}
+ fake_cache.redis_cache.async_increment.assert_not_awaited()
@pytest.mark.asyncio
@@ -772,6 +776,108 @@ async def test_increment_spend_counters_zero_cost_is_noop_finalizes_reservation(
assert reservation == {"finalized": True}
assert fake_cache.redis_cache.async_increment.called is False
+ fake_cache.redis_cache.async_increment_pipeline.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_increment_spend_counters_pipelines_all_scopes_in_one_redis_call(
+ monkeypatch,
+):
+ """Every scope's increment must go out in a single async_increment_pipeline
+ call, not one INCRBYFLOAT round-trip per scope."""
+ counter_cache = ps.DualCache()
+ fake_redis = AsyncMock()
+ fake_redis.async_get_cache = AsyncMock(return_value=1.0) # counters warm
+
+ async def _pipeline(increment_list, **_):
+ return [1.5] * len(increment_list)
+
+ fake_redis.async_increment_pipeline = AsyncMock(side_effect=_pipeline)
+ fake_redis.async_increment = AsyncMock()
+ fake_redis.get_ttl = MagicMock(return_value=None)
+ counter_cache.redis_cache = fake_redis
+ monkeypatch.setattr(ps, "spend_counter_cache", counter_cache)
+ monkeypatch.setattr(ps, "user_api_key_cache", ps.DualCache())
+ monkeypatch.setattr(ps, "prisma_client", None)
+
+ await ps.increment_spend_counters(
+ token="hashed",
+ team_id="team-1",
+ user_id="user-1",
+ response_cost=0.5,
+ org_id="org-1",
+ end_user_id="eu-1",
+ tags=["tag-a", "tag-b"],
+ )
+
+ fake_redis.async_increment_pipeline.assert_awaited_once()
+ assert fake_redis.async_increment.await_count == 0
+ increment_list = fake_redis.async_increment_pipeline.await_args.kwargs["increment_list"]
+ expected_keys = {
+ "spend:key:hashed",
+ "spend:team:team-1",
+ "spend:team_member:user-1:team-1",
+ "spend:user:user-1",
+ "spend:end_user:eu-1",
+ "spend:tag:tag-a",
+ "spend:tag:tag-b",
+ "spend:org:org-1",
+ }
+ assert {op["key"] for op in increment_list} == expected_keys
+ assert all(op["increment_value"] == 0.5 for op in increment_list)
+ for key in expected_keys:
+ assert counter_cache.in_memory_cache.get_cache(key=key) == 1.5
+
+
+@pytest.mark.asyncio
+async def test_increment_spend_counters_pipeline_failure_invalidates_all_counters(
+ monkeypatch,
+):
+ """A failing pipeline must invalidate every pending counter so the next
+ request reseeds from the DB (which already holds this request's cost)
+ instead of trusting a value the write may have partially applied."""
+ from redis.exceptions import MaxConnectionsError
+
+ counter_cache = ps.DualCache()
+ pending_keys = (
+ "spend:key:hashed",
+ "spend:team:team-1",
+ "spend:team_member:user-1:team-1",
+ "spend:user:user-1",
+ "spend:end_user:eu-1",
+ "spend:tag:tag-a",
+ "spend:tag:tag-b",
+ "spend:org:org-1",
+ )
+ for key in pending_keys:
+ counter_cache.in_memory_cache.set_cache(key=key, value=1.0)
+ fake_redis = AsyncMock()
+ fake_redis.async_get_cache = AsyncMock(return_value=1.0) # counters warm
+ fake_redis.async_increment_pipeline = AsyncMock(side_effect=MaxConnectionsError())
+ fake_redis.async_increment = AsyncMock()
+ fake_redis.async_delete_cache = AsyncMock()
+ fake_redis.get_ttl = MagicMock(return_value=None)
+ counter_cache.redis_cache = fake_redis
+ monkeypatch.setattr(ps, "spend_counter_cache", counter_cache)
+ monkeypatch.setattr(ps, "user_api_key_cache", ps.DualCache())
+ monkeypatch.setattr(ps, "prisma_client", None)
+
+ with pytest.raises(MaxConnectionsError):
+ await ps.increment_spend_counters(
+ token="hashed",
+ team_id="team-1",
+ user_id="user-1",
+ response_cost=0.5,
+ org_id="org-1",
+ end_user_id="eu-1",
+ tags=["tag-a", "tag-b"],
+ )
+
+ assert fake_redis.async_increment.await_count == 0
+ deleted_keys = {call.kwargs["key"] for call in fake_redis.async_delete_cache.await_args_list}
+ assert deleted_keys == set(pending_keys)
+ for key in pending_keys:
+ assert counter_cache.in_memory_cache.get_cache(key=key) is None
# ---------------------------------------------------------------------------
@@ -781,9 +887,7 @@ async def test_increment_spend_counters_zero_cost_is_noop_finalizes_reservation(
@pytest.mark.asyncio
async def test_reconcile_budget_reservation_for_counter_update_returns_empty_set_when_none():
- result = await ps._reconcile_budget_reservation_for_counter_update(
- budget_reservation=None, response_cost=1.0
- )
+ result = await ps._reconcile_budget_reservation_for_counter_update(budget_reservation=None, response_cost=1.0)
assert result == set()
@@ -818,179 +922,151 @@ async def test_reconcile_budget_reservation_for_counter_update_failure_invalidat
# ---------------------------------------------------------------------------
-# _increment_end_user_and_tag_spend_counters
+# _prepare_end_user_and_tag_spend_increments
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
-async def test_increment_end_user_and_tag_spend_counters_increments_each_unique_tag(
+async def test_prepare_end_user_and_tag_spend_increments_returns_each_unique_tag(
monkeypatch,
):
- fake_cache = _make_spend_counter_cache(
- redis_get_value=None, redis_increment_value=3.0
- )
+ fake_cache = _make_spend_counter_cache(redis_get_value=1.0)
fake_user_cache = _make_user_api_key_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
- await ps._increment_end_user_and_tag_spend_counters(
+ pending = await ps._prepare_end_user_and_tag_spend_increments(
end_user_id="eu1",
tags=["a", "b", "a", "", None],
response_cost=3.0,
reserved_counter_keys=set(),
)
- observed = {
- "increment_calls": fake_cache.redis_cache.async_increment.call_count,
- "in_memory_set_calls": fake_cache.in_memory_cache.set_cache.call_count,
- "called": fake_cache.redis_cache.async_increment.called,
- }
- assert normalize(observed) == {
- "increment_calls": 3,
- "in_memory_set_calls": 3,
- "called": True,
+ assert {item.counter_key for item in pending} == {
+ "spend:end_user:eu1",
+ "spend:tag:a",
+ "spend:tag:b",
}
+ assert all(item.increment == 3.0 for item in pending)
@pytest.mark.asyncio
-async def test_increment_end_user_and_tag_spend_counters_no_end_user_no_tags_invalid_input_noop(
+async def test_prepare_end_user_and_tag_spend_increments_no_end_user_no_tags_invalid_input_noop(
monkeypatch,
):
fake_cache = _make_spend_counter_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- await ps._increment_end_user_and_tag_spend_counters(
+ pending = await ps._prepare_end_user_and_tag_spend_increments(
end_user_id=None,
tags=None,
response_cost=1.0,
reserved_counter_keys=set(),
)
+ assert pending == ()
assert fake_cache.redis_cache.async_increment.called is False
# ---------------------------------------------------------------------------
-# _increment_org_spend_counter
+# _prepare_org_spend_increment
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
-async def test_increment_org_spend_counter_increments_when_org_present(monkeypatch):
- fake_cache = _make_spend_counter_cache(
- redis_get_value=None, redis_increment_value=10.0
- )
+async def test_prepare_org_spend_increment_returns_pending_when_org_present(monkeypatch):
+ fake_cache = _make_spend_counter_cache(redis_get_value=1.0)
fake_user_cache = _make_user_api_key_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
- await ps._increment_org_spend_counter(
+ pending = await ps._prepare_org_spend_increment(
org_id="org-1",
response_cost=10.0,
reserved_counter_keys=set(),
)
- observed = {
- "increment_called": fake_cache.redis_cache.async_increment.called,
- "increment_calls": fake_cache.redis_cache.async_increment.call_count,
- "counter_key_arg": fake_cache.redis_cache.async_increment.call_args.kwargs[
- "key"
- ],
- }
- assert normalize(observed) == {
- "increment_called": True,
- "increment_calls": 1,
- "counter_key_arg": "spend:org:org-1",
- }
+ assert len(pending) == 1
+ assert pending[0].counter_key == "spend:org:org-1"
+ assert pending[0].increment == 10.0
@pytest.mark.asyncio
-async def test_increment_org_spend_counter_no_org_is_noop_invalid_id(monkeypatch):
+async def test_prepare_org_spend_increment_no_org_is_noop_invalid_id(monkeypatch):
fake_cache = _make_spend_counter_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- await ps._increment_org_spend_counter(
+ pending = await ps._prepare_org_spend_increment(
org_id=None,
response_cost=1.0,
reserved_counter_keys=set(),
)
+ assert pending == ()
assert fake_cache.redis_cache.async_increment.called is False
# ---------------------------------------------------------------------------
-# _init_and_increment_unreserved_spend_counter
+# _prepare_unreserved_spend_counter_increment
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
-async def test_init_and_increment_unreserved_spend_counter_skips_reserved_keys(
+async def test_prepare_unreserved_spend_counter_increment_skips_reserved_keys(
monkeypatch,
):
fake_cache = _make_spend_counter_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- await ps._init_and_increment_unreserved_spend_counter(
+ pending = await ps._prepare_unreserved_spend_counter_increment(
counter_key="spend:tag:x",
source_cache_key="tag:x",
increment=1.0,
reserved_counter_keys={"spend:tag:x"},
)
+ assert pending is None
assert fake_cache.redis_cache.async_increment.called is False
@pytest.mark.asyncio
-async def test_init_and_increment_unreserved_spend_counter_proceeds_when_not_reserved(
+async def test_prepare_unreserved_spend_counter_increment_proceeds_when_not_reserved(
monkeypatch,
):
- fake_cache = _make_spend_counter_cache(
- redis_get_value=None, redis_increment_value=2.0
- )
+ fake_cache = _make_spend_counter_cache(redis_get_value=None)
fake_user_cache = _make_user_api_key_cache()
+ reseed = AsyncMock(return_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", reseed)
- await ps._init_and_increment_unreserved_spend_counter(
+ pending = await ps._prepare_unreserved_spend_counter_increment(
counter_key="spend:tag:y",
source_cache_key="tag:y",
increment=2.0,
reserved_counter_keys=set(),
)
- observed = {
- "increment_called": fake_cache.redis_cache.async_increment.called,
- "redis_get_called": fake_cache.redis_cache.async_get_cache.called,
- "reseed_consulted": True,
- }
- assert observed == {
- "increment_called": True,
- "redis_get_called": True,
- "reseed_consulted": True,
- }
+ assert pending is not None
+ assert pending.counter_key == "spend:tag:y"
+ assert pending.increment == 2.0
+ assert fake_cache.redis_cache.async_get_cache.called is True
+ assert reseed.called is True
# ---------------------------------------------------------------------------
-# _init_and_increment_spend_counter
+# _prepare_spend_counter_increment
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
-async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypatch):
- fake_cache = _make_spend_counter_cache(
- redis_get_value=11.0, redis_increment_value=14.0
- )
+async def test_prepare_spend_counter_increment_warm_cache_skips_reseed(monkeypatch):
+ fake_cache = _make_spend_counter_cache(redis_get_value=11.0)
fake_user_cache = _make_user_api_key_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
@@ -998,12 +1074,14 @@ async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypa
reseed = AsyncMock(return_value=None)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", reseed)
- await ps._init_and_increment_spend_counter(
+ pending = await ps._prepare_spend_counter_increment(
counter_key="spend:key:k",
source_cache_key="k",
increment=3.0,
)
+ assert pending.counter_key == "spend:key:k"
+ assert pending.increment == 3.0
observed = {
"reseed_called": reseed.called,
"increment_called": fake_cache.redis_cache.async_increment.called,
@@ -1011,23 +1089,21 @@ async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypa
}
assert normalize(observed) == {
"reseed_called": False,
- "increment_called": True,
+ "increment_called": False,
"in_memory_seeded_from_redis": True,
}
# ---------------------------------------------------------------------------
-# _init_and_increment_window_spend_counter
+# _prepare_window_spend_counter_increment
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
-async def test_init_and_increment_window_spend_counter_increments_when_initialized(
+async def test_prepare_window_spend_counter_increment_returns_pending_when_initialized(
monkeypatch,
):
- fake_cache = _make_spend_counter_cache(
- redis_get_value=0.0, redis_increment_value=5.0
- )
+ fake_cache = _make_spend_counter_cache(redis_get_value=0.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
@@ -1036,7 +1112,7 @@ async def test_init_and_increment_window_spend_counter_increments_when_initializ
AsyncMock(return_value=0.0),
)
- await ps._init_and_increment_window_spend_counter(
+ pending = await ps._prepare_window_spend_counter_increment(
counter_key="spend:key:k:window:1d",
entity_type="Key",
entity_id="k",
@@ -1045,26 +1121,19 @@ async def test_init_and_increment_window_spend_counter_increments_when_initializ
increment=5.0,
)
- observed = {
- "redis_increment_called": fake_cache.redis_cache.async_increment.called,
- "increment_calls": fake_cache.redis_cache.async_increment.call_count,
- "in_memory_set_calls": fake_cache.in_memory_cache.set_cache.call_count,
- }
- assert normalize(observed) == {
- "redis_increment_called": True,
- "increment_calls": 1,
- "in_memory_set_calls": 2,
- }
+ assert pending is not None
+ assert pending.counter_key == "spend:key:k:window:1d"
+ assert pending.increment == 5.0
@pytest.mark.asyncio
-async def test_init_and_increment_window_spend_counter_missing_window_start_invalid_skips(
+async def test_prepare_window_spend_counter_increment_missing_window_start_invalid_skips(
monkeypatch,
):
fake_cache = _make_spend_counter_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- await ps._init_and_increment_window_spend_counter(
+ pending = await ps._prepare_window_spend_counter_increment(
counter_key="spend:key:k:window:1d",
entity_type="Key",
entity_id="k",
@@ -1073,6 +1142,7 @@ async def test_init_and_increment_window_spend_counter_missing_window_start_inva
increment=5.0,
)
+ assert pending is None
assert fake_cache.redis_cache.async_increment.called is False
@@ -1114,16 +1184,12 @@ async def test_ensure_spend_counter_initialized_warm_skips_reseed_and_source(
async def test_ensure_spend_counter_initialized_cold_seeds_from_source_cache(
monkeypatch,
):
- fake_cache = _make_spend_counter_cache(
- redis_get_value=None, redis_increment_value=7.0
- )
+ fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=7.0)
fake_user_cache = _make_user_api_key_cache(get_value={"spend": 7.0})
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
- monkeypatch.setattr(
- ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
- )
+ monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
await ps._ensure_spend_counter_initialized(
counter_key="spend:user:u",
@@ -1163,9 +1229,7 @@ async def test_get_source_cache_base_spend_reads_first_hit_from_list(monkeypatch
fake_user_cache.async_get_cache = AsyncMock(side_effect=_get)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
- result = await ps._get_source_cache_base_spend(
- source_cache_key=["miss", "hit-obj", "miss2"]
- )
+ result = await ps._get_source_cache_base_spend(source_cache_key=["miss", "hit-obj", "miss2"])
observed = {
"result": result,
@@ -1294,9 +1358,7 @@ async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkey
fake_cache = _make_spend_counter_cache(redis_increment_value=44.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
- result = await ps._increment_spend_counter_cache(
- counter_key="spend:key:k", increment=4.0
- )
+ result = await ps._increment_spend_counter_cache(counter_key="spend:key:k", increment=4.0)
observed = {
"result": result,
@@ -1314,15 +1376,11 @@ async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkey
async def test_increment_spend_counter_cache_redis_error_raises_and_invalidates(
monkeypatch,
):
- fake_cache = _make_spend_counter_cache(
- redis_increment_side_effect=RuntimeError("incr fail")
- )
+ fake_cache = _make_spend_counter_cache(redis_increment_side_effect=RuntimeError("incr fail"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
with pytest.raises(RuntimeError):
- await ps._increment_spend_counter_cache(
- counter_key="spend:key:k", increment=1.0
- )
+ await ps._increment_spend_counter_cache(counter_key="spend:key:k", increment=1.0)
assert fake_cache.in_memory_cache.delete_cache.called is True
assert fake_cache.redis_cache.async_delete_cache.called is True
@@ -1343,9 +1401,7 @@ async def test_invalidate_spend_counter_deletes_in_memory_and_redis(monkeypatch)
observed = {
"in_memory_delete_called": fake_cache.in_memory_cache.delete_cache.called,
"redis_delete_called": fake_cache.redis_cache.async_delete_cache.called,
- "delete_args_key": fake_cache.redis_cache.async_delete_cache.call_args.kwargs[
- "key"
- ],
+ "delete_args_key": fake_cache.redis_cache.async_delete_cache.call_args.kwargs["key"],
}
assert normalize(observed) == {
"in_memory_delete_called": True,
@@ -1357,9 +1413,7 @@ async def test_invalidate_spend_counter_deletes_in_memory_and_redis(monkeypatch)
@pytest.mark.asyncio
async def test_invalidate_spend_counter_swallows_redis_failure_no_raise(monkeypatch):
fake_cache = _make_spend_counter_cache()
- fake_cache.redis_cache.async_delete_cache = AsyncMock(
- side_effect=RuntimeError("redis down")
- )
+ fake_cache.redis_cache.async_delete_cache = AsyncMock(side_effect=RuntimeError("redis down"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
await ps._invalidate_spend_counter(counter_key="spend:key:k")
diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py
index c123eeeed36..6165af4920d 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py
@@ -41,6 +41,15 @@ class _FlakyRedisCache:
self._store[key] = float(value)
return True
+ 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
+
+ def get_ttl(self, **kwargs):
+ return None
+
@pytest.mark.asyncio
async def test_direct_increment_runs_when_reservation_reconcile_hits_redis_failure(
diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py
index 6dab054d8ea..40ebc03781c 100644
--- a/tests/test_litellm/proxy/test_budget_reservation.py
+++ b/tests/test_litellm/proxy/test_budget_reservation.py
@@ -2220,6 +2220,15 @@ class _ExpiringRedisCache:
async def async_delete_cache(self, key: str, *args: object, **kwargs: object) -> None:
self.store.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
+
+ def get_ttl(self, **kwargs) -> None:
+ return None
+
@pytest.mark.asyncio
async def test_reconcile_after_redis_counter_expiry_keeps_request_cost_enforced(
diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py
index 54579e6cb7c..b0f38978727 100644
--- a/tests/test_litellm/proxy/test_proxy_server.py
+++ b/tests/test_litellm/proxy/test_proxy_server.py
@@ -7749,7 +7749,7 @@ async def test_increment_spend_counters_team_and_member():
@pytest.mark.asyncio
-async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss():
+async def test_prepare_spend_counter_increment_reseeds_from_db_on_counter_miss():
"""When the Redis counter is missing, the reseed path reads the
authoritative spend from the DB (not a stale cache), so the next
increment continues from the correct base value."""
@@ -7762,8 +7762,17 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
recorded_increments.append({"key": key, "value": value, "ttl": ttl})
return value
+ async def record_pipeline(increment_list, **kwargs):
+ results = []
+ for op in increment_list:
+ await record_increment(key=op["key"], value=op["increment_value"], ttl=op["ttl"])
+ results.append(op["increment_value"])
+ return results
+
fake_redis = AsyncMock()
fake_redis.async_increment = AsyncMock(side_effect=record_increment)
+ fake_redis.async_increment_pipeline = AsyncMock(side_effect=record_pipeline)
+ fake_redis.get_ttl = MagicMock(return_value=None)
fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing
fake_redis.async_set_cache = AsyncMock(return_value=True) # SET NX wins
counter_cache.redis_cache = fake_redis
@@ -7782,7 +7791,10 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
stale_cache.in_memory_cache.set_cache(key="team_id:team-9", value=stale_team)
import litellm.proxy.proxy_server as ps
- from litellm.proxy.proxy_server import _init_and_increment_spend_counter
+ from litellm.proxy.proxy_server import (
+ _apply_spend_counter_increments,
+ _prepare_spend_counter_increment,
+ )
orig_user, orig_counter, orig_prisma = (
ps.user_api_key_cache,
@@ -7793,11 +7805,12 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
ps.spend_counter_cache = counter_cache
ps.prisma_client = fake_prisma
try:
- await _init_and_increment_spend_counter(
+ pending = await _prepare_spend_counter_increment(
counter_key="spend:team:team-9",
source_cache_key="team_id:team-9",
increment=1.5,
)
+ await _apply_spend_counter_increments(pending=(pending,))
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-9"})
# Seed uses SET NX with db_spend (42) — cross-pod safe, no INCR of 42.
@@ -7976,7 +7989,10 @@ async def test_reseed_spend_from_db_skips_window_variant_keys():
@pytest.mark.asyncio
async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
from litellm.caching.dual_cache import DualCache
- from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
+ from litellm.proxy.proxy_server import (
+ _apply_spend_counter_increments,
+ _prepare_window_spend_counter_increment,
+ )
counter_cache = DualCache()
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
@@ -7992,7 +8008,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
ps.spend_counter_cache = counter_cache
ps.prisma_client = fake_prisma
try:
- await _init_and_increment_window_spend_counter(
+ pending = await _prepare_window_spend_counter_increment(
counter_key="spend:key:key-window:window:1h",
entity_type="Key",
entity_id="key-window",
@@ -8000,6 +8016,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
window_start=window_start,
increment=0.5,
)
+ await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
by=["api_key"],
@@ -8015,7 +8032,10 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
@pytest.mark.asyncio
async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
from litellm.caching.dual_cache import DualCache
- from litellm.proxy.proxy_server import _init_and_increment_spend_counter
+ from litellm.proxy.proxy_server import (
+ _apply_spend_counter_increments,
+ _prepare_spend_counter_increment,
+ )
counter_cache = DualCache()
counter_key = "spend:team:team-stale-local"
@@ -8037,6 +8057,15 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
fake_redis.async_get_cache = AsyncMock(return_value=None)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
+
+ async def redis_increment_pipeline(increment_list, **_):
+ results = []
+ for op in increment_list:
+ results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
+ return results
+
+ fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
+ fake_redis.get_ttl = MagicMock(return_value=None)
counter_cache.redis_cache = fake_redis
db_row = MagicMock()
@@ -8055,11 +8084,12 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
ps.prisma_client = fake_prisma
ps.user_api_key_cache = DualCache()
try:
- await _init_and_increment_spend_counter(
+ pending = await _prepare_spend_counter_increment(
counter_key=counter_key,
source_cache_key="team_id:team-stale-local",
increment=1.5,
)
+ await _apply_spend_counter_increments(pending=(pending,))
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-stale-local"})
# Seed via SET NX (42) + delta via INCRBYFLOAT (1.5) = 43.5.
@@ -8074,7 +8104,10 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
@pytest.mark.asyncio
async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
from litellm.caching.dual_cache import DualCache
- from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
+ from litellm.proxy.proxy_server import (
+ _apply_spend_counter_increments,
+ _prepare_window_spend_counter_increment,
+ )
counter_cache = DualCache()
counter_key = "spend:key:key-window-stale-local:window:1h"
@@ -8097,6 +8130,15 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
fake_redis.async_get_cache = AsyncMock(return_value=None)
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
+
+ async def redis_increment_pipeline(increment_list, **_):
+ results = []
+ for op in increment_list:
+ results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
+ return results
+
+ fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
+ fake_redis.get_ttl = MagicMock(return_value=None)
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
@@ -8111,7 +8153,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
ps.spend_counter_cache = counter_cache
ps.prisma_client = fake_prisma
try:
- await _init_and_increment_window_spend_counter(
+ pending = await _prepare_window_spend_counter_increment(
counter_key=counter_key,
entity_type="Key",
entity_id="key-window-stale-local",
@@ -8119,6 +8161,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
window_start=window_start,
increment=0.5,
)
+ await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
by=["api_key"],
@@ -8138,7 +8181,10 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
@pytest.mark.asyncio
async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed():
from litellm.caching.dual_cache import DualCache
- from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
+ from litellm.proxy.proxy_server import (
+ _apply_spend_counter_increments,
+ _prepare_window_spend_counter_increment,
+ )
counter_cache = DualCache()
counter_key = "spend:key:key-window-concurrent-seed:window:1h"
@@ -8161,6 +8207,15 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache)
fake_redis.async_set_cache = AsyncMock(return_value=False)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
+
+ async def redis_increment_pipeline(increment_list, **_):
+ results = []
+ for op in increment_list:
+ results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
+ return results
+
+ fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
+ fake_redis.get_ttl = MagicMock(return_value=None)
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
@@ -8175,7 +8230,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
ps.spend_counter_cache = counter_cache
ps.prisma_client = fake_prisma
try:
- await _init_and_increment_window_spend_counter(
+ pending = await _prepare_window_spend_counter_increment(
counter_key=counter_key,
entity_type="Key",
entity_id="key-window-concurrent-seed",
@@ -8183,6 +8238,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
window_start=window_start,
increment=0.5,
)
+ await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
fake_redis.async_set_cache.assert_awaited_once_with(
key=counter_key,
@@ -8199,7 +8255,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
@pytest.mark.asyncio
async def test_window_spend_counter_skips_invalid_window_start():
from litellm.caching.dual_cache import DualCache
- from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
+ from litellm.proxy.proxy_server import _prepare_window_spend_counter_increment
counter_cache = DualCache()
@@ -8208,7 +8264,7 @@ async def test_window_spend_counter_skips_invalid_window_start():
orig_counter = ps.spend_counter_cache
ps.spend_counter_cache = counter_cache
try:
- await _init_and_increment_window_spend_counter(
+ pending = await _prepare_window_spend_counter_increment(
counter_key="spend:key:key-invalid-window:window:not-a-duration",
entity_type="Key",
entity_id="key-invalid-window",
@@ -8216,6 +8272,7 @@ async def test_window_spend_counter_skips_invalid_window_start():
window_start=None,
increment=0.5,
)
+ assert pending is None
assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-invalid-window:window:not-a-duration") is None
finally:
@@ -8279,6 +8336,9 @@ async def test_increment_spend_counters_finalizes_after_unreserved_increments():
async def assert_reservation_not_finalized_yet(**kwargs):
assert budget_reservation["finalized"] is False
incremented_counters.append(kwargs["counter_key"])
+ return ps._PendingSpendIncrement(
+ counter_key=kwargs["counter_key"], increment=kwargs["increment"]
+ )
import litellm.proxy.proxy_server as ps
@@ -8287,7 +8347,7 @@ async def test_increment_spend_counters_finalizes_after_unreserved_increments():
ps.user_api_key_cache = DualCache()
try:
with patch(
- "litellm.proxy.proxy_server._init_and_increment_spend_counter",
+ "litellm.proxy.proxy_server._prepare_spend_counter_increment",
new=AsyncMock(side_effect=assert_reservation_not_finalized_yet),
):
await increment_spend_counters(
@@ -8620,7 +8680,7 @@ async def test_get_current_spend_uses_db_zero_over_stale_fallback():
async def test_concurrent_read_and_write_paths_share_one_db_query():
"""
The read path (`get_current_spend`) and the write path
- (`_init_and_increment_spend_counter`) both reseed cold counters from
+ (`_prepare_spend_counter_increment`) both reseed cold counters from
the DB. They must share the per-counter lock so a concurrent pre-call
enforcement read and post-call increment for the same counter collapse
to one DB query, not two.
@@ -8629,7 +8689,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import (
- _init_and_increment_spend_counter,
+ _prepare_spend_counter_increment,
get_current_spend,
)
@@ -8683,7 +8743,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query():
try:
results = await _asyncio.gather(
get_current_spend(counter_key=counter_key, fallback_spend=0.0),
- _init_and_increment_spend_counter(
+ _prepare_spend_counter_increment(
counter_key=counter_key,
source_cache_key="ignored",
increment=1.5,
diff --git a/tests/test_litellm/router_strategy/test_lowest_latency.py b/tests/test_litellm/router_strategy/test_lowest_latency.py
index 812d7bbff32..1a8614e3fca 100644
--- a/tests/test_litellm/router_strategy/test_lowest_latency.py
+++ b/tests/test_litellm/router_strategy/test_lowest_latency.py
@@ -8,11 +8,12 @@ import json
from datetime import datetime, timedelta
import pytest
-
+from pydantic import ValidationError
import litellm
from litellm.caching.caching import DualCache
-from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler
+from litellm.router import Router
+from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler, RoutingArgs
DEPLOYMENT_ID = "9876"
KWARGS = {
@@ -58,9 +59,9 @@ def test_sync_embedding_latency_is_json_serializable():
latencies = _recorded_latencies(cache)
assert latencies, "expected a latency entry to be recorded"
- assert all(
- not isinstance(value, timedelta) for value in latencies
- ), f"raw timedelta leaked into latency list: {latencies}"
+ assert all(not isinstance(value, timedelta) for value in latencies), (
+ f"raw timedelta leaked into latency list: {latencies}"
+ )
assert latencies[-1] == pytest.approx(2.0)
# the exact failure mode from production: redis cache sync json.dumps
json.dumps({"latency": latencies})
@@ -84,9 +85,9 @@ async def test_async_embedding_latency_is_json_serializable():
latencies = _recorded_latencies(cache)
assert latencies, "expected a latency entry to be recorded"
- assert all(
- not isinstance(value, timedelta) for value in latencies
- ), f"raw timedelta leaked into latency list: {latencies}"
+ assert all(not isinstance(value, timedelta) for value in latencies), (
+ f"raw timedelta leaked into latency list: {latencies}"
+ )
assert latencies[-1] == pytest.approx(3.0)
json.dumps({"latency": latencies})
@@ -292,6 +293,85 @@ async def test_streaming_routing_ignores_per_token_ttft_samples_from_older_worke
assert picked["model_info"]["id"] == FAST_TTFT_ID
+@pytest.mark.asyncio
+@pytest.mark.parametrize("sync_mode", [True, False], ids=["sync", "async"])
+@pytest.mark.parametrize(
+ ("ttft_percentile", "first_samples", "second_samples", "expected_id"),
+ [
+ (None, [0.1, 0.1, 1.0], [0.3, 0.3, 0.3], SLOW_TTFT_ID),
+ (0.5, [0.1, 0.1, 1.0], [0.3, 0.3, 0.3], FAST_TTFT_ID),
+ (0.9, [0.1, 0.1, 0.1, 0.1, 1.5], [0.3, 0.3, 0.3, 0.3, 0.3], SLOW_TTFT_ID),
+ ],
+ ids=["default_average", "p50", "p90"],
+)
+async def test_streaming_ttft_ranking_percentile(
+ sync_mode: bool,
+ ttft_percentile: float | None,
+ first_samples: list[float],
+ second_samples: list[float],
+ expected_id: str,
+):
+ cache = DualCache()
+ routing_args = {} if ttft_percentile is None else {"ttft_percentile": ttft_percentile}
+ handler = LowestLatencyLoggingHandler(router_cache=cache, routing_args=routing_args)
+ cache.set_cache(
+ key=f"{MODEL_GROUP}_map",
+ value={
+ FAST_TTFT_ID: {"time_to_first_token_seconds": first_samples},
+ SLOW_TTFT_ID: {"time_to_first_token_seconds": second_samples},
+ },
+ )
+
+ if sync_mode:
+ picked = handler.get_available_deployments(
+ model_group=MODEL_GROUP,
+ healthy_deployments=STREAMING_DEPLOYMENTS,
+ request_kwargs={"stream": True, "metadata": {}},
+ )
+ else:
+ picked = await handler.async_get_available_deployments(
+ model_group=MODEL_GROUP,
+ healthy_deployments=STREAMING_DEPLOYMENTS,
+ request_kwargs={"stream": True, "metadata": {}},
+ )
+
+ assert picked is not None
+ assert picked["model_info"]["id"] == expected_id
+
+
+@pytest.mark.parametrize("ttft_percentile", [0, -0.1, 1.1])
+def test_ttft_percentile_validation(ttft_percentile: float):
+ with pytest.raises(ValidationError):
+ RoutingArgs(ttft_percentile=ttft_percentile)
+
+
+@pytest.mark.parametrize("ttft_percentile", [0.5, 0.9, 0.95, 1.0])
+def test_ttft_percentile_accepts_valid_values(ttft_percentile: float):
+ assert RoutingArgs(ttft_percentile=ttft_percentile).ttft_percentile == ttft_percentile
+
+
+@pytest.mark.asyncio
+async def test_ttft_percentile_does_not_change_non_streaming_routing():
+ cache = DualCache()
+ handler = LowestLatencyLoggingHandler(router_cache=cache, routing_args={"ttft_percentile": 0.9})
+ cache.set_cache(
+ key=f"{MODEL_GROUP}_map",
+ value={
+ FAST_TTFT_ID: {"latency": [1.0], "time_to_first_token_seconds": [0.1]},
+ SLOW_TTFT_ID: {"latency": [0.2], "time_to_first_token_seconds": [1.5]},
+ },
+ )
+
+ picked = await handler.async_get_available_deployments(
+ model_group=MODEL_GROUP,
+ healthy_deployments=STREAMING_DEPLOYMENTS,
+ request_kwargs={"stream": False, "metadata": {}},
+ )
+
+ assert picked is not None
+ assert picked["model_info"]["id"] == SLOW_TTFT_ID
+
+
@pytest.mark.asyncio
@pytest.mark.parametrize(
"cached_entry",
@@ -318,3 +398,81 @@ async def test_async_get_available_deployments_treats_missing_samples_as_zero_la
assert picked is not None
assert picked["model_info"]["id"] == DEPLOYMENT_ID
+
+
+def _latency_router(routing_strategy_args: dict) -> Router:
+ return Router(
+ model_list=[
+ {
+ "model_name": MODEL_GROUP,
+ "litellm_params": {"model": f"openai/{MODEL_GROUP}", "api_key": "sk-fake"},
+ "model_info": {"id": deployment_id},
+ }
+ for deployment_id in (FAST_TTFT_ID, SLOW_TTFT_ID)
+ ],
+ routing_strategy="latency-based-routing",
+ routing_strategy_args=routing_strategy_args,
+ )
+
+
+def _seed_streaming_ttft(router: Router) -> None:
+ router.cache.set_cache(
+ key=f"{MODEL_GROUP}_map",
+ value={
+ FAST_TTFT_ID: {"time_to_first_token_seconds": [0.1, 0.1, 1.0]},
+ SLOW_TTFT_ID: {"time_to_first_token_seconds": [0.3, 0.3, 0.3]},
+ },
+ )
+
+
+async def _pick_streaming(router: Router) -> str:
+ picked = await router.async_get_available_deployment(
+ model=MODEL_GROUP,
+ request_kwargs={"stream": True, "metadata": {}},
+ )
+ return picked["model_info"]["id"]
+
+
+@pytest.mark.asyncio
+async def test_runtime_routing_strategy_args_update_applies_ttft_percentile():
+ """A config reload that adds ttft_percentile must reach the live selector,
+ not sit unused until the proxy restarts."""
+ router = _latency_router({"max_latency_list_size": 50})
+ _seed_streaming_ttft(router)
+
+ assert await _pick_streaming(router) == SLOW_TTFT_ID
+
+ router.update_settings(routing_strategy_args={"max_latency_list_size": 50, "ttft_percentile": 0.5})
+
+ assert await _pick_streaming(router) == FAST_TTFT_ID
+
+
+@pytest.mark.asyncio
+async def test_runtime_routing_strategy_args_update_keeps_previous_args_when_invalid():
+ router = _latency_router({"ttft_percentile": 0.5})
+ _seed_streaming_ttft(router)
+
+ router.update_settings(routing_strategy_args={"ttft_percentile": 5})
+
+ assert await _pick_streaming(router) == FAST_TTFT_ID
+
+
+@pytest.mark.asyncio
+async def test_runtime_routing_strategy_args_update_is_a_noop_without_a_selector():
+ """simple-shuffle has no selector to re-link, so an args update must leave
+ the router alone instead of blowing up on a missing selector attribute."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": MODEL_GROUP,
+ "litellm_params": {"model": f"openai/{MODEL_GROUP}", "api_key": "sk-fake"},
+ "model_info": {"id": FAST_TTFT_ID},
+ }
+ ],
+ routing_strategy="simple-shuffle",
+ )
+
+ router.update_settings(routing_strategy_args={"ttl": 5})
+
+ assert router.routing_strategy_args == {"ttl": 5}
+ assert await _pick_streaming(router) == FAST_TTFT_ID
diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py
index dac991a41c4..3961c8d74c2 100644
--- a/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py
+++ b/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py
@@ -16,12 +16,10 @@ The mechanism works without any cache and supports two encoding strategies:
"""
import time
-from typing import List, Optional
from unittest.mock import AsyncMock, patch
import pytest
-
import litellm
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIResponse
@@ -68,9 +66,7 @@ class TestEncryptedItemIdCodec:
def test_roundtrip(self):
model_id = "deployment-1"
original_item_id = "rs_abc123def456"
- encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(
- model_id, original_item_id
- )
+ encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id)
assert encoded.startswith("encitem_")
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded)
assert decoded is not None
@@ -81,9 +77,7 @@ class TestEncryptedItemIdCodec:
"""Decoding must succeed even if base64 padding (=) was stripped in transit."""
model_id = "gpt-5.1-codex-openai-2"
original_item_id = "rs_0efb96cb222403210069a01d5d52588196a9dc394ffdb89d00"
- encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(
- model_id, original_item_id
- )
+ encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id)
# Strip any trailing '=' to simulate what happens in transit
stripped = encoded.rstrip("=")
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(stripped)
@@ -100,9 +94,7 @@ class TestEncryptedItemIdCodec:
"""item_id values containing ';' must survive the roundtrip."""
model_id = "deployment-1"
original_item_id = "rs_part1;part2;part3"
- encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(
- model_id, original_item_id
- )
+ encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id)
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded)
assert decoded is not None
assert decoded["item_id"] == original_item_id
@@ -118,11 +110,7 @@ class TestUpdateEncryptedContentItemIds:
{"id": "rs_xyz", "type": "reasoning", "encrypted_content": "secret"},
],
}
- result = (
- ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(
- response, model_id
- )
- )
+ result = ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(response, model_id)
# Plain message item untouched
assert result["output"][0]["id"] == "msg_abc"
# Reasoning item with encrypted_content gets encoded
@@ -133,16 +121,8 @@ class TestUpdateEncryptedContentItemIds:
assert decoded["item_id"] == "rs_xyz"
def test_no_op_when_model_id_is_none(self):
- response = {
- "output": [
- {"id": "rs_xyz", "type": "reasoning", "encrypted_content": "secret"}
- ]
- }
- result = (
- ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(
- response, None
- )
- )
+ response = {"output": [{"id": "rs_xyz", "type": "reasoning", "encrypted_content": "secret"}]}
+ result = ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(response, None)
assert result["output"][0]["id"] == "rs_xyz"
@@ -151,9 +131,7 @@ class TestEncryptedContentWrapping:
"""Test wrapping encrypted_content with model_id metadata."""
model_id = "deployment-1"
original_content = "gAAAAABpnW_yEYmSNEyOG_original_encrypted_data"
- wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
- original_content, model_id
- )
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id)
assert wrapped.startswith("litellm_enc:")
assert wrapped != original_content
@@ -170,9 +148,7 @@ class TestEncryptedContentWrapping:
(
model_id,
content,
- ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(
- plain_content
- )
+ ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(plain_content)
assert model_id is None
assert content == plain_content
@@ -189,11 +165,7 @@ class TestEncryptedContentWrapping:
},
],
}
- result = (
- ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(
- response, model_id
- )
- )
+ result = ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(response, model_id)
assert result["output"][0].get("encrypted_content") is None
wrapped = result["output"][1]["encrypted_content"]
assert wrapped.startswith("litellm_enc:")
@@ -210,19 +182,13 @@ class TestRestoreEncryptedContentItemIds:
def test_restores_encoded_ids(self):
model_id = "deployment-1"
original_id = "rs_encrypted_item_456"
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- model_id, original_id
- )
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_id)
request_input = [
{"type": "message", "id": "msg_abc123", "role": "assistant"},
{"type": "reasoning", "id": encoded_id, "encrypted_content": "secret"},
]
- restored = (
- ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
- request_input
- )
- )
+ restored = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input)
assert restored[0]["id"] == "msg_abc123"
assert restored[1]["id"] == original_id
@@ -230,33 +196,21 @@ class TestRestoreEncryptedContentItemIds:
"""Test that wrapped encrypted_content is unwrapped before forwarding."""
model_id = "deployment-1"
original_content = "gAAAAABpnW_yEYmSNEyOG_original"
- wrapped_content = (
- ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
- original_content, model_id
- )
- )
+ wrapped_content = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id)
request_input = [
{"type": "reasoning", "encrypted_content": wrapped_content},
]
- restored = (
- ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
- request_input
- )
- )
+ restored = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input)
assert restored[0]["encrypted_content"] == original_content
def test_no_op_for_plain_string_input(self):
- result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
- "Hello world"
- )
+ result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input("Hello world")
assert result == "Hello world"
def test_no_op_for_unencoded_ids(self):
request_input = [{"type": "message", "id": "msg_plain"}]
- result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
- request_input
- )
+ result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input)
assert result[0]["id"] == "msg_plain"
@@ -283,9 +237,7 @@ async def test_encrypted_content_affinity_tracks_and_routes():
"id": "msg_abc123",
"status": "completed",
"role": "assistant",
- "content": [
- {"type": "output_text", "text": "Hello!", "annotations": []}
- ],
+ "content": [{"type": "output_text", "text": "Hello!", "annotations": []}],
},
{
"type": "reasoning",
@@ -347,9 +299,9 @@ async def test_encrypted_content_affinity_tracks_and_routes():
# The response must have rewritten the encrypted item's ID to encoded form
encoded_item_id = _extract_encoded_item_id(first_response)
- assert encoded_item_id.startswith(
- "encitem_"
- ), f"Expected output item ID to be rewritten to encitem_... but got {encoded_item_id!r}"
+ assert encoded_item_id.startswith("encitem_"), (
+ f"Expected output item ID to be rewritten to encitem_... but got {encoded_item_id!r}"
+ )
# Verify the encoded ID decodes back to the correct deployment + original ID
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded_item_id)
@@ -371,9 +323,9 @@ async def test_encrypted_content_affinity_tracks_and_routes():
)
second_model_id = second_response._hidden_params["model_id"]
- assert (
- second_model_id == first_model_id
- ), f"Expected affinity to route to {first_model_id}, but got {second_model_id}"
+ assert second_model_id == first_model_id, (
+ f"Expected affinity to route to {first_model_id}, but got {second_model_id}"
+ )
@pytest.mark.asyncio
@@ -478,9 +430,7 @@ async def test_encrypted_content_affinity_bypasses_rpm_limits():
# Extract encoded item ID from the first response output
encoded_item_id = _extract_encoded_item_id(first_response)
- assert encoded_item_id.startswith(
- "encitem_"
- ), f"Expected encitem_... but got {encoded_item_id!r}"
+ assert encoded_item_id.startswith("encitem_"), f"Expected encitem_... but got {encoded_item_id!r}"
# Follow-up with the encoded item ID — should pin to same deployment
second_response = await router.aresponses(
@@ -628,17 +578,13 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
if hasattr(first_item, "encrypted_content")
else first_item.get("encrypted_content")
)
- assert wrapped_content.startswith(
- "litellm_enc:"
- ), f"Expected wrapped content but got {wrapped_content[:50]}..."
+ assert wrapped_content.startswith("litellm_enc:"), f"Expected wrapped content but got {wrapped_content[:50]}..."
# Verify we can extract model_id from wrapped content
(
extracted_model_id,
_,
- ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(
- wrapped_content
- )
+ ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(wrapped_content)
assert extracted_model_id == first_model_id
# Second request: use wrapped encrypted_content WITHOUT an ID (Codex behavior)
@@ -653,9 +599,9 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
)
second_model_id = second_response._hidden_params["model_id"]
- assert (
- second_model_id == first_model_id
- ), f"Expected affinity to route to {first_model_id}, but got {second_model_id}"
+ assert second_model_id == first_model_id, (
+ f"Expected affinity to route to {first_model_id}, but got {second_model_id}"
+ )
def test_encrypted_content_wrapping_preserves_original_content():
@@ -664,13 +610,9 @@ def test_encrypted_content_wrapping_preserves_original_content():
This is critical for streaming responses where content must round-trip correctly.
"""
model_id = "test-deployment-1"
- original_encrypted_content = (
- "gAAAAABpnW_yEYmSNEyOG_streaming_test_content_with_special_chars==+/"
- )
+ original_encrypted_content = "gAAAAABpnW_yEYmSNEyOG_streaming_test_content_with_special_chars==+/"
- wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
- original_encrypted_content, model_id
- )
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_encrypted_content, model_id)
assert wrapped.startswith("litellm_enc:")
assert wrapped != original_encrypted_content
@@ -691,9 +633,7 @@ def test_encrypted_content_wrapping_with_multiple_semicolons():
model_id = "deployment-with-semicolons"
original_content = "gAAAAAB;some;content;with;semicolons"
- wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
- original_content, model_id
- )
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id)
(
extracted_model_id,
@@ -764,9 +704,7 @@ async def test_encrypted_content_affinity_preserves_litellm_metadata_for_respons
request_kwargs=request_kwargs,
)
- assert (
- request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] is True
- )
+ assert request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] is True
assert request_kwargs["litellm_metadata"]["model_info"] == {"id": "dep-1"}
@@ -777,9 +715,7 @@ def test_encrypted_content_wrapping_empty_string():
model_id = "test-deployment"
original_content = ""
- wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
- original_content, model_id
- )
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id)
assert wrapped.startswith("litellm_enc:")
@@ -1132,9 +1068,7 @@ def test_boundary_key_accepts_pydantic_litellm_params_instance():
"api_key": "fake-azure-resource-key-a",
}
- pydantic_key = EncryptedContentAffinityCheck._encryption_boundary_key(
- pydantic_params
- )
+ pydantic_key = EncryptedContentAffinityCheck._encryption_boundary_key(pydantic_params)
plain_key = EncryptedContentAffinityCheck._encryption_boundary_key(plain_params)
assert pydantic_key is not None
@@ -1161,18 +1095,8 @@ def test_boundary_key_rejects_non_dict_like_inputs():
for bad in (None, [], "not a dict", 42, object()):
assert EncryptedContentAffinityCheck._encryption_boundary_key(bad) is None
- assert (
- EncryptedContentAffinityCheck._encryption_boundary_key(
- {"api_base": "", "api_key": "k"}
- )
- is None
- )
- assert (
- EncryptedContentAffinityCheck._encryption_boundary_key(
- {"api_base": "https://x"}
- )
- is None
- )
+ assert EncryptedContentAffinityCheck._encryption_boundary_key({"api_base": "", "api_key": "k"}) is None
+ assert EncryptedContentAffinityCheck._encryption_boundary_key({"api_base": "https://x"}) is None
# ---------------------------------------------------------------------------
@@ -1180,10 +1104,11 @@ def test_boundary_key_rejects_non_dict_like_inputs():
# ---------------------------------------------------------------------------
-def _make_originating_mock(api_base: str, api_key: str):
+def _make_originating_mock(api_base: str, api_key: str, model_name: str = "gpt-5.4"):
from unittest.mock import MagicMock
originating = MagicMock()
+ originating.model_name = model_name
originating.litellm_params.model_dump.return_value = {
"api_base": api_base,
"api_key": api_key,
@@ -1192,19 +1117,23 @@ def _make_originating_mock(api_base: str, api_key: str):
def _make_router_mock_with_cooldown(
- originating, cooldown_entries: Optional[List[tuple]] = None
+ originating,
+ cooldown_entries: list[tuple] | None = None,
+ routed_group_model_ids: list[str] | None = None,
):
"""
Build a MagicMock router whose ``cooldown_cache.async_get_active_cooldowns``
- returns ``cooldown_entries`` (defaulting to ``[]`` — no active cooldown).
+ returns ``cooldown_entries`` (defaulting to ``[]`` — no active cooldown), and
+ whose ``get_candidate_model_ids_for_route`` returns ``routed_group_model_ids``
+ (the deployment ids the router resolves for the routed model; defaulting to ``[]``
+ — origin absent from the routed group, i.e. a tier change).
"""
from unittest.mock import AsyncMock, MagicMock
mock_router = MagicMock()
mock_router.get_deployment.return_value = originating
- mock_router.cooldown_cache.async_get_active_cooldowns = AsyncMock(
- return_value=list(cooldown_entries or [])
- )
+ mock_router.cooldown_cache.async_get_active_cooldowns = AsyncMock(return_value=list(cooldown_entries or []))
+ mock_router.get_candidate_model_ids_for_route.return_value = frozenset(routed_group_model_ids or [])
return mock_router
@@ -1235,15 +1164,15 @@ async def test_affinity_raises_service_unavailable_when_origin_cooled_for_non_42
},
)
],
+ routed_group_model_ids=["deployment-a-cooled", "deployment-b"],
)
check = EncryptedContentAffinityCheck(router=mock_router)
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- "deployment-a-cooled", "rs_test"
- )
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a-cooled", "rs_test")
healthy_only_b = [
{
"model_info": {"id": "deployment-b"},
+ "model_name": "gpt-5.4",
"litellm_params": {
"api_base": "https://account-b.openai.azure.com/",
"api_key": "key-b",
@@ -1297,15 +1226,15 @@ async def test_affinity_raises_rate_limit_with_retry_after_when_origin_cooled_fo
},
)
],
+ routed_group_model_ids=["deployment-a-cooled-429", "deployment-b"],
)
check = EncryptedContentAffinityCheck(router=mock_router)
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- "deployment-a-cooled-429", "rs_test"
- )
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a-cooled-429", "rs_test")
healthy_only_b = [
{
"model_info": {"id": "deployment-b"},
+ "model_name": "gpt-5.4",
"litellm_params": {
"api_base": "https://account-b.openai.azure.com/",
"api_key": "key-b",
@@ -1345,15 +1274,16 @@ async def test_affinity_raises_service_unavailable_when_origin_filtered_without_
)
originating = _make_originating_mock("https://account-a.openai.azure.com/", "key-a")
- mock_router = _make_router_mock_with_cooldown(originating, cooldown_entries=[])
+ mock_router = _make_router_mock_with_cooldown(
+ originating, cooldown_entries=[], routed_group_model_ids=["deployment-a-filtered", "deployment-b"]
+ )
check = EncryptedContentAffinityCheck(router=mock_router)
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- "deployment-a-filtered", "rs_test"
- )
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a-filtered", "rs_test")
healthy_only_b = [
{
"model_info": {"id": "deployment-b"},
+ "model_name": "gpt-5.4",
"litellm_params": {
"api_base": "https://account-b.openai.azure.com/",
"api_key": "key-b",
@@ -1377,15 +1307,18 @@ async def test_affinity_raises_service_unavailable_when_origin_filtered_without_
@pytest.mark.asyncio
-async def test_affinity_raises_bad_request_when_origin_removed():
+async def test_affinity_strips_and_dispatches_when_origin_is_unknown_or_removed():
"""
- Originating deployment was removed from the router config and no boundary
- peer is available. This is permanent (the stale encrypted_content cannot
- be honored), so surface a 400 with actionable text.
+ A removed deployment, or a forged/unknown affinity marker, resolves to no
+ originating deployment. It is handled like a cross-group origin: the encrypted
+ reasoning is stripped and the request dispatches with its readable history,
+ rather than returning a distinguishable error. That uniform handling denies an
+ authenticated caller a deployment-id existence oracle, an existing cross-group id
+ and a nonexistent id both strip and proceed, so responses cannot be told apart.
+ The membership lookup is skipped entirely when the origin is unknown.
"""
from unittest.mock import MagicMock
- from litellm.exceptions import BadRequestError
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
EncryptedContentAffinityCheck,
)
@@ -1394,12 +1327,11 @@ async def test_affinity_raises_bad_request_when_origin_removed():
mock_router.get_deployment.return_value = None
check = EncryptedContentAffinityCheck(router=mock_router)
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- "deployment-removed", "rs_test"
- )
- healthy_only_b = [
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-removed")
+ routed_pool = [
{
"model_info": {"id": "deployment-b"},
+ "model_name": "gpt-5.4",
"litellm_params": {
"api_base": "https://account-b.openai.azure.com/",
"api_key": "key-b",
@@ -1408,18 +1340,28 @@ async def test_affinity_raises_bad_request_when_origin_removed():
}
]
request_kwargs = {
- "input": [{"id": encoded_id, "type": "reasoning"}],
+ "litellm_metadata": {},
+ "input": [
+ {"role": "user", "content": "why is the sky blue?"},
+ {
+ "type": "reasoning",
+ "encrypted_content": wrapped,
+ "summary": [{"type": "summary_text", "text": "scattering"}],
+ },
+ {"role": "user", "content": "and sunsets?"},
+ ],
}
- with pytest.raises(BadRequestError) as excinfo:
- await check.async_filter_deployments(
- model="gpt-5.4",
- healthy_deployments=healthy_only_b,
- messages=None,
- request_kwargs=request_kwargs,
- )
+ result = await check.async_filter_deployments(
+ model="gpt-5.4",
+ healthy_deployments=routed_pool,
+ messages=None,
+ request_kwargs=request_kwargs,
+ )
- assert "deployment-removed" not in str(excinfo.value)
+ assert result is routed_pool
+ assert not any(isinstance(item, dict) and item.get("encrypted_content") for item in request_kwargs["input"])
+ mock_router.get_candidate_model_ids_for_route.assert_not_called()
@pytest.mark.asyncio
@@ -1444,9 +1386,7 @@ async def test_affinity_does_not_raise_when_boundary_peer_available():
mock_router.get_deployment.return_value = originating
check = EncryptedContentAffinityCheck(router=mock_router)
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- "deployment-a", "rs_test"
- )
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a", "rs_test")
peer = {
"model_info": {"id": "deployment-a-peer"},
"litellm_params": {
@@ -1490,9 +1430,7 @@ async def test_model_group_affinity_config_enables_encrypted_content_affinity():
},
target_deployment,
]
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- "deployment-b", "rs_test"
- )
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-b", "rs_test")
request_kwargs = {
"input": [{"type": "reasoning", "id": encoded_id}],
"litellm_metadata": {},
@@ -1536,9 +1474,7 @@ async def test_model_group_affinity_config_does_not_disable_global_encrypted_con
},
target_deployment,
]
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- "deployment-b", "rs_test"
- )
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-b", "rs_test")
request_kwargs = {
"input": [{"type": "reasoning", "id": encoded_id}],
"litellm_metadata": {},
@@ -1600,15 +1536,9 @@ async def test_model_group_encrypted_content_affinity_overrides_global_deploymen
try:
callbacks = router.optional_callbacks or []
- deployment_callback = next(
- cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck)
- )
- encrypted_content_callback = next(
- cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck)
- )
- assert callbacks.index(encrypted_content_callback) < callbacks.index(
- deployment_callback
- )
+ deployment_callback = next(cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck))
+ encrypted_content_callback = next(cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck))
+ assert callbacks.index(encrypted_content_callback) < callbacks.index(deployment_callback)
assert encrypted_content_callback.enable_global_affinity is False
cache_key = DeploymentAffinityCheck.get_affinity_cache_key(
@@ -1620,9 +1550,7 @@ async def test_model_group_encrypted_content_affinity_overrides_global_deploymen
value={"model_id": "deployment-a"},
ttl=60,
)
- encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
- "deployment-b", "rs_test"
- )
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-b", "rs_test")
request_kwargs = {
"input": [
{
@@ -1643,16 +1571,301 @@ async def test_model_group_encrypted_content_affinity_overrides_global_deploymen
)
assert after_deployment_affinity == [deployment_a, deployment_b]
- after_encrypted_content_affinity = (
- await encrypted_content_callback.async_filter_deployments(
- model=model_group,
- healthy_deployments=after_deployment_affinity,
- messages=None,
- request_kwargs=request_kwargs,
- )
+ after_encrypted_content_affinity = await encrypted_content_callback.async_filter_deployments(
+ model=model_group,
+ healthy_deployments=after_deployment_affinity,
+ messages=None,
+ request_kwargs=request_kwargs,
)
assert after_encrypted_content_affinity == [deployment_b]
assert request_kwargs.get("_encrypted_content_affinity_pinned") is True
finally:
router.discard()
+
+
+class TestStripEncryptedReasoningFromInput:
+ def test_keeps_summary_and_drops_encrypted_content_and_id(self):
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a")
+ encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a", "rs_1")
+ request_input = [
+ {"role": "user", "content": "first turn"},
+ {
+ "type": "reasoning",
+ "id": encoded_id,
+ "encrypted_content": wrapped,
+ "summary": [{"type": "summary_text", "text": "thought about it"}],
+ },
+ {"type": "reasoning", "id": encoded_id, "encrypted_content": wrapped},
+ {"type": "reasoning", "encrypted_content": wrapped, "summary": []},
+ {"type": "message", "id": "msg_1", "role": "assistant", "content": "hi"},
+ {"role": "user", "content": "second turn"},
+ ]
+ ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
+ assert request_input == [
+ {"role": "user", "content": "first turn"},
+ {
+ "type": "reasoning",
+ "summary": [{"type": "summary_text", "text": "thought about it"}],
+ },
+ {"type": "message", "id": "msg_1", "role": "assistant", "content": "hi"},
+ {"role": "user", "content": "second turn"},
+ ]
+
+ def test_keeps_string_form_summary_when_stripping(self):
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a")
+ request_input = [
+ {"type": "reasoning", "encrypted_content": wrapped, "summary": "plain string thought"},
+ {
+ "type": "reasoning",
+ "encrypted_content": wrapped,
+ "content": [{"type": "output_text", "text": "in content"}],
+ },
+ {"type": "reasoning", "encrypted_content": wrapped, "summary": "", "content": []},
+ ]
+ ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
+ assert request_input == [
+ {"type": "reasoning", "summary": "plain string thought"},
+ {"type": "reasoning", "content": [{"type": "output_text", "text": "in content"}]},
+ ]
+
+ def test_leaves_input_untouched_when_no_encrypted_reasoning(self):
+ request_input = [
+ {"role": "user", "content": "first turn"},
+ {"type": "reasoning", "summary": [{"type": "summary_text", "text": "no blob here"}]},
+ {"role": "user", "content": "second turn"},
+ ]
+ before = [dict(item) for item in request_input]
+ ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
+ assert request_input == before
+
+
+def _cross_group_request_kwargs():
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a")
+ return {
+ "litellm_metadata": {},
+ "input": [
+ {"role": "user", "content": "ZEBRA: why is the sky blue?"},
+ {
+ "type": "reasoning",
+ "encrypted_content": wrapped,
+ "summary": [{"type": "summary_text", "text": "scattering"}],
+ },
+ {"type": "message", "role": "assistant", "content": "Rayleigh scattering."},
+ {"role": "user", "content": "KIWI: and sunsets?"},
+ ],
+ }
+
+
+@pytest.mark.asyncio
+async def test_affinity_strips_encrypted_reasoning_when_routed_to_another_model_group():
+ """
+ An auto-router tier change (or a model switch with no boundary peer): the
+ routed pool holds no deployment of the origin's model group. The origin is
+ healthy, so a 503 would be wrong; the follow-up dispatches to the routed
+ pool with the origin's encrypted reasoning stripped and its summary kept.
+ """
+ from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
+ EncryptedContentAffinityCheck,
+ )
+
+ originating = _make_originating_mock(None, "key-a", model_name="gpt-reasoning-tier")
+ mock_router = _make_router_mock_with_cooldown(
+ originating, cooldown_entries=[], routed_group_model_ids=["deployment-b"]
+ )
+ check = EncryptedContentAffinityCheck(router=mock_router)
+ routed_pool = [
+ {
+ "model_info": {"id": "deployment-b"},
+ "model_name": "gpt-simple-tier",
+ "litellm_params": {
+ "api_base": "https://gateway.example/v1",
+ "api_key": "key-b",
+ "model": "openai/gpt-5-nano",
+ },
+ }
+ ]
+ request_kwargs = _cross_group_request_kwargs()
+ original_input = request_kwargs["input"]
+
+ result = await check.async_filter_deployments(
+ model="gpt-simple-tier",
+ healthy_deployments=routed_pool,
+ messages=None,
+ request_kwargs=request_kwargs,
+ )
+
+ assert result is routed_pool
+ assert "_encrypted_content_affinity_pinned" not in request_kwargs
+ assert request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] is True
+ assert request_kwargs["input"] is original_input
+ assert [item.get("type") or item["role"] for item in original_input] == [
+ "user",
+ "reasoning",
+ "message",
+ "user",
+ ]
+ assert original_input[1] == {
+ "type": "reasoning",
+ "summary": [{"type": "summary_text", "text": "scattering"}],
+ }
+ assert not any(isinstance(item, dict) and item.get("encrypted_content") for item in original_input)
+
+
+@pytest.mark.asyncio
+async def test_affinity_fails_fast_within_the_origins_own_group():
+ """
+ Negative class for the tier-change discriminator: the routed group IS the
+ origin's group (a same-group cooldown, not a tier change), so even with a
+ healthy non-origin sibling that cannot decrypt the content, the request
+ still fails fast and the encrypted reasoning is left intact rather than
+ stripped. Preserves the LIT-3051 cooldown contract.
+ """
+ from litellm.exceptions import ServiceUnavailableError
+ from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
+ EncryptedContentAffinityCheck,
+ )
+
+ originating = _make_originating_mock(
+ "https://account-a.openai.azure.com/", "key-a", model_name="gpt-reasoning-tier"
+ )
+ mock_router = _make_router_mock_with_cooldown(
+ originating, cooldown_entries=[], routed_group_model_ids=["deployment-a", "deployment-a2"]
+ )
+ check = EncryptedContentAffinityCheck(router=mock_router)
+ sibling_pool = [
+ {
+ "model_info": {"id": "deployment-a2"},
+ "model_name": "gpt-reasoning-tier",
+ "litellm_params": {
+ "api_base": "https://account-a2.openai.azure.com/",
+ "api_key": "key-a2",
+ "model": "azure/gpt-5.4",
+ },
+ }
+ ]
+ request_kwargs = _cross_group_request_kwargs()
+
+ with pytest.raises(ServiceUnavailableError):
+ await check.async_filter_deployments(
+ model="gpt-reasoning-tier",
+ healthy_deployments=sibling_pool,
+ messages=None,
+ request_kwargs=request_kwargs,
+ )
+
+ assert request_kwargs["input"][1].get("encrypted_content")
+
+
+@pytest.mark.asyncio
+async def test_affinity_does_not_strip_when_group_is_spelled_differently_but_same_by_id():
+ """
+ The discriminator must key on deployment-id membership, not on the model-group
+ name string. Here the origin's configured group is spelled ``openai/gpt-5.4-mini``
+ while the routed group is the canonical ``gpt-5.4-mini``: same group, different
+ spelling. A name compare (``originating.model_name != model``) would read this as
+ a tier change and strip the reasoning it did not have to. Because the origin's id
+ is a member of the routed group, this is a same-group cooldown instead: the request
+ fails fast and the encrypted reasoning is left intact.
+ """
+ from litellm.exceptions import ServiceUnavailableError
+ from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
+ EncryptedContentAffinityCheck,
+ )
+
+ originating = _make_originating_mock(None, "key-a", model_name="openai/gpt-5.4-mini")
+ mock_router = _make_router_mock_with_cooldown(
+ originating, cooldown_entries=[], routed_group_model_ids=["deployment-mini-a", "deployment-mini-b"]
+ )
+ check = EncryptedContentAffinityCheck(router=mock_router)
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-mini-a")
+ sibling_pool = [
+ {
+ "model_info": {"id": "deployment-mini-b"},
+ "model_name": "gpt-5.4-mini",
+ "litellm_params": {
+ "api_base": "https://gateway.example/v1",
+ "api_key": "key-b",
+ "model": "openai/gpt-5.4-mini",
+ },
+ }
+ ]
+ request_kwargs = {
+ "litellm_metadata": {},
+ "input": [
+ {"role": "user", "content": "why is the sky blue?"},
+ {
+ "type": "reasoning",
+ "encrypted_content": wrapped,
+ "summary": [{"type": "summary_text", "text": "scattering"}],
+ },
+ {"role": "user", "content": "and sunsets?"},
+ ],
+ }
+
+ with pytest.raises(ServiceUnavailableError):
+ await check.async_filter_deployments(
+ model="gpt-5.4-mini",
+ healthy_deployments=sibling_pool,
+ messages=None,
+ request_kwargs=request_kwargs,
+ )
+
+ assert request_kwargs["input"][1].get("encrypted_content")
+
+
+@pytest.mark.asyncio
+async def test_affinity_honors_router_candidate_ids_for_team_and_pattern_routes():
+ """
+ The exact `model_name` index does not include team-public or pattern routes, so a
+ same-group cooldown reached only through one of those would be misread as a tier change
+ and stripped. The check asks the router for the candidate ids it resolves for the route
+ (`get_candidate_model_ids_for_route`), which covers those paths, rather than the bare
+ index. Here that set marks the origin as a candidate, so the request fails fast with its
+ reasoning intact, and the routed group and team are passed through to the router.
+ """
+ from litellm.exceptions import ServiceUnavailableError
+ from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
+ EncryptedContentAffinityCheck,
+ )
+
+ originating = _make_originating_mock(None, "key-a", model_name="model_name_teamA_uuid")
+ mock_router = _make_router_mock_with_cooldown(
+ originating, cooldown_entries=[], routed_group_model_ids=["deployment-team-a", "deployment-team-b"]
+ )
+ check = EncryptedContentAffinityCheck(router=mock_router)
+ wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-team-a")
+ sibling_pool = [
+ {
+ "model_info": {"id": "deployment-team-b"},
+ "model_name": "team-public-model",
+ "litellm_params": {
+ "api_base": "https://gateway.example/v1",
+ "api_key": "key-b",
+ "model": "openai/gpt-5.4-mini",
+ },
+ }
+ ]
+ request_kwargs = {
+ "litellm_metadata": {"user_api_key_team_id": "teamA"},
+ "input": [
+ {"role": "user", "content": "why is the sky blue?"},
+ {
+ "type": "reasoning",
+ "encrypted_content": wrapped,
+ "summary": [{"type": "summary_text", "text": "scattering"}],
+ },
+ {"role": "user", "content": "and sunsets?"},
+ ],
+ }
+
+ with pytest.raises(ServiceUnavailableError):
+ await check.async_filter_deployments(
+ model="team-public-model",
+ healthy_deployments=sibling_pool,
+ messages=None,
+ request_kwargs=request_kwargs,
+ )
+
+ assert request_kwargs["input"][1].get("encrypted_content")
+ mock_router.get_candidate_model_ids_for_route.assert_called_once_with(model="team-public-model", team_id="teamA")
diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py
index f8fa2231597..8f8a7640c08 100644
--- a/tests/test_litellm/test_cost_calculator.py
+++ b/tests/test_litellm/test_cost_calculator.py
@@ -3736,6 +3736,112 @@ def test_completion_cost_logs_reasoning_and_cache_breakdown(_local_model_cost_ma
assert logging_obj.cost_breakdown["cache_read_cost"] == pytest.approx(100 * 3e-08)
+def test_completion_cost_logs_the_rates_it_billed_at(monkeypatch):
+ """A caller reporting the cost lines beside their per-token rates reads both off this one call.
+ completion_cost infers the provider, and xai's inclusive tier thresholds put a request sitting
+ exactly on 200k at the tier rate, which a lookup made without that inferred provider would miss.
+ """
+ from datetime import datetime
+
+ from litellm.litellm_core_utils.litellm_logging import Logging
+
+ monkeypatch.setitem(
+ litellm.model_cost,
+ "xai/tiered-model",
+ {
+ "input_cost_per_token": 3e-6,
+ "output_cost_per_token": 15e-6,
+ "cache_read_input_token_cost": 3e-7,
+ "input_cost_per_token_above_200k_tokens": 6e-6,
+ "output_cost_per_token_above_200k_tokens": 3e-5,
+ "cache_read_input_token_cost_above_200k_tokens": 6e-7,
+ "litellm_provider": "xai",
+ "mode": "chat",
+ },
+ )
+ logging_obj = Logging(
+ model="xai/tiered-model",
+ messages=[{"role": "user", "content": "Hello"}],
+ stream=False,
+ call_type="completion",
+ start_time=datetime.now(),
+ litellm_call_id="billed-rates",
+ function_id="f",
+ )
+ usage = Usage(
+ prompt_tokens=200_000,
+ completion_tokens=1_000,
+ total_tokens=201_000,
+ prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=100_000),
+ )
+
+ litellm.completion_cost(
+ completion_response=ModelResponse(model="xai/tiered-model", usage=usage),
+ model="xai/tiered-model",
+ custom_llm_provider=None,
+ litellm_logging_obj=logging_obj,
+ )
+
+ rates = logging_obj.billed_token_rates
+ assert rates is not None
+ assert rates.input_cost_per_token == pytest.approx(6e-6)
+ assert rates.cache_read_input_token_cost == pytest.approx(6e-7)
+ assert logging_obj.cost_breakdown["cache_read_cost"] == pytest.approx(
+ 100_000 * rates.cache_read_input_token_cost
+ )
+ assert logging_obj.cost_breakdown["output_cost"] == pytest.approx(1_000 * rates.output_cost_per_token)
+
+
+def test_completion_cost_logs_cache_and_reasoning_breakdown_for_custom_pricing():
+ """
+ A custom-priced deployment bills cache tokens at its custom cache rates, but the
+ breakdown stored for the spend logs carried no cache or reasoning lines for it.
+ """
+ from datetime import datetime
+
+ from litellm.litellm_core_utils.litellm_logging import Logging
+ from litellm.types.utils import CompletionTokensDetailsWrapper, CostPerToken
+
+ logging_obj = Logging(
+ model="openai/onprem-model",
+ messages=[{"role": "user", "content": "Hello"}],
+ stream=False,
+ call_type="completion",
+ start_time=datetime.now(),
+ litellm_call_id="custom-pricing-breakdown",
+ function_id="f",
+ )
+ response = ModelResponse(
+ model="openai/onprem-model",
+ usage=Usage(
+ prompt_tokens=1000,
+ completion_tokens=500,
+ total_tokens=1500,
+ prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=800, cache_creation_tokens=100),
+ completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=200),
+ ),
+ )
+
+ total = completion_cost(
+ completion_response=response,
+ model="openai/onprem-model",
+ custom_llm_provider="openai",
+ custom_cost_per_token=CostPerToken(
+ input_cost_per_token=1e-6,
+ output_cost_per_token=2e-6,
+ cache_read_input_token_cost=1e-7,
+ cache_creation_input_token_cost=1.25e-6,
+ ),
+ litellm_logging_obj=logging_obj,
+ )
+
+ assert logging_obj.cost_breakdown is not None
+ assert logging_obj.cost_breakdown["cache_read_cost"] == pytest.approx(800 * 1e-7)
+ assert logging_obj.cost_breakdown["cache_creation_cost"] == pytest.approx(100 * 1.25e-6)
+ assert logging_obj.cost_breakdown["reasoning_cost"] == pytest.approx(200 * 2e-6)
+ assert total == pytest.approx(100 * 1e-6 + 800 * 1e-7 + 100 * 1.25e-6 + 500 * 2e-6)
+
+
def test_cost_per_token_per_second_pricing(monkeypatch):
"""
Models priced by duration (input/output_cost_per_second) with no per-token rates
diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py
index bc79c5f6589..80da922724d 100644
--- a/tests/test_litellm/test_router.py
+++ b/tests/test_litellm/test_router.py
@@ -14728,3 +14728,45 @@ def test_router_stays_quiet_when_a_deployment_drop_params_is_a_flag(value, caplo
)
assert "is not a flag value" not in caplog.text
+
+
+def test_get_candidate_model_ids_for_route_covers_model_name_and_pattern():
+ """
+ get_candidate_model_ids_for_route resolves a route the way the router does, so a
+ pre-call check can tell a genuine cross-group route from same-group unavailability.
+ A concrete model group returns its member ids; a wildcard/pattern deployment is
+ included for a concrete model it matches, which the bare model_name index misses.
+ Regression guard for the LIT-7195 tier-change discriminator's team/pattern gaps.
+ """
+ router = Router(
+ model_list=[
+ {
+ "model_name": "grp",
+ "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-a", "api_base": "https://x.invalid"},
+ "model_info": {"id": "dep-a"},
+ },
+ {
+ "model_name": "grp",
+ "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-b", "api_base": "https://x.invalid"},
+ "model_info": {"id": "dep-b"},
+ },
+ {
+ "model_name": "openai/*",
+ "litellm_params": {"model": "openai/*", "api_key": "sk-c", "api_base": "https://x.invalid"},
+ "model_info": {"id": "dep-wild"},
+ },
+ ]
+ )
+
+ assert router.get_candidate_model_ids_for_route(model="grp") == frozenset({"dep-a", "dep-b"})
+ assert "dep-wild" in router.get_candidate_model_ids_for_route(model="openai/gpt-4o-some-new-model")
+
+
+def test_deployment_ids_stringifies_ids_and_skips_entries_without_a_model_info_id():
+ deployments = (
+ {"model_info": {"id": "a"}},
+ {"model_info": {"id": 2}},
+ {"model_info": {}},
+ {"no_model_info": True},
+ )
+ assert Router._deployment_ids(deployments) == frozenset({"a", "2"})
diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py
index 26f8e9778f5..312dd3d56b1 100644
--- a/tests/test_litellm/test_utils.py
+++ b/tests/test_litellm/test_utils.py
@@ -2,10 +2,12 @@ import asyncio
import json
import logging
import os
+import threading
from datetime import datetime, timedelta, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
+import httpx
import pytest
import respx
from jsonschema import validate
@@ -2384,6 +2386,28 @@ def test_register_model_with_scientific_notation():
_invalidate_model_cost_lowercase_map()
+@respx.mock
+def test_register_model_url_fetch_uses_single_attempt(monkeypatch):
+ monkeypatch.delenv("LITELLM_LOCAL_MODEL_COST_MAP", raising=False)
+ monkeypatch.setattr(litellm, "model_cost", dict(litellm.model_cost))
+ before = dict(litellm.model_cost)
+ threads_before = {thread.name for thread in threading.enumerate()}
+ route = respx.get("https://example.invalid/custom_pricing.json").mock(
+ return_value=httpx.Response(503)
+ )
+
+ litellm.register_model(model_cost="https://example.invalid/custom_pricing.json")
+
+ threads_after = {thread.name for thread in threading.enumerate()}
+ assert route.call_count == 1
+ assert not (threads_after - threads_before) & {"litellm-model-cost-map-retry"}
+ assert not any(
+ thread.name == "litellm-model-cost-map-retry" and thread.is_alive()
+ for thread in threading.enumerate()
+ )
+ assert litellm.model_cost.keys() >= before.keys()
+
+
def test_register_model_openrouter_without_slash():
"""
Test that register_model handles openrouter models without '/' in the name.
diff --git a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.test.tsx b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.test.tsx
index 040fa463f9c..af3f98b746b 100644
--- a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.test.tsx
+++ b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.test.tsx
@@ -89,6 +89,65 @@ describe("PassThroughEndpointsTable", () => {
expect(onDeleteClick).toHaveBeenCalledWith("ep-1");
});
+ it("should disable edit and delete for config-defined endpoints", async () => {
+ const user = userEvent.setup();
+ const onEndpointClick = vi.fn();
+ const onDeleteClick = vi.fn();
+ const configEndpoint: passThroughItem = {
+ id: "ep-config",
+ path: "/from-config",
+ target: "https://config.example.com",
+ headers: {},
+ is_from_config: true,
+ };
+ render(
+ ,
+ );
+
+ await user.click(screen.getByTestId("endpoint-actions-ep-config"));
+ const editItem = await screen.findByTestId("endpoint-action-edit");
+ const deleteItem = await screen.findByTestId("endpoint-action-delete");
+
+ expect(editItem).toHaveAttribute("data-disabled");
+ expect(deleteItem).toHaveAttribute("data-disabled");
+ expect(screen.getByTestId("endpoint-config-hint")).toHaveTextContent(
+ "This endpoint is defined in the config file and cannot be edited or deleted on the dashboard.",
+ );
+
+ await user.click(editItem);
+ await user.click(deleteItem);
+
+ expect(onEndpointClick).not.toHaveBeenCalled();
+ expect(onDeleteClick).not.toHaveBeenCalled();
+ });
+
+ it("should not show the config hint for DB endpoints", async () => {
+ const user = userEvent.setup();
+ render();
+
+ await user.click(screen.getByTestId("endpoint-actions-ep-1"));
+ await screen.findByTestId("endpoint-action-delete");
+ expect(screen.queryByTestId("endpoint-config-hint")).not.toBeInTheDocument();
+ });
+
+ it("should label endpoint source as Config or DB", () => {
+ const configEndpoint: passThroughItem = {
+ id: "ep-config",
+ path: "/from-config",
+ target: "https://config.example.com",
+ headers: {},
+ is_from_config: true,
+ };
+ render();
+ expect(screen.getByText("Config")).toBeInTheDocument();
+ expect(screen.getAllByText("DB")).toHaveLength(2);
+ });
+
it("should disable edit and delete for endpoints without an id", async () => {
const user = userEvent.setup();
const onEndpointClick = vi.fn();
diff --git a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTableColumns.tsx b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTableColumns.tsx
index d22b274861a..5b18685a140 100644
--- a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTableColumns.tsx
+++ b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTableColumns.tsx
@@ -18,6 +18,9 @@ import { cn } from "@/lib/cva.config";
import type { passThroughItem } from "./PassThroughSettings";
+const CONFIG_ENDPOINT_HINT =
+ "This endpoint is defined in the config file and cannot be edited or deleted on the dashboard.";
+
function HeaderWithTooltip({ title, tooltip }: { title: string; tooltip: string }) {
return (
@@ -73,6 +76,7 @@ interface EndpointRowActionsProps {
function EndpointRowActions({ endpoint, onEndpointClick, onDeleteClick }: EndpointRowActionsProps) {
const endpointId = endpoint.id;
+ const isFromConfig = endpoint.is_from_config ?? false;
return (
endpointId && onEndpointClick(endpointId)}
+ disabled={isFromConfig || !endpointId}
+ onClick={() => !isFromConfig && endpointId && onEndpointClick(endpointId)}
>
Edit
@@ -95,12 +99,17 @@ function EndpointRowActions({ endpoint, onEndpointClick, onDeleteClick }: Endpoi
endpointId && onDeleteClick(endpointId)}
+ disabled={isFromConfig || !endpointId}
+ onClick={() => !isFromConfig && endpointId && onDeleteClick(endpointId)}
>
Delete
+ {isFromConfig && (
+
+ {CONFIG_ENDPOINT_HINT}
+
+ )}
);
@@ -124,7 +133,9 @@ export const getPassThroughEndpointsTableColumns = ({
enableSorting: false,
cell: ({ row }) => {
const endpointId = row.original.id;
- if (!endpointId) return
—;
+ if (!endpointId || row.original.is_from_config) {
+ return
—;
+ }
return (
{
+ const isFromConfig = row.original.is_from_config ?? false;
+ return ;
+ },
+ },
{
id: "path",
accessorKey: "path",
diff --git a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.tsx b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.tsx
index 2dc6fdbd32c..8ef2766d412 100644
--- a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.tsx
+++ b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.tsx
@@ -1,4 +1,13 @@
import React, { useState, useEffect } from "react";
+import {
+ AlertDialog,
+ AlertDialogCancel,
+ AlertDialogContent,
+ AlertDialogDescription,
+ AlertDialogFooter,
+ AlertDialogHeader,
+ AlertDialogTitle,
+} from "@/components/ui/alert-dialog";
import { Button } from "@/components/ui/button";
import { deletePassThroughEndpointsCall, getPassThroughEndpointsCall } from "../networking";
import AddPassThroughEndpoint from "../add_pass_through";
@@ -25,6 +34,7 @@ export interface passThroughItem {
methods?: string[];
guardrails?: Record;
default_query_params?: Record;
+ is_from_config?: boolean;
}
const PassThroughSettings: React.FC = ({ accessToken, userRole, userID, premiumUser }) => {
@@ -133,42 +143,22 @@ const PassThroughSettings: React.FC = ({ accessToken,
onDeleteClick={handleDelete}
/>
- {isDeleteModalOpen && (
-
-
-
-
-
-
-
-
-
-
-
-
-
Delete Pass-Through Endpoint
-
-
- Are you sure you want to delete this pass-through endpoint? This action cannot be undone.
-
-
-
-
-
-
-
-
-
-
-
-
- )}
+ !open && cancelDelete()}>
+
+
+ Delete Pass-Through Endpoint
+
+ Are you sure you want to delete this pass-through endpoint? This action cannot be undone.
+
+
+
+ Cancel
+
+
+
+
);
};
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index a77332f1dcb..6acdeca3e75 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -3318,11 +3318,14 @@ export interface paths {
* - model: Model name (e.g., "gpt-4", "claude-3-opus")
* - input_tokens: Expected input tokens per request
* - output_tokens: Expected output tokens per request
+ * - cache_read_input_tokens: Cache-read tokens per request, counted within input_tokens (optional)
+ * - cache_creation_input_tokens: Cache-write tokens per request, counted within input_tokens (optional)
+ * - reasoning_tokens: Reasoning tokens per request, counted within output_tokens (optional)
* - num_requests_per_day: Number of requests per day (optional)
* - num_requests_per_month: Number of requests per month (optional)
*
* Returns cost breakdown including:
- * - Per-request costs (input, output, margin)
+ * - Per-request costs (input, output, margin, plus the cache-read, cache-write and reasoning shares)
* - Daily costs (if num_requests_per_day provided)
* - Monthly costs (if num_requests_per_month provided)
*
@@ -3331,7 +3334,9 @@ export interface paths {
* {
* "model": "gpt-4",
* "input_tokens": 1000,
+ * "cache_read_input_tokens": 800,
* "output_tokens": 500,
+ * "reasoning_tokens": 200,
* "num_requests_per_day": 100,
* "num_requests_per_month": 3000
* }
@@ -26528,6 +26533,18 @@ export interface components {
* @description Request body for /cost/estimate endpoint.
*/
CostEstimateRequest: {
+ /**
+ * Cache Creation Input Tokens
+ * @description Input tokens written to the prompt cache; counted within input_tokens
+ * @default 0
+ */
+ cache_creation_input_tokens: number;
+ /**
+ * Cache Read Input Tokens
+ * @description Input tokens read from the prompt cache; counted within input_tokens
+ * @default 0
+ */
+ cache_read_input_tokens: number;
/**
* Input Tokens
* @description Expected input tokens per request
@@ -26553,17 +26570,65 @@ export interface components {
* @description Expected output tokens per request
*/
output_tokens: number;
+ /**
+ * Reasoning Tokens
+ * @description Reasoning tokens the model emits; counted within output_tokens
+ * @default 0
+ */
+ reasoning_tokens: number;
};
/**
* CostEstimateResponse
* @description Response body for /cost/estimate endpoint.
*/
CostEstimateResponse: {
+ /**
+ * Cache Creation Cost Per Request
+ * @description Cache-write share of input_cost_per_request
+ * @default 0
+ */
+ cache_creation_cost_per_request: number;
+ /**
+ * Cache Creation Input Token Cost
+ * @description Rate billed per cache-write token
+ */
+ cache_creation_input_token_cost?: number | null;
+ /**
+ * Cache Creation Input Tokens
+ * @default 0
+ */
+ cache_creation_input_tokens: number;
+ /**
+ * Cache Read Cost Per Request
+ * @description Cache-read share of input_cost_per_request
+ * @default 0
+ */
+ cache_read_cost_per_request: number;
+ /**
+ * Cache Read Input Token Cost
+ * @description Rate billed per cache-read token
+ */
+ cache_read_input_token_cost?: number | null;
+ /**
+ * Cache Read Input Tokens
+ * @default 0
+ */
+ cache_read_input_tokens: number;
/**
* Cost Per Request
* @description Total cost per request (includes margin)
*/
cost_per_request: number;
+ /**
+ * Daily Cache Creation Cost
+ * @description Cache-write share of daily_input_cost
+ */
+ daily_cache_creation_cost?: number | null;
+ /**
+ * Daily Cache Read Cost
+ * @description Cache-read share of daily_input_cost
+ */
+ daily_cache_read_cost?: number | null;
/**
* Daily Cost
* @description Total daily cost (includes margin)
@@ -26584,12 +26649,20 @@ export interface components {
* @description Daily output token cost
*/
daily_output_cost?: number | null;
+ /**
+ * Daily Reasoning Cost
+ * @description Reasoning share of daily_output_cost
+ */
+ daily_reasoning_cost?: number | null;
/**
* Input Cost Per Request
* @description Input token cost per request (before margin)
*/
input_cost_per_request: number;
- /** Input Cost Per Token */
+ /**
+ * Input Cost Per Token
+ * @description Rate billed per input token
+ */
input_cost_per_token?: number | null;
/** Input Tokens */
input_tokens: number;
@@ -26601,6 +26674,16 @@ export interface components {
margin_cost_per_request: number;
/** Model */
model: string;
+ /**
+ * Monthly Cache Creation Cost
+ * @description Cache-write share of monthly_input_cost
+ */
+ monthly_cache_creation_cost?: number | null;
+ /**
+ * Monthly Cache Read Cost
+ * @description Cache-read share of monthly_input_cost
+ */
+ monthly_cache_read_cost?: number | null;
/**
* Monthly Cost
* @description Total monthly cost (includes margin)
@@ -26621,21 +26704,45 @@ export interface components {
* @description Monthly output token cost
*/
monthly_output_cost?: number | null;
+ /**
+ * Monthly Reasoning Cost
+ * @description Reasoning share of monthly_output_cost
+ */
+ monthly_reasoning_cost?: number | null;
/** Num Requests Per Day */
num_requests_per_day?: number | null;
/** Num Requests Per Month */
num_requests_per_month?: number | null;
+ /**
+ * Output Cost Per Reasoning Token
+ * @description Rate billed per reasoning token
+ */
+ output_cost_per_reasoning_token?: number | null;
/**
* Output Cost Per Request
* @description Output token cost per request (before margin)
*/
output_cost_per_request: number;
- /** Output Cost Per Token */
+ /**
+ * Output Cost Per Token
+ * @description Rate billed per output token
+ */
output_cost_per_token?: number | null;
/** Output Tokens */
output_tokens: number;
/** Provider */
provider?: string | null;
+ /**
+ * Reasoning Cost Per Request
+ * @description Reasoning share of output_cost_per_request
+ * @default 0
+ */
+ reasoning_cost_per_request: number;
+ /**
+ * Reasoning Tokens
+ * @default 0
+ */
+ reasoning_tokens: number;
};
/** CreateCredentialItem */
CreateCredentialItem: {