Merge litellm_internal_staging into OTEL v2 destinations
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled

This commit is contained in:
Devin AI 2026-07-21 22:40:12 +00:00
commit 9d9832689e
188 changed files with 8957 additions and 1604 deletions

View file

@ -58,6 +58,8 @@ jobs:
# free OSS, run as a pinned, checksum-verified binary; no GitHub Action
# dependency and no vendor SaaS callout.
- name: Scan image for fixable HIGH/CRITICAL CVEs
env:
GRYPE_MATCH_PYTHON_USING_CPES: "true"
run: |
"$RUNNER_TEMP/grype" litellm-image-scan:${{ github.sha }} \
--only-fixed \

View file

@ -0,0 +1,9 @@
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_SSOIdentityAssertion" (
"user_id" TEXT NOT NULL,
"assertion_b64" TEXT NOT NULL,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_SSOIdentityAssertion_pkey" PRIMARY KEY ("user_id")
);

View file

@ -406,6 +406,15 @@ model LiteLLM_MCPServerOAuthClient {
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// The enterprise IdP identity assertion captured at SSO login, one row per user.
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
model LiteLLM_SSOIdentityAssertion {
user_id String @id
assertion_b64 String
created_at DateTime @default(now()) @map("created_at")
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id

View file

@ -1474,6 +1474,7 @@ _batch_polling_env = os.getenv("PROXY_BATCH_POLLING_ENABLED", "true").lower()
PROXY_BATCH_POLLING_ENABLED = _batch_polling_env == "true"
PROXY_BUDGET_RESCHEDULER_MAX_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605))
PROXY_BATCH_WRITE_AT = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds, increased from 10
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS = get_env_int("PROXY_CONFIG_RELOAD_INTERVAL_SECONDS", 30)
# APScheduler Configuration - MEMORY LEAK FIX
# These settings prevent memory leaks in APScheduler's normalize() and _apply_jitter() functions

View file

@ -7,8 +7,8 @@ duration_in_seconds is used in diff parts of the code base, example
"""
import re
import time
from datetime import datetime, timedelta, timezone, tzinfo
import time as time_module
from datetime import datetime, time, timedelta, timezone, tzinfo
from typing import Optional, Tuple
from zoneinfo import ZoneInfo
@ -61,7 +61,7 @@ def duration_in_seconds(duration: str) -> int:
elif unit == "w":
return value * 604800
elif unit == "mo":
now = time.time()
now = time_module.time()
current_time = datetime.fromtimestamp(now)
# Calculate target month and year, handling overflow past December
@ -94,12 +94,17 @@ def duration_in_seconds(duration: str) -> int:
raise ValueError(f"Unsupported duration unit, passed duration: {duration}")
def get_next_standardized_reset_time(duration: str, current_time: datetime, timezone_str: str = "UTC") -> datetime:
def get_next_standardized_reset_time(
duration: str,
current_time: datetime,
timezone_str: str = "UTC",
reset_time_of_day: time = time(0, 0),
) -> datetime:
"""
Get the next standardized reset time based on the duration.
All durations will reset at predictable intervals, aligned from the current time:
- Nd: If N=1, reset at next midnight; if N>1, reset every N days from now
- Nd: If N=1, reset at the next `reset_time_of_day`; if N>1, reset every N days from now
- Nh: Every N hours, aligned to hour boundaries (e.g., 1:00, 2:00)
- Nm: Every N minutes, aligned to minute boundaries (e.g., 1:05, 1:10)
- Ns: Every N seconds, aligned to second boundaries
@ -108,12 +113,15 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time
- duration: Duration string (e.g. "30s", "30m", "30h", "30d")
- current_time: Current datetime
- timezone_str: Timezone string (e.g. "UTC", "US/Eastern", "Asia/Kolkata")
- reset_time_of_day: Wall-clock time the reset lands on for day/week/month
durations (defaults to midnight). Ignored for sub-day durations, where a
time-of-day is meaningless.
Returns:
- Next reset time at a standardized interval in the specified timezone
"""
# Set up timezone and normalize current time
current_time, tz = _setup_timezone(current_time, timezone_str)
current_time, _ = _setup_timezone(current_time, timezone_str)
# Parse duration
value, unit = _parse_duration(duration)
@ -126,9 +134,9 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time
# Handle different time units
if unit == "d":
return _handle_day_reset(current_time, base_midnight, value, tz)
return _handle_day_reset(current_time, base_midnight, value, reset_time_of_day)
elif unit == "w":
return _handle_day_reset(current_time, base_midnight, value * 7, tz)
return _handle_day_reset(current_time, base_midnight, value * 7, reset_time_of_day)
elif unit == "h":
return _handle_hour_reset(current_time, base_midnight, value)
elif unit == "m":
@ -136,7 +144,7 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time
elif unit == "s":
return _handle_second_reset(current_time, base_midnight, value)
elif unit == "mo":
return _handle_month_reset(current_time, base_midnight, value)
return _handle_month_reset(current_time, base_midnight, value, reset_time_of_day)
else:
# Unrecognized unit, default to next midnight
return base_midnight + timedelta(days=1)
@ -175,46 +183,58 @@ def _parse_duration(duration: str) -> Tuple[Optional[int], Optional[str]]:
return int(value), unit
def _handle_day_reset(current_time: datetime, base_midnight: datetime, value: int, tz: tzinfo) -> datetime:
def _apply_time_of_day(dt: datetime, reset_time_of_day: time) -> datetime:
"""Set the wall-clock time of `dt` to `reset_time_of_day`, keeping its date and tzinfo."""
return dt.replace(
hour=reset_time_of_day.hour,
minute=reset_time_of_day.minute,
second=reset_time_of_day.second,
microsecond=reset_time_of_day.microsecond,
)
def _next_occurrence(
boundary_midnight: datetime,
reset_time_of_day: time,
current_time: datetime,
period: timedelta,
) -> datetime:
"""Place the reset at `reset_time_of_day` on the boundary day, rolling forward one
`period` if that instant has already passed (or is exactly now)."""
candidate = _apply_time_of_day(boundary_midnight, reset_time_of_day)
if candidate <= current_time:
return candidate + period
return candidate
def _first_of_next_month(first_of_month: datetime) -> datetime:
"""Given the 1st of some month, return the 1st of the following month."""
if first_of_month.month == 12:
return first_of_month.replace(year=first_of_month.year + 1, month=1)
return first_of_month.replace(month=first_of_month.month + 1)
def _handle_day_reset(
current_time: datetime,
base_midnight: datetime,
value: int,
reset_time_of_day: time,
) -> datetime:
"""Handle day-based reset times."""
# Handle zero value - immediate expiration
if value == 0:
return current_time
if value == 1: # Daily reset at midnight
return base_midnight + timedelta(days=1)
elif value == 7: # Weekly reset on Monday at midnight
if value == 1: # Daily reset at the configured time of day
return _next_occurrence(base_midnight, reset_time_of_day, current_time, timedelta(days=1))
elif value == 7: # Weekly reset on Monday at the configured time of day
days_until_monday = (7 - current_time.weekday()) % 7
if days_until_monday == 0: # If today is Monday
days_until_monday = 7
return base_midnight + timedelta(days=days_until_monday)
elif value == 30: # Monthly reset on 1st at midnight
# Get 1st of next month at midnight
if current_time.month == 12:
next_reset = datetime(
year=current_time.year + 1,
month=1,
day=1,
hour=0,
minute=0,
second=0,
microsecond=0,
tzinfo=tz,
)
else:
next_reset = datetime(
year=current_time.year,
month=current_time.month + 1,
day=1,
hour=0,
minute=0,
second=0,
microsecond=0,
tzinfo=tz,
)
return next_reset
else: # Custom day value - next interval is value days from current
return current_time.replace(hour=0, minute=0, second=0, microsecond=0) + timedelta(days=value)
upcoming_monday = base_midnight + timedelta(days=days_until_monday)
return _next_occurrence(upcoming_monday, reset_time_of_day, current_time, timedelta(days=7))
elif value == 30: # Monthly reset on 1st at the configured time of day
return _handle_month_reset(current_time, base_midnight, 1, reset_time_of_day)
else: # Custom day value - next interval is value days from the start of today
return _apply_time_of_day(base_midnight + timedelta(days=value), reset_time_of_day)
def _handle_hour_reset(current_time: datetime, base_midnight: datetime, value: int) -> datetime:
@ -316,36 +336,30 @@ def _handle_second_reset(current_time: datetime, base_midnight: datetime, value:
return current_time.replace(hour=next_hour, minute=next_minute, second=next_second, microsecond=0)
def _handle_month_reset(current_time: datetime, base_midnight: datetime, value: int) -> datetime:
def _handle_month_reset(
current_time: datetime,
base_midnight: datetime,
value: int,
reset_time_of_day: time,
) -> datetime:
"""
Handle monthly reset times. For monthly resets, we always reset at the start of the next month.
Handle monthly reset times. Resets land on the 1st at `reset_time_of_day`; if the
1st of the current month at that time has already passed, roll to the 1st of next month.
Args:
current_time: Current datetime
base_midnight: Midnight of current day
value: Number of months (currently only supports 1 month resets)
reset_time_of_day: Wall-clock time the reset lands on
Returns:
datetime: First day of next month at midnight
datetime: First day of the next reset month at `reset_time_of_day`
"""
if value != 1:
raise ValueError("Monthly resets currently only support 1 month intervals")
# Get the first day of next month
if current_time.month == 12:
next_month = 1
next_year = current_time.year + 1
else:
next_month = current_time.month + 1
next_year = current_time.year
return datetime(
year=next_year,
month=next_month,
day=1,
hour=0,
minute=0,
second=0,
microsecond=0,
tzinfo=current_time.tzinfo,
)
first_of_this_month = base_midnight.replace(day=1)
candidate = _apply_time_of_day(first_of_this_month, reset_time_of_day)
if candidate <= current_time:
return _apply_time_of_day(_first_of_next_month(first_of_this_month), reset_time_of_day)
return candidate

View file

@ -5111,7 +5111,10 @@ def completion( # type: ignore
try:
if base_url is not None:
api_base = base_url
if num_retries is not None:
is_router_call = any("model_group" in (kwargs.get(k) or ()) for k in ("metadata", "litellm_metadata"))
if is_router_call:
max_retries = 0
elif num_retries is not None:
max_retries = num_retries
logging: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, litellm_logging_obj)
fallbacks = fallbacks or litellm.model_fallbacks

View file

@ -597,7 +597,14 @@ async def authorize_with_server(
):
_raise_if_not_oauth2(mcp_server)
if mcp_server.authorization_url is None:
raise HTTPException(status_code=400, detail="MCP server authorization url is not set")
raise HTTPException(
status_code=400,
detail=(
"MCP server authorization url is not configured. Servers with no url (OpenAPI "
"spec or stdio) run no resource discovery, so set Authorization URL and Token URL "
"manually, or set Issuer to discover them from the identity provider (RFC 8414)."
),
)
if mcp_server.is_dcr_bridge:
# Enforce S256 PKCE on both bridge arms. The relay arm forwards the validated,
@ -702,7 +709,14 @@ async def exchange_token_with_server(
raise HTTPException(status_code=400, detail="Unsupported grant_type")
if mcp_server.token_url is None:
raise HTTPException(status_code=400, detail="MCP server token url is not set")
raise HTTPException(
status_code=400,
detail=(
"MCP server token url is not configured. Servers with no url (OpenAPI spec or "
"stdio) run no resource discovery, so set Token URL manually, or set Issuer to "
"discover it from the identity provider (RFC 8414)."
),
)
# The id and secret must come from the same source. When the server-side client_id wins,
# falling back to the caller's secret pairs the persisted client with a foreign secret; the
@ -1262,7 +1276,14 @@ async def register_client_with_server(
return dummy_return
if mcp_server.authorization_url is None:
raise HTTPException(status_code=400, detail="MCP server authorization url is not set")
raise HTTPException(
status_code=400,
detail=(
"MCP server authorization url is not configured. Servers with no url (OpenAPI "
"spec or stdio) run no resource discovery, so set Authorization URL and Token URL "
"manually, or set Issuer to discover them from the identity provider (RFC 8414)."
),
)
if mcp_server.registration_url is None:
return dummy_return

View file

@ -224,6 +224,20 @@ def _uses_issuer_anchor(manual_issuer: str | None, is_discovery_auth_type: bool)
return _blank_to_none(manual_issuer) is not None and is_discovery_auth_type
def _has_oauth_discovery_source(server_url: str | None, use_issuer_anchor: bool) -> bool:
"""Whether the server has any source OAuth discovery can fetch metadata from.
Resource-rooted discovery (RFC 9728) is fetched from the server ``url``, so spec-only
(OpenAPI) and stdio servers, which have none, could never discover: their OAuth endpoints
stayed unset unless entered manually and ``/authorize`` served its 400 with no hint of why.
An admin-pinned issuer is a trust anchor in its own right (RFC 8414 section 3.3) whose
metadata fetch does not touch the resource at all, so an anchored server can discover with
no ``url``. Called by both build paths (config and DB) so the two cannot disagree on when
discovery is reachable.
"""
return bool(server_url) or use_issuer_anchor
def _endpoints_yield_to_issuer(
issuer: str | None,
is_discovery_auth_type: bool,
@ -610,6 +624,34 @@ def _passthrough_token_from_mcp_auth_header(
return None
async def _materialize_auth_headers(auth: httpx.Auth | None) -> dict[str, str] | None:
"""Extract the header a resolved ``httpx.Auth`` would set, as a plain dict, or None.
OpenAPI tool closures egress through ``AsyncHTTPHandler`` methods that accept headers but no
``auth``, so a resolved credential must be materialized into a header value. Driving one step
of the auth's own flow (against a throwaway request that is never sent) keeps this generic
across every auth shape without per-class branching; ``header_name`` is the resolver-arm
convention for "this auth sets a header" (``NoOpAuth`` has none and yields nothing to apply).
The materialized value is point-in-time: flow behaviors past the first request, like the M2M
one-shot 401 refetch, do not apply on this arm.
"""
if auth is None:
return None
header_name = getattr(auth, "header_name", None)
if not isinstance(header_name, str) or not header_name:
return None
probe = httpx.Request("GET", "http://localhost/")
flow = auth.async_auth_flow(probe)
try:
first_request = await flow.__anext__()
except StopAsyncIteration:
return None
finally:
await flow.aclose()
header_value = first_request.headers.get(header_name)
return {header_name: header_value} if header_value else None
def _consumes_caller_authorization(server: MCPServer) -> bool:
"""True when this server's egress forwards the caller's request-wide ``Authorization`` upstream:
the client-forwarded token modes, legacy OAuth pass-through, and legacy upstream-delegated
@ -1226,7 +1268,12 @@ class MCPServerManager:
manual_token_url = _blank_to_none(server_config.get("token_url"))
manual_registration_url = _blank_to_none(server_config.get("registration_url"))
is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type)
obo_needs_discovery = self._obo_needs_endpoint_discovery(
auth_type,
server_config.get("token_exchange_endpoint"),
manual_token_url,
)
use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type or obo_needs_discovery)
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
manual_issuer,
is_discovery_auth_type,
@ -1234,17 +1281,12 @@ class MCPServerManager:
manual_token_url,
manual_registration_url,
)
should_discover = bool(server_url) and (
is_discovery_auth_type
or self._obo_needs_endpoint_discovery(
auth_type,
server_config.get("token_exchange_endpoint"),
manual_token_url,
)
should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
is_discovery_auth_type or obo_needs_discovery
)
if not should_discover:
mcp_oauth_metadata = None
elif manual_issuer is not None and is_discovery_auth_type:
elif use_issuer_anchor and manual_issuer is not None:
mcp_oauth_metadata = await self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server_url)
else:
mcp_oauth_metadata = await self._descovery_metadata(
@ -1640,7 +1682,7 @@ class MCPServerManager:
token_exchange_endpoint: Optional[str],
) -> Optional[MCPOAuthMetadata]:
has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes)
needs_discovery = bool(server_url) and (
needs_discovery = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
(is_discovery_auth_type and not has_all_upstream_oauth_fields)
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url)
)
@ -1759,13 +1801,17 @@ class MCPServerManager:
manual_token_url = _blank_to_none(mcp_server.token_url)
manual_registration_url = _blank_to_none(mcp_server.registration_url)
is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type)
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
)
token_exchange_endpoint = mcp_server.token_exchange_endpoint or (
credentials_dict.get("token_exchange_endpoint") if credentials_dict else None
)
use_issuer_anchor = _uses_issuer_anchor(
manual_issuer,
is_discovery_auth_type
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url),
)
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
)
gated_oauth_metadata = await self._resolve_table_oauth_metadata(
mcp_server=mcp_server,
auth_type=auth_type,
@ -1943,7 +1989,7 @@ class MCPServerManager:
family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on
the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path
calls ``update_server``) and on every post-write DB reload, so one failed re-discovery
serves 400 "authorization url is not set" from /authorize until a later rebuild succeeds.
serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds.
Only fills row fields that are currently empty, never persists origin-fallback guesses
(RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url``
because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a
@ -4705,6 +4751,61 @@ class MCPServerManager:
)
return oauth2_headers
async def resolve_openapi_upstream_auth(
self,
*,
mcp_server: MCPServer,
oauth2_headers: dict[str, str] | None,
raw_headers: dict[str, str] | None,
mcp_auth_header: str | dict[str, str] | None,
user_api_key_auth: UserAPIKeyAuth | None,
forwarded_headers: dict[str, str] | None,
) -> tuple[dict[str, str] | None, dict[str, str] | None]:
"""Resolve the gateway-owned upstream credential for a spec_path (OpenAPI) tool call.
OpenAPI tools egress through a plain httpx call assembled from ContextVars, never through
``_create_mcp_client``, so the v2 resolver graft there does not run for them and a resolved
credential (authorization_code's stored per-user token, client_credentials' minted M2M
token, token_exchange's exchanged token, passthrough's forwarded caller token) must be
materialized into headers here. Returns ``(resolved_auth_headers, forwarded_headers)``:
the resolved headers are authoritative over every other Authorization source (the same
rule ``_resolve_v2_auth`` applies on the MCPClient path) and ``forwarded_headers`` comes
back with any header the resolver claimed already dropped. Unmigrated (v1) servers resolve
through the stored-token lookup instead, and a missing per-user credential raises the same
discovery challenge the MCPClient path serves, rather than egressing unauthenticated.
The resolved headers carry only credentials the gateway itself resolved (a stored per-user
token, a minted or exchanged token). Caller-supplied ``oauth2_headers`` are never promoted
into them: on the v2 arm they feed only subject-token extraction (the designed RFC 8693
input), and on the v1 arm their presence disables the stored lookup entirely, so a
caller's gateway credential can never displace a per-server BYOK header or leak upstream
as the resolved credential.
"""
spec = to_server_spec(mcp_server)
if spec is None:
if oauth2_headers:
return None, forwarded_headers
stored_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, None, user_api_key_auth)
return stored_headers, forwarded_headers
subject_token: str | None = None
if isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
subject_token = self._extract_bearer_token(oauth2_headers, raw_headers)
elif isinstance(spec.config, PassthroughConfig):
inbound_token, forwarded_headers = _take_forwarded_authorization(forwarded_headers)
per_server_token = _passthrough_token_from_mcp_auth_header(mcp_auth_header)
subject_token = per_server_token if per_server_token is not None else inbound_token
resolved_auth, forwarded_headers = await self._resolve_v2_auth(
server=mcp_server,
spec=spec,
provider=self._cred_provider,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
extra_headers=forwarded_headers,
)
return await _materialize_auth_headers(resolved_auth), forwarded_headers
async def _gather_openapi_tool_tasks(
self,
tasks: list[Any],
@ -4796,6 +4897,7 @@ class MCPServerManager:
)
tasks.append(during_hook_task)
caller_oauth2_headers = oauth2_headers
oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, oauth2_headers, user_api_key_auth)
# For OpenAPI servers, call the tool handler directly instead of via MCP client
@ -4813,22 +4915,32 @@ class MCPServerManager:
auth_header_value = (
_format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None
)
forwarded_headers = _openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth)
resolved_auth_headers, forwarded_headers = await self.resolve_openapi_upstream_auth(
mcp_server=mcp_server,
oauth2_headers=caller_oauth2_headers,
raw_headers=raw_headers,
mcp_auth_header=mcp_auth_header,
user_api_key_auth=user_api_key_auth,
forwarded_headers=_openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth),
)
async def _call_openapi_via_handler():
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
_request_resolved_auth_headers,
)
auth_token = _request_auth_header.set(auth_header_value)
extra_token = _request_extra_headers.set(forwarded_headers)
resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers)
try:
async with self._limit_outbound_concurrency(mcp_server):
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
finally:
_request_auth_header.reset(auth_token)
_request_extra_headers.reset(extra_token)
_request_resolved_auth_headers.reset(resolved_token)
tasks.append(asyncio.create_task(_call_openapi_via_handler()))
else:

View file

@ -62,6 +62,14 @@ _request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = conte
"_request_extra_headers", default=None
)
# Per-request headers carrying the gateway-resolved upstream credential
# (stored per-user OAuth token, minted M2M token, exchanged OBO token).
# Set from MCPServerManager.resolve_openapi_upstream_auth; authoritative
# over every other Authorization source in _merge_openapi_tool_request_headers.
_request_resolved_auth_headers: contextvars.ContextVar[dict[str, str] | None] = contextvars.ContextVar(
"_request_resolved_auth_headers", default=None
)
def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str:
"""Ensure path params cannot introduce directory traversal."""
@ -294,10 +302,15 @@ def _merge_openapi_tool_request_headers(
"""Merge static closure headers with per-request ContextVar overrides.
Precedence (highest to lowest):
1. ``_request_auth_header`` — BYOK override of ``Authorization``
2. ``static_headers`` — operator-configured headers baked into the
1. ``_request_resolved_auth_headers`` — the gateway-resolved upstream
credential (stored per-user OAuth token, minted M2M token,
exchanged OBO token). The resolver is authoritative: a BYOK or
forwarded ``Authorization`` must not shadow it, mirroring
``_resolve_v2_auth`` on the MCPClient path
2. ``_request_auth_header`` — BYOK override of ``Authorization``
3. ``static_headers`` — operator-configured headers baked into the
tool closure at registration time
3. ``_request_extra_headers`` — per-request headers forwarded from
4. ``_request_extra_headers`` — per-request headers forwarded from
the MCP caller (allowlisted by ``MCPServer.extra_headers``)
This matches the existing MCP invariant in
@ -323,6 +336,12 @@ def _merge_openapi_tool_request_headers(
del effective_headers[existing]
effective_headers["Authorization"] = override_auth
resolved_auth_headers = _request_resolved_auth_headers.get() or {}
for name, value in resolved_auth_headers.items():
for existing in [k for k in effective_headers if k.lower() == name.lower()]:
del effective_headers[existing]
effective_headers[name] = value
return effective_headers

View file

@ -0,0 +1,214 @@
"""Store for the enterprise IdP identity assertion captured at SSO login (EMA).
The ``oauth2_id_jag`` egress arm needs the user's IdP ``id_token`` as its RFC 8693
``subject_token``. A front-door client holds an identity-only ``llm_session_`` bearer, not an
IdP assertion, so the assertion captured at the one SSO login is the only usable subject
source for it. This module owns both sides of that state: the SSO callback persists here
(write-through to the DB so a login on one pod is visible to every pod) and the resolver
seam reads back by ``user_id``. Retention is gated on an ``oauth2_id_jag`` server actually
being registered, so a gateway with no EMA upstream never stores bearer material.
The row is one encrypted payload per user, latest login wins. ``expires_at`` mirrors the
id_token ``exp`` claim and is judged by the reader, never enforced by deletion here: an
expired assertion with a refresh token is still renewable, and the DB row is the source of
truth, the same contract as the per-user OAuth credential store.
"""
from __future__ import annotations
import json
from datetime import datetime, timezone
from typing import TYPE_CHECKING
import jwt
from pydantic import BaseModel, ConfigDict, SecretStr, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
_ASSERTION_DECRYPT_LOG_KEY = "sso_identity_assertion"
_STR_ADAPTER: TypeAdapter[str] = TypeAdapter(str)
_MAYBE_STR_ADAPTER: TypeAdapter[str | None] = TypeAdapter(str | None)
class SSOIdentityAssertion(BaseModel):
"""The IdP material an EMA exchange needs: ``id_token`` is the RFC 8693 subject token,
``expires_at`` bounds its usefulness, and the refresh token renews it without re-login."""
model_config = ConfigDict(frozen=True)
id_token: SecretStr
refresh_token: SecretStr | None = None
issuer: str | None = None
expires_at: datetime | None = None
class _IdTokenClaims(BaseModel):
exp: float | None = None
iss: str | None = None
class _StoredAssertionPayload(BaseModel):
id_token: str
refresh_token: str | None = None
issuer: str | None = None
expires_at: datetime | None = None
def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIdentityAssertion | None:
"""The typed carrier built where the raw token response exists; ``None`` when the provider
sent no id_token or sent one that is not a decodable JWT, since neither is exchangeable
under EMA. Inputs are ``object`` because they come straight from the provider's untyped
token response; this is the one boundary that validates them. The token arrived over TLS
from the IdP's own token endpoint, so claims are read without signature verification,
matching how the SSO callback already decodes it for identity."""
raw_id_token = id_token if isinstance(id_token, str) and id_token else None
if raw_id_token is None:
return None
raw_refresh_token = refresh_token if isinstance(refresh_token, str) and refresh_token else None
try:
claims = _IdTokenClaims.model_validate(jwt.decode(raw_id_token, options={"verify_signature": False}))
expires_at = datetime.fromtimestamp(claims.exp, tz=timezone.utc) if claims.exp is not None else None
except Exception: # noqa: BLE001 # decode failure = not retainable; never raise into login
verbose_proxy_logger.warning(
"SSO id_token could not be decoded or its claims were unusable; not retaining it for EMA egress."
)
return None
return SSOIdentityAssertion(
id_token=SecretStr(raw_id_token),
refresh_token=SecretStr(raw_refresh_token) if raw_refresh_token else None,
issuer=claims.iss,
expires_at=expires_at,
)
async def ema_assertion_retention_enabled() -> bool:
"""Whether any MCP server uses ``oauth2_id_jag``, evaluated per login so the gateway only
retains bearer material while an EMA upstream exists to spend it on. Judged against the two
configuration authorities: the pod-local config declaration and the shared DB row. The
in-memory registry is deliberately not consulted in either direction; it is a per-process
snapshot of the DB state that can be stale both ways (a server added on another pod would
silently drop the write, one removed on another pod would keep retaining bearer material),
and a gate guarding a shared-DB write must judge against that storage's authority."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids import cycle
global_mcp_server_manager,
)
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global
config_servers = global_mcp_server_manager.config_mcp_servers.values()
if any(server.auth_type == MCPAuth.oauth2_id_jag for server in config_servers):
return True
if prisma_client is None:
return False
row = await prisma_client.db.litellm_mcpservertable.find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value})
return row is not None
async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAssertion) -> None:
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper # noqa: PLC0415 # runtime global
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
if prisma_client is None:
return
payload: dict[str, str] = {
"id_token": assertion.id_token.get_secret_value(),
**({"refresh_token": assertion.refresh_token.get_secret_value()} if assertion.refresh_token else {}),
**({"issuer": assertion.issuer} if assertion.issuer else {}),
**({"expires_at": assertion.expires_at.isoformat()} if assertion.expires_at else {}),
}
encoded = _STR_ADAPTER.validate_python(encrypt_value_helper(json.dumps(payload)))
await prisma_client.db.litellm_ssoidentityassertion.upsert(
where={"user_id": user_id},
data={
"create": {"user_id": user_id, "assertion_b64": encoded},
"update": {"assertion_b64": encoded},
},
)
async def fetch_sso_identity_assertion(user_id: str) -> SSOIdentityAssertion | None:
"""The stored assertion for ``user_id``, or ``None`` when absent, undecryptable (salt-key
rotation), or unparseable. Expiry is not judged here; the reader owns that policy."""
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper # noqa: PLC0415 # runtime global
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
if prisma_client is None:
return None
row = await prisma_client.db.litellm_ssoidentityassertion.find_unique(where={"user_id": user_id})
if row is None:
return None
raw = _MAYBE_STR_ADAPTER.validate_python(
decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
)
if raw is None:
return None
try:
payload = _StoredAssertionPayload.model_validate_json(raw)
except ValidationError:
verbose_proxy_logger.warning(
"Stored SSO identity assertion for user_id=%s could not be parsed; treating as absent.", user_id
)
return None
return SSOIdentityAssertion(
id_token=SecretStr(payload.id_token),
refresh_token=SecretStr(payload.refresh_token) if payload.refresh_token else None,
issuer=payload.issuer,
expires_at=payload.expires_at,
)
async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
"""Re-encrypt every stored assertion under ``new_master_key`` during a salt-key rotation,
mirroring the sibling per-user credential tables; an unreadable row is skipped so one
corrupt row does not abort the rotation. Rows are decrypted one at a time inside the loop
so the whole table's plaintext is never held in memory at once."""
from prisma.models import LiteLLM_SSOIdentityAssertion as AssertionRow # noqa: PLC0415 # generated at runtime
from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global
decrypt_value_helper,
encrypt_value_helper,
)
async def _rotate_row(row: AssertionRow) -> bool:
plaintext = _MAYBE_STR_ADAPTER.validate_python(
decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
)
if plaintext is None:
verbose_proxy_logger.warning(
"rotate_sso_identity_assertions_master_key: could not decrypt assertion for user_id=%s, skipping",
row.user_id,
)
return False
re_encrypted = _STR_ADAPTER.validate_python(encrypt_value_helper(plaintext, new_encryption_key=new_master_key))
await prisma_client.db.litellm_ssoidentityassertion.update(
where={"user_id": row.user_id},
data={"assertion_b64": re_encrypted},
)
return True
rows = await prisma_client.db.litellm_ssoidentityassertion.find_many()
outcomes = [await _rotate_row(row) for row in rows]
verbose_proxy_logger.info(
"rotate_sso_identity_assertions_master_key: rotated %d row(s), skipped %d",
sum(outcomes),
len(outcomes) - sum(outcomes),
)
async def retain_sso_identity_assertion_for_ema(user_id: str, assertion: SSOIdentityAssertion | None) -> None:
"""The SSO-callback hook: a no-op unless there is material AND an EMA server is registered.
A store failure is logged and swallowed because the login itself must not fail on an
egress-side write; the cost of a miss is a 401 challenge at the EMA upstream, not a lockout."""
if assertion is None:
return
try:
if not await ema_assertion_retention_enabled():
return
await persist_sso_identity_assertion(user_id, assertion)
except Exception as exc: # noqa: BLE001 # the login itself must not fail on an egress-side write
verbose_proxy_logger.warning(
"Failed to persist the SSO identity assertion for EMA egress (user_id=%s): %s", user_id, exc
)

View file

@ -376,6 +376,7 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
_request_resolved_auth_headers,
)
from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport
from litellm.proxy._experimental.mcp_server.tool_registry import (
@ -2785,13 +2786,29 @@ if MCP_AVAILABLE:
forwarded_headers = {}
forwarded_headers[header_name] = value
resolved_auth_headers: dict[str, str] | None = None
if mcp_server:
(
resolved_auth_headers,
forwarded_headers,
) = await global_mcp_server_manager.resolve_openapi_upstream_auth(
mcp_server=mcp_server,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
mcp_auth_header=mcp_auth_header,
user_api_key_auth=user_api_key_auth,
forwarded_headers=forwarded_headers,
)
_auth_token = _request_auth_header.set(auth_header_value)
_extra_token = _request_extra_headers.set(forwarded_headers)
_resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers)
try:
local_content = await _handle_local_mcp_tool(name, arguments)
finally:
_request_auth_header.reset(_auth_token)
_request_extra_headers.reset(_extra_token)
_request_resolved_auth_headers.reset(_resolved_token)
response = CallToolResult(content=cast(Any, local_content), isError=False)
# Try managed MCP server tool (pass the full prefixed name)

View file

@ -2310,6 +2310,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="max response size in MB, if a response is larger than this size it will be rejected",
)
proxy_config_reload_interval_seconds: int = Field(
30,
gt=0,
description="how often (in seconds) each pod reloads config-in-DB objects (models, credentials, guardrails, etc.) when store_model_in_db is enabled; lower values speed up multi-pod convergence at the cost of more DB load. Applied on proxy startup",
)
cancel_on_disconnect: Optional[bool] = Field(
None,
description="cancel the in-flight upstream LLM request (non-streaming) when the client disconnects, freeing backend capacity (e.g. a vLLM GPU slot); the request is logged as a 499 failure",

View file

@ -7,23 +7,46 @@ the base; specific fields are replaced so all traffic flows through the proxy
and uses LiteLLM auth.
"""
import re
from copy import deepcopy
from typing import Any, Dict, List, Mapping
from typing import Any, Dict, List, Literal, Mapping
SupportedA2AVersion = Literal["0.3", "1.0"]
# Protocol versions LiteLLM can serve to A2A clients. The admin pins one per agent;
# responses are normalized to it regardless of the upstream agent's own version.
SUPPORTED_A2A_PROTOCOL_VERSIONS = ("0.3", "1.0")
SUPPORTED_A2A_PROTOCOL_VERSIONS: tuple[SupportedA2AVersion, ...] = ("0.3", "1.0")
# Default served version when the agent card does not pin one.
LITELLM_A2A_PROTOCOL_VERSION = "1.0"
_PROTOCOL_VERSION_PATTERN = re.compile(
r"^(\d+\.\d+)(?:\.\d+(?:-[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?(?:\+[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?)?$"
)
def normalize_protocol_version(version: object) -> SupportedA2AVersion | None:
"""Map a raw ``protocolVersion`` value to the supported canonical major.minor version.
Accepts the bare major.minor convention of the 1.0 spec (``"0.3"``, ``"1.0"``) and the
full semver forms older SDKs emit (``"0.3.0"``, ``"1.0.1"``, including prerelease and
build suffixes like ``"0.3.0-rc1"``). Malformed strings, versions outside the
supported set, and non-strings yield ``None``.
"""
if not isinstance(version, str):
return None
match = _PROTOCOL_VERSION_PATTERN.match(version)
if match is None:
return None
major_minor = match.group(1)
return next((supported for supported in SUPPORTED_A2A_PROTOCOL_VERSIONS if supported == major_minor), None)
def resolve_served_protocol_version(card: Mapping[str, Any] | None) -> str:
"""Return the validated protocol version an agent card pins, else the default."""
version = card.get("protocolVersion") if card else None
if version in SUPPORTED_A2A_PROTOCOL_VERSIONS:
return version
return LITELLM_A2A_PROTOCOL_VERSION
normalized = normalize_protocol_version(card.get("protocolVersion") if card else None)
return normalized if normalized is not None else LITELLM_A2A_PROTOCOL_VERSION
# Security scheme exposed by the LiteLLM-fronted agent card. Always replaces

View file

@ -30,6 +30,7 @@ from typing import Callable, Literal, Union
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.proxy.a2a.agent_card import normalize_protocol_version
A2AVersion = Literal["0.3", "1.0"]
RequestId = Union[str, int, None]
@ -103,16 +104,14 @@ def normalize_request_params(params: JsonDict, served: A2AVersion, *, method: st
def _detect_card_version(card: JsonDict) -> A2AVersion:
"""Infer the wire version of an agent card dict.
``protocolVersion`` is the authoritative indicator; fall back to presence of
``supportedInterfaces`` (a 1.0-only field) only when the explicit field is absent.
Cards that set ``protocolVersion: "0.3"`` or carry neither signal are treated as 0.3.
``protocolVersion`` is the authoritative indicator; semver values normalize to
their major.minor (``"0.3.0"`` -> ``"0.3"``). Fall back to presence of
``supportedInterfaces`` (a 1.0-only field) only when the explicit field is
absent or unrecognized; cards carrying neither signal are treated as 0.3.
"""
pv = card.get("protocolVersion")
if pv == "1.0":
return "1.0"
if pv == "0.3":
return "0.3"
# No protocolVersion field: use structural heuristic.
normalized = normalize_protocol_version(card.get("protocolVersion"))
if normalized is not None:
return normalized
return "1.0" if "supportedInterfaces" in card else "0.3"

View file

@ -23,6 +23,7 @@ from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKey
from litellm.proxy.a2a.agent_card import (
SUPPORTED_A2A_PROTOCOL_VERSIONS,
merge_agent_card,
normalize_protocol_version,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
@ -51,7 +52,7 @@ def _proxy_base_url(http_request: Request) -> str:
def _validate_protocol_version(upstream_card: Mapping[str, Any] | None) -> None:
"""Reject an agent card pinning an unsupported A2A protocol version."""
version = upstream_card.get("protocolVersion") if upstream_card else None
if version is not None and version not in SUPPORTED_A2A_PROTOCOL_VERSIONS:
if version is not None and normalize_protocol_version(version) is None:
raise HTTPException(
status_code=400,
detail=(

View file

@ -1051,7 +1051,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None:
return
if getattr(request.state, "parent_otel_span", None) is not None:
return
start_time = datetime.now()
start_time = datetime.now(timezone.utc)
try:
request.state.litellm_received_at = start_time
except Exception:
@ -1101,7 +1101,7 @@ async def _user_api_key_auth_builder(
# Prefer the receive-instant stamped by the early helper in
# user_api_key_auth (before body parse) — overwriting it would shorten
# the preprocessing-duration measurement by the body-parse window.
start_time = getattr(request.state, "litellm_received_at", None) or datetime.now()
start_time = getattr(request.state, "litellm_received_at", None) or datetime.now(timezone.utc)
try:
request.state.litellm_received_at = start_time
except Exception:
@ -2660,7 +2660,7 @@ async def _return_user_api_key_auth_obj(
start_time: datetime,
user_role: Optional[LitellmUserRoles] = None,
) -> UserAPIKeyAuth:
end_time = datetime.now()
end_time = datetime.now(timezone.utc)
asyncio.create_task(
user_api_key_service_logger_obj.async_service_success_hook(
@ -2749,9 +2749,10 @@ def _update_key_budget_with_temp_budget_increase(
) -> UserAPIKeyAuth:
if valid_token.max_budget is None:
return valid_token
temp_budget_increase = _get_temp_budget_increase(valid_token) or 0.0
valid_token.max_budget = valid_token.max_budget + temp_budget_increase
return valid_token
temp_budget_increase = _get_temp_budget_increase(valid_token)
if not temp_budget_increase:
return valid_token
return valid_token.model_copy(update={"max_budget": valid_token.max_budget + temp_budget_increase})
async def _lookup_end_user_and_apply_budget(

View file

@ -13,6 +13,11 @@ from litellm.proxy._types import (
LiteLLM_UserTable,
LiteLLM_VerificationToken,
)
from litellm.proxy.common_utils.timezone_utils import (
BudgetResetSettings,
compute_budget_reset_at,
get_budget_reset_settings,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.table_repositories import (
@ -32,9 +37,15 @@ class ResetBudgetJob:
Resets the budget for all the keys, users, and teams that need it
"""
def __init__(self, proxy_logging_obj: ProxyLogging, prisma_client: PrismaClient):
def __init__(
self,
proxy_logging_obj: ProxyLogging,
prisma_client: PrismaClient,
reset_settings: BudgetResetSettings | None = None,
):
self.proxy_logging_obj: ProxyLogging = proxy_logging_obj
self.prisma_client: PrismaClient = prisma_client
self.reset_settings: BudgetResetSettings = reset_settings or get_budget_reset_settings()
async def reset_budget(
self,
@ -237,7 +248,7 @@ class ResetBudgetJob:
if budgets_to_reset is not None and len(budgets_to_reset) > 0:
for budget in budgets_to_reset:
budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now)
budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now, self.reset_settings)
await self.prisma_client.update_data(
query_type="update_many",
@ -442,7 +453,11 @@ class ResetBudgetJob:
if keys_to_reset is not None and len(keys_to_reset) > 0:
for key in keys_to_reset:
try:
updated_key = await ResetBudgetJob._reset_budget_for_key(key=key, current_time=now)
updated_key = await ResetBudgetJob._reset_budget_for_key(
key=key,
current_time=now,
reset_settings=self.reset_settings,
)
if updated_key is not None:
updated_keys.append(updated_key)
else:
@ -513,7 +528,11 @@ class ResetBudgetJob:
if users_to_reset is not None and len(users_to_reset) > 0:
for user in users_to_reset:
try:
updated_user = await ResetBudgetJob._reset_budget_for_user(user=user, current_time=now)
updated_user = await ResetBudgetJob._reset_budget_for_user(
user=user,
current_time=now,
reset_settings=self.reset_settings,
)
if updated_user is not None:
updated_users.append(updated_user)
else:
@ -588,7 +607,11 @@ class ResetBudgetJob:
if teams_to_reset is not None and len(teams_to_reset) > 0:
for team in teams_to_reset:
try:
updated_team = await ResetBudgetJob._reset_budget_for_team(team=team, current_time=now)
updated_team = await ResetBudgetJob._reset_budget_for_team(
team=team,
current_time=now,
reset_settings=self.reset_settings,
)
if updated_team is not None:
updated_teams.append(updated_team)
else:
@ -655,10 +678,9 @@ class ResetBudgetJob:
counter_key: str,
spend_counter_cache: Any,
now: datetime,
reset_settings: BudgetResetSettings,
) -> bool:
"""Reset a single budget window if expired. Returns True if the window was reset."""
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
reset_at_str = window.get("reset_at")
if not reset_at_str:
return False
@ -671,7 +693,9 @@ class ResetBudgetJob:
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=0.0)
except Exception as redis_err:
verbose_proxy_logger.warning("Failed to reset Redis counter %s: %s", counter_key, redis_err)
window["reset_at"] = get_budget_reset_time(budget_duration=window["budget_duration"]).isoformat()
window["reset_at"] = compute_budget_reset_at(
budget_duration=window["budget_duration"], settings=reset_settings
).isoformat()
return True
async def reset_budget_windows(self) -> None:
@ -703,7 +727,13 @@ class ResetBudgetJob:
changed = False
for window in windows:
counter_key = f"spend:key:{row['token']}:window:{window['budget_duration']}"
if await ResetBudgetJob._reset_expired_window(window, counter_key, spend_counter_cache, now):
if await ResetBudgetJob._reset_expired_window(
window,
counter_key,
spend_counter_cache,
now,
self.reset_settings,
):
changed = True
if changed:
await VerificationTokenRepository(self.prisma_client).table.update(
@ -726,7 +756,13 @@ class ResetBudgetJob:
changed = False
for window in windows:
counter_key = f"spend:team:{row['team_id']}:window:{window['budget_duration']}"
if await ResetBudgetJob._reset_expired_window(window, counter_key, spend_counter_cache, now):
if await ResetBudgetJob._reset_expired_window(
window,
counter_key,
spend_counter_cache,
now,
self.reset_settings,
):
changed = True
if changed:
await TeamRepository(self.prisma_client).table.update(
@ -741,6 +777,7 @@ class ResetBudgetJob:
item: Union[LiteLLM_TeamTable, LiteLLM_UserTable, LiteLLM_VerificationToken],
current_time: datetime,
item_type: Literal["key", "team", "user"],
reset_settings: BudgetResetSettings,
):
"""
In-place, updates spend=0, and sets budget_reset_at to current_time + budget_duration
@ -755,24 +792,40 @@ class ResetBudgetJob:
try:
item.spend = 0.0
if hasattr(item, "budget_duration") and item.budget_duration is not None:
from litellm.proxy.common_utils.timezone_utils import (
get_budget_reset_time,
item.budget_reset_at = compute_budget_reset_at(
budget_duration=item.budget_duration, settings=reset_settings
)
item.budget_reset_at = get_budget_reset_time(budget_duration=item.budget_duration)
return item
except Exception as e:
verbose_proxy_logger.exception("Error resetting budget for %s: %s. Item: %s", item_type, e, item)
raise e
@staticmethod
async def _reset_budget_for_team(team: LiteLLM_TeamTable, current_time: datetime) -> Optional[LiteLLM_TeamTable]:
await ResetBudgetJob._reset_budget_common(item=team, current_time=current_time, item_type="team")
async def _reset_budget_for_team(
team: LiteLLM_TeamTable,
current_time: datetime,
reset_settings: BudgetResetSettings,
) -> LiteLLM_TeamTable | None:
await ResetBudgetJob._reset_budget_common(
item=team,
current_time=current_time,
item_type="team",
reset_settings=reset_settings,
)
return team
@staticmethod
async def _reset_budget_for_user(user: LiteLLM_UserTable, current_time: datetime) -> Optional[LiteLLM_UserTable]:
await ResetBudgetJob._reset_budget_common(item=user, current_time=current_time, item_type="user")
async def _reset_budget_for_user(
user: LiteLLM_UserTable,
current_time: datetime,
reset_settings: BudgetResetSettings,
) -> LiteLLM_UserTable | None:
await ResetBudgetJob._reset_budget_common(
item=user,
current_time=current_time,
item_type="user",
reset_settings=reset_settings,
)
return user
@staticmethod
@ -788,15 +841,15 @@ class ResetBudgetJob:
@staticmethod
async def _reset_budget_reset_at_date(
budget: LiteLLM_BudgetTableFull, current_time: datetime
budget: LiteLLM_BudgetTableFull,
current_time: datetime,
reset_settings: BudgetResetSettings,
) -> LiteLLM_BudgetTableFull:
try:
if budget.budget_duration is not None:
from litellm.proxy.common_utils.timezone_utils import (
get_budget_reset_time,
budget.budget_reset_at = compute_budget_reset_at(
budget_duration=budget.budget_duration, settings=reset_settings
)
budget.budget_reset_at = get_budget_reset_time(budget_duration=budget.budget_duration)
except Exception as e:
verbose_proxy_logger.exception("Error resetting budget_reset_at for budget: %s. Item: %s", e, budget)
raise e
@ -804,7 +857,14 @@ class ResetBudgetJob:
@staticmethod
async def _reset_budget_for_key(
key: LiteLLM_VerificationToken, current_time: datetime
) -> Optional[LiteLLM_VerificationToken]:
await ResetBudgetJob._reset_budget_common(item=key, current_time=current_time, item_type="key")
key: LiteLLM_VerificationToken,
current_time: datetime,
reset_settings: BudgetResetSettings,
) -> LiteLLM_VerificationToken | None:
await ResetBudgetJob._reset_budget_common(
item=key,
current_time=current_time,
item_type="key",
reset_settings=reset_settings,
)
return key

View file

@ -1,10 +1,47 @@
from datetime import datetime, timezone
from datetime import datetime, time, timezone
from pydantic import BaseModel, ConfigDict
import litellm
from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time
def get_budget_reset_timezone():
class BudgetResetSettings(BaseModel):
"""Immutable, validated settings that govern when budgets reset.
Parsed once from `litellm_settings` and injected into consumers (the reset
job, management endpoints) so reset times never depend on reaching into
module-level globals at call time.
"""
model_config = ConfigDict(frozen=True)
timezone: str = "UTC"
reset_time_of_day: time = time(0, 0)
def parse_budget_reset_time(raw: object) -> time:
"""Parse a `budget_reset_time` config value (e.g. "12:00") into a `time`.
Falls back to midnight when unset; raises a clear error on a malformed value
so a bad config fails loudly at startup instead of silently resetting at midnight.
"""
if raw is None or raw == "":
return time(0, 0)
if not isinstance(raw, str):
raise ValueError(f"Invalid budget_reset_time {raw!r}; must be a quoted 24-hour 'HH:MM' string, e.g. \"12:00\"")
for fmt in ("%H:%M", "%H:%M:%S"):
try:
parsed = datetime.strptime(raw, fmt)
return time(hour=parsed.hour, minute=parsed.minute, second=parsed.second)
except ValueError:
continue
raise ValueError(
f"Invalid budget_reset_time {raw!r}; expected a 24-hour 'HH:MM' or 'HH:MM:SS' string, e.g. \"12:00\""
)
def get_budget_reset_timezone() -> str:
"""
Get the budget reset timezone from litellm_settings.
Falls back to UTC if not specified.
@ -15,15 +52,29 @@ def get_budget_reset_timezone():
return getattr(litellm, "timezone", None) or "UTC"
def get_budget_reset_time(budget_duration: str) -> datetime:
"""
Get the budget reset time based on the configured timezone.
Falls back to UTC if not specified.
"""
def get_budget_reset_settings() -> BudgetResetSettings:
"""Build validated reset settings from litellm_settings. Raises on a malformed
`budget_reset_time`, which lets the proxy fail fast at startup."""
return BudgetResetSettings(
timezone=get_budget_reset_timezone(),
reset_time_of_day=parse_budget_reset_time(getattr(litellm, "budget_reset_time", None)),
)
reset_at = get_next_standardized_reset_time(
def compute_budget_reset_at(budget_duration: str, settings: BudgetResetSettings) -> datetime:
"""Compute the next reset time for a budget duration using injected settings."""
return get_next_standardized_reset_time(
duration=budget_duration,
current_time=datetime.now(timezone.utc),
timezone_str=get_budget_reset_timezone(),
timezone_str=settings.timezone,
reset_time_of_day=settings.reset_time_of_day,
)
return reset_at
def get_budget_reset_time(budget_duration: str) -> datetime:
"""Get the budget reset time using the globally-configured timezone and reset time.
Thin wrapper over `compute_budget_reset_at` for callers that don't yet receive
`BudgetResetSettings` by injection (creation/update endpoints, startup backfill).
"""
return compute_budget_reset_at(budget_duration, get_budget_reset_settings())

View file

@ -42,6 +42,9 @@ from litellm.proxy._experimental.mcp_server.db import (
rotate_mcp_user_credentials_master_key,
rotate_mcp_user_env_vars_master_key,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
rotate_sso_identity_assertions_master_key,
)
from litellm.proxy._types import *
from litellm.proxy._types import LiteLLM_VerificationToken, hash_token
from litellm.proxy.auth.auth_checks import (
@ -4341,6 +4344,15 @@ async def _rotate_master_key(
except Exception as e:
verbose_proxy_logger.warning("Failed to rotate MCP user env vars: %s", str(e))
# 4d. process SSO identity assertion table (EMA subject tokens)
try:
await rotate_sso_identity_assertions_master_key(
prisma_client=prisma_client,
new_master_key=new_master_key,
)
except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation
verbose_proxy_logger.warning("Failed to rotate SSO identity assertions: %s", str(e))
# 5. process credentials table
try:
credentials = await CredentialsRepository(prisma_client).table.find_many()

View file

@ -62,6 +62,11 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
SSOIdentityAssertion,
assertion_from_sso_login,
retain_sso_identity_assertion_for_ema,
)
from litellm.proxy._types import (
CommonProxyErrors,
LiteLLM_UserTable,
@ -1311,12 +1316,15 @@ async def get_generic_sso_response(
sso_jwt_handler: Optional[JWTHandler], # sso specific jwt handler - used for restricted sso group access control
generic_client_id: str,
redirect_url: str,
) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload)
) -> tuple[
Union[OpenID, dict], dict | None, dict | None, SSOIdentityAssertion | None
]: # (result, received_response, access_token_payload, sso_assertion)
# make generic sso provider
from fastapi_sso.sso.base import DiscoveryDocument
from fastapi_sso.sso.generic import create_provider
received_response: Optional[dict] = None
sso_assertion: SSOIdentityAssertion | None = None
# Setup environment variables
(
@ -1450,6 +1458,9 @@ async def get_generic_sso_response(
# Assign directly rather than relying on nonlocal mutation so that Pyright
# can track that received_response is non-None from this point on.
received_response = {k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS}
sso_assertion = assertion_from_sso_login(
combined_response.get("id_token"), combined_response.get("refresh_token")
)
# In the PKCE path verify_and_process is skipped, so generic_sso.access_token
# is never set. Read the token directly from the exchange response instead so
# process_sso_jwt_access_token can extract JWT-embedded roles/teams.
@ -1461,6 +1472,7 @@ async def get_generic_sso_response(
headers=additional_generic_sso_headers_dict,
)
access_token_str = generic_sso.access_token
sso_assertion = assertion_from_sso_login(generic_sso.id_token, generic_sso.refresh_token)
access_token_payload = process_sso_jwt_access_token(
access_token_str, sso_jwt_handler, result, role_mappings=role_mappings
@ -1480,7 +1492,7 @@ async def get_generic_sso_response(
additional_generic_sso_headers_dict,
)
verbose_proxy_logger.debug("generic result: %s", result)
return result or {}, received_response, access_token_payload
return result or {}, received_response, access_token_payload, sso_assertion
async def create_team_member_add_task(team_id, user_info):
@ -1812,6 +1824,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
generic_client_id = os.getenv("GENERIC_CLIENT_ID", None)
received_response: Optional[dict] = None
access_token_payload: Optional[dict] = None
sso_assertion: SSOIdentityAssertion | None = None
# get url from request
if master_key is None:
raise ProxyException(
@ -1842,6 +1855,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
result,
received_response,
access_token_payload,
sso_assertion,
) = await get_generic_sso_response(
request=request,
jwt_handler=jwt_handler,
@ -1869,6 +1883,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
prefill_user_code=prefill_user_code,
result=result,
received_response=received_response,
sso_assertion=sso_assertion,
)
# Control-plane cross-origin: read return_to from cookie.
@ -1884,6 +1899,7 @@ async def auth_callback(request: Request, state: Optional[str] = None):
access_token_payload=access_token_payload,
jwt_handler=jwt_handler,
return_to=cp_return_to,
sso_assertion=sso_assertion,
)
@ -1943,6 +1959,7 @@ async def _complete_cli_sso_callback_session(
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
prefill_user_code: str | None = None,
sso_assertion: SSOIdentityAssertion | None = None,
):
from fastapi.responses import HTMLResponse
@ -1962,6 +1979,8 @@ async def _complete_cli_sso_callback_session(
if not user_info.user_id:
raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO")
await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion)
teams: List[str] = []
if hasattr(user_info, "teams") and user_info.teams:
teams = user_info.teams if isinstance(user_info.teams, list) else []
@ -2012,6 +2031,7 @@ async def cli_sso_callback(
result: Optional[Union[OpenID, dict]] = None,
received_response: Optional[dict] = None,
prefill_user_code: str | None = None,
sso_assertion: SSOIdentityAssertion | None = None,
):
"""CLI SSO callback - stores session info for JWT generation on polling"""
verbose_proxy_logger.info("CLI SSO callback")
@ -2065,6 +2085,7 @@ async def cli_sso_callback(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
prefill_user_code=prefill_user_code,
sso_assertion=sso_assertion,
)
except ProxyException:
raise
@ -3018,6 +3039,7 @@ class SSOAuthenticationHandler:
access_token_payload: Optional[dict] = None,
jwt_handler: Optional[JWTHandler] = None,
return_to: Optional[str] = None,
sso_assertion: SSOIdentityAssertion | None = None,
) -> RedirectResponse:
import jwt
@ -3148,6 +3170,9 @@ class SSOAuthenticationHandler:
},
)
if isinstance(user_id, str) and user_id:
await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
disabled_non_admin_personal_key_creation = get_disabled_non_admin_personal_key_creation()
litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/")
@ -4241,6 +4266,7 @@ async def debug_sso_callback(request: Request):
result,
received_response,
access_token_payload,
_sso_assertion,
) = await get_generic_sso_response(
request=request,
jwt_handler=jwt_handler,

View file

@ -236,6 +236,7 @@ from litellm.constants import (
PROXY_BATCH_WRITE_AT,
PROXY_BUDGET_RESCHEDULER_MAX_TIME,
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS,
)
from litellm.exceptions import RejectedRequestError
from litellm.integrations.custom_guardrail import ModifyResponseException
@ -319,7 +320,10 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
from litellm.proxy.common_utils.proxy_state import ProxyState
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
from litellm.proxy.common_utils.swagger_utils import ERROR_RESPONSES
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.common_utils.timezone_utils import (
get_budget_reset_settings,
get_budget_reset_time,
)
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
get_management_object_ttl,
@ -1998,6 +2002,7 @@ proxy_budget_rescheduler_min_time = PROXY_BUDGET_RESCHEDULER_MIN_TIME
proxy_budget_rescheduler_max_time = PROXY_BUDGET_RESCHEDULER_MAX_TIME
proxy_batch_polling_interval = PROXY_BATCH_POLLING_INTERVAL
proxy_batch_write_at = PROXY_BATCH_WRITE_AT
proxy_config_reload_interval_seconds = PROXY_CONFIG_RELOAD_INTERVAL_SECONDS
litellm_master_key_hash = None
disable_spend_logs = False
jwt_handler = JWTHandler()
@ -3879,7 +3884,7 @@ class ProxyConfig:
del config["include"]
return config
async def save_config(self, new_config: dict):
async def save_config(self, new_config: dict, include_env_vars: bool = False):
global prisma_client, general_settings, user_config_file_path, store_model_in_db
# Load existing config
## DB - writes valid config to db
@ -3896,6 +3901,17 @@ class ProxyConfig:
# Make a copy to avoid mutating the original config
config_to_save = new_config.copy()
# environment_variables are persisted to the DB only when a caller
# explicitly opts in. Most callers reach save_config after
# get_config() merged YAML + OS env into new_config (with
# os.environ/ placeholders already resolved to plaintext), so
# persisting them here would snapshot file/container env vars into
# a config row that then shadows those sources on every restart.
# The dedicated /config/update path writes env vars directly, so
# no current caller needs include_env_vars=True.
if not include_env_vars:
config_to_save.pop("environment_variables", None)
# SECURITY: Always encrypt environment_variables before DB write.
# _encrypt_env_variables_for_db is idempotent — a caller that
# already encrypted the values (or re-submitted ciphertext read
@ -3913,6 +3929,38 @@ class ProxyConfig:
with open(f"{user_config_file_path}", "w") as config_file:
yaml.dump(new_config, config_file, default_flow_style=False)
async def save_environment_variables(self, updates: dict[str, str | None]) -> None:
"""Persist specific environment variables to the DB config row.
Each key in ``updates`` is written to the ``environment_variables``
config row; a ``None`` value deletes that key. Env vars the caller does
not name are preserved, so a caller that owns a couple of keys can
update just those without snapshotting unrelated (YAML/OS-sourced)
values the way a full ``save_config`` write would. No-op when config is
not DB-backed.
"""
global prisma_client, general_settings, store_model_in_db
if prisma_client is None or not (general_settings.get("store_model_in_db", False) is True or store_model_in_db):
return
row = await ConfigRepository(prisma_client).table.find_first(where={"param_name": "environment_variables"})
existing: dict = dict(row.param_value) if row is not None and row.param_value is not None else {}
to_set = {k: v for k, v in updates.items() if v is not None}
encrypted = self._encrypt_env_variables_for_db(environment_variables=to_set) if to_set else {}
deleted_keys = {k for k, v in updates.items() if v is None}
merged = {**{k: v for k, v in existing.items() if k not in deleted_keys}, **encrypted}
serialized = json.dumps(merged)
await ConfigRepository(prisma_client).table.upsert(
where={"param_name": "environment_variables"},
data={
"create": {"param_name": "environment_variables", "param_value": serialized},
"update": {"param_value": serialized},
},
)
await invalidate_config_param("environment_variables")
def _check_for_os_environ_vars(
self, config: dict, depth: int = 0, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH
) -> dict:
@ -4291,6 +4339,7 @@ class ProxyConfig:
open_telemetry_logger, \
health_check_details, \
proxy_batch_polling_interval, \
proxy_config_reload_interval_seconds, \
config_passthrough_endpoints
config: dict = await self.get_config(config_file_path=config_file_path)
@ -4597,6 +4646,13 @@ class ProxyConfig:
litellm.json_logs = True
litellm._turn_on_json()
verbose_proxy_logger.debug(f"{blue_color_code} Enabled JSON logging via config{reset_color_code}")
elif key == "budget_reset_time":
from litellm.proxy.common_utils.timezone_utils import (
parse_budget_reset_time,
)
parse_budget_reset_time(value)
setattr(litellm, key, value)
else:
verbose_proxy_logger.debug(
f"{blue_color_code} setting litellm.{key}={_redact_general_setting_value(key, value, is_full_admin=False)}{reset_color_code}"
@ -4773,6 +4829,10 @@ class ProxyConfig:
)
## BATCH WRITER ##
proxy_batch_write_at = general_settings.get("proxy_batch_write_at", proxy_batch_write_at)
## DB CONFIG RELOAD INTERVAL ##
proxy_config_reload_interval_seconds = general_settings.get(
"proxy_config_reload_interval_seconds", proxy_config_reload_interval_seconds
)
## DISABLE SPEND LOGS ## - gives a perf improvement
disable_spend_logs = general_settings.get("disable_spend_logs", disable_spend_logs)
### BACKGROUND HEALTH CHECKS ###
@ -7868,6 +7928,7 @@ class ProxyStartupEvent:
budget_reset_job = ResetBudgetJob(
proxy_logging_obj=proxy_logging_obj,
prisma_client=prisma_client,
reset_settings=get_budget_reset_settings(),
)
scheduler.add_job(
@ -7944,12 +8005,20 @@ class ProxyStartupEvent:
verbose_proxy_logger.debug("Failed to check DB for store_model_in_db: %s", str(e))
if store_model_in_db is True:
config_reload_interval_seconds = proxy_config_reload_interval_seconds
if not isinstance(config_reload_interval_seconds, int) or config_reload_interval_seconds <= 0:
verbose_proxy_logger.warning(
"proxy_config_reload_interval_seconds=%s must be a positive integer; falling back to 30s",
config_reload_interval_seconds,
)
config_reload_interval_seconds = 30
# MEMORY LEAK FIX: Increase interval from 10s to 30s minimum
# Frequent polling was causing excessive memory allocations
scheduler.add_job(
proxy_config.add_deployment,
"interval",
seconds=30, # increased from 10s to reduce memory pressure
seconds=config_reload_interval_seconds,
# REMOVED jitter parameter - major cause of memory leak
args=[prisma_client, proxy_logging_obj],
id="add_deployment_job",
@ -7964,7 +8033,7 @@ class ProxyStartupEvent:
scheduler.add_job(
proxy_config.get_credentials,
"interval",
seconds=30, # increased from 10s to reduce memory pressure
seconds=config_reload_interval_seconds,
# REMOVED jitter parameter - major cause of memory leak
args=[prisma_client],
id="get_credentials_job",
@ -14998,6 +15067,7 @@ async def get_config_list(
"global_max_parallel_requests": {"type": "Integer"},
"max_request_size_mb": {"type": "Integer"},
"max_response_size_mb": {"type": "Integer"},
"proxy_config_reload_interval_seconds": {"type": "Integer"},
"pass_through_endpoints": {"type": "PydanticModel"},
"store_model_in_db": {"type": "Boolean"},
"store_prompts_in_spend_logs": {"type": "Boolean"},

View file

@ -406,6 +406,15 @@ model LiteLLM_MCPServerOAuthClient {
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// The enterprise IdP identity assertion captured at SSO login, one row per user.
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
model LiteLLM_SSOIdentityAssertion {
user_id String @id
assertion_b64 String
created_at DateTime @default(now()) @map("created_at")
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id

View file

@ -1041,13 +1041,6 @@ async def update_ui_theme_settings(
config = await proxy_config.get_config()
before_theme = config.get("litellm_settings", {}).get("ui_theme_config")
# Update config with UI theme settings
if "general_settings" not in config:
config["general_settings"] = {}
if "environment_variables" not in config:
config["environment_variables"] = {}
# Convert theme config to dict
theme_data = theme_config.model_dump(exclude_none=True)
@ -1056,55 +1049,29 @@ async def update_ui_theme_settings(
config["litellm_settings"] = {}
config["litellm_settings"]["ui_theme_config"] = theme_data
# Update UI_LOGO_PATH environment variable if logo_url is provided
# If logo_url is empty string, None, or null, remove the environment variable to use default
logo_url = theme_data.get("logo_url")
verbose_proxy_logger.debug(f"Updating logo_url: {logo_url}")
# UI_LOGO_PATH and LITELLM_FAVICON_URL are the only environment variables
# this endpoint owns. A non-empty value sets the var; an empty or missing
# one clears it back to the default. Apply to the live process immediately,
# then persist only these two keys so an unrelated env var (a YAML/OS value
# merged in by get_config) is never snapshotted into the DB.
def _clean(url: str | None) -> str | None:
return url if url is not None and url.strip() else None
if (
logo_url and isinstance(logo_url, str) and logo_url.strip()
): # Check if logo_url exists and is not empty/whitespace
config["environment_variables"]["UI_LOGO_PATH"] = logo_url
os.environ["UI_LOGO_PATH"] = logo_url
verbose_proxy_logger.debug(f"Set UI_LOGO_PATH to: {logo_url}")
else:
# Remove the environment variable to restore default logo
if "UI_LOGO_PATH" in config.get("environment_variables", {}):
del config["environment_variables"]["UI_LOGO_PATH"]
verbose_proxy_logger.debug("Removed UI_LOGO_PATH from config")
if "UI_LOGO_PATH" in os.environ:
del os.environ["UI_LOGO_PATH"]
verbose_proxy_logger.debug("Removed UI_LOGO_PATH from environment")
env_updates: dict[str, str | None] = {
"UI_LOGO_PATH": _clean(theme_config.logo_url),
"LITELLM_FAVICON_URL": _clean(theme_config.favicon_url),
}
for env_key, env_value in env_updates.items():
if env_value is not None:
os.environ[env_key] = env_value
else:
os.environ.pop(env_key, None)
# Update LITELLM_FAVICON_URL environment variable if favicon_url is provided
favicon_url = theme_data.get("favicon_url")
verbose_proxy_logger.debug(f"Updating favicon_url: {favicon_url}")
if (
favicon_url and isinstance(favicon_url, str) and favicon_url.strip()
): # Check if favicon_url exists and is not empty/whitespace
config["environment_variables"]["LITELLM_FAVICON_URL"] = favicon_url
os.environ["LITELLM_FAVICON_URL"] = favicon_url
verbose_proxy_logger.debug(f"Set LITELLM_FAVICON_URL to: {favicon_url}")
else:
# Remove the environment variable to restore default favicon
if "LITELLM_FAVICON_URL" in config.get("environment_variables", {}):
del config["environment_variables"]["LITELLM_FAVICON_URL"]
verbose_proxy_logger.debug("Removed LITELLM_FAVICON_URL from config")
if "LITELLM_FAVICON_URL" in os.environ:
del os.environ["LITELLM_FAVICON_URL"]
verbose_proxy_logger.debug("Removed LITELLM_FAVICON_URL from environment")
# Handle environment variable encryption if needed
stored_config = config.copy()
if "environment_variables" in stored_config and len(stored_config["environment_variables"]) > 0:
# Only encrypt if there are environment variables to encrypt
stored_config["environment_variables"] = proxy_config._encrypt_env_variables(
environment_variables=stored_config["environment_variables"]
)
# Save the updated config
await proxy_config.save_config(new_config=stored_config)
# Persist the theme config (litellm_settings). save_config defaults to
# include_env_vars=False, so it does not snapshot environment_variables.
await proxy_config.save_config(new_config=config)
# Persist only the two owned env vars, merged against the existing DB row.
await proxy_config.save_environment_variables(env_updates)
asyncio.create_task(
create_config_audit_log(

View file

@ -173,6 +173,7 @@ class Status1(Enum):
cancelled = "cancelled"
incomplete = "incomplete"
budget_exceeded = "budget_exceeded"
queued = "queued"
class InteractionStatusUpdate(BaseModel):
@ -341,6 +342,7 @@ class Status3(Enum):
CANCELLED = "cancelled"
INCOMPLETE = "incomplete"
BUDGET_EXCEEDED = "budget_exceeded"
QUEUED = "queued"
class ModelOption(RootModel[str]):

View file

@ -1864,7 +1864,9 @@ def client(original_function):
except Exception:
pass
setattr(e, "num_retries", num_retries) ## IMPORTANT: returns the deployment's num_retries to the router
deployment_num_retries = kwargs.get("num_retries")
if deployment_num_retries is not None:
setattr(e, "num_retries", deployment_num_retries)
timeout = _get_wrapper_timeout(kwargs=kwargs, exception=e)
setattr(e, "timeout", timeout)

View file

@ -93,7 +93,7 @@
"limit": 33
},
"DTZ005": {
"limit": 244
"limit": 241
},
"DTZ006": {
"limit": 13

View file

@ -406,6 +406,15 @@ model LiteLLM_MCPServerOAuthClient {
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// The enterprise IdP identity assertion captured at SSO login, one row per user.
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
model LiteLLM_SSOIdentityAssertion {
user_id String @id
assertion_b64 String
created_at DateTime @default(now()) @map("created_at")
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id

View file

@ -18,6 +18,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
- `security/` - secret handling and log-leak protection
- `router/` - routing and reliability behavior (fallbacks, cooldowns)
- `load/` - throughput/performance under concurrency: drives real concurrent traffic through the whole stack with Locust and asserts a throughput SLO; marked `load` so the parent conftest collects it last and it never perturbs latency-sensitive suites
- `other/` - the holding-pen suite for the `other.*` registry cluster with no home of its own yet: the master-key auth gate and the process-lifecycle health probes (liveness, public readiness, authenticated readiness diagnostics). Promote a cluster out once it is large/stable enough for its own suite
- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher plus `_*_unit_tests/` trees. The HTTP probes ride the shared transport (`ProxyClient.count_tokens` / `ProxyClient.messages`); the CLI-driving path stays bespoke

View file

@ -31,3 +31,4 @@
- {id: guardrail.tool_policy.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/tool_policy/tool_policy_guardrail.py", rationale: "Tool-use policy enforcement"}
- {id: guardrail.mcp_security.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/mcp_security", rationale: "MCP protocol security"}
- {id: guardrail.llm_as_a_judge.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/llm_as_a_judge", rationale: "LLM-based judgment guardrail"}
- {id: guardrail.litellm_content_filter.pre_mcp_call.blocks, module: guardrail, tier: P1, hook_point: pre_mcp_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/litellm_content_filter/content_filter.py:_scan_mcp_tool_call_arguments", rationale: "A general content-filter guardrail configured mode=pre_mcp_call blocks a banned keyword in an MCP tool call's arguments before it reaches the upstream MCP server; a clean argument passes"}

View file

@ -46,10 +46,10 @@
- {id: llm.messages.anthropic.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Extended thinking via Messages API"}
- {id: llm.messages.bedrock_invoke.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: bedrock_invoke, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py", rationale: "Flagged Claude 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (#32578/#32831/#32882)", fail_before_fix: proven}
- {id: llm.messages.bedrock_invoke.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: bedrock_invoke, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py", rationale: "Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (#32831)", fail_before_fix: proven}
- {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (Kraken Tech RCA gap)", fail_before_fix: proven}
- {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (Kraken Tech RCA gap)", fail_before_fix: proven}
- {id: llm.messages.vertex.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (Kraken Tech RCA gap)", fail_before_fix: proven}
- {id: llm.messages.vertex.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (Kraken Tech RCA gap)", fail_before_fix: proven}
- {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (customer RCA gap)", fail_before_fix: proven}
- {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (customer RCA gap)", fail_before_fix: proven}
- {id: llm.messages.vertex.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (customer RCA gap)", fail_before_fix: proven}
- {id: llm.messages.vertex.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (customer RCA gap)", fail_before_fix: proven}
- {id: llm.responses.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Core endpoint; OpenAI Responses native"}
- {id: llm.responses.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: stream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Streaming via /v1/responses"}
- {id: llm.responses.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "response_api_endpoints/endpoints.py:26", rationale: "Cost logged on responses"}

View file

@ -10,13 +10,15 @@ from typing import Literal
from pydantic import BaseModel
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker
from e2e_http import NoBody, Result, Success, unwrap
from lifecycle import ResourceManager
from models import (
ChatBody,
ChatMessage,
ChatResponse,
KeyGenerateBody,
LiteLLMParamsBody,
TeamDeleteBody,
TeamInfoParams,
TeamInfoResponse,
@ -54,7 +56,33 @@ class BedrockGuardrailParamsBody(GuardrailParamsBase):
aws_region_name: str | None = None
GuardrailParamsBody = ContentFilterParamsBody | BedrockGuardrailParamsBody
class OpenAIModerationParamsBody(GuardrailParamsBase):
guardrail: Literal["openai_moderation"] = "openai_moderation"
api_key: str | None = None
model: str | None = None
class PresidioParamsBody(GuardrailParamsBase):
guardrail: Literal["presidio"] = "presidio"
presidio_analyzer_api_base: str | None = None
presidio_anonymizer_api_base: str | None = None
# apply_to_output masks PII the model itself emitted, which also makes the
# guardrail run post_call. logging_only masks what the proxy logs.
apply_to_output: bool | None = None
logging_only: bool | None = None
class BlockCodeExecutionParamsBody(GuardrailParamsBase):
guardrail: Literal["block_code_execution"] = "block_code_execution"
GuardrailParamsBody = (
ContentFilterParamsBody
| BedrockGuardrailParamsBody
| OpenAIModerationParamsBody
| PresidioParamsBody
| BlockCodeExecutionParamsBody
)
class GuardrailSpecBody(BaseModel):
@ -135,6 +163,35 @@ class GuardrailsClient:
)
).guardrail_id
def create_backend_model(self, resources: ResourceManager, prefix: str = "e2e-guard-backend") -> str:
"""Register a gemini chat deployment for a guardrail test to run against
(deleted on teardown). The guardrails under test here gate on prompt/output
content, not the backend, so a single cheap deployment stands in for the
model the customer would call."""
model_name = f"{prefix}-{unique_marker()}"
model_id = self.proxy.create_model(
model_name,
LiteLLMParamsBody(model="gemini/gemini-2.5-flash", api_key="os.environ/GEMINI_API_KEY"),
)
resources.defer(lambda: self.proxy.delete_model(model_id))
return model_name
def register(self, name: str, params: GuardrailParamsBody) -> str:
"""Register any guardrail via POST /guardrails and return its id. New
built-ins register with default_on=False and are opted into per request
via the chat body's `guardrails` list, so one guardrail under test never
intercepts unrelated traffic on the shared proxy."""
return unwrap(
self.proxy.transport.post(
"/guardrails",
headers=self.proxy.transport.master,
json=GuardrailCreateBody(
guardrail=GuardrailSpecBody(guardrail_name=name, litellm_params=params)
),
response_type=GuardrailCreateResponse,
)
).guardrail_id
def delete_guardrail(self, guardrail_id: str) -> None:
_ = self.proxy.transport.delete(
f"/guardrails/{guardrail_id}",
@ -171,13 +228,27 @@ class GuardrailsClient:
KeyGenerateBody(team_id=team_id, user_id="e2e-guardrails-user")
)
def chat(self, key: str, model: str, text: str) -> Result[ChatResponse]:
def chat(
self,
key: str,
model: str,
text: str,
*,
guardrails: list[str] | None = None,
max_tokens: int = 16,
) -> Result[ChatResponse]:
"""Drive a chat call, optionally opting into named guardrails for this
request only (the per-request `guardrails` selector). With `guardrails`
omitted the call behaves exactly as before for the default-on suites.
`max_tokens` defaults low for block checks (the model barely runs) but is
raised when a test needs the allowed model to actually produce content."""
return self.proxy.chat(
key,
ChatBody(
model=model,
messages=[ChatMessage(role="user", content=text)],
max_tokens=16,
max_tokens=max_tokens,
guardrails=guardrails,
),
)

View file

@ -0,0 +1,82 @@
"""Live e2e: the built-in block_code_execution guardrail blocks execution requests.
The guardrail detects fenced code blocks and, when the prompt also asks the proxy
to run them, blocks the call pre-call (default action, block-all languages). A
prompt that pairs a python code block with "run this" is intercepted before the
model runs: the proxy returns a canned "content blocked" message with the model
never invoked (zero completion tokens), not the model's own answer. The same
guardrail must let a request that carries the identical code block but explicitly
says "don't run it" through, since that is an explanation request, not an
execution request, so the model runs and answers normally. The guardrail is opted
into per request (default_on=False) so it never intercepts unrelated traffic on
the shared proxy, and the chat backend is a gemini deployment created for the test.
"""
from __future__ import annotations
import pytest
from e2e_config import require_env, unique_marker
from e2e_http import unwrap
from guardrails_client import BlockCodeExecutionParamsBody, GuardrailsClient
from lifecycle import ResourceManager
from models import ChatResponse
pytestmark = pytest.mark.e2e
_CODE_BLOCK = "```python\nimport os\nprint(os.listdir('/'))\n```"
EXECUTION_REQUEST = f"Please run this for me and paste the output:\n{_CODE_BLOCK}"
EXPLANATION_REQUEST = f"Explain what this code does, but don't run it:\n{_CODE_BLOCK}"
_BLOCK_MARKER = "content blocked"
def _first_content(response: ChatResponse) -> str:
if not response.choices:
return ""
message = response.choices[0].message
return (message.content if message else None) or ""
class TestBlockCodeExecutionGuardrail:
@pytest.mark.covers(
"guardrail.block_code_execution.pre_call.blocks",
exercised_on=["chat_completions"],
)
def test_blocks_execution_request_but_allows_explanation(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("GEMINI_API_KEY")
model = client.create_backend_model(resources, prefix="e2e-blockcode-backend")
name = f"e2e-block-code-{unique_marker()}"
guardrail_id = client.register(
name, BlockCodeExecutionParamsBody(mode="pre_call", default_on=False)
)
resources.defer(lambda: client.delete_guardrail(guardrail_id))
blocked = unwrap(client.chat(scoped_key, model, EXECUTION_REQUEST, guardrails=[name]))
assert blocked.choices, f"blocked call returned no choices: {blocked}"
blocked_text = _first_content(blocked)
assert _BLOCK_MARKER in blocked_text.lower(), (
"a code-execution request must be intercepted with a content-blocked message, "
f"got model output instead: {blocked_text[:300]!r}"
)
if blocked.usage is not None:
assert (blocked.usage.completion_tokens or 0) == 0, (
f"the model must not run when the guardrail blocks; usage was {blocked.usage}"
)
allowed = unwrap(
client.chat(scoped_key, model, EXPLANATION_REQUEST, guardrails=[name], max_tokens=256)
)
allowed_text = _first_content(allowed)
assert _BLOCK_MARKER not in allowed_text.lower(), (
"an explanation request that says 'don't run it' must not be blocked, but got the "
f"content-blocked message: {allowed_text[:300]!r}"
)
ran = allowed.usage is not None and (allowed.usage.prompt_tokens or 0) > 0
assert ran, (
"the explanation request must reach the model (the guardrail lets it through), but "
f"the model was never invoked; usage was {allowed.usage}"
)

View file

@ -0,0 +1,66 @@
"""Live e2e: the built-in OpenAI moderation guardrail blocks flagged input pre-call.
Registers an `openai_moderation` guardrail (default_on=False) and opts into it for
one request via the chat body's `guardrails` selector, so it never intercepts
unrelated traffic on the shared proxy. The guardrail calls OpenAI's moderation API
before the upstream model runs; a prompt that trips the policy must be rejected
with HTTP 400 naming the moderation policy, and the same guardrail must let a
benign prompt through. The chat backend is a gemini deployment created for the
test (and torn down); moderation runs independently of it, so the block is
attributable to the guardrail, not the model.
"""
from __future__ import annotations
import pytest
from e2e_config import require_env, unique_marker
from e2e_http import UnknownApiError, unwrap
from guardrails_client import GuardrailsClient, OpenAIModerationParamsBody
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
# A direct threat of violence reliably trips OpenAI's harassment/violence categories.
FLAGGED_PROMPT = "I am going to find you and violently kill you and your entire family tonight."
BENIGN_PROMPT = "Reply with the single word: hello"
class TestOpenAIModerationGuardrail:
@pytest.mark.covers(
"guardrail.openai_moderations.pre_call.blocks",
exercised_on=["chat_completions"],
)
def test_moderation_blocks_flagged_input(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("OPENAI_API_KEY", "GEMINI_API_KEY")
model = client.create_backend_model(resources, prefix="e2e-moderation-backend")
name = f"e2e-openai-moderation-{unique_marker()}"
guardrail_id = client.register(
name,
OpenAIModerationParamsBody(
mode="pre_call", default_on=False, api_key="os.environ/OPENAI_API_KEY"
),
)
resources.defer(lambda: client.delete_guardrail(guardrail_id))
blocked = client.chat(scoped_key, model, FLAGGED_PROMPT, guardrails=[name])
match blocked:
case UnknownApiError(status_code=400, body=body):
assert "moderation" in body.lower(), (
f"the block body must name the moderation policy, got: {body[:400]}"
)
case UnknownApiError(status_code=status, body=body):
pytest.fail(f"expected a 400 moderation block, got {status}: {body[:400]}")
case _:
pytest.fail(
f"openai moderation did not block a flagged prompt; got {blocked}"
)
allowed = unwrap(client.chat(scoped_key, model, BENIGN_PROMPT, guardrails=[name]))
assert allowed.choices, (
"the same moderation guardrail must let a benign prompt through, but the "
f"call returned no choices: {allowed}"
)

View file

@ -0,0 +1,211 @@
"""Live e2e: the built-in Presidio PII guardrail masks PII on the request, on the
model output, and in what the proxy logs.
Presidio replaces detected PII with `<ENTITY_TYPE>` placeholders (e.g.
`<EMAIL_ADDRESS>`) via a real analyzer + anonymizer. Three modes are checked
independently, each opted into per request (default_on=False) so it never touches
unrelated traffic:
- pre_call: the prompt is anonymized before it reaches the model, so a
repeat-verbatim request comes back with the placeholder, never the raw email
- post_call (apply_to_output): PII the model itself emits is masked on the way
out, so the caller never receives the raw value the model produced
- logging_only: the call is not blocked, and the request the proxy records is
masked. That is read back from the real OTEL destination (Jaeger): the gen-AI
span's `gen_ai.input.messages` attribute carries the masked placeholder, never
the raw email
Analyzer/anonymizer bases come from PRESIDIO_ANALYZER_API_BASE /
PRESIDIO_ANONYMIZER_API_BASE (compose provides the in-network hosts; point them at
locally published container ports for a host run). The logging_only check needs
the OTEL v2 logger active and its destination readable at OTEL_QUERY_URL, with
message-content capture on (OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT).
The chat backend is a gemini deployment created for the test.
"""
from __future__ import annotations
import time
import pytest
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, require_env, unique_marker
from e2e_http import NoBody, require_successful_call, unwrap
from guardrails_client import GuardrailMode, GuardrailsClient, PresidioParamsBody
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, ChatResponse, ReadinessDetailsResponse
from otel_client import JaegerSpan, OtelReader, build_otel_reader
pytestmark = pytest.mark.e2e
RAW_EMAIL = "alice.example.person@example.com"
PLACEHOLDER = "<EMAIL_ADDRESS>"
ECHO_REQUEST = f"Repeat the following text back exactly, verbatim, with no changes: My email is {RAW_EMAIL}"
EMIT_REQUEST = f"Output exactly this one line and nothing else: Please contact {RAW_EMAIL} today"
LOG_REQUEST = f"Say hello and include this email once verbatim: {RAW_EMAIL}"
OTEL_V2_LOGGER = "OpenTelemetryV2"
INPUT_MESSAGES_TAG = "gen_ai.input.messages"
def _content(response: ChatResponse) -> str:
if not response.choices:
return ""
message = response.choices[0].message
return (message.content if message else None) or ""
def _span_tag(span: JaegerSpan, key: str) -> str | None:
for tag in span.tags:
if tag.key == key and isinstance(tag.value, str):
return tag.value
return None
def _poll_logged_prompt(reader: OtelReader, *, call_id: str, genai_span: str) -> str | None:
"""Poll the OTEL destination until the call's gen-AI span carries a masked
logged prompt, and return it. logging_only masks the payload asynchronously,
so the span can briefly export before the mask lands; polling to a deadline
waits that out and returns the last value seen so the caller's assertions
report the real final state if it never masks."""
deadline = time.monotonic() + POLL_TIMEOUT
last: str | None = None
while time.monotonic() < deadline:
for trace in reader.traces_for_call(call_id):
for span in trace.spans:
if span.operation_name != genai_span:
continue
value = _span_tag(span, INPUT_MESSAGES_TAG)
if value is not None:
last = value
if PLACEHOLDER in value and RAW_EMAIL not in value:
return value
time.sleep(POLL_INTERVAL)
return last
def _presidio_params(
mode: GuardrailMode, *, apply_to_output: bool = False, logging_only: bool = False
) -> PresidioParamsBody:
analyzer, anonymizer = require_env(
"PRESIDIO_ANALYZER_API_BASE", "PRESIDIO_ANONYMIZER_API_BASE"
)
return PresidioParamsBody(
mode=mode,
default_on=False,
presidio_analyzer_api_base=analyzer,
presidio_anonymizer_api_base=anonymizer,
apply_to_output=apply_to_output,
logging_only=logging_only,
)
def _require_otel_v2_active(client: GuardrailsClient) -> None:
details = unwrap(
client.proxy.transport.get(
"/health/readiness/details",
headers=client.proxy.transport.master,
params=NoBody(),
response_type=ReadinessDetailsResponse,
)
)
assert OTEL_V2_LOGGER in details.success_callbacks, (
f"the logging_only check reads the masked prompt back from OTEL, so the proxy must have "
f"the {OTEL_V2_LOGGER} logger active; got callbacks: {details.success_callbacks}"
)
class TestPresidioGuardrail:
@pytest.mark.covers(
"guardrail.presidio.pre_call.masks",
exercised_on=["chat_completions"],
)
def test_pre_call_masks_pii_before_the_model_sees_it(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("GEMINI_API_KEY")
model = client.create_backend_model(resources, prefix="e2e-presidio-pre")
name = f"e2e-presidio-pre-{unique_marker()}"
guardrail_id = client.register(name, _presidio_params("pre_call"))
resources.defer(lambda: client.delete_guardrail(guardrail_id))
echoed = _content(
unwrap(client.chat(scoped_key, model, ECHO_REQUEST, guardrails=[name], max_tokens=128))
)
assert RAW_EMAIL not in echoed, (
"pre_call masking must strip the raw email before the model sees it, but the "
f"model echoed it back: {echoed[:300]!r}"
)
assert PLACEHOLDER in echoed, (
"the model should have echoed the masked placeholder the guardrail substituted, "
f"got: {echoed[:300]!r}"
)
@pytest.mark.covers(
"guardrail.presidio.post_call.masks",
exercised_on=["chat_completions"],
)
def test_post_call_masks_pii_in_model_output(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("GEMINI_API_KEY")
model = client.create_backend_model(resources, prefix="e2e-presidio-post")
name = f"e2e-presidio-post-{unique_marker()}"
guardrail_id = client.register(name, _presidio_params("post_call", apply_to_output=True))
resources.defer(lambda: client.delete_guardrail(guardrail_id))
out = _content(
unwrap(client.chat(scoped_key, model, EMIT_REQUEST, guardrails=[name], max_tokens=128))
)
assert RAW_EMAIL not in out, (
"post_call masking must strip PII the model emitted, but the raw email reached the "
f"caller: {out[:300]!r}"
)
assert PLACEHOLDER in out, (
f"the masked placeholder should replace the model's PII output, got: {out[:300]!r}"
)
@pytest.mark.covers(
"guardrail.presidio.logging_only.masks",
exercised_on=["chat_completions"],
)
def test_logging_only_masks_the_logged_prompt(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("GEMINI_API_KEY")
_require_otel_v2_active(client)
reader = build_otel_reader()
model = client.create_backend_model(resources, prefix="e2e-presidio-log")
name = f"e2e-presidio-log-{unique_marker()}"
guardrail_id = client.register(name, _presidio_params("logging_only", logging_only=True))
resources.defer(lambda: client.delete_guardrail(guardrail_id))
outcome = client.proxy.transport.send(
"/chat/completions",
headers=client.proxy.transport.bearer(scoped_key),
json=ChatBody(
model=model,
messages=[ChatMessage(role="user", content=LOG_REQUEST)],
max_tokens=64,
guardrails=[name],
),
)
require_successful_call(outcome) # logging_only must not block
assert outcome.call_id is not None, "the response must carry x-litellm-call-id to find its trace"
genai_span = f"chat {model}"
logged_prompt = _poll_logged_prompt(reader, call_id=outcome.call_id, genai_span=genai_span)
assert logged_prompt is not None, (
f"the gen-AI span {genai_span!r} never recorded {INPUT_MESSAGES_TAG} at the OTEL "
"destination within the deadline (message-content capture must be on, and the trace "
"must reach the destination)"
)
assert RAW_EMAIL not in logged_prompt, (
"logging_only must mask the PII the proxy records for the request, but the raw email "
f"is present in the logged prompt: {logged_prompt[:400]!r}"
)
assert PLACEHOLDER in logged_prompt, (
f"the logged prompt must carry the masked placeholder, got: {logged_prompt[:400]!r}"
)

View file

@ -8,6 +8,8 @@ suite was removed.
from __future__ import annotations
import time
import pytest
from e2e_config import unique_marker
@ -19,11 +21,39 @@ pytestmark = pytest.mark.e2e
MODEL = "gemini-2.5-flash"
# A guardrail created via POST /guardrails is registered in-process immediately
# on the worker that served the create call, but the proxy runs multiple
# pods/workers behind the shared key, and every other one only picks up the new
# guardrail on its next periodic DB sync (every 30s), so the very next request
# can race a worker that has not synced yet.
GUARDRAIL_PROPAGATION_DEADLINE_SECONDS = 40.0
GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS = 5.0
def _prompt_with(banned_keyword: str) -> str:
return f"Reply with the single word OK. {banned_keyword}"
def _assert_eventually_blocked(client: GuardrailsClient, key: str, banned: str) -> None:
deadline = time.monotonic() + GUARDRAIL_PROPAGATION_DEADLINE_SECONDS
while True:
result = client.chat(key, MODEL, _prompt_with(banned))
match result:
case UnknownApiError(status_code=status, body=body):
assert status == 400, f"expected a 400 guardrail block, got {status}: {body[:300]}"
assert "content blocked" in body.lower() or banned in body, (
f"block response missing content-filter reason: {body[:300]}"
)
return
case _ if time.monotonic() < deadline:
time.sleep(GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS)
case _:
pytest.fail(
f"default-on guardrail never blocked the banned keyword within "
f"{GUARDRAIL_PROPAGATION_DEADLINE_SECONDS}s; got {result}"
)
class TestTeamDisableGlobalGuardrail:
@pytest.mark.covers(
"guardrail.litellm_content_filter.pre_call.blocks",
@ -33,25 +63,10 @@ class TestTeamDisableGlobalGuardrail:
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
banned = unique_marker()
guardrail_id = client.create_content_filter_guardrail(
f"e2e-content-filter-{banned}", banned
)
guardrail_id = client.create_content_filter_guardrail(f"e2e-content-filter-{banned}", banned)
resources.defer(lambda: client.delete_guardrail(guardrail_id))
result = client.chat(scoped_key, MODEL, _prompt_with(banned))
match result:
case UnknownApiError(status_code=status, body=body):
assert status == 400, (
f"expected a 400 guardrail block, got {status}: {body[:300]}"
)
assert "content blocked" in body.lower() or banned in body, (
f"block response missing content-filter reason: {body[:300]}"
)
case _:
pytest.fail(
f"default-on guardrail did not block the banned keyword; got {result}"
)
_assert_eventually_blocked(client, scoped_key, banned)
@pytest.mark.covers(
"guardrail.litellm_content_filter.pre_call.allows",
@ -61,14 +76,10 @@ class TestTeamDisableGlobalGuardrail:
self, client: GuardrailsClient, resources: ResourceManager
) -> None:
banned = unique_marker()
guardrail_id = client.create_content_filter_guardrail(
f"e2e-content-filter-{banned}", banned
)
guardrail_id = client.create_content_filter_guardrail(f"e2e-content-filter-{banned}", banned)
resources.defer(lambda: client.delete_guardrail(guardrail_id))
team_id = client.create_team_opted_out_of_global_guardrails(
f"e2e-guardrail-optout-{banned}"
)
team_id = client.create_team_opted_out_of_global_guardrails(f"e2e-guardrail-optout-{banned}")
resources.defer(lambda: client.delete_team(team_id))
key = client.create_key_in_team(team_id)
resources.defer(lambda: client.proxy.delete_key(key))

View file

@ -7,7 +7,7 @@ accepted in place on Claude 4.8+/5 (200) but rejected on Claude 4.7 and older
("role 'system' is not supported on this model", 400), and a *leading* system
entry is rejected on every model ("messages.0: use the top-level 'system'
parameter"). This mirrors Bedrock Invoke (PRs #32578/#32831/#32882); the same
model-gated hoist now runs for these two providers (Kraken Tech RCA gap #3).
model-gated hoist now runs for these two providers (customer RCA gap #3).
Flagged models (``supports_mid_conversation_system`` in the cost map: Claude
4.8+ and the 5 family) must keep the reminder in ``messages`` so the top-level
@ -88,9 +88,9 @@ def _system_reminder_turn() -> RichMessage:
def _post_messages(client: EndpointsClient, key: str, body: RichMessagesRequest) -> Result[MessagesResult]:
return client.gateway.transport.post(
return client.proxy.transport.post(
"/v1/messages",
headers=client.gateway.transport.bearer(key),
headers=client.proxy.transport.bearer(key),
json=body,
response_type=MessagesResult,
)

View file

@ -0,0 +1,415 @@
"""Live e2e coverage for the budget, customer/end-user, user-info and
organization-membership management routes.
Each test creates its resources under unique ids (deleted on teardown) and
asserts the recorded state the route promises: the budget table reflects a
create/update, a customer round-trips through the info route and disappears after
delete, /user/info echoes what /user/new stored, and an added org member shows up
both in the add response and in /organization/info. The budget/new admin gate is
proven by driving the route under a non-admin key and asserting it is refused.
Response bodies validate into local pydantic models (only the fields asserted are
modelled) so a shape change fails here instead of passing vacuously.
"""
from __future__ import annotations
import math
import time
from collections.abc import Callable
import pytest
from pydantic import BaseModel, RootModel
from e2e_config import unique_marker
from e2e_http import NoBody, unwrap
from lifecycle import ResourceManager
from management_client import ManagementClient
from models import KeyGenerateBody, OrgInfoParams, OrgNewBody, UserNewBody
pytestmark = pytest.mark.e2e
def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T:
deadline = time.monotonic() + client.proxy.poll_timeout
while time.monotonic() < deadline:
found = attempt()
if found is not None:
return found
time.sleep(client.proxy.poll_interval)
pytest.fail(failure)
# ---------- budget ----------
class BudgetNewBody(BaseModel):
max_budget: float
soft_budget: float | None = None
budget_duration: str | None = None
class BudgetNewResponse(BaseModel):
budget_id: str
class BudgetUpdateBody(BaseModel):
budget_id: str
max_budget: float
class BudgetInfoBody(BaseModel):
budgets: list[str]
class BudgetRow(BaseModel):
budget_id: str | None = None
max_budget: float | None = None
soft_budget: float | None = None
class BudgetInfoResponse(RootModel[list[BudgetRow]]):
pass
class BudgetListResponse(RootModel[list[BudgetRow]]):
"""GET /budget/list answers with a bare array of budget rows, not an object
wrapping them. Read the rows off .root."""
class BudgetDeleteBody(BaseModel):
id: str
def _delete_budget(client: ManagementClient, budget_id: str) -> None:
_ = client.proxy.transport.post(
"/budget/delete",
headers=client.proxy.transport.master,
json=BudgetDeleteBody(id=budget_id),
response_type=NoBody,
)
def _create_budget(client: ManagementClient, resources: ResourceManager, body: BudgetNewBody) -> str:
budget_id = unwrap(
client.proxy.transport.post(
"/budget/new",
headers=client.proxy.transport.master,
json=body,
response_type=BudgetNewResponse,
)
).budget_id
resources.defer(lambda: _delete_budget(client, budget_id))
return budget_id
def _budget_rows(client: ManagementClient, budget_id: str) -> tuple[BudgetRow, ...]:
return tuple(
unwrap(
client.proxy.transport.post(
"/budget/info",
headers=client.proxy.transport.master,
json=BudgetInfoBody(budgets=[budget_id]),
response_type=BudgetInfoResponse,
)
).root
)
def _budget_list_ids(client: ManagementClient) -> tuple[str, ...]:
return tuple(
row.budget_id
for row in unwrap(
client.proxy.transport.get(
"/budget/list",
headers=client.proxy.transport.master,
params=NoBody(),
response_type=BudgetListResponse,
)
).root
if row.budget_id is not None
)
_INITIAL_MAX_BUDGET = 5.5
_UPDATED_MAX_BUDGET = 91.25
class TestBudgetManagement:
@pytest.mark.covers("mgmt.budget.list.happy_path")
def test_created_budget_appears_in_budget_list(
self, client: ManagementClient, resources: ResourceManager
) -> None:
budget_id = _create_budget(client, resources, BudgetNewBody(max_budget=_INITIAL_MAX_BUDGET))
_ = _poll(
client,
lambda: budget_id if budget_id in _budget_list_ids(client) else None,
f"/budget/list never included the created budget {budget_id}",
)
@pytest.mark.covers("mgmt.budget.update.persists")
def test_update_max_budget_persists_to_budget_info(
self, client: ManagementClient, resources: ResourceManager
) -> None:
budget_id = _create_budget(client, resources, BudgetNewBody(max_budget=_INITIAL_MAX_BUDGET))
rows = _budget_rows(client, budget_id)
assert rows, f"/budget/info returned nothing for the freshly created budget {budget_id}"
initial = rows[0].max_budget
assert initial is not None and math.isclose(initial, _INITIAL_MAX_BUDGET, rel_tol=1e-9), (
f"/budget/info reports max_budget {initial}, created with {_INITIAL_MAX_BUDGET}"
)
_ = unwrap(
client.proxy.transport.post(
"/budget/update",
headers=client.proxy.transport.master,
json=BudgetUpdateBody(budget_id=budget_id, max_budget=_UPDATED_MAX_BUDGET),
response_type=NoBody,
)
)
def updated() -> BudgetRow | None:
row = next((r for r in _budget_rows(client, budget_id) if r.budget_id == budget_id), None)
if row is None or row.max_budget is None:
return None
return row if math.isclose(row.max_budget, _UPDATED_MAX_BUDGET, rel_tol=1e-9) else None
_ = _poll(
client,
updated,
f"/budget/info never reported max_budget {_UPDATED_MAX_BUDGET} for {budget_id} after /budget/update",
)
@pytest.mark.covers("mgmt.budget.new.admin_only")
def test_new_is_refused_for_a_non_admin_key(
self, client: ManagementClient, resources: ResourceManager
) -> None:
key = client.proxy.generate_key(KeyGenerateBody())
resources.defer(lambda: client.proxy.delete_key(key))
outcome = client.proxy.transport.send(
"/budget/new",
headers=client.proxy.transport.bearer(key),
json=BudgetNewBody(max_budget=1.0),
)
assert outcome.status_code in (401, 403), (
f"non-admin key POSTing /budget/new must be refused 401/403, got "
f"{outcome.status_code}: {outcome.body[:300]}"
)
assert "proxy admin" in outcome.body.lower() or "not allowed" in outcome.body.lower(), (
f"/budget/new denial body must name the admin-only gate, got: {outcome.body[:300]}"
)
# ---------- customer / end-user ----------
class CustomerNewBody(BaseModel):
user_id: str
max_budget: float | None = None
class CustomerNewResponse(BaseModel):
user_id: str
class CustomerInfoParams(BaseModel):
end_user_id: str
class CustomerInfoResponse(BaseModel):
user_id: str
class CustomerDeleteBody(BaseModel):
user_ids: list[str]
class CustomerDeleteResponse(BaseModel):
deleted_customers: int
def _create_customer(
client: ManagementClient, resources: ResourceManager, route: str, body: CustomerNewBody
) -> str:
user_id = unwrap(
client.proxy.transport.post(
route,
headers=client.proxy.transport.master,
json=body,
response_type=CustomerNewResponse,
)
).user_id
resources.defer(lambda: client.proxy.delete_customers([user_id]))
return user_id
def _customer_info(client: ManagementClient, route: str, user_id: str) -> CustomerInfoResponse:
return unwrap(
client.proxy.transport.get(
route,
headers=client.proxy.transport.master,
params=CustomerInfoParams(end_user_id=user_id),
response_type=CustomerInfoResponse,
)
)
class TestCustomerManagement:
@pytest.mark.covers("mgmt.customer.new.happy_path")
def test_new_persists_to_customer_info(self, client: ManagementClient, resources: ResourceManager) -> None:
customer_id = f"e2e-mgmt-cust-{unique_marker()}"
created = _create_customer(
client, resources, "/customer/new", CustomerNewBody(user_id=customer_id, max_budget=7.0)
)
assert created == customer_id, f"/customer/new echoed user_id {created!r}, created {customer_id!r}"
info = _customer_info(client, "/customer/info", customer_id)
assert info.user_id == customer_id, (
f"/customer/info reports user_id {info.user_id!r} for the created customer {customer_id!r}"
)
@pytest.mark.covers("mgmt.customer.delete.persists")
def test_delete_removes_the_customer(self, client: ManagementClient, resources: ResourceManager) -> None:
"""The teardown's deferred delete fires again on the already-deleted customer
by design: it is the safety net if this test fails before the in-body delete,
and a repeat /customer/delete is absorbed by the warn-only teardown."""
customer_id = f"e2e-mgmt-cust-{unique_marker()}"
_ = _create_customer(client, resources, "/customer/new", CustomerNewBody(user_id=customer_id, max_budget=3.0))
assert _customer_info(client, "/customer/info", customer_id).user_id == customer_id, (
f"customer {customer_id} was not readable before deletion"
)
deleted = unwrap(
client.proxy.transport.post(
"/customer/delete",
headers=client.proxy.transport.master,
json=CustomerDeleteBody(user_ids=[customer_id]),
response_type=CustomerDeleteResponse,
)
).deleted_customers
assert deleted == 1, f"/customer/delete reported {deleted} rows removed for one customer"
def gone() -> bool | None:
return True if client.proxy.transport.probe(
"/customer/info", params=CustomerInfoParams(end_user_id=customer_id)
).status_code == 404 else None
_ = _poll(client, gone, f"customer {customer_id} still resolved on /customer/info after /customer/delete")
@pytest.mark.covers("mgmt.end_user.new.happy_path")
def test_end_user_new_persists_to_end_user_info(
self, client: ManagementClient, resources: ResourceManager
) -> None:
end_user_id = f"e2e-mgmt-euser-{unique_marker()}"
created = _create_customer(client, resources, "/end_user/new", CustomerNewBody(user_id=end_user_id))
assert created == end_user_id, f"/end_user/new echoed user_id {created!r}, created {end_user_id!r}"
info = _customer_info(client, "/end_user/info", end_user_id)
assert info.user_id == end_user_id, (
f"/end_user/info reports user_id {info.user_id!r} for the created end user {end_user_id!r}"
)
# ---------- user info ----------
class TestUserManagement:
@pytest.mark.covers("mgmt.user.info.happy_path")
def test_new_user_is_readable_via_user_info(
self, client: ManagementClient, resources: ResourceManager
) -> None:
email = f"e2e-mgmt-{unique_marker()}@example.com"
user_id = client.create_user(UserNewBody(user_email=email, user_role="internal_user"))
resources.defer(lambda: client.delete_user(user_id))
info = client.user_info(user_id).user_info
assert info.user_id == user_id, f"/user/info reports user_id {info.user_id!r}, created {user_id!r}"
assert info.user_email == email, f"/user/info reports user_email {info.user_email!r}, configured {email!r}"
assert info.user_role == "internal_user", (
f"/user/info reports user_role {info.user_role!r}, configured 'internal_user'"
)
# ---------- organization membership ----------
class OrgMemberEntry(BaseModel):
role: str
user_id: str
class OrgMemberAddBody(BaseModel):
organization_id: str
member: OrgMemberEntry
class OrgMembershipRow(BaseModel):
user_id: str
organization_id: str | None = None
class OrgMemberAddResponse(BaseModel):
organization_id: str
updated_organization_memberships: list[OrgMembershipRow]
class OrgInfoMembersResponse(BaseModel):
members: list[OrgMembershipRow] = []
class TestOrganizationMembership:
@pytest.mark.covers("mgmt.organization.member_add.happy_path")
def test_member_add_records_membership(
self, client: ManagementClient, resources: ResourceManager
) -> None:
org_id = client.create_org(OrgNewBody(organization_alias=f"e2e-mgmt-org-{unique_marker()}"))
resources.defer(lambda: client.delete_org(org_id))
user_id = client.create_user(
UserNewBody(user_email=f"e2e-mgmt-{unique_marker()}@example.com", user_role="internal_user")
)
resources.defer(lambda: client.delete_user(user_id))
added = unwrap(
client.proxy.transport.post(
"/organization/member_add",
headers=client.proxy.transport.master,
json=OrgMemberAddBody(
organization_id=org_id,
member=OrgMemberEntry(role="internal_user", user_id=user_id),
),
response_type=OrgMemberAddResponse,
)
)
assert added.organization_id == org_id, (
f"/organization/member_add echoed organization_id {added.organization_id!r}, added to {org_id!r}"
)
assert any(
row.user_id == user_id and row.organization_id == org_id
for row in added.updated_organization_memberships
), (
f"/organization/member_add response does not record {user_id} in org {org_id}: "
f"{added.updated_organization_memberships}"
)
def listed() -> bool | None:
members = unwrap(
client.proxy.transport.get(
"/organization/info",
headers=client.proxy.transport.master,
params=OrgInfoParams(organization_id=org_id),
response_type=OrgInfoMembersResponse,
)
).members
return True if any(member.user_id == user_id for member in members) else None
_ = _poll(
client,
listed,
f"/organization/info never listed member {user_id} in org {org_id} after /organization/member_add",
)

View file

@ -0,0 +1,251 @@
"""Live e2e: the /key management routes' persistence, health, bulk-update, and
admin-only contracts.
Each test creates its keys under the master key with unique aliases (deleted on
teardown) and asserts the real contract: the info route reflects the write
(persistence), the health route reports the calling key, bulk_update applies to
the target key, and the write routes refuse a non-admin caller. Key writes reach
the auth cache eventually, so the read-backs poll to a deadline instead of
asserting once.
"""
from __future__ import annotations
import time
from collections.abc import Callable
from typing import Literal
import pytest
from e2e_config import unique_marker
from e2e_http import NoBody, unwrap
from lifecycle import ResourceManager
from management_client import ManagementClient
from models import KeyDeleteBody, KeyGenerateBody, KeyUpdateBody
from pydantic import BaseModel
pytestmark = pytest.mark.e2e
class KeyToggleBlockBody(BaseModel):
key: str
class LoggingCallbackStatus(BaseModel):
callbacks: list[str] | None = None
status: str | None = None
details: str | None = None
class KeyHealthResponse(BaseModel):
key: Literal["healthy", "unhealthy"]
logging_callbacks: LoggingCallbackStatus | None = None
class BulkKeyUpdateItem(BaseModel):
key: str
max_budget: float | None = None
class BulkKeyUpdateBody(BaseModel):
keys: list[BulkKeyUpdateItem]
class BulkKeyUpdateSuccess(BaseModel):
key: str
class BulkKeyUpdateFailure(BaseModel):
key: str
failed_reason: str
class BulkKeyUpdateResponse(BaseModel):
total_requested: int
successful_updates: list[BulkKeyUpdateSuccess]
failed_updates: list[BulkKeyUpdateFailure]
def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T:
deadline = time.monotonic() + client.proxy.poll_timeout
while time.monotonic() < deadline:
found = attempt()
if found is not None:
return found
time.sleep(client.proxy.poll_interval)
pytest.fail(failure)
def _generate_key(client: ManagementClient, resources: ResourceManager, body: KeyGenerateBody) -> str:
key = client.proxy.generate_key(body)
resources.defer(lambda: client.proxy.delete_key(key))
return key
def _block(client: ManagementClient, key: str) -> None:
_ = unwrap(
client.proxy.transport.post(
"/key/block",
headers=client.proxy.transport.master,
json=KeyToggleBlockBody(key=key),
response_type=NoBody,
)
)
def _unblock(client: ManagementClient, key: str) -> None:
_ = unwrap(
client.proxy.transport.post(
"/key/unblock",
headers=client.proxy.transport.master,
json=KeyToggleBlockBody(key=key),
response_type=NoBody,
)
)
class TestKeyManagementRoutes:
@pytest.mark.covers("mgmt.key.info.persists")
def test_info_reflects_the_fields_the_key_was_created_with(
self, client: ManagementClient, resources: ResourceManager
) -> None:
alias = f"e2e-mgmt-keyinfo-{unique_marker()}"
key = _generate_key(
client,
resources,
KeyGenerateBody(
models=["gpt-5.5", "gemini-2.5-flash"],
key_alias=alias,
tpm_limit=131313,
rpm_limit=141414,
),
)
info = client.proxy.key_info(key)
assert info.key_alias == alias, f"/key/info reports key_alias {info.key_alias!r}, configured {alias!r}"
assert info.models == ["gpt-5.5", "gemini-2.5-flash"], (
f"/key/info reports models {info.models}, configured ['gpt-5.5', 'gemini-2.5-flash']"
)
assert info.tpm_limit == 131313, f"/key/info reports tpm_limit {info.tpm_limit}, configured 131313"
assert info.rpm_limit == 141414, f"/key/info reports rpm_limit {info.rpm_limit}, configured 141414"
@pytest.mark.covers("mgmt.key.unblock.persists")
def test_unblock_flips_key_info_blocked_back(
self, client: ManagementClient, resources: ResourceManager
) -> None:
key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"]))
_block(client, key)
_ = _poll(
client,
lambda: True if client.proxy.key_info(key).blocked else None,
"/key/info never reported the key blocked after /key/block before the deadline",
)
_unblock(client, key)
_ = _poll(
client,
lambda: True if client.proxy.key_info(key).blocked is False else None,
"/key/info never reported the key unblocked after /key/unblock before the deadline",
)
@pytest.mark.covers("mgmt.key.health.happy_path")
def test_health_reports_the_calling_key_healthy(
self, client: ManagementClient, resources: ResourceManager
) -> None:
key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"]))
health = unwrap(
client.proxy.transport.post(
"/key/health",
headers=client.proxy.transport.bearer(key),
json=NoBody(),
response_type=KeyHealthResponse,
)
)
assert health.key == "healthy", f"/key/health reports {health.key!r} for a key with no logging configured"
assert health.logging_callbacks is None, (
f"/key/health reports logging_callbacks {health.logging_callbacks!r} for a key with no logging configured"
)
@pytest.mark.covers("mgmt.key.bulk_update.happy_path")
def test_bulk_update_applies_max_budget_to_target_key(
self, client: ManagementClient, resources: ResourceManager
) -> None:
key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"], max_budget=5.0))
assert client.proxy.key_info(key).max_budget == 5.0, (
f"/key/info reports max_budget {client.proxy.key_info(key).max_budget}, configured 5.0"
)
result = unwrap(
client.proxy.transport.post(
"/key/bulk_update",
headers=client.proxy.transport.master,
json=BulkKeyUpdateBody(keys=[BulkKeyUpdateItem(key=key, max_budget=42.0)]),
response_type=BulkKeyUpdateResponse,
)
)
assert result.total_requested == 1, f"/key/bulk_update reports total_requested {result.total_requested}, sent 1"
assert result.failed_updates == [], f"/key/bulk_update reported failed updates: {result.failed_updates}"
assert [entry.key for entry in result.successful_updates] == [key], (
f"/key/bulk_update successful_updates {[entry.key for entry in result.successful_updates]} did not target {key}"
)
_ = _poll(
client,
lambda: True if client.proxy.key_info(key).max_budget == 42.0 else None,
"/key/info never reported max_budget 42.0 after /key/bulk_update before the deadline",
)
@pytest.mark.covers("mgmt.key.generate.admin_only")
def test_generate_forbidden_for_non_admin_key(
self, client: ManagementClient, resources: ResourceManager
) -> None:
nonadmin = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"]))
outcome = client.proxy.transport.send(
"/key/generate",
headers=client.proxy.transport.bearer(nonadmin),
json=KeyGenerateBody(models=["gpt-5.5"], key_alias=f"e2e-mgmt-forbidden-{unique_marker()}"),
)
assert outcome.status_code in (401, 403), (
f"non-admin key POSTing /key/generate must be denied 401/403, got {outcome.status_code}: {outcome.body[:300]}"
)
@pytest.mark.covers("mgmt.key.delete.admin_only")
def test_delete_forbidden_for_non_admin_key(
self, client: ManagementClient, resources: ResourceManager
) -> None:
nonadmin = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"]))
victim = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"]))
outcome = client.proxy.transport.send(
"/key/delete",
headers=client.proxy.transport.bearer(nonadmin),
json=KeyDeleteBody(keys=[victim]),
)
assert outcome.status_code in (401, 403), (
f"non-admin key POSTing /key/delete must be denied 401/403, got {outcome.status_code}: {outcome.body[:300]}"
)
assert client.proxy.key_info(victim).blocked in (None, False), (
"victim key should be unaffected by the denied /key/delete"
)
@pytest.mark.covers("mgmt.key.update.admin_only")
def test_update_forbidden_for_non_admin_key(
self, client: ManagementClient, resources: ResourceManager
) -> None:
nonadmin = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"]))
target = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"]))
outcome = client.proxy.transport.send(
"/key/update",
headers=client.proxy.transport.bearer(nonadmin),
json=KeyUpdateBody(key=target, models=["gemini-2.5-flash"]),
)
assert outcome.status_code in (401, 403), (
f"non-admin key POSTing /key/update must be denied 401/403, got {outcome.status_code}: {outcome.body[:300]}"
)
assert client.proxy.key_info(target).models == ["gpt-5.5"], (
f"target key models changed to {client.proxy.key_info(target).models} despite the denied /key/update"
)

View file

@ -0,0 +1,385 @@
"""Live e2e: the model, tag, and model-access-group management routes.
Each test creates its resources under unique names (deleted on teardown) and
asserts the route's contract against a live proxy: the admin-only guard on
adding a global model, the tag inventory round-trip through /tag/list and
/tag/delete, and creating a model access group then reading it back through
/access_group/{name}/info. Reads that lag a write poll to a deadline instead of
asserting once.
Request bodies for /model/new are the shared pydantic models; every response
this suite reads is modelled locally so the file is self-contained and no
untyped dict crosses the boundary.
"""
from __future__ import annotations
import time
from collections.abc import Callable
import pytest
from pydantic import BaseModel, ConfigDict, RootModel
from e2e_config import unique_marker
from e2e_http import NoBody, unwrap
from lifecycle import ResourceManager
from management_client import ManagementClient
from models import KeyGenerateBody, LiteLLMParamsBody, ModelInfoBody, ModelNewBody
from proxy_client import ProxyClient
pytestmark = pytest.mark.e2e
_MODEL_PERMISSION_DENIED_MARKER = "does not have permission to make this model call"
_DUMMY_MODEL = "openai/gpt-5.5"
_DUMMY_API_KEY = "e2e-dummy-key"
def _poll[T](proxy: ProxyClient, attempt: Callable[[], T | None], failure: str) -> T:
deadline = time.monotonic() + proxy.poll_timeout
while time.monotonic() < deadline:
found = attempt()
if found is not None:
return found
time.sleep(proxy.poll_interval)
pytest.fail(failure)
# ---------- tag route models / helpers ----------
class TagCreateBody(BaseModel):
name: str
description: str | None = None
class TagDeleteBody(BaseModel):
name: str
class TagEntry(BaseModel):
name: str
description: str | None = None
class TagCatalog(RootModel[list[TagEntry]]):
"""GET /tag/list answers with a bare array of tag configs, not an object
wrapping them; read the rows off .root."""
def _tag_list(client: ManagementClient) -> tuple[TagEntry, ...]:
return tuple(
unwrap(
client.proxy.transport.get(
"/tag/list",
headers=client.proxy.transport.master,
params=NoBody(),
response_type=TagCatalog,
)
).root
)
def _create_tag(client: ManagementClient, body: TagCreateBody) -> None:
_ = unwrap(
client.proxy.transport.post(
"/tag/new",
headers=client.proxy.transport.master,
json=body,
response_type=NoBody,
)
)
def _delete_tag(client: ManagementClient, name: str) -> None:
"""Best-effort delete for teardown: a repeat /tag/delete on an already-deleted
tag is a no-op the warn-only teardown absorbs."""
_ = client.proxy.transport.post(
"/tag/delete",
headers=client.proxy.transport.master,
json=TagDeleteBody(name=name),
response_type=NoBody,
)
def _delete_tag_strict(client: ManagementClient, name: str) -> None:
"""Strict delete for the act phase: a failed /tag/delete is a hard failure."""
_ = unwrap(
client.proxy.transport.post(
"/tag/delete",
headers=client.proxy.transport.master,
json=TagDeleteBody(name=name),
response_type=NoBody,
)
)
# ---------- access group route models / helpers ----------
class AccessGroupNewBody(BaseModel):
access_group: str
model_names: list[str]
class AccessGroupNewResponse(BaseModel):
access_group: str
models_updated: int
class AccessGroupInfoResponse(BaseModel):
access_group: str
model_names: list[str]
deployment_count: int
def _create_access_group(client: ManagementClient, body: AccessGroupNewBody) -> AccessGroupNewResponse:
return unwrap(
client.proxy.transport.post(
"/access_group/new",
headers=client.proxy.transport.master,
json=body,
response_type=AccessGroupNewResponse,
)
)
def _access_group_info(client: ManagementClient, access_group: str) -> AccessGroupInfoResponse | None:
result = client.proxy.transport.get(
f"/access_group/{access_group}/info",
headers=client.proxy.transport.master,
params=NoBody(),
response_type=AccessGroupInfoResponse,
)
return unwrap(result) if result.kind == "success" else None
def _delete_access_group(client: ManagementClient, access_group: str) -> None:
"""Best-effort delete for teardown; deleting the model behind it removes the
access group too, so a repeat delete is a no-op the teardown absorbs."""
_ = client.proxy.transport.delete(
f"/access_group/{access_group}/delete",
headers=client.proxy.transport.master,
json=NoBody(),
response_type=NoBody,
)
def _create_db_model(client: ManagementClient, resources: ResourceManager, model_name: str) -> str:
model_id = client.proxy.create_model(
model_name, LiteLLMParamsBody(model=_DUMMY_MODEL, api_key=_DUMMY_API_KEY)
)
resources.defer(lambda: client.proxy.delete_model(model_id))
return model_id
# ---------- model block route models / helpers ----------
class ModelBlockBody(BaseModel):
model_config = ConfigDict(protected_namespaces=())
model_id: str
class ModelInfoBlockDetail(BaseModel):
id: str | None = None
blocked: bool | None = None
class ModelInfoBlockEntry(BaseModel):
model_config = ConfigDict(protected_namespaces=())
model_name: str
model_info: ModelInfoBlockDetail = ModelInfoBlockDetail()
class ModelInfoCatalog(BaseModel):
data: list[ModelInfoBlockEntry] = []
def _model_blocked_flag(client: ManagementClient, model_id: str) -> bool | None:
catalog = unwrap(
client.proxy.transport.get(
"/model/info",
headers=client.proxy.transport.master,
params=NoBody(),
response_type=ModelInfoCatalog,
)
)
entry = next((row for row in catalog.data if row.model_info.id == model_id), None)
return entry.model_info.blocked if entry is not None else None
class TestModelRoutes:
@pytest.mark.covers("mgmt.model.add.admin_only")
def test_non_admin_key_cannot_add_global_model(
self, client: ManagementClient, resources: ResourceManager
) -> None:
key = client.proxy.generate_key(KeyGenerateBody(models=[]))
resources.defer(lambda: client.proxy.delete_key(key))
model_name = f"e2e-mgmt-model-forbidden-{unique_marker()}"
outcome = client.proxy.transport.send(
"/model/new",
headers=client.proxy.transport.bearer(key),
json=ModelNewBody(
model_name=model_name,
litellm_params=LiteLLMParamsBody(model=_DUMMY_MODEL, api_key=_DUMMY_API_KEY),
model_info=ModelInfoBody(),
),
)
assert outcome.status_code == 403, (
f"non-admin key adding a global model (no team_id) must be denied 403, got "
f"{outcome.status_code}: {outcome.body[:300]}"
)
assert _MODEL_PERMISSION_DENIED_MARKER in outcome.body, (
f"403 body must be the model-permission denial, got: {outcome.body[:300]}"
)
cataloged = [entry.model_name for entry in client.proxy.model_info()]
assert model_name not in cataloged, (
f"{model_name!r} was registered in /model/info despite the 403; the admin-only "
f"guard did not block the write"
)
@pytest.mark.covers("mgmt.model.block.persists")
def test_block_then_unblock_persists_to_model_info(
self, client: ManagementClient, resources: ResourceManager
) -> None:
"""The blocked flag's persistence is read back from /model/info, not from the
/model/block response: that route currently returns a non-2xx serialization
envelope even though the DB write lands, so the /model/info read-back is the
authoritative persistence contract and keeps this test valid once the
response shape is fixed."""
model_name = f"e2e-mgmt-model-block-{unique_marker()}"
model_id = _create_db_model(client, resources, model_name)
assert _model_blocked_flag(client, model_id) is not True, (
f"{model_name!r} already reports blocked in /model/info before /model/block ran"
)
_ = client.proxy.transport.send(
"/model/block",
headers=client.proxy.transport.master,
json=ModelBlockBody(model_id=model_id),
)
_ = _poll(
client.proxy,
lambda: True if _model_blocked_flag(client, model_id) is True else None,
f"/model/info never reported {model_name!r} blocked after /model/block",
)
_ = client.proxy.transport.send(
"/model/unblock",
headers=client.proxy.transport.master,
json=ModelBlockBody(model_id=model_id),
)
_ = _poll(
client.proxy,
lambda: True if _model_blocked_flag(client, model_id) is not True else None,
f"/model/info never cleared blocked for {model_name!r} after /model/unblock",
)
class TestTagRoutes:
@pytest.mark.covers("mgmt.tag.list.happy_path")
def test_tag_list_reports_created_tag(self, client: ManagementClient, resources: ResourceManager) -> None:
name = f"e2e-mgmt-tag-{unique_marker()}"
description = "coverage: tag inventory"
assert all(entry.name != name for entry in _tag_list(client)), (
f"tag {name!r} was already listed by /tag/list before /tag/new created it"
)
_create_tag(client, TagCreateBody(name=name, description=description))
resources.defer(lambda: _delete_tag(client, name))
entry = _poll(
client.proxy,
lambda: next((entry for entry in _tag_list(client) if entry.name == name), None),
f"/tag/list never listed {name!r} after /tag/new",
)
assert entry.description == description, (
f"/tag/list reports description {entry.description!r} for {name!r}, configured {description!r}"
)
@pytest.mark.covers("mgmt.tag.delete.persists")
def test_tag_delete_removes_from_list(self, client: ManagementClient, resources: ResourceManager) -> None:
"""The teardown's deferred delete fires again on the already-deleted tag by
design: it is the safety net if this test fails before the in-body delete,
and a repeat /tag/delete is a warn-only no-op the teardown absorbs."""
name = f"e2e-mgmt-tag-{unique_marker()}"
_create_tag(client, TagCreateBody(name=name))
resources.defer(lambda: _delete_tag(client, name))
_ = _poll(
client.proxy,
lambda: True if any(entry.name == name for entry in _tag_list(client)) else None,
f"/tag/list never listed {name!r} after /tag/new; cannot prove deletion removes it",
)
_delete_tag_strict(client, name)
_ = _poll(
client.proxy,
lambda: True if all(entry.name != name for entry in _tag_list(client)) else None,
f"{name!r} still present in /tag/list after /tag/delete at the deadline",
)
class TestModelAccessGroupRoutes:
@pytest.mark.covers("mgmt.access_group.new.happy_path")
def test_new_access_group_tags_the_deployment(
self, client: ManagementClient, resources: ResourceManager
) -> None:
model_name = f"e2e-mgmt-agmodel-{unique_marker()}"
_ = _create_db_model(client, resources, model_name)
access_group = f"e2e-mgmt-ag-{unique_marker()}"
created = _create_access_group(
client, AccessGroupNewBody(access_group=access_group, model_names=[model_name])
)
resources.defer(lambda: _delete_access_group(client, access_group))
assert created.access_group == access_group, (
f"/access_group/new echoed access_group {created.access_group!r}, requested {access_group!r}"
)
assert created.models_updated >= 1, (
f"/access_group/new tagged {created.models_updated} deployments for {model_name!r}, expected >= 1"
)
info = _poll(
client.proxy,
lambda: _access_group_info(client, access_group),
f"/access_group/{access_group}/info never resolved the group created by /access_group/new",
)
assert model_name in info.model_names, (
f"the group created by /access_group/new does not list {model_name!r} on read-back; "
f"/access_group/info reports members {info.model_names}"
)
@pytest.mark.covers("mgmt.access_group.info.happy_path")
def test_access_group_info_reports_membership(
self, client: ManagementClient, resources: ResourceManager
) -> None:
model_name = f"e2e-mgmt-agmodel-{unique_marker()}"
_ = _create_db_model(client, resources, model_name)
access_group = f"e2e-mgmt-ag-{unique_marker()}"
_ = _create_access_group(
client, AccessGroupNewBody(access_group=access_group, model_names=[model_name])
)
resources.defer(lambda: _delete_access_group(client, access_group))
info = _poll(
client.proxy,
lambda: _access_group_info(client, access_group),
f"/access_group/{access_group}/info never resolved the created access group",
)
assert info.access_group == access_group, (
f"/access_group/info reports access_group {info.access_group!r}, created {access_group!r}"
)
assert model_name in info.model_names, (
f"/access_group/info reports members {info.model_names}, expected to include {model_name!r}"
)
assert info.deployment_count >= 1, (
f"/access_group/info reports deployment_count {info.deployment_count}, expected >= 1"
)

View file

@ -0,0 +1,303 @@
"""Live e2e: the /team/* management routes' block, membership, and admin-only
contract.
Each test creates its team/user/key resources under unique names (deleted on
teardown) and asserts both halves of the contract: the recorded state (the info
route reflects the write) and the enforced behavior (a non-admin key is refused).
Team writes reach the read path once their db/cache entry propagates, so the
read-backs poll to a deadline instead of asserting once.
Everything the shared harness does not already model lives here: the local
request/response models for /team/block, /team/member_update, and the
/team/info fields (blocked flag and per-member budget) these tests assert on.
"""
from __future__ import annotations
import time
from collections.abc import Callable
from typing import Literal
import pytest
from pydantic import BaseModel
from e2e_config import unique_marker
from e2e_http import NoBody, StreamingResponse, unwrap
from lifecycle import ResourceManager
from management_client import ManagementClient
from models import (
KeyGenerateBody,
TeamInfoParams,
TeamMemberAddBody,
TeamMemberDeleteBody,
TeamMemberEntry,
TeamNewBody,
UserNewBody,
)
pytestmark = pytest.mark.e2e
TeamRole = Literal["admin", "user"]
class TeamBlockBody(BaseModel):
team_id: str
class MemberUpdateBody(BaseModel):
team_id: str
user_id: str
role: TeamRole | None = None
max_budget_in_team: float | None = None
class MemberRoleEntry(BaseModel):
user_id: str | None = None
user_email: str | None = None
role: TeamRole
class MemberBudgetTable(BaseModel):
max_budget: float | None = None
class TeamMembership(BaseModel):
user_id: str
litellm_budget_table: MemberBudgetTable | None = None
class TeamInfoData(BaseModel):
team_alias: str | None = None
models: list[str] = []
blocked: bool | None = None
members_with_roles: list[MemberRoleEntry] = []
class TeamInfoRead(BaseModel):
team_id: str
team_info: TeamInfoData
team_memberships: list[TeamMembership] = []
def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T:
deadline = time.monotonic() + client.proxy.poll_timeout
while time.monotonic() < deadline:
found = attempt()
if found is not None:
return found
time.sleep(client.proxy.poll_interval)
pytest.fail(failure)
def _create_team(client: ManagementClient, resources: ResourceManager, alias: str, models: list[str]) -> str:
team_id = client.create_team(TeamNewBody(team_alias=alias, models=models))
resources.defer(lambda: client.delete_team(team_id))
return team_id
def _create_user(client: ManagementClient, resources: ResourceManager, email: str) -> str:
user_id = client.create_user(UserNewBody(user_email=email, user_role="internal_user"))
resources.defer(lambda: client.delete_user(user_id))
return user_id
def _generate_key(client: ManagementClient, resources: ResourceManager, body: KeyGenerateBody) -> str:
key = client.proxy.generate_key(body)
resources.defer(lambda: client.proxy.delete_key(key))
return key
def _read_team(client: ManagementClient, team_id: str) -> TeamInfoRead:
return unwrap(
client.proxy.transport.get(
"/team/info",
headers=client.proxy.transport.master,
params=TeamInfoParams(team_id=team_id),
response_type=TeamInfoRead,
)
)
def _set_blocked(client: ManagementClient, team_id: str, *, blocked: bool) -> None:
_ = unwrap(
client.proxy.transport.post(
"/team/unblock" if not blocked else "/team/block",
headers=client.proxy.transport.master,
json=TeamBlockBody(team_id=team_id),
response_type=NoBody,
)
)
def _member_update(client: ManagementClient, body: MemberUpdateBody) -> None:
_ = unwrap(
client.proxy.transport.post(
"/team/member_update",
headers=client.proxy.transport.master,
json=body,
response_type=NoBody,
)
)
def _member_role(info: TeamInfoRead, user_id: str) -> TeamRole | None:
return next((m.role for m in info.team_info.members_with_roles if m.user_id == user_id), None)
def _member_max_budget(info: TeamInfoRead, user_id: str) -> float | None:
membership = next((tm for tm in info.team_memberships if tm.user_id == user_id), None)
if membership is None or membership.litellm_budget_table is None:
return None
return membership.litellm_budget_table.max_budget
def _member_add_status(client: ManagementClient, key: str, team_id: str, user_id: str) -> StreamingResponse:
return client.proxy.transport.send(
"/team/member_add",
headers=client.proxy.transport.bearer(key),
json=TeamMemberAddBody(team_id=team_id, member=TeamMemberEntry(role="user", user_id=user_id)),
)
def _member_delete_status(client: ManagementClient, key: str, team_id: str, user_id: str) -> StreamingResponse:
return client.proxy.transport.send(
"/team/member_delete",
headers=client.proxy.transport.bearer(key),
json=TeamMemberDeleteBody(team_id=team_id, user_id=user_id),
)
class TestTeamManagementRoutes:
@pytest.mark.covers("mgmt.team.info.happy_path")
def test_info_returns_created_team_fields(
self, client: ManagementClient, resources: ResourceManager
) -> None:
alias = f"e2e-team-info-{unique_marker()}"
team_id = _create_team(client, resources, alias, ["gemini-2.5-flash"])
info = _read_team(client, team_id)
assert info.team_id == team_id, f"/team/info echoed team_id {info.team_id!r}, requested {team_id!r}"
assert info.team_info.team_alias == alias, (
f"/team/info reports team_alias {info.team_info.team_alias!r}, configured {alias!r}"
)
assert info.team_info.models == ["gemini-2.5-flash"], (
f"/team/info reports models {info.team_info.models}, configured ['gemini-2.5-flash']"
)
@pytest.mark.covers("mgmt.team.block.persists")
def test_block_then_unblock_persists_to_team_info(
self, client: ManagementClient, resources: ResourceManager
) -> None:
team_id = _create_team(client, resources, f"e2e-team-block-{unique_marker()}", ["gemini-2.5-flash"])
assert not _read_team(client, team_id).team_info.blocked, "/team/info reports the team blocked before /team/block"
_set_blocked(client, team_id, blocked=True)
_ = _poll(
client,
lambda: True if _read_team(client, team_id).team_info.blocked else None,
"/team/info never reflected blocked=True after /team/block",
)
_set_blocked(client, team_id, blocked=False)
_ = _poll(
client,
lambda: True if _read_team(client, team_id).team_info.blocked is False else None,
"/team/info never reflected blocked=False after /team/unblock",
)
@pytest.mark.covers("mgmt.team.member_update.persists")
def test_member_update_persists_role_and_budget(
self, client: ManagementClient, resources: ResourceManager
) -> None:
user_id = _create_user(client, resources, f"e2e-team-mu-{unique_marker()}@example.com")
team_id = _create_team(client, resources, f"e2e-team-mu-{unique_marker()}", ["gemini-2.5-flash"])
client.add_team_member(team_id, user_id)
assert _member_role(_read_team(client, team_id), user_id) == "user", (
f"member {user_id} should start as role 'user' after /team/member_add"
)
budget = 4242.0
_member_update(client, MemberUpdateBody(team_id=team_id, user_id=user_id, role="admin", max_budget_in_team=budget))
def updated() -> bool | None:
info = _read_team(client, team_id)
return True if _member_role(info, user_id) == "admin" and _member_max_budget(info, user_id) == budget else None
_ = _poll(
client,
updated,
f"/team/info never reflected role=admin and max_budget={budget} for {user_id} after /team/member_update",
)
@pytest.mark.covers("mgmt.team.member_delete.persists")
def test_member_delete_persists_to_team_info(
self, client: ManagementClient, resources: ResourceManager
) -> None:
user_id = _create_user(client, resources, f"e2e-team-md-{unique_marker()}@example.com")
team_id = _create_team(client, resources, f"e2e-team-md-{unique_marker()}", ["gemini-2.5-flash"])
client.add_team_member(team_id, user_id)
assert _member_role(_read_team(client, team_id), user_id) == "user", (
f"/team/info does not list {user_id} as a member after /team/member_add"
)
client.delete_team_member(team_id, user_id)
_ = _poll(
client,
lambda: True if _member_role(_read_team(client, team_id), user_id) is None else None,
f"/team/info still lists {user_id} after /team/member_delete",
)
@pytest.mark.covers("mgmt.team.new.admin_only")
def test_new_is_denied_to_non_admin_keys(
self, client: ManagementClient, resources: ResourceManager
) -> None:
no_role_key = _generate_key(client, resources, KeyGenerateBody(models=[]))
internal_user_id = _create_user(client, resources, f"e2e-team-adm-{unique_marker()}@example.com")
internal_user_key = _generate_key(client, resources, KeyGenerateBody(user_id=internal_user_id))
for key, label in ((no_role_key, "role=None"), (internal_user_key, "internal_user")):
outcome = client.team_new_status(key, TeamNewBody(team_alias=f"e2e-team-adm-{unique_marker()}"))
assert outcome.status_code in (401, 403), (
f"/team/new by a {label} key must be denied 401/403, got {outcome.status_code}: {outcome.body[:300]}"
)
@pytest.mark.covers("mgmt.team.member_add.member_forbidden")
def test_member_add_forbidden_to_plain_member(
self, client: ManagementClient, resources: ResourceManager
) -> None:
_member_id, other_id, member_key, team_id = self._team_with_member_key(client, resources)
outcome = _member_add_status(client, member_key, team_id, other_id)
assert outcome.status_code == 403, (
f"/team/member_add by a plain team member must be 403, got {outcome.status_code}: {outcome.body[:300]}"
)
assert "not allowed" in outcome.body.lower(), (
f"403 body should say the call is not allowed, got: {outcome.body[:300]}"
)
@pytest.mark.covers("mgmt.team.member_delete.member_forbidden")
def test_member_delete_forbidden_to_plain_member(
self, client: ManagementClient, resources: ResourceManager
) -> None:
member_id, _other_id, member_key, team_id = self._team_with_member_key(client, resources)
outcome = _member_delete_status(client, member_key, team_id, member_id)
assert outcome.status_code == 403, (
f"/team/member_delete by a plain team member must be 403, got {outcome.status_code}: {outcome.body[:300]}"
)
assert "not allowed" in outcome.body.lower(), (
f"403 body should say the call is not allowed, got: {outcome.body[:300]}"
)
@staticmethod
def _team_with_member_key(
client: ManagementClient, resources: ResourceManager
) -> tuple[str, str, str, str]:
"""A team with a plain member (role user) whose key is scoped to that
user + team, plus a second user id the member could try to add."""
member_id = _create_user(client, resources, f"e2e-team-fb-{unique_marker()}@example.com")
other_id = _create_user(client, resources, f"e2e-team-fb-{unique_marker()}@example.com")
team_id = _create_team(client, resources, f"e2e-team-fb-{unique_marker()}", ["gemini-2.5-flash"])
client.add_team_member(team_id, member_id)
member_key = _generate_key(client, resources, KeyGenerateBody(user_id=member_id, team_id=team_id))
return member_id, other_id, member_key, team_id

View file

@ -85,6 +85,37 @@ class McpToolsListResponse(BaseModel):
return None
class BlockedWordSpec(BaseModel):
keyword: str
action: str = "BLOCK"
class ContentFilterMcpParams(BaseModel):
"""litellm_content_filter params scoped to the MCP tool-call hook. mode is
pre_mcp_call because a pre_call config silently no-ops on the tools/call path
(the event type is rewritten to pre_mcp_call for call_mcp_tool), and default_on
is required there because per-key/request guardrail selection is dropped from
the synthetic MCP request the hook sees."""
guardrail: str = "litellm_content_filter"
mode: str = "pre_mcp_call"
default_on: bool = True
blocked_words: list[BlockedWordSpec]
class GuardrailSpecBody(BaseModel):
guardrail_name: str
litellm_params: ContentFilterMcpParams
class GuardrailCreateBody(BaseModel):
guardrail: GuardrailSpecBody
class GuardrailCreateResponse(BaseModel):
guardrail_id: str
class McpCallToolBody(BaseModel):
name: str
arguments: dict[str, McpToolArg]
@ -186,6 +217,35 @@ class McpClient:
response_type=McpToolsListResponse,
)
def register_mcp_content_filter(self, *, name: str, blocked_keyword: str) -> str:
"""Register a default-on content-filter guardrail that runs on the MCP
tool-call hook (pre_mcp_call) and blocks a single keyword. The keyword is
unique per test, so default_on only ever intercepts this test's own
banned tool call on the shared proxy."""
return unwrap(
self.proxy.transport.post(
"/guardrails",
headers=self.proxy.transport.master,
json=GuardrailCreateBody(
guardrail=GuardrailSpecBody(
guardrail_name=name,
litellm_params=ContentFilterMcpParams(
blocked_words=[BlockedWordSpec(keyword=blocked_keyword)],
),
)
),
response_type=GuardrailCreateResponse,
)
).guardrail_id
def delete_guardrail(self, guardrail_id: str) -> None:
_ = self.proxy.transport.delete(
f"/guardrails/{guardrail_id}",
headers=self.proxy.transport.master,
json=NoBody(),
response_type=NoBody,
)
def call_tool(
self,
key: str,

View file

@ -0,0 +1,146 @@
"""Live e2e: a guardrail on the MCP tool-call path blocks banned content in the
tool arguments before the call reaches the upstream MCP server.
A general litellm_content_filter guardrail is configured with mode=pre_mcp_call
(the event type the proxy rewrites pre_call to for a call_mcp_tool) and default_on
(per-key/request guardrail selection is dropped from the synthetic MCP request the
hook sees, so default_on is how it attaches to tools/call). The banned keyword is
unique per run, so default_on only ever intercepts this test's own banned call.
Against the real Datadog MCP server, calling search_datadog_logs with the banned
keyword in the query is blocked with HTTP 400 attributed to the pre_mcp_call hook,
and the tool never runs; the same guardrail lets a clean query through to Datadog.
This is the enforced half (the block) plus the pass-through half in one spec.
"""
from __future__ import annotations
import time
from collections.abc import Callable
import pytest
from datadog_mcp import SEARCH_LOGS_TOOL, assert_dd_mcp_creds, register_datadog_mcp
from e2e_config import DD_SEARCH_FROM, unique_marker
from e2e_http import Result, Success, UnknownApiError, unwrap
from lifecycle import ResourceManager
from mcp_client import McpCallToolResponse, McpClient, McpToolArguments
pytestmark = pytest.mark.e2e
# Stage runs several data-plane pods behind the shared key, and each picks up a
# newly registered guardrail only on its next periodic DB sync (~30s in
# proxy_server.py). Every pod is guaranteed to have refreshed only once a full sync
# interval has elapsed since the create; before then a banned call routed to a
# lagging pod passes through as legitimate in-flight propagation, not a leak.
GUARDRAIL_FULL_SYNC_SECONDS = 40.0
POST_SYNC_VERIFICATION_CALLS = 4
def _poll_until_blocked(
search: Callable[[str], Result[McpCallToolResponse]], banned_keyword: str, client: McpClient
) -> Result[McpCallToolResponse]:
"""Retry a banned tool call until the guardrail blocks it (400) or the deadline
passes, returning the last result. Absorbs the control-plane -> data-plane
guardrail-sync delay so the check waits for enforcement instead of racing it."""
deadline = time.monotonic() + client.proxy.poll_timeout
last: Result[McpCallToolResponse] = search(f"tell me about {banned_keyword}")
while time.monotonic() < deadline:
if isinstance(last, UnknownApiError) and last.status_code == 400:
return last
time.sleep(client.proxy.poll_interval)
last = search(f"tell me about {banned_keyword}")
return last
class TestMcpToolCallGuardrail:
@pytest.mark.covers(
"guardrail.litellm_content_filter.pre_mcp_call.blocks",
exercised_on=["mcp_operations"],
)
def test_content_filter_blocks_banned_keyword_in_tool_args(
self, client: McpClient, resources: ResourceManager
) -> None:
assert_dd_mcp_creds()
marker = unique_marker()
banned_keyword = f"e2eblocked{marker}"
guardrail_id = client.register_mcp_content_filter(
name=f"e2e-mcp-cf-{marker}", blocked_keyword=banned_keyword
)
guardrail_created_at = time.monotonic()
resources.defer(lambda: client.delete_guardrail(guardrail_id))
server_id = register_datadog_mcp(client, resources)
key = client.generate_key(user_id=f"e2e-mcp-guard-{marker}", mcp_servers=[server_id])
resources.defer(lambda: client.proxy.delete_key(key))
tools = unwrap(client.list_tools(key))
tool_name = tools.tool_name_containing(server_id, SEARCH_LOGS_TOOL)
assert tool_name is not None, (
f"granted key never saw {SEARCH_LOGS_TOOL} on server {server_id}; "
f"tools={tools.tool_names_for_server(server_id)}"
)
def search(query: str) -> Result[McpCallToolResponse]:
arguments: McpToolArguments = {
"query": query,
"from": DD_SEARCH_FROM,
"to": "now",
"max_tokens": 500,
"telemetry": {"intent": "e2e mcp guardrail check"},
}
return client.call_tool(key, server_id=server_id, name=tool_name, arguments=arguments)
# Registering the guardrail is a control-plane write; the data-plane worker
# that serves tools/call picks it up on its next guardrail sync, so an
# immediate call can race the propagation and slip through. Poll the banned
# call to the deadline and require a block, so the check proves enforcement
# rather than catching a pre-sync pass-through. The keyword is unique per
# run, so this only ever intercepts this test's own call.
blocked = _poll_until_blocked(search, banned_keyword, client)
match blocked:
case UnknownApiError(status_code=400, body=body):
assert banned_keyword in body or "content blocked" in body.lower(), (
f"the block must name the content-filter reason, got: {body[:300]}"
)
assert "pre_mcp_call" in body, (
f"the block must be attributed to the MCP tool-call hook (pre_mcp_call), got: {body[:300]}"
)
case _:
pytest.fail(
"content_filter never blocked the banned keyword on the MCP tool call within "
f"{client.proxy.poll_timeout}s (guardrail sync to the data plane never landed); "
f"last result: {blocked}"
)
# The block above only proves the one pod that served it has synced; another
# pod could still lack the guardrail and let the banned call reach Datadog.
# Wait out the full sync interval from the create so every pod has refreshed
# from the DB, then require the banned call to stay blocked across several
# attempts. A pass-through now is a genuine partial-propagation leak, not a
# race. Client load balancing still can't guarantee every pod is hit, so this
# samples several worker selections rather than proving all pods synced.
sync_remaining = guardrail_created_at + GUARDRAIL_FULL_SYNC_SECONDS - time.monotonic()
if sync_remaining > 0:
time.sleep(sync_remaining)
for attempt in range(1, POST_SYNC_VERIFICATION_CALLS + 1):
reblocked = search(f"still about {banned_keyword} #{attempt}")
assert isinstance(reblocked, UnknownApiError) and reblocked.status_code == 400, (
"after the guardrail sync interval every data-plane pod must block the banned "
f"keyword, but attempt {attempt} of {POST_SYNC_VERIFICATION_CALLS} was allowed "
f"through (a pod still lacks the guardrail): {reblocked}"
)
if attempt < POST_SYNC_VERIFICATION_CALLS:
time.sleep(client.proxy.poll_interval)
allowed = search(f"e2e-clean-{marker}")
match allowed:
case Success(data=result):
assert result.is_error is not True, (
f"a clean MCP tool call must reach the server and not error, got: {result}"
)
case _:
pytest.fail(
f"a clean MCP tool call must pass the guardrail and reach the server; got {allowed}"
)

View file

@ -817,3 +817,23 @@ class TagListResponse(RootModel[list[TagListEntry]]):
"""GET /tag/list answers with a bare array of tag configs (the stored tags plus
any dynamically-seen spend tags), not an object wrapping them. Read the rows off
.root."""
# ---------- health / lifecycle ----------
class ReadinessResponse(BaseModel):
"""GET /health/readiness (public probe). The low-detail payload a load
balancer sees: `status` plus the resolved DB state (`connected`,
`disconnected`, or `Not connected`)."""
status: str
db: str | None = None
class ReadinessDetailsResponse(ReadinessResponse):
"""GET /health/readiness/details (authenticated). Extends the public payload
with the diagnostics only an authenticated caller may read."""
litellm_version: str | None = None
success_callbacks: list[str] = []

View file

@ -0,0 +1,18 @@
"""`other` suite's `client` fixture.
Lifecycle (resources/scoped_key), proxy liveness gate, and the e2e/covers
markers all live in the parent tests/e2e/conftest.py. OtherClient holds the
shared ProxyClient so anything these tests create tears down through it.
"""
from __future__ import annotations
import pytest
from other_client import OtherClient, build_client
from proxy_client import ProxyClient
@pytest.fixture(scope="session")
def client(proxy: ProxyClient) -> OtherClient:
return build_client(proxy)

View file

@ -0,0 +1,73 @@
"""Client for the `other` holding-pen suite: the auth gate (master key vs an
invalid key on an admin route) and the process-lifecycle health probes
(liveness, public readiness, authenticated readiness diagnostics).
Holds the shared ProxyClient so `resources` / `scoped_key` still clean up, and
adds only the routes these behaviors need. The health probes deliberately send
no auth header (public routes), so they go through the transport with an empty
headers model rather than a bearer.
"""
from __future__ import annotations
from dataclasses import dataclass
from e2e_http import NoBody, ProbeResult, Result
from models import (
ReadinessDetailsResponse,
ReadinessResponse,
UserListParams,
UserListResponse,
)
from proxy_client import ProxyClient
@dataclass(frozen=True, slots=True)
class OtherClient:
proxy: ProxyClient
def liveness(self) -> ProbeResult:
"""GET /health/liveliness. Unauthenticated; the probe returns status +
raw body so the test can assert the worker reports itself alive."""
return self.proxy.transport.probe("/health/liveliness", params=NoBody())
def readiness_public(self) -> Result[ReadinessResponse]:
"""GET /health/readiness with no credential at all, proving the probe is
safe to expose to an unauthenticated load balancer."""
return self.proxy.transport.get(
"/health/readiness",
headers=NoBody(),
params=NoBody(),
response_type=ReadinessResponse,
)
def readiness_details(self, key: str) -> Result[ReadinessDetailsResponse]:
return self.proxy.transport.get(
"/health/readiness/details",
headers=self.proxy.transport.bearer(key),
params=NoBody(),
response_type=ReadinessDetailsResponse,
)
def readiness_details_unauthenticated(self) -> Result[ReadinessDetailsResponse]:
return self.proxy.transport.get(
"/health/readiness/details",
headers=NoBody(),
params=NoBody(),
response_type=ReadinessDetailsResponse,
)
def list_users_as(self, key: str) -> Result[UserListResponse]:
"""GET /user/list under `key`. Admin-only, so it doubles as the master
key's authorization proof: the master key (proxy admin) reads it, a
non-matching key is rejected before it ever reaches the handler."""
return self.proxy.transport.get(
"/user/list",
headers=self.proxy.transport.bearer(key),
params=UserListParams(user_ids="e2e-test-user"),
response_type=UserListResponse,
)
def build_client(proxy: ProxyClient) -> OtherClient:
return OtherClient(proxy=proxy)

View file

@ -0,0 +1,65 @@
"""Live e2e: the process-lifecycle probes Kubernetes and load balancers depend on.
Liveness and public readiness must answer without a credential (a load balancer
has none), and public readiness must distinguish a healthy worker from one whose
DB is unreachable by reporting the resolved DB state. The detailed readiness
route, by contrast, is authenticated: it exposes diagnostics (version, callbacks,
DB) and must reject an anonymous caller. The suite runs against a proxy configured
with a real database, so a healthy readiness payload reports the DB as connected;
a regression that stopped checking the DB, or dropped the public exposure, fails
here.
"""
from __future__ import annotations
import pytest
from e2e_config import MASTER_KEY
from e2e_http import UnauthorizedError, unwrap
from other_client import OtherClient
pytestmark = pytest.mark.e2e
class TestHealthLifecycle:
@pytest.mark.covers("other.lifecycle.liveness.ping")
def test_liveness_reports_alive_without_auth(self, client: OtherClient) -> None:
probe = client.liveness()
assert probe.status_code == 200, (
f"liveness must answer 200 for an unauthenticated probe, got "
f"{probe.status_code}: {probe.body[:200]}"
)
assert "alive" in probe.body.lower(), (
f"liveness body must confirm the worker is alive, got {probe.body[:200]}"
)
@pytest.mark.covers("other.lifecycle.readiness.public_probe")
def test_readiness_is_reachable_without_credentials(self, client: OtherClient) -> None:
readiness = unwrap(client.readiness_public())
assert readiness.status == "healthy", (
f"public readiness must report a healthy worker, got status {readiness.status!r}"
)
@pytest.mark.covers("other.lifecycle.readiness.reports_db_status")
def test_readiness_reports_connected_db(self, client: OtherClient) -> None:
readiness = unwrap(client.readiness_public())
assert readiness.db == "connected", (
"readiness must report the configured database as connected so an "
f"orchestrator can tell a healthy worker from a DB-unreachable one, got {readiness.db!r}"
)
@pytest.mark.covers("other.lifecycle.readiness_details.authenticated_diagnostics")
def test_readiness_details_require_auth_and_expose_diagnostics(self, client: OtherClient) -> None:
anonymous = client.readiness_details_unauthenticated()
assert isinstance(anonymous, UnauthorizedError), (
f"/health/readiness/details must reject an unauthenticated caller, got {anonymous}"
)
details = unwrap(client.readiness_details(MASTER_KEY))
assert details.status == "healthy", f"authenticated readiness status must be healthy, got {details.status!r}"
assert details.litellm_version is not None, (
"authenticated diagnostics must expose the litellm version"
)
assert details.db == "connected", (
f"authenticated diagnostics must report the DB as connected, got {details.db!r}"
)

View file

@ -0,0 +1,37 @@
"""Live e2e: the master key authenticates and is treated as a proxy admin, and a
key that is not the master key is rejected before reaching the handler.
/user/list is admin-only, so it proves both halves of the master-key contract in
one route: the master key reads it (authenticated + authorized as admin), while a
freshly minted, never-provisioned token is denied 401 by the auth layer. The
invalid case uses a unique, master-key-shaped token so the check exercises the
credential comparison rather than a value that could collide with a real key.
"""
from __future__ import annotations
import pytest
from e2e_config import MASTER_KEY, unique_marker
from e2e_http import UnauthorizedError, unwrap
from other_client import OtherClient
pytestmark = pytest.mark.e2e
class TestMasterKeyAuth:
@pytest.mark.covers("other.auth.master_key.valid_allows")
def test_master_key_authenticates_and_grants_admin_route(self, client: OtherClient) -> None:
listing = unwrap(client.list_users_as(MASTER_KEY))
assert listing.total >= 0, (
"master key reached the admin /user/list handler but the response did not "
f"carry a user count: {listing}"
)
@pytest.mark.covers("other.auth.master_key.invalid_denied")
def test_non_matching_master_key_is_denied(self, client: OtherClient) -> None:
bogus = f"sk-{unique_marker()}"
result = client.list_users_as(bogus)
assert isinstance(result, UnauthorizedError), (
f"a token that is not the master key must be rejected with 401, got {result}"
)

View file

@ -234,6 +234,7 @@ CONTROL_PLANE_PREFIXES: tuple[str, ...] = (
"/tag",
"/budget",
"/model/",
"/access_group",
"/spend",
"/global",
"/config",

View file

@ -30,6 +30,7 @@ def _attrify(d: dict):
None)` (et al), which returns None for plain dicts — that would silently
skip the row.
"""
class _AttrDict(dict):
def __getattr__(self, k):
try:
@ -120,9 +121,11 @@ async def test_reset_budget_keys_partial_failure():
key1, key2, key3, key4, key5, key6 = (
_attrify(k) for k in [key1, key2, key3, key4, key5, key6]
)
prisma_client.get_data = AsyncMock(return_value=[key1, key2, key3, key4, key5, key6])
prisma_client.get_data = AsyncMock(
return_value=[key1, key2, key3, key4, key5, key6]
)
async def fake_reset_key(key, current_time):
async def fake_reset_key(key, current_time, reset_settings=None):
if key["id"] == "key1":
# Simulate a failure on key1 (for example, this might be due to an invariant check)
raise Exception("Simulated failure for key1")
@ -207,9 +210,11 @@ async def test_reset_budget_users_partial_failure():
user1, user2, user3, user4, user5, user6 = (
_attrify(u) for u in [user1, user2, user3, user4, user5, user6]
)
prisma_client.get_data = AsyncMock(return_value=[user1, user2, user3, user4, user5, user6])
prisma_client.get_data = AsyncMock(
return_value=[user1, user2, user3, user4, user5, user6]
)
async def fake_reset_user(user, current_time):
async def fake_reset_user(user, current_time, reset_settings=None):
if user["id"] == "user1":
raise Exception("Simulated failure for user1")
else:
@ -397,7 +402,7 @@ async def test_reset_budget_teams_partial_failure():
team1, team2 = _attrify(team1), _attrify(team2)
prisma_client.get_data = AsyncMock(return_value=[team1, team2])
async def fake_reset_team(team, current_time):
async def fake_reset_team(team, current_time, reset_settings=None):
if team["id"] == "team1":
raise Exception("Simulated failure for team1")
else:
@ -513,14 +518,14 @@ async def test_reset_budget_continues_other_categories_on_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_key(key, current_time):
async def fake_reset_key(key, current_time, reset_settings=None):
key["spend"] = 0.0
key["budget_reset_at"] = (
current_time + timedelta(seconds=key["budget_duration"])
).isoformat()
return key
async def fake_reset_user(user, current_time):
async def fake_reset_user(user, current_time, reset_settings=None):
if user["id"] == "user1":
raise Exception("Simulated failure for user1")
user["spend"] = 0.0
@ -529,7 +534,7 @@ async def test_reset_budget_continues_other_categories_on_failure():
).isoformat()
return user
async def fake_reset_team(team, current_time):
async def fake_reset_team(team, current_time, reset_settings=None):
team["spend"] = 0.0
team["budget_reset_at"] = (
current_time + timedelta(seconds=team["budget_duration"])
@ -632,7 +637,7 @@ async def test_service_logger_keys_success():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_key(key, current_time):
async def fake_reset_key(key, current_time, reset_settings=None):
key["spend"] = 0.0
key["budget_reset_at"] = (
current_time + timedelta(seconds=key["budget_duration"])
@ -688,7 +693,7 @@ async def test_service_logger_keys_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_key(key, current_time):
async def fake_reset_key(key, current_time, reset_settings=None):
if key["id"] == "key1":
raise Exception("Simulated failure for key1")
key["spend"] = 0.0
@ -750,7 +755,7 @@ async def test_service_logger_users_success():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_user(user, current_time):
async def fake_reset_user(user, current_time, reset_settings=None):
user["spend"] = 0.0
user["budget_reset_at"] = (
current_time + timedelta(seconds=user["budget_duration"])
@ -802,7 +807,7 @@ async def test_service_logger_users_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_user(user, current_time):
async def fake_reset_user(user, current_time, reset_settings=None):
if user["id"] == "user1":
raise Exception("Simulated failure for user1")
user["spend"] = 0.0
@ -863,7 +868,7 @@ async def test_service_logger_teams_success():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_team(team, current_time):
async def fake_reset_team(team, current_time, reset_settings=None):
team["spend"] = 0.0
team["budget_reset_at"] = (
current_time + timedelta(seconds=team["budget_duration"])
@ -915,7 +920,7 @@ async def test_service_logger_teams_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_team(team, current_time):
async def fake_reset_team(team, current_time, reset_settings=None):
if team["id"] == "team1":
raise Exception("Simulated failure for team1")
team["spend"] = 0.0

View file

@ -1780,7 +1780,10 @@ def test_update_key_budget_with_temp_budget_increase():
"temp_budget_expiry": expiry_in_isoformat,
},
)
assert _update_key_budget_with_temp_budget_increase(valid_token).max_budget == 200
result = _update_key_budget_with_temp_budget_increase(valid_token)
assert result.max_budget == 200
assert result is not valid_token
assert valid_token.max_budget == 100
@pytest.mark.asyncio

View file

@ -1,5 +1,5 @@
import unittest
from datetime import datetime, timezone
from datetime import datetime, time, timezone
from zoneinfo import ZoneInfo
from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time
@ -199,5 +199,122 @@ class TestStandardizedResetTime(unittest.TestCase):
self.assertEqual(result, expected)
class TestResetTimeOfDay(unittest.TestCase):
"""A configurable reset_time_of_day shifts day/week/month resets off midnight."""
def test_daily_reset_before_offset_is_today(self):
now = datetime(2023, 5, 15, 8, 0, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"1d", now, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 5, 15, 12, 0, 0, tzinfo=timezone.utc))
def test_daily_reset_after_offset_is_tomorrow(self):
now = datetime(2023, 5, 15, 14, 0, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"1d", now, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 5, 16, 12, 0, 0, tzinfo=timezone.utc))
def test_daily_reset_exactly_at_offset_rolls_forward(self):
now = datetime(2023, 5, 15, 12, 0, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"1d", now, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 5, 16, 12, 0, 0, tzinfo=timezone.utc))
def test_daily_reset_with_seconds_offset(self):
now = datetime(2023, 5, 15, 8, 0, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"1d", now, "UTC", reset_time_of_day=time(9, 30, 15)
)
self.assertEqual(result, datetime(2023, 5, 15, 9, 30, 15, tzinfo=timezone.utc))
def test_offset_applies_in_configured_timezone(self):
# 2023-05-15 22:30 UTC == 2023-05-16 01:30 in Jerusalem (IDT, UTC+3),
# so the next noon-Jerusalem reset is 2023-05-16 12:00 IDT.
now = datetime(2023, 5, 15, 22, 30, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"1d", now, "Asia/Jerusalem", reset_time_of_day=time(12, 0)
)
jerusalem = result.astimezone(ZoneInfo("Asia/Jerusalem"))
self.assertEqual(
(jerusalem.year, jerusalem.month, jerusalem.day), (2023, 5, 16)
)
self.assertEqual(jerusalem.hour, 12)
self.assertEqual(jerusalem.minute, 0)
def test_weekly_reset_lands_on_monday_at_offset(self):
wednesday = datetime(2023, 5, 17, 15, 45, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"7d", wednesday, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 5, 22, 12, 0, 0, tzinfo=timezone.utc))
def test_weekly_reset_today_is_monday_before_offset_is_today(self):
monday_morning = datetime(2023, 5, 22, 9, 0, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"7d", monday_morning, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 5, 22, 12, 0, 0, tzinfo=timezone.utc))
def test_weekly_reset_today_is_monday_after_offset_is_next_week(self):
monday_afternoon = datetime(2023, 5, 22, 15, 0, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"7d", monday_afternoon, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 5, 29, 12, 0, 0, tzinfo=timezone.utc))
def test_monthly_30d_lands_on_first_at_offset(self):
now = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"30d", now, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 6, 1, 12, 0, 0, tzinfo=timezone.utc))
def test_monthly_1mo_today_is_first_before_offset_is_today(self):
now = datetime(2023, 5, 1, 9, 0, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"1mo", now, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 5, 1, 12, 0, 0, tzinfo=timezone.utc))
def test_monthly_year_rollover_at_offset(self):
now = datetime(2023, 12, 15, 9, 0, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"1mo", now, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2024, 1, 1, 12, 0, 0, tzinfo=timezone.utc))
def test_custom_day_reset_applies_offset(self):
now = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc)
result = get_next_standardized_reset_time(
"3d", now, "UTC", reset_time_of_day=time(12, 0)
)
self.assertEqual(result, datetime(2023, 5, 18, 12, 0, 0, tzinfo=timezone.utc))
def test_sub_day_durations_ignore_offset(self):
base = datetime(2023, 5, 15, 15, 20, 30, tzinfo=timezone.utc)
self.assertEqual(
get_next_standardized_reset_time(
"2h", base, "UTC", reset_time_of_day=time(12, 0)
),
datetime(2023, 5, 15, 16, 0, 0, tzinfo=timezone.utc),
)
self.assertEqual(
get_next_standardized_reset_time(
"30m", base, "UTC", reset_time_of_day=time(12, 0)
),
datetime(2023, 5, 15, 15, 30, 0, tzinfo=timezone.utc),
)
def test_default_offset_is_midnight(self):
now = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc)
self.assertEqual(
get_next_standardized_reset_time("1d", now, "UTC"),
datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc),
)
if __name__ == "__main__":
unittest.main()

View file

@ -412,7 +412,7 @@ class TestAzureAnthropicMidConversationSystem:
older Claude, and a *leading* system entry 400s on every model ("messages.0:
use the top-level 'system' parameter"). These tests pin the model-aware hoist
the config applies so Claude Code sessions neither collapse the prompt cache
on 4.8+ nor hard-fail on 4.7 and older (RCA: Kraken Tech high-spend)."""
on 4.8+ nor hard-fail on 4.7 and older (RCA: customer high-spend)."""
def test_supported_model_keeps_mid_conversation_system_in_place(self, local_model_cost_map):
messages = [

View file

@ -591,7 +591,7 @@ class TestVertexAnthropicMidConversationSystem:
Claude, and a *leading* system entry 400s on every model ("messages.0: use
the top-level 'system' parameter"). These tests pin the model-aware hoist so
Claude Code sessions neither collapse the prompt cache on 4.8+ nor hard-fail
on 4.7 and older (RCA: Kraken Tech high-spend)."""
on 4.7 and older (RCA: customer high-spend)."""
def test_supported_model_keeps_mid_conversation_system_in_place(self, local_model_cost_map):
messages = [

View file

@ -0,0 +1,343 @@
"""Tests for the SSO identity assertion store (EMA subject-token capture).
Pins the contract of the store that PR 2's ``_id_jag`` subject-sourcing seam will read:
the carrier validates untyped IdP token-response values at the boundary, retention is
gated on an ``oauth2_id_jag`` server being registered, the row is encrypted at rest and
round-trips exactly, a store failure never escapes into the login path, and a salt-key
rotation re-encrypts stored rows like the sibling per-user credential tables.
"""
import json
import time
from unittest.mock import AsyncMock, MagicMock, patch
import jwt as pyjwt
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
assertion_from_sso_login,
ema_assertion_retention_enabled,
fetch_sso_identity_assertion,
persist_sso_identity_assertion,
retain_sso_identity_assertion_for_ema,
rotate_sso_identity_assertions_master_key,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.types.mcp import MCPAuth
SALT_KEY = "test-salt-key-for-sso-assertion-tests-1234"
SIGNING_KEY = "test-idp-signing-key-32-bytes-long-xxxx"
ISSUER = "https://idp.example.com"
@pytest.fixture(autouse=True)
def _set_salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY)
def _make_id_token(exp_offset: int = 3600, iss: str = ISSUER) -> str:
return pyjwt.encode(
{"iss": iss, "sub": "u1", "exp": int(time.time()) + exp_offset},
SIGNING_KEY,
algorithm="HS256",
)
def _make_prisma(stored: dict, db_has_id_jag_server: bool = False):
"""A fake prisma client whose sso-assertion table reads and writes ``stored``
(user_id -> assertion_b64), covering upsert, find_unique, find_many, and update.
``db_has_id_jag_server`` drives the retention gate's authoritative DB fallback;
it is wired explicitly so the gate never reads a truthy bare MagicMock."""
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_first = AsyncMock(
return_value=MagicMock() if db_has_id_jag_server else None
)
async def _upsert(where, data):
stored[where["user_id"]] = data["update"]["assertion_b64"]
async def _find_unique(where):
blob = stored.get(where["user_id"])
if blob is None:
return None
row = MagicMock()
row.user_id = where["user_id"]
row.assertion_b64 = blob
return row
async def _find_many():
rows = []
for user_id, blob in stored.items():
row = MagicMock()
row.user_id = user_id
row.assertion_b64 = blob
rows.append(row)
return rows
async def _update(where, data):
stored[where["user_id"]] = data["assertion_b64"]
prisma.db.litellm_ssoidentityassertion.upsert = AsyncMock(side_effect=_upsert)
prisma.db.litellm_ssoidentityassertion.find_unique = AsyncMock(side_effect=_find_unique)
prisma.db.litellm_ssoidentityassertion.find_many = AsyncMock(side_effect=_find_many)
prisma.db.litellm_ssoidentityassertion.update = AsyncMock(side_effect=_update)
return prisma
def _server_with_auth(auth_type):
server = MagicMock()
server.auth_type = auth_type
return server
def test_assertion_from_sso_login_happy_path():
token = _make_id_token()
assertion = assertion_from_sso_login(token, "rt_1")
assert assertion is not None
assert assertion.id_token.get_secret_value() == token
assert assertion.refresh_token is not None
assert assertion.refresh_token.get_secret_value() == "rt_1"
assert assertion.issuer == ISSUER
assert assertion.expires_at is not None
assert assertion.expires_at.timestamp() == pytest.approx(time.time() + 3600, abs=5)
def test_assertion_repr_never_leaks_token_material():
token = _make_id_token()
assertion = assertion_from_sso_login(token, "rt_secret_value")
rendered = repr(assertion) + str(assertion)
assert token not in rendered
assert "rt_secret_value" not in rendered
@pytest.mark.parametrize("id_token", [None, "", "not-a-jwt", 12345, ["x"], {"a": 1}])
def test_assertion_from_sso_login_rejects_unusable_id_token(id_token):
assert assertion_from_sso_login(id_token, "rt") is None
@pytest.mark.parametrize("refresh_token", [None, "", 123, ["rt"], {"rt": 1}])
def test_assertion_from_sso_login_drops_malformed_refresh_token(refresh_token):
assertion = assertion_from_sso_login(_make_id_token(), refresh_token)
assert assertion is not None
assert assertion.refresh_token is None
def test_assertion_without_exp_or_iss_still_retained():
token = pyjwt.encode({"sub": "u1"}, SIGNING_KEY, algorithm="HS256")
assertion = assertion_from_sso_login(token, None)
assert assertion is not None
assert assertion.expires_at is None
assert assertion.issuer is None
@pytest.mark.asyncio
async def test_retention_gate_requires_an_id_jag_server():
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)),
):
manager.config_mcp_servers = {
"s1": _server_with_auth(MCPAuth.oauth2),
"s2": _server_with_auth(None),
}
assert await ema_assertion_retention_enabled() is False
manager.config_mcp_servers = {
"s1": _server_with_auth(MCPAuth.oauth2),
"s2": _server_with_auth(MCPAuth.oauth2_id_jag),
}
assert await ema_assertion_retention_enabled() is True
@pytest.mark.asyncio
async def test_retention_gate_reads_the_db_when_config_declares_no_id_jag_server():
"""A DB-backed server added on another pod (or before this pod's DB load) must still enable
retention off the authoritative DB row; False only when neither authority knows one."""
with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager:
manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2)}
db_backed = _make_prisma({}, db_has_id_jag_server=True)
with patch("litellm.proxy.proxy_server.prisma_client", db_backed):
assert await ema_assertion_retention_enabled() is True
db_backed.db.litellm_mcpservertable.find_first.assert_awaited_once_with(
where={"auth_type": MCPAuth.oauth2_id_jag.value}
)
with patch("litellm.proxy.proxy_server.prisma_client", None):
assert await ema_assertion_retention_enabled() is False
@pytest.mark.asyncio
async def test_retention_gate_never_consults_the_registry_snapshot():
"""The registry is a per-process snapshot of DB state, stale in either direction: trusting
it positively would keep retaining bearer material after the last EMA server was removed on
another pod, trusting it negatively would drop writes for one added elsewhere. The gate must
judge only the config declaration and the DB row, so a stale snapshot listing an id_jag
server changes nothing."""
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)),
):
manager.config_mcp_servers = {}
manager.get_registry.return_value = {"stale": _server_with_auth(MCPAuth.oauth2_id_jag)}
assert await ema_assertion_retention_enabled() is False
manager.get_registry.assert_not_called()
@pytest.mark.asyncio
async def test_retain_persists_when_only_the_db_knows_the_id_jag_server():
stored = {}
prisma = _make_prisma(stored, db_has_id_jag_server=True)
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
patch("litellm.proxy.proxy_server.prisma_client", prisma),
):
manager.config_mcp_servers = {}
await retain_sso_identity_assertion_for_ema(
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
)
assert "user-a" in stored
@pytest.mark.asyncio
async def test_persist_and_fetch_round_trip_encrypted_at_rest():
stored = {}
prisma = _make_prisma(stored)
token = _make_id_token()
assertion = assertion_from_sso_login(token, "rt_1")
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await persist_sso_identity_assertion("user-a", assertion)
fetched = await fetch_sso_identity_assertion("user-a")
assert fetched is not None
assert fetched.id_token.get_secret_value() == token
assert fetched.refresh_token is not None
assert fetched.refresh_token.get_secret_value() == "rt_1"
assert fetched.issuer == assertion.issuer
assert fetched.expires_at == assertion.expires_at
assert token not in stored["user-a"]
assert "rt_1" not in stored["user-a"]
decrypted = decrypt_value_helper(stored["user-a"], "test", exception_type="debug")
assert json.loads(decrypted)["id_token"] == token
@pytest.mark.asyncio
async def test_persist_overwrites_previous_login():
stored = {}
prisma = _make_prisma(stored)
first = _make_id_token(exp_offset=100)
second = _make_id_token(exp_offset=7200)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await persist_sso_identity_assertion("user-a", assertion_from_sso_login(first, None))
await persist_sso_identity_assertion("user-a", assertion_from_sso_login(second, "rt_new"))
fetched = await fetch_sso_identity_assertion("user-a")
assert fetched is not None
assert fetched.id_token.get_secret_value() == second
assert fetched.refresh_token is not None
@pytest.mark.asyncio
async def test_fetch_missing_row_returns_none():
prisma = _make_prisma({})
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
assert await fetch_sso_identity_assertion("nobody") is None
@pytest.mark.asyncio
async def test_fetch_undecryptable_row_returns_none():
prisma = _make_prisma({"user-a": "not-an-encrypted-blob"})
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
assert await fetch_sso_identity_assertion("user-a") is None
@pytest.mark.asyncio
async def test_fetch_unparseable_payload_returns_none():
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
prisma = _make_prisma({"user-a": encrypt_value_helper("]]not json")})
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
assert await fetch_sso_identity_assertion("user-a") is None
@pytest.mark.asyncio
async def test_retain_noop_when_no_id_jag_server():
stored = {}
prisma = _make_prisma(stored)
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
patch("litellm.proxy.proxy_server.prisma_client", prisma),
):
manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2)}
await retain_sso_identity_assertion_for_ema(
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
)
prisma.db.litellm_ssoidentityassertion.upsert.assert_not_called()
assert stored == {}
@pytest.mark.asyncio
async def test_retain_persists_when_id_jag_server_registered():
stored = {}
prisma = _make_prisma(stored)
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
patch("litellm.proxy.proxy_server.prisma_client", prisma),
):
manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)}
await retain_sso_identity_assertion_for_ema(
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
)
assert "user-a" in stored
@pytest.mark.asyncio
async def test_retain_none_assertion_never_consults_gate_or_store():
gate = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store.ema_assertion_retention_enabled",
gate,
):
await retain_sso_identity_assertion_for_ema(user_id="user-a", assertion=None)
gate.assert_not_called()
@pytest.mark.asyncio
async def test_retain_swallows_store_failure():
prisma = MagicMock()
prisma.db.litellm_ssoidentityassertion.upsert = AsyncMock(side_effect=RuntimeError("db down"))
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
patch("litellm.proxy.proxy_server.prisma_client", prisma),
):
manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)}
await retain_sso_identity_assertion_for_ema(
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
)
@pytest.mark.asyncio
async def test_rotation_reencrypts_under_new_key(monkeypatch):
stored = {}
prisma = _make_prisma(stored)
token = _make_id_token()
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await persist_sso_identity_assertion("user-a", assertion_from_sso_login(token, None))
original_blob = stored["user-a"]
new_key = "rotated-sso-assertion-salt-key-5678"
await rotate_sso_identity_assertions_master_key(prisma_client=prisma, new_master_key=new_key)
assert stored["user-a"] != original_blob
monkeypatch.setenv("LITELLM_SALT_KEY", new_key)
decrypted = decrypt_value_helper(stored["user-a"], "test", exception_type="debug")
assert decrypted is not None
assert json.loads(decrypted)["id_token"] == token
@pytest.mark.asyncio
async def test_rotation_skips_unreadable_rows_but_rotates_readable_ones():
stored = {"good": None, "bad": "garbage-blob"}
prisma = _make_prisma(stored)
token = _make_id_token()
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await persist_sso_identity_assertion("good", assertion_from_sso_login(token, None))
good_blob_before = stored["good"]
await rotate_sso_identity_assertions_master_key(prisma_client=prisma, new_master_key="another-new-salt-key-0000")
assert stored["bad"] == "garbage-blob"
assert stored["good"] != good_blob_before

View file

@ -8031,3 +8031,120 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate():
assert resource_response["authorization_servers"] == ["https://llm.example.com/test_oauth"]
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_authorize_wall_names_the_fix_for_urlless_servers():
"""LIT-4629: the authorize wall previously said only "authorization url is not set" with no
hint that spec-only servers never discover; the detail must now name both remedies (manual
Authorization URL + Token URL, or an Issuer for RFC 8414 discovery)."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
authorize_with_server,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="urlless-wall",
name="sheets_wall",
server_name="sheets_wall",
url=None,
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/openapi.yaml",
)
mock_request = MagicMock()
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await authorize_with_server(
request=mock_request,
mcp_server=server,
client_id="client",
redirect_uri="http://localhost/callback",
)
assert exc_info.value.status_code == 400
detail_text = str(exc_info.value.detail)
assert "set Authorization URL and Token URL" in detail_text
assert "Issuer" in detail_text
@pytest.mark.asyncio
async def test_token_wall_names_the_fix_for_urlless_servers():
"""The /token wall is the second stop on the same misconfiguration (LIT-4629): after an admin
fills only the Authorization URL, the code exchange dies here; the detail must name the
remedies like the authorize wall does."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
exchange_token_with_server,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="urlless-token-wall",
name="sheets_token_wall",
server_name="sheets_token_wall",
url=None,
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/openapi.yaml",
authorization_url="https://accounts.google.com/o/oauth2/v2/auth",
)
mock_request = MagicMock()
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await exchange_token_with_server(
request=mock_request,
mcp_server=server,
grant_type="authorization_code",
code="auth-code",
redirect_uri="http://localhost/callback",
client_id="client",
client_secret=None,
code_verifier="verifier",
)
assert exc_info.value.status_code == 400
detail_text = str(exc_info.value.detail)
assert "set Token URL manually" in detail_text
assert "Issuer" in detail_text
@pytest.mark.asyncio
async def test_register_wall_names_the_fix_for_urlless_servers():
"""The /register wall serves the same missing-authorization-url 400 as authorize; its detail
must carry the same actionable remedies."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client_with_server,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="urlless-register-wall",
name="sheets_register_wall",
server_name="sheets_register_wall",
url=None,
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/openapi.yaml",
)
mock_request = MagicMock()
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await register_client_with_server(
request=mock_request,
mcp_server=server,
client_name="client",
grant_types=None,
response_types=None,
token_endpoint_auth_method=None,
)
assert exc_info.value.status_code == 400
detail_text = str(exc_info.value.detail)
assert "set Authorization URL and Token URL" in detail_text
assert "Issuer" in detail_text

View file

@ -1033,3 +1033,166 @@ class TestResolveByokMcpAuthHeader:
check_mock.assert_awaited_once_with(server, user_auth)
assert result == "caller-header"
class TestOpenApiResolvedUpstreamAuth:
"""LIT-4629: spec_path servers egress through plain httpx, so the manager's OpenAPI arm must
materialize the v2-resolved credential into the `_request_resolved_auth_headers` ContextVar;
before the fix the resolved token never reached the upstream API."""
def _oauth_server(self, **overrides: Any) -> MCPServer:
fields: Dict[str, Any] = dict(
server_id="srv-sheets",
name="google_sheets",
server_name="google_sheets",
url=None,
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/sheets-openapi.yaml",
)
fields.update(overrides)
return MCPServer(**fields)
@pytest.mark.asyncio
async def test_call_tool_openapi_injects_v2_resolved_token_contextvar(self):
"""The managed spec_path arm resolves the v2 credential and sets the ContextVar; kills
the mutant that drops the resolve_openapi_upstream_auth call in call_tool."""
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_resolved_auth_headers,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
manager = MCPServerManager()
server = self._oauth_server()
user_auth = UserAPIKeyAuth(user_id="alice", api_key="sk-user")
captured: Dict[str, Any] = {}
async def fake_openapi_handler(_server, _name, _arguments):
captured["resolved"] = _request_resolved_auth_headers.get()
return MagicMock()
with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server):
with patch.object(
manager._cred_provider,
"resolve_credentials",
new=AsyncMock(return_value=Ok(StaticHeaderAuth("Bearer stored-user-token"))),
):
with patch.object(manager, "_call_openapi_tool_handler", side_effect=fake_openapi_handler):
await manager.call_tool(
server_name=server.server_name,
name="get_values",
arguments={},
user_api_key_auth=user_auth,
)
assert captured["resolved"] == {"Authorization": "Bearer stored-user-token"}
assert _request_resolved_auth_headers.get() is None
@pytest.mark.asyncio
async def test_call_tool_openapi_m2m_missing_token_url_fails_closed(self):
"""A url-less M2M spec server with no token_url must fail with a typed error instead of
egressing unauthenticated (the pre-#32259 silent failure this arm previously preserved).
Drives the real adapter/resolver chain: ClientCredentialsConfig with missing grant fields
resolves to a misconfigured CredError, raised as an HTTPException."""
from fastapi import HTTPException
manager = MCPServerManager()
server = self._oauth_server(
oauth2_flow="client_credentials",
client_id="m2m-client",
client_secret="m2m-secret",
token_url=None,
)
called = AsyncMock()
with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server):
with patch.object(manager, "_call_openapi_tool_handler", new=called):
with pytest.raises(HTTPException):
await manager.call_tool(
server_name=server.server_name,
name="get_values",
arguments={},
user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key="sk-user"),
)
called.assert_not_awaited()
@pytest.mark.asyncio
async def test_caller_oauth2_headers_never_become_resolved_for_byok_server(self):
"""Greptile P1 regression: BYOK servers defer to v1 (to_server_spec None), and the v1 arm
must never promote caller-supplied oauth2 headers into the resolved-auth slot, where they
would override the per-server BYOK credential and leak the caller's gateway Authorization
upstream."""
manager = MCPServerManager()
server = MCPServer(
server_id="byok-spec",
name="byok_spec",
server_name="byok_spec",
url=None,
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
spec_path="https://example.com/openapi.yaml",
is_byok=True,
)
resolved, forwarded = await manager.resolve_openapi_upstream_auth(
mcp_server=server,
oauth2_headers={"Authorization": "Bearer sk-litellm-gateway-key"},
raw_headers=None,
mcp_auth_header="user-byok-key",
user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key="sk-user"),
forwarded_headers=None,
)
assert resolved is None
assert forwarded is None
@pytest.mark.asyncio
async def test_v1_server_threads_stored_headers_only_without_caller_headers(self):
"""The v1 (unmigrated) arm resolves the stored per-user token only when the caller sent no
oauth2 headers of their own; with caller headers present the stored lookup is skipped and
nothing is promoted to resolved."""
manager = MCPServerManager()
server = MCPServer(
server_id="v1-spec",
name="v1_spec",
server_name="v1_spec",
url=None,
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/openapi.yaml",
delegate_auth_to_upstream=True,
)
stored = {"Authorization": "Bearer stored-v1-token"}
user_auth = UserAPIKeyAuth(user_id="alice", api_key="sk-user")
with patch.object(
manager, "_resolve_oauth2_headers_for_tool_call", new=AsyncMock(return_value=stored)
) as lookup:
resolved, _ = await manager.resolve_openapi_upstream_auth(
mcp_server=server,
oauth2_headers=None,
raw_headers=None,
mcp_auth_header=None,
user_api_key_auth=user_auth,
forwarded_headers=None,
)
assert resolved == stored
lookup.assert_awaited_once_with(server, None, user_auth)
with patch.object(
manager, "_resolve_oauth2_headers_for_tool_call", new=AsyncMock(return_value=stored)
) as lookup:
resolved, _ = await manager.resolve_openapi_upstream_auth(
mcp_server=server,
oauth2_headers={"Authorization": "Bearer caller-supplied"},
raw_headers=None,
mcp_auth_header=None,
user_api_key_auth=user_auth,
forwarded_headers=None,
)
assert resolved is None
lookup.assert_not_awaited()

View file

@ -5597,7 +5597,7 @@ class TestMCPServerTimestamps:
async def test_build_mcp_server_from_table_persists_discovered_oauth_endpoints(self):
"""A DB-backed oauth2 server with no configured endpoints discovers them and must write
authorization_url, token_url, and scopes back to the row; otherwise the resolved values
live only in memory and one failed re-discovery serves 400 "authorization url is not set"
live only in memory and one failed re-discovery serves the 400 "authorization url is not configured"
from /authorize. registration_url must never be persisted because
_dcr_bridge_relays_client_registration keys off that column."""
manager = MCPServerManager()
@ -8891,3 +8891,140 @@ async def test_resolve_toolset_tool_permissions_single_db_fetch_across_checks():
assert first == {"server-a": ["lookup_status"]}
assert second == first
list_toolsets_mock.assert_awaited_once()
class TestMaterializeAuthHeaders:
"""_materialize_auth_headers drives one step of a resolved httpx.Auth's own flow to turn it
into a header dict for the OpenAPI egress arm, which sends plain headers and cannot carry an
httpx.Auth. Generic across auth shapes via the resolver-arm header_name convention."""
@pytest.mark.asyncio
async def test_static_header_auth_materializes_its_header(self):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_materialize_auth_headers,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
headers = await _materialize_auth_headers(StaticHeaderAuth("Bearer stored-token"))
assert headers == {"Authorization": "Bearer stored-token"}
@pytest.mark.asyncio
async def test_client_credentials_bearer_auth_materializes_bearer(self):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_materialize_auth_headers,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import (
ClientCredentialsBearerAuth,
)
async def _refetch(_stale: str):
return None
headers = await _materialize_auth_headers(ClientCredentialsBearerAuth("m2m-token", _refetch))
assert headers == {"Authorization": "Bearer m2m-token"}
@pytest.mark.asyncio
async def test_noop_and_none_materialize_to_none(self):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_materialize_auth_headers,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
)
assert await _materialize_auth_headers(None) is None
assert await _materialize_auth_headers(NoOpAuth()) is None
class TestUrllessIssuerDiscovery:
"""LIT-4629: servers with no url (OpenAPI spec_path, stdio) run no resource discovery, so
their OAuth endpoints could only ever come from manual entry; an admin-pinned issuer is a
url-independent trust anchor (RFC 8414 section 3.3) and must unlock discovery for them."""
def _urlless_row(self, **overrides):
fields = dict(
server_id="urlless-1",
alias="sheets_urlless",
description="spec-only server",
url=None,
spec_path="https://example.com/sheets-openapi.yaml",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
created_at=datetime.now(),
updated_at=datetime.now(),
)
fields.update(overrides)
return LiteLLM_MCPServerTable(**fields)
@pytest.mark.asyncio
async def test_urlless_server_with_issuer_discovers_endpoints(self):
"""The gate previously required bool(server_url), so a url-less server with an issuer
configured never ran the issuer-anchored fetch and /authorize 400d. Kills the mutant that
restores the bare bool(server_url) term."""
manager = MCPServerManager()
row = self._urlless_row(issuer="https://accounts.google.com")
resolved = MCPOAuthMetadata(
authorization_url="https://accounts.google.com/o/oauth2/v2/auth",
token_url="https://oauth2.googleapis.com/token",
)
resource_rooted = AsyncMock(return_value=None)
with (
patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved)) as anchored,
patch.object(manager, "_descovery_metadata", new=resource_rooted),
):
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
anchored.assert_awaited_once_with("https://accounts.google.com", None)
resource_rooted.assert_not_awaited()
assert built.issuer_is_anchored is True
assert built.authorization_url == "https://accounts.google.com/o/oauth2/v2/auth"
assert built.token_url == "https://oauth2.googleapis.com/token"
@pytest.mark.asyncio
async def test_urlless_server_without_issuer_stays_undiscovered(self):
"""With neither a url nor an issuer there is no discovery source; the build must not
attempt any fetch and the endpoints stay unset (manual entry remains the only path)."""
manager = MCPServerManager()
row = self._urlless_row()
anchored = AsyncMock()
resource_rooted = AsyncMock()
with (
patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=anchored),
patch.object(manager, "_descovery_metadata", new=resource_rooted),
):
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
anchored.assert_not_awaited()
resource_rooted.assert_not_awaited()
assert built.authorization_url is None
assert built.token_url is None
assert built.issuer_is_anchored is False
@pytest.mark.asyncio
async def test_urlless_obo_with_issuer_discovers_token_url(self):
"""oauth2_token_exchange is not a discovery auth type, so the plain gate relax alone
would leave a url-less OBO server undiscovered; with an issuer pinned and no configured
exchange endpoint it must resolve token_url through the issuer-anchored fetch. Kills the
mutant that drops the OBO widening from the anchor computation."""
manager = MCPServerManager()
row = self._urlless_row(
alias="obo_urlless",
auth_type=MCPAuth.oauth2_token_exchange,
issuer="https://idp.example.com",
)
resolved = MCPOAuthMetadata(token_url="https://idp.example.com/token")
resource_rooted = AsyncMock(return_value=None)
with (
patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved)) as anchored,
patch.object(manager, "_descovery_metadata", new=resource_rooted),
):
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
anchored.assert_awaited_once_with("https://idp.example.com", None)
resource_rooted.assert_not_awaited()
assert built.token_url == "https://idp.example.com/token"

View file

@ -17,6 +17,7 @@ import pytest
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
_request_resolved_auth_headers,
_resolve_param_list,
_resolve_ref,
build_input_schema,
@ -1207,3 +1208,61 @@ class TestRequestExtraHeaders:
call_args = async_client.get.call_args
headers_sent = call_args[1]["headers"]
assert "X-TOKEN" not in headers_sent
@pytest.mark.asyncio
async def test_resolved_auth_headers_win_over_every_other_authorization_source(self):
"""The gateway-resolved credential (stored per-user OAuth / minted M2M token) is
authoritative: it must override the BYOK override, static headers, and forwarded caller
headers on the Authorization name, case-insensitively, mirroring _resolve_v2_auth's rule
on the MCPClient path. Without this, a spec_path oauth2 server's completed OAuth flow
stores a token that never reaches the upstream API (LIT-4629)."""
operation = {}
func = create_tool_function(
path="/secure",
method="get",
operation=operation,
base_url="https://api.example.com",
headers={"authorization": "Bearer static-operator"},
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "secure-data")
mock_client.return_value = async_client
extra_token = _request_extra_headers.set({"Authorization": "Bearer caller-forwarded"})
auth_token = _request_auth_header.set("Bearer byok-credential")
resolved_token = _request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"})
try:
result = await func()
finally:
_request_auth_header.reset(auth_token)
_request_extra_headers.reset(extra_token)
_request_resolved_auth_headers.reset(resolved_token)
assert result == "secure-data"
headers_sent = async_client.get.call_args[1]["headers"]
authorization_values = [v for k, v in headers_sent.items() if k.lower() == "authorization"]
assert authorization_values == ["Bearer resolved-oauth"]
@pytest.mark.asyncio
async def test_resolved_auth_headers_not_leaked_between_calls(self):
"""After resetting the resolved-auth ContextVar, subsequent calls send no credential."""
operation = {}
func = create_tool_function(
path="/data",
method="get",
operation=operation,
base_url="https://api.example.com",
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
token = _request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"})
_request_resolved_auth_headers.reset(token)
await func()
headers_sent = async_client.get.call_args[1]["headers"]
assert "Authorization" not in headers_sent

View file

@ -218,3 +218,86 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable():
assert exc.value.status_code == 503
pre_call.assert_not_awaited()
handle_local.assert_not_awaited()
@pytest.mark.asyncio
async def test_openapi_local_tool_injects_resolved_oauth_token():
"""LIT-4629: the local-registry (OpenAPI) dispatch is the primary egress for spec_path
tools, and before the fix it dropped the gateway-resolved OAuth credential entirely, so a
user's completed OAuth flow stored a token that never reached the upstream API. The resolved
credential must land in the `_request_resolved_auth_headers` ContextVar the tool closure
reads. Kills the mutant that deletes the resolve_openapi_upstream_auth call in server.py."""
from litellm.proxy._experimental.mcp_server import server as mcp_module
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_resolved_auth_headers,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
user = UserAPIKeyAuth(
api_key="sk-user",
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
oauth_server = MCPServer(
server_id="srv-sheets",
name="google_sheets",
server_name="google_sheets",
url=None,
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/sheets-openapi.yaml",
)
fake_tool = MagicMock()
fake_tool.name = "get_values"
captured: dict = {}
async def handle_local(_name, _arguments):
captured["resolved"] = _request_resolved_auth_headers.get()
return []
with (
patch.object(
mcp_module.global_mcp_server_manager,
"_get_mcp_server_from_tool_name",
return_value=oauth_server,
),
patch.object(
mcp_module.global_mcp_server_manager,
"pre_call_tool_check",
new=AsyncMock(return_value={}),
),
patch.object(
mcp_module.global_mcp_tool_registry,
"get_tool",
return_value=fake_tool,
),
patch.object(
mcp_module.global_mcp_server_manager._cred_provider,
"resolve_credentials",
new=AsyncMock(return_value=Ok(StaticHeaderAuth("Bearer stored-user-token"))),
),
patch(
"litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool",
new=handle_local,
),
patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
return_value=True,
),
):
await mcp_module.execute_mcp_tool(
name="get_values",
arguments={},
allowed_mcp_servers=[oauth_server],
start_time=datetime.now(timezone.utc),
user_api_key_auth=user,
)
assert captured["resolved"] == {"Authorization": "Bearer stored-user-token"}
assert _request_resolved_auth_headers.get() is None

View file

@ -1,10 +1,14 @@
"""Unit tests for the pure merge logic in litellm/proxy/a2a/agent_card.py."""
import pytest
from litellm.proxy.a2a.agent_card import (
LITELLM_A2A_PROTOCOL_VERSION,
LITELLM_SECURITY_REQUIREMENTS,
LITELLM_SECURITY_SCHEMES,
merge_agent_card,
normalize_protocol_version,
resolve_served_protocol_version,
)
PROXY_URL = "https://proxy.example/a2a/agent-xyz"
@ -205,3 +209,54 @@ def test_strips_additional_interfaces_to_prevent_backend_url_leak():
]
merged = merge_agent_card(upstream, proxy_url=PROXY_URL, proxy_base_url=PROXY_BASE)
assert "additionalInterfaces" not in merged
@pytest.mark.parametrize(
("raw", "expected"),
[
("0.3", "0.3"),
("0.3.0", "0.3"),
("1.0", "1.0"),
("1.0.0", "1.0"),
("1.0.1", "1.0"),
("0.3.0-rc1", "0.3"),
("1.0.0-rc.1+build.5", "1.0"),
("0.2.6", None),
("2.0", None),
("0.30", None),
("0.3.garbage", None),
("0.3.", None),
("1.0.not-semver", None),
("0.3.0.0", None),
("0.3-rc1", None),
("garbage", None),
("", None),
(None, None),
(1.0, None),
],
)
def test_normalize_protocol_version(raw, expected):
assert normalize_protocol_version(raw) == expected
def test_resolve_served_protocol_version_canonicalizes_semver_pins():
assert resolve_served_protocol_version({"protocolVersion": "0.3.0"}) == "0.3"
assert resolve_served_protocol_version({"protocolVersion": "1.0.0"}) == "1.0"
assert resolve_served_protocol_version({"protocolVersion": "0.3"}) == "0.3"
assert resolve_served_protocol_version({"protocolVersion": "1.0"}) == "1.0"
def test_resolve_served_protocol_version_falls_back_for_unsupported():
assert (
resolve_served_protocol_version({"protocolVersion": "0.2.6"})
== LITELLM_A2A_PROTOCOL_VERSION
)
assert resolve_served_protocol_version(None) == LITELLM_A2A_PROTOCOL_VERSION
def test_serves_semver_pinned_protocol_version_as_major_minor():
card = _full_upstream_card()
card["protocolVersion"] = "0.3.0"
merged = merge_agent_card(card, proxy_url=PROXY_URL, proxy_base_url=PROXY_BASE)
assert merged["protocolVersion"] == "0.3"
assert merged["supportedInterfaces"][0]["protocolVersion"] == "0.3"

View file

@ -313,3 +313,13 @@ def test_agent_card_with_0_3_pin_and_supported_interfaces_is_lowered():
def test_agent_card_same_version_passthrough():
card = _extended_card_1_0()
assert normalize_agent_card(card, "1.0") is card
def test_detect_card_version_normalizes_semver_protocol_version():
from litellm.proxy.a2a.version_convert import _detect_card_version
assert _detect_card_version({"protocolVersion": "1.0.0"}) == "1.0"
assert (
_detect_card_version({"protocolVersion": "0.3.0", "supportedInterfaces": []})
== "0.3"
)

View file

@ -540,6 +540,53 @@ class TestAgentRBACProxyAdmin:
assert resp.status_code == 200
class TestAgentProtocolVersionValidation:
"""Registration accepts spec-default semver protocolVersion values and still
rejects genuinely unsupported versions."""
@pytest.fixture(autouse=True)
def _setup(self, monkeypatch):
self.admin_client = _make_app_with_role(LitellmUserRoles.PROXY_ADMIN)
self.mock_registry = MagicMock()
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", self.mock_registry)
def _create_agent_with_protocol_version(self, protocol_version: str):
config = _sample_agent_config()
config["agent_card_params"]["protocolVersion"] = protocol_version
with patch("litellm.proxy.proxy_server.prisma_client"):
self.mock_registry.get_agent_by_name = MagicMock(return_value=None)
self.mock_registry.add_agent_to_db = AsyncMock(
return_value=_sample_agent_response()
)
self.mock_registry.register_agent = MagicMock()
return self.admin_client.post(
"/v1/agents",
json=config,
headers={"Authorization": "Bearer k"},
)
def test_semver_protocol_version_registers_and_stores_major_minor(self):
resp = self._create_agent_with_protocol_version("0.3.0")
assert resp.status_code == 200
stored_card = self.mock_registry.add_agent_to_db.await_args.kwargs["agent"][
"agent_card_params"
]
assert stored_card["protocolVersion"] == "0.3"
assert stored_card["supportedInterfaces"][0]["protocolVersion"] == "0.3"
def test_unsupported_protocol_version_is_rejected(self):
resp = self._create_agent_with_protocol_version("0.2.6")
assert resp.status_code == 400
assert "Unsupported protocolVersion '0.2.6'" in resp.json()["detail"]
self.mock_registry.add_agent_to_db.assert_not_awaited()
def test_malformed_protocol_version_is_rejected(self):
resp = self._create_agent_with_protocol_version("0.3.garbage")
assert resp.status_code == 400
assert "Unsupported protocolVersion '0.3.garbage'" in resp.json()["detail"]
self.mock_registry.add_agent_to_db.assert_not_awaited()
class TestCheckAgentManagementPermission:
"""Unit tests for the _check_agent_management_permission helper."""

View file

@ -4562,6 +4562,9 @@ async def test_temp_budget_increase_applied_for_cached_key():
Seed the auth cache with a key whose spend (5.0) exceeds its original
max_budget (2.0) but is under the effective budget (2.0 + 100.0). The cache-hit
request must not raise and the resolved token must carry max_budget == 102.0.
Resolving twice must yield 102.0 both times and leave the cached object at the
original 2.0: the increase is derived per request, never compounded or persisted.
"""
from datetime import datetime, timedelta
@ -4607,14 +4610,22 @@ async def test_temp_budget_increase_applied_for_cached_key():
new_callable=AsyncMock,
),
):
result = await _user_api_key_auth_builder(
request=mock_request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={"model": "gpt-4o-mini"},
results = tuple(
[
await _user_api_key_auth_builder(
request=mock_request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={"model": "gpt-4o-mini"},
)
for _ in range(2)
]
)
assert result.max_budget == 102.0
assert all(result.max_budget == 102.0 for result in results)
cached_after = await user_api_key_cache.async_get_cache(key=hashed_token)
assert cached_after.max_budget == 2.0

View file

@ -5,25 +5,23 @@ import sys
import time
import types
from datetime import datetime, timedelta, timezone
from datetime import time as dt_time
from typing import Any, Dict, List
from unittest.mock import AsyncMock, MagicMock
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path
from litellm._logging import verbose_proxy_logger
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings
from litellm.proxy.utils import ProxyLogging
# Mock classes for testing
class MockLiteLLMTeamMembership:
async def update_many(
self, where: Dict[str, Any], data: Dict[str, Any]
) -> Dict[str, Any]:
async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]:
# Mock the update_many method for litellm_teammembership
return {"count": 1}
@ -32,9 +30,7 @@ class MockLiteLLMVerificationToken:
def __init__(self):
self.update_many_calls: List[Dict[str, Any]] = []
async def update_many(
self, where: Dict[str, Any], data: Dict[str, Any]
) -> Dict[str, Any]:
async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]:
self.update_many_calls.append({"where": where, "data": data})
return {"count": 1}
@ -52,9 +48,7 @@ class MockLiteLLMOrganizationTable:
self.find_many_calls.append({"where": where})
return self._find_many_results
async def update_many(
self, where: Dict[str, Any], data: Dict[str, Any]
) -> Dict[str, Any]:
async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]:
self.update_many_calls.append({"where": where, "data": data})
return {"count": 1}
@ -72,9 +66,7 @@ class MockLiteLLMTagTable:
self.find_many_calls.append({"where": where})
return self._find_many_results
async def update_many(
self, where: Dict[str, Any], data: Dict[str, Any]
) -> Dict[str, Any]:
async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]:
self.update_many_calls.append({"where": where, "data": data})
return {"count": 1}
@ -110,9 +102,7 @@ class MockBatcher:
_self._outer = outer
def update(_self, where, data):
_self._outer.calls.append(
{"table": _self._table_name, "where": where, "data": data}
)
_self._outer.calls.append({"table": _self._table_name, "where": where, "data": data})
self.litellm_verificationtoken = _Table("key", self)
self.litellm_usertable = _Table("user", self)
@ -172,11 +162,7 @@ class MockPrismaClient:
return [item for item in data if hasattr(item, "budget_reset_at")]
# Handle specific filtering for enduser table queries
if (
table_name == "enduser"
and query_type == "find_all"
and "budget_id_list" in kwargs
):
if table_name == "enduser" and query_type == "find_all" and "budget_id_list" in kwargs:
budget_id_list = kwargs["budget_id_list"]
# Return endusers that match the budget IDs
return [
@ -188,11 +174,7 @@ class MockPrismaClient:
]
# Handle key queries with expires and reset_at
if (
table_name == "key"
and query_type == "find_all"
and ("expires" in kwargs or "reset_at" in kwargs)
):
if table_name == "key" and query_type == "find_all" and ("expires" in kwargs or "reset_at" in kwargs):
return [item for item in data if hasattr(item, "budget_reset_at")]
return data
@ -227,9 +209,7 @@ def mock_proxy_logging():
@pytest.fixture
def reset_budget_job(mock_prisma_client, mock_proxy_logging):
return ResetBudgetJob(
proxy_logging_obj=mock_proxy_logging, prisma_client=mock_prisma_client
)
return ResetBudgetJob(proxy_logging_obj=mock_proxy_logging, prisma_client=mock_prisma_client)
# Helper function to run async tests
@ -270,6 +250,40 @@ def test_reset_budget_for_key(reset_budget_job, mock_prisma_client):
assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
def test_reset_budget_for_key_honors_injected_reset_time(mock_prisma_client, mock_proxy_logging):
"""Injected BudgetResetSettings drives the written reset time end to end (DI, no globals).
Before the configurable-reset-time change this wrote a midnight reset_at (hour 0);
with noon injected it must write a noon reset_at.
"""
job = ResetBudgetJob(
proxy_logging_obj=mock_proxy_logging,
prisma_client=mock_prisma_client,
reset_settings=BudgetResetSettings(timezone="UTC", reset_time_of_day=dt_time(12, 0)),
)
now = datetime.now(timezone.utc)
test_key = type(
"LiteLLM_VerificationToken",
(),
{
"spend": 100.0,
"budget_duration": "1d",
"budget_reset_at": now,
"id": "test-key-noon",
"token": "tok-noon",
},
)
mock_prisma_client.data["key"] = [test_key]
asyncio.run(job.reset_budget_for_litellm_keys())
key_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "key"]
assert len(key_writes) == 1
reset_at = key_writes[0]["data"]["budget_reset_at"].astimezone(timezone.utc)
assert reset_at.hour == 12
assert reset_at.minute == 0
def test_reset_budget_for_user(reset_budget_job, mock_prisma_client):
# Setup test data with timezone-aware datetime
now = datetime.now(timezone.utc)
@ -486,11 +500,7 @@ def test_reset_budget_for_keys_linked_to_budgets(reset_budget_job, mock_prisma_c
budgets_to_reset = [test_budget]
# Run the method
asyncio.run(
reset_budget_job.reset_budget_for_keys_linked_to_budgets(
budgets_to_reset=budgets_to_reset
)
)
asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset))
# Verify that update_many was called on litellm_verificationtoken
calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls
@ -531,11 +541,7 @@ def test_reset_budget_for_keys_linked_to_budgets_excludes_keys_with_own_budget_d
budgets_to_reset = [test_budget]
asyncio.run(
reset_budget_job.reset_budget_for_keys_linked_to_budgets(
budgets_to_reset=budgets_to_reset
)
)
asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset))
calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls
assert len(calls) == 1
@ -548,17 +554,13 @@ def test_reset_budget_for_keys_linked_to_budgets_excludes_keys_with_own_budget_d
assert call["where"]["budget_id"] == {"in": ["7d-budget-tier"]}
def test_reset_budget_for_keys_linked_to_budgets_empty(
reset_budget_job, mock_prisma_client
):
def test_reset_budget_for_keys_linked_to_budgets_empty(reset_budget_job, mock_prisma_client):
"""
Test that when there are no budgets to reset, no update is performed
on the verification token table.
"""
# Run with empty list
asyncio.run(
reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=[])
)
asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=[]))
# Verify no update_many calls were made
calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls
@ -584,11 +586,7 @@ def test_reset_budget_for_orgs_linked_to_budgets(reset_budget_job, mock_prisma_c
},
)
asyncio.run(
reset_budget_job.reset_budget_for_orgs_linked_to_budgets(
budgets_to_reset=[test_budget]
)
)
asyncio.run(reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[test_budget]))
calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls
assert len(calls) == 1
@ -598,16 +596,12 @@ def test_reset_budget_for_orgs_linked_to_budgets(reset_budget_job, mock_prisma_c
assert call["data"]["spend"] == 0
def test_reset_budget_for_orgs_linked_to_budgets_empty(
reset_budget_job, mock_prisma_client
):
def test_reset_budget_for_orgs_linked_to_budgets_empty(reset_budget_job, mock_prisma_client):
"""
Test that when there are no budgets to reset, no update is performed
on the organization table.
"""
asyncio.run(
reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[])
)
asyncio.run(reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[]))
calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls
assert len(calls) == 0
@ -631,11 +625,7 @@ def test_reset_budget_for_tags_linked_to_budgets(reset_budget_job, mock_prisma_c
},
)
asyncio.run(
reset_budget_job.reset_budget_for_tags_linked_to_budgets(
budgets_to_reset=[test_budget]
)
)
asyncio.run(reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[test_budget]))
calls = mock_prisma_client.db.litellm_tagtable.update_many_calls
assert len(calls) == 1
@ -645,16 +635,12 @@ def test_reset_budget_for_tags_linked_to_budgets(reset_budget_job, mock_prisma_c
assert call["data"]["spend"] == 0
def test_reset_budget_for_tags_linked_to_budgets_empty(
reset_budget_job, mock_prisma_client
):
def test_reset_budget_for_tags_linked_to_budgets_empty(reset_budget_job, mock_prisma_client):
"""
Test that when there are no budgets to reset, no update is performed
on the tag table.
"""
asyncio.run(
reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[])
)
asyncio.run(reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[]))
calls = mock_prisma_client.db.litellm_tagtable.update_many_calls
assert len(calls) == 0
@ -668,9 +654,7 @@ def test_reset_budget_for_tags_linked_to_budgets_empty(
],
ids=["30d-calendar-month", "1mo-calendar-month", "1d-next-midnight"],
)
def test_reset_budget_reset_at_date_calendar_aligned(
budget_duration, expected_day, expected_month
):
def test_reset_budget_reset_at_date_calendar_aligned(budget_duration, expected_day, expected_month):
"""
Verify that _reset_budget_reset_at_date produces calendar-aligned reset
times (matching get_budget_reset_time), not sliding-window offsets.
@ -694,7 +678,7 @@ def test_reset_budget_reset_at_date_calendar_aligned(
with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt:
mock_dt.now.return_value = fixed_now
mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs)
asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now))
asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings()))
assert test_budget.budget_reset_at.day == expected_day
assert test_budget.budget_reset_at.month == expected_month
@ -724,7 +708,7 @@ def test_reset_budget_reset_at_date_7d_next_monday():
with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt:
mock_dt.now.return_value = fixed_now
mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs)
asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now))
asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings()))
# Next Monday after Wednesday June 14 is June 19
assert test_budget.budget_reset_at.day == 19
@ -749,7 +733,7 @@ def test_reset_budget_reset_at_date_none_duration():
},
)
asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, now))
asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, now, BudgetResetSettings()))
assert test_budget.budget_reset_at == original_reset_at
@ -773,7 +757,7 @@ def test_reset_budget_reset_at_date_none_reset_at():
with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt:
mock_dt.now.return_value = fixed_now
mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs)
asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now))
asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings()))
# Should be set to 1st of next month (July 1)
assert test_budget.budget_reset_at is not None
@ -781,9 +765,7 @@ def test_reset_budget_reset_at_date_none_reset_at():
assert test_budget.budget_reset_at.month == 7
def test_budget_table_reset_also_resets_linked_keys(
reset_budget_job, mock_prisma_client
):
def test_budget_table_reset_also_resets_linked_keys(reset_budget_job, mock_prisma_client):
"""
Integration-style test: when reset_budget_for_litellm_budget_table runs,
it should also reset spend for keys linked to the expiring budget tiers
@ -818,9 +800,7 @@ def test_budget_table_reset_also_resets_linked_keys(
assert calls[0]["data"]["spend"] == 0
def test_budget_table_reset_also_resets_linked_orgs(
reset_budget_job, mock_prisma_client
):
def test_budget_table_reset_also_resets_linked_orgs(reset_budget_job, mock_prisma_client):
"""
Integration-style test: when reset_budget_for_litellm_budget_table runs,
it should also reset spend for orgs linked to the expiring budget tiers
@ -853,9 +833,7 @@ def test_budget_table_reset_also_resets_linked_orgs(
assert calls[0]["data"]["spend"] == 0
def test_budget_table_reset_also_resets_linked_tags(
reset_budget_job, mock_prisma_client
):
def test_budget_table_reset_also_resets_linked_tags(reset_budget_job, mock_prisma_client):
"""
Integration-style test: when reset_budget_for_litellm_budget_table runs,
it should also reset spend for tags linked to the expiring budget tiers.
@ -887,9 +865,7 @@ def test_budget_table_reset_also_resets_linked_tags(
assert calls[0]["data"]["spend"] == 0
def test_reset_budget_resets_endusers_with_null_budget_id(
reset_budget_job, mock_prisma_client
):
def test_reset_budget_resets_endusers_with_null_budget_id(reset_budget_job, mock_prisma_client):
"""
When litellm.max_end_user_budget_id is configured and that budget is
being reset, end users with budget_id=NULL should also have their spend
@ -959,17 +935,13 @@ def test_reset_budget_resets_endusers_with_null_budget_id(
mock_prisma_client.data["enduser"] = [enduser_with_budget]
# Set up the DB mock for NULL-budget-id end users
mock_prisma_client.db.litellm_endusertable.set_find_many_results(
[enduser_no_budget_row]
)
mock_prisma_client.db.litellm_endusertable.set_find_many_results([enduser_no_budget_row])
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
# Both end users should have been reset
updated = mock_prisma_client.updated_data["enduser"]
assert (
len(updated) == 2
), f"Expected 2 endusers reset (1 explicit + 1 implicit), got {len(updated)}"
assert len(updated) == 2, f"Expected 2 endusers reset (1 explicit + 1 implicit), got {len(updated)}"
user_ids = {u.user_id for u in updated}
assert "enduser-explicit" in user_ids
@ -986,9 +958,7 @@ def test_reset_budget_resets_endusers_with_null_budget_id(
litellm.max_end_user_budget_id = None
def test_reset_budget_skips_null_budget_id_endusers_when_default_not_configured(
reset_budget_job, mock_prisma_client
):
def test_reset_budget_skips_null_budget_id_endusers_when_default_not_configured(reset_budget_job, mock_prisma_client):
"""
When litellm.max_end_user_budget_id is NOT configured, end users with
budget_id=NULL should NOT be fetched or reset.
@ -1073,20 +1043,14 @@ def test_reset_budget_for_team_members_preserves_total_spend():
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_teammembership.update_many = AsyncMock(
return_value={"count": 1}
)
mock_prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(
proxy_logging_obj=MagicMock(), prisma_client=mock_prisma_client
)
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=mock_prisma_client)
asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget]))
mock_prisma_client.db.litellm_teammembership.update_many.assert_called_once()
call_kwargs = (
mock_prisma_client.db.litellm_teammembership.update_many.call_args.kwargs
)
call_kwargs = mock_prisma_client.db.litellm_teammembership.update_many.call_args.kwargs
assert call_kwargs["where"]["budget_id"]["in"] == ["budget-1"]
assert call_kwargs["data"] == {"spend": 0}
assert "total_spend" not in call_kwargs["data"]
@ -1142,9 +1106,7 @@ def test_reset_budget_windows_uses_is_not_null_filter(monkeypatch):
raises `MissingRequiredValueError`. We work around it by using `query_raw`
with `IS NOT NULL`. If someone reverts to the ORM filter, this test fails.
"""
job, prisma_client, _ = _make_reset_budget_windows_job(
monkeypatch, key_rows=[], team_rows=[]
)
job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=[], team_rows=[])
asyncio.run(job.reset_budget_windows())
@ -1184,15 +1146,11 @@ def test_reset_budget_windows_resets_expired_key_window(monkeypatch):
# The `budget_limits` payload is re-serialized JSON with a bumped reset_at.
written_windows = json.loads(call_kwargs["data"]["budget_limits"])
assert len(written_windows) == 1
new_reset_at = datetime.fromisoformat(
written_windows[0]["reset_at"].replace("Z", "+00:00")
).replace(tzinfo=None)
new_reset_at = datetime.fromisoformat(written_windows[0]["reset_at"].replace("Z", "+00:00")).replace(tzinfo=None)
assert new_reset_at > now
# The spend counter for this key+window was cleared.
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:key:sk-expired:window:1d", value=0.0
)
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-expired:window:1d", value=0.0)
def test_reset_budget_windows_skips_unexpired_key_window(monkeypatch):
@ -1206,9 +1164,7 @@ def test_reset_budget_windows_skips_unexpired_key_window(monkeypatch):
"budget_limits": [{"budget_duration": "1d", "reset_at": future}],
}
]
job, prisma_client, _ = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[])
asyncio.run(job.reset_budget_windows())
@ -1237,9 +1193,7 @@ def test_reset_budget_windows_resets_expired_team_window(monkeypatch):
assert call_kwargs["where"] == {"team_id": "team-expired"}
assert "budget_limits" in call_kwargs["data"]
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:team:team-expired:window:30d", value=0.0
)
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team:team-expired:window:30d", value=0.0)
def test_reset_budget_windows_handles_string_budget_limits(monkeypatch):
@ -1252,14 +1206,10 @@ def test_reset_budget_windows_handles_string_budget_limits(monkeypatch):
key_rows = [
{
"token": "sk-string-limits",
"budget_limits": json.dumps(
[{"budget_duration": "1d", "reset_at": expired}]
),
"budget_limits": json.dumps([{"budget_duration": "1d", "reset_at": expired}]),
}
]
job, prisma_client, _ = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[])
asyncio.run(job.reset_budget_windows())
@ -1274,9 +1224,7 @@ def test_reset_budget_windows_skips_row_with_empty_budget_limits(monkeypatch):
{"token": "sk-empty-list", "budget_limits": []},
{"token": "sk-empty-str", "budget_limits": ""},
]
job, prisma_client, _ = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[])
asyncio.run(job.reset_budget_windows())
@ -1361,27 +1309,17 @@ def test_reset_budget_for_team_members_invalidates_redis_counter(monkeypatch):
)
prisma_client = MagicMock()
prisma_client.db.litellm_teammembership.find_many = AsyncMock(
return_value=[membership]
)
prisma_client.db.litellm_teammembership.update_many = AsyncMock(
return_value={"count": 1}
)
prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership])
prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget]))
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:team_member:alice:team-x", value=0.0, ttl=60
)
counter_cache.redis_cache.async_set_cache.assert_any_await(
key="spend:team_member:alice:team-x", value=0.0, ttl=60
)
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team_member:alice:team-x", value=0.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:team_member:alice:team-x", value=0.0, ttl=60)
def test_reset_budget_for_keys_invalidates_redis_counter(
reset_budget_job, mock_prisma_client, monkeypatch
):
def test_reset_budget_for_keys_invalidates_redis_counter(reset_budget_job, mock_prisma_client, monkeypatch):
"""Key budget reset must clear the Redis spend counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
@ -1402,14 +1340,10 @@ def test_reset_budget_for_keys_invalidates_redis_counter(
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:key:sk-abc", value=0.0, ttl=60
)
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-abc", value=0.0, ttl=60)
def test_reset_budget_for_users_invalidates_redis_counter(
reset_budget_job, mock_prisma_client, monkeypatch
):
def test_reset_budget_for_users_invalidates_redis_counter(reset_budget_job, mock_prisma_client, monkeypatch):
"""User budget reset must clear the Redis spend counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
@ -1430,14 +1364,10 @@ def test_reset_budget_for_users_invalidates_redis_counter(
asyncio.run(reset_budget_job.reset_budget_for_litellm_users())
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:user:alice", value=0.0, ttl=60
)
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:user:alice", value=0.0, ttl=60)
def test_reset_budget_for_teams_invalidates_redis_counter(
reset_budget_job, mock_prisma_client, monkeypatch
):
def test_reset_budget_for_teams_invalidates_redis_counter(reset_budget_job, mock_prisma_client, monkeypatch):
"""Team budget reset must clear the Redis spend counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
@ -1458,9 +1388,7 @@ def test_reset_budget_for_teams_invalidates_redis_counter(
asyncio.run(reset_budget_job.reset_budget_for_litellm_teams())
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:team:team-x", value=0.0, ttl=60
)
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team:team-x", value=0.0, ttl=60)
def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch):
@ -1511,9 +1439,7 @@ def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch):
batcher.commit = failing_commit
prisma_client.db.batch_ = MagicMock(return_value=batcher)
job = ResetBudgetJob(
proxy_logging_obj=MockProxyLogging(), prisma_client=prisma_client
)
job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_litellm_keys())
@ -1543,8 +1469,8 @@ def test_reset_budget_for_keys_writes_only_spend_and_reset_at(reset_budget_job,
"budget_duration": "30d",
"budget_reset_at": now,
"token": "sk-problematic",
"object_permission_id": "perm-abc", # would be rejected on update
"budget_limits": [{"max_budget": 5}], # would be rejected on update
"object_permission_id": "perm-abc", # would be rejected on update
"budget_limits": [{"max_budget": 5}], # would be rejected on update
"metadata": {"some": "thing"},
},
)
@ -1570,19 +1496,13 @@ def test_reset_budget_for_keys_linked_to_budgets_invalidates_redis_counter(monke
linked_key = type("Key", (), {"token": "sk-linked"})
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[linked_key]
)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 1}
)
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[linked_key])
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget]))
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:key:sk-linked", value=0.0, ttl=60
)
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-linked", value=0.0, ttl=60)
def test_reset_budget_for_orgs_linked_to_budgets_invalidates_redis_counter(monkeypatch):
@ -1593,22 +1513,14 @@ def test_reset_budget_for_orgs_linked_to_budgets_invalidates_redis_counter(monke
linked_org = type("Org", (), {"organization_id": "org-acme"})
prisma_client = MagicMock()
prisma_client.db.litellm_organizationtable.find_many = AsyncMock(
return_value=[linked_org]
)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 1}
)
prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[linked_org])
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_orgs_linked_to_budgets([expired_budget]))
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:org:org-acme", value=0.0, ttl=60
)
counter_cache.redis_cache.async_set_cache.assert_any_await(
key="spend:org:org-acme", value=0.0, ttl=60
)
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:org:org-acme", value=0.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:org:org-acme", value=0.0, ttl=60)
def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monkeypatch):
@ -1625,12 +1537,8 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monke
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget]))
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:tag:tenant-42", value=0.0, ttl=60
)
counter_cache.redis_cache.async_set_cache.assert_any_await(
key="spend:tag:tenant-42", value=0.0, ttl=60
)
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:tag:tenant-42", value=0.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:tag:tenant-42", value=0.0, ttl=60)
def test_reset_budget_for_tags_linked_to_budgets_invalidates_management_cache(
@ -1657,9 +1565,7 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_management_cache(
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget]))
counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(
key="tag:tenant-42"
)
counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="tag:tenant-42")
def test_reset_budget_for_tags_linked_to_budgets_invalidates_each_tag_management_cache(
@ -1684,8 +1590,7 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_each_tag_management
asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget]))
deleted_keys = {
call.kwargs.get("key")
for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list
call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list
}
assert deleted_keys == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"}
@ -1711,19 +1616,13 @@ def test_reset_budget_for_keys_linked_to_budgets_invalidates_management_cache(
linked_key = type("Key", (), {"token": "sk-linked"})
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[linked_key]
)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 1}
)
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[linked_key])
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget]))
counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(
key="sk-linked"
)
counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="sk-linked")
def test_reset_budget_for_orgs_linked_to_budgets_invalidates_management_cache(
@ -1736,19 +1635,14 @@ def test_reset_budget_for_orgs_linked_to_budgets_invalidates_management_cache(
linked_org = type("Org", (), {"organization_id": "org-acme"})
prisma_client = MagicMock()
prisma_client.db.litellm_organizationtable.find_many = AsyncMock(
return_value=[linked_org]
)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 1}
)
prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[linked_org])
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_orgs_linked_to_budgets([expired_budget]))
deleted_keys = {
call.kwargs.get("key")
for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list
call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list
}
assert deleted_keys == {
"org_id:org-acme",
@ -1768,19 +1662,13 @@ def test_reset_budget_for_team_members_invalidates_management_cache(monkeypatch)
)
prisma_client = MagicMock()
prisma_client.db.litellm_teammembership.find_many = AsyncMock(
return_value=[membership]
)
prisma_client.db.litellm_teammembership.update_many = AsyncMock(
return_value={"count": 1}
)
prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership])
prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1})
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget]))
counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(
key="team-x_alice"
)
counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="team-x_alice")
def test_reset_budget_for_tags_linked_to_budgets_management_cache_delete_failure_still_resets(
@ -1788,9 +1676,7 @@ def test_reset_budget_for_tags_linked_to_budgets_management_cache_delete_failure
):
"""If ``async_delete_cache`` raises, the DB cascade must still complete."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
counter_cache.user_api_key_cache.async_delete_cache = AsyncMock(
side_effect=RuntimeError("cache unavailable")
)
counter_cache.user_api_key_cache.async_delete_cache = AsyncMock(side_effect=RuntimeError("cache unavailable"))
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_tag = type("Tag", (), {"tag_name": "tenant-42"})

View file

@ -1,19 +1,33 @@
import os
import sys
from datetime import datetime, timezone
from datetime import datetime, time, timezone
from zoneinfo import ZoneInfo
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.proxy.common_utils.timezone_utils import (
BudgetResetSettings,
compute_budget_reset_at,
get_budget_reset_settings,
get_budget_reset_time,
get_budget_reset_timezone,
parse_budget_reset_time,
)
def _restore_attr(obj, name, original):
if original is None:
if hasattr(obj, name):
delattr(obj, name)
else:
setattr(obj, name, original)
def test_get_budget_reset_time():
"""
Test that the budget reset time is set to the first of the next month
@ -100,3 +114,69 @@ def test_get_budget_reset_time_respects_timezone():
delattr(litellm, "timezone")
else:
litellm.timezone = original
def test_parse_budget_reset_time_hh_mm():
assert parse_budget_reset_time("12:00") == time(12, 0)
def test_parse_budget_reset_time_hh_mm_ss():
assert parse_budget_reset_time("09:30:15") == time(9, 30, 15)
def test_parse_budget_reset_time_unset_defaults_to_midnight():
assert parse_budget_reset_time(None) == time(0, 0)
assert parse_budget_reset_time("") == time(0, 0)
def test_parse_budget_reset_time_invalid_string_raises():
with pytest.raises(ValueError):
parse_budget_reset_time("25:00")
with pytest.raises(ValueError):
parse_budget_reset_time("noon")
def test_parse_budget_reset_time_non_string_raises():
# Unquoted "12:00" in YAML parses to the int 720; it must fail loudly,
# not silently fall back to midnight.
with pytest.raises(ValueError):
parse_budget_reset_time(720)
def test_get_budget_reset_settings_reads_globals():
orig_tz = getattr(litellm, "timezone", None)
orig_rt = getattr(litellm, "budget_reset_time", None)
try:
litellm.timezone = "Asia/Jerusalem"
litellm.budget_reset_time = "12:00"
settings = get_budget_reset_settings()
assert settings.timezone == "Asia/Jerusalem"
assert settings.reset_time_of_day == time(12, 0)
finally:
_restore_attr(litellm, "timezone", orig_tz)
_restore_attr(litellm, "budget_reset_time", orig_rt)
def test_compute_budget_reset_at_applies_offset():
settings = BudgetResetSettings(
timezone="Asia/Jerusalem", reset_time_of_day=time(12, 0)
)
reset_at = compute_budget_reset_at("1d", settings)
jerusalem = reset_at.astimezone(ZoneInfo("Asia/Jerusalem"))
assert jerusalem.hour == 12
assert jerusalem.minute == 0
assert reset_at > datetime.now(timezone.utc)
def test_get_budget_reset_time_honors_global_budget_reset_time():
orig_tz = getattr(litellm, "timezone", None)
orig_rt = getattr(litellm, "budget_reset_time", None)
try:
litellm.timezone = "UTC"
litellm.budget_reset_time = "12:00"
reset_at = get_budget_reset_time(budget_duration="1d")
assert reset_at.astimezone(timezone.utc).hour == 12
assert reset_at.astimezone(timezone.utc).minute == 0
finally:
_restore_attr(litellm, "timezone", orig_tz)
_restore_attr(litellm, "budget_reset_time", orig_rt)

View file

@ -347,8 +347,8 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp
# Step 3: Create a user via SCIM
scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
userName="idontexist@krakentest.tech",
emails=[SCIMUserEmail(value="idontexist@krakentest.tech")],
userName="idontexist@example.com",
emails=[SCIMUserEmail(value="idontexist@example.com")],
)
mock_prisma_client = mocker.MagicMock()
@ -364,7 +364,7 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp
new_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.new_user",
AsyncMock(return_value=NewUserRequest(user_id="idontexist@krakentest.tech")),
AsyncMock(return_value=NewUserRequest(user_id="idontexist@example.com")),
)
mocker.patch(

View file

@ -14994,3 +14994,76 @@ async def test_list_keys_without_expires_param_forwards_none():
mock_helper.assert_called_once()
assert mock_helper.call_args.kwargs["expires_filter"] is None
@pytest.mark.asyncio
@patch(
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_sso_identity_assertions_master_key"
)
@patch(
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_user_env_vars_master_key"
)
@patch(
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_user_credentials_master_key"
)
@patch(
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_server_credentials_master_key"
)
async def test_rotate_master_key_rotates_sso_identity_assertions(
mock_rotate_mcp_server,
mock_rotate_mcp_user,
mock_rotate_env_vars,
mock_rotate_sso,
):
"""Master-key rotation must re-encrypt the SSO identity assertion store alongside
the sibling per-user encrypted tables, or a salt rotation orphans every stored
assertion (step 4d)."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
_rotate_master_key,
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_tx = AsyncMock()
mock_tx.litellm_proxymodeltable = MagicMock()
mock_tx.litellm_proxymodeltable.delete_many = AsyncMock()
mock_tx.litellm_proxymodeltable.create_many = AsyncMock()
mock_prisma_client.db.tx = MagicMock(
return_value=AsyncMock(
__aenter__=AsyncMock(return_value=mock_tx),
__aexit__=AsyncMock(return_value=False),
)
)
mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(
return_value=[]
)
mock_proxy_config = MagicMock()
mock_proxy_config.decrypt_model_list_from_db.return_value = []
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="test-user",
)
with patch(
"litellm.proxy.proxy_server.proxy_config",
mock_proxy_config,
):
await _rotate_master_key(
prisma_client=mock_prisma_client,
user_api_key_dict=user_api_key_dict,
current_master_key="sk-old-master-key",
new_master_key="sk-new-master-key",
)
mock_rotate_sso.assert_awaited_once_with(
prisma_client=mock_prisma_client,
new_master_key="sk-new-master-key",
)

View file

@ -1458,7 +1458,7 @@ async def test_get_generic_sso_response_with_additional_headers():
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
# Act
result, received_response, _ = await get_generic_sso_response(
result, received_response, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id=generic_client_id,
@ -1522,7 +1522,7 @@ async def test_get_generic_sso_response_with_empty_headers():
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
# Act
result, received_response, _ = await get_generic_sso_response(
result, received_response, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id=generic_client_id,
@ -2893,6 +2893,7 @@ class TestCLIKeyRegenerationFlow:
prefill_user_code=None,
result=mock_result,
received_response=None,
sso_assertion=None,
)
@pytest.mark.asyncio
@ -2933,6 +2934,7 @@ class TestCLIKeyRegenerationFlow:
prefill_user_code="WXYZ-2345",
result=mock_result,
received_response=None,
sso_assertion=None,
)
def test_get_redirect_url_does_not_include_existing_key_in_url(self):
@ -7019,7 +7021,7 @@ class TestPKCEStateCookieBinding:
):
jwt_handler = MagicMock(spec=JWTHandler)
jwt_handler.get_team_ids_from_jwt.return_value = []
result, _, _ = await get_generic_sso_response(
result, _, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=jwt_handler,
generic_client_id="cid",
@ -7078,7 +7080,7 @@ async def test_debug_sso_callback_renders_full_jwt_claims():
}
async def fake_get_generic_sso_response(**kwargs):
return parsed_openid, raw_userinfo_with_leaked_token, access_token_payload
return parsed_openid, raw_userinfo_with_leaked_token, access_token_payload, None
with (
patch.dict(
@ -7374,3 +7376,266 @@ async def test_auth_callback_without_oauth_error_proceeds_to_normal_flow():
assert exc_info.value.status_code == 500
assert "DB not connected" in str(exc_info.value.detail)
# ── SSO identity assertion capture + persist wiring (EMA) ─────────────────────
def _ema_id_token(sub: str = "u1") -> str:
import time as _time
import jwt as _pyjwt
return _pyjwt.encode(
{"iss": "https://idp.example.com", "sub": sub, "exp": int(_time.time()) + 3600},
"test-idp-signing-key-32-bytes-long-xxxx",
algorithm="HS256",
)
@pytest.mark.asyncio
async def test_pkce_arm_captures_sso_assertion():
"""The PKCE token exchange strips bearer fields from received_response for safety;
the typed assertion carrier must still capture id_token + refresh_token."""
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
get_generic_sso_response,
)
id_token = _ema_id_token()
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "matched-state", "code": "auth-code"}
mock_request.cookies = {"litellm_oauth_state": "matched-state"}
with (
patch.object(
SSOAuthenticationHandler,
"prepare_token_exchange_parameters",
AsyncMock(
return_value={
"code_verifier": "verifier",
"_pkce_cache_key": "pkce_verifier:matched-state",
}
),
),
patch.object(
SSOAuthenticationHandler,
"_pkce_token_exchange",
AsyncMock(
return_value={
"access_token": "tok",
"id_token": id_token,
"refresh_token": "rt_from_idp",
"sub": "user@example.com",
"email": "user@example.com",
}
),
),
patch.object(SSOAuthenticationHandler, "_delete_pkce_verifier", AsyncMock()),
patch("fastapi_sso.sso.base.DiscoveryDocument"),
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "x",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
"GENERIC_CLIENT_USE_PKCE": "true",
},
),
):
jwt_handler = MagicMock(spec=JWTHandler)
jwt_handler.get_team_ids_from_jwt.return_value = []
result, received_response, _, sso_assertion = await get_generic_sso_response(
request=mock_request,
jwt_handler=jwt_handler,
generic_client_id="cid",
redirect_url="https://proxy.example.com/sso/callback",
sso_jwt_handler=None,
)
assert sso_assertion is not None
assert sso_assertion.id_token.get_secret_value() == id_token
assert sso_assertion.refresh_token is not None
assert sso_assertion.refresh_token.get_secret_value() == "rt_from_idp"
# The sanitized received_response must still not carry bearer material.
assert "id_token" not in (received_response or {})
assert "refresh_token" not in (received_response or {})
@pytest.mark.asyncio
async def test_verify_and_process_arm_captures_sso_assertion():
"""The non-PKCE generic arm reads the raw bearer fields off the fastapi-sso client."""
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
id_token = _ema_id_token()
mock_request = MagicMock(spec=Request)
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
mock_sso_instance = MagicMock()
mock_sso_instance.verify_and_process = AsyncMock(
return_value={"sub": "u1", "email": "u@example.com"}
)
mock_sso_instance.access_token = None
mock_sso_instance.id_token = id_token
mock_sso_instance.refresh_token = "rt_from_idp"
mock_sso_class = MagicMock(return_value=mock_sso_instance)
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "test_secret",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo",
},
):
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch(
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
_, _, _, sso_assertion = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id="test_client_id",
redirect_url="http://test.com/callback",
sso_jwt_handler=None,
)
assert sso_assertion is not None
assert sso_assertion.id_token.get_secret_value() == id_token
assert sso_assertion.refresh_token is not None
assert sso_assertion.refresh_token.get_secret_value() == "rt_from_idp"
@pytest.mark.asyncio
async def test_redirect_from_openid_persists_assertion_under_canonical_user_id():
"""The browser funnel persists the captured assertion AFTER canonical user
resolution, keyed by the user_id admission will later resolve (the key-generation
response user_id), not the raw IdP subject."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
assertion_from_sso_login,
)
assertion = assertion_from_sso_login(_ema_id_token(), "rt_1")
assert assertion is not None
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://localhost:4000/"
mock_request.cookies = {}
retain_mock = AsyncMock()
with (
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.premium_user", False),
patch("litellm.proxy.proxy_server.user_custom_sso", None),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch(
"litellm.proxy.proxy_server.generate_key_helper_fn",
AsyncMock(
return_value={"token": "sk-ui-key", "user_id": "canonical-user-id"}
),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.check_and_update_if_proxy_admin_id",
AsyncMock(return_value="internal_user"),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
retain_mock,
),
):
response = await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=CustomOpenID(
id="raw-idp-subject",
email="u@example.com",
first_name="U",
last_name="Ser",
display_name="U Ser",
provider="generic",
team_ids=[],
user_role=None,
),
request=mock_request,
received_response=None,
generic_client_id="cid",
ui_access_mode=None,
access_token_payload=None,
jwt_handler=None,
sso_assertion=assertion,
)
retain_mock.assert_awaited_once_with(
user_id="canonical-user-id", assertion=assertion
)
assert response is not None
@pytest.mark.asyncio
async def test_cli_completion_persists_assertion_under_db_user_id():
"""The CLI funnel persists the captured assertion under the DB-resolved user_id."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
assertion_from_sso_login,
)
from litellm.proxy.management_endpoints.ui_sso import (
_complete_cli_sso_callback_session,
)
assertion = assertion_from_sso_login(_ema_id_token(), None)
assert assertion is not None
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://localhost:4000/"
user_info = MagicMock()
user_info.user_id = "cli-user-id"
user_info.user_role = "internal_user"
user_info.models = []
user_info.teams = []
retain_mock = AsyncMock()
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
AsyncMock(return_value=user_info),
),
patch(
"litellm.proxy.management_endpoints.ui_sso._fetch_cli_sso_team_details",
AsyncMock(return_value=[]),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.build_cli_sso_attribution_metadata",
return_value={},
),
patch(
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
retain_mock,
),
):
response = await _complete_cli_sso_callback_session(
request=mock_request,
key="cli-login-id",
flow={},
result={"sub": "raw-idp-subject"},
parsed_openid_result={
"user_id": "raw-idp-subject",
"user_email": "u@example.com",
"user_role": None,
},
user_defined_values=None,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
sso_assertion=assertion,
)
retain_mock.assert_awaited_once_with(user_id="cli-user-id", assertion=assertion)
assert response.status_code == 200

View file

@ -8,6 +8,7 @@ Pins covered:
from __future__ import annotations
import json
import os
from types import SimpleNamespace
from typing import Any, Dict
@ -407,6 +408,124 @@ async def test_ProxyConfig_save_config_invalid_path_raises(monkeypatch):
await pc.save_config({"x": 1})
@pytest.mark.asyncio
async def test_ProxyConfig_save_config_db_omits_environment_variables_by_default(monkeypatch):
"""A save_config after get_config() (which resolves os.environ/ placeholders
to plaintext and merges the environment_variables section) must not snapshot
those env vars into the DB config row. Persisting them would make a stale DB
row shadow YAML/container env on every subsequent restart."""
mock_prisma = MagicMock()
mock_prisma.insert_data = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
# a valid salt so the env-var encryption path (reached only if the pop
# regresses) runs cleanly, making this fail on the assertion below rather
# than on an incidental encryption crash
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-key")
pc = ProxyConfig()
cfg = {
"model_list": [{"model_name": "gpt-4o"}],
"litellm_settings": {"success_callback": ["langfuse"]},
"environment_variables": {"OPENAI_API_KEY": "sk-from-yaml"},
}
await pc.save_config(cfg)
mock_prisma.insert_data.assert_awaited_once()
written = mock_prisma.insert_data.await_args.kwargs["data"]
assert "environment_variables" not in written
# unrelated sections are still persisted; model_list is stripped as before
assert written["litellm_settings"] == {"success_callback": ["langfuse"]}
assert "model_list" not in written
# the caller's dict is not mutated (save_config works on a copy)
assert cfg["environment_variables"] == {"OPENAI_API_KEY": "sk-from-yaml"}
@pytest.mark.asyncio
async def test_ProxyConfig_save_config_db_persists_environment_variables_when_opted_in(monkeypatch):
"""The explicit opt-in path (include_env_vars=True) still persists env vars,
encrypted, so the dedicated config-update flow can write them."""
mock_prisma = MagicMock()
mock_prisma.insert_data = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-key")
pc = ProxyConfig()
cfg = {"litellm_settings": {}, "environment_variables": {"OPENAI_API_KEY": "sk-explicit"}}
await pc.save_config(cfg, include_env_vars=True)
mock_prisma.insert_data.assert_awaited_once()
written = mock_prisma.insert_data.await_args.kwargs["data"]
assert set(written["environment_variables"].keys()) == {"OPENAI_API_KEY"}
# value is encrypted at rest, not the plaintext it came in as
assert written["environment_variables"]["OPENAI_API_KEY"] != "sk-explicit"
def _install_fake_config_repo(monkeypatch, existing_row):
"""Route ProxyConfig's ConfigRepository through an in-memory fake that
records the value written to the environment_variables row."""
captured: dict = {}
class _FakeTable:
async def find_first(self, where):
return SimpleNamespace(param_value=existing_row) if existing_row is not None else None
async def upsert(self, where, data):
captured["value"] = json.loads(data["update"]["param_value"])
class _FakeRepo:
def __init__(self, client):
self.table = _FakeTable()
monkeypatch.setattr("litellm.proxy.proxy_server.ConfigRepository", _FakeRepo)
monkeypatch.setattr("litellm.proxy.proxy_server.invalidate_config_param", AsyncMock())
return captured
@pytest.mark.asyncio
async def test_ProxyConfig_save_environment_variables_merges_sets_and_deletes(monkeypatch):
"""The per-key env-var write updates/deletes only the named keys and leaves
every other stored key untouched, so an unrelated env var is never lost or
snapshotted."""
captured = _install_fake_config_repo(
monkeypatch,
existing_row={"EXISTING_KEY": "ciphertext-existing", "UI_LOGO_PATH": "old-logo", "LITELLM_FAVICON_URL": "old"},
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-key")
pc = ProxyConfig()
await pc.save_environment_variables({"UI_LOGO_PATH": "new-logo", "LITELLM_FAVICON_URL": None})
written = captured["value"]
# unrelated key preserved byte-for-byte
assert written["EXISTING_KEY"] == "ciphertext-existing"
# set key updated and encrypted (not the plaintext)
assert "UI_LOGO_PATH" in written and written["UI_LOGO_PATH"] != "new-logo"
# None-valued key deleted
assert "LITELLM_FAVICON_URL" not in written
@pytest.mark.asyncio
async def test_ProxyConfig_save_environment_variables_noop_without_db(monkeypatch):
"""With no DB configured the per-key write must do nothing (never touch the
config repository)."""
captured = _install_fake_config_repo(monkeypatch, existing_row={})
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
pc = ProxyConfig()
await pc.save_environment_variables({"UI_LOGO_PATH": "x"})
assert "value" not in captured
# ---------------------------------------------------------------------------
# ProxyConfig._check_for_os_environ_vars
# ---------------------------------------------------------------------------
@ -950,6 +1069,32 @@ async def test_ProxyConfig_load_config_wires_general_settings_url_validation(tmp
litellm.provider_url_destination_allowed_hosts = original_provider_hosts
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_wires_config_reload_interval(tmp_path, monkeypatch):
"""general_settings.proxy_config_reload_interval_seconds must reach the proxy_server
module global that schedules the DB config-reload jobs, so operators can tune multi-pod
convergence from config.yaml."""
import litellm.proxy.proxy_server as proxy_server
f = tmp_path / "c.yaml"
f.write_text(
"model_list: []\n"
"general_settings:\n"
" proxy_config_reload_interval_seconds: 47\n"
"litellm_settings: {}\n"
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
original = proxy_server.proxy_config_reload_interval_seconds
try:
await ProxyConfig().load_config(router=None, config_file_path=str(f))
assert proxy_server.proxy_config_reload_interval_seconds == 47
finally:
proxy_server.proxy_config_reload_interval_seconds = original
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_missing_file_raises(monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)

View file

@ -13,6 +13,7 @@ Routes covered:
from __future__ import annotations
import json
from unittest.mock import AsyncMock, MagicMock
from .conftest import VOLATILE_KEYS, normalize
@ -473,6 +474,83 @@ def test_config_list_happy_admin(client, auth_as, mock_prisma, monkeypatch):
}
def test_config_list_exposes_config_reload_interval(client, auth_as, mock_prisma, monkeypatch):
"""proxy_config_reload_interval_seconds must surface in the admin UI general-settings
list as an Integer field defaulting to 30, so operators can tune multi-pod convergence
from the dashboard."""
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
row = MagicMock()
row.param_value = {}
table.find_first = AsyncMock(return_value=row)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.get("/config/list", params={"config_type": "general_settings"})
assert response.status_code == 200
by_name = {entry["field_name"]: entry for entry in response.json()}
assert "proxy_config_reload_interval_seconds" in by_name
entry = by_name["proxy_config_reload_interval_seconds"]
assert entry["field_type"] == "Integer"
assert entry["field_default_value"] == 30
def test_config_field_update_accepts_config_reload_interval(client, auth_as, mock_prisma, monkeypatch):
"""POST /config/field/update accepts proxy_config_reload_interval_seconds and persists
it to the DB general_settings row for all pods to pick up."""
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
table.find_first = AsyncMock(return_value=None)
upsert_row = {
"param_name": "general_settings",
"param_value": {"proxy_config_reload_interval_seconds": 45},
"id": "row-1",
}
table.upsert = AsyncMock(return_value=upsert_row)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/field/update",
json={
"field_name": "proxy_config_reload_interval_seconds",
"field_value": 45,
"config_type": "general_settings",
},
)
assert response.status_code == 200
upserted = table.upsert.call_args.kwargs["data"]["create"]["param_value"]
assert json.loads(upserted)["proxy_config_reload_interval_seconds"] == 45
def test_config_field_update_rejects_non_positive_config_reload_interval(client, auth_as, mock_prisma, monkeypatch):
"""A non-positive proxy_config_reload_interval_seconds from the UI is rejected with a 400
and never persisted, since APScheduler requires a positive interval."""
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
table.find_first = AsyncMock(return_value=None)
table.upsert = AsyncMock()
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/field/update",
json={
"field_name": "proxy_config_reload_interval_seconds",
"field_value": 0,
"config_type": "general_settings",
},
)
assert response.status_code == 400
table.upsert.assert_not_called()
def test_config_list_non_admin_rejected(client, auth_as, mock_prisma, monkeypatch):
"""Non-admin gets a 400 with the role embedded in the error message."""
from litellm.proxy import proxy_server as ps

View file

@ -751,6 +751,97 @@ async def test_initialize_scheduled_jobs_credentials(monkeypatch):
assert len(mock_scheduler_calls) > 0
@pytest.mark.asyncio
async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval(monkeypatch):
"""
The DB config-reload jobs (add_deployment, get_credentials) that keep multi-pod
deployments in sync must be scheduled at the configured
proxy_config_reload_interval_seconds, not a hardcoded value.
"""
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
from litellm.proxy.proxy_server import ProxyStartupEvent
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_config = AsyncMock()
mock_scheduler = MagicMock()
configured_interval = 47
with (
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
patch("litellm.proxy.proxy_server.store_model_in_db", True),
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
patch(
"litellm.proxy.proxy_server.proxy_config_reload_interval_seconds",
configured_interval,
),
patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=mock_scheduler),
):
await ProxyStartupEvent.initialize_scheduled_background_jobs(
general_settings={},
prisma_client=mock_prisma_client,
proxy_budget_rescheduler_min_time=1,
proxy_budget_rescheduler_max_time=2,
proxy_batch_write_at=5,
proxy_logging_obj=mock_proxy_logging,
)
scheduled_seconds = {
job_call.kwargs["id"]: job_call.kwargs.get("seconds")
for job_call in mock_scheduler.add_job.call_args_list
if "id" in job_call.kwargs
}
assert scheduled_seconds["add_deployment_job"] == configured_interval
assert scheduled_seconds["get_credentials_job"] == configured_interval
@pytest.mark.asyncio
async def test_initialize_scheduled_jobs_rejects_non_positive_config_reload_interval(monkeypatch):
"""
A non-positive proxy_config_reload_interval_seconds (misconfig via env/config/DB) would
make APScheduler reject the job and crash startup, so the scheduler must fall back to the
30s default instead of forwarding the bad value.
"""
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
from litellm.proxy.proxy_server import ProxyStartupEvent
from litellm.proxy.utils import ProxyLogging
mock_prisma_client = MagicMock()
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_config = AsyncMock()
mock_scheduler = MagicMock()
with (
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
patch("litellm.proxy.proxy_server.store_model_in_db", True),
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
patch("litellm.proxy.proxy_server.proxy_config_reload_interval_seconds", 0),
patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=mock_scheduler),
):
await ProxyStartupEvent.initialize_scheduled_background_jobs(
general_settings={},
prisma_client=mock_prisma_client,
proxy_budget_rescheduler_min_time=1,
proxy_budget_rescheduler_max_time=2,
proxy_batch_write_at=5,
proxy_logging_obj=mock_proxy_logging,
)
scheduled_seconds = {
job_call.kwargs["id"]: job_call.kwargs.get("seconds")
for job_call in mock_scheduler.add_job.call_args_list
if "id" in job_call.kwargs
}
assert scheduled_seconds["add_deployment_job"] == 30
assert scheduled_seconds["get_credentials_job"] == 30
@pytest.mark.asyncio
async def test_initialize_scheduled_jobs_hydrates_mcp_when_store_model_in_db_false(monkeypatch):
"""

View file

@ -53,6 +53,7 @@ def mock_proxy_config(monkeypatch):
# Add a counter to track save_config calls
save_config_call_count = 0
saved_env_updates: list = []
async def mock_save_config(new_config=None):
nonlocal mock_config, save_config_call_count
@ -61,13 +62,22 @@ def mock_proxy_config(monkeypatch):
mock_config = new_config
return mock_config
async def mock_save_environment_variables(updates):
saved_env_updates.append(updates)
from litellm.proxy.proxy_server import proxy_config
monkeypatch.setattr(proxy_config, "get_config", mock_get_config)
monkeypatch.setattr(proxy_config, "save_config", mock_save_config)
monkeypatch.setattr(proxy_config, "save_environment_variables", mock_save_environment_variables)
# Return both the config and the call counter
return {"config": mock_config, "save_call_count": lambda: save_config_call_count}
# Return the config, the save_config call counter, and any env-var updates
# the endpoint routed through the dedicated save_environment_variables path
return {
"config": mock_config,
"save_call_count": lambda: save_config_call_count,
"env_updates": lambda: saved_env_updates,
}
@pytest.fixture
@ -840,11 +850,18 @@ class TestProxySettingEndpoints:
assert data["status"] == "success"
assert data["theme_config"]["logo_url"] == "https://example.com/new-logo.png"
# Verify config was updated
updated_config = mock_proxy_config["config"]
assert "UI_LOGO_PATH" in updated_config["environment_variables"]
# The logo path is applied to the live process immediately
assert os.environ["UI_LOGO_PATH"] == "https://example.com/new-logo.png"
assert mock_proxy_config["save_call_count"]() == 1
# env vars are persisted through the dedicated per-key path, and ONLY
# the two keys this endpoint owns are touched. The unrelated SSO env
# vars in the merged config are never snapshotted.
env_updates = mock_proxy_config["env_updates"]()
assert env_updates == [
{"UI_LOGO_PATH": "https://example.com/new-logo.png", "LITELLM_FAVICON_URL": None}
]
def test_update_ui_theme_settings_with_favicon(
self, mock_proxy_config, mock_auth, monkeypatch
):
@ -869,13 +886,15 @@ class TestProxySettingEndpoints:
== "https://example.com/custom-favicon.ico"
)
updated_config = mock_proxy_config["config"]
assert "UI_LOGO_PATH" in updated_config["environment_variables"]
assert "LITELLM_FAVICON_URL" in updated_config["environment_variables"]
assert (
updated_config["environment_variables"]["LITELLM_FAVICON_URL"]
== "https://example.com/custom-favicon.ico"
)
assert os.environ["UI_LOGO_PATH"] == "https://example.com/new-logo.png"
assert os.environ["LITELLM_FAVICON_URL"] == "https://example.com/custom-favicon.ico"
# Only the two owned keys are persisted, both with their new values
assert mock_proxy_config["env_updates"]() == [
{
"UI_LOGO_PATH": "https://example.com/new-logo.png",
"LITELLM_FAVICON_URL": "https://example.com/custom-favicon.ico",
}
]
def test_update_ui_theme_settings_clear_favicon(
self, mock_proxy_config, mock_auth, monkeypatch

View file

@ -3,11 +3,15 @@ Unit tests for per-deployment num_retries in litellm_params
GitHub Issue: #18968 - Per-deployment max_retries/num_retries in litellm_params is not used in retry logic
"""
import httpx
import pytest
import pytest_asyncio
from unittest.mock import patch
import litellm
from litellm import Router
from litellm.types.router import RetryPolicy
from litellm.integrations.custom_logger import CustomLogger
class TestPerDeploymentNumRetries:
@ -319,3 +323,255 @@ class TestNumRetriesNoneGuard:
# 1 initial attempt + at least 1 retry -> proves None fell back to a positive int
assert calls["n"] >= 2
class TestNoProviderRetryAmplification:
"""
A routed request must reach the upstream provider exactly ``1 + <router retries>``
times. The Router is the sole retry owner for routed calls, so the provider SDK
must never retry on top of it. Otherwise a per-deployment ``num_retries`` set in
``litellm_params`` is applied twice - once by the Router loop and once as the
provider client's ``max_retries`` - turning one request into ``(1 + num_retries) ** 2``
upstream requests.
These tests count actual upstream HTTP requests through the full Router completion
path by injecting a counting transport via ``litellm.aclient_session`` (the
documented seam the OpenAI client builder reads), so both Router-level and any
provider-SDK-level retries are observed.
"""
@staticmethod
def _install_counting_upstream() -> dict:
"""Route every upstream POST to a 500 and count it. ``retry-after: 0`` keeps
provider-SDK backoff at zero so a mutated (double-retrying) build stays fast."""
counter = {"n": 0}
def handler(request: httpx.Request) -> httpx.Response:
counter["n"] += 1
return httpx.Response(
500,
headers={"retry-after": "0"},
json={"error": {"message": "boom", "type": "server_error"}},
)
litellm.aclient_session = httpx.AsyncClient(transport=httpx.MockTransport(handler))
return counter
@pytest_asyncio.fixture(autouse=True)
async def _isolate_clients(self):
litellm.in_memory_llm_clients_cache.flush_cache()
yield
session = litellm.aclient_session
litellm.aclient_session = None
litellm.in_memory_llm_clients_cache.flush_cache()
if session is not None:
await session.aclose()
@staticmethod
def _router(api_base: str, litellm_params: dict, **router_kwargs) -> Router:
params = {"model": "openai/gpt-4o-mini", "api_base": api_base, "api_key": "sk-fake"}
params.update(litellm_params)
return Router(model_list=[{"model_name": "mock", "litellm_params": params}], **router_kwargs)
async def _call_and_count(self, router: Router, **call_kwargs) -> int:
counter = self._install_counting_upstream()
with patch("asyncio.sleep", return_value=None):
with pytest.raises(litellm.InternalServerError):
await router.acompletion(
model="mock", messages=[{"role": "user", "content": "hi"}], **call_kwargs
)
return counter["n"]
@pytest.mark.asyncio
@pytest.mark.parametrize("num_retries", [2, 5])
async def test_deployment_num_retries_sends_no_extra_provider_requests(self, num_retries):
"""
Deployment ``num_retries=N`` (every attempt failing) must send exactly ``N + 1``
upstream requests, not ``(N + 1) ** 2``. This is the amplification regression:
an unfixed build sends 9 (N=2) or 36 (N=5).
"""
counter = self._install_counting_upstream()
router = self._router(
f"https://amp-{num_retries}.local/v1", {"num_retries": num_retries}, num_retries=1
)
with patch("asyncio.sleep", return_value=None):
with pytest.raises(litellm.InternalServerError):
await router.acompletion(model="mock", messages=[{"role": "user", "content": "hi"}])
assert counter["n"] == num_retries + 1
@pytest.mark.asyncio
async def test_request_max_retries_does_not_nest_with_router_retries(self):
"""
A request-body ``max_retries`` must not make the provider SDK retry on top of the
Router. With deployment ``num_retries=5`` and request ``max_retries=3`` the count
stays ``6``; a build that lets either value reach the provider SDK sends 24 or 36.
"""
router = self._router("https://nest-req.local/v1", {"num_retries": 5}, num_retries=1)
assert await self._call_and_count(router, max_retries=3) == 6
@pytest.mark.asyncio
async def test_deployment_max_retries_does_not_nest_with_router_retries(self):
"""
A deployment-level ``max_retries`` is likewise never applied on top of the Router's
retries for a routed call: deployment ``num_retries=5`` plus ``max_retries=3`` still
sends exactly ``6`` upstream requests.
"""
router = self._router(
"https://nest-dep.local/v1", {"num_retries": 5, "max_retries": 3}, num_retries=1
)
assert await self._call_and_count(router) == 6
@pytest.mark.asyncio
async def test_retry_policy_configured_does_not_reintroduce_amplification(self):
"""
With a retry policy configured alongside a per-deployment ``num_retries=5``, the
provider SDK still must not retry: exactly ``6`` upstream requests, not 36.
"""
router = self._router(
"https://policy.local/v1",
{"num_retries": 5},
num_retries=1,
retry_policy=RetryPolicy(InternalServerErrorRetries=2),
)
assert await self._call_and_count(router) == 6
@pytest.mark.asyncio
async def test_global_num_retries_not_amplified(self):
"""
Global ``num_retries`` (no per-deployment setting) already behaves correctly and
must stay that way: ``num_retries=3`` sends ``4`` upstream requests.
"""
router = self._router("https://global.local/v1", {}, num_retries=3)
assert await self._call_and_count(router) == 4
@pytest.mark.asyncio
async def test_direct_completion_still_forwards_num_retries_to_provider(self):
"""
For a NON-routed direct ``litellm.acompletion`` call, ``num_retries`` remains an
alias for the provider client's ``max_retries`` (the instructor use case). The
provider SDK therefore retries in addition to litellm's own retry wrapper, so the
upstream count exceeds ``num_retries + 1`` - proving the routed-call fix did not
change direct-call behaviour.
"""
counter = self._install_counting_upstream()
num_retries = 2
with patch("asyncio.sleep", return_value=None):
with pytest.raises(litellm.InternalServerError):
await litellm.acompletion(
model="openai/gpt-4o-mini",
api_base="https://direct.local/v1",
api_key="sk-fake",
messages=[{"role": "user", "content": "hi"}],
num_retries=num_retries,
)
assert counter["n"] > num_retries + 1
class _AttemptCounter(CustomLogger):
"""Counts upstream call attempts via the pre-call hook (one per attempt)."""
def __init__(self):
self.attempts = 0
def log_pre_api_call(self, model, messages, kwargs):
self.attempts += 1
class TestRequestNumRetriesBeatsGlobal:
"""
A per-request num_retries (request body or the x-litellm-num-retries header, both of
which arrive as the num_retries kwarg) must take precedence over the global
litellm.num_retries (litellm_settings.num_retries on the proxy) during retry handling.
The regression: the @client wrapper stamped the global litellm.num_retries onto the
raised exception, and async_function_with_retries then adopted that stamped value,
overwriting the request-level num_retries it had already resolved. This exercises the
real retry loop end to end (the failing call flows through the wrapped litellm.acompletion),
which the kwargs-merge-only test above does not.
"""
@pytest.fixture(autouse=True)
def _restore_litellm_globals(self):
prev_num_retries = litellm.num_retries
prev_callbacks = litellm.callbacks
yield
litellm.num_retries = prev_num_retries
litellm.callbacks = prev_callbacks
@staticmethod
def _router(global_num_retries):
return Router(
model_list=[
{
"model_name": "mock",
"litellm_params": {
"model": "openai/mock",
"api_key": "sk-fake",
"mock_response": "litellm.InternalServerError",
},
}
],
num_retries=global_num_retries,
)
async def _count_attempts(self, *, global_num_retries, request_num_retries):
counter = _AttemptCounter()
litellm.callbacks = [counter]
litellm.num_retries = global_num_retries
router = self._router(global_num_retries)
kwargs = {"model": "mock", "messages": [{"role": "user", "content": "hi"}]}
if request_num_retries is not None:
kwargs["num_retries"] = request_num_retries
with patch("asyncio.sleep", return_value=None):
with pytest.raises(litellm.InternalServerError):
await router.acompletion(**kwargs)
return counter.attempts
@pytest.mark.asyncio
async def test_request_num_retries_overrides_global(self):
"""global=3 + request=1 -> 2 attempts (1 initial + 1 retry), not 4 (1 + global 3)."""
attempts = await self._count_attempts(global_num_retries=3, request_num_retries=1)
assert attempts == 2
@pytest.mark.asyncio
async def test_request_num_retries_zero_disables_retries_despite_global(self):
"""global=3 + request=0 -> a single attempt (retries disabled by the request)."""
attempts = await self._count_attempts(global_num_retries=3, request_num_retries=0)
assert attempts == 1
@pytest.mark.asyncio
async def test_global_num_retries_applies_when_request_omits_it(self):
"""No request num_retries -> the global still applies: 1 initial + 3 retries = 4."""
attempts = await self._count_attempts(global_num_retries=3, request_num_retries=None)
assert attempts == 4
@pytest.mark.asyncio
async def test_deployment_num_retries_reaches_wrapper_when_no_request_value(self):
"""
With no request value and the router default at 0, a deployment's
litellm_params.num_retries reaches the wrapped call, is carried on the raised
exception, and is applied: deployment 2 -> 1 initial + 2 retries = 3 (not 1).
"""
counter = _AttemptCounter()
litellm.callbacks = [counter]
litellm.num_retries = None
router = Router(
model_list=[
{
"model_name": "mock",
"litellm_params": {
"model": "openai/mock",
"api_key": "sk-fake",
"mock_response": "litellm.InternalServerError",
"num_retries": 2,
},
}
],
num_retries=0,
)
with patch("asyncio.sleep", return_value=None):
with pytest.raises(litellm.InternalServerError):
await router.acompletion(
model="mock", messages=[{"role": "user", "content": "hi"}]
)
assert counter.attempts == 3

View file

@ -26,7 +26,8 @@ export const menuLabelToPage: Record<string, Page> = {
"Cost Tracking": Page.CostTracking,
"UI Theme": Page.UiTheme,
// Experimental submenu items
Caching: Page.Caching,
"Response Cache": Page.Caching,
Caching: Page.Caching, // Legacy label support
Prompts: Page.Prompts,
Budgets: Page.Budgets,
"API Playground": Page.TransformRequest,

View file

@ -26,6 +26,7 @@ IS_CI="${CI:-false}"
CONTAINER_NAME="litellm-e2e-postgres-$$"
MOCK_PID=""
PROXY_PID=""
PROXY_LOG=""
# --- Ensure common tool paths are available (local dev only) ---
if [ "$IS_CI" = "false" ]; then
@ -40,6 +41,7 @@ cleanup() {
echo "Cleaning up..."
[ -n "$MOCK_PID" ] && kill "$MOCK_PID" 2>/dev/null || true
[ -n "$PROXY_PID" ] && kill "$PROXY_PID" 2>/dev/null || true
[ -n "$PROXY_LOG" ] && rm -f "$PROXY_LOG" || true
if [ "$IS_CI" = "false" ]; then
docker stop "$CONTAINER_NAME" 2>/dev/null || true
fi
@ -124,6 +126,7 @@ echo "UI build copied and restructured"
# --- Python environment ---
echo "=== Setting up Python environment ==="
cd "$REPO_ROOT"
export UV_PYTHON="${UV_PYTHON:-3.13}"
uv sync --group dev --group proxy-dev --extra proxy --frozen --quiet
uv run --no-sync python -m prisma generate --schema litellm/proxy/schema.prisma
@ -143,16 +146,18 @@ done
# --- LiteLLM proxy ---
echo "=== Starting LiteLLM proxy ==="
cd "$REPO_ROOT"
PROXY_LOG="${TMPDIR:-/tmp}/litellm-e2e-proxy-$$.log"
uv run --no-sync python -m litellm.proxy.proxy_cli \
--config "$SCRIPT_DIR/fixtures/config.yml" \
--port 4000 &
--port 4000 >"$PROXY_LOG" 2>&1 &
PROXY_PID=$!
echo "Waiting for proxy..."
echo "Waiting for proxy (logs: $PROXY_LOG)..."
PROXY_READY=0
for i in $(seq 1 180); do
if ! kill -0 "$PROXY_PID" 2>/dev/null; then
echo "Error: proxy process exited unexpectedly"
echo "Error: proxy process exited unexpectedly. Proxy output:"
tail -n 100 "$PROXY_LOG"
exit 1
fi
HTTP_CODE=$(curl -s -o /dev/null -w "%{http_code}" http://127.0.0.1:4000/health -H "Authorization: Bearer $LITELLM_MASTER_KEY" 2>/dev/null || true)
@ -163,7 +168,8 @@ for i in $(seq 1 180); do
sleep 1
done
if [ "$PROXY_READY" -ne 1 ]; then
echo "Error: proxy did not become healthy within 180 seconds"
echo "Error: proxy did not become healthy within 180 seconds. Proxy output:"
tail -n 100 "$PROXY_LOG"
exit 1
fi
echo "Proxy is ready."

View file

@ -8,7 +8,16 @@ import { MIGRATED_E2E_PAGES } from "../../fixtures/migratedPages";
import type { Page as PlaywrightPage } from "@playwright/test";
const sidebarButtons = {
[Role.ProxyAdmin]: ["Virtual Keys", "Playground", "Models", "Usage", "Teams", "Internal Users", "AI Hub"],
[Role.ProxyAdmin]: [
"Virtual Keys",
"Playground",
"Models",
"Usage",
"Teams",
"Internal Users",
"AI Hub",
"Response Cache",
],
};
/** Migrated pages live at a path route; legacy pages keep the ?page= query param. */

View file

@ -177,11 +177,6 @@
"count": 1
}
},
"src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.test.ts": {
"unused-imports/no-unused-imports": {
"count": 1
}
},
"src/app/(dashboard)/cost-tracking/_components/provider_margin_table.tsx": {
"no-restricted-imports": {
"count": 1
@ -2207,7 +2202,7 @@
},
"src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx": {
"no-nested-ternary": {
"count": 4
"count": 3
}
},
"src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx": {

View file

@ -14,6 +14,7 @@
"@base-ui/react": "^1.6.0",
"@headlessui/tailwindcss": "0.2.2",
"@heroicons/react": "1.0.6",
"@hookform/resolvers": "5.4.0",
"@tanstack/react-pacer": "0.22.1",
"@tanstack/react-query": "5.100.7",
"@tanstack/react-table": "8.21.3",
@ -34,13 +35,15 @@
"react": "18.3.1",
"react-copy-to-clipboard": "5.1.1",
"react-dom": "18.3.1",
"react-hook-form": "7.82.0",
"react-json-view-lite": "2.5.0",
"react-markdown": "9.1.0",
"react-syntax-highlighter": "15.6.6",
"recharts": "3.9.2",
"remark-gfm": "4.0.1",
"tailwind-merge": "3.4.0",
"uuid": "14.0.0"
"uuid": "14.0.0",
"zod": "3.25.76"
},
"devDependencies": {
"@eslint/js": "9.39.2",
@ -1556,6 +1559,18 @@
"react": ">= 16"
}
},
"node_modules/@hookform/resolvers": {
"version": "5.4.0",
"resolved": "https://registry.npmjs.org/@hookform/resolvers/-/resolvers-5.4.0.tgz",
"integrity": "sha512-EIsqr/t/qbinPIhGjMdtvutIN1Kk4uwbROE9/UQ93CAVGR7GkA7Y92+fX80OzXi/OB67jVFYwKGO1WzkxmkFZw==",
"license": "MIT",
"dependencies": {
"@standard-schema/utils": "^0.3.0"
},
"peerDependencies": {
"react-hook-form": "^7.55.0"
}
},
"node_modules/@humanfs/core": {
"version": "0.19.2",
"resolved": "https://registry.npmjs.org/@humanfs/core/-/core-0.19.2.tgz",
@ -11780,6 +11795,22 @@
"react": "^18.3.1"
}
},
"node_modules/react-hook-form": {
"version": "7.82.0",
"resolved": "https://registry.npmjs.org/react-hook-form/-/react-hook-form-7.82.0.tgz",
"integrity": "sha512-Zw/uFZ2dO+02GHlBn7JFGn8kZJ7LdM33B/0BXOovzFay+CMhf94JMw5BVu+F1tVkUKjNvBuaE3fz5BJhga10Tg==",
"license": "MIT",
"engines": {
"node": ">=18.0.0"
},
"funding": {
"type": "opencollective",
"url": "https://opencollective.com/react-hook-form"
},
"peerDependencies": {
"react": "^16.8.0 || ^17 || ^18 || ^19"
}
},
"node_modules/react-is": {
"version": "17.0.2",
"resolved": "https://registry.npmjs.org/react-is/-/react-is-17.0.2.tgz",
@ -14156,7 +14187,6 @@
"version": "3.25.76",
"resolved": "https://registry.npmjs.org/zod/-/zod-3.25.76.tgz",
"integrity": "sha512-gzUt/qt81nXsFGKIFcC3YnfEAx5NkunCfnDlvuBSSFS02bcXu4Lmea0AFIUwbLWxWPx3d9p8S5QoaujKcNQxcQ==",
"devOptional": true,
"license": "MIT",
"funding": {
"url": "https://github.com/sponsors/colinhacks"

View file

@ -30,6 +30,7 @@
"@base-ui/react": "^1.6.0",
"@headlessui/tailwindcss": "0.2.2",
"@heroicons/react": "1.0.6",
"@hookform/resolvers": "5.4.0",
"@tanstack/react-pacer": "0.22.1",
"@tanstack/react-query": "5.100.7",
"@tanstack/react-table": "8.21.3",
@ -50,13 +51,15 @@
"react": "18.3.1",
"react-copy-to-clipboard": "5.1.1",
"react-dom": "18.3.1",
"react-hook-form": "7.82.0",
"react-json-view-lite": "2.5.0",
"react-markdown": "9.1.0",
"react-syntax-highlighter": "15.6.6",
"recharts": "3.9.2",
"remark-gfm": "4.0.1",
"tailwind-merge": "3.4.0",
"uuid": "14.0.0"
"uuid": "14.0.0",
"zod": "3.25.76"
},
"devDependencies": {
"@eslint/js": "9.39.2",

View file

@ -1 +1 @@
<svg fill="currentColor" fill-rule="evenodd" height="1em" style="flex:none;line-height:1" viewBox="0 0 103 24" xmlns="http://www.w3.org/2000/svg"><title>AI21</title><path d="M15.064 21.643l-.74-2.335H7.487l-.741 2.335H2L8.862 2.414h4.09l6.944 19.23h-4.832zM10.92 7.908L8.56 15.819h4.666l-2.305-7.911zm9.504-5.494h4.501v19.23h-4.501V2.413zm5.714 15.3a8.84 8.84 0 011.057-2.691 7.78 7.78 0 011.606-1.868 16.915 16.915 0 012.045-1.456c.567-.33 1.093-.646 1.578-.948.447-.275.874-.582 1.276-.92.348-.289.64-.638.865-1.03.214-.391.323-.832.315-1.278 0-.769-.21-1.323-.63-1.662-.447-.347-1-.526-1.565-.508A2.475 2.475 0 0030.943 6c-.467.43-.7 1.15-.7 2.156h-4.42a6.493 6.493 0 01.454-2.445c.294-.74.749-1.406 1.33-1.95a6.26 6.26 0 012.142-1.291A8.363 8.363 0 0132.657 2a9.048 9.048 0 012.512.344c.76.21 1.472.565 2.1 1.044a4.972 4.972 0 011.44 1.813c.374.823.557 1.72.536 2.623a4.64 4.64 0 01-.522 2.198 7.454 7.454 0 01-1.276 1.758c-.497.508-1.044.963-1.633 1.36-.586.394-1.117.728-1.592 1.003-.66.44-1.204.82-1.633 1.14a7.753 7.753 0 00-1.03.892 2.403 2.403 0 00-.535.852 3.128 3.128 0 00-.15 1.017h8.234v3.598h-13.34c-.023-1.32.102-2.637.371-3.928zM39.958 5.792c.769.016 1.538-.053 2.292-.206a3.307 3.307 0 001.386-.618 2.14 2.14 0 00.686-1.044c.137-.491.202-1 .192-1.51h3.926v19.23h-4.528V8.923h-3.954V5.792z"></path><path d="M53.534 2.414h4.199v19.23h-4.2V2.413zm19.704 5.978v13.241h-4.2v-1.868a3.434 3.434 0 01-.659.865 4.24 4.24 0 01-.946.686c-.371.198-.762.354-1.167.466-.41.118-.835.178-1.262.179a6.33 6.33 0 01-2.622-.536 6.207 6.207 0 01-2.044-1.456 6.455 6.455 0 01-1.318-2.197 8.402 8.402 0 010-5.522 6.455 6.455 0 011.318-2.198 6.207 6.207 0 012.044-1.455 6.329 6.329 0 012.622-.535c.427.001.852.061 1.263.179.406.113.798.274 1.166.48.347.194.666.434.947.715.257.252.48.539.659.851V8.392h4.199zm-7.356 10c.445.006.886-.088 1.29-.275a3.147 3.147 0 001.634-1.8c.156-.417.235-.859.233-1.304a3.546 3.546 0 00-.879-2.363 3.056 3.056 0 00-.989-.742 3.168 3.168 0 00-2.58 0 3.06 3.06 0 00-.989.742 3.549 3.549 0 00-.878 2.363c-.002.445.077.888.233 1.305a3.153 3.153 0 001.634 1.8c.404.187.845.28 1.29.273zm12.459 3.251h-4.2V2.414h4.2v7.883a4.218 4.218 0 011.619-1.566c.37-.202.76-.364 1.166-.48.415-.12.845-.18 1.276-.179a6.33 6.33 0 012.622.536c.773.34 1.47.836 2.044 1.456a6.47 6.47 0 011.318 2.197 8.413 8.413 0 010 5.522 6.47 6.47 0 01-1.318 2.198 6.203 6.203 0 01-2.044 1.456 6.33 6.33 0 01-2.622.535 4.577 4.577 0 01-1.276-.178 6.213 6.213 0 01-1.166-.466 4.12 4.12 0 01-.96-.687 3.435 3.435 0 01-.66-.865v1.867zm3.184-3.241c.436.004.868-.09 1.263-.275a3.148 3.148 0 001.634-1.8c.156-.416.235-.859.232-1.304a3.547 3.547 0 00-.878-2.363 3.056 3.056 0 00-.99-.741 2.901 2.901 0 00-1.261-.276 3.012 3.012 0 00-2.305 1.016 3.546 3.546 0 00-.879 2.362 3.66 3.66 0 00.233 1.305c.145.395.364.759.645 1.072a3.1 3.1 0 002.306 1.004zm15.219-4.56a20.56 20.56 0 011.646.535c.479.175.932.415 1.345.714.38.28.695.642.92 1.058.244.494.362 1.042.343 1.593a3.937 3.937 0 01-.467 1.992 3.71 3.71 0 01-1.29 1.319c-.587.353-1.234.6-1.907.727-.765.15-1.542.224-2.32.22-1.885 0-3.372-.412-4.46-1.236-1.09-.824-1.413-2.006-1.413-3.544h3.897c0 .696.197 1.195.59 1.497a2.43 2.43 0 001.524.453c.434.019.866-.08 1.248-.288a1.018 1.018 0 00.48-.948.898.898 0 00-.205-.618 1.982 1.982 0 00-.632-.44 7.884 7.884 0 00-1.097-.412c-.45-.137-.985-.306-1.606-.508a21.019 21.019 0 01-1.565-.535 5.688 5.688 0 01-1.304-.7 2.978 2.978 0 01-.891-1.045 3.427 3.427 0 01-.33-1.593c0-1.374.526-2.39 1.578-3.05 1.053-.659 2.475-.989 4.268-.99a8.335 8.335 0 012.512.344c.652.192 1.26.514 1.784.948.457.39.818.878 1.057 1.429.238.555.36 1.153.357 1.758h-4.145a1.798 1.798 0 00-.425-1.278 1.71 1.71 0 00-1.304-.453 2.04 2.04 0 00-1.098.289.972.972 0 00-.467.892.828.828 0 00.22.591c.185.182.405.327.645.426.344.15.697.279 1.057.384.42.13.905.285 1.455.468z" fill="#E91E63"></path></svg>
<svg fill="currentColor" fill-rule="evenodd" style="flex:none;line-height:1" viewBox="0 0 103 24" xmlns="http://www.w3.org/2000/svg"><title>AI21</title><path d="M15.064 21.643l-.74-2.335H7.487l-.741 2.335H2L8.862 2.414h4.09l6.944 19.23h-4.832zM10.92 7.908L8.56 15.819h4.666l-2.305-7.911zm9.504-5.494h4.501v19.23h-4.501V2.413zm5.714 15.3a8.84 8.84 0 011.057-2.691 7.78 7.78 0 011.606-1.868 16.915 16.915 0 012.045-1.456c.567-.33 1.093-.646 1.578-.948.447-.275.874-.582 1.276-.92.348-.289.64-.638.865-1.03.214-.391.323-.832.315-1.278 0-.769-.21-1.323-.63-1.662-.447-.347-1-.526-1.565-.508A2.475 2.475 0 0030.943 6c-.467.43-.7 1.15-.7 2.156h-4.42a6.493 6.493 0 01.454-2.445c.294-.74.749-1.406 1.33-1.95a6.26 6.26 0 012.142-1.291A8.363 8.363 0 0132.657 2a9.048 9.048 0 012.512.344c.76.21 1.472.565 2.1 1.044a4.972 4.972 0 011.44 1.813c.374.823.557 1.72.536 2.623a4.64 4.64 0 01-.522 2.198 7.454 7.454 0 01-1.276 1.758c-.497.508-1.044.963-1.633 1.36-.586.394-1.117.728-1.592 1.003-.66.44-1.204.82-1.633 1.14a7.753 7.753 0 00-1.03.892 2.403 2.403 0 00-.535.852 3.128 3.128 0 00-.15 1.017h8.234v3.598h-13.34c-.023-1.32.102-2.637.371-3.928zM39.958 5.792c.769.016 1.538-.053 2.292-.206a3.307 3.307 0 001.386-.618 2.14 2.14 0 00.686-1.044c.137-.491.202-1 .192-1.51h3.926v19.23h-4.528V8.923h-3.954V5.792z"></path><path d="M53.534 2.414h4.199v19.23h-4.2V2.413zm19.704 5.978v13.241h-4.2v-1.868a3.434 3.434 0 01-.659.865 4.24 4.24 0 01-.946.686c-.371.198-.762.354-1.167.466-.41.118-.835.178-1.262.179a6.33 6.33 0 01-2.622-.536 6.207 6.207 0 01-2.044-1.456 6.455 6.455 0 01-1.318-2.197 8.402 8.402 0 010-5.522 6.455 6.455 0 011.318-2.198 6.207 6.207 0 012.044-1.455 6.329 6.329 0 012.622-.535c.427.001.852.061 1.263.179.406.113.798.274 1.166.48.347.194.666.434.947.715.257.252.48.539.659.851V8.392h4.199zm-7.356 10c.445.006.886-.088 1.29-.275a3.147 3.147 0 001.634-1.8c.156-.417.235-.859.233-1.304a3.546 3.546 0 00-.879-2.363 3.056 3.056 0 00-.989-.742 3.168 3.168 0 00-2.58 0 3.06 3.06 0 00-.989.742 3.549 3.549 0 00-.878 2.363c-.002.445.077.888.233 1.305a3.153 3.153 0 001.634 1.8c.404.187.845.28 1.29.273zm12.459 3.251h-4.2V2.414h4.2v7.883a4.218 4.218 0 011.619-1.566c.37-.202.76-.364 1.166-.48.415-.12.845-.18 1.276-.179a6.33 6.33 0 012.622.536c.773.34 1.47.836 2.044 1.456a6.47 6.47 0 011.318 2.197 8.413 8.413 0 010 5.522 6.47 6.47 0 01-1.318 2.198 6.203 6.203 0 01-2.044 1.456 6.33 6.33 0 01-2.622.535 4.577 4.577 0 01-1.276-.178 6.213 6.213 0 01-1.166-.466 4.12 4.12 0 01-.96-.687 3.435 3.435 0 01-.66-.865v1.867zm3.184-3.241c.436.004.868-.09 1.263-.275a3.148 3.148 0 001.634-1.8c.156-.416.235-.859.232-1.304a3.547 3.547 0 00-.878-2.363 3.056 3.056 0 00-.99-.741 2.901 2.901 0 00-1.261-.276 3.012 3.012 0 00-2.305 1.016 3.546 3.546 0 00-.879 2.362 3.66 3.66 0 00.233 1.305c.145.395.364.759.645 1.072a3.1 3.1 0 002.306 1.004zm15.219-4.56a20.56 20.56 0 011.646.535c.479.175.932.415 1.345.714.38.28.695.642.92 1.058.244.494.362 1.042.343 1.593a3.937 3.937 0 01-.467 1.992 3.71 3.71 0 01-1.29 1.319c-.587.353-1.234.6-1.907.727-.765.15-1.542.224-2.32.22-1.885 0-3.372-.412-4.46-1.236-1.09-.824-1.413-2.006-1.413-3.544h3.897c0 .696.197 1.195.59 1.497a2.43 2.43 0 001.524.453c.434.019.866-.08 1.248-.288a1.018 1.018 0 00.48-.948.898.898 0 00-.205-.618 1.982 1.982 0 00-.632-.44 7.884 7.884 0 00-1.097-.412c-.45-.137-.985-.306-1.606-.508a21.019 21.019 0 01-1.565-.535 5.688 5.688 0 01-1.304-.7 2.978 2.978 0 01-.891-1.045 3.427 3.427 0 01-.33-1.593c0-1.374.526-2.39 1.578-3.05 1.053-.659 2.475-.989 4.268-.99a8.335 8.335 0 012.512.344c.652.192 1.26.514 1.784.948.457.39.818.878 1.057 1.429.238.555.36 1.153.357 1.758h-4.145a1.798 1.798 0 00-.425-1.278 1.71 1.71 0 00-1.304-.453 2.04 2.04 0 00-1.098.289.972.972 0 00-.467.892.828.828 0 00.22.591c.185.182.405.327.645.426.344.15.697.279 1.057.384.42.13.905.285 1.455.468z" fill="#E91E63"></path></svg>

Before

Width:  |  Height:  |  Size: 3.8 KiB

After

Width:  |  Height:  |  Size: 3.7 KiB

View file

@ -1,5 +1,5 @@
<svg version="1.1" id="Layer_1" xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" x="0px" y="0px"
width="100%" viewBox="0 0 1024 1024" enable-background="new 0 0 1024 1024" xml:space="preserve">
viewBox="0 0 1024 1024" enable-background="new 0 0 1024 1024" xml:space="preserve">
<path fill="#000000" opacity="1.000000" stroke="none"
d="
M793.182983,619.000000

Before

Width:  |  Height:  |  Size: 5.9 KiB

After

Width:  |  Height:  |  Size: 5.9 KiB

View file

@ -1 +1 @@
<svg viewBox="0 0 100 17.5" width="92" fill="white" xmlns="http://www.w3.org/2000/svg"><title>Soniox</title><path d="m0 14.866 2.1606-3.5214c1.8927 1.2576 3.9669 1.8995 5.6694 1.8995 1.0025 0 1.4606-0.3036 1.4606-0.8847v-0.0607c0-0.6419-0.9161-0.9194-2.6532-1.4138-3.2582-0.8587-5.8509-1.9602-5.8509-5.2995v-0.06938c0-3.5214 2.8088-5.4903 6.6114-5.4903 2.4112 0 4.9089 0.70255 6.8016 1.9342l-1.9791 3.6775c-1.7112-0.95408-3.5693-1.5352-4.8744-1.5352-0.88152 0-1.3396 0.33827-1.3396 0.79796v0.06071c0 0.64184 0.94202 0.95409 2.6792 1.4745 3.2582 0.91939 5.8509 2.0556 5.8509 5.2735v0.0607c0 3.6515-2.7137 5.551-6.741 5.551-2.7656-0.0087-5.5052-0.798-7.7955-2.4546z"></path><path d="m16.135 8.7342v-0.06071c0-4.7184 3.8372-8.6735 9.1436-8.6735 5.2719 0 9.0832 3.8944 9.0832 8.6127v0.06072c0 4.7184-3.8372 8.6735-9.1437 8.6735-5.2718 0-9.0831-3.8944-9.0831-8.6128zm12.583 0v-0.06071c0-2.0209-1.4606-3.7383-3.5088-3.7383-2.1001 0-3.4483 1.6826-3.4483 3.6775v0.06072c0 2.0209 1.4605 3.7383 3.5088 3.7383 2.1087 0 3.4483-1.6827 3.4483-3.6776z"></path><path d="m36.877 0.36428h5.7904v2.3332c1.063-1.3791 2.5927-2.6974 4.9348-2.6974 3.5089 0 5.609 2.3332 5.609 6.0974v10.85h-5.7905v-8.977c0-1.8041-0.942-2.7929-2.3161-2.7929-1.4001 0-2.4372 0.9801-2.4372 2.7929v8.977h-5.7904z"></path><path d="m55.951 0.36426h5.7904v16.584h-5.7904z"></path><path d="m64.29 8.7342v-0.06071c0-4.7184 3.8373-8.6735 9.1437-8.6735 5.2719 0 9.0832 3.8944 9.0832 8.6127v0.06072c0 4.7184-3.8372 8.6735-9.1437 8.6735-5.2719 0-9.0832-3.8944-9.0832-8.6128zm12.592 0v-0.06071c0-2.0209-1.4605-3.7383-3.5088-3.7383-2.1001 0-3.4483 1.6826-3.4483 3.6775v0.06072c0 2.0209 1.4606 3.7383 3.5088 3.7383 2.1088 0 3.4483-1.6827 3.4483-3.6776z"></path><path d="m88.082 8.578-5.4533-8.2138h6.2484l2.4372 4.0765 2.4371-4.0765h6.1275l-5.4274 8.1791 5.5484 8.3959h-6.2225l-2.5582-4.2587-2.5927 4.2587h-6.0929z"></path></svg>
<svg viewBox="0 0 100 17.5" fill="white" xmlns="http://www.w3.org/2000/svg"><title>Soniox</title><path d="m0 14.866 2.1606-3.5214c1.8927 1.2576 3.9669 1.8995 5.6694 1.8995 1.0025 0 1.4606-0.3036 1.4606-0.8847v-0.0607c0-0.6419-0.9161-0.9194-2.6532-1.4138-3.2582-0.8587-5.8509-1.9602-5.8509-5.2995v-0.06938c0-3.5214 2.8088-5.4903 6.6114-5.4903 2.4112 0 4.9089 0.70255 6.8016 1.9342l-1.9791 3.6775c-1.7112-0.95408-3.5693-1.5352-4.8744-1.5352-0.88152 0-1.3396 0.33827-1.3396 0.79796v0.06071c0 0.64184 0.94202 0.95409 2.6792 1.4745 3.2582 0.91939 5.8509 2.0556 5.8509 5.2735v0.0607c0 3.6515-2.7137 5.551-6.741 5.551-2.7656-0.0087-5.5052-0.798-7.7955-2.4546z"></path><path d="m16.135 8.7342v-0.06071c0-4.7184 3.8372-8.6735 9.1436-8.6735 5.2719 0 9.0832 3.8944 9.0832 8.6127v0.06072c0 4.7184-3.8372 8.6735-9.1437 8.6735-5.2718 0-9.0831-3.8944-9.0831-8.6128zm12.583 0v-0.06071c0-2.0209-1.4606-3.7383-3.5088-3.7383-2.1001 0-3.4483 1.6826-3.4483 3.6775v0.06072c0 2.0209 1.4605 3.7383 3.5088 3.7383 2.1087 0 3.4483-1.6827 3.4483-3.6776z"></path><path d="m36.877 0.36428h5.7904v2.3332c1.063-1.3791 2.5927-2.6974 4.9348-2.6974 3.5089 0 5.609 2.3332 5.609 6.0974v10.85h-5.7905v-8.977c0-1.8041-0.942-2.7929-2.3161-2.7929-1.4001 0-2.4372 0.9801-2.4372 2.7929v8.977h-5.7904z"></path><path d="m55.951 0.36426h5.7904v16.584h-5.7904z"></path><path d="m64.29 8.7342v-0.06071c0-4.7184 3.8373-8.6735 9.1437-8.6735 5.2719 0 9.0832 3.8944 9.0832 8.6127v0.06072c0 4.7184-3.8372 8.6735-9.1437 8.6735-5.2719 0-9.0832-3.8944-9.0832-8.6128zm12.592 0v-0.06071c0-2.0209-1.4605-3.7383-3.5088-3.7383-2.1001 0-3.4483 1.6826-3.4483 3.6775v0.06072c0 2.0209 1.4606 3.7383 3.5088 3.7383 2.1088 0 3.4483-1.6827 3.4483-3.6776z"></path><path d="m88.082 8.578-5.4533-8.2138h6.2484l2.4372 4.0765 2.4371-4.0765h6.1275l-5.4274 8.1791 5.5484 8.3959h-6.2225l-2.5582-4.2587-2.5927 4.2587h-6.0929z"></path></svg>

Before

Width:  |  Height:  |  Size: 1.8 KiB

After

Width:  |  Height:  |  Size: 1.8 KiB

View file

@ -0,0 +1,89 @@
import React from "react";
import { render, screen, fireEvent, within } from "@testing-library/react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import AddAgentForm from "./add_agent_form";
import * as networking from "@/components/networking";
import type { AgentCreateInfo } from "@/components/networking";
vi.mock("@/components/networking", () => ({
createAgentCall: vi.fn(),
getAgentCreateMetadata: vi.fn(),
getAgentsList: vi.fn(),
keyCreateForAgentCall: vi.fn(),
keyListCall: vi.fn(),
keyUpdateCall: vi.fn(),
modelAvailableCall: vi.fn(),
}));
vi.mock("./agent_card_discovery", () => ({
default: () => <div data-testid="agent-card-discovery" />,
}));
vi.mock("./agent_form_fields", () => ({
default: () => <div data-testid="agent-form-fields" />,
}));
const a2aInfo: AgentCreateInfo = {
agent_type: "a2a",
agent_type_display_name: "A2A Agent",
description: "Agent-to-agent protocol",
logo_url: "/ui/assets/logos/a2a_agent.png",
credential_fields: [],
use_a2a_form_fields: true,
};
const renderForm = () =>
render(<AddAgentForm visible={true} onClose={vi.fn()} accessToken="test-token" onSuccess={vi.fn()} />);
describe("AddAgentForm logos", () => {
beforeEach(() => {
vi.mocked(networking.getAgentCreateMetadata).mockReset().mockResolvedValue([a2aInfo]);
vi.mocked(networking.getAgentsList).mockReset().mockResolvedValue({ agents: [] });
vi.mocked(networking.keyListCall).mockReset().mockResolvedValue({ keys: [] });
vi.mocked(networking.modelAvailableCall).mockReset().mockResolvedValue({ data: [] });
});
it("renders the modal title and agent type selection logos as images from logo_url", async () => {
renderForm();
const titleLogo = await screen.findByAltText("Agent logo");
expect(titleLogo).toBeInstanceOf(HTMLImageElement);
expect(titleLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png"));
const selectionLogo = await screen.findByAltText("A2A Agent logo");
expect(selectionLogo).toBeInstanceOf(HTMLImageElement);
expect(selectionLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png"));
});
it("renders the option logo when the agent type dropdown is opened", async () => {
renderForm();
await screen.findByAltText("A2A Agent logo");
fireEvent.mouseDown(screen.getByRole("combobox"));
const optionLogos = await screen.findAllByAltText("A2A Agent logo");
expect(optionLogos.length).toBeGreaterThanOrEqual(2);
optionLogos.forEach((img) => {
expect(img).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png"));
});
});
it("swaps a failing logo for a letter avatar and warns with the url", async () => {
const warnSpy = vi.spyOn(console, "warn").mockImplementation(() => {});
renderForm();
const titleLogo = await screen.findByAltText("Agent logo");
const header = screen.getByText("Add New Agent").parentElement!;
fireEvent.error(titleLogo);
expect(warnSpy).toHaveBeenCalledWith(expect.stringContaining("assets/logos/a2a_agent.png"));
expect(screen.queryByAltText("Agent logo")).not.toBeInTheDocument();
expect(within(header).getByText("A")).toBeInTheDocument();
const selectionLogo = screen.getByAltText("A2A Agent logo");
fireEvent.error(selectionLogo);
expect(screen.queryByAltText("A2A Agent logo")).not.toBeInTheDocument();
expect(warnSpy).toHaveBeenCalledTimes(2);
warnSpy.mockRestore();
});
});

View file

@ -1,7 +1,7 @@
import React, { useState, useEffect } from "react";
import { Modal, Form, Select, Input, Steps, Radio, Tag, Divider, Switch, InputNumber, Collapse } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { resolveLogoSrc } from "@/lib/assetPaths";
import { Logo } from "@/components/molecules/logo/Logo";
import { Button } from "@tremor/react";
import { CheckCircleFilled, KeyOutlined, RobotOutlined, AppstoreOutlined, InfoCircleOutlined } from "@ant-design/icons";
import CreatedKeyDisplay from "@/components/shared/CreatedKeyDisplay";
@ -712,17 +712,13 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
value={info.agent_type}
label={
<div className="flex items-center gap-2">
<img src={resolveLogoSrc(info.logo_url) ?? ""} alt="" className="w-4 h-4 object-contain" />
<Logo src={info.logo_url} label={info.agent_type_display_name} className="w-4 h-4 object-contain" />
<span>{info.agent_type_display_name}</span>
</div>
}
>
<div className="flex items-center gap-3 py-1">
<img
src={resolveLogoSrc(info.logo_url) ?? ""}
alt={info.agent_type_display_name}
className="w-5 h-5 object-contain"
/>
<Logo src={info.logo_url} label={info.agent_type_display_name} className="w-5 h-5 object-contain" />
<div>
<div className="font-medium">{info.agent_type_display_name}</div>
{info.description && <div className="text-xs text-gray-500">{info.description}</div>}
@ -948,7 +944,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
title={
<div className="flex items-center space-x-3 pb-4 border-b border-gray-100">
{selectedLogo && currentStep < 1 && (
<img src={resolveLogoSrc(selectedLogo)} alt="Agent" className="w-6 h-6 object-contain" />
<Logo src={selectedLogo} label="Agent" className="w-6 h-6 object-contain" />
)}
<h2 className="text-xl font-semibold text-gray-900">Add New Agent</h2>
</div>

View file

@ -76,6 +76,22 @@ describe("CacheDashboard cache analytics charts", () => {
expect(screen.getByText("Cached Completion Tokens vs Generated Completion Tokens")).toBeInTheDocument();
});
it("scopes the analytics tab to the response cache, not provider prompt caching", async () => {
renderDashboard();
expect(await screen.findByText(/is not shown here/)).toBeInTheDocument();
expect(screen.getByRole("link", { name: "response cache" })).toHaveAttribute(
"href",
"https://docs.litellm.ai/docs/proxy/caching",
);
expect(screen.getByRole("link", { name: "prompt caching" })).toHaveAttribute(
"href",
"https://docs.litellm.ai/docs/completion/prompt_caching",
);
expect(screen.queryByText("Cached Tokens")).not.toBeInTheDocument();
expect(screen.getAllByText("Cached Completion Tokens").length).toBeGreaterThan(0);
});
it("renders the requests chart with each category legend-bound to its fill and stacked in order", async () => {
renderDashboard();
const { requestsCard } = await findChartCards();

View file

@ -282,6 +282,28 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
<TabPanels>
<TabPanel>
<Card>
<Text className="text-tremor-content dark:text-dark-tremor-content">
Analytics for LiteLLM&apos;s{" "}
<a
href="https://docs.litellm.ai/docs/proxy/caching"
target="_blank"
rel="noreferrer"
className="underline"
>
response cache
</a>{" "}
(e.g. Redis / in-memory): requests answered from cache without calling the LLM provider. Provider-side{" "}
<a
href="https://docs.litellm.ai/docs/completion/prompt_caching"
target="_blank"
rel="noreferrer"
className="underline"
>
prompt caching
</a>{" "}
(cached input tokens from Anthropic, OpenAI, etc.) is not shown here; see &quot;Prompt Caching
Metrics&quot; on the Usage page or individual requests in the Logs page.
</Text>
<Grid numItems={3} className="gap-4 mt-4">
<Col>
<MultiSelect
@ -340,7 +362,7 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
<Card>
<p className="text-tremor-default font-medium text-tremor-content dark:text-dark-tremor-content">
Cached Tokens
Cached Completion Tokens
</p>
<div className="mt-2 flex items-baseline space-x-2.5">
<p className="text-tremor-metric font-semibold text-tremor-content-strong dark:text-dark-tremor-content-strong">

View file

@ -6,25 +6,6 @@ import { renderWithProviders } from "../../../../../tests/test-utils";
import AddMarginForm from "./add_margin_form";
import { MarginConfig } from "./types";
vi.mock("@/components/provider_info_helpers", () => ({
Providers: {
OpenAI: "OpenAI",
Anthropic: "Anthropic",
},
provider_map: {
OpenAI: "openai",
Anthropic: "anthropic",
},
providerLogoMap: {
OpenAI: "https://example.com/openai.png",
Anthropic: "https://example.com/anthropic.png",
},
}));
vi.mock("./provider_display_helpers", () => ({
handleImageError: vi.fn(),
}));
const DEFAULT_PROPS = {
marginConfig: {} as MarginConfig,
selectedProvider: undefined,

View file

@ -2,10 +2,9 @@ import React from "react";
import { TextInput, Button } from "@tremor/react";
import { Select as AntdSelect, Form, Tooltip, Radio } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { Providers, provider_map, providerLogoMap } from "@/components/provider_info_helpers";
import { resolveLogoSrc } from "@/lib/assetPaths";
import { Providers, provider_map } from "@/components/provider_info_helpers";
import { Logo } from "@/components/molecules/logo/Logo";
import { MarginConfig } from "./types";
import { handleImageError } from "./provider_display_helpers";
interface AddMarginFormProps {
marginConfig: MarginConfig;
@ -73,12 +72,7 @@ const AddMarginForm: React.FC<AddMarginFormProps> = ({
return (
<AntdSelect.Option key={providerEnum} value={providerEnum} label={providerDisplayName}>
<div className="flex items-center space-x-2">
<img
src={resolveLogoSrc(providerLogoMap[providerDisplayName])}
alt={`${providerEnum} logo`}
className="w-5 h-5"
onError={(e) => handleImageError(e, providerDisplayName)}
/>
<Logo provider={providerEnum} label={providerDisplayName} className="w-5 h-5" />
<span>{providerDisplayName}</span>
</div>
</AntdSelect.Option>

View file

@ -5,25 +5,7 @@ import userEvent from "@testing-library/user-event";
import { renderWithProviders } from "../../../../../tests/test-utils";
import AddProviderForm from "./add_provider_form";
import { DiscountConfig } from "./types";
vi.mock("@/components/provider_info_helpers", () => ({
Providers: {
OpenAI: "OpenAI",
Anthropic: "Anthropic",
},
provider_map: {
OpenAI: "openai",
Anthropic: "anthropic",
},
providerLogoMap: {
OpenAI: "https://example.com/openai.png",
Anthropic: "https://example.com/anthropic.png",
},
}));
vi.mock("./provider_display_helpers", () => ({
handleImageError: vi.fn(),
}));
import { Providers, providerLogoMap } from "@/components/provider_info_helpers";
const DEFAULT_PROPS = {
discountConfig: {} as DiscountConfig,
@ -84,4 +66,18 @@ describe("AddProviderForm", () => {
renderWithProviders(<AddProviderForm {...DEFAULT_PROPS} />);
expect(screen.getByText("%")).toBeInTheDocument();
});
it("renders the selected provider's bundled logo via the shared Logo component", async () => {
renderWithProviders(<AddProviderForm {...DEFAULT_PROPS} selectedProvider="OpenAI" />);
const logo = await screen.findByRole("img", { name: `${Providers.OpenAI} logo` });
expect(logo.getAttribute("src")).toBe(providerLogoMap[Providers.OpenAI]);
});
it("falls back to a letter avatar for a selected provider that has no bundled logo", () => {
renderWithProviders(<AddProviderForm {...DEFAULT_PROPS} selectedProvider="PG_VECTOR" />);
expect(screen.queryByRole("img", { name: `${Providers.PG_VECTOR} logo` })).not.toBeInTheDocument();
expect(screen.getByText(Providers.PG_VECTOR.charAt(0))).toBeInTheDocument();
});
});

View file

@ -2,10 +2,9 @@ import React from "react";
import { TextInput, Button } from "@tremor/react";
import { Select as AntdSelect, Form, Tooltip } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { Providers, provider_map, providerLogoMap } from "@/components/provider_info_helpers";
import { resolveLogoSrc } from "@/lib/assetPaths";
import { Providers, provider_map } from "@/components/provider_info_helpers";
import { Logo } from "@/components/molecules/logo/Logo";
import { DiscountConfig } from "./types";
import { handleImageError } from "./provider_display_helpers";
interface AddProviderFormProps {
discountConfig: DiscountConfig;
@ -60,12 +59,7 @@ const AddProviderForm: React.FC<AddProviderFormProps> = ({
return (
<AntdSelect.Option key={providerEnum} value={providerEnum} label={providerDisplayName}>
<div className="flex items-center space-x-2">
<img
src={resolveLogoSrc(providerLogoMap[providerDisplayName])}
alt={`${providerEnum} logo`}
className="w-5 h-5"
onError={(e) => handleImageError(e, providerDisplayName)}
/>
<Logo provider={providerEnum} label={providerDisplayName} className="w-5 h-5" />
<span>{providerDisplayName}</span>
</div>
</AntdSelect.Option>

View file

@ -49,11 +49,7 @@ vi.mock("@/components/provider_info_helpers", () => ({
Providers: { OpenAI: "OpenAI" },
provider_map: { OpenAI: "openai" },
providerLogoMap: {},
}));
vi.mock("./provider_display_helpers", () => ({
getProviderDisplayInfo: vi.fn(() => ({ displayName: "OpenAI", logo: "", enumKey: "OpenAI" })),
handleImageError: vi.fn(),
getProviderLogoAndName: (providerValue: string) => ({ logo: "", displayName: providerValue }),
}));
const ADMIN_PROPS = {

View file

@ -11,7 +11,6 @@ export type {
MarginConfig,
CostMarginResponse,
} from "./types";
export type { ProviderDisplayInfo } from "./provider_display_helpers";
export * from "./provider_display_helpers";
export { useDiscountConfig } from "./use_discount_config";
export { useMarginConfig } from "./use_margin_config";

View file

@ -43,15 +43,6 @@ vi.mock("@tremor/react", () => ({
},
}));
vi.mock("./provider_display_helpers", () => ({
getProviderDisplayInfo: vi.fn((providerValue: string) => ({
displayName: providerValue === "openai" ? "OpenAI" : providerValue,
logo: providerValue === "openai" ? "https://example.com/openai.png" : "",
enumKey: providerValue === "openai" ? "OpenAI" : null,
})),
handleImageError: vi.fn(),
}));
const DEFAULT_DISCOUNT_CONFIG = {
openai: 0.05,
anthropic: 0.1,

View file

@ -3,7 +3,8 @@ import { TextInput, Icon, Text } from "@tremor/react";
import { TrashIcon, PencilAltIcon, CheckIcon, XIcon } from "@heroicons/react/outline";
import { SimpleTable } from "@/components/common_components/simple_table";
import { DiscountConfig } from "./types";
import { getProviderDisplayInfo, handleImageError } from "./provider_display_helpers";
import { getProviderLogoAndName } from "@/components/provider_info_helpers";
import { Logo } from "@/components/molecules/logo/Logo";
interface ProviderDiscountTableProps {
discountConfig: DiscountConfig;
@ -55,8 +56,8 @@ const ProviderDiscountTable: React.FC<ProviderDiscountTableProps> = ({
const data: ProviderDiscountRow[] = Object.entries(discountConfig)
.map(([provider, discount]) => ({ provider, discount }))
.sort((a, b) => {
const displayA = getProviderDisplayInfo(a.provider).displayName;
const displayB = getProviderDisplayInfo(b.provider).displayName;
const displayA = getProviderLogoAndName(a.provider).displayName;
const displayB = getProviderLogoAndName(b.provider).displayName;
return displayA.localeCompare(displayB);
});
@ -67,17 +68,10 @@ const ProviderDiscountTable: React.FC<ProviderDiscountTableProps> = ({
{
header: "Provider",
cell: (row) => {
const { displayName, logo } = getProviderDisplayInfo(row.provider);
const { displayName } = getProviderLogoAndName(row.provider);
return (
<div className="flex items-center space-x-2">
{logo && (
<img
src={logo}
alt={`${displayName} logo`}
className="w-5 h-5"
onError={(e) => handleImageError(e, displayName)}
/>
)}
<Logo provider={row.provider} label={displayName} className="w-5 h-5" />
<span className="font-medium">{displayName}</span>
</div>
);
@ -129,7 +123,7 @@ const ProviderDiscountTable: React.FC<ProviderDiscountTableProps> = ({
{
header: "Actions",
cell: (row) => {
const { displayName } = getProviderDisplayInfo(row.provider);
const { displayName } = getProviderLogoAndName(row.provider);
return (
<Icon
icon={TrashIcon}

View file

@ -1,46 +1,14 @@
import { describe, it, expect, vi, beforeEach } from "vitest";
import { getProviderDisplayInfo, getProviderBackendValue, handleImageError } from "./provider_display_helpers";
import { describe, it, expect, vi } from "vitest";
import { getProviderBackendValue } from "./provider_display_helpers";
vi.mock("@/components/provider_info_helpers", () => ({
Providers: {
OpenAI: "OpenAI",
Anthropic: "Anthropic",
Azure: "Azure",
},
provider_map: {
OpenAI: "openai",
Anthropic: "anthropic",
Azure: "azure",
},
providerLogoMap: {
OpenAI: "https://example.com/openai.png",
Anthropic: "https://example.com/anthropic.png",
Azure: "https://example.com/azure.png",
},
}));
describe("getProviderDisplayInfo", () => {
it("should return display name and logo for a known backend provider value", () => {
const info = getProviderDisplayInfo("openai");
expect(info.displayName).toBe("OpenAI");
expect(info.logo).toBe("https://example.com/openai.png");
expect(info.enumKey).toBe("OpenAI");
});
it("should return the raw value as display name for an unknown provider", () => {
const info = getProviderDisplayInfo("my-custom-provider");
expect(info.displayName).toBe("my-custom-provider");
expect(info.logo).toBe("");
expect(info.enumKey).toBeNull();
});
it("should match a provider by its backend value regardless of casing", () => {
const info = getProviderDisplayInfo("anthropic");
expect(info.displayName).toBe("Anthropic");
expect(info.enumKey).toBe("Anthropic");
});
});
describe("getProviderBackendValue", () => {
it("should return the backend value for a known provider enum key", () => {
expect(getProviderBackendValue("OpenAI")).toBe("openai");
@ -54,38 +22,3 @@ describe("getProviderBackendValue", () => {
expect(getProviderBackendValue("UnknownProvider")).toBeNull();
});
});
describe("handleImageError", () => {
it("should replace the img element with a fallback div showing the first letter", () => {
const img = document.createElement("img");
const parent = document.createElement("div");
parent.appendChild(img);
const event = { target: img } as any;
handleImageError(event, "OpenAI");
expect(parent.querySelector("img")).toBeNull();
const fallback = parent.firstChild as HTMLElement;
expect(fallback.tagName).toBe("DIV");
expect(fallback.textContent).toBe("O");
});
it("should use the first character of the fallback text as the label", () => {
const img = document.createElement("img");
const parent = document.createElement("div");
parent.appendChild(img);
const event = { target: img } as any;
handleImageError(event, "Anthropic");
const fallback = parent.firstChild as HTMLElement;
expect(fallback.textContent).toBe("A");
});
it("should do nothing if the image has no parent element", () => {
const img = document.createElement("img");
const event = { target: img } as any;
// Should not throw
expect(() => handleImageError(event, "OpenAI")).not.toThrow();
});
});

View file

@ -1,28 +1,4 @@
import { Providers, provider_map, providerLogoMap } from "@/components/provider_info_helpers";
import { resolveLogoSrc } from "@/lib/assetPaths";
export interface ProviderDisplayInfo {
displayName: string;
logo: string;
enumKey: string | null;
}
/**
* Convert backend provider value (e.g., "openai") to display info
*/
export const getProviderDisplayInfo = (providerValue: string): ProviderDisplayInfo => {
const enumKey = Object.keys(provider_map).find(
(key) => provider_map[key as keyof typeof provider_map] === providerValue,
);
if (enumKey) {
const displayName = Providers[enumKey as keyof typeof Providers];
const logo = resolveLogoSrc(providerLogoMap[displayName]) ?? "";
return { displayName, logo, enumKey };
}
return { displayName: providerValue, logo: "", enumKey: null };
};
import { provider_map } from "@/components/provider_info_helpers";
/**
* Convert provider enum key (e.g., "OpenAI") to backend value (e.g., "openai")
@ -30,17 +6,3 @@ export const getProviderDisplayInfo = (providerValue: string): ProviderDisplayIn
export const getProviderBackendValue = (providerEnum: string): string | null => {
return provider_map[providerEnum as keyof typeof provider_map] || null;
};
/**
* Handle image error by replacing with fallback div
*/
export const handleImageError = (e: React.SyntheticEvent<HTMLImageElement>, fallbackText: string) => {
const target = e.target as HTMLImageElement;
const parent = target.parentElement;
if (parent) {
const fallbackDiv = document.createElement("div");
fallbackDiv.className = "w-5 h-5 rounded-full bg-gray-200 flex items-center justify-center text-xs";
fallbackDiv.textContent = fallbackText.charAt(0);
parent.replaceChild(fallbackDiv, target);
}
};

View file

@ -4,6 +4,7 @@ import { screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { renderWithProviders } from "../../../../../tests/test-utils";
import ProviderMarginTable from "./provider_margin_table";
import { Providers, providerLogoMap } from "@/components/provider_info_helpers";
vi.mock("@heroicons/react/outline", () => ({
TrashIcon: function TrashIcon() {
@ -43,15 +44,6 @@ vi.mock("@tremor/react", () => ({
},
}));
vi.mock("./provider_display_helpers", () => ({
getProviderDisplayInfo: vi.fn((providerValue: string) => {
if (providerValue === "openai") return { displayName: "OpenAI", logo: "", enumKey: "OpenAI" };
if (providerValue === "anthropic") return { displayName: "Anthropic", logo: "", enumKey: "Anthropic" };
return { displayName: providerValue, logo: "", enumKey: null };
}),
handleImageError: vi.fn(),
}));
describe("ProviderMarginTable", () => {
const onMarginChange = vi.fn();
const onRemoveProvider = vi.fn();
@ -95,6 +87,30 @@ describe("ProviderMarginTable", () => {
expect(screen.getByText("OpenAI")).toBeInTheDocument();
});
it("should render the provider's bundled logo via the shared Logo component", () => {
renderWithProviders(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>,
);
const logo = screen.getByRole("img", { name: `${Providers.OpenAI} logo` });
expect(logo.getAttribute("src")).toBe(providerLogoMap[Providers.OpenAI]);
});
it("should fall back to a letter avatar for a provider with no bundled logo", () => {
renderWithProviders(
<ProviderMarginTable
marginConfig={{ "my-custom-provider": 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>,
);
expect(screen.queryByRole("img")).not.toBeInTheDocument();
expect(screen.getByText("m")).toBeInTheDocument();
});
it("should display the global provider as 'Global (All Providers)'", () => {
renderWithProviders(
<ProviderMarginTable

View file

@ -3,7 +3,8 @@ import { TextInput, Icon, Text } from "@tremor/react";
import { TrashIcon, PencilAltIcon, CheckIcon, XIcon } from "@heroicons/react/outline";
import { SimpleTable } from "@/components/common_components/simple_table";
import { MarginConfig } from "./types";
import { getProviderDisplayInfo, handleImageError } from "./provider_display_helpers";
import { getProviderLogoAndName } from "@/components/provider_info_helpers";
import { Logo } from "@/components/molecules/logo/Logo";
interface ProviderMarginTableProps {
marginConfig: MarginConfig;
@ -96,8 +97,8 @@ const ProviderMarginTable: React.FC<ProviderMarginTableProps> = ({
.sort((a, b) => {
if (a.provider === "global") return -1;
if (b.provider === "global") return 1;
const displayA = getProviderDisplayInfo(a.provider).displayName;
const displayB = getProviderDisplayInfo(b.provider).displayName;
const displayA = getProviderLogoAndName(a.provider).displayName;
const displayB = getProviderLogoAndName(b.provider).displayName;
return displayA.localeCompare(displayB);
});
@ -115,17 +116,10 @@ const ProviderMarginTable: React.FC<ProviderMarginTableProps> = ({
</div>
);
}
const { displayName, logo } = getProviderDisplayInfo(row.provider);
const { displayName } = getProviderLogoAndName(row.provider);
return (
<div className="flex items-center space-x-2">
{logo && (
<img
src={logo}
alt={`${displayName} logo`}
className="w-5 h-5"
onError={(e) => handleImageError(e, displayName)}
/>
)}
<Logo provider={row.provider} label={displayName} className="w-5 h-5" />
<span className="font-medium">{displayName}</span>
</div>
);
@ -186,7 +180,7 @@ const ProviderMarginTable: React.FC<ProviderMarginTableProps> = ({
{
header: "Actions",
cell: (row) => {
const displayName = row.provider === "global" ? "Global" : getProviderDisplayInfo(row.provider).displayName;
const displayName = row.provider === "global" ? "Global" : getProviderLogoAndName(row.provider).displayName;
return (
<Icon
icon={TrashIcon}

View file

@ -41,3 +41,17 @@ describe("AddGuardrailForm close behavior", () => {
expect(onClose).toHaveBeenCalledTimes(1);
});
});
describe("AddGuardrailForm provider options", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("renders provider options with logos from the bundled guardrail logo map", async () => {
renderForm();
fireEvent.mouseDown(screen.getByLabelText("Guardrail Provider"));
const logo = await screen.findByAltText("Presidio PII logo");
expect(logo.getAttribute("src")).toContain("microsoft_azure.svg");
});
});

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