diff --git a/.github/workflows/image-scan.yml b/.github/workflows/image-scan.yml index 90ede5a653f..8d791ca5bc7 100644 --- a/.github/workflows/image-scan.yml +++ b/.github/workflows/image-scan.yml @@ -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 \ diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql new file mode 100644 index 00000000000..95412df0a96 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql @@ -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") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index da403156874..5f8e6280445 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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 diff --git a/litellm/constants.py b/litellm/constants.py index b08d16e0606..9dd80750f95 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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 diff --git a/litellm/litellm_core_utils/duration_parser.py b/litellm/litellm_core_utils/duration_parser.py index 438ff5600ba..79036367652 100644 --- a/litellm/litellm_core_utils/duration_parser.py +++ b/litellm/litellm_core_utils/duration_parser.py @@ -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 diff --git a/litellm/main.py b/litellm/main.py index fb05a375111..dc3ec469a1b 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 882c34dbd6a..9a1b5cf4864 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8f30071eb5d..90b70dd01f2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 1ee300be718..0b795057837 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py new file mode 100644 index 00000000000..e0927cc4f64 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py @@ -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 + ) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index a8ab0937124..396dd6c7dc7 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 56beaca3344..26053822644 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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", diff --git a/litellm/proxy/a2a/agent_card.py b/litellm/proxy/a2a/agent_card.py index e97ab4a01ae..29a689a32de 100644 --- a/litellm/proxy/a2a/agent_card.py +++ b/litellm/proxy/a2a/agent_card.py @@ -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 diff --git a/litellm/proxy/a2a/version_convert.py b/litellm/proxy/a2a/version_convert.py index e8f49e6f6a9..9de33a0966a 100644 --- a/litellm/proxy/a2a/version_convert.py +++ b/litellm/proxy/a2a/version_convert.py @@ -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" diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index a7ceffed97b..2421f270974 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -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=( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index c1d152a6605..4846f6e00c0 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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( diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index e758420ee37..23a5b8f9c53 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -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 diff --git a/litellm/proxy/common_utils/timezone_utils.py b/litellm/proxy/common_utils/timezone_utils.py index 32f9f47d519..a50daf40144 100644 --- a/litellm/proxy/common_utils/timezone_utils.py +++ b/litellm/proxy/common_utils/timezone_utils.py @@ -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()) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index b32b26138a1..d5695eb67b9 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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() diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 6c2e06a418c..3c8444ecf26 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3b40abed19e..6de3e43fc1a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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"}, diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index da403156874..5f8e6280445 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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 diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index a8926d26047..48fa4bebaa3 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -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( diff --git a/litellm/types/interactions/generated.py b/litellm/types/interactions/generated.py index 793cc02ff17..4a1ef5ed696 100644 --- a/litellm/types/interactions/generated.py +++ b/litellm/types/interactions/generated.py @@ -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]): diff --git a/litellm/utils.py b/litellm/utils.py index 174bed09396..a11c5500503 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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) diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 609279b0232..a5950677739 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -93,7 +93,7 @@ "limit": 33 }, "DTZ005": { - "limit": 244 + "limit": 241 }, "DTZ006": { "limit": 13 diff --git a/schema.prisma b/schema.prisma index da403156874..5f8e6280445 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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 diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 47f3c74d7f1..bf34a896771 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -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 diff --git a/tests/e2e/coverage_registry/guardrail.yaml b/tests/e2e/coverage_registry/guardrail.yaml index 68722fbbb96..d54c12ba6dc 100644 --- a/tests/e2e/coverage_registry/guardrail.yaml +++ b/tests/e2e/coverage_registry/guardrail.yaml @@ -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"} diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 3b1aff80024..26280d35da0 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -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"} diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index d24cd36c2fd..53f2e4480df 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -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, ), ) diff --git a/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py b/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py new file mode 100644 index 00000000000..e36fc7c3f9d --- /dev/null +++ b/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py @@ -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}" + ) diff --git a/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py b/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py new file mode 100644 index 00000000000..4e2fcbf8fba --- /dev/null +++ b/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py @@ -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}" + ) diff --git a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py new file mode 100644 index 00000000000..a911f387382 --- /dev/null +++ b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py @@ -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 `` placeholders (e.g. +``) 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 = "" + +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}" + ) diff --git a/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py b/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py index cd32a19d54a..db917d6ede9 100644 --- a/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py @@ -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)) diff --git a/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py b/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py index 7a04b044634..97d24e0564b 100644 --- a/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py +++ b/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py @@ -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, ) diff --git a/tests/e2e/management/test_budget_customer_user_org_e2e.py b/tests/e2e/management/test_budget_customer_user_org_e2e.py new file mode 100644 index 00000000000..54cc18b228b --- /dev/null +++ b/tests/e2e/management/test_budget_customer_user_org_e2e.py @@ -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", + ) diff --git a/tests/e2e/management/test_key_management_e2e.py b/tests/e2e/management/test_key_management_e2e.py new file mode 100644 index 00000000000..711175abb0d --- /dev/null +++ b/tests/e2e/management/test_key_management_e2e.py @@ -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" + ) diff --git a/tests/e2e/management/test_model_tag_accessgroup_e2e.py b/tests/e2e/management/test_model_tag_accessgroup_e2e.py new file mode 100644 index 00000000000..e6a187ae105 --- /dev/null +++ b/tests/e2e/management/test_model_tag_accessgroup_e2e.py @@ -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" + ) diff --git a/tests/e2e/management/test_team_management_e2e.py b/tests/e2e/management/test_team_management_e2e.py new file mode 100644 index 00000000000..108aeaad21b --- /dev/null +++ b/tests/e2e/management/test_team_management_e2e.py @@ -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 diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py index f68fdf63b3f..b0aa4c68e3a 100644 --- a/tests/e2e/mcp/mcp_client.py +++ b/tests/e2e/mcp/mcp_client.py @@ -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, diff --git a/tests/e2e/mcp/test_mcp_guardrail_e2e.py b/tests/e2e/mcp/test_mcp_guardrail_e2e.py new file mode 100644 index 00000000000..63239444454 --- /dev/null +++ b/tests/e2e/mcp/test_mcp_guardrail_e2e.py @@ -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}" + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 28cc7984598..37438010c3b 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -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] = [] diff --git a/tests/e2e/logging/otel_client.py b/tests/e2e/otel_client.py similarity index 100% rename from tests/e2e/logging/otel_client.py rename to tests/e2e/otel_client.py diff --git a/tests/e2e/other/conftest.py b/tests/e2e/other/conftest.py new file mode 100644 index 00000000000..9141b6e364e --- /dev/null +++ b/tests/e2e/other/conftest.py @@ -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) diff --git a/tests/e2e/other/other_client.py b/tests/e2e/other/other_client.py new file mode 100644 index 00000000000..1aa83ac42c7 --- /dev/null +++ b/tests/e2e/other/other_client.py @@ -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) diff --git a/tests/e2e/other/test_health_lifecycle_e2e.py b/tests/e2e/other/test_health_lifecycle_e2e.py new file mode 100644 index 00000000000..2551352e8fa --- /dev/null +++ b/tests/e2e/other/test_health_lifecycle_e2e.py @@ -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}" + ) diff --git a/tests/e2e/other/test_master_key_auth_e2e.py b/tests/e2e/other/test_master_key_auth_e2e.py new file mode 100644 index 00000000000..6ab33c9b62a --- /dev/null +++ b/tests/e2e/other/test_master_key_auth_e2e.py @@ -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}" + ) diff --git a/tests/e2e/transport.py b/tests/e2e/transport.py index 005b49272e8..ce33face1d2 100644 --- a/tests/e2e/transport.py +++ b/tests/e2e/transport.py @@ -234,6 +234,7 @@ CONTROL_PLANE_PREFIXES: tuple[str, ...] = ( "/tag", "/budget", "/model/", + "/access_group", "/spend", "/global", "/config", diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py index 5c96eb619bf..44da3ea06a0 100644 --- a/tests/litellm_utils_tests/test_proxy_budget_reset.py +++ b/tests/litellm_utils_tests/test_proxy_budget_reset.py @@ -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 diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index d22b343d843..ee18c96c393 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_duration_parser.py b/tests/test_litellm/litellm_core_utils/test_duration_parser.py index 3e4446c6672..b6b617610a8 100644 --- a/tests/test_litellm/litellm_core_utils/test_duration_parser.py +++ b/tests/test_litellm/litellm_core_utils/test_duration_parser.py @@ -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() diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py index 25d24cfc3ac..1e1b98861b4 100644 --- a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py +++ b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py @@ -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 = [ diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py index 2d09cc0ed32..292bddf1274 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py @@ -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 = [ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py new file mode 100644 index 00000000000..a3f46a49ba9 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 0489b197652..636c7fbd3d5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 73486fe0b6a..b56a12db5b1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 03f91260955..a5cb16822cf 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index 39f3c767220..7bcacb3ff4a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index 3ad01e9c3ec..1e4349c3143 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -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 diff --git a/tests/test_litellm/proxy/a2a/test_agent_card.py b/tests/test_litellm/proxy/a2a/test_agent_card.py index d302bde7895..dfa848e335e 100644 --- a/tests/test_litellm/proxy/a2a/test_agent_card.py +++ b/tests/test_litellm/proxy/a2a/test_agent_card.py @@ -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" diff --git a/tests/test_litellm/proxy/a2a/test_version_convert.py b/tests/test_litellm/proxy/a2a/test_version_convert.py index f3c51ca6b72..7eb5debb792 100644 --- a/tests/test_litellm/proxy/a2a/test_version_convert.py +++ b/tests/test_litellm/proxy/a2a/test_version_convert.py @@ -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" + ) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py index 3740c01b7fc..bcd3333baf9 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py @@ -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.""" diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index b0f2b0fe297..e4c24bf1136 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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 diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 5e348b1bb7e..be5bc74c385 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -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"}) diff --git a/tests/test_litellm/proxy/common_utils/test_timezone_utils.py b/tests/test_litellm/proxy/common_utils/test_timezone_utils.py index 80b813226df..7f686c53c95 100644 --- a/tests/test_litellm/proxy/common_utils/test_timezone_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_timezone_utils.py @@ -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) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index f27f1197090..3ff8e2a6886 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -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( diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index dffca3093fa..51f72f91dc3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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", + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 5631aa69102..e1856860c8a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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 diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 93b21c9d3c1..8d1d8185e4d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -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) diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_config.py b/tests/test_litellm/proxy/proxy_server/test_routes_config.py index 4ac6fc46a61..ad3c470acf3 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_config.py @@ -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 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index a100e7837f4..99a6c946a1e 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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): """ diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 69845ec59c2..671a6cd39ba 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -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 diff --git a/tests/test_litellm/test_router_per_deployment_num_retries.py b/tests/test_litellm/test_router_per_deployment_num_retries.py index af2372616a6..25574fcb268 100644 --- a/tests/test_litellm/test_router_per_deployment_num_retries.py +++ b/tests/test_litellm/test_router_per_deployment_num_retries.py @@ -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 + `` + 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 diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts b/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts index 4a4bb64c8ed..d6e7ea86982 100644 --- a/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts +++ b/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts @@ -26,7 +26,8 @@ export const menuLabelToPage: Record = { "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, diff --git a/ui/litellm-dashboard/e2e_tests/run_e2e.sh b/ui/litellm-dashboard/e2e_tests/run_e2e.sh index ed0641d04e6..ea95f18890c 100755 --- a/ui/litellm-dashboard/e2e_tests/run_e2e.sh +++ b/ui/litellm-dashboard/e2e_tests/run_e2e.sh @@ -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." diff --git a/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts index 7e42d07ae7c..b220dc09ae2 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts @@ -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. */ diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index d596897c4c9..c8a2883e729 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -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": { diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 7a65b63b33c..49d289879d3 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -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" diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index b5e93d175bf..c29b53cd818 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -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", diff --git a/ui/litellm-dashboard/public/assets/logos/ai21.svg b/ui/litellm-dashboard/public/assets/logos/ai21.svg index 7e62a9517af..3c8c75e6d6f 100644 --- a/ui/litellm-dashboard/public/assets/logos/ai21.svg +++ b/ui/litellm-dashboard/public/assets/logos/ai21.svg @@ -1 +1 @@ -AI21 \ No newline at end of file +AI21 \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/logos/promptguard.svg b/ui/litellm-dashboard/public/assets/logos/promptguard.svg index 44cdd52eae3..4b2fd3c386e 100644 --- a/ui/litellm-dashboard/public/assets/logos/promptguard.svg +++ b/ui/litellm-dashboard/public/assets/logos/promptguard.svg @@ -1,5 +1,5 @@ + viewBox="0 0 1024 1024" enable-background="new 0 0 1024 1024" xml:space="preserve"> Soniox +Soniox diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx new file mode 100644 index 00000000000..767e7c2ae5f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx @@ -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: () =>
, +})); + +vi.mock("./agent_form_fields", () => ({ + default: () =>
, +})); + +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(); + +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(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx index 8ca2b5afe16..e35388b78da 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx @@ -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 = ({ visible, onClose, accessTok value={info.agent_type} label={
- + {info.agent_type_display_name}
} >
- {info.agent_type_display_name} +
{info.agent_type_display_name}
{info.description &&
{info.description}
} @@ -948,7 +944,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok title={
{selectedLogo && currentStep < 1 && ( - Agent + )}

Add New Agent

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx index 17d14cd7fac..13472a3d1df 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx @@ -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(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx index b8e8dc8adb1..51c0b85cedb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx @@ -282,6 +282,28 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole + + Analytics for LiteLLM's{" "} + + response cache + {" "} + (e.g. Redis / in-memory): requests answered from cache without calling the LLM provider. Provider-side{" "} + + prompt caching + {" "} + (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. + = ({ accessToken, token, userRole

- Cached Tokens + Cached Completion Tokens

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.test.tsx index 21ee41936c1..1ededd9e4b1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.test.tsx @@ -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, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx index f2c06387301..a17b7fc4ac3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx @@ -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 = ({ return (

- {`${providerEnum} handleImageError(e, providerDisplayName)} - /> + {providerDisplayName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx index 48d23d4645d..08fb63c32b9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx @@ -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(); expect(screen.getByText("%")).toBeInTheDocument(); }); + + it("renders the selected provider's bundled logo via the shared Logo component", async () => { + renderWithProviders(); + + 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(); + + expect(screen.queryByRole("img", { name: `${Providers.PG_VECTOR} logo` })).not.toBeInTheDocument(); + expect(screen.getByText(Providers.PG_VECTOR.charAt(0))).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx index c4961263533..0fdaed8814b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx @@ -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 = ({ return (
- {`${providerEnum} handleImageError(e, providerDisplayName)} - /> + {providerDisplayName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.test.tsx index 0e1c7da92ba..0dae83ba808 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.test.tsx @@ -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 = { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/index.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/index.ts index 8de7fdd7271..90701dd8f1f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/index.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/index.ts @@ -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"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx index c1c43ebdb4f..2e8dbb429f0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx @@ -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, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx index d802f6d83dd..8727d6cb33c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx @@ -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 = ({ 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 = ({ { header: "Provider", cell: (row) => { - const { displayName, logo } = getProviderDisplayInfo(row.provider); + const { displayName } = getProviderLogoAndName(row.provider); return (
- {logo && ( - {`${displayName} handleImageError(e, displayName)} - /> - )} + {displayName}
); @@ -129,7 +123,7 @@ const ProviderDiscountTable: React.FC = ({ { header: "Actions", cell: (row) => { - const { displayName } = getProviderDisplayInfo(row.provider); + const { displayName } = getProviderLogoAndName(row.provider); return ( ({ - 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(); - }); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.ts index 5489eb12487..ed98ba3586b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.ts @@ -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, 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); - } -}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.test.tsx index e1b17dea23d..170e61141b6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.test.tsx @@ -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( + , + ); + 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( + , + ); + expect(screen.queryByRole("img")).not.toBeInTheDocument(); + expect(screen.getByText("m")).toBeInTheDocument(); + }); + it("should display the global provider as 'Global (All Providers)'", () => { renderWithProviders( = ({ .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 = ({
); } - const { displayName, logo } = getProviderDisplayInfo(row.provider); + const { displayName } = getProviderLogoAndName(row.provider); return (
- {logo && ( - {`${displayName} handleImageError(e, displayName)} - /> - )} + {displayName}
); @@ -186,7 +180,7 @@ const ProviderMarginTable: React.FC = ({ { 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 ( { 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"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 202568b478d..17331014c57 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -12,10 +12,10 @@ import { type CompetitorIntentConfig } from "./content_filter/CompetitorIntentCo import { choiceToSkipSystemForCreate, choiceToSkipToolForCreate, + getGuardrailLogo, getGuardrailProviders, getSupportedModesForProvider, guardrail_provider_map, - guardrailLogoMap, populateGuardrailProviderMap, populateGuardrailProviders, shouldRenderContentFilterConfigSettings, @@ -23,7 +23,7 @@ import { shouldRenderPIIConfigSettings, toModeArray, } from "./guardrail_info_helpers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import GuardrailOptionalParams from "./guardrail_optional_params"; import GuardrailProviderFields from "./guardrail_provider_fields"; import LLMJudgeFields from "./llm_judge/LLMJudgeFields"; @@ -725,53 +725,19 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a dropdownRender={(menu) => menu} showSearch={true} > - {Object.entries(getGuardrailProviders()).map(([key, value]) => ( -
- } - > + {Object.entries(getGuardrailProviders()).map(([key, value]) => { + const optionContent = (
- {guardrailLogoMap[value] && ( - { - // Hide broken image icon if image fails to load - e.currentTarget.style.display = "none"; - }} - /> - )} + {value}
- - ))} + ); + return ( + + ); + })} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx index 9ceb6ba244b..ec3d05a6907 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx @@ -16,6 +16,7 @@ import { import { cn } from "@/lib/cva.config"; import { getGuardrailLogoAndName } from "./guardrail_info_helpers"; +import { Logo } from "@/components/molecules/logo/Logo"; const CONFIG_DELETE_HINT = "Config guardrails are defined in the config file and cannot be deleted from the dashboard."; @@ -23,16 +24,7 @@ function GuardrailProviderCell({ provider }: { provider: string }) { const { logo, displayName } = getGuardrailLogoAndName(provider); return (
- {logo ? ( - { - (event.currentTarget as HTMLImageElement).style.display = "none"; - }} - /> - ) : null} + {displayName}
); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.test.tsx index 0fa5d2ffcd2..2d1f35e456c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.test.tsx @@ -53,15 +53,28 @@ describe("GuardrailCard", () => { expect(screen.queryByText(/F1:/)).not.toBeInTheDocument(); }); + it("should render the logo through the shared Logo component with the card src", () => { + render(); + const img = screen.getByAltText("Test Guardrail logo"); + expect(img.getAttribute("src")).toContain("/logos/test.svg"); + }); + + it("should pass a bundled static-import src through unchanged", () => { + const bundledCard: GuardrailCardInfo = { ...baseCard, logo: "/_next/static/media/akto.svg" }; + render(); + expect(screen.getByAltText("Test Guardrail logo")).toHaveAttribute("src", "/_next/static/media/akto.svg"); + }); + it("should show fallback initial when logo fails to load", () => { render(); - const img = screen.getByRole("presentation"); + const img = screen.getByAltText("Test Guardrail logo"); act(() => { fireEvent.error(img); }); expect(screen.getByText("T")).toBeInTheDocument(); + expect(screen.queryByAltText("Test Guardrail logo")).not.toBeInTheDocument(); }); it("should show fallback initial when logo src is empty", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx index 8e9fcc21dfe..53abf3eb81c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx @@ -1,42 +1,7 @@ import React, { useState } from "react"; import { CheckCircleFilled } from "@ant-design/icons"; import { GuardrailCardInfo } from "./guardrail_garden_data"; -import { resolveLogoSrc } from "@/lib/assetPaths"; - -const LogoWithFallback: React.FC<{ src: string; name: string }> = ({ src, name }) => { - const [hasError, setHasError] = useState(false); - - if (hasError || !src) { - return ( -
- {name?.charAt(0) || "?"} -
- ); - } - - return ( - setHasError(true)} - /> - ); -}; +import { Logo } from "@/components/molecules/logo/Logo"; const GuardrailCard: React.FC<{ card: GuardrailCardInfo; onClick: () => void }> = ({ card, onClick }) => { const [hovered, setHovered] = useState(false); @@ -61,7 +26,7 @@ const GuardrailCard: React.FC<{ card: GuardrailCardInfo; onClick: () => void }> > {/* Icon + Name row */}
- + {card.name}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts new file mode 100644 index 00000000000..13909e48185 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts @@ -0,0 +1,54 @@ +import { describe, expect, it } from "vitest"; +import { ALL_CARDS, LITELLM_CONTENT_FILTER_CARDS, PARTNER_GUARDRAIL_CARDS } from "./guardrail_garden_data"; + +const EXPECTED_PARTNER_LOGO_FILES: Record = { + presidio: "microsoft_azure.svg", + bedrock: "bedrock.svg", + lakera: "lakeraai.jpeg", + openai_moderation: "openai_small.svg", + google_model_armor: "google.svg", + guardrails_ai: "guardrails_ai.jpeg", + zscaler: "zscaler.svg", + panw: "palo_alto_networks.jpeg", + cisco_ai_defense: "cisco.png", + noma: "noma_security.png", + aporia: "aporia.png", + aim: "aim_security.jpeg", + cato_networks: "cato_networks.svg", + prompt_security: "prompt_security.png", + lasso: "lasso.png", + pangea: "pangea.png", + enkryptai: "enkrypt_ai.avif", + javelin: "javelin.png", + pillar: "pillar.jpeg", + akto: "akto.svg", + promptguard: "promptguard.svg", + xecguard: "xecguard.svg", + deepkeep: "deepkeep.svg", + repelloai: "repelloai.png", + straiker: "straiker.svg", +}; + +describe("guardrail_garden_data logos", () => { + it("points every partner card at its own provider's bundled logo file", () => { + expect(new Set(PARTNER_GUARDRAIL_CARDS.map((card) => card.id))).toEqual( + new Set(Object.keys(EXPECTED_PARTNER_LOGO_FILES)), + ); + for (const card of PARTNER_GUARDRAIL_CARDS) { + expect(card.logo, `card ${card.id}`).toContain(EXPECTED_PARTNER_LOGO_FILES[card.id]); + } + }); + + it("uses the LiteLLM logo for every content filter card", () => { + for (const card of LITELLM_CONTENT_FILTER_CARDS) { + expect(card.logo, `card ${card.id}`).toContain("litellm_logo.jpg"); + } + }); + + it("bundles every card logo instead of referencing runtime /ui asset paths", () => { + for (const card of ALL_CARDS) { + expect(card.logo, `card ${card.id}`).not.toBe(""); + expect(card.logo, `card ${card.id}`).not.toContain("/ui/assets/logos/"); + } + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts index a29d12f53f9..744af89a357 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts @@ -1,3 +1,5 @@ +import { guardrailLogoMap } from "./guardrail_info_helpers"; + export interface GuardrailCardInfo { id: string; name: string; @@ -16,7 +18,7 @@ export interface GuardrailCardInfo { providerKey?: string; } -const ASSET_PREFIX = "/ui/assets/logos/"; +const litellmContentFilterLogo = guardrailLogoMap["LiteLLM Content Filter"]; export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ { @@ -26,7 +28,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Detects requests for personalized financial advice, investment recommendations, or financial planning.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Topic Blocker"], eval: { f1: 100.0, @@ -42,7 +44,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects insults, name-calling, and personal attacks directed at the chatbot, staff, or other people.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Topic Blocker"], eval: { f1: 100.0, @@ -58,7 +60,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects requests for unauthorized legal advice, case analysis, or legal recommendations.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Topic Blocker"], }, { @@ -67,7 +69,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects requests for medical diagnosis, treatment recommendations, or health advice.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Topic Blocker"], }, { @@ -76,7 +78,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects content related to violence, criminal planning, attacks, and violent threats.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Safety"], }, { @@ -85,7 +87,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects content related to self-harm, suicide, and dangerous self-destructive behavior.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Safety"], }, { @@ -94,7 +96,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects content that could endanger child safety or exploit minors.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Safety"], }, { @@ -103,7 +105,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects content related to illegal weapons manufacturing, distribution, or acquisition.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Safety"], }, { @@ -112,7 +114,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects gender-based discrimination, stereotypes, and biased language.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Bias"], }, { @@ -121,7 +123,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects racial discrimination, stereotypes, and racially biased content.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Bias"], }, { @@ -130,7 +132,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects religious discrimination, intolerance, and religiously biased content.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Bias"], }, { @@ -139,7 +141,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects discrimination based on sexual orientation and related biased content.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Bias"], }, { @@ -148,7 +150,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects jailbreak attempts designed to bypass AI safety guidelines and restrictions.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -157,7 +159,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects attempts to extract sensitive data through prompt manipulation.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -166,7 +168,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects SQL injection attempts embedded in prompts.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -175,7 +177,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects attempts to inject malicious code through prompts.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -184,7 +186,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects attempts to extract or override system prompts.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -193,7 +195,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects toxic, abusive, and hateful language across multiple languages (EN, AU, DE, ES, FR).", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Toxicity"], }, { @@ -203,7 +205,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Detect and block sensitive data patterns like SSNs, credit card numbers, API keys, and custom regex patterns.", category: "litellm", subcategory: "Patterns", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["PII", "Regex", "Data Protection"], }, { @@ -213,7 +215,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Block or mask content containing specific keywords or phrases. Upload custom word lists or add individual terms.", category: "litellm", subcategory: "Keywords", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Keywords", "Blocklist"], }, { @@ -223,7 +225,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Detects markdown fenced code blocks in requests and responses. Block or mask executable code (e.g. Python, JavaScript, Bash) by language with configurable confidence.", category: "litellm", subcategory: "Code Safety", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Code", "Safety", "Prompt Injection"], }, { @@ -233,7 +235,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Block or reframe competitor comparison and ranking intent. Detect when users ask to compare or recommend competitors (airline or generic competitor lists).", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Competitor", "Topic Blocker"], }, ]; @@ -245,7 +247,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "Microsoft Presidio for PII detection and anonymization. Supports 30+ entity types with configurable actions.", category: "partner", - logo: `${ASSET_PREFIX}microsoft_azure.svg`, + logo: guardrailLogoMap["Presidio PII"], tags: ["PII", "Microsoft"], providerKey: "PresidioPII", }, @@ -254,7 +256,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Bedrock Guardrail", description: "AWS Bedrock Guardrails for content filtering, topic avoidance, and sensitive information detection.", category: "partner", - logo: `${ASSET_PREFIX}bedrock.svg`, + logo: guardrailLogoMap["Bedrock Guardrail"], tags: ["AWS", "Content Safety"], providerKey: "Bedrock", }, @@ -263,7 +265,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Lakera", description: "AI security platform protecting against prompt injections, data leakage, and harmful content.", category: "partner", - logo: `${ASSET_PREFIX}lakeraai.jpeg`, + logo: guardrailLogoMap["Lakera"], tags: ["Security", "Prompt Injection"], providerKey: "Lakera", }, @@ -272,7 +274,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "OpenAI Moderation", description: "OpenAI's content moderation API for detecting harmful content across multiple categories.", category: "partner", - logo: `${ASSET_PREFIX}openai_small.svg`, + logo: guardrailLogoMap["OpenAI Moderation"], tags: ["Content Moderation", "OpenAI"], }, { @@ -280,7 +282,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Google Cloud Model Armor", description: "Google Cloud's model protection service for safe and responsible AI deployments.", category: "partner", - logo: `${ASSET_PREFIX}google.svg`, + logo: guardrailLogoMap["Google Cloud Model Armor"], tags: ["Google Cloud", "Safety"], }, { @@ -288,7 +290,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Guardrails AI", description: "Open-source framework for adding structural, type, and quality guarantees to LLM outputs.", category: "partner", - logo: `${ASSET_PREFIX}guardrails_ai.jpeg`, + logo: guardrailLogoMap["Guardrails AI"], tags: ["Open Source", "Validation"], }, { @@ -296,7 +298,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Zscaler AI Guard", description: "Enterprise AI security from Zscaler for monitoring and protecting AI/ML workloads.", category: "partner", - logo: `${ASSET_PREFIX}zscaler.svg`, + logo: guardrailLogoMap["Zscaler AI Guard"], tags: ["Enterprise", "Security"], }, { @@ -304,7 +306,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "PANW Prisma AIRS", description: "Palo Alto Networks Prisma AI Runtime Security for securing AI applications in production.", category: "partner", - logo: `${ASSET_PREFIX}palo_alto_networks.jpeg`, + logo: guardrailLogoMap["PANW Prisma AIRS"], tags: ["Enterprise", "Security"], }, { @@ -313,7 +315,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "Cisco AI Defense Inspection API for runtime protection: prompt injection, PII/PCI/PHI, harassment, hate speech, profanity, violence, and code detection.", category: "partner", - logo: `${ASSET_PREFIX}cisco.png`, + logo: guardrailLogoMap["Cisco AI Defense"], tags: ["Enterprise", "Security", "Prompt Injection", "PII"], providerKey: "CiscoAiDefense", }, @@ -322,7 +324,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Noma Security", description: "AI security platform for detecting and preventing AI-specific threats and vulnerabilities.", category: "partner", - logo: `${ASSET_PREFIX}noma_security.png`, + logo: guardrailLogoMap["Noma Security"], tags: ["Security", "Threat Detection"], }, { @@ -330,7 +332,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Aporia AI", description: "Real-time AI guardrails for hallucination detection, topic control, and policy enforcement.", category: "partner", - logo: `${ASSET_PREFIX}aporia.png`, + logo: guardrailLogoMap["Aporia AI"], tags: ["Hallucination", "Policy"], }, { @@ -338,7 +340,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "AIM Guardrail", description: "AIM Security guardrails for comprehensive AI threat detection and mitigation.", category: "partner", - logo: `${ASSET_PREFIX}aim_security.jpeg`, + logo: guardrailLogoMap["AIM Guardrail"], tags: ["Security", "Threat Detection"], }, { @@ -346,7 +348,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Cato Networks Guardrail", description: "Cato Networks guardrails for comprehensive AI threat detection and mitigation.", category: "partner", - logo: `${ASSET_PREFIX}cato_networks.svg`, + logo: guardrailLogoMap["Cato Networks Guardrail"], tags: ["Security", "Threat Detection"], }, { @@ -354,7 +356,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Prompt Security", description: "Protect against prompt injection attacks, data leakage, and other LLM security threats.", category: "partner", - logo: `${ASSET_PREFIX}prompt_security.png`, + logo: guardrailLogoMap["Prompt Security"], tags: ["Prompt Injection", "Security"], }, { @@ -362,7 +364,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Lasso Guardrail", description: "Content moderation and safety guardrails for responsible AI deployments.", category: "partner", - logo: `${ASSET_PREFIX}lasso.png`, + logo: guardrailLogoMap["Lasso Guardrail"], tags: ["Content Moderation"], }, { @@ -370,7 +372,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Pangea Guardrail", description: "Pangea's AI guardrails for secure, compliant, and trustworthy AI applications.", category: "partner", - logo: `${ASSET_PREFIX}pangea.png`, + logo: guardrailLogoMap["Pangea Guardrail"], tags: ["Compliance", "Security"], }, { @@ -378,7 +380,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "EnkryptAI", description: "AI security and governance platform for enterprise AI safety and compliance.", category: "partner", - logo: `${ASSET_PREFIX}enkrypt_ai.avif`, + logo: guardrailLogoMap["EnkryptAI"], tags: ["Enterprise", "Governance"], }, { @@ -386,7 +388,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Javelin Guardrails", description: "AI gateway with built-in guardrails for secure and compliant AI operations.", category: "partner", - logo: `${ASSET_PREFIX}javelin.png`, + logo: guardrailLogoMap["Javelin Guardrails"], tags: ["Gateway", "Security"], }, { @@ -394,7 +396,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Pillar Guardrail", description: "AI safety platform for monitoring, testing, and securing AI systems.", category: "partner", - logo: `${ASSET_PREFIX}pillar.jpeg`, + logo: guardrailLogoMap["Pillar Guardrail"], tags: ["Monitoring", "Safety"], }, { @@ -402,7 +404,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Akto Guardrail", description: "AI security platform from Akto.io with automatic monitoring and guardrails for AI/ML applications.", category: "partner", - logo: `${ASSET_PREFIX}akto.svg`, + logo: guardrailLogoMap["Akto"], tags: ["Security", "Safety", "Monitoring"], }, { @@ -411,7 +413,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "AI security gateway with prompt injection detection, PII redaction, topic filtering, entity blocklists, and hallucination detection. Self-hostable with drop-in proxy integration.", category: "partner", - logo: `${ASSET_PREFIX}promptguard.svg`, + logo: guardrailLogoMap["PromptGuard"], tags: ["Security", "Prompt Injection", "PII"], providerKey: "Promptguard", eval: { @@ -428,7 +430,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "CyCraft XecGuard AI security gateway. Multi-policy scanning (prompt injection, harmful content, PII, system-prompt enforcement) plus RAG context grounding.", category: "partner", - logo: `${ASSET_PREFIX}xecguard.svg`, + logo: guardrailLogoMap["XecGuard"], tags: ["Security", "Policy", "Grounding", "RAG"], providerKey: "Xecguard", }, @@ -438,7 +440,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "DeepKeep AI Firewall for comprehensive LLM security — prompt injection detection, PII protection, content moderation, and policy enforcement with configurable guardrail pipelines.", category: "partner", - logo: `${ASSET_PREFIX}deepkeep.svg`, + logo: guardrailLogoMap["DeepKeep AI Firewall"], tags: ["Security", "Prompt Injection", "PII", "Firewall"], providerKey: "Deepkeep", }, @@ -448,7 +450,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "RepelloAI Argus scans prompts and responses against policies configured per asset in the Repello dashboard.", category: "partner", - logo: `${ASSET_PREFIX}repelloai.png`, + logo: guardrailLogoMap["RepelloAI Argus"], tags: ["Security", "Policy", "Prompt Injection"], providerKey: "Repelloai", }, @@ -458,7 +460,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "Defend AI Agentic Guardrails: Indirect/Direct Prompt Injection, Tool Misuse, Malicious MCP and Skills", category: "partner", - logo: `${ASSET_PREFIX}straiker.svg`, + logo: guardrailLogoMap["Straiker"], tags: ["Agentic", "Prompt Injection", "Tool Misuse", "MCP", "Skills"], providerKey: "Straiker", }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.test.tsx new file mode 100644 index 00000000000..e17e739267d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.test.tsx @@ -0,0 +1,32 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it, vi } from "vitest"; +import GuardrailDetailView from "./guardrail_garden_detail"; +import type { GuardrailCardInfo } from "./guardrail_garden_data"; + +vi.mock("./add_guardrail_form", () => ({ default: () => null })); + +const makeCard = (overrides: Partial = {}): GuardrailCardInfo => ({ + id: "bedrock", + name: "Bedrock Guardrail", + description: "AWS Bedrock Guardrails for content filtering.", + category: "partner", + logo: "/_next/static/media/bedrock.svg", + tags: ["AWS"], + ...overrides, +}); + +const renderDetail = (card: GuardrailCardInfo) => + render(); + +describe("GuardrailDetailView logo", () => { + it("renders the card logo through the shared Logo component with the bundled src", () => { + renderDetail(makeCard()); + expect(screen.getByAltText("Bedrock Guardrail logo")).toHaveAttribute("src", "/_next/static/media/bedrock.svg"); + }); + + it("falls back to a letter avatar when the card has no logo", () => { + renderDetail(makeCard({ logo: "" })); + expect(screen.queryByAltText("Bedrock Guardrail logo")).not.toBeInTheDocument(); + expect(screen.getByText("B")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx index c92486bbad9..71c7a527614 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx @@ -2,7 +2,7 @@ import React, { useState } from "react"; import { Button } from "antd"; import { ArrowLeftOutlined } from "@ant-design/icons"; import AddGuardrailForm from "./add_guardrail_form"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import { GUARDRAIL_PRESETS } from "./guardrail_garden_configs"; import { GuardrailCardInfo } from "./guardrail_garden_data"; @@ -60,14 +60,7 @@ const GuardrailDetailView: React.FC = ({ card, onBack, {/* ── Header block (Vertex-style) ── */}
- { - (e.target as HTMLImageElement).style.display = "none"; - }} - /> +

{card.name}

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx index c89fe7277c9..7bb7737e152 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx @@ -81,6 +81,37 @@ describe("Guardrail Info", () => { expect(getByText("Settings")).toBeInTheDocument(); }); + it("should render the provider logo from the bundled guardrail logo map", async () => { + vi.mocked(networking.getGuardrailInfo).mockResolvedValue({ + guardrail_id: "123", + guardrail_name: "Test Guardrail", + litellm_params: { + guardrail: "presidio", + mode: "pre_call", + default_on: true, + }, + created_at: "2024-01-01T00:00:00Z", + updated_at: "2024-01-01T00:00:00Z", + guardrail_definition_location: "database", + }); + + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + supported_entities: [], + supported_actions: [], + pii_entity_categories: [], + supported_modes: ["pre_call", "post_call"], + }); + + vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({}); + + const { findByAltText } = render( + {}} accessToken="123" isAdmin={true} />, + ); + + const logo = await findByAltText("Presidio PII logo"); + expect(logo.getAttribute("src")).toContain("microsoft_azure.svg"); + }); + it("should not render the edit button for config guardrails", async () => { // Mock the network responses vi.mocked(networking.getGuardrailInfo).mockResolvedValue({ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx index 1941ec94a60..07df6ff15d9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx @@ -12,6 +12,7 @@ import { Button, Divider, Form, Input, Select, Tooltip } from "antd"; import { CheckIcon, CopyIcon } from "lucide-react"; import React, { useCallback, useEffect, useState } from "react"; import NotificationsManager from "@/components/molecules/notifications_manager"; +import { Logo } from "@/components/molecules/logo/Logo"; import ContentFilterManager, { formatContentFilterDataForAPI } from "./content_filter/ContentFilterManager"; import CustomCodeModal, { EditGuardrailData } from "./custom_code/CustomCodeModal"; import { @@ -524,17 +525,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, Provider
- {logo && ( - {`${displayName} { - // Hide broken image - (e.target as HTMLImageElement).style.display = "none"; - }} - /> - )} + {displayName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index dd70bb2cf51..12aaba0d696 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -1,4 +1,30 @@ -import { resolveLogoSrc } from "@/lib/assetPaths"; +import aimSecurityLogo from "../../../../../public/assets/logos/aim_security.jpeg"; +import aktoLogo from "../../../../../public/assets/logos/akto.svg"; +import aporiaLogo from "../../../../../public/assets/logos/aporia.png"; +import bedrockLogo from "../../../../../public/assets/logos/bedrock.svg"; +import catoNetworksLogo from "../../../../../public/assets/logos/cato_networks.svg"; +import ciscoLogo from "../../../../../public/assets/logos/cisco.png"; +import deepkeepLogo from "../../../../../public/assets/logos/deepkeep.svg"; +import enkryptAiLogo from "../../../../../public/assets/logos/enkrypt_ai.avif"; +import googleLogo from "../../../../../public/assets/logos/google.svg"; +import guardrailsAiLogo from "../../../../../public/assets/logos/guardrails_ai.jpeg"; +import javelinLogo from "../../../../../public/assets/logos/javelin.png"; +import lakeraAiLogo from "../../../../../public/assets/logos/lakeraai.jpeg"; +import lassoLogo from "../../../../../public/assets/logos/lasso.png"; +import litellmLogo from "../../../../../public/assets/logos/litellm_logo.jpg"; +import microsoftAzureLogo from "../../../../../public/assets/logos/microsoft_azure.svg"; +import nomaSecurityLogo from "../../../../../public/assets/logos/noma_security.png"; +import openaiSmallLogo from "../../../../../public/assets/logos/openai_small.svg"; +import paloAltoNetworksLogo from "../../../../../public/assets/logos/palo_alto_networks.jpeg"; +import pangeaLogo from "../../../../../public/assets/logos/pangea.png"; +import pillarLogo from "../../../../../public/assets/logos/pillar.jpeg"; +import promptSecurityLogo from "../../../../../public/assets/logos/prompt_security.png"; +import promptguardLogo from "../../../../../public/assets/logos/promptguard.svg"; +import qohashLogo from "../../../../../public/assets/logos/qohash.jpg"; +import repelloAiLogo from "../../../../../public/assets/logos/repelloai.png"; +import straikerLogo from "../../../../../public/assets/logos/straiker.svg"; +import xecguardLogo from "../../../../../public/assets/logos/xecguard.svg"; +import zscalerLogo from "../../../../../public/assets/logos/zscaler.svg"; // Legacy enum - keeping for backward compatibility export enum GuardrailProviders { @@ -136,40 +162,43 @@ export const shouldRenderLLMJudgeFields = (provider: string | null) => { return guardrail_provider_map[provider] === "llm_as_a_judge"; }; -const asset_logos_folder = "/ui/assets/logos/"; +export const guardrailLogoMap = { + "Zscaler AI Guard": zscalerLogo.src, + "Presidio PII": microsoftAzureLogo.src, + "Bedrock Guardrail": bedrockLogo.src, + Lakera: lakeraAiLogo.src, + "Azure Content Safety Prompt Shield": microsoftAzureLogo.src, + "Azure Content Safety Text Moderation": microsoftAzureLogo.src, + "Aporia AI": aporiaLogo.src, + "PANW Prisma AIRS": paloAltoNetworksLogo.src, + "Cisco AI Defense": ciscoLogo.src, + "Noma Security": nomaSecurityLogo.src, + "Javelin Guardrails": javelinLogo.src, + "Pillar Guardrail": pillarLogo.src, + "Google Cloud Model Armor": googleLogo.src, + "Guardrails AI": guardrailsAiLogo.src, + "Lasso Guardrail": lassoLogo.src, + "Pangea Guardrail": pangeaLogo.src, + "AIM Guardrail": aimSecurityLogo.src, + "Cato Networks Guardrail": catoNetworksLogo.src, + "OpenAI Moderation": openaiSmallLogo.src, + EnkryptAI: enkryptAiLogo.src, + "Prompt Security": promptSecurityLogo.src, + PromptGuard: promptguardLogo.src, + XecGuard: xecguardLogo.src, + "LiteLLM Content Filter": litellmLogo.src, + "LiteLLM LLM as a Judge": litellmLogo.src, + Akto: aktoLogo.src, + "DeepKeep AI Firewall": deepkeepLogo.src, + "Qostodian Nexus": qohashLogo.src, + "RepelloAI Argus": repelloAiLogo.src, + Straiker: straikerLogo.src, +} satisfies Record; -export const guardrailLogoMap: Record = { - "Zscaler AI Guard": `${asset_logos_folder}zscaler.svg`, - "Presidio PII": `${asset_logos_folder}microsoft_azure.svg`, - "Bedrock Guardrail": `${asset_logos_folder}bedrock.svg`, - Lakera: `${asset_logos_folder}lakeraai.jpeg`, - "Azure Content Safety Prompt Shield": `${asset_logos_folder}microsoft_azure.svg`, - "Azure Content Safety Text Moderation": `${asset_logos_folder}microsoft_azure.svg`, - "Aporia AI": `${asset_logos_folder}aporia.png`, - "PANW Prisma AIRS": `${asset_logos_folder}palo_alto_networks.jpeg`, - "Cisco AI Defense": `${asset_logos_folder}cisco.png`, - "Noma Security": `${asset_logos_folder}noma_security.png`, - "Javelin Guardrails": `${asset_logos_folder}javelin.png`, - "Pillar Guardrail": `${asset_logos_folder}pillar.jpeg`, - "Google Cloud Model Armor": `${asset_logos_folder}google.svg`, - "Guardrails AI": `${asset_logos_folder}guardrails_ai.jpeg`, - "Lasso Guardrail": `${asset_logos_folder}lasso.png`, - "Pangea Guardrail": `${asset_logos_folder}pangea.png`, - "AIM Guardrail": `${asset_logos_folder}aim_security.jpeg`, - "Cato Networks Guardrail": `${asset_logos_folder}cato_networks.svg`, - "OpenAI Moderation": `${asset_logos_folder}openai_small.svg`, - EnkryptAI: `${asset_logos_folder}enkrypt_ai.avif`, - "Prompt Security": `${asset_logos_folder}prompt_security.png`, - PromptGuard: `${asset_logos_folder}promptguard.svg`, - XecGuard: `${asset_logos_folder}xecguard.svg`, - "LiteLLM Content Filter": `${asset_logos_folder}litellm_logo.jpg`, - "LiteLLM LLM as a Judge": `${asset_logos_folder}litellm_logo.jpg`, - Akto: `${asset_logos_folder}akto.svg`, - "DeepKeep AI Firewall": `${asset_logos_folder}deepkeep.svg`, - "Qostodian Nexus": `${asset_logos_folder}qohash.jpg`, - "RepelloAI Argus": `${asset_logos_folder}repelloai.png`, - Straiker: `${asset_logos_folder}straiker.svg`, -}; +export const getGuardrailLogo = (displayName: string): string | undefined => + Object.prototype.hasOwnProperty.call(guardrailLogoMap, displayName) + ? guardrailLogoMap[displayName as keyof typeof guardrailLogoMap] + : undefined; export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string; displayName: string } => { if (!guardrailValue) { @@ -188,7 +217,7 @@ export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string; // Get the display name from current GuardrailProviders and logo from map const currentProviders = getGuardrailProviders(); const displayName = currentProviders[enumKey as keyof typeof currentProviders]; - const logo = resolveLogoSrc(guardrailLogoMap[displayName as keyof typeof guardrailLogoMap]) ?? ""; + const logo = getGuardrailLogo(displayName ?? "") ?? ""; return { logo, displayName: displayName || guardrailValue }; }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx index 4f556e74c16..7612b702391 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx @@ -30,6 +30,22 @@ describe("GuardrailTable", () => { } }); + it("renders the provider logo from the bundled guardrail logo map", () => { + render(); + const logo = screen.getByAltText("Presidio PII logo"); + expect(logo.getAttribute("src")).toContain("microsoft_azure.svg"); + }); + + it("falls back to a letter avatar for an unknown provider slug", () => { + const guardrail = makeGuardrail({ + litellm_params: { guardrail: "mystery_guard", mode: "pre_call", default_on: false }, + }); + render(); + expect(screen.getByText("mystery_guard")).toBeInTheDocument(); + expect(screen.queryByAltText("mystery_guard logo")).not.toBeInTheDocument(); + expect(screen.getByText("m")).toBeInTheDocument(); + }); + it("deletes a DB guardrail through the actions menu", async () => { const user = userEvent.setup(); const onDeleteClick = vi.fn(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.test.ts new file mode 100644 index 00000000000..5eb3bdc105d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.test.ts @@ -0,0 +1,103 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { renderHook, waitFor } from "@testing-library/react"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import React, { ReactNode } from "react"; +import { useSetKeyBlockedState, setKeyBlockedState } from "./useSetKeyBlockedState"; +import { apiClient } from "@/components/networking"; + +vi.mock("@/components/networking", () => ({ + apiClient: { post: vi.fn() }, +})); + +const mockUseAuthorized = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => mockUseAuthorized(), +})); + +const mockPost = vi.mocked(apiClient.post); + +const createWrapper = () => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false }, mutations: { retry: false } } }); + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + return { queryClient, wrapper }; +}; + +describe("setKeyBlockedState", () => { + beforeEach(() => { + mockPost.mockReset(); + }); + + it("POSTs the key hash to /key/block when blocking", async () => { + mockPost.mockResolvedValueOnce({ blocked: true }); + + const result = await setKeyBlockedState("sk-access", { keyToken: "hashed-token", blocked: true }); + + expect(mockPost).toHaveBeenCalledWith("/key/block", { + accessToken: "sk-access", + body: { key: "hashed-token" }, + }); + expect(result).toEqual({ blocked: true }); + }); + + it("POSTs the key hash to /key/unblock when unblocking", async () => { + mockPost.mockResolvedValueOnce({ blocked: false }); + + const result = await setKeyBlockedState("sk-access", { keyToken: "hashed-token", blocked: false }); + + expect(mockPost).toHaveBeenCalledWith("/key/unblock", { + accessToken: "sk-access", + body: { key: "hashed-token" }, + }); + expect(result).toEqual({ blocked: false }); + }); + + it("falls back to the requested state when the response has no blocked field", async () => { + mockPost.mockResolvedValueOnce(null); + + const result = await setKeyBlockedState("sk-access", { keyToken: "hashed-token", blocked: true }); + + expect(result).toEqual({ blocked: true }); + }); +}); + +describe("useSetKeyBlockedState", () => { + beforeEach(() => { + mockPost.mockReset(); + mockUseAuthorized.mockReturnValue({ accessToken: "sk-access" }); + }); + + it("invalidates key queries after a successful mutation", async () => { + mockPost.mockResolvedValueOnce({ blocked: true }); + const { queryClient, wrapper } = createWrapper(); + const invalidateSpy = vi.spyOn(queryClient, "invalidateQueries"); + + const { result } = renderHook(() => useSetKeyBlockedState(), { wrapper }); + result.current.mutate({ keyToken: "hashed-token", blocked: true }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(invalidateSpy).toHaveBeenCalledWith({ queryKey: ["keys"] }); + }); + + it("surfaces request failures as mutation errors", async () => { + mockPost.mockRejectedValueOnce(new Error("Key not found.")); + const { wrapper } = createWrapper(); + + const { result } = renderHook(() => useSetKeyBlockedState(), { wrapper }); + result.current.mutate({ keyToken: "missing", blocked: true }); + + await waitFor(() => expect(result.current.isError).toBe(true)); + expect(result.current.error?.message).toBe("Key not found."); + }); + + it("errors without an access token", async () => { + mockUseAuthorized.mockReturnValue({ accessToken: null }); + const { wrapper } = createWrapper(); + + const { result } = renderHook(() => useSetKeyBlockedState(), { wrapper }); + result.current.mutate({ keyToken: "hashed-token", blocked: true }); + + await waitFor(() => expect(result.current.isError).toBe(true)); + expect(mockPost).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.ts new file mode 100644 index 00000000000..792ef567f99 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.ts @@ -0,0 +1,45 @@ +import { useMutation, useQueryClient } from "@tanstack/react-query"; +import { apiClient } from "@/components/networking"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { keyKeys } from "./useKeys"; + +export interface SetKeyBlockedStateInput { + keyToken: string; + blocked: boolean; +} + +export interface SetKeyBlockedStateResult { + blocked: boolean; +} + +interface BlockKeyResponse { + blocked?: boolean | null; +} + +export const setKeyBlockedState = async ( + accessToken: string, + { keyToken, blocked }: SetKeyBlockedStateInput, +): Promise => { + const response = await apiClient.post(blocked ? "/key/block" : "/key/unblock", { + accessToken, + body: { key: keyToken }, + }); + return { blocked: response?.blocked ?? blocked }; +}; + +export const useSetKeyBlockedState = () => { + const { accessToken } = useAuthorized(); + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: async (input) => { + if (!accessToken) { + throw new Error("Access token is required"); + } + return setKeyBlockedState(accessToken, input); + }, + onSuccess: () => { + queryClient.invalidateQueries({ queryKey: keyKeys.all }); + }, + }); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx index 94b9058b372..67b5d6bfe92 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx @@ -52,4 +52,23 @@ describe("MCPLogoSelector", () => { await user.click(githubButton); expect(onChange).toHaveBeenCalledWith(undefined); }); + + it("should render grid logos from bundled static assets instead of public paths", () => { + render(); + const src = screen.getByAltText("GitHub").getAttribute("src"); + expect(src).toMatch(/^\/_next\//); + expect(src).toContain("github.svg"); + }); + + it("should preview a stored well-known path via its bundled asset", () => { + render(); + const src = screen.getByAltText("Selected logo").getAttribute("src"); + expect(src).toMatch(/^\/_next\//); + expect(src).toContain("github.svg"); + }); + + it("should preview a custom external URL untouched", () => { + render(); + expect(screen.getByAltText("Selected logo").getAttribute("src")).toBe("https://cdn.example.com/logo.png"); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx index 6f626a1a70b..a67a0dc882d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx @@ -1,31 +1,51 @@ -import React, { useState } from "react"; +import React from "react"; import { Input, Tooltip } from "antd"; import { InfoCircleOutlined, LinkOutlined } from "@ant-design/icons"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; +import githubLogo from "../../../../../public/assets/logos/github.svg"; +import slackLogo from "../../../../../public/assets/logos/slack.svg"; +import notionLogo from "../../../../../public/assets/logos/notion.svg"; +import linearLogo from "../../../../../public/assets/logos/linear.svg"; +import jiraLogo from "../../../../../public/assets/logos/jira.svg"; +import figmaLogo from "../../../../../public/assets/logos/figma.svg"; +import gmailLogo from "../../../../../public/assets/logos/gmail.svg"; +import googleDriveLogo from "../../../../../public/assets/logos/google_drive.svg"; +import stripeLogo from "../../../../../public/assets/logos/stripe.svg"; +import shopifyLogo from "../../../../../public/assets/logos/shopify.svg"; +import salesforceLogo from "../../../../../public/assets/logos/salesforce.svg"; +import hubspotLogo from "../../../../../public/assets/logos/hubspot.svg"; +import twilioLogo from "../../../../../public/assets/logos/twilio.svg"; +import cloudflareLogo from "../../../../../public/assets/logos/cloudflare.svg"; +import sentryLogo from "../../../../../public/assets/logos/sentry.svg"; +import postgresqlLogo from "../../../../../public/assets/logos/postgresql.svg"; +import snowflakeLogo from "../../../../../public/assets/logos/snowflake.svg"; +import zapierLogo from "../../../../../public/assets/logos/zapier.svg"; +import googleLogo from "../../../../../public/assets/logos/google.svg"; +import gitlabLogo from "../../../../../public/assets/logos/gitlab.svg"; const logos = "/ui/assets/logos/"; -const WELL_KNOWN_LOGOS: { name: string; url: string }[] = [ - { name: "GitHub", url: `${logos}github.svg` }, - { name: "Slack", url: `${logos}slack.svg` }, - { name: "Notion", url: `${logos}notion.svg` }, - { name: "Linear", url: `${logos}linear.svg` }, - { name: "Jira", url: `${logos}jira.svg` }, - { name: "Figma", url: `${logos}figma.svg` }, - { name: "Gmail", url: `${logos}gmail.svg` }, - { name: "Google Drive", url: `${logos}google_drive.svg` }, - { name: "Stripe", url: `${logos}stripe.svg` }, - { name: "Shopify", url: `${logos}shopify.svg` }, - { name: "Salesforce", url: `${logos}salesforce.svg` }, - { name: "HubSpot", url: `${logos}hubspot.svg` }, - { name: "Twilio", url: `${logos}twilio.svg` }, - { name: "Cloudflare", url: `${logos}cloudflare.svg` }, - { name: "Sentry", url: `${logos}sentry.svg` }, - { name: "PostgreSQL", url: `${logos}postgresql.svg` }, - { name: "Snowflake", url: `${logos}snowflake.svg` }, - { name: "Zapier", url: `${logos}zapier.svg` }, - { name: "Google", url: `${logos}google.svg` }, - { name: "GitLab", url: `${logos}gitlab.svg` }, +const WELL_KNOWN_LOGOS: { name: string; url: string; src: string }[] = [ + { name: "GitHub", url: `${logos}github.svg`, src: githubLogo.src }, + { name: "Slack", url: `${logos}slack.svg`, src: slackLogo.src }, + { name: "Notion", url: `${logos}notion.svg`, src: notionLogo.src }, + { name: "Linear", url: `${logos}linear.svg`, src: linearLogo.src }, + { name: "Jira", url: `${logos}jira.svg`, src: jiraLogo.src }, + { name: "Figma", url: `${logos}figma.svg`, src: figmaLogo.src }, + { name: "Gmail", url: `${logos}gmail.svg`, src: gmailLogo.src }, + { name: "Google Drive", url: `${logos}google_drive.svg`, src: googleDriveLogo.src }, + { name: "Stripe", url: `${logos}stripe.svg`, src: stripeLogo.src }, + { name: "Shopify", url: `${logos}shopify.svg`, src: shopifyLogo.src }, + { name: "Salesforce", url: `${logos}salesforce.svg`, src: salesforceLogo.src }, + { name: "HubSpot", url: `${logos}hubspot.svg`, src: hubspotLogo.src }, + { name: "Twilio", url: `${logos}twilio.svg`, src: twilioLogo.src }, + { name: "Cloudflare", url: `${logos}cloudflare.svg`, src: cloudflareLogo.src }, + { name: "Sentry", url: `${logos}sentry.svg`, src: sentryLogo.src }, + { name: "PostgreSQL", url: `${logos}postgresql.svg`, src: postgresqlLogo.src }, + { name: "Snowflake", url: `${logos}snowflake.svg`, src: snowflakeLogo.src }, + { name: "Zapier", url: `${logos}zapier.svg`, src: zapierLogo.src }, + { name: "Google", url: `${logos}google.svg`, src: googleLogo.src }, + { name: "GitLab", url: `${logos}gitlab.svg`, src: gitlabLogo.src }, ]; interface MCPLogoSelectorProps { @@ -34,16 +54,12 @@ interface MCPLogoSelectorProps { } const MCPLogoSelector: React.FC = ({ value, onChange }) => { - const [imgErrors, setImgErrors] = useState>(new Set()); + const selectedWellKnown = WELL_KNOWN_LOGOS.find((l) => l.url === value); const handleSelect = (url: string) => { onChange?.(value === url ? undefined : url); }; - const handleImgError = (url: string) => { - setImgErrors((prev) => new Set(prev).add(url)); - }; - return (
@@ -56,13 +72,10 @@ const MCPLogoSelector: React.FC = ({ value, onChange }) => {/* Preview */} {value && (
- Selected logo { - (e.target as HTMLImageElement).style.display = "none"; - }} />
{value}
@@ -81,8 +94,6 @@ const MCPLogoSelector: React.FC = ({ value, onChange }) =>
{WELL_KNOWN_LOGOS.map((logo) => { const isSelected = value === logo.url; - const hasFailed = imgErrors.has(logo.url); - if (hasFailed) return null; return ( ); @@ -112,7 +118,7 @@ const MCPLogoSelector: React.FC = ({ value, onChange }) => } placeholder="Or paste a custom logo URL..." - value={value && !WELL_KNOWN_LOGOS.some((l) => l.url === value) ? value : ""} + value={value && !selectedWellKnown ? value : ""} onChange={(e) => { const v = e.target.value.trim(); onChange?.(v || undefined); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx index a0998b587fb..d6343afe219 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx @@ -1,8 +1,9 @@ import React from "react"; import { render, screen } from "@testing-library/react"; -import { describe, it, expect, vi } from "vitest"; +import { describe, it, expect, vi, afterEach } from "vitest"; import MCPServerCard from "./MCPServerCard"; import type { MCPServer } from "@/components/mcp_tools/types"; +import { setServerRootPath } from "@/lib/serverRootPath"; const baseServer: MCPServer = { server_id: "srv-1", @@ -43,3 +44,26 @@ describe("MCPServerCard OAuth flow indicator", () => { expect(screen.queryByText("OAuth flow not set")).not.toBeInTheDocument(); }); }); + +describe("MCPServerCard logo", () => { + afterEach(() => { + setServerRootPath("/"); + }); + + it("passes an external logo_url through untouched", () => { + renderCard({ mcp_info: { server_name: "demo_server", logo_url: "https://cdn.example.com/logo.png" } }); + expect(screen.getByAltText("demo_server logo").getAttribute("src")).toBe("https://cdn.example.com/logo.png"); + }); + + it("prefixes a stored asset path with the server root path under a non-root mount", () => { + setServerRootPath("/litellm"); + renderCard({ mcp_info: { server_name: "demo_server", logo_url: "/ui/assets/logos/github.svg" } }); + expect(screen.getByAltText("demo_server logo").getAttribute("src")).toBe("/litellm/ui/assets/logos/github.svg"); + }); + + it("renders a letter avatar when no logo_url is set", () => { + renderCard({ mcp_info: { server_name: "demo_server" } }); + expect(screen.queryByAltText("demo_server logo")).not.toBeInTheDocument(); + expect(screen.getByText("DE")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx index 4282cdba278..c7dd6e47f76 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx @@ -1,4 +1,4 @@ -import { useState, type FC, type KeyboardEvent, type MouseEvent } from "react"; +import { type FC, type KeyboardEvent, type MouseEvent } from "react"; import { Dropdown, Tooltip, Typography, Tag } from "antd"; import type { MenuProps } from "antd"; import { @@ -9,6 +9,7 @@ import { ThunderboltOutlined, } from "@ant-design/icons"; import { AUTH_TYPE, type MCPServer } from "@/components/mcp_tools/types"; +import { Logo } from "@/components/molecules/logo/Logo"; import { getMaskedAndFullUrl } from "./utils"; const { Text } = Typography; @@ -52,8 +53,6 @@ const MCPServerCard: FC = ({ const name = server.server_name || alias || server.server_id; // Logo is sourced exclusively from the admin-set `mcp_info.logo_url`. const candidateLogo = server.mcp_info?.logo_url ?? undefined; - const [failedLogoUrl, setFailedLogoUrl] = useState(null); - const logoUrl = candidateLogo && failedLogoUrl !== candidateLogo ? candidateLogo : undefined; const transport = server.transport || "http"; const displayTransport = server.spec_path && transport !== "stdio" ? "openapi" : transport; const authType = server.auth_type || "none"; @@ -148,13 +147,8 @@ const MCPServerCard: FC = ({ className={`group relative flex h-full cursor-pointer flex-col gap-3 rounded-lg p-4 transition-all duration-150 focus:outline-hidden focus-visible:ring-2 focus-visible:ring-blue-400 ${cardClass}`} >
- {logoUrl ? ( - {`${name} setFailedLogoUrl(logoUrl)} - /> + {candidateLogo ? ( + ) : (
{(name || "?").slice(0, 2).toUpperCase()} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx index 9f7639d00c7..b21a5218c20 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx @@ -39,10 +39,9 @@ import NotificationsManager from "@/components/molecules/notifications_manager"; import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; import { useTestMCPConnection } from "@/hooks/useTestMCPConnection"; import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import mcpLogo from "../../../../../public/assets/logos/mcp_logo.png"; -const asset_logos_folder = "/ui/assets/logos/"; -export const mcpLogoImg = `${asset_logos_folder}mcp_logo.png`; +export const mcpLogoImg = mcpLogo.src; interface CreateMCPServerProps { userRole: string; @@ -791,7 +790,7 @@ const CreateMCPServer: React.FC = ({ )} MCP Logo { + it("renders the tavily logo from the static bundle, untouched by server-root prefixing", () => { + render(); + const img = screen.getByRole("img", { name: "Tavily logo" }); + expect(img).toHaveAttribute("src", "/_next/static/media/tavily.png"); + }); + + it("renders the exa_ai logo file for the exa_ai slug", () => { + render(); + const img = screen.getByRole("img", { name: "Exa AI logo" }); + expect(img.getAttribute("src")).toContain("exa_ai.png"); + }); + + it("renders the google_pse logo file for the google_pse slug", () => { + render(); + expect(screen.getByRole("img", { name: "Google PSE logo" }).getAttribute("src")).toContain("google_pse.png"); + }); + + it("falls back to a letter avatar for a provider with no bundled logo", () => { + render(); + expect(screen.queryByRole("img")).toBeNull(); + expect(screen.getByText("B")).toBeInTheDocument(); + expect(screen.getByText("Brave Search")).toBeInTheDocument(); + }); + + it("does not guess a legacy /ui/assets/logos/.png url for unknown providers", () => { + const { container } = render(); + expect(container.querySelector("img")).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx index b1cb5eb5581..1eeff00cb1b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx @@ -4,44 +4,37 @@ import { useQuery } from "@tanstack/react-query"; import { Button, TextInput } from "@tremor/react"; import { Form, Input, Modal, Select, Tooltip, Typography } from "antd"; import React, { useState } from "react"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { createSearchTool, fetchAvailableSearchProviders } from "@/components/networking"; import SearchConnectionTest from "./SearchConnectionTest"; import { AvailableSearchProvider, SearchTool } from "./types"; +import dataforseoLogo from "../../../../../public/assets/logos/dataforseo.png"; +import exaAiLogo from "../../../../../public/assets/logos/exa_ai.png"; +import googlePseLogo from "../../../../../public/assets/logos/google_pse.png"; +import parallelAiLogo from "../../../../../public/assets/logos/parallel_ai.png"; +import perplexityLogo from "../../../../../public/assets/logos/perplexity.png"; +import tavilyLogo from "../../../../../public/assets/logos/tavily.png"; const { TextArea } = Input; -// Search provider logos folder path (matches existing provider logo pattern) -const searchProviderLogosFolder = "/ui/assets/logos/"; - -// Helper function to get logo path for a search provider -const getSearchProviderLogo = (providerName: string): string => { - return `${searchProviderLogosFolder}${providerName}.png`; +const searchProviderLogoMap: Record = { + perplexity: perplexityLogo.src, + tavily: tavilyLogo.src, + parallel_ai: parallelAiLogo.src, + exa_ai: exaAiLogo.src, + google_pse: googlePseLogo.src, + dataforseo: dataforseoLogo.src, }; -// Component to display search provider logo and name interface SearchProviderLabelProps { providerName: string; displayName: string; } -const SearchProviderLabel: React.FC = ({ providerName, displayName }) => ( -
- {/* eslint-disable-next-line @next/next/no-img-element */} - { - e.currentTarget.style.display = "none"; - }} - /> +export const SearchProviderLabel: React.FC = ({ providerName, displayName }) => ( +
+ {displayName}
); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index a5e0488cbb9..cbf3a2cc1f6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -765,4 +765,37 @@ describe("EntityUsage", () => { }); expect(screen.queryByText(userUuid)).not.toBeInTheDocument(); }); + + it("renders the provider spend table logo from the bundled provider map", async () => { + render(); + + const logo = await screen.findByAltText("openai logo"); + expect(logo.getAttribute("src")).toContain("openai_small"); + }); + + it("renders a letter avatar instead of an img for an unknown provider slug", async () => { + const spendDataUnknownProvider = { + ...mockSpendData, + results: [ + { + ...mockSpendData.results[0], + breakdown: { + ...mockSpendData.results[0].breakdown, + providers: { + "zzz-internal": mockSpendData.results[0].breakdown.providers.openai, + }, + }, + }, + ], + }; + mockTagDailyActivityCall.mockResolvedValue(spendDataUnknownProvider); + + render(); + + await waitFor(() => { + expect(screen.getAllByText("zzz-internal").length).toBeGreaterThan(0); + }); + expect(screen.queryByAltText("zzz-internal logo")).not.toBeInTheDocument(); + expect(screen.getByText("z")).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 94a9cedbdf5..534e2be7fe8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -38,7 +38,7 @@ import { teamDailyActivityCall, userDailyActivityCall, } from "@/components/networking"; -import { getProviderLogoAndName } from "@/components/provider_info_helpers"; +import { Logo } from "@/components/molecules/logo/Logo"; import { usePaginatedDailyActivity } from "../../hooks/usePaginatedDailyActivity"; import { BreakdownMetrics, @@ -774,24 +774,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti
- {provider.provider && ( - {`${provider.provider} { - const target = e.target as HTMLImageElement; - const parent = target.parentElement; - if (parent) { - const fallbackDiv = document.createElement("div"); - fallbackDiv.className = - "w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs"; - fallbackDiv.textContent = provider.provider?.charAt(0) || "-"; - parent.replaceChild(fallbackDiv, target); - } - }} - /> - )} + {provider.provider && } {provider.provider}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx index b5f4b56986d..79ba9b77c14 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx @@ -14,7 +14,7 @@ import { getProviderSpecificFields, VectorStoreFieldConfig, } from "@/components/vector_store_providers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import NotificationsManager from "@/components/molecules/notifications_manager"; import S3VectorsConfig from "./S3VectorsConfig"; @@ -294,22 +294,10 @@ const CreateVectorStore: React.FC = ({ accessToken, onSu return (
- {`${providerEnum} { - // Create a div with provider initial as fallback - 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 = providerDisplayName.charAt(0); - parent.replaceChild(fallbackDiv, target); - } - }} /> {providerDisplayName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.test.tsx index e371cbfba42..3eb013d47ad 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.test.tsx @@ -1,27 +1,34 @@ import { render, screen } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; import { CredentialItem } from "@/components/networking"; +import { Providers, providerLogoMap } from "@/components/provider_info_helpers"; +import { VectorStoreProviders } from "@/components/vector_store_providers"; import VectorStoreForm from "./VectorStoreForm"; vi.mock("@/components/networking"); +const renderForm = () => + render( + , + ); + describe("VectorStoreForm", () => { it("should render the form when visible", () => { - const mockOnCancel = vi.fn(); - const mockOnSuccess = vi.fn(); - const mockAccessToken = "test-token"; - const mockCredentials: CredentialItem[] = []; - - render( - , - ); + renderForm(); expect(screen.getByText("Add New Vector Store")).toBeInTheDocument(); }); + + it("renders the default provider's bundled logo via the shared Logo component", () => { + renderForm(); + + const logo = screen.getByRole("img", { name: `${VectorStoreProviders.Bedrock} logo` }); + expect(logo.getAttribute("src")).toBe(providerLogoMap[Providers.Bedrock]); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx index 82417286738..6cdb895b98d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx @@ -10,7 +10,7 @@ import { getProviderSpecificFields, VectorStoreFieldConfig, } from "@/components/vector_store_providers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models"; import NotificationsManager from "@/components/molecules/notifications_manager"; @@ -130,22 +130,10 @@ const VectorStoreForm: React.FC = ({ return (
- {`${providerEnum} { - // Create a div with provider initial as fallback - 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 = providerDisplayName.charAt(0); - parent.replaceChild(fallbackDiv, target); - } - }} /> {providerDisplayName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx index 7c27b347eb1..ec20a3fd318 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx @@ -10,8 +10,9 @@ import { CredentialItem, } from "@/components/networking"; import { VectorStore } from "@/components/vector_store_management/types"; -import { Providers, providerLogoMap, provider_map } from "@/components/provider_info_helpers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Providers, provider_map } from "@/components/provider_info_helpers"; +import { getVectorStoreProviderLogoAndName } from "@/components/vector_store_providers"; +import { Logo } from "@/components/molecules/logo/Logo"; import VectorStoreTester from "./VectorStoreTester"; import NotificationsManager from "@/components/molecules/notifications_manager"; @@ -181,23 +182,7 @@ const VectorStoreInfoView: React.FC = ({ return (
- {`${providerEnum} { - // Create a div with provider initial as fallback - 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 = providerDisplayName.charAt(0); - parent.replaceChild(fallbackDiv, target); - } - }} - /> + {providerDisplayName}
@@ -292,43 +277,11 @@ const VectorStoreInfoView: React.FC = ({
{(() => { const provider = vectorStoreDetails.custom_llm_provider || "bedrock"; - const { displayName, logo } = (() => { - // Find the enum key by matching provider_map values - const enumKey = Object.keys(provider_map).find( - (key) => provider_map[key].toLowerCase() === provider.toLowerCase(), - ); - - if (!enumKey) { - return { displayName: provider, logo: "" }; - } - - // Get the display name from Providers enum and logo from map - const displayName = Providers[enumKey as keyof typeof Providers]; - const logo = resolveLogoSrc(providerLogoMap[displayName]) ?? ""; - - return { displayName, logo }; - })(); + const { displayName, logo } = getVectorStoreProviderLogoAndName(provider); return ( <> - {logo && ( - {`${displayName} { - 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 = displayName.charAt(0); - parent.replaceChild(fallbackDiv, target); - } - }} - /> - )} + {displayName} ); diff --git a/ui/litellm-dashboard/src/components/SSOModals.test.tsx b/ui/litellm-dashboard/src/components/SSOModals.test.tsx index 365d23f4036..792a964f01c 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.test.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.test.tsx @@ -472,4 +472,43 @@ describe("SSOModals", () => { expect(NotificationsManager.success).toHaveBeenCalledWith("SSO settings cleared successfully"); expect(mockHandleAddSSOOk).toHaveBeenCalled(); }); + + it("renders provider logos in the SSO provider dropdown", async () => { + const TestWrapper = () => { + const [form] = Form.useForm(); + return ( + {}} + handleAddSSOCancel={() => {}} + handleShowInstructions={() => {}} + handleInstructionsOk={() => {}} + handleInstructionsCancel={() => {}} + form={form} + accessToken={null} + ssoConfigured={false} + /> + ); + }; + + render(); + + fireEvent.mouseDown(screen.getByLabelText("SSO Provider")); + + await waitFor(() => { + expect(screen.getAllByAltText("Google SSO logo").length).toBeGreaterThan(0); + }); + + expect(screen.getAllByAltText("Google SSO logo")[0]).toHaveAttribute("src", expect.stringContaining("google.svg")); + expect(screen.getAllByAltText("Microsoft SSO logo")[0]).toHaveAttribute( + "src", + expect.stringContaining("microsoft_azure.svg"), + ); + expect(screen.getAllByAltText("Okta / Auth0 SSO logo")[0]).toHaveAttribute( + "src", + expect.stringContaining("https://www.okta.com/"), + ); + expect(screen.queryByAltText("Generic SSO logo")).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/SSOModals.tsx b/ui/litellm-dashboard/src/components/SSOModals.tsx index 9b6cd40e0c6..88ce72573d7 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.tsx @@ -4,6 +4,8 @@ import { Text, TextInput } from "@tremor/react"; import { getSSOSettings, updateSSOSettings } from "./networking"; import NotificationsManager from "./molecules/notifications_manager"; import { parseErrorMessage } from "./shared/errorUtils"; +import { Logo } from "@/components/molecules/logo/Logo"; +import { ssoProviderDisplayNames, ssoProviderLogoMap } from "./Settings/AdminSettings/SSOSettings/constants"; interface SSOModalsProps { isAddSSOModalVisible: boolean; @@ -18,13 +20,6 @@ interface SSOModalsProps { ssoConfigured?: boolean; // Add optional prop to indicate if SSO is configured } -const ssoProviderLogoMap: Record = { - google: "https://artificialanalysis.ai/img/logos/google_small.svg", - microsoft: "https://upload.wikimedia.org/wikipedia/commons/a/a8/Microsoft_Azure_Logo.svg", - okta: "https://www.okta.com/sites/default/files/Okta_Logo_BrightBlue_Medium.png", - generic: "", -}; - // Define the SSO provider configuration type interface SSOProviderConfig { envVarMap: Record; @@ -340,17 +335,14 @@ const SSOModals: React.FC = ({
{logo && ( - {value} )} - {value.toLowerCase() === "okta" - ? "Okta / Auth0" - : value.charAt(0).toUpperCase() + value.slice(1)}{" "} - SSO + {ssoProviderDisplayNames[value] || value.charAt(0).toUpperCase() + value.slice(1) + " SSO"}
diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx index c68e2716f5b..21132fff63b 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx @@ -293,4 +293,40 @@ describe("renderProviderFields", () => { expect(result).not.toBeNull(); expect(result?.length).toBe(5); }); + + it("renders provider logos in the dropdown and falls back to a letter avatar on load error", async () => { + const TestWrapper = () => { + const [form] = Form.useForm(); + return ; + }; + + renderWithProviders(); + + await act(async () => { + fireEvent.mouseDown(screen.getByLabelText("SSO Provider")); + }); + + await waitFor(() => { + expect(screen.getAllByAltText("Google SSO logo").length).toBeGreaterThan(0); + }); + + expect(screen.getAllByAltText("Google SSO logo")[0]).toHaveAttribute("src", expect.stringContaining("google.svg")); + expect(screen.getAllByAltText("Microsoft SSO logo")[0]).toHaveAttribute( + "src", + expect.stringContaining("microsoft_azure.svg"), + ); + expect(screen.queryByAltText("Generic SSO logo")).not.toBeInTheDocument(); + + const oktaLogo = screen.getAllByAltText("Okta / Auth0 SSO logo")[0]; + expect(oktaLogo).toHaveAttribute("src", expect.stringContaining("https://www.okta.com/")); + + await act(async () => { + fireEvent.error(oktaLogo); + }); + + await waitFor(() => { + expect(screen.queryByAltText("Okta / Auth0 SSO logo")).not.toBeInTheDocument(); + expect(screen.getByText("O")).toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx index d16b04466e0..6971c107a73 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx @@ -4,6 +4,7 @@ import { TextInput } from "@tremor/react"; import { Checkbox, Form, Input, Select } from "antd"; import React from "react"; import { ssoProviderLogoMap, ssoProviderDisplayNames } from "../constants"; +import { Logo } from "@/components/molecules/logo/Logo"; export interface BaseSSOSettingsFormProps { form: any; // Replace with proper Form type if available @@ -117,10 +118,10 @@ const BaseSSOSettingsForm: React.FC = ({ form, onFormS
{logo && ( - {value} )} diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.test.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.test.tsx index 5e7908a872b..e585bec4fd5 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.test.tsx @@ -1,14 +1,13 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { render, screen } from "@testing-library/react"; -import { describe, expect, it, vi } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; import SSOSettings from "./SSOSettings"; +const mockUseSSOSettings = vi.fn(); + // Mock the useSSOSettings hook vi.mock("@/app/(dashboard)/hooks/sso/useSSOSettings", () => ({ - useSSOSettings: () => ({ - data: null, - refetch: vi.fn(), - }), + useSSOSettings: () => mockUseSSOSettings(), })); const createQueryClient = () => @@ -21,17 +20,61 @@ const createQueryClient = () => }, }); -describe("SSOSettings", () => { - it("should render", () => { - const queryClient = createQueryClient(); +const renderSSOSettings = () => { + const queryClient = createQueryClient(); - render( - - - , - ); + return render( + + + , + ); +}; + +const googleConfiguredValues = { + google_client_id: "google-client-id", + google_client_secret: "google-client-secret", + microsoft_client_id: null, + microsoft_client_secret: null, + microsoft_tenant: null, + generic_client_id: null, + generic_client_secret: null, + generic_authorization_endpoint: null, + generic_token_endpoint: null, + generic_userinfo_endpoint: null, + proxy_base_url: null, + user_email: null, + ui_access_mode: null, + role_mappings: null, + team_mappings: null, +}; + +describe("SSOSettings", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockUseSSOSettings.mockReturnValue({ + data: null, + isLoading: false, + refetch: vi.fn(), + }); + }); + + it("should render", () => { + renderSSOSettings(); expect(screen.getByText("SSO Configuration")).toBeInTheDocument(); expect(screen.getByText("Manage Single Sign-On authentication settings")).toBeInTheDocument(); }); + + it("shows the local google logo asset for a google-configured settings payload", () => { + mockUseSSOSettings.mockReturnValue({ + data: { values: googleConfiguredValues }, + isLoading: false, + refetch: vi.fn(), + }); + + renderSSOSettings(); + + const logo = screen.getByAltText("Google SSO logo"); + expect(logo).toHaveAttribute("src", expect.stringContaining("google.svg")); + }); }); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx index 053da380103..e3361050422 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx @@ -4,6 +4,7 @@ import { useSSOSettings, type SSOSettingsValues } from "@/app/(dashboard)/hooks/ import { Button, Card, Descriptions, Space, Tag, Typography } from "antd"; import { Edit, Shield, Trash2 } from "lucide-react"; import { useState } from "react"; +import { Logo } from "@/components/molecules/logo/Logo"; import { ssoProviderDisplayNames, ssoProviderLogoMap } from "./constants"; import AddSSOSettingsModal from "./Modals/AddSSOSettingsModal"; import DeleteSSOSettingsModal from "./Modals/DeleteSSOSettingsModal"; @@ -166,10 +167,10 @@ export default function SSOSettings() {
{ssoProviderLogoMap[selectedProvider] && ( - {selectedProvider} )} {config.providerText} diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/constants.ts b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/constants.ts index e2aa21e4b25..b5f5ccb1b8c 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/constants.ts +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/constants.ts @@ -1,7 +1,10 @@ +import googleLogo from "../../../../../public/assets/logos/google.svg"; +import microsoftAzureLogo from "../../../../../public/assets/logos/microsoft_azure.svg"; + // SSO Provider logos export const ssoProviderLogoMap: Record = { - google: "https://artificialanalysis.ai/img/logos/google_small.svg", - microsoft: "https://upload.wikimedia.org/wikipedia/commons/a/a8/Microsoft_Azure_Logo.svg", + google: googleLogo.src, + microsoft: microsoftAzureLogo.src, okta: "https://www.okta.com/sites/default/files/Okta_Logo_BrightBlue_Medium.png", generic: "", }; diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx index 513054aae7a..95f45ea199e 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx @@ -174,6 +174,13 @@ it("should render VirtualKeysTable component", () => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); }); +it("shows the Budget Reset column by default", async () => { + renderWithProviders(); + await waitFor(() => { + expect(screen.getByText("Budget Reset")).toBeInTheDocument(); + }); +}); + it("left-anchors the create-key CTA below the title, between the header and the table toolbar", () => { renderWithProviders(Create New Key} />); @@ -498,8 +505,13 @@ describe("Status column reflects blocked / expiry / scim metadata", () => { renderWithProviders(); + const tag = await screen.findByTestId(`key-status-${mockKey.token_id}`); + expect(tag).toHaveTextContent("Active"); + + const user = userEvent.setup(); + await user.hover(tag); await waitFor(() => { - expect(screen.getByTestId(`key-status-${mockKey.token_id}`)).toHaveTextContent("Active"); + expect(screen.getByText(/not blocked and has not expired/i)).toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx index 133ff89a898..fdbc07ee020 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx @@ -46,7 +46,11 @@ const getKeyStatus = (key: KeyResponse): KeyStatus => { if (!Number.isNaN(expiresAt) && expiresAt < Date.now()) { return { tone: "warning", label: "Expired", tooltip: "This key has passed its expiry date." }; } - return { tone: "success", label: "Active" }; + return { + tone: "success", + label: "Active", + tooltip: "This key is not blocked and has not expired.", + }; }; const UserPopoverCell = ({ @@ -359,6 +363,5 @@ export const KEY_TABLE_HIDDEN_COLUMNS: Record = { created_by: false, updated_at: false, expires: false, - budget_reset_at: false, rate_limits: false, }; diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index 73cfdce5263..a99f8048dca 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -10,7 +10,7 @@ import React, { useEffect, useMemo, useState } from "react"; import TeamDropdown from "../common_components/team_dropdown"; import type { Team } from "../key_team_helpers/key_list"; import { type CredentialItem, type ProviderCreateInfo, modelAvailableCall } from "../networking"; -import { Providers, providerLogoMap } from "../provider_info_helpers"; +import { Providers } from "../provider_info_helpers"; import { ProviderLogo } from "../molecules/models/ProviderLogo"; import AdvancedSettings from "./advanced_settings"; import ConditionalPublicModelName from "./conditional_public_model_name"; @@ -181,7 +181,6 @@ const AddModelForm: React.FC = ({ {sortedProviderMetadata.map((providerInfo) => { const displayName = providerInfo.provider_display_name; const providerKey = providerInfo.provider; - const logoSrc = providerLogoMap[displayName] ?? ""; return ( diff --git a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx index ab09968c4dd..9e4a76d859d 100644 --- a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx @@ -1,19 +1,28 @@ +import arizeLogo from "../../public/assets/logos/arize.png"; +import awsLogo from "../../public/assets/logos/aws.svg"; +import braintrustLogo from "../../public/assets/logos/braintrust.png"; +import datadogLogo from "../../public/assets/logos/datadog.png"; +import galileoLogo from "../../public/assets/logos/galileo.ico"; +import lagoLogo from "../../public/assets/logos/lago.svg"; +import langfuseLogo from "../../public/assets/logos/langfuse.png"; +import langsmithLogo from "../../public/assets/logos/langsmith.png"; +import openmeterLogo from "../../public/assets/logos/openmeter.png"; +import otelLogo from "../../public/assets/logos/otel.png"; + interface CallbackConfig { id: string; displayName: string; - logo: string; + logo?: string; supports_key_team_logging: boolean; dynamic_params: Record; description: string; } -const asset_logos_folder = "/ui/assets/logos/"; - export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "arize", displayName: "Arize", - logo: `${asset_logos_folder}arize.png`, + logo: arizeLogo.src, // OTEL v2 destination: assigned per identity via the "Logging Exporters" field // (metadata.logging_exporters), not configured as a per-team callback here. supports_key_team_logging: false, @@ -23,7 +32,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "braintrust", displayName: "Braintrust", - logo: `${asset_logos_folder}braintrust.png`, + logo: braintrustLogo.src, supports_key_team_logging: false, dynamic_params: { braintrust_api_key: "password", @@ -34,7 +43,6 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "custom_callback_api", displayName: "Custom Callback API", - logo: `${asset_logos_folder}custom.svg`, supports_key_team_logging: true, dynamic_params: { custom_callback_api_url: "text", @@ -45,7 +53,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "galileo", displayName: "Galileo", - logo: `${asset_logos_folder}galileo.ico`, + logo: galileoLogo.src, supports_key_team_logging: false, dynamic_params: { GALILEO_API_KEY: "password", @@ -60,7 +68,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "datadog", displayName: "Datadog", - logo: `${asset_logos_folder}datadog.png`, + logo: datadogLogo.src, supports_key_team_logging: false, dynamic_params: { dd_api_key: "password", @@ -71,7 +79,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "lago", displayName: "Lago", - logo: `${asset_logos_folder}lago.svg`, + logo: lagoLogo.src, supports_key_team_logging: false, dynamic_params: { lago_api_url: "text", @@ -82,7 +90,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "langfuse", displayName: "Langfuse", - logo: `${asset_logos_folder}langfuse.png`, + logo: langfuseLogo.src, supports_key_team_logging: true, dynamic_params: { langfuse_public_key: "text", @@ -94,7 +102,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "langfuse_otel", displayName: "Langfuse OTEL", - logo: `${asset_logos_folder}langfuse.png`, + logo: langfuseLogo.src, // OTEL v2 destination: assigned per identity via the "Logging Exporters" field // (metadata.logging_exporters), not configured as a per-team callback here. supports_key_team_logging: false, @@ -104,7 +112,6 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "weave_otel", displayName: "Weave OTEL", - logo: `${asset_logos_folder}weave.png`, // OTEL v2 destination: assigned per identity via the "Logging Exporters" field // (metadata.logging_exporters), not configured as a per-team callback here. supports_key_team_logging: false, @@ -114,7 +121,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "langsmith", displayName: "LangSmith", - logo: `${asset_logos_folder}langsmith.png`, + logo: langsmithLogo.src, supports_key_team_logging: true, dynamic_params: { langsmith_api_key: "password", @@ -127,7 +134,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "openmeter", displayName: "OpenMeter", - logo: `${asset_logos_folder}openmeter.png`, + logo: openmeterLogo.src, supports_key_team_logging: false, dynamic_params: { openmeter_api_key: "password", @@ -138,7 +145,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "otel", displayName: "Open Telemetry", - logo: `${asset_logos_folder}otel.png`, + logo: otelLogo.src, supports_key_team_logging: false, dynamic_params: { otel_endpoint: "text", @@ -149,7 +156,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "s3", displayName: "S3", - logo: `${asset_logos_folder}aws.svg`, + logo: awsLogo.src, supports_key_team_logging: false, dynamic_params: { s3_bucket_name: "text", @@ -162,7 +169,7 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ { id: "SQS", displayName: "SQS", - logo: `${asset_logos_folder}aws.svg`, + logo: awsLogo.src, supports_key_team_logging: false, dynamic_params: { sqs_queue_url: "text", diff --git a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.test.tsx b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.test.tsx new file mode 100644 index 00000000000..656ef157363 --- /dev/null +++ b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.test.tsx @@ -0,0 +1,88 @@ +import React from "react"; +import { render, screen, fireEvent } from "@testing-library/react"; +import { describe, it, expect, vi, afterEach } from "vitest"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import MCPAppsPanel from "./MCPAppsPanel"; +import { fetchMCPServers, listMCPTools } from "../networking"; +import type { MCPServer } from "../mcp_tools/types"; +import { setServerRootPath } from "@/lib/serverRootPath"; + +vi.mock("../networking", () => ({ + fetchMCPServers: vi.fn(), + getMCPOAuthUserCredentialStatus: vi.fn(), + listMCPTools: vi.fn(), + deleteMCPOAuthUserCredential: vi.fn(), +})); + +vi.mock("@/hooks/useUserMcpOAuthFlow", () => ({ + useUserMcpOAuthFlow: () => ({ startOAuthFlow: vi.fn(), status: "idle" }), +})); + +const servers = [ + { + server_id: "s-ext", + server_name: "external_logo", + auth_type: "none", + mcp_info: { server_name: "external_logo", logo_url: "https://cdn.example.com/ext.png" }, + }, + { + server_id: "s-local", + server_name: "local_logo", + auth_type: "none", + mcp_info: { server_name: "local_logo", logo_url: "/ui/assets/logos/github.svg" }, + }, + { + server_id: "s-none", + server_name: "no_logo", + auth_type: "none", + }, +] as MCPServer[]; + +const renderPanel = () => + render( + + + , + ); + +describe("MCPAppsPanel logos", () => { + afterEach(() => { + setServerRootPath("/"); + }); + + it("resolves backend logo_url values in the server grid", async () => { + setServerRootPath("/litellm"); + vi.mocked(fetchMCPServers).mockResolvedValue(servers); + vi.mocked(listMCPTools).mockResolvedValue({ tools: [] }); + + renderPanel(); + + expect(await screen.findByText("external_logo")).toBeInTheDocument(); + expect(screen.getByAltText("external_logo logo").getAttribute("src")).toBe("https://cdn.example.com/ext.png"); + expect(screen.getByAltText("local_logo logo").getAttribute("src")).toBe("/litellm/ui/assets/logos/github.svg"); + }); + + it("renders a colored letter avatar for servers without logo_url", async () => { + vi.mocked(fetchMCPServers).mockResolvedValue(servers); + vi.mocked(listMCPTools).mockResolvedValue({ tools: [] }); + + renderPanel(); + + expect(await screen.findByText("no_logo")).toBeInTheDocument(); + expect(screen.queryByAltText("no_logo logo")).not.toBeInTheDocument(); + expect(screen.getByText("N")).toBeInTheDocument(); + }); + + it("resolves the logo_url in the detail header", async () => { + setServerRootPath("/litellm"); + vi.mocked(fetchMCPServers).mockResolvedValue(servers); + vi.mocked(listMCPTools).mockResolvedValue({ tools: [] }); + + renderPanel(); + + fireEvent.click(await screen.findByText("local_logo")); + + expect(await screen.findByRole("heading", { name: "local_logo" })).toBeInTheDocument(); + expect(screen.getByAltText("local_logo logo").getAttribute("src")).toBe("/litellm/ui/assets/logos/github.svg"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx index 867522090d2..25ced2d62c3 100644 --- a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx +++ b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx @@ -14,6 +14,7 @@ import { listMCPTools, } from "../networking"; import { AUTH_TYPE, MCPServer, MCPTool, handleTransport } from "../mcp_tools/types"; +import { Logo } from "@/components/molecules/logo/Logo"; import MessageManager from "@/components/molecules/message_manager"; import { useUserMcpOAuthFlow } from "@/hooks/useUserMcpOAuthFlow"; @@ -270,26 +271,19 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange
{detailServer.mcp_info?.logo_url ? ( - {`${name} { - const el = e.target as HTMLImageElement; - el.style.display = "none"; - if (el.nextElementSibling) (el.nextElementSibling as HTMLElement).style.display = "flex"; - }} /> - ) : null} -
- {name.charAt(0).toUpperCase()} -
+ ) : ( +
+ {name.charAt(0).toUpperCase()} +
+ )}

{name}

{detailServer.description ?? "MCP server"}

@@ -478,26 +472,19 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange } ${Math.floor(idx / 2) < Math.floor((filtered.length - 1) / 2) ? "border-b" : ""}`} > {server.mcp_info?.logo_url ? ( - {`${name} { - const el = e.target as HTMLImageElement; - el.style.display = "none"; - if (el.nextElementSibling) (el.nextElementSibling as HTMLElement).style.display = "flex"; - }} /> - ) : null} -
- {name.charAt(0).toUpperCase()} -
+ ) : ( +
+ {name.charAt(0).toUpperCase()} +
+ )}
{name}
diff --git a/ui/litellm-dashboard/src/components/chat/MCPConnectPicker.test.tsx b/ui/litellm-dashboard/src/components/chat/MCPConnectPicker.test.tsx new file mode 100644 index 00000000000..f2912d460f0 --- /dev/null +++ b/ui/litellm-dashboard/src/components/chat/MCPConnectPicker.test.tsx @@ -0,0 +1,55 @@ +import React from "react"; +import { render, screen } from "@testing-library/react"; +import { describe, it, expect, vi, afterEach } from "vitest"; +import MCPConnectPicker from "./MCPConnectPicker"; +import { fetchMCPServers } from "../networking"; +import type { MCPServer } from "../mcp_tools/types"; +import { setServerRootPath } from "@/lib/serverRootPath"; + +vi.mock("../networking", () => ({ + fetchMCPServers: vi.fn(), + listMCPTools: vi.fn(), +})); + +const servers = [ + { + server_id: "s-ext", + server_name: "external_logo", + mcp_info: { server_name: "external_logo", logo_url: "https://cdn.example.com/ext.png" }, + }, + { + server_id: "s-local", + server_name: "local_logo", + mcp_info: { server_name: "local_logo", logo_url: "/ui/assets/logos/github.svg" }, + }, + { + server_id: "s-none", + server_name: "no_logo", + }, +] as MCPServer[]; + +describe("MCPConnectPicker logos", () => { + afterEach(() => { + setServerRootPath("/"); + }); + + it("resolves backend logo_url values through the Logo component", async () => { + setServerRootPath("/litellm"); + vi.mocked(fetchMCPServers).mockResolvedValue(servers); + + render(); + + expect(await screen.findByText("external_logo")).toBeInTheDocument(); + expect(screen.getByAltText("external_logo logo").getAttribute("src")).toBe("https://cdn.example.com/ext.png"); + expect(screen.getByAltText("local_logo logo").getAttribute("src")).toBe("/litellm/ui/assets/logos/github.svg"); + }); + + it("renders no logo at all for servers without logo_url", async () => { + vi.mocked(fetchMCPServers).mockResolvedValue(servers); + + render(); + + expect(await screen.findByText("no_logo")).toBeInTheDocument(); + expect(screen.queryByAltText("no_logo logo")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/chat/MCPConnectPicker.tsx b/ui/litellm-dashboard/src/components/chat/MCPConnectPicker.tsx index abeccb0f041..a353457946c 100644 --- a/ui/litellm-dashboard/src/components/chat/MCPConnectPicker.tsx +++ b/ui/litellm-dashboard/src/components/chat/MCPConnectPicker.tsx @@ -3,6 +3,7 @@ import { Loader2 } from "lucide-react"; import { Switch } from "@/components/ui/switch"; import { Skeleton } from "@/components/ui/skeleton"; import MessageManager from "@/components/molecules/message_manager"; +import { Logo } from "@/components/molecules/logo/Logo"; import { fetchMCPServers, listMCPTools } from "../networking"; import { MCPServer } from "../mcp_tools/types"; @@ -98,13 +99,10 @@ const MCPConnectPicker: React.FC = ({ accessToken, selectedServers, onCha return (
{server.mcp_info?.logo_url && ( - {`${name} { - (e.target as HTMLImageElement).style.display = "none"; - }} /> )}
diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index b0d534e9b7c..355db8bdfc7 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -244,7 +244,13 @@ const menuGroups: MenuGroup[] = [ icon: , external_url: "https://models.litellm.ai/cookbook", }, - { key: "caching", page: "caching", label: "Caching", icon: , roles: all_admin_roles }, + { + key: "caching", + page: "caching", + label: "Response Cache", + icon: , + roles: all_admin_roles, + }, { key: "experimental", page: "experimental", diff --git a/ui/litellm-dashboard/src/components/logging_settings_view.test.tsx b/ui/litellm-dashboard/src/components/logging_settings_view.test.tsx new file mode 100644 index 00000000000..b0a1b54bac7 --- /dev/null +++ b/ui/litellm-dashboard/src/components/logging_settings_view.test.tsx @@ -0,0 +1,40 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; +import { LoggingSettingsView } from "./logging_settings_view"; + +describe("LoggingSettingsView logos", () => { + it("renders the bundled logo for a known logging integration", () => { + render( + , + ); + + expect(screen.getByAltText("Langfuse logo")).toHaveAttribute("src", "/_next/static/media/langfuse.png"); + }); + + it("renders the bundled logo for a disabled callback given by internal slug", () => { + render(); + + expect(screen.getByAltText("Datadog logo")).toHaveAttribute("src", "/_next/static/media/datadog.png"); + }); + + it("renders a letter avatar for an unknown callback name", () => { + render( + , + ); + + expect(document.querySelector("img")).toBeNull(); + expect(screen.getByText("m")).toBeInTheDocument(); + expect(screen.getByText("mystery_callback")).toBeInTheDocument(); + }); + + it("renders a letter avatar for the custom callback API, which has no bundled logo", () => { + render(); + + expect(screen.queryByAltText("Custom Callback API logo")).toBeNull(); + expect(screen.getByText("C")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/logging_settings_view.tsx b/ui/litellm-dashboard/src/components/logging_settings_view.tsx index 2116a08ed23..9417634ab8e 100644 --- a/ui/litellm-dashboard/src/components/logging_settings_view.tsx +++ b/ui/litellm-dashboard/src/components/logging_settings_view.tsx @@ -2,7 +2,7 @@ import React from "react"; import { Tag } from "antd"; import { CogIcon, BanIcon } from "@heroicons/react/outline"; import { callbackInfo, callback_map, reverse_callback_map } from "./callback_info_helpers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; interface LoggingConfig { callback_name: string; @@ -117,7 +117,6 @@ export function LoggingSettingsView({
{loggingConfigs.map((config, index) => { const displayName = getLoggingDisplayName(config.callback_name); - const logoUrl = resolveLogoSrc(callbackInfo[displayName]?.logo); return (
- {logoUrl ? ( - {displayName} - ) : ( - - )} +
{displayName} @@ -163,7 +162,6 @@ export function LoggingSettingsView({ {disabledCallbacks.map((callbackName, index) => { // Handle both display names and internal values const displayName = reverse_callback_map[callbackName] || callbackName; - const logoUrl = resolveLogoSrc(callbackInfo[displayName]?.logo); return (
- {logoUrl ? ( - {displayName} - ) : ( - - )} +
{displayName} Disabled for this key diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx index c92a4a90578..534e06dfe52 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx @@ -4,8 +4,8 @@ import type { UploadProps } from "antd/es/upload"; import { useState } from "react"; import ProviderSpecificFields from "../add_model/provider_specific_fields"; import { CredentialItem } from "../networking"; -import { Providers, providerLogoMap } from "../provider_info_helpers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Providers } from "../provider_info_helpers"; +import { Logo } from "@/components/molecules/logo/Logo"; import { resetCredentialFormOnProviderChange } from "./credential_form_helpers"; const { Link } = Typography; @@ -92,22 +92,7 @@ export default function CredentialModal({ {Object.entries(Providers).map(([providerEnum, providerDisplayName]) => (
- {`${providerEnum} { - 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 = providerDisplayName.charAt(0); - parent.replaceChild(fallbackDiv, target); - } - }} - /> + {providerDisplayName}
diff --git a/ui/litellm-dashboard/src/components/model_info_view.test.tsx b/ui/litellm-dashboard/src/components/model_info_view.test.tsx index a496bb05b91..6cf06d759f2 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.test.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.test.tsx @@ -916,4 +916,37 @@ describe("ModelInfoView", () => { expect(screen.getByText(/Created By/)).toBeInTheDocument(); }); }); + + it("renders the provider card logo from the bundled provider map", async () => { + render(, { wrapper }); + + const logo = await screen.findByAltText("openai logo"); + expect(logo.getAttribute("src")).toContain("openai_small"); + }); + + it("renders a letter avatar instead of an img for an unknown provider slug", async () => { + mockUseModelsInfo.mockReturnValue({ + data: { + data: [ + { + ...defaultModelData, + litellm_params: { + ...defaultModelData.litellm_params, + custom_llm_provider: "zzz-internal", + }, + }, + ], + }, + isLoading: false, + error: null, + }); + + render(, { wrapper }); + + await waitFor(() => { + expect(screen.getAllByText("zzz-internal").length).toBeGreaterThan(0); + }); + expect(screen.queryByAltText("zzz-internal logo")).not.toBeInTheDocument(); + expect(screen.getByText("z")).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/model_info_view.tsx b/ui/litellm-dashboard/src/components/model_info_view.tsx index 8aaabdc50a2..fe28e0fb40c 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.tsx @@ -43,7 +43,7 @@ import { tagListCall, testConnectionRequest, } from "./networking"; -import { getProviderLogoAndName } from "./provider_info_helpers"; +import { Logo } from "@/components/molecules/logo/Logo"; import UpdateModelCredentialsModal from "./update_model_credentials_modal"; import NumericalInput from "./shared/numerical_input"; import { Tag } from "./tag_management/types"; @@ -660,30 +660,7 @@ export default function ModelInfoView({ Provider
- {modelData.provider && ( - {`${modelData.provider} { - const target = e.currentTarget as HTMLImageElement; - const parent = target.parentElement; - if (!parent || !parent.contains(target)) { - return; - } - - try { - const fallbackDiv = document.createElement("div"); - fallbackDiv.className = - "w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs"; - fallbackDiv.textContent = modelData.provider?.charAt(0) || "-"; - parent.replaceChild(fallbackDiv, target); - } catch (error) { - console.error("Failed to replace provider logo fallback:", error); - } - }} - /> - )} + {modelData.provider && } {modelData.provider || "Not Set"}
diff --git a/ui/litellm-dashboard/src/components/molecules/logo/Logo.test.tsx b/ui/litellm-dashboard/src/components/molecules/logo/Logo.test.tsx new file mode 100644 index 00000000000..95221f35b8d --- /dev/null +++ b/ui/litellm-dashboard/src/components/molecules/logo/Logo.test.tsx @@ -0,0 +1,73 @@ +import React from "react"; +import { describe, expect, it, vi } from "vitest"; +import { act, fireEvent, render, screen } from "@testing-library/react"; +import { Logo } from "./Logo"; +import { Providers, providerLogoMap } from "@/components/provider_info_helpers"; + +vi.mock("@/lib/serverRootPath", () => ({ serverRootPath: "/litellm" })); + +describe("Logo", () => { + it("renders the bundled logo untouched by the server root path for a known provider", () => { + render(); + const img = screen.getByRole("img", { name: "openai logo" }); + expect(img.getAttribute("src")).toBe(providerLogoMap[Providers.OpenAI]); + expect(img.getAttribute("src")).toContain("openai_small"); + }); + + it("renders a letter avatar and no img for an unknown provider", () => { + render(); + expect(screen.getByText("u")).toBeInTheDocument(); + expect(screen.queryByRole("img")).not.toBeInTheDocument(); + }); + + it("renders a dash avatar when neither provider nor src is given", () => { + render(); + expect(screen.getByText("-")).toBeInTheDocument(); + expect(screen.queryByRole("img")).not.toBeInTheDocument(); + }); + + it("resolves a backend asset path through the server root path in src mode", () => { + render(); + const img = screen.getByRole("img", { name: "GitHub logo" }); + expect(img.getAttribute("src")).toBe("/litellm/ui/assets/logos/github.svg"); + }); + + it("passes an external https URL through untouched in src mode", () => { + render(); + expect(screen.getByRole("img").getAttribute("src")).toBe("https://cdn.example.com/logo.png"); + }); + + it("swaps to the letter avatar and warns with the failing URL on image error", () => { + const warnSpy = vi.spyOn(console, "warn").mockImplementation(() => {}); + render(); + const img = screen.getByRole("img", { name: "GitHub logo" }); + + act(() => { + fireEvent.error(img); + }); + + expect(screen.queryByRole("img")).not.toBeInTheDocument(); + expect(screen.getByText("G")).toBeInTheDocument(); + expect(warnSpy).toHaveBeenCalledWith(expect.stringContaining("/litellm/ui/assets/logos/github.svg")); + warnSpy.mockRestore(); + }); + + it("retries with a new src after a previous src errored", () => { + const warnSpy = vi.spyOn(console, "warn").mockImplementation(() => {}); + const { rerender } = render(); + + act(() => { + fireEvent.error(screen.getByRole("img", { name: "Agent logo" })); + }); + expect(screen.queryByRole("img")).not.toBeInTheDocument(); + + rerender(); + const img = screen.getByRole("img", { name: "Agent logo" }); + expect(img.getAttribute("src")).toBe("/litellm/ui/assets/logos/github.svg"); + + rerender(); + expect(screen.queryByRole("img")).not.toBeInTheDocument(); + expect(screen.getByText("A")).toBeInTheDocument(); + warnSpy.mockRestore(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/molecules/logo/Logo.tsx b/ui/litellm-dashboard/src/components/molecules/logo/Logo.tsx new file mode 100644 index 00000000000..6f18c89ae8e --- /dev/null +++ b/ui/litellm-dashboard/src/components/molecules/logo/Logo.tsx @@ -0,0 +1,36 @@ +import React, { useState } from "react"; +import { getProviderLogoAndName } from "@/components/provider_info_helpers"; +import { resolveLogoSrc } from "@/lib/assetPaths"; + +interface LogoProps { + provider?: string; + src?: string | null; + label?: string; + className?: string; +} + +export const Logo: React.FC = ({ provider, src, label, className = "w-4 h-4" }) => { + const [erroredSrc, setErroredSrc] = useState(null); + const resolvedSrc = provider !== undefined ? getProviderLogoAndName(provider).logo : resolveLogoSrc(src) ?? ""; + const name = label ?? provider ?? ""; + + if (erroredSrc === resolvedSrc || !resolvedSrc) { + return ( +
+ {name.charAt(0) || "-"} +
+ ); + } + + return ( + {`${name { + console.warn(`Logo failed to load: ${resolvedSrc}`); + setErroredSrc(resolvedSrc); + }} + /> + ); +}; diff --git a/ui/litellm-dashboard/src/components/molecules/models/ProviderLogo.tsx b/ui/litellm-dashboard/src/components/molecules/models/ProviderLogo.tsx index 4a9da15e333..4bc11126ae4 100644 --- a/ui/litellm-dashboard/src/components/molecules/models/ProviderLogo.tsx +++ b/ui/litellm-dashboard/src/components/molecules/models/ProviderLogo.tsx @@ -1,24 +1,11 @@ -import React, { useState } from "react"; -import { getProviderLogoAndName } from "../../provider_info_helpers"; +import React from "react"; +import { Logo } from "@/components/molecules/logo/Logo"; interface ProviderLogoProps { provider: string; className?: string; } -export const ProviderLogo: React.FC = ({ provider, className = "w-4 h-4" }) => { - const [hasError, setHasError] = useState(false); - const { logo } = getProviderLogoAndName(provider); - - const showFallback = hasError || !logo; - - if (showFallback) { - return ( -
- {provider?.charAt(0) || "-"} -
- ); - } - - return {`${provider} setHasError(true)} />; -}; +export const ProviderLogo: React.FC = ({ provider, className = "w-4 h-4" }) => ( + +); diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index ea39dab3317..748f3ab19fd 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -389,7 +389,7 @@ export function getGlobalLitellmHeaderName(): string { return globalLitellmHeaderName; } -const apiClient = createApiClient({ +export const apiClient = createApiClient({ getBaseUrl: getProxyBaseUrl, getAuthHeaderName: getGlobalLitellmHeaderName, onError: handleError, diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx index 6e19de4a6cc..966bd4e262c 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx @@ -113,24 +113,20 @@ describe("provider_info_helpers", () => { }); }); - describe("provider logo asset paths", () => { - // Regression: a relative "../ui/assets/logos/" base resolved to - // "/ui/ui/assets/logos/..." (404) on the public model hub at - // /ui/model_hub_table/, which sits a level below the /ui/ SPA. Root-absolute - // paths resolve correctly at any route depth. - it("should expose every provider logo as a root-absolute /ui path", () => { + describe("provider logo bundled assets", () => { + it("should expose every provider logo as a truthy bundled URL, never a raw /ui/assets path", () => { const logos = Object.values(providerLogoMap); expect(logos.length).toBeGreaterThan(0); logos.forEach((logo) => { - expect(logo.startsWith("/ui/assets/logos/")).toBe(true); - expect(logo).not.toContain("../"); + expect(typeof logo).toBe("string"); + expect(logo.length).toBeGreaterThan(0); + expect(logo.startsWith("/ui/assets/")).toBe(false); }); }); - it("should resolve a provider logo to a root-absolute path via getProviderLogoAndName", () => { + it("should resolve a provider to its own bundled logo via getProviderLogoAndName", () => { const { logo } = getProviderLogoAndName("openai"); - expect(logo.startsWith("/ui/assets/logos/")).toBe(true); - expect(logo).not.toContain("../"); + expect(logo).toContain("openai_small"); }); }); @@ -430,20 +426,19 @@ describe("getProviderLogoAndName under a custom server_root_path", () => { vi.doUnmock("@/lib/serverRootPath"); }); - // Regression: under SERVER_ROOT_PATH=/litellm the logo must be requested at - // /litellm/ui/assets/logos/... A bare /ui/... path is served off the root and - // 404s behind the reverse proxy. - it("prefixes the server root path onto the resolved logo", async () => { + it("returns the bundled logo URL untouched under a sub-path mount", async () => { vi.resetModules(); vi.doMock("@/lib/serverRootPath", () => ({ serverRootPath: "/litellm" })); - const { getProviderLogoAndName } = await import("./provider_info_helpers"); - expect(getProviderLogoAndName("openai").logo).toBe("/litellm/ui/assets/logos/openai_small.svg"); + const helpers = await import("./provider_info_helpers"); + const { logo } = helpers.getProviderLogoAndName("openai"); + expect(logo).toBe(helpers.providerLogoMap[helpers.Providers.OpenAI]); + expect(logo.startsWith("/litellm")).toBe(false); }); - it("leaves the logo at /ui/... when mounted at the root", async () => { + it("returns the bundled logo URL untouched at the root mount", async () => { vi.resetModules(); vi.doMock("@/lib/serverRootPath", () => ({ serverRootPath: "/" })); - const { getProviderLogoAndName } = await import("./provider_info_helpers"); - expect(getProviderLogoAndName("openai").logo).toBe("/ui/assets/logos/openai_small.svg"); + const helpers = await import("./provider_info_helpers"); + expect(helpers.getProviderLogoAndName("openai").logo).toBe(helpers.providerLogoMap[helpers.Providers.OpenAI]); }); }); diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx index c955087b40a..f8be4511e3a 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx @@ -1,4 +1,67 @@ import { resolveLogoSrc } from "@/lib/assetPaths"; +import a2aAgentLogo from "../../public/assets/logos/a2a_agent.png"; +import ai21Logo from "../../public/assets/logos/ai21.svg"; +import aimlApiLogo from "../../public/assets/logos/aiml_api.svg"; +import anthropicLogo from "../../public/assets/logos/anthropic.svg"; +import assemblyaiSmallLogo from "../../public/assets/logos/assemblyai_small.png"; +import basetenLogo from "../../public/assets/logos/baseten.svg"; +import bedrockLogo from "../../public/assets/logos/bedrock.svg"; +import cerebrasLogo from "../../public/assets/logos/cerebras.svg"; +import cloudflareLogo from "../../public/assets/logos/cloudflare.svg"; +import cohereLogo from "../../public/assets/logos/cohere.svg"; +import cometapiLogo from "../../public/assets/logos/cometapi.svg"; +import cursorLogo from "../../public/assets/logos/cursor.svg"; +import databricksLogo from "../../public/assets/logos/databricks.svg"; +import deepgramLogo from "../../public/assets/logos/deepgram.png"; +import deepinfraLogo from "../../public/assets/logos/deepinfra.png"; +import deepseekLogo from "../../public/assets/logos/deepseek.svg"; +import elevenlabsLogo from "../../public/assets/logos/elevenlabs.png"; +import falAiLogo from "../../public/assets/logos/fal_ai.jpg"; +import featherlessLogo from "../../public/assets/logos/featherless.svg"; +import fireworksLogo from "../../public/assets/logos/fireworks.svg"; +import friendliLogo from "../../public/assets/logos/friendli.svg"; +import githubCopilotLogo from "../../public/assets/logos/github_copilot.svg"; +import googleLogo from "../../public/assets/logos/google.svg"; +import groqLogo from "../../public/assets/logos/groq.svg"; +import huggingfaceLogo from "../../public/assets/logos/huggingface.svg"; +import hyperbolicLogo from "../../public/assets/logos/hyperbolic.svg"; +import infinityLogo from "../../public/assets/logos/infinity.png"; +import jinaLogo from "../../public/assets/logos/jina.png"; +import lambdaLogo from "../../public/assets/logos/lambda.svg"; +import lmstudioLogo from "../../public/assets/logos/lmstudio.svg"; +import metaLlamaLogo from "../../public/assets/logos/meta_llama.svg"; +import microsoftAzureLogo from "../../public/assets/logos/microsoft_azure.svg"; +import minimaxLogo from "../../public/assets/logos/minimax.svg"; +import mistralLogo from "../../public/assets/logos/mistral.svg"; +import moonshotLogo from "../../public/assets/logos/moonshot.svg"; +import morphLogo from "../../public/assets/logos/morph.svg"; +import nebiusLogo from "../../public/assets/logos/nebius.svg"; +import novitaLogo from "../../public/assets/logos/novita.svg"; +import nvidiaNimLogo from "../../public/assets/logos/nvidia_nim.svg"; +import nvidiaTritonLogo from "../../public/assets/logos/nvidia_triton.png"; +import ollamaLogo from "../../public/assets/logos/ollama.svg"; +import openaiSmallLogo from "../../public/assets/logos/openai_small.svg"; +import openrouterLogo from "../../public/assets/logos/openrouter.svg"; +import oracleLogo from "../../public/assets/logos/oracle.svg"; +import perplexityAiLogo from "../../public/assets/logos/perplexity-ai.svg"; +import qwenLogo from "../../public/assets/logos/qwen.png"; +import recraftLogo from "../../public/assets/logos/recraft.svg"; +import replicateLogo from "../../public/assets/logos/replicate.svg"; +import runwayLogo from "../../public/assets/logos/runway.png"; +import sambanovaLogo from "../../public/assets/logos/sambanova.svg"; +import sapLogo from "../../public/assets/logos/sap.png"; +import snowflakeLogo from "../../public/assets/logos/snowflake.svg"; +import sonioxLogo from "../../public/assets/logos/soniox.svg"; +import togetheraiLogo from "../../public/assets/logos/togetherai.svg"; +import topazLogo from "../../public/assets/logos/topaz.svg"; +import v0Logo from "../../public/assets/logos/v0.svg"; +import vercelLogo from "../../public/assets/logos/vercel.svg"; +import vllmLogo from "../../public/assets/logos/vllm.png"; +import volcengineLogo from "../../public/assets/logos/volcengine.png"; +import voyageLogo from "../../public/assets/logos/voyage.webp"; +import watsonxLogo from "../../public/assets/logos/watsonx.svg"; +import xaiLogo from "../../public/assets/logos/xai.svg"; +import xinferenceLogo from "../../public/assets/logos/xinference.svg"; export enum Providers { A2A_Agent = "A2A Agent", @@ -220,94 +283,91 @@ export const provider_map: Record = { const standaloneSubproviderSlugs = new Set(["bedrock_mantle"]); -const asset_logos_folder = "/ui/assets/logos/"; - export const providerLogoMap: Record = { - [Providers.A2A_Agent]: `${asset_logos_folder}a2a_agent.png`, - [Providers.AI21]: `${asset_logos_folder}ai21.svg`, - [Providers.AI21_CHAT]: `${asset_logos_folder}ai21.svg`, - [Providers.AIML]: `${asset_logos_folder}aiml_api.svg`, - [Providers.AIOHTTP_OPENAI]: `${asset_logos_folder}openai_small.svg`, - [Providers.Anthropic]: `${asset_logos_folder}anthropic.svg`, - [Providers.ANTHROPIC_TEXT]: `${asset_logos_folder}anthropic.svg`, - [Providers.AssemblyAI]: `${asset_logos_folder}assemblyai_small.png`, - [Providers.Azure]: `${asset_logos_folder}microsoft_azure.svg`, - [Providers.Azure_AI_Studio]: `${asset_logos_folder}microsoft_azure.svg`, - [Providers.AZURE_TEXT]: `${asset_logos_folder}microsoft_azure.svg`, - [Providers.BASETEN]: `${asset_logos_folder}baseten.svg`, - [Providers.Bedrock]: `${asset_logos_folder}bedrock.svg`, - [Providers.BedrockMantle]: `${asset_logos_folder}bedrock.svg`, - [Providers.SageMaker]: `${asset_logos_folder}bedrock.svg`, - [Providers.Cerebras]: `${asset_logos_folder}cerebras.svg`, - [Providers.CLOUDFLARE]: `${asset_logos_folder}cloudflare.svg`, - [Providers.CODESTRAL]: `${asset_logos_folder}mistral.svg`, - [Providers.Cohere]: `${asset_logos_folder}cohere.svg`, - [Providers.COHERE_CHAT]: `${asset_logos_folder}cohere.svg`, - [Providers.COMETAPI]: `${asset_logos_folder}cometapi.svg`, - [Providers.Cursor]: `${asset_logos_folder}cursor.svg`, - [Providers.Databricks]: `${asset_logos_folder}databricks.svg`, - [Providers.Dashscope]: `${asset_logos_folder}dashscope.svg`, - [Providers.Deepseek]: `${asset_logos_folder}deepseek.svg`, - [Providers.Deepgram]: `${asset_logos_folder}deepgram.png`, - [Providers.DeepInfra]: `${asset_logos_folder}deepinfra.png`, - [Providers.ElevenLabs]: `${asset_logos_folder}elevenlabs.png`, - [Providers.FalAI]: `${asset_logos_folder}fal_ai.jpg`, - [Providers.FEATHERLESS_AI]: `${asset_logos_folder}featherless.svg`, - [Providers.FireworksAI]: `${asset_logos_folder}fireworks.svg`, - [Providers.FRIENDLIAI]: `${asset_logos_folder}friendli.svg`, - [Providers.GITHUB_COPILOT]: `${asset_logos_folder}github_copilot.svg`, - [Providers.Google_AI_Studio]: `${asset_logos_folder}google.svg`, - [Providers.GradientAI]: `${asset_logos_folder}gradientai.svg`, - [Providers.Groq]: `${asset_logos_folder}groq.svg`, - [Providers.Hosted_Vllm]: `${asset_logos_folder}vllm.png`, - [Providers.HUGGINGFACE]: `${asset_logos_folder}huggingface.svg`, - [Providers.HYPERBOLIC]: `${asset_logos_folder}hyperbolic.svg`, - [Providers.Infinity]: `${asset_logos_folder}infinity.png`, - [Providers.JinaAI]: `${asset_logos_folder}jina.png`, - [Providers.LAMBDA_AI]: `${asset_logos_folder}lambda.svg`, - [Providers.LM_STUDIO]: `${asset_logos_folder}lmstudio.svg`, - [Providers.LLAMA]: `${asset_logos_folder}meta_llama.svg`, - [Providers.MiniMax]: `${asset_logos_folder}minimax.svg`, - [Providers.MistralAI]: `${asset_logos_folder}mistral.svg`, - [Providers.MOONSHOT]: `${asset_logos_folder}moonshot.svg`, - [Providers.MORPH]: `${asset_logos_folder}morph.svg`, - [Providers.NEBIUS]: `${asset_logos_folder}nebius.svg`, - [Providers.NOVITA]: `${asset_logos_folder}novita.svg`, - [Providers.NVIDIA_NIM]: `${asset_logos_folder}nvidia_nim.svg`, - [Providers.Ollama]: `${asset_logos_folder}ollama.svg`, - [Providers.OLLAMA_CHAT]: `${asset_logos_folder}ollama.svg`, - [Providers.OOBABOOGA]: `${asset_logos_folder}openai_small.svg`, - [Providers.OpenAI]: `${asset_logos_folder}openai_small.svg`, - [Providers.OPENAI_LIKE]: `${asset_logos_folder}openai_small.svg`, - [Providers.OpenAI_Text]: `${asset_logos_folder}openai_small.svg`, - [Providers.OpenAI_Text_Compatible]: `${asset_logos_folder}openai_small.svg`, - [Providers.OpenAI_Compatible]: `${asset_logos_folder}openai_small.svg`, - [Providers.Openrouter]: `${asset_logos_folder}openrouter.svg`, - [Providers.Oracle]: `${asset_logos_folder}oracle.svg`, - [Providers.Perplexity]: `${asset_logos_folder}perplexity-ai.svg`, - [Providers.RECRAFT]: `${asset_logos_folder}recraft.svg`, - [Providers.REPLICATE]: `${asset_logos_folder}replicate.svg`, - [Providers.RunwayML]: `${asset_logos_folder}runwayml.png`, - [Providers.SAGEMAKER_LEGACY]: `${asset_logos_folder}bedrock.svg`, - [Providers.Sambanova]: `${asset_logos_folder}sambanova.svg`, - [Providers.SAP]: `${asset_logos_folder}sap.png`, - [Providers.Snowflake]: `${asset_logos_folder}snowflake.svg`, - [Providers.Soniox]: `${asset_logos_folder}soniox.svg`, - [Providers.TEXT_COMPLETION_CODESTRAL]: `${asset_logos_folder}mistral.svg`, - [Providers.TogetherAI]: `${asset_logos_folder}togetherai.svg`, - [Providers.TOPAZ]: `${asset_logos_folder}topaz.svg`, - [Providers.Triton]: `${asset_logos_folder}nvidia_triton.png`, - [Providers.V0]: `${asset_logos_folder}v0.svg`, - [Providers.VERCEL_AI_GATEWAY]: `${asset_logos_folder}vercel.svg`, - [Providers.Vertex_AI]: `${asset_logos_folder}google.svg`, - [Providers.VERTEX_AI_BETA]: `${asset_logos_folder}google.svg`, - [Providers.VLLM]: `${asset_logos_folder}vllm.png`, - [Providers.VolcEngine]: `${asset_logos_folder}volcengine.png`, - [Providers.Voyage]: `${asset_logos_folder}voyage.webp`, - [Providers.WATSONX]: `${asset_logos_folder}watsonx.svg`, - [Providers.WATSONX_TEXT]: `${asset_logos_folder}watsonx.svg`, - [Providers.xAI]: `${asset_logos_folder}xai.svg`, - [Providers.XINFERENCE]: `${asset_logos_folder}xinference.svg`, + [Providers.A2A_Agent]: a2aAgentLogo.src, + [Providers.AI21]: ai21Logo.src, + [Providers.AI21_CHAT]: ai21Logo.src, + [Providers.AIML]: aimlApiLogo.src, + [Providers.AIOHTTP_OPENAI]: openaiSmallLogo.src, + [Providers.Anthropic]: anthropicLogo.src, + [Providers.ANTHROPIC_TEXT]: anthropicLogo.src, + [Providers.AssemblyAI]: assemblyaiSmallLogo.src, + [Providers.Azure]: microsoftAzureLogo.src, + [Providers.Azure_AI_Studio]: microsoftAzureLogo.src, + [Providers.AZURE_TEXT]: microsoftAzureLogo.src, + [Providers.BASETEN]: basetenLogo.src, + [Providers.Bedrock]: bedrockLogo.src, + [Providers.BedrockMantle]: bedrockLogo.src, + [Providers.SageMaker]: bedrockLogo.src, + [Providers.Cerebras]: cerebrasLogo.src, + [Providers.CLOUDFLARE]: cloudflareLogo.src, + [Providers.CODESTRAL]: mistralLogo.src, + [Providers.Cohere]: cohereLogo.src, + [Providers.COHERE_CHAT]: cohereLogo.src, + [Providers.COMETAPI]: cometapiLogo.src, + [Providers.Cursor]: cursorLogo.src, + [Providers.Databricks]: databricksLogo.src, + [Providers.Dashscope]: qwenLogo.src, + [Providers.Deepseek]: deepseekLogo.src, + [Providers.Deepgram]: deepgramLogo.src, + [Providers.DeepInfra]: deepinfraLogo.src, + [Providers.ElevenLabs]: elevenlabsLogo.src, + [Providers.FalAI]: falAiLogo.src, + [Providers.FEATHERLESS_AI]: featherlessLogo.src, + [Providers.FireworksAI]: fireworksLogo.src, + [Providers.FRIENDLIAI]: friendliLogo.src, + [Providers.GITHUB_COPILOT]: githubCopilotLogo.src, + [Providers.Google_AI_Studio]: googleLogo.src, + [Providers.Groq]: groqLogo.src, + [Providers.Hosted_Vllm]: vllmLogo.src, + [Providers.HUGGINGFACE]: huggingfaceLogo.src, + [Providers.HYPERBOLIC]: hyperbolicLogo.src, + [Providers.Infinity]: infinityLogo.src, + [Providers.JinaAI]: jinaLogo.src, + [Providers.LAMBDA_AI]: lambdaLogo.src, + [Providers.LM_STUDIO]: lmstudioLogo.src, + [Providers.LLAMA]: metaLlamaLogo.src, + [Providers.MiniMax]: minimaxLogo.src, + [Providers.MistralAI]: mistralLogo.src, + [Providers.MOONSHOT]: moonshotLogo.src, + [Providers.MORPH]: morphLogo.src, + [Providers.NEBIUS]: nebiusLogo.src, + [Providers.NOVITA]: novitaLogo.src, + [Providers.NVIDIA_NIM]: nvidiaNimLogo.src, + [Providers.Ollama]: ollamaLogo.src, + [Providers.OLLAMA_CHAT]: ollamaLogo.src, + [Providers.OOBABOOGA]: openaiSmallLogo.src, + [Providers.OpenAI]: openaiSmallLogo.src, + [Providers.OPENAI_LIKE]: openaiSmallLogo.src, + [Providers.OpenAI_Text]: openaiSmallLogo.src, + [Providers.OpenAI_Text_Compatible]: openaiSmallLogo.src, + [Providers.OpenAI_Compatible]: openaiSmallLogo.src, + [Providers.Openrouter]: openrouterLogo.src, + [Providers.Oracle]: oracleLogo.src, + [Providers.Perplexity]: perplexityAiLogo.src, + [Providers.RECRAFT]: recraftLogo.src, + [Providers.REPLICATE]: replicateLogo.src, + [Providers.RunwayML]: runwayLogo.src, + [Providers.SAGEMAKER_LEGACY]: bedrockLogo.src, + [Providers.Sambanova]: sambanovaLogo.src, + [Providers.SAP]: sapLogo.src, + [Providers.Snowflake]: snowflakeLogo.src, + [Providers.Soniox]: sonioxLogo.src, + [Providers.TEXT_COMPLETION_CODESTRAL]: mistralLogo.src, + [Providers.TogetherAI]: togetheraiLogo.src, + [Providers.TOPAZ]: topazLogo.src, + [Providers.Triton]: nvidiaTritonLogo.src, + [Providers.V0]: v0Logo.src, + [Providers.VERCEL_AI_GATEWAY]: vercelLogo.src, + [Providers.Vertex_AI]: googleLogo.src, + [Providers.VERTEX_AI_BETA]: googleLogo.src, + [Providers.VLLM]: vllmLogo.src, + [Providers.VolcEngine]: volcengineLogo.src, + [Providers.Voyage]: voyageLogo.src, + [Providers.WATSONX]: watsonxLogo.src, + [Providers.WATSONX_TEXT]: watsonxLogo.src, + [Providers.xAI]: xaiLogo.src, + [Providers.XINFERENCE]: xinferenceLogo.src, }; export const getProviderLogoAndName = (providerValue: string): { logo: string; displayName: string } => { diff --git a/ui/litellm-dashboard/src/components/settings.test.tsx b/ui/litellm-dashboard/src/components/settings.test.tsx index 7998b54d23e..862530b71b2 100644 --- a/ui/litellm-dashboard/src/components/settings.test.tsx +++ b/ui/litellm-dashboard/src/components/settings.test.tsx @@ -1,9 +1,10 @@ -import { act, render, screen, waitFor } from "@testing-library/react"; +import { act, fireEvent, render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; +import { Form } from "antd"; import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { alertingSettingsCall, getCallbackConfigsCall, getCallbacksCall } from "./networking"; -import Settings from "./settings"; +import Settings, { backendCallbackLogoSrc, CallbackSelector } from "./settings"; type SettingsTestProps = { accessToken: string | null; @@ -249,3 +250,44 @@ describe("Settings", () => { expect(getByText("CloudZero Cost Tracking")).toBeInTheDocument(); }); }); + +describe("backendCallbackLogoSrc", () => { + it("prefixes bare filenames with the assets logo folder", () => { + expect(backendCallbackLogoSrc("datadog.png")).toBe("/ui/assets/logos/datadog.png"); + }); + + it("passes through urls, data uris, and paths untouched", () => { + expect(backendCallbackLogoSrc("https://logos.example.com/x.png")).toBe("https://logos.example.com/x.png"); + expect(backendCallbackLogoSrc("data:image/png;base64,abc")).toBe("data:image/png;base64,abc"); + expect(backendCallbackLogoSrc("/custom/path.png")).toBe("/custom/path.png"); + }); + + it("returns undefined when the backend provides no logo", () => { + expect(backendCallbackLogoSrc(undefined)).toBeUndefined(); + expect(backendCallbackLogoSrc(null)).toBeUndefined(); + expect(backendCallbackLogoSrc("")).toBeUndefined(); + }); +}); + +describe("CallbackSelector logos", () => { + it("resolves backend logos per entry: bare filename, external url, and missing logo", async () => { + const callbackConfigs = [ + { id: "langfuse", displayName: "Langfuse", logo: "langfuse.png" }, + { id: "hosted", displayName: "Hosted", logo: "https://logos.example.com/hosted.png" }, + { id: "nologo", displayName: "NoLogo" }, + ]; + + render( +
+ + , + ); + + fireEvent.mouseDown(screen.getByRole("combobox")); + + expect(await screen.findByAltText("Langfuse logo")).toHaveAttribute("src", "/ui/assets/logos/langfuse.png"); + expect(screen.getByAltText("Hosted logo")).toHaveAttribute("src", "https://logos.example.com/hosted.png"); + expect(screen.queryByAltText("NoLogo logo")).toBeNull(); + expect(screen.getByText("N")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/settings.tsx b/ui/litellm-dashboard/src/components/settings.tsx index 453e47e8f30..567e86e57ee 100644 --- a/ui/litellm-dashboard/src/components/settings.tsx +++ b/ui/litellm-dashboard/src/components/settings.tsx @@ -22,7 +22,7 @@ import React, { useEffect, useState } from "react"; import { Button as Button2, Form, Input, Modal, Select, Typography } from "antd"; import EmailSettings from "./email_settings"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import NotificationsManager from "./molecules/notifications_manager"; const { Title, Paragraph } = Typography; @@ -70,6 +70,12 @@ interface genericCallbackParams { const assetsLogoFolder = "/ui/assets/logos/"; +export const backendCallbackLogoSrc = (logo: string | null | undefined): string | undefined => { + if (!logo) return undefined; + if (logo.includes("/") || logo.startsWith("data:") || logo.startsWith("http")) return logo; + return `${assetsLogoFolder}${logo}`; +}; + interface DynamicParamsFieldsProps { params: string[]; callbackConfigs: any[]; @@ -145,7 +151,7 @@ interface CallbackSelectorProps { disabled?: boolean; } -const CallbackSelector: React.FC = ({ +export const CallbackSelector: React.FC = ({ callbackConfigs, selectedCallback, onCallbackChange, @@ -170,25 +176,14 @@ const CallbackSelector: React.FC = ({ onChange={onCallbackChange} > {callbackConfigs.map((callbackConfig) => { - const logo = callbackConfig.logo; - const logoSrc = resolveLogoSrc( - logo && (logo.includes("/") || logo.startsWith("data:") || logo.startsWith("http")) - ? logo - : `${assetsLogoFolder}${logo}`, - ); - return (
- {/* eslint-disable-next-line @next/next/no-img-element */} - {`${callbackConfig.displayName} { - e.currentTarget.style.display = "none"; - }} />
{callbackConfig.displayName} diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx index fa0e672026e..c4799594465 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx @@ -18,6 +18,7 @@ import { type OnChangeFn, type Row, type RowData, + type RowSelectionState, type Table, type TableOptions, useReactTable, @@ -70,6 +71,8 @@ export function validateDataTableConfig( const bothSortingSources = props.defaultSorting !== undefined && props.sorting !== undefined; const bothFilterSources = props.defaultColumnFilters !== undefined && props.columnFilters !== undefined; + const controlledSelectionIncomplete = props.rowSelection !== undefined && props.onRowSelectionChange === undefined; + return [ serverSortingIncomplete ? "sortingMode='server' requires both `sorting` and `onSortingChange`." : null, serverPaginationIncomplete @@ -80,6 +83,9 @@ export function validateDataTableConfig( bothFilterSources ? "Provide either `defaultColumnFilters` (uncontrolled) or `columnFilters` (controlled), not both." : null, + controlledSelectionIncomplete + ? "Controlled `rowSelection` requires `onRowSelectionChange`; without it selection changes are dropped." + : null, ].filter((message): message is string => message !== null); } @@ -448,6 +454,9 @@ function useDataTableInstance(props: DataTablePro renderSubComponent, expanded, onExpandedChange, + enableRowSelection, + rowSelection, + onRowSelectionChange, } = props; const sortingState = useControllable(sorting, onSortingChange, defaultSorting ?? []); @@ -462,6 +471,7 @@ function useDataTableInstance(props: DataTablePro ); const globalFilterState = useControllable(globalFilter, onGlobalFilterChange, ""); const expandedState = useControllable(expanded, onExpandedChange, {}); + const rowSelectionState = useControllable(rowSelection, onRowSelectionChange, {}); const [columnVisibility, setColumnVisibility] = useState(defaultColumnVisibility ?? {}); const [columnSizing, setColumnSizing] = useState({}); const columnPinning = React.useMemo(() => derivePinning(columns), [columns]); @@ -476,6 +486,7 @@ function useDataTableInstance(props: DataTablePro columnFilters: filterState.value, globalFilter: globalFilterState.value, expanded: expandedState.value, + rowSelection: rowSelectionState.value, columnVisibility, columnSizing, }, @@ -491,11 +502,13 @@ function useDataTableInstance(props: DataTablePro onColumnFiltersChange: filterState.onChange, onGlobalFilterChange: globalFilterState.onChange, onExpandedChange: expandedState.onChange, + onRowSelectionChange: rowSelectionState.onChange, onColumnVisibilityChange: setColumnVisibility, onColumnSizingChange: setColumnSizing, getCoreRowModel: getCoreRowModel(), ...buildRowModels(sortingMode, paginationMode, filterMode, expansionGuard), ...(getRowId !== undefined ? { getRowId } : {}), + ...(enableRowSelection !== undefined ? { enableRowSelection } : {}), ...(paginationMode === "server" && rowCount !== undefined ? { rowCount } : {}), }; diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableRowSelection.test.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableRowSelection.test.tsx new file mode 100644 index 00000000000..2464fb93309 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableRowSelection.test.tsx @@ -0,0 +1,137 @@ +import type { ColumnDef, RowSelectionState } from "@tanstack/react-table"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { useState } from "react"; +import { describe, expect, it } from "vitest"; + +import { createSelectionColumn, DataTable, validateDataTableConfig } from "./index"; + +interface Model { + id: string; + name: string; +} + +const data: Model[] = [ + { id: "m1", name: "Alpha" }, + { id: "m2", name: "Beta" }, + { id: "m3", name: "Gamma" }, +]; + +const columns: ColumnDef[] = [ + createSelectionColumn({ rowAriaLabel: (row) => `Select ${row.original.name}` }), + { id: "name", accessorKey: "name", header: "Name", enableSorting: false }, +]; + +const selectAll = () => screen.getByTestId("datatable-select-all"); +const rowBox = (id: string) => screen.getByTestId(`datatable-select-row-${id}`); +const selectedCount = () => screen.getByTestId("count"); + +function ControlledHarness() { + const [rowSelection, setRowSelection] = useState({}); + + return ( + <> + + {Object.keys(rowSelection) + .filter((key) => rowSelection[key]) + .sort() + .join(",")} + + + row.id} + rowSelection={rowSelection} + onRowSelectionChange={setRowSelection} + /> + + ); +} + +describe("DataTable row selection", () => { + it("supports uncontrolled per-row toggle, select-all, and indeterminate", async () => { + const user = userEvent.setup(); + + render( + row.id} + toolbar={(table) => {table.getSelectedRowModel().rows.length}} + />, + ); + + expect(selectedCount()).toHaveTextContent("0"); + + await user.click(rowBox("m1")); + expect(selectedCount()).toHaveTextContent("1"); + expect(selectAll()).toHaveAttribute("aria-checked", "mixed"); + + await user.click(selectAll()); + expect(selectedCount()).toHaveTextContent("3"); + expect(selectAll()).toHaveAttribute("aria-checked", "true"); + + await user.click(selectAll()); + expect(selectedCount()).toHaveTextContent("0"); + }); + + it("keys controlled selection by getRowId so the parent can map back to entities", async () => { + const user = userEvent.setup(); + render(); + + await user.click(rowBox("m2")); + expect(screen.getByTestId("keys")).toHaveTextContent("m2"); + + await user.click(rowBox("m3")); + expect(screen.getByTestId("keys")).toHaveTextContent("m2,m3"); + }); + + it("lets the parent clear the selection, the pattern an external pager needs", async () => { + const user = userEvent.setup(); + render(); + + await user.click(selectAll()); + expect(screen.getByTestId("keys")).toHaveTextContent("m1,m2,m3"); + + await user.click(screen.getByTestId("clear")); + expect(screen.getByTestId("keys")).toBeEmptyDOMElement(); + expect(rowBox("m1")).toHaveAttribute("aria-checked", "false"); + }); + + it("respects an enableRowSelection predicate", async () => { + const user = userEvent.setup(); + + render( + row.id} + enableRowSelection={(row) => row.original.id !== "m2"} + toolbar={(table) => {table.getSelectedRowModel().rows.length}} + />, + ); + + expect(rowBox("m2")).toHaveAttribute("aria-disabled", "true"); + + await user.click(rowBox("m2")); + expect(selectedCount()).toHaveTextContent("0"); + + await user.click(rowBox("m1")); + expect(selectedCount()).toHaveTextContent("1"); + }); + + it("rejects controlled rowSelection without onRowSelectionChange", () => { + const errors = validateDataTableConfig({ data, columns, rowSelection: { m1: true } }); + + expect(errors).toContain( + "Controlled `rowSelection` requires `onRowSelectionChange`; without it selection changes are dropped.", + ); + }); + + it("does not complain when selection is left uncontrolled", () => { + expect(validateDataTableConfig({ data, columns })).toHaveLength(0); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableSelectionColumn.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableSelectionColumn.tsx new file mode 100644 index 00000000000..da32a01ab0e --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableSelectionColumn.tsx @@ -0,0 +1,53 @@ +"use client"; + +import type { ColumnDef, Row, RowData, Table } from "@tanstack/react-table"; + +import { Checkbox } from "@/components/ui/checkbox"; + +interface SelectionColumnOptions { + rowAriaLabel?: (row: Row) => string; +} + +function SelectAllCheckbox({ table }: { table: Table }) { + const allSelected = table.getIsAllPageRowsSelected(); + const someSelected = table.getIsSomePageRowsSelected(); + + return ( + table.toggleAllPageRowsSelected(Boolean(checked))} + /> + ); +} + +function SelectRowCheckbox({ row, label }: { row: Row; label: string }) { + return ( + row.toggleSelected(Boolean(checked))} + /> + ); +} + +export function createSelectionColumn( + options: SelectionColumnOptions = {}, +): ColumnDef { + const { rowAriaLabel } = options; + + return { + id: "select", + size: 44, + enableSorting: false, + enableHiding: false, + enableResizing: false, + meta: { title: "Select", className: "w-11", headerClassName: "w-11" }, + header: ({ table }) => , + cell: ({ row }) => , + }; +} diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/index.ts b/ui/litellm-dashboard/src/components/shared/DataTable/index.ts index 1ee1eed1258..62ddd1b0742 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/index.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/index.ts @@ -3,6 +3,7 @@ import "./columnMeta"; export { DataTable, DataTableConfigError, validateDataTableConfig } from "./DataTable"; export { DataTableFilterDrawer, DataTableFilterField, type FilterDraft } from "./DataTableFilterDrawer"; export { DataTablePagination, DEFAULT_PAGE_SIZE_OPTIONS } from "./DataTablePagination"; +export { createSelectionColumn } from "./DataTableSelectionColumn"; export { DataTableToolbar } from "./DataTableToolbar"; export { DataTableViewOptions } from "./DataTableViewOptions"; export { diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts index 672ab512ef4..40f3a4df204 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts @@ -6,6 +6,7 @@ import type { PaginationState, Row, RowData, + RowSelectionState, SortingState, Table, VisibilityState, @@ -59,6 +60,10 @@ export interface DataTableProps { expanded?: ExpandedState; onExpandedChange?: OnChangeFn; + enableRowSelection?: boolean | ((row: Row) => boolean); + rowSelection?: RowSelectionState; + onRowSelectionChange?: OnChangeFn; + onRowClick?: (row: TData) => void; rowClassName?: (row: Row) => string; diff --git a/ui/litellm-dashboard/src/components/shared/form/FormField.test.tsx b/ui/litellm-dashboard/src/components/shared/form/FormField.test.tsx new file mode 100644 index 00000000000..af9122e2bd5 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/form/FormField.test.tsx @@ -0,0 +1,180 @@ +import { zodResolver } from "@hookform/resolvers/zod"; +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import * as React from "react"; +import { useForm } from "react-hook-form"; +import { describe, expect, it, vi } from "vitest"; +import { z } from "zod/v4"; + +import { Input } from "@/components/ui/input"; + +import { FormField } from "./FormField"; + +const schema = z.object({ + team_alias: z.string().min(1, "Please input a team name"), + owner: z.string(), +}); + +type FormInput = z.input; + +const TestForm = ({ + onSubmit, + defaultValues = { team_alias: "team-a", owner: "" }, + description, +}: { + onSubmit: (values: z.output) => void; + defaultValues?: FormInput; + description?: React.ReactNode; +}) => { + const form = useForm>({ + resolver: zodResolver(schema), + defaultValues, + }); + + return ( +
+ + {(field) => } + + +
+ ); +}; + +describe("FormField", () => { + it("associates the label with the control so it is reachable by its accessible name", () => { + render(); + + expect(screen.getByLabelText("Team Name")).toHaveValue("team-a"); + }); + + it("gives each field instance a unique control id", () => { + const Harness = () => { + const form = useForm({ defaultValues: { team_alias: "", owner: "" } }); + return ( + <> + + {(field) => } + + + {(field) => } + + + ); + }; + render(); + + expect(screen.getByLabelText("One").id).not.toBe(screen.getByLabelText("Two").id); + }); + + it("feeds edits back into form state and submits the parsed output", async () => { + const user = userEvent.setup(); + const onSubmit = vi.fn(); + render(); + + await user.clear(screen.getByLabelText("Team Name")); + await user.type(screen.getByLabelText("Team Name"), "team-b"); + await user.click(screen.getByRole("button", { name: "Save" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalledTimes(1)); + expect(onSubmit.mock.calls[0][0]).toEqual({ team_alias: "team-b", owner: "" }); + }); + + it("renders the zod message and blocks submit when validation fails", async () => { + const user = userEvent.setup(); + const onSubmit = vi.fn(); + render(); + + await user.clear(screen.getByLabelText("Team Name")); + await user.click(screen.getByRole("button", { name: "Save" })); + + expect(await screen.findByRole("alert")).toHaveTextContent("Please input a team name"); + expect(onSubmit).not.toHaveBeenCalled(); + }); + + it("marks the control invalid and points aria-describedby at the message", async () => { + const user = userEvent.setup(); + render(); + + await user.clear(screen.getByLabelText("Team Name")); + await user.click(screen.getByRole("button", { name: "Save" })); + + const control = await screen.findByLabelText("Team Name"); + await waitFor(() => expect(control).toHaveAttribute("aria-invalid", "true")); + expect(control.getAttribute("aria-describedby")).toBe(screen.getByRole("alert").id); + }); + + it("leaves a valid control free of aria-invalid", () => { + render(); + + expect(screen.getByLabelText("Team Name")).not.toHaveAttribute("aria-invalid"); + }); + + it("clears the message once the value becomes valid again", async () => { + const user = userEvent.setup(); + render(); + + await user.clear(screen.getByLabelText("Team Name")); + await user.click(screen.getByRole("button", { name: "Save" })); + expect(await screen.findByRole("alert")).toBeInTheDocument(); + + await user.type(screen.getByLabelText("Team Name"), "team-c"); + + await waitFor(() => expect(screen.queryByRole("alert")).not.toBeInTheDocument()); + }); + + it("describes the control by its description when there is no error", () => { + render(); + + const control = screen.getByLabelText("Team Name"); + const describedBy = control.getAttribute("aria-describedby"); + + expect(describedBy).not.toBeNull(); + expect(document.getElementById(describedBy!)).toHaveTextContent("Shown to team members"); + }); + + it("describes the control by both description and error while invalid", async () => { + const user = userEvent.setup(); + render(); + + await user.clear(screen.getByLabelText("Team Name")); + await user.click(screen.getByRole("button", { name: "Save" })); + await screen.findByRole("alert"); + + const ids = screen.getByLabelText("Team Name").getAttribute("aria-describedby")?.split(" ") ?? []; + + expect(ids).toHaveLength(2); + expect(ids).toContain(screen.getByRole("alert").id); + }); + + it("omits aria-describedby entirely when there is no description and no error", () => { + render(); + + expect(screen.getByLabelText("Team Name")).not.toHaveAttribute("aria-describedby"); + }); + + it("hands the control a value and onChange so non-native widgets can be wired", async () => { + const user = userEvent.setup(); + const seen: unknown[] = []; + const Harness = () => { + const form = useForm({ defaultValues: { team_alias: "team-a", owner: "" } }); + return ( + + {(field) => { + seen.push(field.value); + return ( + + ); + }} + + ); + }; + render(); + + await user.click(screen.getByRole("button", { name: "widget" })); + + await waitFor(() => expect(seen.at(-1)).toBe("from-widget")); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/form/FormField.tsx b/ui/litellm-dashboard/src/components/shared/form/FormField.tsx new file mode 100644 index 00000000000..3b9783333cc --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/form/FormField.tsx @@ -0,0 +1,75 @@ +"use client"; + +import * as React from "react"; +import { + Controller, + type Control, + type ControllerRenderProps, + type FieldPath, + type FieldValues, +} from "react-hook-form"; + +import { Field, FieldDescription, FieldError, FieldLabel } from "./field"; + +export type FormFieldControlProps< + TFieldValues extends FieldValues, + TName extends FieldPath, +> = ControllerRenderProps & { + id: string; + "aria-invalid": true | undefined; + "aria-describedby": string | undefined; +}; + +export interface FormFieldProps> { + control: Control; + name: TName; + label?: React.ReactNode; + description?: React.ReactNode; + orientation?: "vertical" | "horizontal" | "responsive"; + className?: string; + children: (control: FormFieldControlProps) => React.ReactNode; +} + +export const FormField = >({ + control, + name, + label, + description, + orientation, + className, + children, +}: FormFieldProps) => { + const reactId = React.useId(); + const controlId = `${reactId}-control`; + const descriptionId = `${reactId}-description`; + const errorId = `${reactId}-error`; + + return ( + { + const invalid = fieldState.error !== undefined; + const describedBy = + [description !== undefined ? descriptionId : undefined, invalid ? errorId : undefined] + .filter((id): id is string => id !== undefined) + .join(" ") || undefined; + const controlProps: FormFieldControlProps = { + ...field, + id: controlId, + "aria-invalid": invalid || undefined, + "aria-describedby": describedBy, + }; + + return ( + + {label !== undefined && {label}} + {children(controlProps)} + {description !== undefined && {description}} + + + ); + }} + /> + ); +}; diff --git a/ui/litellm-dashboard/src/components/shared/form/field.test.tsx b/ui/litellm-dashboard/src/components/shared/form/field.test.tsx new file mode 100644 index 00000000000..54b589ce2f4 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/form/field.test.tsx @@ -0,0 +1,125 @@ +import { render, screen } from "@testing-library/react"; +import * as React from "react"; +import { describe, expect, it } from "vitest"; + +import { + Field, + FieldContent, + FieldDescription, + FieldError, + FieldGroup, + FieldLabel, + FieldLegend, + FieldSeparator, + FieldSet, + FieldTitle, +} from "./field"; + +describe("FieldError", () => { + it("renders nothing when there are no errors and no children", () => { + const { container } = render(); + + expect(container).toBeEmptyDOMElement(); + }); + + it("renders nothing when every error entry is undefined", () => { + const { container } = render(); + + expect(container).toBeEmptyDOMElement(); + }); + + it("renders a single message as plain text, not a list", () => { + render(); + + expect(screen.getByRole("alert")).toHaveTextContent("Required"); + expect(screen.queryByRole("listitem")).not.toBeInTheDocument(); + }); + + it("collapses duplicate messages to a single entry", () => { + render(); + + expect(screen.getByRole("alert")).toHaveTextContent("Required"); + expect(screen.queryByRole("listitem")).not.toBeInTheDocument(); + }); + + it("renders distinct messages as a list", () => { + render(); + + const items = screen.getAllByRole("listitem"); + expect(items.map((item) => item.textContent)).toEqual(["Too short", "Must be lowercase"]); + }); + + it("prefers explicit children over the errors prop", () => { + render(from children); + + expect(screen.getByRole("alert")).toHaveTextContent("from children"); + expect(screen.getByRole("alert")).not.toHaveTextContent("from errors"); + }); + + it("exposes the message to assistive tech via role=alert", () => { + render(); + + expect(screen.getByRole("alert")).toBeInTheDocument(); + }); +}); + +describe("Field", () => { + it("marks itself invalid so descendants can style off it", () => { + render( + + child + , + ); + + expect(screen.getByRole("group")).toHaveAttribute("data-invalid", "true"); + }); + + it("defaults to vertical orientation", () => { + render(); + + expect(screen.getByRole("group")).toHaveAttribute("data-orientation", "vertical"); + }); + + it("honours an explicit orientation", () => { + render(); + + expect(screen.getByRole("group")).toHaveAttribute("data-orientation", "horizontal"); + }); +}); + +describe("field primitives forward refs to their DOM node", () => { + it.each([ + ["Field", Field, HTMLDivElement], + ["FieldContent", FieldContent, HTMLDivElement], + ["FieldDescription", FieldDescription, HTMLParagraphElement], + ["FieldGroup", FieldGroup, HTMLDivElement], + ["FieldLabel", FieldLabel, HTMLLabelElement], + ["FieldSeparator", FieldSeparator, HTMLDivElement], + ["FieldTitle", FieldTitle, HTMLDivElement], + ])("%s", (_name, Component, expected) => { + const ref = React.createRef(); + render(React.createElement(Component as React.ElementType, { ref })); + + expect(ref.current).toBeInstanceOf(expected); + }); + + it("FieldSet and FieldLegend", () => { + const fieldSet = React.createRef(); + const legend = React.createRef(); + render( +
+ Legend +
, + ); + + expect(fieldSet.current).toBeInstanceOf(HTMLFieldSetElement); + expect(legend.current).toBeInstanceOf(HTMLLegendElement); + }); + + it("FieldError", () => { + const ref = React.createRef(); + render(); + + expect(ref.current).toBeInstanceOf(HTMLDivElement); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/form/field.tsx b/ui/litellm-dashboard/src/components/shared/form/field.tsx new file mode 100644 index 00000000000..36ce691827c --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/form/field.tsx @@ -0,0 +1,223 @@ +"use client"; + +import * as React from "react"; +import { type VariantProps } from "cva"; + +import { Label } from "@/components/ui/label"; +import { Separator } from "@/components/ui/separator"; +import { cn, cva } from "@/lib/cva.config"; + +const FieldSet = React.forwardRef>( + ({ className, ...props }, ref) => ( +
[data-slot=checkbox-group]]:gap-3 has-[>[data-slot=radio-group]]:gap-3", + className, + )} + {...props} + /> + ), +); +FieldSet.displayName = "FieldSet"; + +const FieldLegend = React.forwardRef< + HTMLLegendElement, + React.ComponentPropsWithoutRef<"legend"> & { variant?: "legend" | "label" } +>(({ className, variant = "legend", ...props }, ref) => ( + +)); +FieldLegend.displayName = "FieldLegend"; + +const FieldGroup = React.forwardRef>( + ({ className, ...props }, ref) => ( +
+ ), +); +FieldGroup.displayName = "FieldGroup"; + +const fieldVariants = cva({ + base: "group/field flex w-full gap-3 data-[invalid=true]:text-destructive", + variants: { + orientation: { + vertical: "flex-col *:w-full [&>.sr-only]:w-auto", + horizontal: + "flex-row items-center has-[>[data-slot=field-content]]:items-start *:data-[slot=field-label]:flex-auto has-[>[data-slot=field-content]]:[&>[role=checkbox],[role=radio]]:mt-px", + responsive: + "flex-col *:w-full @md/field-group:flex-row @md/field-group:items-center @md/field-group:*:w-auto @md/field-group:has-[>[data-slot=field-content]]:items-start @md/field-group:*:data-[slot=field-label]:flex-auto [&>.sr-only]:w-auto @md/field-group:has-[>[data-slot=field-content]]:[&>[role=checkbox],[role=radio]]:mt-px", + }, + }, + defaultVariants: { + orientation: "vertical", + }, +}); + +const Field = React.forwardRef< + HTMLDivElement, + React.ComponentPropsWithoutRef<"div"> & VariantProps +>(({ className, orientation = "vertical", ...props }, ref) => ( +
+)); +Field.displayName = "Field"; + +const FieldContent = React.forwardRef>( + ({ className, ...props }, ref) => ( +
+ ), +); +FieldContent.displayName = "FieldContent"; + +const FieldLabel = React.forwardRef>( + ({ className, ...props }, ref) => ( +