mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge litellm_internal_staging into OTEL v2 destinations
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled
This commit is contained in:
commit
9d9832689e
188 changed files with 8957 additions and 1604 deletions
2
.github/workflows/image-scan.yml
vendored
2
.github/workflows/image-scan.yml
vendored
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
);
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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=(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@
|
|||
"limit": 33
|
||||
},
|
||||
"DTZ005": {
|
||||
"limit": 244
|
||||
"limit": 241
|
||||
},
|
||||
"DTZ006": {
|
||||
"limit": 13
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
66
tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py
Normal file
66
tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py
Normal 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}"
|
||||
)
|
||||
211
tests/e2e/guardrails/test_presidio_guardrail_e2e.py
Normal file
211
tests/e2e/guardrails/test_presidio_guardrail_e2e.py
Normal 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}"
|
||||
)
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
415
tests/e2e/management/test_budget_customer_user_org_e2e.py
Normal file
415
tests/e2e/management/test_budget_customer_user_org_e2e.py
Normal 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",
|
||||
)
|
||||
251
tests/e2e/management/test_key_management_e2e.py
Normal file
251
tests/e2e/management/test_key_management_e2e.py
Normal 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"
|
||||
)
|
||||
385
tests/e2e/management/test_model_tag_accessgroup_e2e.py
Normal file
385
tests/e2e/management/test_model_tag_accessgroup_e2e.py
Normal 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"
|
||||
)
|
||||
303
tests/e2e/management/test_team_management_e2e.py
Normal file
303
tests/e2e/management/test_team_management_e2e.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
146
tests/e2e/mcp/test_mcp_guardrail_e2e.py
Normal file
146
tests/e2e/mcp/test_mcp_guardrail_e2e.py
Normal 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}"
|
||||
)
|
||||
|
|
@ -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] = []
|
||||
|
|
|
|||
18
tests/e2e/other/conftest.py
Normal file
18
tests/e2e/other/conftest.py
Normal 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)
|
||||
73
tests/e2e/other/other_client.py
Normal file
73
tests/e2e/other/other_client.py
Normal 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)
|
||||
65
tests/e2e/other/test_health_lifecycle_e2e.py
Normal file
65
tests/e2e/other/test_health_lifecycle_e2e.py
Normal 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}"
|
||||
)
|
||||
37
tests/e2e/other/test_master_key_auth_e2e.py
Normal file
37
tests/e2e/other/test_master_key_auth_e2e.py
Normal 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}"
|
||||
)
|
||||
|
|
@ -234,6 +234,7 @@ CONTROL_PLANE_PREFIXES: tuple[str, ...] = (
|
|||
"/tag",
|
||||
"/budget",
|
||||
"/model/",
|
||||
"/access_group",
|
||||
"/spend",
|
||||
"/global",
|
||||
"/config",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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. */
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
34
ui/litellm-dashboard/package-lock.json
generated
34
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 |
|
|
@ -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 |
|
|
@ -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 |
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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'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 "Prompt Caching
|
||||
Metrics" 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">
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue