From cfbcef319c61ad29fbb2caf9a16ac5419006bb2b Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 21 Jul 2026 22:24:14 -0700 Subject: [PATCH 01/25] fix(mcp): log actionable OAuth discovery failures for misconfigured server urls A typo'd MCP server url failed OAuth endpoint discovery silently: every failure died at debug level, the config loader warned nothing, and the /authorize 400 blamed "servers with no url" even when a url was set. _descovery_metadata now records each attempt's outcome and, when a total failure would leave the server's flow without a needed endpoint, logs one warning with the trail (urls origin-only, exception text url-stripped). Both server loaders warn which endpoints stayed unresolved for the server's flow (client_credentials never needs authorization_url, OBO needs only token_url) with the remedies; this replaces the DB path's reason-less warning and closes the config path's no-warning gap. The authorize/token/register 400 details branch on server shape via one shared helper and point at the proxy logs. _redact_mcp_resource_url moves to oauth_utils.py so the manager can import it without a cycle. Resolves LIT-4658 --- .../mcp_server/discoverable_endpoints.py | 56 ++- .../mcp_server/mcp_server_manager.py | 355 ++++++++++++++---- .../_experimental/mcp_server/oauth_utils.py | 23 +- .../proxy/_experimental/mcp_server/server.py | 25 +- .../mcp_server/test_discoverable_endpoints.py | 114 ++++++ .../mcp_server/test_mcp_server_manager.py | 175 ++++++++- 6 files changed, 623 insertions(+), 125 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 9a1b5cf4864..bc0822a71d4 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -499,6 +499,35 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None: ) +def _endpoint_not_configured_detail( + mcp_server: MCPServer, + endpoint_label: str, + manual_remedy: str, + issuer_remedy: str, +) -> str: + """The 400 detail for an unresolved OAuth endpoint, naming the likely cause for this server's + shape (LIT-4658): an anchored issuer whose metadata fell short, a configured (possibly + misconfigured) server url whose discovery failed, or no discovery source at all. Kept free of + URLs and issuer values because these endpoints are reachable pre-auth.""" + if mcp_server.issuer_is_anchored: + return ( + f"MCP server {endpoint_label} is not configured. Endpoint discovery anchored on the configured " + f"Issuer (RFC 8414) failed or its metadata did not include this endpoint; check the proxy logs " + f"for 'MCP OAuth' warnings from server load, verify the Issuer, or {manual_remedy}." + ) + if mcp_server.url: + return ( + f"MCP server {endpoint_label} is not configured. OAuth endpoint discovery against the configured " + f"server url did not resolve it; the url may be misconfigured. Check the proxy logs for " + f"'MCP OAuth' warnings from server load, verify the server url, or {manual_remedy}, or " + f"{issuer_remedy}." + ) + return ( + f"MCP server {endpoint_label} is not configured. Servers with no url (OpenAPI spec or stdio) run no " + f"resource discovery, so {manual_remedy}, or {issuer_remedy}." + ) + + def _raise_unless_oauth2_discovery_server( mcp_server: Optional[MCPServer], mcp_server_name: Optional[str], @@ -599,10 +628,11 @@ async def authorize_with_server( if mcp_server.authorization_url is None: 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)." + detail=_endpoint_not_configured_detail( + mcp_server, + "authorization url", + "set Authorization URL and Token URL manually", + "set Issuer to discover them from the identity provider (RFC 8414)", ), ) @@ -711,10 +741,11 @@ async def exchange_token_with_server( if mcp_server.token_url is None: 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)." + detail=_endpoint_not_configured_detail( + mcp_server, + "token url", + "set Token URL manually", + "set Issuer to discover it from the identity provider (RFC 8414)", ), ) @@ -1278,10 +1309,11 @@ async def register_client_with_server( if mcp_server.authorization_url is None: 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)." + detail=_endpoint_not_configured_detail( + mcp_server, + "authorization url", + "set Authorization URL and Token URL manually", + "set Issuer to discover them from the identity provider (RFC 8414)", ), ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 90b70dd01f2..4faf6da340f 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -13,6 +13,7 @@ import json import os import re import time +from collections.abc import Sequence from contextlib import asynccontextmanager from typing import Any, AsyncIterator, Callable, Literal, Optional, Union, cast from urllib.parse import urlparse @@ -50,6 +51,9 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) +from litellm.proxy._experimental.mcp_server.elicitation_handler import ( + MCP_ELICITATION_AVAILABLE, +) from litellm.proxy._experimental.mcp_server.exceptions import ( MCPServerListError, MCPUpstreamAuthError, @@ -59,17 +63,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( raise_classified_list_failure, upstream_auth_challenge, ) -from litellm.proxy._experimental.mcp_server.elicitation_handler import ( - MCP_ELICITATION_AVAILABLE, -) -from litellm.proxy._experimental.mcp_server.sampling_handler import ( - MCP_SAMPLING_AVAILABLE, -) from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( MCPPerUserTokenCache, mcp_per_user_token_cache, resolve_mcp_auth, ) +from litellm.proxy._experimental.mcp_server.oauth_utils import ( + _redact_mcp_resource_url, +) from litellm.proxy._experimental.mcp_server.outbound_credentials import ( Error, Ok, @@ -100,6 +101,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( ServerSpec, TokenExchangeConfig, ) +from litellm.proxy._experimental.mcp_server.sampling_handler import ( + MCP_SAMPLING_AVAILABLE, +) from litellm.proxy._experimental.mcp_server.utils import ( MCP_TOOL_PREFIX_SEPARATOR, MCPMissingUserEnvVarsError, @@ -143,11 +147,9 @@ from litellm.types.mcp_server.mcp_server_manager import ( from litellm.types.utils import CallTypes try: - from mcp.shared.tool_name_validation import ( - validate_tool_name, # pyright: ignore[reportAssignmentType] - ) from mcp.shared.tool_name_validation import ( SEP_986_URL, + validate_tool_name, # pyright: ignore[reportAssignmentType] ) except ImportError: from pydantic import BaseModel @@ -408,6 +410,88 @@ def _restrict_discovery_to_corroborated_authorization_server( return metadata.model_copy(update={"token_url": None, "registration_url": None}) +def _redacted_origin_list(urls: Sequence[str]) -> str: + return ", ".join(_redact_mcp_resource_url(url) or "" for url in urls) + + +def _sanitized_error_text(exc: Exception) -> str: + return re.sub(r"https?://\S+", "", str(exc))[:200] + + +def _discovery_failure_leaves_needs_unresolved( + *, + needs_authorization_url: bool, + needs_token_url: bool, + manual_authorization_url: str | None, + manual_token_url: str | None, +) -> bool: + return (needs_authorization_url and not manual_authorization_url) or (needs_token_url and not manual_token_url) + + +def _warn_oauth_endpoints_unresolved( + *, + server_ref: str, + server_url: str | None, + discovery_attempted: bool, + issuer_anchored: bool, + metadata: MCPOAuthMetadata | None, + needs_authorization_url: bool, + needs_token_url: bool, + manual_authorization_url: str | None, + manual_token_url: str | None, +) -> None: + """Log one actionable warning when a server that depends on OAuth endpoint discovery finishes a + build without the endpoints that its flows need (LIT-4658). + + This is the operator-facing signal for a misconfigured server url: discovery failures themselves + are logged where they happen (``_descovery_metadata``), and this names WHICH server is affected, + which endpoints stayed unresolved after manual configuration was considered, and the remedies. + Scopes never trigger the warning on their own: scope-less metadata is normal for many servers and + warning on it every rebuild would be noise. Callers own the per-flow policy of which endpoints + are needed (client_credentials never needs authorization_url; OBO needs only token_url); the + issuer-anchored arm is excluded here because it has its own RFC 8414 §3.3 warning. + """ + if issuer_anchored: + return + unresolved = tuple( + field + for field, needed, value in ( + ( + "authorization_url", + needs_authorization_url, + manual_authorization_url or (metadata.authorization_url if metadata else None), + ), + ( + "token_url", + needs_token_url, + manual_token_url or (metadata.token_url if metadata else None), + ), + ) + if needed and not value + ) + if not unresolved: + return + if discovery_attempted: + verbose_logger.warning( + "MCP server %s: OAuth endpoint discovery left %s unresolved (server url origin: %s). OAuth flows " + "that need them will fail with 'not configured' errors until they resolve. Check the preceding " + "'MCP OAuth' log lines for why discovery failed, verify the configured server url, or set the " + "unresolved endpoint urls manually, or set issuer to discover them from the identity provider " + "(RFC 8414)", + server_ref, + ", ".join(unresolved), + _redact_mcp_resource_url(server_url) or "", + ) + return + verbose_logger.warning( + "MCP server %s uses OAuth but has no discovery source (no server url or pinned issuer), and %s not " + "set manually. Set the missing endpoint urls on the server, or set issuer to discover them from the " + "identity provider (RFC 8414)", + server_ref, + " and ".join(unresolved) + (" is" if len(unresolved) == 1 else " are"), + ) + + def invalidate_user_env_vars_cache(user_id: str, server_id: str) -> None: """Drop a cached entry after the user stores or clears their env var values so the next request reads the fresh value instead of a stale one.""" @@ -884,10 +968,10 @@ def _create_sampling_callback(user_api_key_auth: Optional[Any] = None): return None async def _sampling_callback(context, params): + import litellm from litellm.proxy._experimental.mcp_server.sampling_handler import ( handle_sampling_create_message, ) - import litellm from litellm.proxy._experimental.mcp_server.server import ( get_active_auth_context, ) @@ -1284,6 +1368,15 @@ class MCPServerManager: should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and ( is_discovery_auth_type or obo_needs_discovery ) + config_oauth2_flow = server_config.get("oauth2_flow", None) + needs_authorization_url = is_discovery_auth_type and config_oauth2_flow != "client_credentials" + needs_token_url = is_discovery_auth_type or obo_needs_discovery + warn_on_empty_discovery = _discovery_failure_leaves_needs_unresolved( + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) if not should_discover: mcp_oauth_metadata = None elif use_issuer_anchor and manual_issuer is not None: @@ -1292,6 +1385,7 @@ class MCPServerManager: mcp_oauth_metadata = await self._descovery_metadata( server_url=server_url, allow_origin_fallback=is_discovery_auth_type, + warn_when_no_metadata=warn_on_empty_discovery, ) if use_issuer_anchor: @@ -1326,7 +1420,6 @@ class MCPServerManager: ) effective_issuer = manual_issuer or discovered_issuer - config_oauth2_flow = server_config.get("oauth2_flow", None) if auth_type == MCPAuth.oauth2 and config_oauth2_flow not in ( "client_credentials", "authorization_code", @@ -1358,6 +1451,18 @@ class MCPServerManager: "authorization-code flow." ) + _warn_oauth_endpoints_unresolved( + server_ref=server_name or server_id, + server_url=server_url, + discovery_attempted=should_discover, + issuer_anchored=use_issuer_anchor, + metadata=gated_oauth_metadata, + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) + new_server = MCPServer( server_id=server_id, name=name_for_prefix, @@ -1485,14 +1590,12 @@ class MCPServerManager: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( build_input_schema, create_tool_function, + load_openapi_spec_async, + resolve_operation_params, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( get_base_url as get_openapi_base_url, ) - from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - load_openapi_spec_async, - resolve_operation_params, - ) from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) @@ -1681,10 +1784,20 @@ class MCPServerManager: scopes: Optional[list[str]], token_exchange_endpoint: Optional[str], ) -> Optional[MCPOAuthMetadata]: + obo_needs_discovery = self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url) + needs_authorization_url = ( + is_discovery_auth_type and getattr(mcp_server, "oauth2_flow", None) != "client_credentials" + ) + needs_token_url = is_discovery_auth_type or obo_needs_discovery + warn_on_empty_discovery = _discovery_failure_leaves_needs_unresolved( + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes) 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) + (is_discovery_auth_type and not has_all_upstream_oauth_fields) or obo_needs_discovery ) if not needs_discovery: mcp_oauth_metadata: Optional[MCPOAuthMetadata] = None @@ -1694,24 +1807,32 @@ class MCPServerManager: mcp_oauth_metadata = await self._descovery_metadata( server_url=server_url, # type: ignore[arg-type] allow_origin_fallback=is_discovery_auth_type, - ) - if needs_discovery and not use_issuer_anchor and mcp_oauth_metadata is None: - verbose_logger.warning( - "MCP OAuth discovery yielded no metadata for server %s (%s); " - "OAuth endpoints/scopes stay unresolved until a rebuild succeeds", - mcp_server.server_id, - server_url, + warn_when_no_metadata=warn_on_empty_discovery, ) if use_issuer_anchor: return mcp_oauth_metadata - if is_discovery_auth_type: - return _restrict_discovery_to_corroborated_authorization_server( + gated_metadata = ( + _restrict_discovery_to_corroborated_authorization_server( mcp_oauth_metadata, manual_authorization_url, mcp_server.server_id, bool(getattr(mcp_server, "dcr_bridge", None)), ) - return mcp_oauth_metadata + if is_discovery_auth_type + else mcp_oauth_metadata + ) + _warn_oauth_endpoints_unresolved( + server_ref=mcp_server.alias or mcp_server.server_name or mcp_server.server_id, + server_url=server_url, + discovery_attempted=needs_discovery, + issuer_anchored=False, + metadata=gated_metadata, + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) + return gated_metadata async def build_mcp_server_from_table( self, @@ -3430,6 +3551,7 @@ class MCPServerManager: server_url: str, *, allow_origin_fallback: bool = True, + warn_when_no_metadata: bool = False, ) -> Optional[MCPOAuthMetadata]: """Discover OAuth metadata by following RFC 9728 (protected resource metadata discovery). @@ -3438,8 +3560,32 @@ class MCPServerManager: it (a human sees the redirect), but token_exchange (OBO) sets it False so the gateway never exchanges a subject token against an endpoint it inferred rather than one explicitly configured or authoritatively advertised via RFC 9728 / RFC 8414. - """ + ``warn_when_no_metadata`` makes an all-empty result log one WARNING with the per-step attempt + outcomes (LIT-4658), so a misconfigured server url is diagnosable from default-level logs. The + server loaders set it; the issuer-anchored resource-scopes lookup keeps it off because empty + scopes are not a fault there. + """ + metadata, attempts = await self._discover_metadata_recording_attempts( + server_url, allow_origin_fallback=allow_origin_fallback + ) + if metadata is None and warn_when_no_metadata: + verbose_logger.warning( + "MCP OAuth endpoint discovery against %s found no authorization server metadata. Attempts: %s. " + "The MCP server url may be misconfigured, or the upstream may not support OAuth discovery " + "(RFC 9728 / RFC 8414)", + _redact_mcp_resource_url(server_url) or "", + "; ".join(attempts) if attempts else "none recorded", + ) + return metadata + + async def _discover_metadata_recording_attempts( + self, + server_url: str, + *, + allow_origin_fallback: bool, + ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: + origin = _redact_mcp_resource_url(server_url) or "" try: client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) response = await client.get(server_url) @@ -3452,67 +3598,112 @@ class MCPServerManager: if metadata is None and not resource_scopes and authorization_servers and response.status_code == 200: verbose_logger.warning( "MCP OAuth discovery for %s received 200 OK without RFC 9728 challenge and no discoverable authorization metadata.", - server_url, + origin, ) + attempts = ( + f"GET {origin}: HTTP {response.status_code} (no RFC 9728 challenge)", + *( + ("well-known protected-resource lookup found no authorization servers",) + if not authorization_servers + else () + ), + *( + (f"authorization server metadata fetch failed for: {_redacted_origin_list(authorization_servers)}",) + if authorization_servers and metadata is None + else () + ), + ) if metadata is None and resource_scopes: - return MCPOAuthMetadata(scopes=resource_scopes) + return MCPOAuthMetadata(scopes=resource_scopes), attempts if metadata is not None and resource_scopes: metadata.scopes = resource_scopes - return metadata + return metadata, attempts except HTTPStatusError as exc: - verbose_logger.debug( - "MCP OAuth discovery for %s received status error: %s", - server_url, - exc, - ) - - header_value: Optional[str] = None - if exc.response is not None: - header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get( - "www-authenticate" - ) - - resource_metadata_url, scopes = self._parse_www_authenticate_header(header_value) - - authorization_servers = [] - resource_scopes = None - if resource_metadata_url: - ( - authorization_servers, - resource_scopes, - ) = await self._fetch_oauth_metadata_from_resource(resource_metadata_url, server_url) - else: - ( - authorization_servers, - resource_scopes, - ) = await self._attempt_well_known_discovery(server_url) - - metadata = None - used_origin_fallback = False - if allow_origin_fallback and not authorization_servers: - try: - parsed_url = urlparse(server_url) - if parsed_url.scheme and parsed_url.netloc: - authorization_servers = [f"{parsed_url.scheme}://{parsed_url.netloc}"] - used_origin_fallback = True - except Exception: - authorization_servers = [] - - if authorization_servers: - metadata = await self._fetch_authorization_server_metadata(authorization_servers, server_url) - if metadata is not None and used_origin_fallback: - metadata.from_origin_fallback = True - - preferred_scopes = scopes or resource_scopes - if metadata is None and preferred_scopes: - metadata = MCPOAuthMetadata(scopes=preferred_scopes) - elif metadata is not None and preferred_scopes: - metadata.scopes = preferred_scopes - - return metadata + return await self._discover_after_status_error(server_url, exc, allow_origin_fallback=allow_origin_fallback) except Exception as exc: # pragma: no cover - network/transient issues verbose_logger.debug("MCP OAuth discovery failed for %s: %s", server_url, exc) - return None + return None, (f"GET {origin}: {type(exc).__name__}: {_sanitized_error_text(exc)}",) + + async def _discover_after_status_error( + self, + server_url: str, + exc: HTTPStatusError, + *, + allow_origin_fallback: bool, + ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: + origin = _redact_mcp_resource_url(server_url) or "" + verbose_logger.debug( + "MCP OAuth discovery for %s received status error: %s", + server_url, + exc, + ) + + header_value: Optional[str] = None + if exc.response is not None: + header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get("www-authenticate") + status_attempt = ( + f"GET {origin}: HTTP {exc.response.status_code}" + if exc.response is not None + else f"GET {origin}: status error" + ) + + resource_metadata_url, scopes = self._parse_www_authenticate_header(header_value) + + authorization_servers = [] + resource_scopes = None + if resource_metadata_url: + ( + authorization_servers, + resource_scopes, + ) = await self._fetch_oauth_metadata_from_resource(resource_metadata_url, server_url) + lookup_attempt = ( + None + if authorization_servers + else "challenge-advertised resource metadata yielded no authorization servers" + ) + else: + ( + authorization_servers, + resource_scopes, + ) = await self._attempt_well_known_discovery(server_url) + lookup_attempt = ( + None + if authorization_servers + else "no challenge-advertised resource metadata; well-known protected-resource lookup found no authorization servers" + ) + + metadata = None + used_origin_fallback = False + if allow_origin_fallback and not authorization_servers: + try: + parsed_url = urlparse(server_url) + if parsed_url.scheme and parsed_url.netloc: + authorization_servers = [f"{parsed_url.scheme}://{parsed_url.netloc}"] + used_origin_fallback = True + except Exception: + authorization_servers = [] + + fallback_attempt = None + if authorization_servers: + metadata = await self._fetch_authorization_server_metadata(authorization_servers, server_url) + if metadata is not None and used_origin_fallback: + metadata.from_origin_fallback = True + if metadata is None: + fallback_attempt = ( + f"origin fallback: no authorization server metadata at {origin}" + if used_origin_fallback + else f"authorization server metadata fetch failed for: {_redacted_origin_list(authorization_servers)}" + ) + + attempts = tuple(entry for entry in (status_attempt, lookup_attempt, fallback_attempt) if entry) + + preferred_scopes = scopes or resource_scopes + if metadata is None and preferred_scopes: + return MCPOAuthMetadata(scopes=preferred_scopes), attempts + if metadata is not None and preferred_scopes: + metadata.scopes = preferred_scopes + + return metadata, attempts def _parse_www_authenticate_header(self, header_value: Optional[str]) -> tuple[Optional[str], Optional[list[str]]]: if not header_value: diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 53686e329bb..2f92f75a352 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -4,7 +4,7 @@ import os from ipaddress import ip_address from typing import Any, Dict, List, NoReturn, Optional -from urllib.parse import ParseResult, urlparse, urlunparse +from urllib.parse import ParseResult, urlparse, urlsplit, urlunparse, urlunsplit from fastapi import HTTPException, Request @@ -70,6 +70,27 @@ def _origin_label(scheme: str, netloc: str) -> str: return f"{scheme}://{netloc}" if netloc else f"{scheme}://" +def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]: + """Reduce an MCP server URL to its origin (scheme + host + port) for logging. + + Everything else is dropped: userinfo (``user:pass@``), the query string, the + fragment, and the path, because hosted MCP servers routinely embed the + credential in the path (e.g. ``/mcp/s/``) and this value is persisted + in spend-log metadata that a caller who can invoke the tool can read back. + Returns None when the URL has no host to identify (nothing safe to log). + """ + if not isinstance(url, str) or not url: + return None + try: + parts = urlsplit(url) + except ValueError: + return None + if not parts.hostname: + return None + netloc = f"{parts.hostname}:{parts.port}" if parts.port else parts.hostname + return urlunsplit((parts.scheme, netloc, "", "", "")) or None + + def _resolve_proxy_base_url_env() -> Optional[str]: global _warned_invalid_proxy_base_url configured = os.environ.get("PROXY_BASE_URL", "").strip() diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 396dd6c7dc7..56e8ec25076 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -27,7 +27,6 @@ from typing import ( Union, cast, ) -from urllib.parse import urlsplit, urlunsplit import httpx from fastapi import FastAPI, HTTPException @@ -59,6 +58,9 @@ from litellm.proxy._experimental.mcp_server.mcp_context import ( _mcp_gateway_server_name, ) from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug +from litellm.proxy._experimental.mcp_server.oauth_utils import ( + _redact_mcp_resource_url, +) from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_DESCRIPTION, LITELLM_MCP_SERVER_NAME, @@ -106,27 +108,6 @@ _MAX_STATEFUL_SESSIONS_PER_OWNER = 100 _MCP_ROUTING_PEEK_MAX_BYTES = 4096 -def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]: - """Reduce an MCP server URL to its origin (scheme + host + port) for logging. - - Everything else is dropped: userinfo (``user:pass@``), the query string, the - fragment, and the path, because hosted MCP servers routinely embed the - credential in the path (e.g. ``/mcp/s/``) and this value is persisted - in spend-log metadata that a caller who can invoke the tool can read back. - Returns None when the URL has no host to identify (nothing safe to log). - """ - if not isinstance(url, str) or not url: - return None - try: - parts = urlsplit(url) - except ValueError: - return None - if not parts.hostname: - return None - netloc = f"{parts.hostname}:{parts.port}" if parts.port else parts.hostname - return urlunsplit((parts.scheme, netloc, "", "", "")) or None - - def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: """Remove a (user_id, server_id) entry from the BYOK credential cache. 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 636c7fbd3d5..eeeea82b647 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 @@ -8148,3 +8148,117 @@ async def test_register_wall_names_the_fix_for_urlless_servers(): 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_authorize_wall_points_at_discovery_failure_for_url_servers(): + """LIT-4658: a server WITH a url that still has no authorization_url got here because OAuth + discovery against that url failed (typically a misconfigured url); the old detail blamed + "servers with no url", sending the operator down the wrong path. The detail must now name the + discovery failure and point at the proxy logs where LIT-4658's warnings carry the reason.""" + 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="typo-url-wall", + name="typo_wall", + server_name="typo_wall", + url="https://typo-host.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + 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 "may be misconfigured" in detail_text + assert "proxy logs" in detail_text + assert "Servers with no url" not in detail_text + assert "typo-host.example.com" not in detail_text + + +@pytest.mark.asyncio +async def test_token_wall_points_at_discovery_failure_for_url_servers(): + 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="typo-url-token-wall", + name="typo_token_wall", + server_name="typo_token_wall", + url="https://typo-host.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + authorization_url="https://idp.example.com/authorize", + ) + 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 "token url is not configured" in detail_text + assert "may be misconfigured" in detail_text + assert "Servers with no url" not in detail_text + + +@pytest.mark.asyncio +async def test_authorize_wall_names_the_issuer_for_anchored_servers(): + 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="anchored-wall", + name="anchored_wall", + server_name="anchored_wall", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + issuer_is_anchored=True, + ) + 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 "verify the Issuer" in detail_text + assert "Servers with no url" not in detail_text + assert "idp.example.com" not in detail_text 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 a5cb16822cf..77f072b81d5 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 @@ -3031,7 +3031,7 @@ class TestMCPServerManager: registration_url="https://discovered.example.com/register", ) - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): assert server_url == "https://example.com/mcp" # oauth2 (browser flow) keeps the origin fallback; only OBO disables it. assert allow_origin_fallback is True @@ -5426,7 +5426,7 @@ class TestMCPServerTimestamps: manager = MCPServerManager() calls: list[bool] = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): calls.append(allow_origin_fallback) return MCPOAuthMetadata( scopes=None, @@ -5461,7 +5461,7 @@ class TestMCPServerTimestamps: manager = MCPServerManager() calls: list[str] = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): calls.append(server_url) raise AssertionError("discovery must not run when token_exchange_endpoint is configured") @@ -5491,7 +5491,7 @@ class TestMCPServerTimestamps: back to the row, so the next rebuild skips discovery instead of re-running it every time.""" manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): assert server_url == "https://example.com/mcp" assert allow_origin_fallback is False # OBO never guesses the origin return MCPOAuthMetadata( @@ -5602,7 +5602,7 @@ class TestMCPServerTimestamps: _dcr_bridge_relays_client_registration keys off that column.""" manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): assert allow_origin_fallback is True return MCPOAuthMetadata( scopes=["mcp.read", "mcp.write"], @@ -5817,7 +5817,7 @@ class TestMCPServerTimestamps: persist_discovered_endpoints=False neither the oauth2 nor the OBO write-back may fire.""" manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): return MCPOAuthMetadata( scopes=["s1"], authorization_url="https://idp.example.com/authorize", @@ -8539,7 +8539,7 @@ class TestOBOEndpointDiscovery: ) seen = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): seen.append((server_url, allow_origin_fallback)) return discovered @@ -8567,7 +8567,7 @@ class TestOBOEndpointDiscovery: async def test_config_obo_with_configured_endpoint_skips_discovery(self): manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): raise AssertionError("discovery must not run when the endpoint is configured") manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] @@ -9028,3 +9028,162 @@ class TestUrllessIssuerDiscovery: anchored.assert_awaited_once_with("https://idp.example.com", None) resource_rooted.assert_not_awaited() assert built.token_url == "https://idp.example.com/token" + + +class TestDiscoveryFailureLogging: + """LIT-4658: a misconfigured MCP server url must be diagnosable from default-level server logs. + + Discovery failures used to die at debug level and the config-load path emitted no warning at + all, so the only operator-facing signal was the bare 400 at /authorize.""" + + def _connect_error_client(self, url: str) -> MagicMock: + client = MagicMock() + client.get = AsyncMock( + side_effect=httpx.ConnectError(f"[Errno 8] nodename nor servname provided for {url}") + ) + return client + + @pytest.mark.asyncio + async def test_descovery_metadata_warns_with_redacted_attempts_on_connect_error(self, caplog): + manager = MCPServerManager() + secret_url = "https://typo-host.example.com/mcp/s/PATHSECRET/mcp" + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=self._connect_error_client(secret_url), + ), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + result = await manager._descovery_metadata(secret_url, warn_when_no_metadata=True) + assert result is None + assert "found no authorization server metadata" in caplog.text + assert "ConnectError" in caplog.text + assert "https://typo-host.example.com" in caplog.text + # hosted MCP urls embed credentials in the path; neither the url nor the exception + # text may leak it into warning-level logs + assert "PATHSECRET" not in caplog.text + + @pytest.mark.asyncio + async def test_descovery_metadata_stays_silent_without_warn_flag(self, caplog): + manager = MCPServerManager() + url = "https://typo-host.example.com/mcp" + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=self._connect_error_client(url), + ), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + result = await manager._descovery_metadata(url) + assert result is None + assert "found no authorization server metadata" not in caplog.text + + @pytest.mark.asyncio + async def test_descovery_metadata_attempt_trail_names_each_failed_step(self, caplog): + manager = MCPServerManager() + url = "https://real-host.example.com/mcp-typo" + client = MagicMock() + client.get = AsyncMock( + return_value=httpx.Response(404, request=httpx.Request("GET", url)) + ) + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=client, + ), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + result = await manager._descovery_metadata(url, warn_when_no_metadata=True) + assert result is None + assert "HTTP 404" in caplog.text + assert "well-known protected-resource lookup found no authorization servers" in caplog.text + assert "origin fallback" in caplog.text + + @pytest.mark.asyncio + async def test_load_servers_from_config_warns_when_endpoints_unresolved(self, caplog): + manager = MCPServerManager() + manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined] + config = { + "typo_server": { + "url": "https://typo.example.com/mcp", + "transport": MCPTransport.http, + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "authorization_code", + } + } + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.load_servers_from_config(config) + assert "typo_server" in caplog.text + assert "authorization_url, token_url" in caplog.text + assert "unresolved" in caplog.text + assert "verify the configured server url" in caplog.text + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "extra_config", + [ + { + "authorization_url": "https://idp.example.com/auth", + "token_url": "https://idp.example.com/token", + }, + { + "oauth2_flow": "client_credentials", + "token_url": "https://idp.example.com/token", + "client_id": "cid", + "client_secret": "csec", + }, + ], + ) + async def test_load_servers_from_config_silent_when_flow_needs_covered(self, caplog, extra_config): + """Manually covered endpoints and M2M servers (which never need authorization_url) must not + warn on every reload; the warning is a misconfiguration signal, not discovery telemetry.""" + manager = MCPServerManager() + manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined] + config = { + "covered_server": { + "url": "https://up.example.com/mcp", + "transport": MCPTransport.http, + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "authorization_code", + **extra_config, + } + } + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.load_servers_from_config(config) + assert "unresolved" not in caplog.text + assert "no discovery source" not in caplog.text + + @pytest.mark.asyncio + async def test_config_server_without_discovery_source_warns_about_missing_endpoints(self, caplog): + manager = MCPServerManager() + manager._register_openapi_tools = AsyncMock() # type: ignore[attr-defined] + config = { + "spec_only": { + "spec_path": "https://example.com/openapi.yaml", + "transport": MCPTransport.http, + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "authorization_code", + } + } + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.load_servers_from_config(config) + assert "no discovery source" in caplog.text + assert "authorization_url and token_url are not set manually" in caplog.text + + @pytest.mark.asyncio + async def test_db_build_warns_when_discovery_fails_for_oauth2_row(self, caplog): + manager = MCPServerManager() + manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined] + record = LiteLLM_MCPServerTable( + server_id="typo-row-1", + server_name="typo_row", + url="https://typo.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) + assert "typo_row" in caplog.text + assert "authorization_url, token_url" in caplog.text + assert "unresolved" in caplog.text From df75d298ec25db8cb69560b7ca3c266ba84636ea Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 21 Jul 2026 22:45:39 -0700 Subject: [PATCH 02/25] fix(mcp): keep url redaction total when the port is malformed urlsplit validates the port lazily, so a non-numeric port raised ValueError out of _redact_mcp_resource_url after the urlsplit try had already passed; the server loaders now call the helper while warning about typo'd urls, which would have turned the warning into a load failure. Resolve hostname and port inside the guard and pin the malformed-port case in the redaction test --- litellm/proxy/_experimental/mcp_server/oauth_utils.py | 6 ++++-- .../proxy/_experimental/mcp_server/test_mcp_server.py | 4 ++++ 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 2f92f75a352..8a5f398003b 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -83,11 +83,13 @@ def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]: return None try: parts = urlsplit(url) + hostname = parts.hostname + port = parts.port except ValueError: return None - if not parts.hostname: + if not hostname: return None - netloc = f"{parts.hostname}:{parts.port}" if parts.port else parts.hostname + netloc = f"{hostname}:{port}" if port else hostname return urlunsplit((parts.scheme, netloc, "", "", "")) or None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ae4f12fc1e1..dff1f1d87c7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -7375,6 +7375,10 @@ async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow(): ("", None), ("not a url", None), ("http://[::1", None), + # urlsplit validates the port lazily on attribute access, so a malformed port must not + # raise out of the helper: the server loaders call it while warning about exactly this + # kind of typo'd url (LIT-4658) + ("https://example.com:bad/mcp", None), ], ) def test_redact_mcp_resource_url_strips_credentials(url, expected): From 3cdd6ab9a126b6e9084060e3643edd27a6a16724 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Fri, 17 Jul 2026 15:39:16 -0700 Subject: [PATCH 03/25] test(e2e): drive a real Linear OAuth MCP through chat completions under both ingress headers --- pyproject.toml | 1 + .../check_e2e_no_raw_requests.py | 7 +- tests/e2e/CLAUDE.md | 3 +- tests/e2e/e2e_config.py | 3 + tests/e2e/mcp/linear_session_capture.py | 58 ++++ tests/e2e/mcp/oauth_chat_client.py | 271 ++++++++++++++++++ .../mcp/test_mcp_chat_completion_oauth_e2e.py | 197 +++++++++++++ tests/e2e/models.py | 82 +++++- uv.lock | 2 + 9 files changed, 620 insertions(+), 4 deletions(-) create mode 100644 tests/e2e/mcp/linear_session_capture.py create mode 100644 tests/e2e/mcp/oauth_chat_client.py create mode 100644 tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py diff --git a/pyproject.toml b/pyproject.toml index 62bd37c3db6..a448ab042b8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -198,6 +198,7 @@ e2e-dev = [ "playwright==1.61.0", "websockets>=15.0.1,<16.0", "locust==2.45.0", + "mcp>=1.28.1,<2.0", ] proxy-dev = [ "prisma==0.11.0", diff --git a/tests/code_coverage_tests/check_e2e_no_raw_requests.py b/tests/code_coverage_tests/check_e2e_no_raw_requests.py index e70e83652d1..fe6a77fc26c 100644 --- a/tests/code_coverage_tests/check_e2e_no_raw_requests.py +++ b/tests/code_coverage_tests/check_e2e_no_raw_requests.py @@ -2,8 +2,10 @@ raw HTTP client imports (requests, urllib.request, httpx, aiohttp, http.client) are banned in suite code. Importing requests' exception types for catching is fine anywhere; a small allowlist grandfathers the files that legitimately make raw calls -(the transport itself, the root conftest liveness probe, and the claude_code version -resolver's constant registry URL fetch). Referenced by tests/e2e/CLAUDE.md.""" +(the transport itself, the root conftest liveness probe, the claude_code version +resolver's constant registry URL fetch, and the mcp OAuth client, whose httpx +client is the object the official mcp SDK's streamable_http_client requires and so +cannot go through the sync requests transport). Referenced by tests/e2e/CLAUDE.md.""" from __future__ import annotations @@ -19,6 +21,7 @@ ALLOWED_RAW_CLIENT_FILES = { "e2e_http.py": ("requests",), "conftest.py": ("requests",), "claude_code/pr_gate_version_resolver.py": ("urllib.request",), + "mcp/oauth_chat_client.py": ("httpx",), } EXCEPTION_ONLY_NAMES = frozenset({"RequestException", "ConnectionError", "Timeout", "HTTPError"}) diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 0e39664e358..17aee22560c 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -14,7 +14,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family - `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`) - `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials (API surface; not Playwright) - `a2a/` - the A2A (agent-to-agent) surface: admin registration via `/v1/agents`, proxy-fronted card discovery at `/.well-known/agent-card.json`, and JSON-RPC `message/send` invocation, driving agents backed by the litellm completion bridge (a real provider) and asserting protocol-version normalization (0.3 vs 1.0) -- `mcp/` - the MCP server surface over api_key auth against the real Datadog remote MCP server only (see "MCP suite: real Datadog only" below) +- `mcp/` - the MCP server surface over api_key auth against the real Datadog remote MCP server (see "MCP suite: real Datadog only" below); plus the gateway-managed OAuth (authorization_code) path exercised through `/chat/completions`, the one behavior Datadog's static-header auth cannot reach, seeding the per-user upstream token via the interactive authorize dance driven with the mcp SDK's own OAuth client (headless-browser consent from a saved session) and asserting the completion lists and executes the server's tools with the stored per-user token - `logging/` - logging-integration delivery (datadog and friends) - `security/` - secret handling and log-leak protection - `router/` - routing and reliability behavior (fallbacks, cooldowns) @@ -33,6 +33,7 @@ Every test under `tests/e2e/mcp/` must exercise the proxy against the real Datad - Prefer calling real Datadog tools that prove the product path (e.g. `search_datadog_logs` for list/call and permission denials). Seed a unique marker (`e2e-datadog-mcp-*`) in a chat completion when you need a log the tool can find; dual-read with `dd_logs` from conftest when delivery matters - Delete the MCP server (and any keys) through `resources.defer` the same way every other suite tears down - If a new MCP behavior cannot be covered with Datadog's tool surface, say so in the PR and get agreement before inventing another upstream; the default is always Datadog +- The one standing exception is `test_mcp_chat_completion_oauth_e2e.py`. Datadog authenticates with the static `DD-API-KEY` / `DD-APPLICATION-KEY` headers and exposes no authorize/token dance at all, so it cannot exercise gateway-managed OAuth or per-user token seeding in any form. That test drives a real Linear MCP server instead; it is still a real remote upstream, so the no-mock, no-fixture rule above holds unchanged ## Lay the pattern down in a class diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index e7c48690c0a..4ecc215a22d 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -38,6 +38,9 @@ UI_BASE_URL = os.environ.get("E2E_UI_BASE_URL", PROXY_BASE_URL).rstrip("/") CHEAP_ANTHROPIC_MODEL = os.environ.get("E2E_CHEAP_ANTHROPIC_MODEL", "claude-haiku-4-5") CHEAP_OPENAI_MODEL = os.environ.get("E2E_CHEAP_OPENAI_MODEL", "gpt-5.5") +LINEAR_MCP_URL = os.environ.get("E2E_LINEAR_MCP_URL", "https://mcp.linear.app/mcp") +LINEAR_STORAGE_STATE = os.environ.get("E2E_LINEAR_STORAGE_STATE", "") + # Jaeger query API of the compose stack's OTEL trace destination (the `jaeger` # service in docker-compose.yml maps it to host 16686). Trace-completeness tests # read exported spans back through it. diff --git a/tests/e2e/mcp/linear_session_capture.py b/tests/e2e/mcp/linear_session_capture.py new file mode 100644 index 00000000000..1c867e17e22 --- /dev/null +++ b/tests/e2e/mcp/linear_session_capture.py @@ -0,0 +1,58 @@ +"""One-time helper to capture a logged-in Linear browser session for the +real-Linear MCP e2e test. + +The real-Linear test drives the genuine gateway-managed authorization_code +dance against ``mcp.linear.app``. The only step that cannot be scripted is +Linear's login (magic link / SSO), so a human authenticates once here and the +resulting session (cookies + local storage) is persisted to disk. The e2e test +then loads that session in a headless Playwright context and clicks Approve on +Linear's consent screen every run, with no human and no login automation. + +Run it with the e2e venv, log into Linear in the window that opens, then return +to the terminal and press Enter: + + LITELLM=~/litellm-mcpe2e + "$LITELLM"/.venv/bin/python "$LITELLM"/tests/e2e/mcp/linear_session_capture.py + +The session is written to ``E2E_LINEAR_STORAGE_STATE`` (default +``~/.litellm-e2e/linear_storage_state.json``), outside the repo. It is a +secret: never commit it. Re-run this whenever Linear expires the session. +""" + +from __future__ import annotations + +import os +from pathlib import Path + +from playwright.sync_api import sync_playwright + +DEFAULT_STATE_PATH = Path.home() / ".litellm-e2e" / "linear_storage_state.json" + + +def capture(state_path: Path) -> None: + """Open a headed browser at Linear, wait for the human to log in, then save + the authenticated session to ``state_path``.""" + state_path.parent.mkdir(parents=True, exist_ok=True) + with sync_playwright() as playwright: + browser = playwright.chromium.launch(headless=False) + context = browser.new_context() + page = context.new_page() + page.goto("https://linear.app/login", wait_until="domcontentloaded") + print("\n" + "=" * 72) + print("Log into Linear in the browser window that just opened.") + print("If Linear emails you a magic link, paste the link into THIS window's") + print("address bar (opening it in your default browser won't capture the") + print("session). Google SSO works too as long as you complete it here.") + print("When your Linear workspace has loaded, come back and press Enter.") + print("=" * 72) + input("Press Enter once you are logged in... ") + page.goto("https://mcp.linear.app/", wait_until="domcontentloaded") + context.storage_state(path=str(state_path)) + browser.close() + print(f"\nSaved Linear session to {state_path}") + print("Point the e2e test at it with:") + print(f' export E2E_LINEAR_STORAGE_STATE="{state_path}"') + + +if __name__ == "__main__": + capture(Path(os.environ.get("E2E_LINEAR_STORAGE_STATE", str(DEFAULT_STATE_PATH)))) diff --git a/tests/e2e/mcp/oauth_chat_client.py b/tests/e2e/mcp/oauth_chat_client.py new file mode 100644 index 00000000000..2eaf512cfa5 --- /dev/null +++ b/tests/e2e/mcp/oauth_chat_client.py @@ -0,0 +1,271 @@ +"""Client for the mcp chat-completion OAuth e2e suite. + +Registers a gateway-managed OAuth (authorization_code) MCP server, seeds the +per-user upstream token by driving the interactive authorize dance with the +official mcp SDK's OAuthClientProvider (the browser leg is a headless Chromium +primed with a human's saved Linear session), then exercises the server through +/chat/completions, where the gateway lists and executes its tools with the +stored per-user token. + +Management routes (/v1/mcp/server CRUD, /chat/completions) go through the +shared ProxyClient transport. The MCP protocol used to seed the token goes through +the mcp SDK, the same library production MCP hosts run. +""" + +from __future__ import annotations + +import asyncio +import re +import time +from dataclasses import dataclass +from typing import TYPE_CHECKING +from urllib.parse import parse_qsl + +import httpx +import pytest +from mcp import ClientSession +from mcp.client.auth import OAuthClientProvider +from mcp.client.streamable_http import streamable_http_client +from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken + +from e2e_config import PROXY_BASE_URL, REQUEST_TIMEOUT +from proxy_client import ProxyClient +from e2e_http import AuthHeaders, NoBody, unwrap +from models import ChatBody, ChatResponse, McpServerCreateBody, McpServerInfo + +if TYPE_CHECKING: + from playwright.async_api import Route + +# Where the "browser" lands at the end of the authorize dance. Nothing listens +# here: the route interceptor short-circuits the final redirect and reads the +# code/state off its query string, exactly like a desktop MCP host intercepting +# its loopback redirect. +OAUTH_CLIENT_REDIRECT_URI = "http://127.0.0.1:53682/e2e/callback" +BROWSER_CONSENT_TIMEOUT = 60.0 + + +def _mcp_url(alias: str) -> str: + return f"{PROXY_BASE_URL}/{alias}/mcp" + + +class InMemoryTokenStorage: + """The mcp SDK's TokenStorage protocol, in memory for one dance: the + DCR-registered client and the gateway tokens minted for it.""" + + def __init__(self) -> None: + self._tokens: OAuthToken | None = None + self._client_info: OAuthClientInformationFull | None = None + + async def get_tokens(self) -> OAuthToken | None: + return self._tokens + + async def set_tokens(self, tokens: OAuthToken) -> None: + self._tokens = tokens + + async def get_client_info(self) -> OAuthClientInformationFull | None: + return self._client_info + + async def set_client_info(self, client_info: OAuthClientInformationFull) -> None: + self._client_info = client_info + + +async def _browser_follow_authorize(start_url: str, storage_state_path: str) -> tuple[str, str | None]: + """Play the browser's role for a real upstream whose authorize endpoint + serves an interactive consent page (Linear). A headless Chromium primed + with a human's saved Linear session opens the gateway authorize URL and + clicks through Linear's consent screens (the mcp.linear.app Approve form, + then the linear.app workspace-selection page), riding the rest of the chain + (Linear -> gateway callback -> host redirect_uri). The final hop is + intercepted and short-circuited, since nothing listens there, and its + code/state are read off the query string.""" + from playwright.async_api import async_playwright + + captured: dict[str, str] = {} # mutable-ok: hand-off from the request listener + trail: list[str] = [] # mutable-ok: navigation diagnostics for a failed dance + + def _note_request(request: object) -> None: + url = getattr(request, "url", "") + if url.startswith(OAUTH_CLIENT_REDIRECT_URI) and "url" not in captured: + captured["url"] = url + + async def _swallow_redirect(route: "Route") -> None: + await route.fulfill(status=200, content_type="text/plain", body="ok") + + async with async_playwright() as playwright: + browser = await playwright.chromium.launch(headless=True) + context = await browser.new_context(storage_state=storage_state_path) + await context.route(re.compile(re.escape(OAUTH_CLIENT_REDIRECT_URI) + r".*"), _swallow_redirect) + page = await context.new_page() + page.on("request", _note_request) + page.on("framenavigated", lambda frame: trail.append(frame.url.split("?", 1)[0])) + await page.goto(start_url, wait_until="domcontentloaded") + deadline = time.monotonic() + BROWSER_CONSENT_TIMEOUT + while "url" not in captured and time.monotonic() < deadline: + try: + await page.wait_for_load_state("networkidle", timeout=8000) + except Exception: # noqa: BLE001 - a busy consent page never idles; fall through and try to advance it + pass + if "url" in captured: + break + control = page.locator( + 'button[name="action"][value="approve"], button:has-text("Authorize"), ' + 'button:has-text("Allow"), button:has-text("@"), a:has-text("@")' + ).first + try: + await control.click(timeout=5000) + except Exception: # noqa: BLE001 - nothing to advance yet; loop and re-check + await asyncio.sleep(0.5) + final_url = page.url + await browser.close() + + landing = captured.get("url") + assert landing is not None, ( + f"consent flow never reached {OAUTH_CLIENT_REDIRECT_URI}; " + f"final={final_url.split('?', 1)[0]!r}; trail={trail[-6:]}" + ) + params = dict(parse_qsl(httpx.URL(landing).query.decode())) + assert "code" in params, f"client redirect_uri carried no code: {landing}" + return params["code"], params.get("state") + + +def _oauth_provider(url: str, storage: InMemoryTokenStorage, storage_state_path: str) -> OAuthClientProvider: + """The SDK's real OAuth machinery (RFC 9728/8414 discovery, RFC 7591 DCR, + PKCE, token exchange) with the browser leg driven by Playwright against the + upstream's consent screen.""" + code_holder: dict[str, str | None] = {} # mutable-ok: hand-off between the two SDK callbacks + + async def redirect_handler(authorize_url: str) -> None: + code, state = await _browser_follow_authorize(authorize_url, storage_state_path) + code_holder["code"] = code + code_holder["state"] = state + + async def callback_handler() -> tuple[str, str | None]: + code = code_holder.get("code") + assert code is not None, "callback_handler ran before the authorize redirect completed" + return code, code_holder.get("state") + + return OAuthClientProvider( + server_url=url, + client_metadata=OAuthClientMetadata.model_validate( + { + "redirect_uris": [OAUTH_CLIENT_REDIRECT_URI], + "token_endpoint_auth_method": "none", + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + "client_name": "e2e-mcp-host", + } + ), + storage=storage, + redirect_handler=redirect_handler, + callback_handler=callback_handler, + ) + + +class _HeaderInjectingTransport(httpx.AsyncBaseTransport): + """Adds the caller's LiteLLM key header to every outgoing SDK request + (discovery, DCR, token exchange), so the gateway resolves which user to + store the upstream token for from the key on the token exchange, exactly + like a production MCP host configured with a LiteLLM key header.""" + + def __init__(self, inner: httpx.AsyncBaseTransport, headers: dict[str, str]) -> None: + self._inner = inner + self._headers = headers + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + for name, value in self._headers.items(): + if name not in request.headers: + request.headers[name] = value + return await self._inner.handle_async_request(request) + + +def _oauth_http_client(headers: dict[str, str], auth: OAuthClientProvider) -> httpx.AsyncClient: + return httpx.AsyncClient( + headers=headers, + auth=auth, + timeout=httpx.Timeout(REQUEST_TIMEOUT), + follow_redirects=True, + transport=_HeaderInjectingTransport(httpx.AsyncHTTPTransport(), headers), + ) + + +async def _seed_via_dance( + url: str, headers: dict[str, str], storage: InMemoryTokenStorage, storage_state_path: str +) -> tuple[str, ...]: + async with _oauth_http_client(headers, _oauth_provider(url, storage, storage_state_path)) as http_client: + async with streamable_http_client(url, http_client=http_client) as (read, write, _): + async with ClientSession(read, write) as session: + await session.initialize() + listed = await session.list_tools() + return tuple(sorted(tool.name for tool in listed.tools)) + + +@dataclass(frozen=True, slots=True) +class ChatMcpClient: + proxy: ProxyClient + + def create_server(self, body: McpServerCreateBody) -> McpServerInfo: + return unwrap( + self.proxy.transport.post( + "/v1/mcp/server", + headers=self.proxy.transport.master, + json=body, + response_type=McpServerInfo, + ) + ) + + def server_info(self, server_id: str) -> McpServerInfo: + return unwrap( + self.proxy.transport.get( + f"/v1/mcp/server/{server_id}", + headers=self.proxy.transport.master, + params=NoBody(), + response_type=McpServerInfo, + ) + ) + + def delete_server(self, server_id: str) -> None: + _ = self.proxy.transport.delete( + f"/v1/mcp/server/{server_id}", + headers=self.proxy.transport.master, + json=NoBody(), + response_type=NoBody, + ) + + def seed_user_token(self, alias: str, key: str, storage_state_path: str) -> tuple[str, ...]: + """Drive the interactive authorize dance for `key`'s user so the gateway + stores their upstream token, retried to the shared deadline since the + just-created server and key propagate asynchronously. The LiteLLM key + rides x-litellm-api-key so the gateway binds the token to that user. + Returns the upstream tool names the dance listed, proof the token works.""" + headers = {"x-litellm-api-key": f"Bearer {key}"} + storage = InMemoryTokenStorage() + deadline = time.monotonic() + self.proxy.poll_timeout + last_error: Exception | None = None + while time.monotonic() < deadline: + try: + return asyncio.run(_seed_via_dance(_mcp_url(alias), headers, storage, storage_state_path)) + except Exception as exc: # noqa: BLE001 - retried to the deadline; the last error surfaces below + last_error = exc + time.sleep(self.proxy.poll_interval) + pytest.fail( + f"authorize dance for {alias!r} never completed within {self.proxy.poll_timeout}s; " + f"last error: {last_error!r}" + ) + + def chat_with_mcp(self, headers: AuthHeaders, body: ChatBody) -> ChatResponse: + """POST /chat/completions carrying the LiteLLM key in `headers` (either + ingress form) with an MCP server attached in `body.tools`. The gateway + resolves the user from the key and lists/executes the server's tools + with that user's stored upstream token.""" + return unwrap( + self.proxy.transport.post( + "/chat/completions", + headers=headers, + json=body, + response_type=ChatResponse, + ) + ) + + +def build_chat_client(proxy: ProxyClient) -> ChatMcpClient: + return ChatMcpClient(proxy=proxy) diff --git a/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py b/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py new file mode 100644 index 00000000000..01e94f7b86f --- /dev/null +++ b/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py @@ -0,0 +1,197 @@ +"""On-demand e2e: a chat completion drives a gateway-managed OAuth MCP server. + +The real end-user flow for MCP over an OAuth server: a user registers a Linear +authorization_code server, authorizes it once so the gateway stores their +upstream token, then sends a normal /chat/completions request with the Linear +MCP attached. The gateway resolves the user from the LiteLLM key, lists Linear's +tools with the stored per-user token, lets the model call one, executes it +upstream with that token, and returns the answer. This is proven against the +real Linear MCP server (mcp.linear.app) and a real Anthropic model, once per +documented ingress header (x-litellm-api-key and Authorization). + +The authorize dance is seeded through the mcp SDK's OAuthClientProvider; the one +step Linear cannot auto-approve is the human consent, so it is captured once out +of band (mcp/linear_session_capture.py) into a saved browser session and a +headless Chromium clicks Approve every run. The test therefore skips unless +E2E_LINEAR_STORAGE_STATE points at that session, so it never runs on the per-PR +CI path; it is a nightly/on-demand real-server smoke test. + +Fail-before-fix: without the stored per-user token the gateway lists no Linear +tools, so mcp_list_tools comes back empty, nothing is called, and the +assertions fail; a served, called, non-empty Linear tool proves the gateway +pulled and used the user's token. +""" + +from __future__ import annotations + +import os + +import pytest + +from e2e_config import CHEAP_ANTHROPIC_MODEL, LINEAR_MCP_URL, LINEAR_STORAGE_STATE, unique_marker +from e2e_http import AuthHeaders +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, KeyGenerateBody, McpChatTool, McpServerCreateBody, ObjectPermission +from proxy_client import ProxyClient + +pytest.importorskip("mcp", reason="mcp SDK not installed; run `uv sync --inexact --group e2e-dev`") +pytest.importorskip( + "playwright.async_api", + reason="playwright not installed; run `uv pip install playwright` and `playwright install chromium`", +) + +from oauth_chat_client import ChatMcpClient, build_chat_client # noqa: E402 # imports follow the importorskip guards + +pytestmark = [ + pytest.mark.e2e, + pytest.mark.skipif( + not LINEAR_STORAGE_STATE or not os.path.exists(LINEAR_STORAGE_STATE), + reason="set E2E_LINEAR_STORAGE_STATE to a Linear session captured via mcp/linear_session_capture.py", + ), +] + +# Pinned from a live dance during verification (never guessed); the gateway +# prefixes every upstream tool name with the server alias. list_teams is a +# read-only Linear tool that takes no arguments and returns the caller's teams. +LINEAR_READONLY_TOOL = "list_teams" +LINEAR_PROMPT = "Use the list_teams tool to list my Linear teams, then reply with the name of one of them." + + +@pytest.fixture(scope="session") +def chat_client(proxy: ProxyClient) -> ChatMcpClient: + return build_chat_client(proxy) + + +class TestMcpChatCompletionOauth: + """A scoped internal-user key on a real Linear authorization_code server, + used through /chat/completions once per ingress header: the gateway pulls + the user's stored upstream token, lists and executes Linear's tools during + the completion, and returns the answer.""" + + @pytest.mark.covers("mcp.list_tools.oauth.succeeds") + @pytest.mark.covers("mcp.call_tool.oauth.succeeds") + def test_chat_completion_uses_linear_with_x_litellm_api_key_header( + self, chat_client: ChatMcpClient, resources: ResourceManager + ) -> None: + marker = unique_marker() + alias = f"e2elinear{marker}" + created = chat_client.create_server( + McpServerCreateBody( + alias=alias, + url=LINEAR_MCP_URL, + allow_all_keys=False, + auth_type="oauth2", + oauth2_flow="authorization_code", + ) + ) + resources.defer(lambda: chat_client.delete_server(created.server_id)) + + stored = chat_client.server_info(created.server_id) + assert stored.auth_type == "oauth2" + assert stored.oauth2_flow == "authorization_code" + assert stored.allow_all_keys is False + + key = chat_client.proxy.generate_key( + KeyGenerateBody( + user_id="e2e-test-user", + object_permission=ObjectPermission(mcp_servers=[created.server_id]), + ) + ) + resources.defer(lambda: chat_client.proxy.delete_key(key)) + + seeded = chat_client.seed_user_token(alias, key, LINEAR_STORAGE_STATE) + assert f"{alias}-{LINEAR_READONLY_TOOL}" in seeded, ( + f"the authorize dance listed {seeded}, expected it to include {alias}-{LINEAR_READONLY_TOOL}" + ) + + response = chat_client.chat_with_mcp( + AuthHeaders.model_validate({"x-litellm-api-key": f"Bearer {key}"}), + ChatBody( + model=CHEAP_ANTHROPIC_MODEL, + messages=[ChatMessage(role="user", content=LINEAR_PROMPT)], + tools=[ + McpChatTool( + server_url=f"litellm_proxy/mcp/{alias}", + server_label=alias, + require_approval="never", + ) + ], + ), + ) + + message = response.choices[0].message + assert message is not None and message.content, f"completion returned no answer: {response}" + meta = message.provider_specific_fields + assert meta is not None, f"no MCP metadata on the completion: {response}" + listed = {t.function.name for t in (meta.mcp_list_tools or []) if t.function} + assert f"{alias}-{LINEAR_READONLY_TOOL}" in listed, ( + f"the gateway listed {sorted(listed)}, expected the stored token to surface {alias}-{LINEAR_READONLY_TOOL}" + ) + results = [r for r in (meta.mcp_call_results or []) if r.name == f"{alias}-{LINEAR_READONLY_TOOL}"] + assert results and results[0].result, ( + f"Linear tool {alias}-{LINEAR_READONLY_TOOL} was not executed with a result: {meta.mcp_call_results}" + ) + + @pytest.mark.covers("mcp.list_tools.oauth.succeeds") + @pytest.mark.covers("mcp.call_tool.oauth.succeeds") + def test_chat_completion_uses_linear_with_authorization_bearer_header( + self, chat_client: ChatMcpClient, resources: ResourceManager + ) -> None: + marker = unique_marker() + alias = f"e2elinear{marker}" + created = chat_client.create_server( + McpServerCreateBody( + alias=alias, + url=LINEAR_MCP_URL, + allow_all_keys=False, + auth_type="oauth2", + oauth2_flow="authorization_code", + ) + ) + resources.defer(lambda: chat_client.delete_server(created.server_id)) + + stored = chat_client.server_info(created.server_id) + assert stored.auth_type == "oauth2" + assert stored.oauth2_flow == "authorization_code" + assert stored.allow_all_keys is False + + key = chat_client.proxy.generate_key( + KeyGenerateBody( + user_id="e2e-test-user", + object_permission=ObjectPermission(mcp_servers=[created.server_id]), + ) + ) + resources.defer(lambda: chat_client.proxy.delete_key(key)) + + seeded = chat_client.seed_user_token(alias, key, LINEAR_STORAGE_STATE) + assert f"{alias}-{LINEAR_READONLY_TOOL}" in seeded, ( + f"the authorize dance listed {seeded}, expected it to include {alias}-{LINEAR_READONLY_TOOL}" + ) + + response = chat_client.chat_with_mcp( + AuthHeaders.model_validate({"authorization": f"Bearer {key}"}), + ChatBody( + model=CHEAP_ANTHROPIC_MODEL, + messages=[ChatMessage(role="user", content=LINEAR_PROMPT)], + tools=[ + McpChatTool( + server_url=f"litellm_proxy/mcp/{alias}", + server_label=alias, + require_approval="never", + ) + ], + ), + ) + + message = response.choices[0].message + assert message is not None and message.content, f"completion returned no answer: {response}" + meta = message.provider_specific_fields + assert meta is not None, f"no MCP metadata on the completion: {response}" + listed = {t.function.name for t in (meta.mcp_list_tools or []) if t.function} + assert f"{alias}-{LINEAR_READONLY_TOOL}" in listed, ( + f"the gateway listed {sorted(listed)}, expected the stored token to surface {alias}-{LINEAR_READONLY_TOOL}" + ) + results = [r for r in (meta.mcp_call_results or []) if r.name == f"{alias}-{LINEAR_READONLY_TOOL}"] + assert results and results[0].result, ( + f"Linear tool {alias}-{LINEAR_READONLY_TOOL} was not executed with a result: {meta.mcp_call_results}" + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index b3ea9346180..d21920cf848 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -6,6 +6,7 @@ response validates without mirroring every proxy field. No untyped dicts. from __future__ import annotations +from collections.abc import Sequence from datetime import datetime from typing import Literal @@ -196,6 +197,19 @@ class ChatTool(BaseModel): function: ChatToolFunction +class McpChatTool(BaseModel): + """An MCP server attached to a chat completion (OpenAI `type: "mcp"` tool). + `server_url` selects the gateway-registered server by its alias suffix; with + `require_approval="never"` the gateway lists, calls, and feeds the server's + tools back to the model in one agentic turn.""" + + type: Literal["mcp"] = "mcp" + server_url: str + require_approval: str + server_label: str | None = None + allowed_tools: list[str] | None = None + + class ChatBody(BaseModel): model: str messages: list[ChatMessage] @@ -206,7 +220,7 @@ class ChatBody(BaseModel): reasoning_effort: str | None = None thinking: ThinkingParam | None = None service_tier: str | None = None - tools: list[ChatTool] | None = None + tools: Sequence[ChatTool | McpChatTool] | None = None tool_choice: str | None = None guardrails: list[str] | None = None response_format: dict[str, object] | None = None @@ -242,10 +256,46 @@ class ToolCall(BaseModel): function: ToolCallFunction = ToolCallFunction() +class McpToolFunctionRef(BaseModel): + name: str + + +class McpListedTool(BaseModel): + """One entry of `mcp_list_tools`: a tool the gateway listed from the + attached MCP server and exposed to the model, in OpenAI function shape.""" + + function: McpToolFunctionRef | None = None + + +class McpToolCall(BaseModel): + """One entry of `mcp_tool_calls`: a tool the model asked the gateway to run.""" + + function: McpToolFunctionRef | None = None + + +class McpCallResult(BaseModel): + """One entry of `mcp_call_results`: what the gateway got back from executing + a tool upstream on the caller's behalf.""" + + name: str | None = None + result: str | None = None + + +class McpResponseMetadata(BaseModel): + """`choices[].message.provider_specific_fields` MCP section: which tools the + gateway listed from the attached server, which the model called, and their + results. Populated only when the completion drove an MCP server.""" + + mcp_list_tools: list[McpListedTool] | None = None + mcp_tool_calls: list[McpToolCall] | None = None + mcp_call_results: list[McpCallResult] | None = None + + class OutMessage(BaseModel): content: str | None = None reasoning_content: str | None = None tool_calls: list[ToolCall] | None = None + provider_specific_fields: McpResponseMetadata | None = None class ChatChoice(BaseModel): @@ -355,6 +405,36 @@ class CountTokensResponse(BaseModel): input_tokens: int +# ---------- mcp servers ---------- + + +class McpServerCreateBody(BaseModel): + """POST /v1/mcp/server. For a gateway-managed OAuth server, `auth_type` is + `oauth2` and `oauth2_flow` is `authorization_code`; the upstream endpoints + are discovered and registered via DCR when left unset. `allow_all_keys` + false scopes the server to keys granted it through object_permission.""" + + alias: str + url: str + transport: str = "http" + allow_all_keys: bool = True + auth_type: str | None = None + oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None + authorization_url: str | None = None + token_url: str | None = None + + +class McpServerInfo(BaseModel): + """Response of POST /v1/mcp/server and GET /v1/mcp/server/{server_id}.""" + + server_id: str + alias: str | None = None + url: str | None = None + auth_type: str | None = None + oauth2_flow: str | None = None + allow_all_keys: bool | None = None + + class EmbedBody(BaseModel): model: str input: str diff --git a/uv.lock b/uv.lock index b3d6fccff26..9c60ca1ee48 100644 --- a/uv.lock +++ b/uv.lock @@ -4301,6 +4301,7 @@ dev = [ ] e2e-dev = [ { name = "locust" }, + { name = "mcp" }, { name = "playwright" }, { name = "websockets" }, ] @@ -4478,6 +4479,7 @@ dev = [ ] e2e-dev = [ { name = "locust", specifier = "==2.45.0" }, + { name = "mcp", specifier = ">=1.28.1,<2.0" }, { name = "playwright", specifier = "==1.61.0" }, { name = "websockets", specifier = ">=15.0.1,<16.0" }, ] From a78130461f52c0edbb14e5102ef7717dbbb11b4f Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 21 Jul 2026 17:53:23 -0700 Subject: [PATCH 04/25] feat(mcp): gateway DCR session admission at the aggregate /mcp endpoint (LIT-3637) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Admits a keyless SSO user (no virtual key) at the aggregate /mcp endpoint from a gateway DCR session bearer, resolving team/org/SCIM/budget authorization fresh on every call. - Aggregate DCR front door: stateless /register (sealed llm_dcrc_ client ids), SSO-backed /authorize + /authorize/complete, and /token minting identity-only session tokens with PKCE, single-use codes/flows, and rotating refresh tokens. - Admission: a session-shaped Authorization at the aggregate scope opens via _admit_gateway_session, reloads the live user, and runs the centralized policy gate; failures return the RFC 9728 invalid_token challenge. Gated on the un-forgeable, server-only mcp_admitted_user_subject marker, so virtual-key and JWT auth are unchanged. - Authorization model: an admitted subject is resolved as one plain UserAPIKeyAuth per grant source (its own grants, plus each team it is a live roster member of), each answered by the SAME resolver virtual keys use, then unioned. That branch is the FIRST statement of BOTH public resolvers, so no single-credential prelude runs for it and a fault in a lookup it never uses cannot deny its grants. A source team counts only while it is a live grantor: roster membership, not blocked, and neither the team nor its owning org over budget (enforced through the SAME _team_max_budget_check / _organization_max_budget_check owners common_checks uses for keys). Each team source carries that team's own org, so the existing org ceiling caps it; for a keyless source the org list only ever intersects (a ceiling must not become a grant) and an unresolvable ceiling denies rather than silently uncapping, on both the server and tool axes. _roster_team_object is the single owner of "which teams count": a team whose roster no longer lists the user neither grants servers nor throttles, in one place. - Rate limits: the subject is bounded by its user rpm/tpm AND by the per-server mcp_rpm_limit of the team a call is ATTRIBUTED to — the same single source billing charges, from the same owner. A key charges its one pinned team's bucket; a keyless subject has no team_id, so admission stamps each granting team's limit map onto the auth (server-only field, stripped from validated input like the marker) and the limiter emits that team's mcp_per_team descriptor. Charging every granting team instead would let one cross-team user drain several teams' SHARED buckets on a single call and block their other members; and a server the user's OWN grant reaches charges no team bucket at all, because no team provided it. Per-KEY MCP limits do not apply because there is no key. - Wrapper channels: the manager-level union treats the admitted subject by the same grant model. The admin-role short-circuit and the absolute no_mcp_servers early-return are key-credential rules and never apply to it (a session bearer is a third-party client credential, not the dashboard, and the subject's opt-out silences only its own source). Operator-open channels (allow_all_keys, the user's own BYOM submissions) are owned by one operator_open_server_ids helper that BOTH the server union and the admitted tool resolution consult (suppress-BYOM-when- explicitly-scoped is a key-credential rule and never applies to the subject, whose user row carries the DB-default empty mcp_servers), so an open-channel server is default-open for tools instead of listable but uninvokable. - Redirect URIs: one owner, validate_redirect_uri_shape, decides redirect-URI hygiene (bad scheme, fragment, missing host, userinfo, backslash host) and resolves allowlisted native callbacks, shared by DCR registration and the OAuth endpoints. Registration keeps a deliberately wider trust policy than validate_trusted_redirect_uri: public dynamic registration accepts any https client, and its controls are mandatory S256 PKCE plus the consent screen. - Egress leak-defense: a gateway admission credential (session bearer / bridge envelope) is scrubbed from EVERY egress header context, anchored to the credential shape, so it can never be forwarded upstream and replayed. - Single-use guard: auth-code, refresh and connect-flow claims resolve the proxy's cross-worker redis cache themselves rather than trusting the cache passed in, and fail CLOSED on a Redis fault instead of falling back to a per-worker count that a captured id could replay through another worker. - Sign-in return_to: one shared, never-raising helper persists a safe return_to for every sign-in branch (SSO/Okta/generic and username/password), and every branch RESUMES through the same _sso_return_to_redirect the SSO callback uses, so however a deployment signs in the stored value is honored identically (same-origin path directly; control_plane_url via the one-time login-code handoff). A stale cookie is ignored rather than failing a completed sign-in. - Budgets, both halves: ENFORCEMENT (an already over-budget team or its owning org stops being a grantor, in the source gate) and ACCOUNTING (a team-derived tool call is billed to the granting team and ITS org, so that budget accumulates and the right organization is charged). A server the user's own grant reaches bills the user; when several teams grant one server the pick is the lowest team_id, stable and auditable. Billing rides a COPY, so authorization still sees the full union, and it is inert when the target server cannot be resolved from the tool name. Deferred (tracked): client-selected server scoping of the session token (LIT-4680). Co-Authored-By: Claude Opus 4.8 --- .../mcp_server/auth/user_api_key_auth_mcp.py | 951 +++++++++-- .../mcp_server/discoverable_endpoints.py | 81 +- .../mcp_server/gateway_dcr_flow.py | 637 ++++++++ .../mcp_server/mcp_server_manager.py | 94 +- .../_experimental/mcp_server/oauth_utils.py | 42 +- .../session_credentials.py | 3 +- .../outbound_credentials/session_token.py | 8 +- .../proxy/_experimental/mcp_server/server.py | 12 +- litellm/proxy/_types.py | 16 + litellm/proxy/auth/auth_checks.py | 44 +- litellm/proxy/auth/login_utils.py | 26 + .../hooks/parallel_request_limiter_v3.py | 46 +- litellm/proxy/management_endpoints/ui_sso.py | 145 +- litellm/proxy/proxy_server.py | 65 +- .../auth/test_user_api_key_auth_mcp.py | 1431 ++++++++++++++++- .../mcp_server/test_discoverable_endpoints.py | 178 +- .../mcp_server/test_gateway_dcr_flow.py | 590 +++++++ .../proxy/auth/test_login_utils.py | 79 +- .../proxy/management_endpoints/test_ui_sso.py | 77 + .../proxy_server/test_routes_login_sso.py | 129 +- tests/test_litellm/proxy/test_proxy_server.py | 14 +- ui/litellm-dashboard/eslint-suppressions.json | 2 +- .../src/app/chat/integrations/page.tsx | 16 +- .../chat/ConnectFlowBanner.test.tsx | 51 + .../src/components/chat/ConnectFlowBanner.tsx | 59 + .../src/components/chat/MCPAppsPanel.tsx | 112 +- .../src/components/mcp_tools/types.test.tsx | 21 + .../src/components/mcp_tools/types.tsx | 9 + 28 files changed, 4545 insertions(+), 393 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py create mode 100644 ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.test.tsx create mode 100644 ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.tsx diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index f1fcc95c532..01e8b1490b2 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -25,6 +25,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credenti from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( EnvelopeIdentity, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + is_session_bearer_shaped, +) from litellm.proxy._types import ( UI_TEAM_ID, LiteLLM_TeamTable, @@ -124,6 +127,29 @@ def _has_client_supplied_mcp_auth( return bool(mcp_auth_header) or bool(mcp_server_auth_headers) +def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool: + """True when this auth is a keyless subject admitted by the gateway session / bridge user + path, as opposed to a JWT or other keyless auth that merely lacks a ``team_id``. + + Reads the server-only ``UserAPIKeyAuth.mcp_admitted_user_subject`` field, set exclusively by + ``_reload_admitted_user`` at admission. It is deliberately NOT a ``metadata`` key: virtual-key + metadata is caller-controlled at key creation, so a metadata marker could be forged on a + personal key to gain the team-inherited grant union or to dodge the caller-Authorization + egress scrub. This field cannot be set from caller input.""" + return user_api_key_auth is not None and user_api_key_auth.mcp_admitted_user_subject is True + + +def _is_aggregate_mcp_scope(route: str, mcp_servers: list[str] | None) -> bool: + """True when a request targets the aggregate ``/mcp`` endpoint rather than any named + server. Named targets arrive either through ``x-mcp-servers`` (``mcp_servers``) or a + path segment (``/mcp/{server}`` / ``/{server}/mcp``); the aggregate scope has neither. + The gateway-DCR session arm and challenge fire only here, so a per-server flow is never + affected.""" + if mcp_servers: + return False + return len(MCPRequestHandler._extract_target_server_names_from_path(route)) == 0 + + def _is_aggregate_gateway_dcr_challenge_scope( route: str, mcp_servers: list[str] | None, @@ -141,11 +167,9 @@ def _is_aggregate_gateway_dcr_challenge_scope( client. Fails closed to the original admission error otherwise.""" if not _is_litellm_auth_admission_error(exc): return False - if mcp_servers: - return False if _has_client_supplied_mcp_auth(mcp_auth_header, mcp_server_auth_headers): return False - return len(MCPRequestHandler._extract_target_server_names_from_path(route)) == 0 + return _is_aggregate_mcp_scope(route, mcp_servers) def _aggregate_gateway_dcr_challenge(request: Request, invalid_token: bool) -> HTTPException: @@ -362,6 +386,21 @@ class MCPRequestHandler: request=request, route=request_route, ) + elif ( + _is_aggregate_mcp_scope(request_route, mcp_servers) + and oauth2_headers + and is_session_bearer_shaped(oauth2_headers["Authorization"]) + ): + # A gateway DCR session bearer at the aggregate /mcp scope: open the + # identity-only session token and admit under the live litellm user it + # references. A session-shaped bearer that does not open fails closed with + # the aggregate invalid_token challenge; a non-session bearer never reaches + # here (is_session_bearer_shaped is false) and falls through to the oauth2 arm. + validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session( + authorization_value=oauth2_headers["Authorization"], + request=request, + route=request_route, + ) elif oauth2_headers: # Authorization on a non-delegated server: the bearer must be a real # LiteLLM credential, so a failed validation is a genuine 401/403 and @@ -392,15 +431,87 @@ class MCPRequestHandler: bearer_presented=False, ) + # Leak-defense (single chokepoint): a gateway admission credential — the session bearer or the + # bridge envelope — is NEVER a valid upstream MCP token. Scrub it from EVERY egress header context + # (top-level Authorization, the deprecated `x-mcp-auth`, and per-server `x-mcp-{alias}-authorization`) + # so no client-forwarded, OBO-subject, or passthrough path can send it upstream, where a hostile + # server could capture and replay it against the aggregate endpoint as this user. Anchored to the + # credential SHAPE, so a legitimate upstream/passthrough token (never session- or envelope-shaped) + # is forwarded unchanged; per-server vaulted credentials (resolved at egress) are unaffected. + raw_headers = dict(headers) + ( + oauth2_headers, + raw_headers, + mcp_auth_header, + mcp_server_auth_headers, + ) = MCPRequestHandler._scrub_gateway_admission_credentials( + admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth), + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + ) + return ( validated_user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, oauth2_headers, - dict(headers), + raw_headers, ) + @staticmethod + def _is_gateway_admission_credential(value: str | None) -> bool: + """True when a header value is a gateway admission credential — a session bearer (``llm_session_`` / + ``llm_srefresh_``) or a bridge envelope. Such a value proves who signed in to the GATEWAY; it is + never a valid credential for an UPSTREAM MCP server, so it must never be forwarded, where a hostile + upstream could capture and replay it against the aggregate ``/mcp`` endpoint as this user.""" + return value is not None and (is_session_bearer_shaped(value) or is_bridge_envelope_shaped(value)) + + @staticmethod + def _scrub_gateway_admission_credentials( + admitted: bool, + oauth2_headers: dict[str, str] | None, + raw_headers: dict[str, str], + mcp_auth_header: str | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + ) -> tuple[dict[str, str] | None, dict[str, str], str | None, dict[str, dict[str, str]] | None]: + """Remove any gateway admission credential from EVERY egress header context, keyed on the credential + SHAPE: the top-level ``Authorization`` (``oauth2_headers`` + ``raw_headers``), the deprecated + ``x-mcp-auth`` (``mcp_auth_header``), and per-server ``x-mcp-{alias}-authorization`` + (``mcp_server_auth_headers``). A legitimate upstream/passthrough token is never session- or + envelope-shaped, so it is forwarded unchanged; the per-server token the bridge arm injects is the + real upstream credential (also not gateway-shaped), so it survives. An admitted subject's top-level + Authorization IS the admission bearer, so it is dropped unconditionally as defense-in-depth even + though it is already gateway-shaped.""" + cred = MCPRequestHandler._is_gateway_admission_credential + + # 1. Top-level Authorization → oauth2_headers. + authz = oauth2_headers.get("Authorization") if oauth2_headers else None + if admitted or cred(authz): + oauth2_headers = None + + # 2. raw_headers: drop the admitted subject's Authorization, and ANY header whose value is a + # gateway credential (covers x-mcp-auth and x-mcp-{alias}-authorization in their raw form). + raw_headers = { + k: v for k, v in raw_headers.items() if not ((admitted and k.lower() == "authorization") or cred(v)) + } + + # 3. Deprecated x-mcp-auth value. + if cred(mcp_auth_header): + mcp_auth_header = None + + # 4. Per-server x-mcp-{alias}-authorization values (drop the value, then any now-empty server dict). + if mcp_server_auth_headers: + stripped = { + alias: {h: val for h, val in hdrs.items() if not cred(val)} + for alias, hdrs in mcp_server_auth_headers.items() + } + mcp_server_auth_headers = {alias: hdrs for alias, hdrs in stripped.items() if hdrs} + + return oauth2_headers, raw_headers, mcp_auth_header, mcp_server_auth_headers + @staticmethod def _extract_target_server_names_from_path(path: str) -> List[str]: """ @@ -626,6 +737,71 @@ class MCPRequestHandler: case _: assert_never(result) + @staticmethod + async def _admit_gateway_session( + authorization_value: str, + request: Request, + route: str, + ) -> UserAPIKeyAuth: + """Open a gateway DCR session bearer and admit the live litellm user it references. + + The custody sibling of :meth:`_admit_dcr_bridge_delegate`: the session token seals + no upstream credential (those are vaulted per user and resolved at egress), so this + admits identity only and injects no per-server header. The token's signature proves + the user signed in when it was minted, but authorization is resolved fresh here, the + sealed ``user_id`` reloads the current user record through the SAME + :meth:`_reload_admitted_user` the bridge user-subject path uses, and the admitted + identity runs through the centralized policy gate, so the user's present team, org, + budget, and SCIM state gate the request rather than a snapshot frozen at mint time. + + Fails closed with the aggregate ``invalid_token`` challenge on an expired, tampered, + or foreign token, on a refresh token presented at the tool edge, and when the + referenced user is missing, deactivated, or rejected by the policy gate. The + pre-DB gates (size, IP, route allowlist) run first, mirroring the bridge arm and the + standard pipeline, so a caller blocked by IP or route is turned away before any + crypto or DB read.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + NotSessionBearer, + SessionBearerAdmitted, + SessionBearerInvalid, + resolve_session_bearer, + session_keys_from_master_key, + ) + from litellm.proxy.proxy_server import master_key + + if not master_key: + raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set") + + await MCPRequestHandler._run_pre_db_read_auth_checks(request=request, route=route) + + keys = session_keys_from_master_key(master_key) + result = resolve_session_bearer(authorization_value, keys, datetime.now(timezone.utc)) + match result: + case SessionBearerAdmitted(): + try: + admitted = await MCPRequestHandler._reload_admitted_user(result.principal.user_id) + await MCPRequestHandler._enforce_admitted_live_policy( + admitted=admitted, request=request, route=route + ) + except HTTPException as exc: + # A cryptographically valid bearer whose referenced user is now missing or + # SCIM-deactivated is an invalid_token at the aggregate scope: relay the RFC 9728 + # challenge so the DCR client re-authorizes, matching the SessionBearerInvalid + # arm, instead of a bare 401 with no WWW-Authenticate. A 503 (DB outage) is a + # transient availability failure, not an auth failure, so it passes through. + if exc.status_code == 401: + raise _aggregate_gateway_dcr_challenge(request, invalid_token=True) from exc + raise + return admitted + case SessionBearerInvalid(): + raise _aggregate_gateway_dcr_challenge(request, invalid_token=True) + case NotSessionBearer(): + # Unreachable: the arm is entered only for an is_session_bearer_shaped + # value. Kept for match exhaustiveness and fails closed regardless. + raise _aggregate_gateway_dcr_challenge(request, invalid_token=True) + case _: + assert_never(result) + @staticmethod async def _run_pre_db_read_auth_checks(request: Request, route: str) -> None: """Run the proxy-wide gates ``user_api_key_auth`` applies before any key lookup: the @@ -671,14 +847,18 @@ class MCPRequestHandler: The DCR client authenticates via SSO at the bridged authorize, which yields a user subject rather than a virtual key, so the envelope admits under the user's own - identity: the reloaded ``user_id`` and the user's own MCP object permission ride on the - returned ``UserAPIKeyAuth``, and the SAME ``get_allowed_mcp_servers`` the key path uses then - computes which servers the user may reach, so the user's litellm MCP grants and access groups - gate the request exactly as a key's do. Only the user's OWN object permission is bound: a - ``UserAPIKeyAuth`` carries a single ``team_id`` while a user may belong to many teams, so - team-inherited MCP grants for a user are a follow-up (they need a many-teams union - ``get_allowed_mcp_servers`` does not do off one auth object). The caller's centralized policy - gate enforces the user's live budget and org state, and a SCIM-deactivated owner fails closed. + identity: the reloaded ``user_id``, the user's own MCP object permission, and the user's + ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the SAME ``get_allowed_mcp_servers`` + the key path uses then computes which servers the user may reach, so the user's litellm MCP + grants and access groups gate the request exactly as a key's do. Because the returned auth is + stamped ``mcp_admitted_user_subject`` (below), ``get_allowed_mcp_servers`` unions the servers the + user reaches through ANY of their teams on top of these direct grants — a ``UserAPIKeyAuth`` + pins one ``team_id`` but a user belongs to many, so the team fan-out happens off the marker, not + the single ``team_id``. Each source is bounded by ITS OWN org: the user's direct grants by the + bound ``org_id`` (their primary org), and each team's grant by that team's owning org inside + ``_allowed_mcp_servers_for_single_team`` — so a user who spans organizations does not leak one + org's servers past another org's ceiling. The caller's centralized policy gate enforces the + user's live budget and org state, and a SCIM-deactivated owner fails closed. Error handling mirrors the key path's retryable-503 contract, but ``get_user_object`` defeats a type-based check: where ``get_key_object`` raises a typed ``ProxyException`` for a missing key @@ -721,12 +901,92 @@ class MCPRequestHandler: raise HTTPException(status_code=401, detail="Invalid or expired credential") if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False: raise HTTPException(status_code=401, detail="Invalid or expired credential") - return UserAPIKeyAuth( + admitted = UserAPIKeyAuth( user_id=user_object.user_id, user_role=user_object.user_role, + org_id=user_object.organization_id, object_permission=object_permission, object_permission_id=user_object.object_permission_id, + # Copy the live user's rate limits, exactly as the standard user-subject auth path does + # (user_api_key_auth.py). The parallel limiter reads these off the auth object rather than + # re-fetching, and treats None as sys.maxsize (unlimited), so a keyless admitted user with + # them unset would invoke tools past their configured user RPM/TPM. + # + # Rate-limit model for the keyless admitted subject: bounded by their USER rpm/tpm + # (copied here; those descriptors key off user_id, which is set) AND by the per-server + # mcp_rpm_limit of EVERY team it reaches servers through, stamped below. Per-KEY MCP + # limits genuinely do not apply, because there is no key. + user_tpm_limit=user_object.tpm_limit, + user_rpm_limit=user_object.rpm_limit, ) + # Set the server-only admission marker AFTER construction: the before-validator strips it + # from any validated input, so a post-construction assignment is the only way to set it, and + # caller-supplied data (key metadata, JWT claims) can never forge it. + admitted.mcp_admitted_user_subject = True + # Carry each granting team's per-server MCP rpm limit. A key is pinned to one team so the + # limiter reads team_metadata directly; this subject reaches servers through several teams + # under its own identity, so without this the team ceiling silently does not apply to it and + # a cross-team user outruns every team's mcp_rpm_limit. Resolved from the same roster-checked + # sources the grant union uses, so a team can only throttle what it actually granted. + admitted.mcp_source_team_rpm_limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(admitted) + return admitted + + @staticmethod + async def _admitted_subject_team_rpm_limits(auth: UserAPIKeyAuth) -> dict[str, dict[str, int]] | None: + """``team_id -> mcp_rpm_limit`` for every team this subject reaches servers through, with each + team's map filtered to the servers THAT team's grant actually reaches. + + A limit rides the same scope as the access it bounds: a team's throttle exists to cap usage of + the access the team granted, so a roster team whose grant does not reach a server (not granted, + blocked, org-forbidden, opted out) must not be charged when the user reaches that server + through a DIFFERENT team — otherwise this user's calls drain a bucket shared by that team's own + keys for access the team never provided. The grant scope comes from the SAME + ``get_allowed_mcp_servers(source)`` call authorization uses, so the throttle scope cannot + diverge from the access scope. Limit maps are keyed by server name/alias (the limiter matches + on the called server's name) while grants are ids, so each key is resolved through + ``expand_permission_list`` — the one existing name->id owner — before the membership check. + + Returns None when no team contributes an applicable limit, so the limiter adds no descriptors + rather than empty ones. A lookup failure narrows to None rather than raising: rate limiting + must not be able to deny a request that authorization already allowed, and the user's own + rpm/tpm still bounds them.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + try: + limits: dict[str, dict[str, int]] = {} + source_grants = await MCPRequestHandler.admitted_source_grants(auth) + for source, granted_ids in source_grants: + if not source.team_id: + continue + team_obj = await MCPRequestHandler._roster_team_object(source.team_id, auth) + team_limit = (team_obj.metadata or {}).get("mcp_rpm_limit") if team_obj is not None else None + if not isinstance(team_limit, dict) or not team_limit: + continue + applicable: dict[str, int] = {} + for server_name, rpm in team_limit.items(): + for server_id in global_mcp_server_manager.expand_permission_list([server_name]): + if server_id not in granted_ids: + continue + # Charge ONLY the source the call is attributed to — the same single source + # billing picks, from the same owner. Adding a descriptor for every granting + # team let one cross-team user drain several teams' SHARED buckets at once, + # blocking their other members for access those teams did not provide on + # this call; and when the user's OWN grant reaches the server, no team + # provided it, so no team bucket is charged at all. + attributed = await MCPRequestHandler.attributing_source_for_server( + auth, server_id, source_grants=source_grants + ) + if attributed is not None and attributed.team_id == source.team_id: + applicable[server_name] = rpm + break + if applicable: + limits[source.team_id] = applicable + return limits or None + except Exception as e: # noqa: BLE001 # throttling metadata must never fail an allowed request + verbose_logger.warning(f"Failed to resolve per-team MCP rpm limits for admitted subject: {str(e)}") + return None @staticmethod async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth: @@ -1075,6 +1335,8 @@ class MCPRequestHandler: @staticmethod async def get_allowed_mcp_servers( user_api_key_auth: Optional[UserAPIKeyAuth] = None, + *, + keyless_source: bool = False, ) -> List[str]: """ Get list of allowed MCP servers for the given user/key based on permissions. @@ -1096,6 +1358,14 @@ class MCPRequestHandler: from litellm.proxy.proxy_server import general_settings try: + # A keyless admitted subject is resolved entirely per source, BEFORE any single-source + # rule runs here. Ordering matters: the no_mcp_servers opt-out below reads the caller's + # own object_permission, so leaving it above this branch let a user's own opt-out zero + # their TEAMS' grants too — the sources are independent, and an opt-out on one of them + # must silence only that one (it is applied per source, inside the recursive call). + if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None: + return await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth) + # Get allowed servers from key and team allowed_mcp_servers_for_key = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) @@ -1125,8 +1395,18 @@ class MCPRequestHandler: # team's by default. With require_key_mcp_access_defined the # team is a ceiling rather than a default, so the key must # grant servers explicitly (or via an access group) to reach - # any — it inherits none. - base = set() if general_settings.get("require_key_mcp_access_defined", False) else team_set + # any — it inherits none. That ceiling is for VIRTUAL KEYS that + # can declare their own access; a keyless gateway/bridge-admitted + # user has no key to declare access on — team membership IS their + # only access path — so the flag must not zero their team grants. + # A keyless admitted subject returned above and never reaches this virtual-key ceiling, + # so require_key_mcp_access_defined can only ever zero a real key's inherited team grants. + # ``keyless_source`` marks one grant source of an admitted subject, which has no key + # to declare access on, so the flag must not zero its team grants. + require_key_access = ( + general_settings.get("require_key_mcp_access_defined", False) and not keyless_source + ) + base = team_set if not require_key_access else set() else: base = key_set & team_set # both restrict → intersect @@ -1185,24 +1465,331 @@ class MCPRequestHandler: ######################################################### # Apply org-level ceiling if org_id is set ######################################################### - if user_api_key_auth and user_api_key_auth.org_id: - allowed_mcp_servers_for_org = await MCPRequestHandler._get_allowed_mcp_servers_for_org( - user_api_key_auth - ) - if len(allowed_mcp_servers_for_org) > 0: - if has_lower_level_mcp_restrictions: - # Lower-level restrictions exist, so org can only cap them. - allowed_mcp_servers = [s for s in allowed_mcp_servers if s in allowed_mcp_servers_for_org] - else: - # No lower-level restrictions → org list becomes the ceiling - allowed_mcp_servers = allowed_mcp_servers_for_org - verbose_logger.debug(f"Applied org ceiling filter. Final allowed servers: {allowed_mcp_servers}") + allowed_mcp_servers = await MCPRequestHandler._apply_primary_org_ceiling( + allowed_mcp_servers, + user_api_key_auth, + has_lower_level_mcp_restrictions, + keyless_source=keyless_source, + ) return list(set(allowed_mcp_servers)) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}") return [] + @staticmethod + async def _apply_primary_org_ceiling( + allowed_mcp_servers: list[str], + user_api_key_auth: UserAPIKeyAuth | None, + has_lower_level_mcp_restrictions: bool, + keyless_source: bool = False, + ) -> list[str]: + """Cap the resolved server list by this caller's org ceiling. If the org names an explicit MCP + list, lower-level restrictions are intersected with it, else the org list becomes the ceiling. + No org, or an empty org list, leaves the result unchanged. + + ``keyless_source`` marks one grant source of a keyless admitted subject and governs BOTH + org divergences, because they are the same fact about that caller shape. + + First, what an UNRESOLVABLE ceiling means. A virtual key keeps the + long-standing fail-open behavior (a DB blip must not lock working keys out mid-incident). A + keyless admitted subject fails CLOSED, because its only org bound is this ceiling: silently + dropping it on a transient fault would widen a cross-org user to servers their team's org + forbids, which is a privilege escalation rather than an availability blip. + + Second, whether the org list may SUBSTITUTE for absent lower-level grants. For a key it may + (that is the key ceiling model). For a source it may only ever intersect, because the + admitted model is a union of grants and a ceiling that grants is not a ceiling.""" + if not (user_api_key_auth and user_api_key_auth.org_id): + return allowed_mcp_servers + allowed_mcp_servers_for_org = await MCPRequestHandler._get_allowed_mcp_servers_for_org(user_api_key_auth) + if allowed_mcp_servers_for_org is None: + verbose_logger.warning( + f"MCP org ceiling unresolved for org_id={user_api_key_auth.org_id!r}; " + f"{'denying (keyless admitted subject)' if keyless_source else 'leaving uncapped (key auth)'}" + ) + return [] if keyless_source else allowed_mcp_servers + if len(allowed_mcp_servers_for_org) == 0: + return allowed_mcp_servers + if has_lower_level_mcp_restrictions or keyless_source: + # Lower-level restrictions exist, so org can only cap them. + # + # A keyless admitted source ALWAYS takes this arm: its model is a union of GRANTS, so an + # org list may only narrow what a source already grants, never become one. Letting it + # substitute would hand every admitted user with an org_id that org's whole server list + # without any direct or team grant — a ceiling silently acting as a grant. + capped = [s for s in allowed_mcp_servers if s in allowed_mcp_servers_for_org] + else: + # No lower-level restrictions → org list becomes the ceiling. + capped = allowed_mcp_servers_for_org + verbose_logger.debug(f"Applied org ceiling filter. Final allowed servers: {capped}") + return capped + + @staticmethod + def _scoped_source_auth( + auth: UserAPIKeyAuth, + *, + team_id: str | None, + org_id: str | None, + carry_user_grants: bool, + ) -> UserAPIKeyAuth: + """A plain, UNMARKED auth describing ONE grant source of an admitted subject. + + Only the fields the resolver actually consults are carried. Everything else is left at its + default on purpose: ``api_key``/``token`` stay unset (this is not a key), budget, spend and + rate-limit fields stay unset because the admitted subject's own user-level limits are what + the request is metered against and cloning them per source would show the limiter N copies of + the same descriptor, and ``user_role`` stays unset because an admin role would grant every + server if this auth ever reached the server-manager wrapper. The admission marker cannot be + set through the constructor at all (a before-validator pops it), so each source is resolved + as an ordinary caller and cannot re-enter the admitted path. + """ + scoped = UserAPIKeyAuth( + user_id=auth.user_id, + team_id=team_id, + org_id=org_id, + parent_otel_span=auth.parent_otel_span, + ) + if carry_user_grants: + # The user's OWN grants. A team source deliberately carries none of these: the resolver + # loads that team's object_permission and access groups from team_id itself, and mixing + # the user's in would widen the team source with grants the team never made. + scoped.object_permission = auth.object_permission + scoped.object_permission_id = auth.object_permission_id + scoped.access_group_ids = auth.access_group_ids + return scoped + + @staticmethod + async def _admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + """The independent sources a keyless admitted subject reaches MCP servers through: their own + direct grants, plus every team they are a live roster member of. + + Each team source carries that TEAM's org as its ``org_id``, which is what makes the canonical + resolver apply the team's OWN owning-org ceiling to it — a cross-org user's teams are each + bounded by their own org rather than by the caller's home org. A team with no organization + falls back to the user's org so it is bounded rather than unbounded. + + Roster membership is checked HERE because it is a property of the source list, not of any one + resolution: a key is structurally pinned to a team it belongs to, while a user's cached + ``teams`` array can name a team whose ``members_with_roles`` no longer contains them (SCIM + group sync, or cache lag after a team_member_delete), and JWT auth can rewrite that array + outright. The roster is the source of truth for revocation. + """ + from litellm.proxy.proxy_server import prisma_client + + sources = [ + MCPRequestHandler._scoped_source_auth(auth, team_id=None, org_id=auth.org_id, carry_user_grants=True) + ] + if not auth.user_id or prisma_client is None: + return sources + for team_id in await MCPRequestHandler._resolve_user_team_ids(auth.user_id, auth): + team_obj = await MCPRequestHandler._roster_team_object(team_id, auth) + if team_obj is None: + continue + sources.append( + MCPRequestHandler._scoped_source_auth( + auth, + team_id=team_id, + org_id=team_obj.organization_id or auth.org_id, + carry_user_grants=False, + ) + ) + return sources + + @staticmethod + async def _roster_team_object(team_id: str, auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None: + """The team row for ``team_id``, but ONLY when ``auth``'s user is a live roster member of it. + + The single owner of "is this team really one of this subject's sources", so the grant union + and the per-team rate limits cannot disagree about which teams count. A team lingering in the + user's cached ``teams`` array whose ``members_with_roles`` no longer lists them returns None + here, which is what revokes both its grants and its throttle in one place.""" + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + if prisma_client is None or not auth.user_id: + return None + try: + team_obj: LiteLLM_TeamTable | None = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others + # The unit of fault isolation is the SOURCE: a team that cannot be resolved contributes + # nothing this request (fail closed for that team alone — access only ever narrows), + # while the user's own grants and every other resolvable team stand. Raising here + # instead would collapse the whole union to deny-all because one team's row was + # momentarily unreadable, on the servers, tools and throttle axes alike. + verbose_logger.warning(f"MCP admitted-subject source team {team_id!r} unresolvable, skipping: {str(e)}") + return None + if team_obj is None: + return None + member_user_ids = {getattr(m, "user_id", None) for m in (team_obj.members_with_roles or [])} - {None} + if auth.user_id not in member_user_ids: + return None + # A team over its own max budget — or owned by an org over ITS budget — is not a live + # grantor, exactly as it is not for a virtual key pinned to it (common_checks rejects that + # key outright). Enforced with the SAME owners the key path uses (_team_max_budget_check / + # _organization_max_budget_check, cross-pod Redis-first spend), targeted at the TEAM's org + # via the scoped source view, so a cross-org team is judged by its own org's budget. This is + # budget ENFORCEMENT of an already-exceeded state; ATTRIBUTION of new spend stays with the + # user (documented deferral) — the two are different questions. Sitting here, no consumer of + # the source list (servers, tools, throttle stamping) can ever see an over-budget team. + from litellm.exceptions import BudgetExceededError + from litellm.proxy.auth.auth_checks import ( + _organization_max_budget_check, + _team_max_budget_check, + ) + + source_view = MCPRequestHandler._scoped_source_auth( + auth, team_id=team_id, org_id=team_obj.organization_id or auth.org_id, carry_user_grants=False + ) + try: + await _team_max_budget_check( + team_object=team_obj, valid_token=source_view, proxy_logging_obj=proxy_logging_obj + ) + await _organization_max_budget_check( + valid_token=source_view, + team_object=team_obj, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + except BudgetExceededError as e: + verbose_logger.info(f"MCP admitted-subject source team {team_id!r} over budget, not a grantor: {str(e)}") + return None + except Exception as e: # noqa: BLE001 # per-source isolation: a budget-check fault narrows, never raises + verbose_logger.warning(f"MCP budget check failed for source team {team_id!r}, skipping source: {str(e)}") + return None + return team_obj + + @staticmethod + async def admitted_source_grants(auth: UserAPIKeyAuth) -> list[tuple[UserAPIKeyAuth, set[str]]]: + """``(source, the servers that source grants)`` for every source of an admitted subject. + + THE owner of "which source reaches which server". The reachable union, the per-team throttle + scope, the tool union and billing attribution are all just different reads of this one + answer — computing it separately per consumer is how they drift (a throttle map scoped by + roster instead of by grant charged unrelated teams' buckets).""" + return [ + (source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True))) + for source in await MCPRequestHandler._admitted_subject_sources(auth) + ] + + @staticmethod + async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]: + """Union of what each of the admitted subject's sources reaches, each answered by the + canonical resolver so no rule is reimplemented for this caller shape.""" + reachable: set[str] = set() + for _source, granted in await MCPRequestHandler.admitted_source_grants(auth): + reachable.update(granted) + return list(reachable) + + @staticmethod + async def billing_auth_for_tool_call(auth: UserAPIKeyAuth, tool_name: str) -> UserAPIKeyAuth: + """The auth object a tool call's SPEND should be recorded against. + + Returns ``auth`` unchanged for every caller that is not a keyless admitted subject, so key + and JWT billing is byte-identical. For an admitted subject whose call is reached through a + team's grant, returns a copy carrying that team's ``team_id`` and its owning ``org_id`` so + the team's budget accumulates and the correct organization is charged. + + Inert rather than wrong when the target server cannot be resolved from the tool name (a + display-name override, or a REST caller passing server_id with an unprefixed name): billing + then falls back to today's user-level attribution instead of guessing a team. Resolution + reuses the manager's own tool-name lookup rather than re-deriving prefix rules that live + there.""" + if not _is_mcp_admitted_user_subject(auth): + return auth + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + server = global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name) + if server is None: + return auth + source = await MCPRequestHandler.attributing_source_for_server(auth, server.server_id) + if source is None or not source.team_id: + return auth + billed = auth.model_copy() + billed.team_id = source.team_id + billed.org_id = source.org_id + return billed + except Exception as e: # noqa: BLE001 # attribution must never fail an authorized call + verbose_logger.warning(f"MCP billing attribution failed for {tool_name!r}, billing the user: {str(e)}") + return auth + + @staticmethod + async def attributing_source_for_server( + auth: UserAPIKeyAuth, + server_id: str, + source_grants: list[tuple[UserAPIKeyAuth, set[str]]] | None = None, + ) -> UserAPIKeyAuth | None: + """The source a billable call to ``server_id`` is attributed to, or None to bill the caller + as themselves (their own grant reaches it, or nothing does). + + A keyless admitted subject carries no ``team_id``, so downstream spend skipped team updates + entirely and charged the user's PRIMARY org — a team-derived call neither accumulated its + team's budget (so that budget could never begin to block) nor charged the org that owns the + granting team. Attribution restores both. + + The rule: a user's OWN grant is not "through a team", so it bills the user. Otherwise the + call is billed to a granting team — deterministically the lowest ``team_id`` when several + grant the same server, so the choice is stable, reproducible and auditable rather than + dependent on dict ordering. Reads the one grant owner, so the team that gets billed is + always a team that actually granted the server.""" + source_grants = source_grants or await MCPRequestHandler.admitted_source_grants(auth) + granting = [(source, granted) for source, granted in source_grants if server_id in granted] + if not granting: + return None + for source, _granted in granting: + if source.team_id is None: + return None # the user's own grant reaches it: their spend, their org + return min((source for source, _ in granting), key=lambda s: s.team_id or "") + + @staticmethod + async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None: + """Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the + sources that actually grant that server. + + A source that does not grant the server contributes nothing, so its tool rules cannot leak + onto a server reached through a different source. A source that grants the server with no + tool restriction means the user can use every tool on it, so allow-all wins the union. When + no source grants the server the result is ``[]`` — deny all, fail closed.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + # An OPEN channel (operator-opened allow_all_keys, the user's own BYOM submission) makes the + # server REACHABLE through the user themselves — no grant source names it, so without this + # the union below would return [] and leave it listable but uninvokable. Reachability is ALL + # it confers: it is not a waiver of the ceilings that bound the server. The user's own + # mcp_tool_permissions and their org's tool ceiling still bind, which is what a virtual key + # on the same allow_all server gets (its key_tools and _apply_agent_and_org_tool_ceilings + # both run). Returning None here instead skipped both and let a session holder invoke tools + # their own or their org's policy excludes. + reachable_via_open_channel = server_id in await global_mcp_server_manager.operator_open_server_ids(auth) + + allowed: set[str] = set() + for source, granted in await MCPRequestHandler.admitted_source_grants(auth): + # The open channel is evaluated against the user's OWN source (team_id is None), so that + # source's restrictions apply to it; a team's rules never ride an open-channel server. + if server_id not in granted and not (reachable_via_open_channel and source.team_id is None): + continue + tools = await MCPRequestHandler.get_allowed_tools_for_server(server_id, source, keyless_source=True) + if tools is None: + return None + allowed.update(tools) + return sorted(allowed) + @staticmethod def _get_key_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, @@ -1262,6 +1849,8 @@ class MCPRequestHandler: async def get_allowed_tools_for_server( server_id: str, user_api_key_auth: Optional[UserAPIKeyAuth] = None, + *, + keyless_source: bool = False, ) -> Optional[List[str]]: """ Get list of allowed tool names for a specific server based on key/team permissions. @@ -1278,6 +1867,15 @@ class MCPRequestHandler: return None try: + # FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject is + # resolved per grant source and shares NOTHING with the single-credential prelude below. + # Ordering is the invariant, not a nicety — when this branch sat after the prelude, a + # fault in a lookup the subject never uses (its own mcp_toolsets, its team_obj_perm) hit + # the fail-closed handler and denied tools its teams did grant. Nothing that resolves a + # single credential's scope may run before this line. + if _is_mcp_admitted_user_subject(user_api_key_auth): + return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth) + # Get key and team object permissions (already loaded in main auth flow) key_obj_perm = MCPRequestHandler._get_key_object_permission(user_api_key_auth) team_obj_perm = await MCPRequestHandler._get_team_object_permission(user_api_key_auth) @@ -1319,6 +1917,9 @@ class MCPRequestHandler: else None ) + # A keyless gateway/bridge-admitted user has no single team_id, so team_obj_perm above is + # None and the single-team lookup yields allow-all — silently dropping every team's + # per-server tool exclusions. Resolve it as the union over the sources that grant the # Apply same inheritance logic as get_allowed_mcp_servers if team_tools: if key_tools: @@ -1331,42 +1932,82 @@ class MCPRequestHandler: # No team restrictions → use key restrictions allowed_tools = cast(List[str], key_tools) - # Intersect with agent's tool permissions if agent_id is set - if user_api_key_auth.agent_id: - # Pre-fetch agent object_permission once to avoid duplicate DB query - agent_obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) - agent_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( - server_id=server_id, - user_api_key_auth=user_api_key_auth, - agent_object_permission=agent_obj_perm, - ) - if agent_tools is not None: - if allowed_tools is not None: - allowed_tools = list(set(allowed_tools) & set(agent_tools)) - else: - allowed_tools = agent_tools - - # Apply org-level tool ceiling if org_id is set - if user_api_key_auth.org_id: - # _get_org_object_permission uses user_api_key_cache, so this is not a - # fresh DB round-trip when get_allowed_mcp_servers was already called. - org_obj_perm = await MCPRequestHandler._get_org_object_permission(user_api_key_auth) - org_tools = ( - global_mcp_server_manager.expand_tool_permissions(org_obj_perm.mcp_tool_permissions).get(server_id) - if org_obj_perm and org_obj_perm.mcp_tool_permissions - else None - ) - if org_tools is not None: - if allowed_tools is not None: - allowed_tools = list(set(allowed_tools) & set(org_tools)) - else: - allowed_tools = list(org_tools) - - return allowed_tools + return await MCPRequestHandler._apply_agent_and_org_tool_ceilings( + allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source + ) except Exception as e: verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}") - return None + # Fail CLOSED for a keyless admitted subject: ANY error resolving the tool allowlist + # (multi-team fan-out, org/agent lookups) must deny the server's tools ([]) for this + # request rather than collapse to allow-all (None), mirroring the fail-closed server + # path. Key/JWT auth keeps its prior allow-all-on-error behavior. + # + # keyless_source matters as much as the marker: each source of an admitted subject is + # resolved through an UNMARKED auth, so without it a fault under a source returned None, + # and None wins the union as allow-all — dropping every team and org tool ceiling on a + # blip. The marker alone only covers a fault raised before the fan-out. + return [] if (keyless_source or _is_mcp_admitted_user_subject(user_api_key_auth)) else None + + @staticmethod + async def _apply_agent_and_org_tool_ceilings( + allowed_tools: list[str] | None, + server_id: str, + user_api_key_auth: UserAPIKeyAuth, + keyless_source: bool = False, + ) -> list[str] | None: + """Narrow a key/team tool allowlist by the agent's tool permissions and the caller's org tool + ceiling. Each level only ever intersects, and None at a level means "no restriction from this + level". + + An UNRESOLVABLE org ceiling (``_get_org_object_permission`` raises: the org names a permission + that cannot be loaded) is decided here, per caller shape, mirroring the servers axis: a + virtual key keeps its long-standing fail-open — the org step is skipped and the key/team/agent + restrictions already computed STAND (letting the raise escape would collapse them to + allow-all, which is fail-open WIDER than before the fault). A keyless admitted source + re-raises, and the outer handler denies tools for that one source while the subject's other + sources stand — its only org bound is this ceiling, so skipping it would widen access.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + if user_api_key_auth.agent_id: + # Pre-fetch agent object_permission once to avoid a duplicate DB query. + agent_obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + agent_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( + server_id=server_id, + user_api_key_auth=user_api_key_auth, + agent_object_permission=agent_obj_perm, + ) + if agent_tools is not None: + allowed_tools = ( + list(set(allowed_tools) & set(agent_tools)) if allowed_tools is not None else agent_tools + ) + + if user_api_key_auth.org_id: + # _get_org_object_permission uses user_api_key_cache, so this is not a fresh DB round-trip + # when get_allowed_mcp_servers was already called. + try: + org_obj_perm = await MCPRequestHandler._get_org_object_permission(user_api_key_auth) + except Exception as e: # noqa: BLE001 # unresolvable org ceiling, decided per caller shape + if keyless_source: + raise + verbose_logger.warning( + f"MCP org tool ceiling unresolvable for org_id={user_api_key_auth.org_id!r}; " + f"skipping org intersect, key/team/agent restrictions stand: {str(e)}" + ) + return allowed_tools + org_tools = ( + global_mcp_server_manager.expand_tool_permissions(org_obj_perm.mcp_tool_permissions).get(server_id) + if org_obj_perm and org_obj_perm.mcp_tool_permissions + else None + ) + if org_tools is not None: + allowed_tools = ( + list(set(allowed_tools) & set(org_tools)) if allowed_tools is not None else list(org_tools) + ) + + return allowed_tools @staticmethod async def is_tool_allowed_for_server( @@ -1537,10 +2178,96 @@ class MCPRequestHandler: @staticmethod async def _get_allowed_mcp_servers_for_team( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: + """Get allowed MCP servers a caller inherits from the team it is pinned to. + + Exactly one team, or none. A subject that reaches servers through SEVERAL teams does not + fan out here: it is resolved one source per team in ``_resolve_admitted_subject_servers``, + and each of those sources pins a single ``team_id`` before reaching this point. Keeping the + fan-out here as well would be a second multi-team path to drift from that one. """ - Get allowed MCP servers for a team. + team_ids = await MCPRequestHandler._team_ids_for_mcp_grant(user_api_key_auth) + if not team_ids: + return [] + return await MCPRequestHandler._allowed_mcp_servers_for_single_team(team_ids[0], user_api_key_auth) + + @staticmethod + async def _team_ids_for_mcp_grant(user_api_key_auth: UserAPIKeyAuth | None) -> list[str]: + """The team ids whose MCP grants a caller inherits. + + A caller with an explicit ``team_id`` uses that single team; every other caller inherits no + team grants. That covers key auth and JWT auth (a keyless ``user_id`` auth with no team_id, + which must NOT silently gain the union across every team the user belongs to), and it covers + each single-source auth an admitted subject fans out into — those pin a team_id, so they land + on the first branch. The admitted subject itself never reaches here: it resolves per source + in ``_resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel + resolves to no teams exactly as before.""" + if user_api_key_auth is None or not user_api_key_auth.team_id: + return [] + return [] if user_api_key_auth.team_id == UI_TEAM_ID else [user_api_key_auth.team_id] + + @staticmethod + async def _resolve_user_team_ids(user_id: str, user_api_key_auth: UserAPIKeyAuth) -> list[str]: + """The distinct team ids a user belongs to, from the live user record. Returns [] on + no DB, a missing user, or any resolution failure so a lookup blip narrows access + rather than raising; the caller's direct grants still apply.""" + from litellm.proxy.auth.auth_checks import get_user_object + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + if prisma_client is None: + return [] + try: + user_object = await get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + parent_otel_span=user_api_key_auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises + verbose_logger.warning(f"Failed to resolve user teams for MCP grant: {str(e)}") + return [] + if user_object is None or not user_object.teams: + return [] + return list(dict.fromkeys(t for t in user_object.teams if t and t != UI_TEAM_ID)) + + @staticmethod + async def _team_granted_servers(team_obj: LiteLLM_TeamTable, team_access_group_servers: list[str]) -> set[str]: + """The raw MCP-server set a team grants (before any org ceiling): its object_permission (direct + ``mcp_servers``, the ``all_proxy_servers`` sentinel → the full registry, legacy access groups, + tool-perm-referenced servers) unioned with its unified ``access_group_ids`` servers.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + object_permissions = team_obj.object_permission + if object_permissions is None: + return set(team_access_group_servers) + if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []): + return set(global_mcp_server_manager.get_registry().keys()) + legacy_access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups( + object_permissions.mcp_access_groups or [] + ) + return ( + set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])) + | set(legacy_access_group_servers) + | set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()) + | set(team_access_group_servers) + ) + + @staticmethod + async def _allowed_mcp_servers_for_single_team( + team_id: str, + user_api_key_auth: UserAPIKeyAuth | None, + ) -> list[str]: + """Allowed MCP servers granted by ONE team (its raw grant, then capped by the team's own org + for a keyless admitted subject). Unions two sources: - Legacy team.object_permission (mcp_servers, mcp_access_groups, @@ -1551,9 +2278,6 @@ class MCPRequestHandler: the gate (no assigned_team_ids check needed here). """ try: - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) from litellm.proxy.auth.auth_checks import ( _get_mcp_server_ids_from_access_groups, get_team_object, @@ -1564,22 +2288,24 @@ class MCPRequestHandler: user_api_key_cache, ) - if user_api_key_auth is None or not user_api_key_auth.team_id or prisma_client is None: + if not team_id or team_id == UI_TEAM_ID or prisma_client is None: return [] - if user_api_key_auth.team_id == UI_TEAM_ID: - return [] - - team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( - team_id=user_api_key_auth.team_id, + parent_otel_span = user_api_key_auth.parent_otel_span if user_api_key_auth is not None else None + team_obj: LiteLLM_TeamTable | None = await get_team_object( + team_id=team_id, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, + parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) if team_obj is None: return [] - + if team_obj.blocked: + # A blocked team grants nothing. The central policy gate enforces this for a key + # pinned to a single team_id, but a keyless admitted identity (no team_id) unions + # across all of its teams and would otherwise inherit a blocked team's MCP grants. + return [] team_access_group_servers = await _get_mcp_server_ids_from_access_groups( access_group_ids=team_obj.access_group_ids or [], prisma_client=prisma_client, @@ -1587,27 +2313,8 @@ class MCPRequestHandler: proxy_logging_obj=proxy_logging_obj, ) - object_permissions = team_obj.object_permission - if object_permissions is None: - return list(set(team_access_group_servers)) - - if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []): - return list(global_mcp_server_manager.get_registry().keys()) - - direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []) - - legacy_access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups( - object_permissions.mcp_access_groups or [] - ) - - tool_perm_servers = list( - global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys() - ) - - all_servers = ( - direct_mcp_servers + legacy_access_group_servers + tool_perm_servers + team_access_group_servers - ) - return list(set(all_servers)) + servers = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers) + return list(servers) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers for team: {str(e)}") return [] @@ -1621,7 +2328,11 @@ class MCPRequestHandler: ``get_object_permission`` helpers so MCP requests share the same ``user_api_key_cache`` entries as the rest of the proxy. """ - from litellm.proxy.auth.auth_checks import get_object_permission, get_org_object + from litellm.proxy.auth.auth_checks import ( + OrganizationNotFoundError, + get_object_permission, + get_org_object, + ) from litellm.proxy.proxy_server import ( prisma_client, proxy_logging_obj, @@ -1635,6 +2346,9 @@ class MCPRequestHandler: verbose_logger.debug("prisma_client is None") return None + # An ABSENT org is a determinate fact, not a failure: a team's organization_id can point at a + # row that was deleted or has not synced yet, and get_org_object raises for that. It places no + # ceiling, exactly as a key with a dangling org_id is not locked out. try: org_obj = await get_org_object( org_id=user_api_key_auth.org_id, @@ -1643,21 +2357,35 @@ class MCPRequestHandler: parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) - - if org_obj is None or not org_obj.object_permission_id: - return None - - return await get_object_permission( - object_permission_id=org_obj.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - except Exception as e: - verbose_logger.warning(f"Failed to get org object permission: {str(e)}") + except OrganizationNotFoundError as e: + # CONFIRMED absent (deleted org, not-yet-synced organization_id): a determinate fact, so + # it places no ceiling. Every OTHER exception is an operational failure and propagates — + # caught upstream as an unresolvable ceiling, which denies for a keyless source and stays + # fail-open for a key. Catching bare Exception here treated a DB outage as "no org", which + # silently dropped a real org's ceiling for exactly as long as the outage lasted. + verbose_logger.debug(f"MCP org ceiling: org {user_api_key_auth.org_id!r} does not exist: {e}") return None + if org_obj is None or not org_obj.object_permission_id: + return None + + # From here the org NAMES a permission. Failing to read it is INDETERMINATE, so it must not + # collapse into the same None that means "no ceiling" -- that is what would silently drop a + # real ceiling on a transient fault. Raise and let each caller pick fail-open or fail-closed. + object_permission = await get_object_permission( + object_permission_id=org_obj.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=user_api_key_auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + if object_permission is None: + raise ValueError( + f"org {user_api_key_auth.org_id!r} names object_permission_id " + f"{org_obj.object_permission_id!r} which could not be loaded" + ) + return object_permission + @staticmethod async def _get_allowed_mcp_servers_for_org( user_api_key_auth: Optional[UserAPIKeyAuth] = None, @@ -1692,8 +2420,11 @@ class MCPRequestHandler: all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers return list(set(all_servers)) except Exception as e: + # None = the org ceiling could NOT be resolved, which is not the same fact as [] = the + # org places no restriction. Collapsing the two is what let a transient DB fault silently + # remove an org's ceiling; the caller picks fail-open or fail-closed from this signal. verbose_logger.warning(f"Failed to get allowed MCP servers for org: {str(e)}") - return [] + return None @staticmethod async def _get_allowed_mcp_servers_for_end_user( diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ee3196b539c..6fe4cc64d12 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -33,6 +33,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( _finish_bridge_mint, _prepare_bridge_mint, _prepare_bridge_refresh, + _reload_active_user_by_id, ) from litellm.proxy._experimental.mcp_server.faults import ( CallerRejected, @@ -43,6 +44,14 @@ from litellm.proxy._experimental.mcp_server.faults import ( dcr_fault_detail, render_token_fault, ) +from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( + aggregate_authorize, + aggregate_token, + complete_connect_flow, + is_gateway_dcr_client_id, + register_aggregate_client, + relative_request_url, +) from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, get_request_base_url, @@ -324,14 +333,25 @@ def redeem_passthrough_authorization_code( return sealed +def _session_cookie_user_id(request: Request) -> str | None: + """The signed-in litellm user for a browser request, or ``None``. Thin wrapper so the + aggregate DCR flow's verbs receive the identity as a plain value instead of parsing + cookies themselves.""" + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # circular import at module load + _user_id_from_session_cookie, + ) + + return _user_id_from_session_cookie(request) + + def _redirect_to_litellm_login(request: Request) -> RedirectResponse: """Send an unauthenticated browser through litellm login before the interactive bridge authorize can capture its identity. The bridge oauth_delegate flow seals the SSO user into the gateway code, - so a session is required; without one there is nothing to bind. After login the user re-initiates - the connection, which then finds the session cookie (the seamless return-to round-trip, which is - origin-validated against the control-plane URL, is a follow-up).""" + so a session is required; without one there is nothing to bind. A same-origin relative + ``return_to`` (honored by the SSO callback) brings the browser straight back to this authorize + request after login instead of stranding it on the dashboard.""" base_url = get_request_base_url(request) - return RedirectResponse(f"{base_url}/sso/key/generate") + return RedirectResponse(f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}") # LIT-4197: some upstream authorization servers reject an over-long ``state`` @@ -1601,6 +1621,18 @@ async def authorize( global_mcp_server_manager, ) + if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id): + return aggregate_authorize( + request=request, + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=code_challenge, + code_challenge_method=code_challenge_method, + response_type=response_type, + session_user_id=_session_cookie_user_id(request), + ) + lookup_name: Optional[str] = mcp_server_name or client_id client_ip = IPAddressUtils.get_mcp_client_ip(request) mcp_server = ( @@ -1664,6 +1696,25 @@ async def token_endpoint( global_mcp_server_manager, ) + if mcp_server_name is None and is_gateway_dcr_client_id(client_id): + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load + master_key, + user_api_key_cache, + ) + + return await aggregate_token( + request=request, + grant_type=grant_type, + code=code, + redirect_uri=redirect_uri, + client_id=client_id, + code_verifier=code_verifier, + refresh_token=refresh_token, + master_key=master_key, + reload_user=_reload_active_user_by_id, + cache=user_api_key_cache, + ) + lookup_name = mcp_server_name or client_id client_ip = IPAddressUtils.get_mcp_client_ip(request) mcp_server = global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) @@ -1685,6 +1736,21 @@ async def token_endpoint( ) +@router.post("/authorize/complete") +async def authorize_complete(request: Request, flow: str = Form(...)): + """Finish an aggregate connect flow: mint the gateway authorization code for the + signed-in user and redirect back to the DCR client. POST plus the per-flow HttpOnly + cookie set at /authorize; an anonymous or bad-flow request just 400s.""" + from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # circular import at module load + + return await complete_connect_flow( + request=request, + flow_handle=flow, + session_user_id=_session_cookie_user_id(request), + cache=user_api_key_cache, + ) + + # Per RFC 6749 §4.1.2.1, an IdP that rejects an OAuth authorization request # redirects back to the configured redirect URI with ``error`` / # ``error_description`` / ``error_uri`` query params and no ``code``. The MCP @@ -2422,6 +2488,13 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non } client_ip = IPAddressUtils.get_mcp_client_ip(request) if not mcp_server_name: + # A real DCR request carries redirect_uris (RFC 7591): route it to the aggregate DCR + # endpoint the aggregate authorization-server metadata advertises. A single-server + # deployment registers at /{server}/register instead (its bare-origin discovery + # advertises that), so this does not affect it. A request without redirect_uris is not + # a DCR request, so the legacy single-server-or-dummy fallback is kept for it. + if data.get("redirect_uris"): + return await register_aggregate_client(request=request, request_body=data) resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) if resolved: return await register_client_with_server( diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py new file mode 100644 index 00000000000..58233c4c9e5 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -0,0 +1,637 @@ +"""The gateway-level DCR flow for the aggregate ``/mcp`` endpoint (``mcp_gateway_dcr``). + +An OAuth-only DCR client (Claude Desktop, Claude Code, MCP Inspector) pointed at the +aggregate ``/mcp`` endpoint discovers the gateway as its authorization server (PR 1 of +this track) and then walks the flow implemented here: + +1. ``POST /register``: stateless dynamic client registration. The ``client_id`` IS the + registration: the client's redirect URIs are sealed into it with the repo's + authenticated symmetric helper, so nothing is persisted and a forged or tampered + client_id simply fails to open. Clients are always public (``token_endpoint_auth_method + "none"``); PKCE S256 is what protects the code. +2. ``GET /authorize``: validates the client and redirect URI, requires S256 PKCE, and + interposes LiteLLM sign-in. Without a session cookie the browser is sent through + ``/sso/key/generate`` with a same-origin ``return_to`` so it lands back here after + login. With a session, the flow parameters and the SSO user are sealed into a per-flow + HttpOnly cookie (the same pattern as the upstream OAuth state relay) and the browser is + sent to the connect page, where the user authorizes individual servers (vaulting those + tokens server-side) before finishing. +3. ``POST /authorize/complete``: the deliberate finish step. A POST (not GET) bound to the + SameSite=Lax flow cookie, so a cross-site link cannot silently mint a code with the + victim's session, and the signed-in user must match the user sealed into the flow. + Mints a short-lived, single-use, gateway-sealed authorization code and redirects to the + client's registered redirect URI. +4. ``POST /token``: exchanges the code (PKCE-verified, client- and redirect-bound, + single-use) for the identity-only session tokens of + :mod:`.outbound_credentials.session_token`, re-validating that the litellm user is + still active first; the ``refresh_token`` grant rotates the pair the same way. + +Nothing here stores state server-side except the single-use code guard (a TTL cache +entry). Every sealed value is authenticated encryption over the proxy salt/master key +family, opened totally (bad input maps to an OAuth error, never a raise), and every +identity is a stable reference re-validated live at mint, refresh, and (in the admission +PR) tool-call time. Upstream server credentials never appear anywhere in this flow; they +are vaulted per user by the existing ``/v1/mcp`` authorize endpoints and resolved at +egress by user id. +""" + +from __future__ import annotations + +import hashlib +import hmac +import secrets +from base64 import urlsafe_b64encode +from collections.abc import Mapping +from datetime import datetime, timezone +from typing import Awaitable, Callable, Literal, TypeVar +from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse + +from fastapi import HTTPException, Request +from fastapi.responses import JSONResponse, RedirectResponse, Response +from pydantic import BaseModel, ConfigDict, Field, ValidationError +from typing_extensions import assert_never + +from litellm._logging import verbose_logger +from litellm.caching.caching import DualCache +from litellm.proxy._experimental.mcp_server.oauth_utils import ( + TOKEN_NO_CACHE_HEADERS, + get_request_base_url, + is_loopback_redirect_host, + validate_redirect_uri_shape, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + SessionRefreshOpened, + open_session_refresh_bearer, + session_keys_from_master_key, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + SESSION_REFRESH_TTL_SECONDS, + MintedSessionToken, + SessionKeys, + SessionPrincipal, + mint_session_refresh_token, + mint_session_token, +) +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) + +GATEWAY_DCR_CLIENT_ID_PREFIX = "llm_dcrc_" +"""Marker prefix on every gateway-issued DCR client_id so the root authorize/token +endpoints can route an aggregate-flow request without decrypting, and existing per-server +flows (whose client_ids are upstream-issued) are never captured by the aggregate arm.""" + +GATEWAY_AUTH_CODE_PREFIX = "llm_gcode_" +"""Marker prefix on the gateway-sealed authorization code, distinct from the bridge +``llm_bcode_`` so neither flow can consume the other's codes.""" + +CONNECT_FLOW_COOKIE_PREFIX = "mcp_connect_flow_" +"""Per-flow HttpOnly cookie holding the sealed connect flow, keyed by a short random +handle carried in the connect-page URL (the same handle-plus-cookie pattern as the +``mcp_oauth_state_`` upstream relay, for the same reasons: replica-safe with no +server-side session store, and the sealed value never appears in a URL).""" + +CONNECT_FLOW_TTL_SECONDS = 600 +GATEWAY_AUTH_CODE_TTL_SECONDS = 120 +_CLAIM_TTL_BUFFER_SECONDS = 60 +_USED_CODE_CACHE_PREFIX = "mcp_gateway_dcr_code_used:" +_USED_FLOW_CACHE_PREFIX = "mcp_gateway_dcr_flow_used:" +_USED_REFRESH_CACHE_PREFIX = "mcp_gateway_dcr_refresh_used:" + +MAX_REDIRECT_URIS = 3 +MAX_REDIRECT_URI_LENGTH = 256 +MAX_CLIENT_ID_LENGTH = 2048 +"""Registration bounds. They exist to bound the sealed client_id, which rides inside +every session-token claim set: 3 URIs of 256 bytes seal to roughly 1.2KB, comfortably +under this cap and under the session token's own 4KB ceiling. Claude Desktop and MCP +Inspector register one or two redirect URIs.""" + +MAX_STATE_LENGTH = 1024 +"""Bound on the client ``state`` sealed into the flow cookie and echoed on the auth-code +redirect. An unbounded ``state`` can push the sealed cookie past the browser's ~4KB cap +(silently dropped, breaking the flow); spec clients send a short opaque value.""" + +MIN_CODE_VERIFIER_LENGTH = 43 +MAX_CODE_VERIFIER_LENGTH = 128 +"""RFC 7636 section 4.1 bounds for the PKCE ``code_verifier``. Enforced so an out-of-range +verifier gets a clean ``invalid_request`` instead of an opaque PKCE-mismatch.""" + +_UNPREFIXED = "" +"""Prefix for a sealed value that carries no wire marker because it is never routed by +prefix (the connect flow lives only in its own per-handle cookie, opened by that one +handle). Named so the empty-string argument to ``_seal`` / ``_open_sealed`` reads as +deliberate rather than a typo.""" + +_CLIENT_RECORD_DEBUG_KEY = "gateway_dcr_client" +_CONNECT_FLOW_DEBUG_KEY = "gateway_connect_flow" +_AUTH_CODE_DEBUG_KEY = "gateway_authorization_code" + +ReloadUserFailure = Literal["unresolvable", "unavailable", "no_active_key"] +ReloadUser = Callable[[str], Awaitable[ReloadUserFailure | None]] +"""Injected live-user revalidation (the token endpoint's mirror of admission): +``None`` means the user is active; ``unavailable`` is a retryable DB outage; anything +else fails the grant closed.""" + + +class GatewayDcrClient(BaseModel): + """The registration record sealed into a gateway DCR ``client_id``. + + ``extra="forbid"`` so a sealed value of another type (an auth code, a connect flow) + that happened to decrypt under the shared key can never validate as a client record: + cross-type confusion is rejected at the model boundary, not left to differing required + fields.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + redirect_uris: tuple[str, ...] = Field(min_length=1, max_length=MAX_REDIRECT_URIS) + iat: int + + +class _ConnectFlow(BaseModel): + """One in-flight authorize: the SSO user it belongs to and the client parameters + needed to mint the code at the finish step. Sealed into the per-flow cookie. ``jti`` + makes the flow single-use at complete; ``extra="forbid"`` rejects cross-type + confusion.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + user_id: str = Field(min_length=1) + client_id: str = Field(min_length=1) + redirect_uri: str = Field(min_length=1) + state: str + code_challenge: str = Field(min_length=1) + jti: str = Field(min_length=1) + exp: int + + +class _GatewayAuthCode(BaseModel): + """The gateway-sealed authorization code: the user consent it represents and the + bindings the token endpoint must verify (client, redirect URI, PKCE challenge), + plus a ``jti`` for the single-use guard. ``extra="forbid"`` rejects cross-type + confusion.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + user_id: str = Field(min_length=1) + client_id: str = Field(min_length=1) + redirect_uri: str = Field(min_length=1) + code_challenge: str = Field(min_length=1) + jti: str = Field(min_length=1) + iat: int + exp: int + + +def is_gateway_dcr_client_id(client_id: str | None) -> bool: + """Cheap prefix routing test so the root endpoints only enter the aggregate arm for + clients this flow registered; every other client_id keeps today's behavior.""" + return client_id is not None and client_id.startswith(GATEWAY_DCR_CLIENT_ID_PREFIX) + + +def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse: + """RFC 6749 section 5.2 / RFC 7591 section 3.2.2 error body. Descriptions carry no + token, code, or URL material so they are safe to relay to any client.""" + return JSONResponse( + status_code=status_code, + content={"error": error, "error_description": description}, + headers=TOKEN_NO_CACHE_HEADERS, + ) + + +def _seal(prefix: str, payload: BaseModel) -> str: + return prefix + encrypt_value_helper(payload.model_dump_json()) + + +_SealedModelT = TypeVar("_SealedModelT", bound=BaseModel) + + +def _open_sealed(value: str, prefix: str, model: type[_SealedModelT], debug_key: str) -> _SealedModelT | None: + """Open a sealed value totally: anything that is not prefix-shaped, does not decrypt, + or does not validate returns ``None`` for the caller to map onto an OAuth error.""" + if not value.startswith(prefix): + return None + decrypted = decrypt_value_helper(value[len(prefix) :], debug_key, return_original_value=False) + if not isinstance(decrypted, str): + return None + try: + return model.model_validate_json(decrypted) + except ValidationError: + return None + + +def open_gateway_dcr_client(client_id: str) -> GatewayDcrClient | None: + return _open_sealed(client_id, GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient, _CLIENT_RECORD_DEBUG_KEY) + + +async def register_aggregate_client(request: Request, request_body: Mapping[str, object]) -> Response: + """RFC 7591 dynamic registration against the gateway itself, statelessly. + + Only ``redirect_uris`` is authoritative; every client is registered as a public + ``token_endpoint_auth_method "none"`` client regardless of what it asked for (RFC + 7591 lets the server override metadata), because the gateway never issues client + secrets: possession of a secret would add nothing over the mandatory S256 PKCE, and a + stateless registration has nowhere to keep one. Nothing is persisted, so open + registration cannot be used to fill storage. + + Redirect-URI *hygiene* is not decided here: :func:`validate_redirect_uri_shape` is + the single owner of that rule across the MCP OAuth surface, so allowlisted native + callbacks (``cursor://``) are accepted and fragments, missing hosts, userinfo + (``https://claude.ai@attacker.example/cb``) and backslash hosts are rejected exactly + as they are on /authorize and /callback. + + What this endpoint does decide is its own trust policy, which is deliberately wider + than :func:`validate_trusted_redirect_uri`'s: registration is *public*, so any https + client may register (that is what lets a hosted MCP client register at all), and the + controls are mandatory S256 PKCE plus the consent screen showing the client origin. + http is confined to loopback per RFC 8252 section 7.3. + """ + raw_uris = request_body.get("redirect_uris") + if not isinstance(raw_uris, list) or not raw_uris or len(raw_uris) > MAX_REDIRECT_URIS: + return _oauth_error( + 400, + "invalid_redirect_uri", + f"redirect_uris must be a list of 1 to {MAX_REDIRECT_URIS} URIs", + ) + if not all(isinstance(uri, str) and len(uri) <= MAX_REDIRECT_URI_LENGTH for uri in raw_uris): + return _oauth_error( + 400, + "invalid_redirect_uri", + f"each redirect URI must be a string of at most {MAX_REDIRECT_URI_LENGTH} characters", + ) + for uri in raw_uris: + parsed = urlparse(uri) + try: + if validate_redirect_uri_shape(parsed): + continue # allowlisted native callback, e.g. cursor:// + except HTTPException as exc: + # The shared validator speaks HTTP; RFC 7591 registration answers with an OAuth + # error object, so translate the shape without re-deciding the rule. + return _oauth_error(400, "invalid_redirect_uri", str(exc.detail)) + if parsed.scheme == "https" or (parsed.scheme == "http" and is_loopback_redirect_host(parsed)): + continue + return _oauth_error( + 400, + "invalid_redirect_uri", + "each redirect URI must be https, http on a loopback host, or a registered native callback", + ) + now = datetime.now(timezone.utc) + client_id = _seal( + GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient(redirect_uris=tuple(raw_uris), iat=int(now.timestamp())) + ) + if len(client_id) > MAX_CLIENT_ID_LENGTH: + return _oauth_error(400, "invalid_client_metadata", "registered metadata is too large") + return JSONResponse( + status_code=201, + content={ + "client_id": client_id, + "client_id_issued_at": int(now.timestamp()), + "redirect_uris": list(raw_uris), + "token_endpoint_auth_method": "none", + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + }, + ) + + +def _flow_cookie_name(handle: str) -> str: + return f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}" + + +def _cookie_path_and_secure(request: Request) -> tuple[str, bool]: + parsed = urlparse(get_request_base_url(request)) + return parsed.path or "/", parsed.scheme == "https" + + +def _append_query_params(url: str, params: dict[str, str]) -> str: + parsed = urlparse(url) + query = parse_qsl(parsed.query, keep_blank_values=True) + list(params.items()) + return urlunparse(parsed._replace(query=urlencode(query))) + + +def relative_request_url(request: Request) -> str: + """The request's own path and query as a same-origin ``return_to`` target for the + login round-trip; relative by construction, so it can never leave the gateway.""" + path = request.url.path + return f"{path}?{request.url.query}" if request.url.query else path + + +def aggregate_authorize( + request: Request, + client_id: str, + redirect_uri: str, + state: str, + code_challenge: str | None, + code_challenge_method: str | None, + response_type: str | None, + session_user_id: str | None, +) -> Response: + """The aggregate authorize verb: validate the client, require S256 PKCE, interpose + LiteLLM sign-in, and hand the browser to the connect page with the flow sealed into a + per-flow cookie. + + Validation failures respond directly with 400 and never redirect: per RFC 6749 + section 4.1.2.1 an unvalidated redirect URI must not receive an error redirect, and + once the client is at fault there is no trusted place to send the browser. + """ + client = open_gateway_dcr_client(client_id) + if client is None: + return _oauth_error(400, "invalid_client", "unknown or malformed client_id") + if redirect_uri not in client.redirect_uris: + return _oauth_error(400, "invalid_request", "redirect_uri is not registered for this client") + if response_type != "code": + return _oauth_error(400, "unsupported_response_type", "response_type must be 'code'") + if not code_challenge or code_challenge_method != "S256": + return _oauth_error( + 400, + "invalid_request", + "PKCE is required: send code_challenge with code_challenge_method=S256", + ) + if len(state) > MAX_STATE_LENGTH: + return _oauth_error(400, "invalid_request", f"state must be at most {MAX_STATE_LENGTH} characters") + base_url = get_request_base_url(request) + if session_user_id is None: + login_url = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}" + return RedirectResponse(login_url, status_code=303) + now = datetime.now(timezone.utc) + handle = secrets.token_urlsafe(24) + flow = _ConnectFlow( + user_id=session_user_id, + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=code_challenge, + jti=secrets.token_urlsafe(24), + exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS, + ) + connect_url = _append_query_params( + f"{base_url}/ui/chat/integrations", + {"connect_flow": handle, "connect_client": _origin_only(redirect_uri)}, + ) + response = RedirectResponse(connect_url, status_code=303) + path, secure = _cookie_path_and_secure(request) + response.set_cookie( + key=_flow_cookie_name(handle), + value=_seal(_UNPREFIXED, flow), + max_age=CONNECT_FLOW_TTL_SECONDS, + path=path, + secure=secure, + httponly=True, + samesite="lax", + ) + return response + + +def _origin_only(url: str) -> str: + """Scheme+host for display on the connect page; never the full redirect URI, whose + path or query could carry values that do not belong in a page URL or logs.""" + parsed = urlparse(url) + return f"{parsed.scheme}://{parsed.netloc}" if parsed.netloc else "" + + +async def complete_connect_flow( + request: Request, + flow_handle: str, + session_user_id: str | None, + cache: DualCache, +) -> Response: + """The deliberate finish step of the connect flow: mint the gateway authorization + code and send the browser back to the client. + + Reached by POST so a cross-site GET cannot trigger it, and bound to the HttpOnly + per-flow cookie plus an exact match between the signed-in user and the user sealed + into the flow: a link crafted by another party dies here with ``access_denied`` + instead of minting a code for the victim's identity. The flow is single-use (an atomic + claim on its ``jti``), so a double-submit cannot mint two codes from one sign-in. + """ + sealed_flow = request.cookies.get(_flow_cookie_name(flow_handle)) + if sealed_flow is None: + return _oauth_error(400, "invalid_request", "unknown or expired connect flow") + flow = _open_sealed(sealed_flow, _UNPREFIXED, _ConnectFlow, _CONNECT_FLOW_DEBUG_KEY) + if flow is None: + return _oauth_error(400, "invalid_request", "unknown or expired connect flow") + now = datetime.now(timezone.utc) + if now.timestamp() >= flow.exp: + return _oauth_error(400, "invalid_request", "the connect flow has expired; restart the connection") + if session_user_id is None: + return _oauth_error(401, "login_required", "sign in to LiteLLM to finish connecting") + if session_user_id != flow.user_id: + return _oauth_error(403, "access_denied", "the signed-in user does not match this connect flow") + if not await _SingleUseGuard(cache).claim( + f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + ): + return _oauth_error(400, "invalid_request", "this connect flow was already completed; restart the connection") + code = _seal( + GATEWAY_AUTH_CODE_PREFIX, + _GatewayAuthCode( + user_id=flow.user_id, + client_id=flow.client_id, + redirect_uri=flow.redirect_uri, + code_challenge=flow.code_challenge, + jti=secrets.token_urlsafe(24), + iat=int(now.timestamp()), + exp=int(now.timestamp()) + GATEWAY_AUTH_CODE_TTL_SECONDS, + ), + ) + params = {"code": code, **({"state": flow.state} if flow.state else {})} + response = RedirectResponse(_append_query_params(flow.redirect_uri, params), status_code=303) + path, secure = _cookie_path_and_secure(request) + response.delete_cookie(key=_flow_cookie_name(flow_handle), path=path, secure=secure, httponly=True, samesite="lax") + return response + + +def _pkce_verifier_matches(code_verifier: str, code_challenge: str) -> bool: + """RFC 7636 S256 verification, total over hostile input. The comparison is over bytes + so a non-ASCII ``code_challenge`` (which reaches here unvalidated from the client's + authorize request) simply fails to match instead of raising ``TypeError`` the way + ``hmac.compare_digest`` does on two ``str`` with non-ASCII content. The verifier is + ASCII per spec; a compliant client's challenge is base64url and matches.""" + digest = hashlib.sha256(code_verifier.encode("ascii", "replace")).digest() + computed = urlsafe_b64encode(digest).rstrip(b"=") + return hmac.compare_digest(computed, code_challenge.encode("utf-8")) + + +class _SingleUseGuard: + """Atomic single-use claim for a one-time id (an auth-code, connect-flow ``jti``, or refresh-token + ``jti``) over the injected proxy cache. + + Uses an atomic increment rather than a get-then-set: two concurrent redemptions of the same id + cannot both observe "unused", because exactly one increment returns 1. The claim IS the gate, so it + fails closed. Crucially, the increment must be recorded in a backend SHARED across replicas, or the + single-use property is per-worker only (each replica's in-memory counter returns 1, so a captured + id replays through a different worker): + + - When a Redis backend is configured it is the SOLE authority: the claim goes straight to Redis + (``INCR`` is atomic across replicas), and any Redis fault fails the claim CLOSED — it never falls + back to the per-worker in-memory count (``DualCache.async_increment_cache`` does fall back, which + is exactly the replay window this avoids). + - With no Redis configured (single-replica) the in-memory increment is authoritative within the one + process. A multi-worker deployment must run Redis for the guarantee to hold across workers. + + The id's own TTL is the outer bound. For the auth code, PKCE binding is the primary defense against + interception; this makes the RFC 6749 4.1.2 single-use property reliable on top of it.""" + + def __init__(self, cache: DualCache) -> None: + self._cache = cache + + async def claim(self, key: str, ttl_seconds: int) -> bool: + """Atomically claim ``key``. ``True`` iff this caller is the first (increment to 1); ``False`` + on a replay (>1) or when the claim could not be recorded in the shared backend (fail closed).""" + from litellm.proxy.proxy_server import redis_usage_cache # noqa: PLC0415 # circular import at module load + + # Resolve the shared authority HERE rather than trusting the injected cache: callers pass + # user_api_key_cache, which only carries a redis_cache when enable_redis_auth_cache is set + # (off by default), so a guard that read its injected cache silently degraded every claim to + # a per-worker count on a stock multi-worker deployment. redis_usage_cache is the store the + # proxy already treats as cross-worker, so no call site can wire the guarantee away. + redis_cache = redis_usage_cache or getattr(self._cache, "redis_cache", None) + if redis_cache is not None: + # Shared, atomic authority for multi-replica deployments. Claim ONLY against Redis and fail + # CLOSED on any Redis fault (async_increment re-raises) rather than fall back to the + # per-worker in-memory count, which would let each replica observe count==1 and replay the id. + try: + count = await redis_cache.async_increment(key, 1, ttl=ttl_seconds) + except Exception as e: # noqa: BLE001 # ANY Redis fault fails the single-use claim closed + verbose_logger.warning( + "mcp gateway single-use claim: shared cache backend unavailable, failing closed: %s", e + ) + return False + return count == 1 + # No shared backend configured (single-replica): the in-memory increment is authoritative. + count = await self._cache.async_increment_cache(key, 1, ttl=ttl_seconds, local_only=True) + return count == 1 + + +def _session_token_pair(principal: SessionPrincipal, keys: SessionKeys, now: datetime) -> Response: + access = mint_session_token(principal, keys, now) + refresh = mint_session_refresh_token(principal, keys, now) + if not isinstance(access, MintedSessionToken) or not isinstance(refresh, MintedSessionToken): + return _oauth_error(500, "server_error", "failed to mint the session credential") + return JSONResponse( + status_code=200, + content={ + "access_token": access.token.get_secret_value(), + "token_type": "Bearer", + "expires_in": int((access.expires_at - now).total_seconds()), + "refresh_token": refresh.token.get_secret_value(), + }, + headers=TOKEN_NO_CACHE_HEADERS, + ) + + +def _reload_failure_response(failure: ReloadUserFailure) -> Response: + """Map the live-user revalidation failure onto its OAuth error, exhaustively, so a new + ``ReloadUserFailure`` member is a type error here rather than silently 400ing.""" + match failure: + case "unavailable": + return _oauth_error(503, "temporarily_unavailable", "the gateway database is unavailable; retry") + case "unresolvable": + return _oauth_error(500, "server_error", "the gateway is not configured to resolve users") + case "no_active_key": + return _oauth_error(400, "invalid_grant", "the user for this grant is no longer active") + case _: + assert_never(failure) + + +async def aggregate_token( + request: Request, + grant_type: str, + code: str | None, + redirect_uri: str | None, + client_id: str, + code_verifier: str | None, + refresh_token: str | None, + master_key: str | None, + reload_user: ReloadUser, + cache: DualCache, +) -> Response: + """The aggregate token verb: authorization_code and refresh_token grants for the + identity-only session pair. Every path re-validates the litellm user live before + minting, so a deactivated user cannot obtain or renew a session.""" + if master_key is None: + verbose_logger.error("mcp_gateway_dcr token grant rejected: no master_key configured") + return _oauth_error(500, "server_error", "the gateway has no master key configured") + keys = session_keys_from_master_key(master_key) + now = datetime.now(timezone.utc) + if grant_type == "authorization_code": + return await _authorization_code_grant( + code=code, + redirect_uri=redirect_uri, + client_id=client_id, + code_verifier=code_verifier, + keys=keys, + now=now, + reload_user=reload_user, + guard=_SingleUseGuard(cache), + ) + if grant_type == "refresh_token": + return await _refresh_token_grant( + refresh_token=refresh_token, + client_id=client_id, + keys=keys, + now=now, + reload_user=reload_user, + guard=_SingleUseGuard(cache), + ) + return _oauth_error(400, "unsupported_grant_type", "grant_type must be authorization_code or refresh_token") + + +async def _authorization_code_grant( + code: str | None, + redirect_uri: str | None, + client_id: str, + code_verifier: str | None, + keys: SessionKeys, + now: datetime, + reload_user: ReloadUser, + guard: _SingleUseGuard, +) -> Response: + if not code or not redirect_uri or not code_verifier: + return _oauth_error(400, "invalid_request", "code, redirect_uri, and code_verifier are required") + if not MIN_CODE_VERIFIER_LENGTH <= len(code_verifier) <= MAX_CODE_VERIFIER_LENGTH: + return _oauth_error(400, "invalid_request", "code_verifier must be 43 to 128 characters (RFC 7636)") + parsed = _open_sealed(code, GATEWAY_AUTH_CODE_PREFIX, _GatewayAuthCode, _AUTH_CODE_DEBUG_KEY) + if parsed is None: + return _oauth_error(400, "invalid_grant", "the authorization code is invalid") + if now.timestamp() >= parsed.exp: + return _oauth_error(400, "invalid_grant", "the authorization code has expired") + if client_id != parsed.client_id or redirect_uri != parsed.redirect_uri: + return _oauth_error(400, "invalid_grant", "the authorization code was issued to a different client") + if not _pkce_verifier_matches(code_verifier, parsed.code_challenge): + return _oauth_error(400, "invalid_grant", "PKCE verification failed") + # Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable + # 503) does not consume a still-valid code and force the client to restart sign-in. + failure = await reload_user(parsed.user_id) + if failure is not None: + return _reload_failure_response(failure) + # Atomic single-use claim is the gate: on a concurrent double-redeem exactly one caller + # wins, and a claim that cannot be recorded fails closed. + if not await guard.claim( + f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", GATEWAY_AUTH_CODE_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + ): + return _oauth_error(400, "invalid_grant", "the authorization code was already used") + return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now) + + +async def _refresh_token_grant( + refresh_token: str | None, + client_id: str, + keys: SessionKeys, + now: datetime, + reload_user: ReloadUser, + guard: _SingleUseGuard, +) -> Response: + if not refresh_token: + return _oauth_error(400, "invalid_request", "refresh_token is required") + opened = open_session_refresh_bearer(refresh_token, keys, now, expected_client_id=client_id) + if not isinstance(opened, SessionRefreshOpened): + return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client") + failure = await reload_user(opened.principal.user_id) + if failure is not None: + return _reload_failure_response(failure) + # Refresh-token rotation (OAuth 2.0 Security BCP section 4.13): the presented refresh token is + # single-use. Claim its jti before issuing the replacement pair, so a captured or replayed + # refresh token cannot mint a second pair after the legitimate holder rotated. Claimed AFTER + # user revalidation so a transient DB 503 does not burn a still-valid token; a claim that + # cannot be recorded fails closed, exactly like the authorization-code path. + if not await guard.claim( + f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + ): + return _oauth_error(400, "invalid_grant", "the refresh token was already used") + return _session_token_pair(opened.principal, keys, now) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 90b70dd01f2..bc4f5d60589 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -49,6 +49,7 @@ from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, + _is_mcp_admitted_user_subject, ) from litellm.proxy._experimental.mcp_server.exceptions import ( MCPServerListError, @@ -2197,6 +2198,56 @@ class MCPServerManager: return [server_id for server_id in submitted_server_ids if self.get_mcp_server_by_id(server_id) is not None] + async def operator_open_server_ids( + self, + user_api_key_auth: UserAPIKeyAuth | None = None, + *, + allow_all_server_ids: list[str] | None = None, + submitted_server_ids: list[str] | None = None, + ) -> set: + """Servers reachable through OPEN channels rather than a grant: operator-opened + ``allow_all_keys`` servers, plus the caller's own active BYOM submissions when the caller + carries no explicit ``mcp_servers`` scope. + + The single owner of that question for BOTH axes. The server union in + ``get_allowed_mcp_servers`` adds these ids, and the admitted subject's tool resolution asks + the same question to treat an open-channel server as default-open for tools — exactly how a + virtual key experiences it. Encoding the channel membership twice is how a server ends up + listable but uninvokable. + + Empty inside a toolset scope: toolset_mcp_route / dynamic_mcp_route set + ``_mcp_active_toolset_id`` before calling the handler, pinning the request to the toolset's + own servers (checking op.mcp_toolsets==[] instead would false-positive on DB-default rows + where Postgres initialises the column to ARRAY[]::TEXT[]). + + ``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union, + which precomputes both for its fallback path, does not compute them twice.""" + from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415 + _mcp_active_toolset_id, + ) + + if _mcp_active_toolset_id.get() is not None: + return set() + if allow_all_server_ids is None: + allow_all_server_ids = self.get_allow_all_keys_server_ids() + open_ids = set(allow_all_server_ids) + key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None + # "Explicitly scoped, so do not widen with BYOM" is a rule about a CREDENTIAL that carries + # its own mcp_servers list. It does not describe a keyless admitted subject: its + # object_permission is the user's own row, whose mcp_servers column is [] by DB default, so + # applying this rule would hide almost every admitted user's OWN submitted servers. Their + # submissions are theirs by authorship, and their scope comes from the per-source union. + has_explicit_object_permission = ( + not _is_mcp_admitted_user_subject(user_api_key_auth) + and key_object_permission is not None + and (key_object_permission.mcp_servers is not None) + ) + if not has_explicit_object_permission: + if submitted_server_ids is None: + submitted_server_ids = await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth) + open_ids.update(submitted_server_ids) + return open_ids + async def get_allowed_mcp_servers(self, user_api_key_auth: Optional[UserAPIKeyAuth] = None) -> list[str]: """ Get the allowed MCP Servers for the user. @@ -2210,11 +2261,22 @@ class MCPServerManager: allow_all_server_ids = self.get_allow_all_keys_server_ids() + # A keyless admitted subject is resolved per grant source, and channel decisions that are + # absolute for a scoped KEY credential are not absolute for it: its own opt-out silences its + # own source (handled per source in the resolver), never its teams' grants, and its admin + # role does not swallow the grant model — a session bearer is a third-party client + # credential, not the dashboard, so an admin signing in through the connect flow gets their + # grants like anyone else rather than handing the client the full registry ahead of every + # per-team org ceiling. + is_admitted_subject = _is_mcp_admitted_user_subject(user_api_key_auth) + # The key explicitly opted out of every MCP server. Return zero before # layering on allow_all_keys or submitted servers so the opt-out is absolute. key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None - if key_object_permission is not None and ( - SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or []) + if ( + not is_admitted_subject + and key_object_permission is not None + and (SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or [])) ): return [] @@ -2234,8 +2296,14 @@ class MCPServerManager: ) try: - # If admin but NO explicit object permission, get all servers - if user_api_key_auth and _user_has_admin_view(user_api_key_auth) and not has_explicit_object_permission: + # If admin but NO explicit object permission, get all servers (never for an admitted + # subject — see is_admitted_subject above) + if ( + user_api_key_auth + and not is_admitted_subject + and _user_has_admin_view(user_api_key_auth) + and not has_explicit_object_permission + ): verbose_logger.debug("Admin user without explicit object_permission - returning all servers") return list(self.get_registry().keys()) @@ -2243,20 +2311,14 @@ class MCPServerManager: allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) verbose_logger.debug(f"Allowed MCP Servers for user api key auth: {allowed_mcp_servers}") combined_servers = set(allowed_mcp_servers) - # Only skip allow_all_keys servers when the request is inside a toolset - # scope. toolset_mcp_route / dynamic_mcp_route set _mcp_active_toolset_id - # before calling the handler — that ContextVar is the reliable signal. - # Using op.mcp_toolsets==[] would false-positive on DB-default rows where - # Postgres initialises the column to ARRAY[]::TEXT[]. - from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415 - _mcp_active_toolset_id, + combined_servers.update( + await self.operator_open_server_ids( + user_api_key_auth, + allow_all_server_ids=allow_all_server_ids, + submitted_server_ids=submitted_server_ids, + ) ) - in_toolset_scope = _mcp_active_toolset_id.get() is not None - if not in_toolset_scope: - combined_servers.update(allow_all_server_ids) - combined_servers.update(submitted_server_ids) - # For anonymous callers (no user_id, no role), also surface any # servers the operator has opted into upstream-delegated auth. # These servers handle their own auth at the upstream level, so diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 53686e329bb..0f5bde908c0 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -343,8 +343,36 @@ def _parse_redirect_uri_for_validation(redirect_uri: str) -> ParseResult: ) -def _validate_trusted_http_redirect_shape(parsed: ParseResult) -> bool: - """Return True when ``parsed`` is an allowlisted native callback (caller may return).""" +def is_loopback_redirect_host(parsed: ParseResult) -> bool: + """True when the redirect host is loopback (RFC 8252 section 7.3). + + Shared by every redirect-URI policy in the MCP OAuth surface so that none of them + hand-rolls its own host list: a literal ``("localhost", "127.0.0.1", "::1")`` tuple + silently misses the rest of 127.0.0.0/8 and IPv6-mapped forms. + """ + host = (parsed.hostname or "").lower() + if host == "localhost": + return True + try: + return ip_address(host).is_loopback + except ValueError: + return False + + +def validate_redirect_uri_shape(parsed: ParseResult) -> bool: + """Validate redirect-URI *hygiene* and resolve allowlisted native callbacks. + + Returns True when ``parsed`` is an allowlisted native callback (the caller may accept + it outright); returns False for http/https, leaving the trust decision to the caller; + raises for a URI that no policy should ever accept (bad scheme, fragment, missing + host, userinfo, backslash in the host). + + This is deliberately separate from :func:`validate_trusted_redirect_uri`, which adds + the *first-party* trust policy (same-origin, loopback, ops allowlist) appropriate to + the proxy's own OAuth endpoints. Public dynamic-client registration accepts any https + client and relies on PKCE plus the consent screen instead, so it shares this hygiene + rule but not that trust policy. + """ if parsed.scheme not in ("http", "https"): if _matches_trusted_native_redirect_uri(parsed): return True @@ -396,14 +424,8 @@ def _trusted_redirect_uri_is_allowed( ): return True - host = (parsed.hostname or "").lower() - if host == "localhost": + if is_loopback_redirect_host(parsed): return True - try: - if ip_address(host).is_loopback: - return True - except ValueError: - pass if parsed.scheme == "https": for entry in _parse_trusted_redirect_origins(): @@ -522,7 +544,7 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None: :func:`validate_loopback_redirect_uri`. """ parsed = _parse_redirect_uri_for_validation(redirect_uri) - if _validate_trusted_http_redirect_shape(parsed): + if validate_redirect_uri_shape(parsed): return redirect_netloc = _strip_default_port(parsed.scheme, parsed.netloc) proxy_base = _resolve_proxy_base_for_redirect(request) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py index 08d5cc8b1f1..8844d8c8ad0 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py @@ -149,6 +149,7 @@ class SessionRefreshOpened(BaseModel): model_config = ConfigDict(frozen=True) tag: Literal["opened"] = "opened" principal: SessionPrincipal + jti: str class SessionRefreshInvalid(BaseModel): @@ -187,4 +188,4 @@ def open_session_refresh_bearer( return SessionRefreshInvalid() if opened.principal.client_id != expected_client_id: return SessionRefreshInvalid() - return SessionRefreshOpened(principal=opened.principal) + return SessionRefreshOpened(principal=opened.principal, jti=opened.jti) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py index 9325428f049..4ccbcd1a511 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py @@ -113,10 +113,12 @@ class MintedSessionToken(BaseModel): class OpenedSessionToken(BaseModel): - """A validated session token of either kind: the principal it was minted for.""" + """A validated session token of either kind: the principal it was minted for, plus the + ``jti`` so the token endpoint can enforce single-use rotation on a refresh token.""" model_config = ConfigDict(frozen=True) principal: SessionPrincipal + jti: str class SessionTokenTooLarge(BaseModel): @@ -320,7 +322,9 @@ def _open( return SessionMalformed() if now.timestamp() >= claims.exp: return SessionExpired() - return OpenedSessionToken(principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id)) + return OpenedSessionToken( + principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id), jti=claims.jti + ) def _decode_claims( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 396dd6c7dc7..f135fa5e5b4 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -977,7 +977,17 @@ if MCP_AVAILABLE: data = await add_litellm_data_to_request( data=body_data, request=request, - user_api_key_dict=user_api_key_auth, + # Bill a team-derived call to the team that granted it. A keyless admitted + # subject carries no team_id, so spend skipped team updates entirely and + # charged the user's PRIMARY org — the granting team's budget never + # accumulated (so it could never begin to block) and, cross-org, the wrong + # organization was charged. This is the ACCOUNTING half; the enforcement + # half (an already-over-budget team stops granting) lives in the source gate. + # Authorization is unaffected: it ran before this, and the union is resolved + # from the untouched auth object passed to call_mcp_tool below. + user_api_key_dict=await MCPRequestHandler.billing_auth_for_tool_call( + user_api_key_auth, tool_name=name + ), proxy_config=proxy_config, ) else: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b73841c4793..444e5ba0731 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2605,6 +2605,17 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob user_max_budget: Optional[float] = None request_route: Optional[str] = None is_session_token: bool = False + # Server-only marker set exclusively by the MCP gateway admission path + # (_reload_admitted_user) for a keyless user-subject admitted via a gateway DCR session + # bearer or bridge envelope. Not a DB column and never populated from caller-controlled key + # metadata or JWT claims, so it cannot be forged to gain the team-inherited MCP grant union + # or to escape the caller-Authorization egress scrub. exclude=True keeps it out of serialization. + mcp_admitted_user_subject: bool = Field(default=False, exclude=True) + # team_id -> that team's mcp_rpm_limit map, for a keyless admitted subject that reaches MCP + # servers through several teams at once and therefore has no single team_id for the limiter to + # key off. Server-only and stripped from validated input for the same reason as the marker + # above: a forged entry would let a caller pick which team's rpm bucket it is charged against. + mcp_source_team_rpm_limits: dict[str, dict[str, int]] | None = Field(default=None, exclude=True) budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True) budget_throttle_pct: Optional[float] = Field(default=None, exclude=True) user: Optional[Any] = None # Expanded user object when expand=user is used @@ -2625,6 +2636,11 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # If values is already an instance (not a dict), return it as-is if not isinstance(values, dict): return values + # mcp_admitted_user_subject is a server-only marker, set ONLY by the MCP gateway admission + # path via post-construction assignment. Strip it from any validated input (constructor + # kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data. + values.pop("mcp_admitted_user_subject", None) + values.pop("mcp_source_team_rpm_limits", None) if values.get("api_key") is not None: values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}) if isinstance(values.get("api_key"), str): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 99a867a5d07..ce82ca74267 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2661,6 +2661,15 @@ async def get_managed_vector_store_rows_by_uuids( return result +class OrganizationNotFoundError(Exception): + """The organization row is CONFIRMED absent, as opposed to a lookup that failed. + + Subclasses Exception so every existing except Exception caller keeps its current + behavior; it exists so a caller that wants to treat "no such org" as "no restriction" can do + that WITHOUT also swallowing an outage and silently dropping a real org ceiling. + """ + + @log_db_metrics async def get_org_object( org_id: str, @@ -2707,25 +2716,30 @@ async def get_org_object( query_kwargs["include"] = {"litellm_budget_table": True} response = await OrganizationRepository(prisma_client).table.find_unique(**query_kwargs) - - if response is None: - raise Exception - - _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) - # Cache the result - await user_api_key_cache.async_set_cache( - key=cache_key, - value=_org_obj, - model_type=LiteLLM_OrganizationTable, - ttl=DEFAULT_IN_MEMORY_TTL, - ) - - return _org_obj except Exception: - raise Exception( + # An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed + # missing row, and relabelling it as "doesn't exist" made every caller unable to tell them + # apart — a caller that treats absence as "this org places no restriction" then drops a real + # org ceiling during an outage. Propagate the real error; callers that already catch + # Exception are unaffected. + raise + + if response is None: + raise OrganizationNotFoundError( f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call." ) + _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) + # Cache the result + await user_api_key_cache.async_set_cache( + key=cache_key, + value=_org_obj, + model_type=LiteLLM_OrganizationTable, + ttl=DEFAULT_IN_MEMORY_TTL, + ) + + return _org_obj + async def _get_resources_from_access_groups( access_group_ids: List[str], diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index 11f12e597b9..f35d94c986e 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -7,12 +7,15 @@ login endpoints (e.g., /login and /v2/login). import os import secrets +from datetime import datetime, timedelta, timezone from typing import Literal, Optional, cast +import jwt from fastapi import HTTPException import litellm from litellm.constants import LITELLM_PROXY_ADMIN_NAME, LITELLM_UI_SESSION_DURATION +from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy._types import ( LiteLLM_UserTable, LitellmUserRoles, @@ -313,6 +316,29 @@ async def authenticate_user( ) +def _ui_session_exp_timestamp() -> int: + """The ``exp`` claim (unix seconds) for a UI session cookie, ``LITELLM_UI_SESSION_DURATION`` + from now. The virtual key sealed inside the cookie already expires after this same + duration; stamping the JWT itself gives the cookie the bounded lifetime the dashboard's + client-side expiry check and the server-side session-cookie readers both assume, instead + of a token that stays signature-valid until the master key rotates.""" + ttl_seconds = duration_in_seconds(LITELLM_UI_SESSION_DURATION) + return int((datetime.now(timezone.utc) + timedelta(seconds=ttl_seconds)).timestamp()) + + +def encode_ui_session_jwt(returned_ui_token_object: ReturnedUITokenObject, master_key: str) -> str: + """Encode a UI session cookie JWT with a bounded ``exp``. + + The single choke point every UI login path (SSO and username/password /login, /v2, + /v3) uses to mint the ``token`` cookie, so the cookie's lifetime is set in exactly one + place and cannot drift between paths. Without the ``exp`` the cookie is valid until the + master key rotates, and the session-cookie readers that require a bounded lifetime + (the MCP interactive sign-in) reject it. + """ + claims = {**cast(dict, returned_ui_token_object), "exp": _ui_session_exp_timestamp()} + return jwt.encode(claims, master_key, algorithm="HS256") + + def create_ui_token_object( login_result: LoginResult, general_settings: dict, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 22ea9fe176a..b2216488db2 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -1781,28 +1781,38 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ from litellm.proxy.auth.auth_utils import get_team_mcp_rpm_limit - if not mcp_server_name or not user_api_key_dict.team_id: + if not mcp_server_name: return - mcp_rpm_limit = get_team_mcp_rpm_limit(user_api_key_dict) - if not mcp_rpm_limit: - return + # Which teams' buckets does this call charge? A key is pinned to exactly one team. A keyless + # MCP-admitted subject reaches servers through SEVERAL teams at once and has no team_id, so + # without the second source below its calls charged no team bucket at all and it outran every + # team's mcp_rpm_limit. Every applicable team is charged rather than one being picked: the + # limiter enforces all descriptors, so each team's own ceiling binds on a call made through + # its grant, and there is no arbitrary attribution when several teams grant the same server. + team_limits: list[tuple[str | None, dict[str, int] | None]] = [] + if user_api_key_dict.team_id: + team_limits.append((user_api_key_dict.team_id, get_team_mcp_rpm_limit(user_api_key_dict))) + for source_team_id, source_limit in (user_api_key_dict.mcp_source_team_rpm_limits or {}).items(): + team_limits.append((source_team_id, source_limit)) - server_rpm_limit = mcp_rpm_limit.get(mcp_server_name) - if server_rpm_limit is None: - return - - descriptors.append( - RateLimitDescriptor( - key="mcp_per_team", - value=f"{user_api_key_dict.team_id}:{mcp_server_name}", - rate_limit={ - "requests_per_unit": server_rpm_limit, - "tokens_per_unit": None, - "window_size": self.window_size, - }, + for team_id, mcp_rpm_limit in team_limits: + if not team_id or not mcp_rpm_limit: + continue + server_rpm_limit = mcp_rpm_limit.get(mcp_server_name) + if server_rpm_limit is None: + continue + descriptors.append( + RateLimitDescriptor( + key="mcp_per_team", + value=f"{team_id}:{mcp_server_name}", + rate_limit={ + "requests_per_unit": server_rpm_limit, + "tokens_per_unit": None, + "window_size": self.window_size, + }, + ) ) - ) def _should_enforce_rate_limit( self, diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index de988a0140f..31b98bf20e4 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -36,7 +36,7 @@ if TYPE_CHECKING: import httpx import jwt -from fastapi import APIRouter, Depends, Header, HTTPException, Request, status +from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status from fastapi.responses import RedirectResponse import litellm @@ -965,15 +965,8 @@ async def google_login( state=cli_state, request=request, ) - if return_to is not None and sso_redirect is not None: - if SSOAuthenticationHandler._validate_return_to(return_to): - sso_redirect.set_cookie( - key="litellm_cp_return_to", - value=return_to, - max_age=600, - httponly=True, - samesite="lax", - ) + if sso_redirect is not None: + _persist_return_to_cookie(sso_redirect, return_to) return sso_redirect from fastapi.responses import HTMLResponse @@ -982,13 +975,19 @@ async def google_login( os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true" or general_settings.get("hide_default_credentials_hint", False) is True ) - return HTMLResponse( + form_response = HTMLResponse( content=build_ui_login_form( show_deprecation_banner=True, hide_default_credentials_hint=hide_default_credentials_hint, ), status_code=200, ) + # Preserve return_to across the username/password sign-in too, via the SAME shared, never-raising + # helper the SSO branch uses, so /login can resume the connect flow instead of dead-ending at the + # dashboard. One implementation → the two sign-in branches cannot diverge (and the login form always + # renders, since the helper never raises on a bad return_to). + _persist_return_to_cookie(form_response, return_to) + return form_response def generic_response_convertor( @@ -2418,6 +2417,92 @@ async def sso_readiness(): ) +def _is_same_origin_return_path(return_to: str) -> bool: + """True for a strictly relative return path that stays on the gateway's own origin by + construction, and is therefore safe to honor without a configured ``control_plane_url``. + Used by the MCP gateway DCR authorize round-trip so a browser sent through login lands + back on the authorize request. + + Requires a single leading ``/`` (not protocol-relative ``//``), no backslash (browsers + fold ``\\`` to ``/``, so ``/\\evil.com`` would escape the origin), and no control or + whitespace characters. Rejecting control chars keeps a ``\\r\\n``/tab-bearing value out + of the redirect ``Location`` and the ``litellm_cp_return_to`` cookie entirely, rather + than relying on downstream header encoding to neutralize it.""" + if not return_to.startswith("/") or return_to.startswith("//") or "\\" in return_to: + return False + return not any(ord(ch) < 0x20 or ch in (" ", "\x7f") for ch in return_to) + + +async def _sso_return_to_redirect( + return_to: str | None, + jwt_token: str, + redis_usage_cache, + user_api_key_cache, +) -> RedirectResponse | None: + """Resolve the post-SSO redirect for a ``return_to``, or None to fall through to the dashboard. + + Two arms, both clearing the one-shot ``litellm_cp_return_to`` cookie: + - **Same-origin relative path** (the MCP gateway DCR authorize round-trip): set the session cookie + exactly like the dashboard path, then send the browser back where it came from. + - **Control-plane cross-origin** (``control_plane_url``): stash the JWT behind a single-use opaque + code (60s TTL) so the token never lands in browser history/logs; the control plane redeems it via + ``POST /v3/login/exchange``. + + Extracted from ``get_redirect_response_from_openid`` to keep that method inside the complexity + budget; behavior is identical to the inline arms it replaces (including letting + ``_validate_return_to`` raise for a mismatched absolute return_to, as before).""" + if return_to is None: + return None + + if _is_same_origin_return_path(return_to): + redirect_response = RedirectResponse(url=return_to, status_code=303) + redirect_response.set_cookie(key="token", value=jwt_token) + redirect_response.delete_cookie("litellm_cp_return_to") + return redirect_response + + if SSOAuthenticationHandler._validate_return_to(return_to): + code = secrets.token_urlsafe(32) + cache_key = f"login_code:{code}" + cache_value = {"token": jwt_token, "redirect_url": return_to} + if redis_usage_cache is not None: + await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) + else: + await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) + + separator = "&" if "?" in return_to else "?" + redirect_url = return_to + separator + urlencode({"login": "success", "code": code}) + verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code") + redirect_response = RedirectResponse(url=redirect_url, status_code=303) + redirect_response.delete_cookie("litellm_cp_return_to") + return redirect_response + + return None + + +def _persist_return_to_cookie(response: Response, return_to: str | None) -> None: + """Best-effort: persist a SAFE ``return_to`` on ``response`` as the one-shot ``litellm_cp_return_to`` + cookie so ANY sign-in path — SSO / Okta / generic OR the username/password form — can resume there + afterwards. THIS is the single source of truth, called by every sign-in branch so they cannot + diverge (a per-branch reimplementation is exactly how the two drifted before). Honors a strictly + relative same-origin path, and (when ``control_plane_url`` is configured) a return_to matching that + origin. It NEVER raises: a mismatched or invalid ``return_to`` is simply not stored, so it can never + block sign-in — the login entrypoint must always render.""" + if return_to is None: + return + try: + safe = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler._validate_return_to(return_to) + except HTTPException: + return # a non-matching absolute return_to is ignored, never blocks sign-in + if safe: + response.set_cookie( + key="litellm_cp_return_to", + value=return_to, + max_age=600, + httponly=True, + samesite="lax", + ) + + class SSOAuthenticationHandler: """ Handler for SSO Authentication across all SSO providers @@ -3055,7 +3140,6 @@ class SSOAuthenticationHandler: return_to: Optional[str] = None, sso_assertion: SSOIdentityAssertion | None = None, ) -> RedirectResponse: - import jwt from litellm.proxy.proxy_server import ( general_settings, @@ -3219,30 +3303,21 @@ class SSOAuthenticationHandler: server_root_path=get_server_root_path(), ) - jwt_token = jwt.encode( - cast(dict, returned_ui_token_object), - master_key or "", - algorithm="HS256", + from litellm.proxy.auth.login_utils import encode_ui_session_jwt + + jwt_token = encode_ui_session_jwt(returned_ui_token_object, master_key or "") + + # Post-SSO return_to handling (the same-origin DCR round-trip and the control-plane + # cross-origin code exchange) lives in one shared helper so this method stays inside the + # complexity budget. None falls through to the dashboard redirect below. + return_to_redirect = await _sso_return_to_redirect( + return_to=return_to, + jwt_token=jwt_token, + redis_usage_cache=redis_usage_cache, + user_api_key_cache=user_api_key_cache, ) - - # Control-plane cross-origin: store JWT behind a single-use opaque - # code (60s TTL) so the token never appears in browser history / logs. - # The control plane redeems it via POST /v3/login/exchange. - if return_to is not None and SSOAuthenticationHandler._validate_return_to(return_to): - code = secrets.token_urlsafe(32) - cache_key = f"login_code:{code}" - cache_value = {"token": jwt_token, "redirect_url": return_to} - if redis_usage_cache is not None: - await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) - else: - await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) - - separator = "&" if "?" in return_to else "?" - redirect_url = return_to + separator + urlencode({"login": "success", "code": code}) - verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code") - redirect_response = RedirectResponse(url=redirect_url, status_code=303) - redirect_response.delete_cookie("litellm_cp_return_to") - return redirect_response + if return_to_redirect is not None: + return return_to_redirect if user_id is not None and isinstance(user_id, str): litellm_dashboard_ui += "?login=success" diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 32845763f22..a20b557e38b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -13472,7 +13472,7 @@ async def fallback_login(request: Request): @router.post("/login", include_in_schema=False) # hidden since this is a helper for UI sso login async def login(request: Request): global premium_user, general_settings, master_key - from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object + from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt from litellm.proxy.utils import get_custom_url form = await request.form() @@ -13495,13 +13495,7 @@ async def login(request: Request): ) # Generate JWT token - import jwt - - jwt_token = jwt.encode( - cast(dict, returned_ui_token_object), - cast(str, master_key), - algorithm="HS256", - ) + jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key)) # Build redirect URL litellm_dashboard_ui = get_custom_url(str(request.base_url)) @@ -13511,16 +13505,51 @@ async def login(request: Request): litellm_dashboard_ui += "/ui/" litellm_dashboard_ui += "?login=success" + # Honor a same-origin return_to preserved by the sign-in page (e.g. the aggregate DCR connect flow's + # authorize round-trip), mirroring the SSO callback; otherwise land on the dashboard. Gated by + # _is_same_origin_return_path (strictly relative path) so it can never be an open redirect, and the + # one-shot cookie is cleared after use. + from litellm.proxy.management_endpoints.ui_sso import _sso_return_to_redirect + + # Resume through the SAME resumer the SSO callback uses, rather than a second, narrower arm. + # _persist_return_to_cookie stores both shapes it accepts (a relative same-origin path AND a + # control_plane_url-matching absolute URL); honoring only the relative one here silently dropped + # the control-plane case, landing the user on the dashboard. One function decides how a stored + # return_to is honored for EVERY sign-in branch, so the write and read sets cannot diverge: it + # sets the token cookie on the same-origin arm and hands off via a one-time login code on the + # cross-origin arm, and clears the one-shot cookie in both. + cp_return_to = request.cookies.get("litellm_cp_return_to") + if cp_return_to: + try: + resumed = await _sso_return_to_redirect( + return_to=cp_return_to, + jwt_token=jwt_token, + redis_usage_cache=redis_usage_cache, + user_api_key_cache=user_api_key_cache, + ) + except Exception: # noqa: BLE001 # resuming must NEVER block a completed sign-in + # The symmetric half of _persist_return_to_cookie's "never raises" contract. The resumer + # rejects a return_to that no longer matches control_plane_url (a config change between + # the cookie's write and this read), and the user has ALREADY authenticated here — + # failing their login over a stale one-shot cookie is the worst possible outcome. Land + # on the dashboard instead; the cookie is cleared below either way. + verbose_proxy_logger.info("Ignoring stale litellm_cp_return_to cookie; landing on dashboard") + resumed = None + if resumed is not None: + return resumed + # Create redirect response with cookie redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303) redirect_response.set_cookie(key="token", value=jwt_token) + if cp_return_to: + redirect_response.delete_cookie(key="litellm_cp_return_to") return redirect_response @router.post("/v2/login", include_in_schema=False) # hidden helper for UI logins via API async def login_v2(request: Request): global premium_user, general_settings, master_key - from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object + from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt from litellm.proxy.utils import get_custom_url try: @@ -13541,13 +13570,7 @@ async def login_v2(request: Request): premium_user=premium_user, ) - import jwt - - jwt_token = jwt.encode( - cast(dict, returned_ui_token_object), - cast(str, master_key), - algorithm="HS256", - ) + jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key)) litellm_dashboard_ui = get_custom_url(str(request.base_url)) if litellm_dashboard_ui.endswith("/"): @@ -13591,7 +13614,7 @@ async def login_v2(request: Request): ) # control-plane login — always returns token in body for cross-origin use async def login_v3(request: Request): global premium_user, general_settings, master_key - from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object + from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt from litellm.proxy.utils import get_custom_url try: @@ -13620,13 +13643,7 @@ async def login_v3(request: Request): premium_user=premium_user, ) - import jwt - - jwt_token = jwt.encode( - cast(dict, returned_ui_token_object), - cast(str, master_key), - algorithm="HS256", - ) + jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key)) litellm_dashboard_ui = get_custom_url(str(request.base_url)) if litellm_dashboard_ui.endswith("/"): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 7b05b8c9dd0..5d32a9c740d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -15,6 +15,7 @@ from starlette.datastructures import Headers from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, + _is_mcp_admitted_user_subject, ) from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -206,6 +207,31 @@ class TestMCPRequestHandler: assert sorted(result) == sorted(expected) + async def test_admitted_subject_not_zeroed_by_require_key_mcp_access_defined(self): + """10x-flow regression: with require_key_mcp_access_defined ON (team = ceiling for keys), a + keyless gateway/bridge-admitted subject whose ONLY access path is team membership must still + inherit the team's servers. The flag zeros empty *virtual keys* that must declare their own + access; a keyless admitted user has no key to declare it on, so it must not be zeroed.""" + auth = UserAPIKeyAuth(api_key=None, user_id="sso-user") + auth.mcp_admitted_user_subject = True + with ( + patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_key", new_callable=AsyncMock, return_value=[] + ), + patch.object( + MCPRequestHandler, + "_get_allowed_mcp_servers_for_team", + new_callable=AsyncMock, + return_value=["team_server1", "team_server2"], + ), + patch.object( + MCPRequestHandler, "_get_key_access_group_mcp_server_extras", new_callable=AsyncMock, return_value=[] + ), + patch("litellm.proxy.proxy_server.general_settings", {"require_key_mcp_access_defined": True}), + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert sorted(result) == ["team_server1", "team_server2"] + @pytest.mark.parametrize( "key_servers,grants,expected,scenario", [ @@ -5302,6 +5328,7 @@ class TestMCPDcrBridgeDelegateAdmission: self._patch_user_reload( return_value=MagicMock( user_id="sso-user-7", + organization_id=None, metadata={"scim_active": True}, user_role=None, object_permission=None, @@ -5344,6 +5371,7 @@ class TestMCPDcrBridgeDelegateAdmission: self._patch_user_reload( return_value=MagicMock( user_id="sso-user-7", + organization_id=None, metadata={"scim_active": True}, user_role=None, object_permission=object_permission, @@ -5421,7 +5449,9 @@ class TestMCPDcrBridgeDelegateAdmission: with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - self._patch_user_reload(return_value=MagicMock(user_id="offboarded-user", metadata={"scim_active": False})), + self._patch_user_reload( + return_value=MagicMock(user_id="offboarded-user", organization_id=None, metadata={"scim_active": False}) + ), ): mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() with pytest.raises(HTTPException) as exc_info: @@ -6204,7 +6234,9 @@ class TestAggregateGatewayDcrChallenge: with pytest.raises(HTTPException) as exc_info: await MCPRequestHandler.process_mcp_request(self._scope()) www_authenticate = (exc_info.value.headers or {})["WWW-Authenticate"] - assert 'resource_metadata="http://testserver/.well-known/oauth-protected-resource/litellm/mcp"' in www_authenticate + assert ( + 'resource_metadata="http://testserver/.well-known/oauth-protected-resource/litellm/mcp"' in www_authenticate + ) async def test_no_challenge_for_explicit_litellm_key(self): """An explicit x-litellm-api-key declares a litellm-key client; a typo @@ -6225,9 +6257,7 @@ class TestAggregateGatewayDcrChallenge: patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()), ): with pytest.raises(ProxyException): - await MCPRequestHandler.process_mcp_request( - self._scope(extra_headers=((b"x-mcp-servers", b"github"),)) - ) + await MCPRequestHandler.process_mcp_request(self._scope(extra_headers=((b"x-mcp-servers", b"github"),))) async def test_no_challenge_for_path_named_server(self): """/mcp/{server} targets one server; the aggregate challenge must not @@ -6261,3 +6291,1394 @@ class TestAggregateGatewayDcrChallenge: with pytest.raises(ProxyException) as exc_info: await MCPRequestHandler.process_mcp_request(self._scope()) assert str(exc_info.value.code) == "500" + + +@pytest.mark.asyncio +class TestGatewaySessionAdmission: + """The aggregate /mcp session-bearer admission arm (mcp_gateway_dcr). A valid session + token admits under the LIVE litellm user it references; an invalid/expired/refresh/foreign + token fails closed with the aggregate invalid_token challenge; the arm fires ONLY at the + aggregate scope, never for named servers or per-server flows.""" + + _MASTER_KEY = "sk-gateway-session-admission-master-key" + + def _session_bearer(self, user_id="sso-user-42", client_id="llm_dcrc_abc"): + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + session_keys_from_master_key, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + SessionPrincipal, + mint_session_token, + mint_session_refresh_token, + ) + + keys = session_keys_from_master_key(self._MASTER_KEY) + principal = SessionPrincipal(user_id=user_id, client_id=client_id) + return mint_session_token, mint_session_refresh_token, principal, keys + + def _access_token(self, **kw): + from datetime import datetime, timezone + + mint, _refresh, principal, keys = self._session_bearer(**kw) + return mint(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + + def _scope(self, bearer, path="/mcp", extra_headers=()): + return { + "type": "http", + "method": "POST", + "path": path, + "headers": [(b"host", b"testserver"), (b"authorization", f"Bearer {bearer}".encode()), *extra_headers], + } + + @staticmethod + @contextlib.contextmanager + def _patch_user_reload(*, user_id, active=True, organization_id=None, tpm_limit=None, rpm_limit=None): + get_user_object = AsyncMock( + return_value=MagicMock( + user_id=user_id, + organization_id=organization_id, + metadata={"scim_active": active} if not active else {"scim_active": True}, + user_role=None, + object_permission=None, + object_permission_id=None, + tpm_limit=tpm_limit, + rpm_limit=rpm_limit, + ) + ) + with ( + patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + ): + yield get_user_object + + async def test_session_admission_binds_org_id_so_the_org_ceiling_applies(self): + """The admitted auth carries the user's org_id, so get_allowed_mcp_servers keeps the + org-level MCP ceiling in force for a gateway session instead of skipping it.""" + token = self._access_token(user_id="org-user") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="org-user", organization_id="org-123"), + ): + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert auth_result.org_id == "org-123" + + async def test_session_admission_copies_user_rate_limits(self): + """Security regression: the reconstructed auth must carry the live user's RPM/TPM, exactly as + the standard user-subject path does. The parallel limiter reads these off the auth object and + treats None as unlimited, so a keyless subject with them unset would invoke tools past their + configured user rate limits.""" + token = self._access_token(user_id="rl-user") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="rl-user", tpm_limit=1000, rpm_limit=50), + ): + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert auth_result.user_tpm_limit == 1000 + assert auth_result.user_rpm_limit == 50 + + async def test_valid_session_admits_under_live_user_at_aggregate_scope(self): + token = self._access_token(user_id="sso-user-42") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ) as mock_auth, + self._patch_user_reload(user_id="sso-user-42") as get_user_object, + ): + auth_result, _h, _servers, mcp_server_auth_headers, _o, _r = await MCPRequestHandler.process_mcp_request( + self._scope(token) + ) + assert get_user_object.await_args.kwargs["user_id"] == "sso-user-42" + assert auth_result.user_id == "sso-user-42" + mock_auth.assert_not_called() + # Identity-only admission injects no per-server upstream credential (unlike the + # bridge envelope arm); the headers dict is whatever the request carried, here empty. + assert not mcp_server_auth_headers + + async def test_expired_session_fails_closed_with_invalid_token_challenge(self): + from datetime import datetime, timezone + + mint, _refresh, principal, keys = self._session_bearer() + token = mint(principal, keys, datetime(2020, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + ): + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert exc_info.value.status_code == 401 + assert 'error="invalid_token"' in (exc_info.value.headers or {})["WWW-Authenticate"] + + async def test_tampered_session_fails_closed(self): + token = self._access_token() + tampered = token[:-3] + ("aaa" if not token.endswith("aaa") else "bbb") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + ): + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(self._scope(tampered)) + assert exc_info.value.status_code == 401 + + async def test_deactivated_user_fails_with_invalid_token_challenge(self): + """A cryptographically valid bearer whose referenced user is SCIM-deactivated must fail with + the aggregate invalid_token challenge (WWW-Authenticate), matching the expired/tampered arms, + so the DCR client re-authorizes instead of getting a bare 401 with no challenge.""" + token = self._access_token(user_id="offboarded-user") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="offboarded-user", active=False), + ): + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert exc_info.value.status_code == 401 + assert 'error="invalid_token"' in (exc_info.value.headers or {})["WWW-Authenticate"] + + async def test_session_bearer_scrubbed_from_egress_header_contexts(self): + """Security regression (credential leak): after a keyless session admission, the session + bearer must be removed from BOTH returned egress header contexts (oauth2_headers and the raw + headers) so no passthrough/OBO egress can forward it upstream for replay as this user.""" + token = self._access_token(user_id="sso-user-42") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="sso-user-42"), + ): + _auth, _h, _servers, _msah, oauth2_headers, raw_headers = await MCPRequestHandler.process_mcp_request( + self._scope(token) + ) + # the request carried "Authorization: Bearer "; both egress contexts must be scrubbed + assert oauth2_headers is None + assert not any(k.lower() == "authorization" for k in (raw_headers or {})) + + async def test_refresh_token_is_not_admitted_at_the_tool_edge(self): + from datetime import datetime, timezone + + _mint, refresh, principal, keys = self._session_bearer() + refresh_token = refresh(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + ): + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(self._scope(refresh_token)) + assert exc_info.value.status_code == 401 + + async def test_foreign_key_session_fails_closed(self): + token = self._access_token() + with ( + patch("litellm.proxy.proxy_server.master_key", "sk-a-totally-different-master-key"), + ): + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert exc_info.value.status_code == 401 + + async def test_arm_does_not_fire_for_named_server(self): + """A session-shaped bearer aimed at a named server (path scope) does not enter the + aggregate arm; it is treated as an ordinary bearer on that server.""" + token = self._access_token() + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + side_effect=ProxyException(message="bad key", type="auth_error", param="api_key", code=401), + ) as mock_auth, + ): + with pytest.raises((HTTPException, ProxyException)): + await MCPRequestHandler.process_mcp_request(self._scope(token, path="/mcp/github")) + mock_auth.assert_called_once() + + +@pytest.mark.asyncio +class TestUserSubjectTeamUnion: + """_get_allowed_mcp_servers_for_team unions across ALL a user's teams for a keyless + user-subject caller (the gateway DCR session bearer and bridge user-envelope), while a + key-based caller keeps its single-team behavior byte-identically.""" + + def _team(self, team_id, mcp_servers, members=("sso-user",), tool_perms=None): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member + + return LiteLLM_TeamTable( + team_id=team_id, + members_with_roles=[Member(user_id=u, role="user") for u in members], + access_group_ids=[], + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id=f"op-{team_id}", mcp_servers=mcp_servers, mcp_tool_permissions=tool_perms + ), + ) + + @contextlib.contextmanager + def _patch(self, *, teams_by_id, user_teams=None, orgs_by_id=None): + async def _get_team_object(team_id, **kw): + return teams_by_id.get(team_id) + + async def _get_user_object(user_id, **kw): + return MagicMock(user_id=user_id, teams=user_teams or []) + + async def _get_org_object(org_id, **kw): + return (orgs_by_id or {}).get(org_id) + + async def _spend_from_fallback(counter_key, fallback_spend, max_budget=None, **kw): + # The budget owners read cross-pod spend Redis-first with the row's spend as fallback; + # unit tests have no Redis, so the fallback IS the spend. + return fallback_spend + + with ( + patch("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object), + patch("litellm.proxy.auth.auth_checks.get_user_object", _get_user_object), + patch("litellm.proxy.auth.auth_checks.get_org_object", _get_org_object), + patch("litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[])), + patch("litellm.proxy.proxy_server.get_current_spend", _spend_from_fallback), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ): + yield + + @staticmethod + def _admitted_subject(user_id): + auth = UserAPIKeyAuth(user_id=user_id, api_key=None) + auth.mcp_admitted_user_subject = True + return auth + + async def test_keyless_user_unions_servers_across_all_their_teams(self): + teams = {"team-a": self._team("team-a", ["srv1", "srv2"]), "team-b": self._team("team-b", ["srv2", "srv3"])} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1", "srv2", "srv3"} + + async def test_key_based_caller_uses_single_team_only(self): + """A key-based caller (api_key set) with a team_id sees ONLY that team, even though the + same user belongs to other teams: key auth must be byte-identical to before.""" + teams = {"team-a": self._team("team-a", ["srv1"]), "team-b": self._team("team-b", ["srv2", "srv3"])} + auth = UserAPIKeyAuth(user_id="sso-user", api_key="sk-hash", team_id="team-a") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert set(result) == {"srv1"} + + async def test_keyless_user_with_explicit_team_id_uses_that_team_only(self): + """A keyless caller that already pins a team_id (not the user-subject fan-out shape) + resolves only that team; the union is strictly for the no-team-id user-subject case.""" + teams = {"team-a": self._team("team-a", ["srv1"]), "team-b": self._team("team-b", ["srv2"])} + auth = UserAPIKeyAuth(user_id="sso-user", api_key=None, team_id="team-a") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert set(result) == {"srv1"} + + async def test_keyless_user_with_no_teams_gets_nothing_from_teams(self): + auth = self._admitted_subject("lonely-user") + with self._patch(teams_by_id={}, user_teams=[]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert result == [] + + async def test_ui_session_team_id_still_resolves_to_nothing(self): + from litellm.proxy._types import UI_TEAM_ID + + auth = UserAPIKeyAuth(user_id="dash-user", api_key="sk-hash", team_id=UI_TEAM_ID) + with self._patch(teams_by_id={}, user_teams=["team-a"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert result == [] + + async def test_team_ids_helper_gates_on_shape(self): + from litellm.proxy._types import UI_TEAM_ID + + # key-based with team -> that team + assert await MCPRequestHandler._team_ids_for_mcp_grant( + UserAPIKeyAuth(api_key="sk", team_id="t1", user_id="u") + ) == ["t1"] + # An admitted subject never fans out HERE: it resolves one source per team first, and each of + # those pins a team_id, so this helper only ever answers the single-team question. The fan-out + # itself is _admitted_subject_sources' job, asserted below. + with self._patch(teams_by_id={}, user_teams=["t2", "t3"]): + assert await MCPRequestHandler._team_ids_for_mcp_grant(self._admitted_subject("u")) == [] + # keyless, no user_id -> nothing + assert await MCPRequestHandler._team_ids_for_mcp_grant(UserAPIKeyAuth(api_key=None)) == [] + # keyless with a user_id but NOT admission-marked (JWT auth) -> nothing (unchanged behavior) + with self._patch(teams_by_id={}, user_teams=["t2", "t3"]): + assert ( + await MCPRequestHandler._team_ids_for_mcp_grant(UserAPIKeyAuth(api_key=None, user_id="jwt-user")) == [] + ) + # UI sentinel -> nothing + assert ( + await MCPRequestHandler._team_ids_for_mcp_grant( + UserAPIKeyAuth(api_key="sk", team_id=UI_TEAM_ID, user_id="u") + ) + == [] + ) + + async def test_org_outage_is_not_treated_as_a_missing_org(self): + """A CONFIRMED-absent org places no ceiling; a FAILED lookup must not be read as the same + fact. get_org_object used to relabel every error as "doesn't exist", so a DB outage silently + dropped a real org's ceiling for as long as it lasted. Absent -> the team's grant stands; + outage -> the keyless source denies.""" + from litellm.proxy.auth.auth_checks import OrganizationNotFoundError + + teams = {"t1": self._team("t1", ["srv1"])} + teams["t1"].organization_id = "org-a" + auth = self._admitted_subject("sso-user") + + absent = AsyncMock(side_effect=OrganizationNotFoundError("Organization doesn't exist in db.")) + with self._patch(teams_by_id=teams, user_teams=["t1"]): + with patch("litellm.proxy.auth.auth_checks.get_org_object", absent): + reachable = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(reachable) == {"srv1"}, "a deleted org places no ceiling" + + outage = AsyncMock(side_effect=RuntimeError("connection reset by peer")) + with self._patch(teams_by_id=teams, user_teams=["t1"]): + with patch("litellm.proxy.auth.auth_checks.get_org_object", outage): + reachable = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert reachable == [], "an unresolvable ceiling must deny a keyless source, not be skipped" + + async def test_org_ceiling_fault_fails_closed_for_admitted_but_open_for_keys(self): + """An unresolvable org ceiling is NOT the same fact as "this org places no restriction". + + For a virtual key the ceiling is one of several bounds and a DB blip must not lock working + keys out, so it stays fail-open. For a keyless admitted subject the per-source org ceiling is + the ONLY org bound, so dropping it on a fault would widen a cross-org user to servers their + team's org forbids. That is escalation, not an availability blip, so it fails closed.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + boom = AsyncMock(side_effect=RuntimeError("org lookup exploded")) + # The subject must actually REACH something, or the assertion passes either way and pins + # nothing (a fail-open mutant survived an earlier version of this test for exactly that). + auth = self._admitted_subject("sso-user") + auth.org_id = "org-a" + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) + with self._patch(teams_by_id={}, user_teams=[]): + assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"srv1"} # control + with patch.object(MCPRequestHandler, "_get_org_object_permission", boom): + admitted = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert admitted == [], "admitted subject must fail CLOSED when its org ceiling cannot resolve" + + key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", team_id="t1", org_id="org-a") + with self._patch(teams_by_id={"t1": self._team("t1", ["srv1"])}, user_teams=[]): + with patch.object(MCPRequestHandler, "_get_org_object_permission", boom): + keyed = await MCPRequestHandler.get_allowed_mcp_servers(key_auth) + assert set(keyed) == {"srv1"}, "key auth must keep its long-standing fail-open behavior" + + async def test_only_the_attributing_team_bucket_is_charged(self): + """A team's mcp_rpm_limit bounds that team's SHARED bucket. Charging every granting team let + one cross-team user drain several teams' buckets on a single call, blocking their other + members for access those teams did not provide. Exactly one source is charged, and it is the + SAME source billing picks — one owner for both, so they cannot disagree.""" + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + + t1 = self._team("t1", ["srv1"]) + t1.metadata = {"mcp_rpm_limit": {"srv1": 5}} + t2 = self._team("t2", ["srv1"]) + t2.metadata = {"mcp_rpm_limit": {"srv1": 9}} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id={"t1": t1, "t2": t2}, user_teams=["t1", "t2"]): + auth.mcp_source_team_rpm_limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + billed = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + assert auth.mcp_source_team_rpm_limits == {"t1": {"srv1": 5}}, "t2's shared bucket is untouched" + assert billed is not None and billed.team_id == "t1", "throttling and billing pick the same source" + + descriptors: list = [] + limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) + limiter._add_mcp_per_team_rate_limit_descriptor(auth, "srv1", descriptors) + charged = {d["value"]: d["rate_limit"]["requests_per_unit"] for d in descriptors} + assert charged == {"t1:srv1": 5}, "only the attributing team's bucket is charged" + + async def test_direct_user_grant_charges_no_team_bucket(self): + """When the user's OWN grant reaches the server, no team provided the access, so no team + bucket may be charged — the user's own rpm/tpm is what bounds them. Mirrors billing, which + bills the user and their own org for exactly this case.""" + t1 = self._team("t1", ["srv1"]) + t1.metadata = {"mcp_rpm_limit": {"srv1": 5}} + auth = self._admitted_subject("sso-user") + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) + with self._patch(teams_by_id={"t1": t1}, user_teams=["t1"]): + limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + billed = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + assert limits is None, "a direct user grant must not charge any team's shared bucket" + assert billed is None, "and billing agrees: the user is billed, not a team" + + def _manager_with(self, server_ids, allow_all=()): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.types.mcp import MCPTransport + + manager = MCPServerManager() + for sid in server_ids: + manager.registry[sid] = MCPServer( + server_id=sid, + name=sid, + server_name=sid, + url="https://example.com/mcp", + transport=MCPTransport.http, + allow_all_keys=sid in allow_all, + ) + manager._get_active_submitted_mcp_server_ids_for_user = AsyncMock(return_value=[]) + return manager + + async def test_team_derived_call_bills_the_granting_team_and_its_org(self): + """ACCOUNTING half of team budgets. Without attribution the admitted auth kept team_id=None, + so spend skipped team updates (the team's budget never accumulated, so it could never begin + to block) and charged the user's PRIMARY org rather than the org owning the granting team.""" + t_grant = self._team("t-grant", ["srv1"]) + t_grant.organization_id = "org-team" + auth = self._admitted_subject("sso-user") + auth.org_id = "org-user-primary" + with self._patch(teams_by_id={"t-grant": t_grant}, user_teams=["t-grant"]): + source = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + assert source is not None and source.team_id == "t-grant" + assert source.org_id == "org-team", "the granting team's org is charged, not the user's primary" + assert auth.team_id is None and auth.org_id == "org-user-primary", "authz object untouched" + + async def test_billing_auth_carries_team_and_org_onto_the_spend_object(self): + """Asserted on billing_auth_for_tool_call itself, not on the source it picks: the source + already carries the team's org by construction, so asserting there leaves the copy step + unpinned (a mutant dropping org_id survived exactly that). This is the object spend reads.""" + t_grant = self._team("t-grant", ["srv1"]) + t_grant.organization_id = "org-team" + auth = self._admitted_subject("sso-user") + auth.org_id = "org-user-primary" + server = MagicMock(server_id="srv1") + with self._patch(teams_by_id={"t-grant": t_grant}, user_teams=["t-grant"]): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager._get_mcp_server_from_tool_name", + MagicMock(return_value=server), + ): + billed = await MCPRequestHandler.billing_auth_for_tool_call(auth, tool_name="t-grant/tool_a") + assert (billed.team_id, billed.org_id) == ("t-grant", "org-team") + assert (auth.team_id, auth.org_id) == (None, "org-user-primary"), "authz object must be untouched" + + async def test_own_grant_bills_the_user_not_a_team(self): + """A server the user's OWN grant reaches is not reached "through a team", so it bills the + user and their own org — attributing it to an unrelated team the user happens to belong to + would charge that team for access it never provided.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + t_other = self._team("t-other", ["srv1"]) + auth = self._admitted_subject("sso-user") + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) + with self._patch(teams_by_id={"t-other": t_other}, user_teams=["t-other"]): + assert await MCPRequestHandler.attributing_source_for_server(auth, "srv1") is None + + async def test_billing_attribution_is_deterministic_across_several_granting_teams(self): + """When several teams grant the same server the pick must be stable and reproducible rather + than dependent on dict/roster ordering, or the same call bills different teams run to run.""" + teams = {"t-b": self._team("t-b", ["srv1"]), "t-a": self._team("t-a", ["srv1"])} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["t-b", "t-a"]): + first = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + with self._patch(teams_by_id=teams, user_teams=["t-a", "t-b"]): + second = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + assert first is not None and first.team_id == "t-a" + assert second is not None and second.team_id == "t-a", "roster order must not change who is billed" + + async def test_billing_auth_leaves_non_admitted_callers_untouched(self): + """Key and JWT billing must be byte-identical: the attribution wrapper returns the very same + object for anything that is not a keyless admitted subject.""" + key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", team_id="t1", org_id="org-a") + assert await MCPRequestHandler.billing_auth_for_tool_call(key_auth, tool_name="srv1-tool") is key_auth + + async def test_admitted_tools_never_run_the_single_credential_prelude(self): + """ORDERING is the invariant: the admitted branch is the FIRST statement of the tools + resolver, exactly as in the servers resolver. A fault in a lookup the subject never uses + (its own mcp_toolsets) must not reach it at all — when this branch sat after the prelude, + such a fault hit the fail-closed handler and denied tools its teams did grant.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"t1": self._team("t1", ["srv1"], tool_perms={"srv1": ["read"]})} + auth = self._admitted_subject("sso-user") + # The subject must carry a toolset, or the prelude never resolves one and the fault below is + # unreachable — the branch could sit anywhere and the test would still pass (it did). + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_toolsets=["ts-1"]) + boom = AsyncMock(side_effect=RuntimeError("toolset resolution exploded")) + with self._patch(teams_by_id=teams, user_teams=["t1"]): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.resolve_toolset_tool_permissions", + boom, + ): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + # The fault DOES fire, correctly, inside the subject's own source (which carries its + # toolsets) — that source contributes nothing. What must not happen is the top-level + # prelude running it first and denying the team's grant through the fail-closed handler. + assert tools == ["read"], "a fault in the subject's own toolsets must not deny its team's tools" + + async def test_admitted_own_byom_servers_stay_open(self): + """BYOM suppression-by-explicit-scope is a rule about a CREDENTIAL carrying its own + mcp_servers list. An admitted subject's object_permission is the user's own row, whose + mcp_servers column is [] by DB default — applying the rule would hide almost every admitted + user's OWN submitted servers. A key with an explicit scope still gets no BYOM widening.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + manager = self._manager_with(["srv-byom"]) + manager._get_active_submitted_mcp_server_ids_for_user = AsyncMock(return_value=["srv-byom"]) + db_default_perm = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=[]) + + admitted = self._admitted_subject("sso-user") + admitted.object_permission = db_default_perm + scoped_key = UserAPIKeyAuth(user_id="u", api_key="sk-hash", object_permission=db_default_perm) + + assert await manager.operator_open_server_ids(admitted) == {"srv-byom"} + assert await manager.operator_open_server_ids(scoped_key) == set(), "explicit key scope still suppresses BYOM" + + async def test_admitted_admin_is_scoped_to_grants_not_full_registry(self): + """The wrapper's admin short-circuit hands the FULL registry to any admin-role auth before + the grant union or the per-team org ceilings run. A session bearer is a third-party client + credential, not the dashboard: an admin signing in through the connect flow gets their + grants like anyone else. A real admin key keeps the dashboard behavior unchanged.""" + from litellm.proxy._types import LitellmUserRoles + + manager = self._manager_with(["srv-granted", "srv-secret"]) + admitted = self._admitted_subject("admin-user") + admitted.user_role = LitellmUserRoles.PROXY_ADMIN + with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-granted"])): + admitted_view = set(await manager.get_allowed_mcp_servers(admitted)) + key_admin_view = set( + await manager.get_allowed_mcp_servers( + UserAPIKeyAuth(user_id="admin-user", api_key="sk-hash", user_role=LitellmUserRoles.PROXY_ADMIN) + ) + ) + assert admitted_view == {"srv-granted"}, "an admitted admin gets their grants, not the registry" + assert key_admin_view == {"srv-granted", "srv-secret"}, "admin KEY behavior must be unchanged" + + async def test_admitted_opt_out_via_wrapper_keeps_team_servers(self): + """The wrapper's no_mcp_servers early-return is a KEY rule (a scoped credential's opt-out is + absolute). The admitted subject's opt-out silences only its own source, which the resolver + enforces per source — the wrapper must defer to it, or the resolver-level rule is dead code + on the production path.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, SpecialMCPServerNames + + manager = self._manager_with(["srv-team"]) + opt_out = LiteLLM_ObjectPermissionTable( + object_permission_id="op-u", mcp_servers=[SpecialMCPServerNames.no_mcp_servers.value] + ) + admitted = self._admitted_subject("sso-user") + admitted.object_permission = opt_out + with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-team"])): + admitted_view = set(await manager.get_allowed_mcp_servers(admitted)) + key_view = await manager.get_allowed_mcp_servers( + UserAPIKeyAuth(user_id="u", api_key="sk-hash", object_permission=opt_out) + ) + assert "srv-team" in admitted_view, "user opt-out must not zero team grants on the wrapper path" + assert key_view == [], "a key's opt-out stays absolute" + + async def test_open_channel_confers_reachability_not_a_ceiling_waiver(self): + """An open channel (allow_all_keys / own BYOM) makes a server REACHABLE. It is not a waiver + of the ceilings that bound it: the user's own mcp_tool_permissions still apply, exactly as a + virtual key's key_tools do on the same allow_all server. Returning None outright let a + session holder invoke tools their own policy excludes.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + auth = self._admitted_subject("sso-user") + # The user is restricted to `read` on srv-open, and NO grant source names that server — + # it is reachable only through the open channel, which is exactly the bypass path. + auth.object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="op-u", mcp_servers=[], mcp_tool_permissions={"srv-open": ["read"]} + ) + open_ids = AsyncMock(return_value={"srv-open"}) + with self._patch(teams_by_id={}, user_teams=[]): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.operator_open_server_ids", + open_ids, + ): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-open", auth) + assert tools == ["read"], "the user's own tool policy must still bind on an open-channel server" + + async def test_open_channel_server_gets_default_open_tools_for_admitted(self): + """A server reachable through an open channel (allow_all_keys / own BYOM) is granted by NO + source, so the source union alone returns [] — listable but uninvokable. The tools axis asks + the same open-channel owner the server union uses, so the server is default-open for tools + exactly as a virtual key experiences it.""" + auth = self._admitted_subject("sso-user") + open_ids = AsyncMock(return_value={"srv-open"}) + with self._patch(teams_by_id={}, user_teams=[]): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.operator_open_server_ids", + open_ids, + ): + open_tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-open", auth) + closed_tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-ungranted", auth) + assert open_tools is None, "open-channel server must be default-open for tools" + assert closed_tools == [], "a server no source or channel grants stays deny-all" + + async def test_over_budget_team_grants_nothing_and_healthy_team_stands(self): + """Budget ENFORCEMENT is the sibling of blocked: a team that has already exceeded its + max_budget is rejected outright for a virtual key pinned to it (common_checks), so it must + not keep granting servers, tools or throttle scope to a keyless union subject either. + Enforced through the SAME owner the key path uses (_team_max_budget_check). Distinct from + budget ATTRIBUTION of new spend, which stays with the user (documented deferral).""" + t_over = self._team("t-over", ["srv1"]) + t_over.max_budget = 10.0 + t_over.spend = 11.0 + t_ok = self._team("t-ok", ["srv2"]) + t_ok.max_budget = 10.0 + t_ok.spend = 1.0 + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id={"t-over": t_over, "t-ok": t_ok}, user_teams=["t-over", "t-ok"]): + servers = set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) + limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + assert servers == {"srv2"}, "an over-budget team must stop granting; the healthy team stands" + assert limits is None, "an over-budget team is not a source, so it stamps no throttle either" + + async def test_team_in_over_budget_org_grants_nothing(self): + """The org axis of the same rule, judged against the TEAM's own org (not the caller's + primary): a team owned by an org over its budget grants nothing, exactly as a key in that + org is rejected by _organization_max_budget_check.""" + t_in_broke_org = self._team("t-b", ["srv1"]) + t_in_broke_org.organization_id = "org-broke" + # object_permission_id=None: the org has NO MCP ceiling, so the source is denied by the + # budget gate alone. A truthy auto-Mock id here made an earlier version of this test pass + # through the org-CEILING fault path with the budget gate deleted — vacuous. + org = MagicMock(object_permission_id=None, litellm_budget_table=MagicMock(max_budget=5.0), spend=9.0) + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id={"t-b": t_in_broke_org}, user_teams=["t-b"], orgs_by_id={"org-broke": org}): + servers = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert servers == [], "a team in an over-budget org must not grant through the union" + + async def test_one_faulting_team_does_not_deny_the_other_sources(self): + """The unit of fault isolation is the SOURCE. One team's row being momentarily unreadable + contributes nothing for THAT team (access only narrows) while the user's own grants and every + other resolvable team stand — it must not collapse the whole union to deny-all on either + axis.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + t_ok = self._team("t-ok", ["srv1"], tool_perms={"srv1": ["read"]}) + auth = self._admitted_subject("sso-user") + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv-own"]) + teams = {"t-ok": t_ok} # t-boom absent from the map -> our patched get_team_object RAISES for it + + async def _team_or_boom(team_id, **kw): + if team_id not in teams: + raise RuntimeError(f"transient DB blip loading {team_id}") + return teams[team_id] + + with self._patch(teams_by_id=teams, user_teams=["t-boom", "t-ok"]): + with patch("litellm.proxy.auth.auth_checks.get_team_object", _team_or_boom): + servers = set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert servers == {"srv-own", "srv1"}, "healthy sources must stand when one team faults" + assert tools == ["read"], "the healthy team's tool grant must survive the other team's fault" + + async def test_key_org_tool_ceiling_fault_keeps_key_restrictions(self): + """Virtual-key tools axis mirrors its servers axis on an unresolvable org ceiling: the org + intersect is SKIPPED and the key's own tool restrictions stand. Letting the fault escape + collapsed the whole resolution to None (allow-all), which is fail-open WIDER than before the + fault — key restrictions must never be dropped by an org lookup blip.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", org_id="org-a") + key_auth.object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="op-k", mcp_servers=["srv1"], mcp_tool_permissions={"srv1": ["read"]} + ) + boom = AsyncMock(side_effect=RuntimeError("org permission load exploded")) + with self._patch(teams_by_id={}, user_teams=[]): + with patch.object(MCPRequestHandler, "_get_org_object_permission", boom): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", key_auth) + assert tools == ["read"], "key tool restrictions must survive an unresolvable org ceiling" + + async def test_team_rpm_limit_binds_only_within_that_teams_grant_scope(self): + """A limit rides the same scope as the access it bounds. A roster team is charged ONLY for + servers its own grant reaches: not for a server the user reaches through a DIFFERENT team + (else this user's calls drain a bucket shared by that team's keys for access the team never + provided), not for map entries beyond its grant, and never when the team is blocked.""" + # t-granting grants srv1 and limits it; also names srv9 in its map, which it does NOT grant. + t_granting = self._team("t-granting", ["srv1"]) + t_granting.metadata = {"mcp_rpm_limit": {"srv1": 5, "srv9": 7}} + # t-other grants only srv2 but retains limit metadata for srv1 -> must not be charged for it. + t_other = self._team("t-other", ["srv2"]) + t_other.metadata = {"mcp_rpm_limit": {"srv1": 3}} + # t-blocked grants srv1 and limits it, but is blocked -> grants nothing, charges nothing. + t_blocked = self._team("t-blocked", ["srv1"]) + t_blocked.metadata = {"mcp_rpm_limit": {"srv1": 2}} + t_blocked.blocked = True + + auth = self._admitted_subject("sso-user") + teams = {"t-granting": t_granting, "t-other": t_other, "t-blocked": t_blocked} + with self._patch(teams_by_id=teams, user_teams=["t-granting", "t-other", "t-blocked"]): + limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + + assert limits == {"t-granting": {"srv1": 5}}, ( + "only the granting team's bucket, and only for the server it grants" + ) + + async def test_non_roster_team_rpm_limit_does_not_apply(self): + """The roster gates grants and throttles through one owner, so a team the user was removed + from neither grants servers nor gets charged for their calls.""" + stale = self._team("t-stale", ["srv1"], members=("someone-else",)) + stale.metadata = {"mcp_rpm_limit": {"srv1": 1}} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id={"t-stale": stale}, user_teams=["t-stale"]): + limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + assert limits is None + + async def test_org_list_caps_a_source_but_never_becomes_a_grant(self): + """The admitted model is a union of GRANTS, so an org allowlist may only narrow what a source + already grants. For a virtual key with no lower-level restriction the org list legitimately + BECOMES the allowed set, and inheriting that arm would hand every admitted user with an + org_id their whole org's server list with no direct or team grant behind it.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + auth = self._admitted_subject("sso-user") + auth.org_id = "org-a" # org allows srv1+srv2; the user and their teams grant NOTHING + org_perm = AsyncMock( + return_value=LiteLLM_ObjectPermissionTable(object_permission_id="op-org-a", mcp_servers=["srv1", "srv2"]) + ) + with self._patch(teams_by_id={}, user_teams=[]): + with patch.object(MCPRequestHandler, "_get_org_object_permission", org_perm): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert result == [], "an org ceiling must not grant servers the user was never granted" + + async def test_tool_ceiling_fails_closed_when_a_SOURCE_faults(self): + """Each source is resolved through an UNMARKED auth, so a fault under a source must still + deny. Returning None there would win the union as allow-all and drop every team/org tool + ceiling on a DB blip -- the marker alone only covers faults raised before the fan-out.""" + auth = self._admitted_subject("sso-user") + teams = {"t1": self._team("t1", ["srv1"])} + # Fault INSIDE the tool resolution only. Faulting something the server path also uses would + # make the source grant nothing, so the union would return [] without the tool path ever + # running -- the test would pass while pinning nothing (an earlier version did exactly that). + boom = AsyncMock(side_effect=RuntimeError("org tool ceiling exploded")) + with self._patch(teams_by_id=teams, user_teams=["t1"]): + assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["srv1"] # control: granted + with patch.object(MCPRequestHandler, "_apply_agent_and_org_tool_ceilings", boom): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert tools == [], "a source-level fault must deny tools, never collapse to allow-all" + + async def test_own_opt_out_silences_only_that_source_not_the_teams(self): + """no_mcp_servers on the USER's own grants opts that source out. It must not zero the teams: + the sources are independent, so an opt-out on one silences one. (The same sentinel on a + virtual KEY still overrides team inheritance -- that is the key ceiling model, unchanged.)""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, SpecialMCPServerNames + + auth = self._admitted_subject("sso-user") + auth.object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="op-user", mcp_servers=[SpecialMCPServerNames.no_mcp_servers.value] + ) + with self._patch(teams_by_id={"t1": self._team("t1", ["srv1"])}, user_teams=["t1"]): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1"}, "the user's own opt-out must not zero their team's grants" + + async def test_sources_fan_out_per_team_and_drop_non_roster_teams(self): + """The fan-out lives here now. One source per grant source: the user's own grants (no team_id, + carrying their object_permission) plus each team they are a LIVE roster member of. A team that + lingers in the user's cached `teams` array but no longer lists them in members_with_roles is + dropped, which is what revokes access after a team_member_delete the user row hasn't caught up + on. Each team source carries that team's own org, which is what makes the shared resolver apply + the team's owning-org ceiling rather than the caller's home org.""" + teams = { + "t-member": self._team("t-member", ["srv1"], members=("sso-user",)), + "t-stale": self._team("t-stale", ["srv2"], members=("someone-else",)), + } + teams["t-member"].organization_id = "org-a" + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["t-member", "t-stale"]): + sources = await MCPRequestHandler._admitted_subject_sources(auth) + + assert [(s.team_id, s.org_id) for s in sources] == [(None, None), ("t-member", "org-a")] + # The user's own source carries their grants; a team source must NOT, or the team would be + # widened by grants the team never made. + assert sources[0].object_permission is auth.object_permission + assert sources[1].object_permission is None + # Every source is an ordinary caller, so it cannot re-enter the admitted fan-out. + assert all(not s.mcp_admitted_user_subject for s in sources) + # Nothing that meters or elevates the request may ride along onto a per-source clone. + assert all(s.api_key is None and s.user_role is None for s in sources) + + async def test_jwt_keyless_user_without_team_claim_does_not_union(self): + """Regression for the review finding: a JWT-authenticated caller is also keyless with a + user_id and (with no team claim) no team_id, but it is NOT admission-marked, so it must + keep its prior behavior of inheriting no team grants rather than silently gaining the + union across every team the user belongs to.""" + teams = {"team-a": self._team("team-a", ["srv1"]), "team-b": self._team("team-b", ["srv2"])} + jwt_auth = UserAPIKeyAuth(user_id="jwt-user", api_key=None) # no admission marker + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(jwt_auth) + assert result == [] + + async def test_forged_metadata_marker_on_a_real_key_grants_no_union(self): + """Security regression (forged admission marker): the admitted-subject marker is a + server-only ``UserAPIKeyAuth`` field, NOT a metadata key, precisely because virtual-key + metadata is caller-controlled at key creation. A user who sets + ``mcp_admitted_user_subject: true`` in their own key's metadata (api_key present, no + team_id) must NOT be treated as an admitted subject and must gain no cross-team union.""" + teams = {"team-a": self._team("team-a", ["srv1"]), "team-b": self._team("team-b", ["srv2"])} + forged = UserAPIKeyAuth( + user_id="attacker", + api_key="sk-real-key", + metadata={"mcp_admitted_user_subject": True}, # caller-forged marker in key metadata + ) + assert _is_mcp_admitted_user_subject(forged) is False + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + assert await MCPRequestHandler._team_ids_for_mcp_grant(forged) == [] + assert await MCPRequestHandler._get_allowed_mcp_servers_for_team(forged) == [] + + async def test_admitted_subject_team_tool_restriction_binds(self): + """Security regression (team tool restrictions bypassed): a keyless admitted subject whose + granting team restricts ``srv1`` to ``{tool_a}`` must NOT receive allow-all on srv1. The + single-team-id tool lookup returns None (allow-all) for a keyless multi-team user, dropping + the exclusion; the union across granting teams restores it.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member + + team = LiteLLM_TeamTable( + team_id="team-a", + members_with_roles=[Member(user_id="sso-user", role="user")], + access_group_ids=[], + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-team-a", + mcp_servers=["srv1"], + mcp_tool_permissions={"srv1": ["tool_a"]}, + ), + ) + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id={"team-a": team}, user_teams=["team-a"]): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert tools == ["tool_a"] + + async def test_blocked_team_grants_no_servers_to_admitted_subject(self): + """Security regression: a blocked team grants nothing. The central policy gate enforces this + for a key pinned to a single team_id, but a keyless admitted subject unions across ALL its + teams (no team_id), so a blocked team's MCP grants must be dropped at the per-team resolver.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member + + blocked = LiteLLM_TeamTable( + team_id="team-blocked", + blocked=True, + members_with_roles=[Member(user_id="sso-user", role="user")], + access_group_ids=[], + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-blk", mcp_servers=["srv-secret"]), + ) + teams = {"team-ok": self._team("team-ok", ["srv-ok"]), "team-blocked": blocked} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-ok", "team-blocked"]): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv-ok"} + + async def test_admitted_subject_not_on_team_roster_gets_no_grant(self): + """Security regression (membership containment): a keyless subject whose user_id is NOT on a + team's roster inherits nothing from it, even when the team id lingers in the user's (stale or + cached) teams array. The team roster is the source of truth, so a removed or foreign + membership revokes access at the union rather than granting it.""" + teams = {"team-x": self._team("team-x", ["srv-x"], members=("someone-else",))} + auth = self._admitted_subject("sso-user") # in user.teams for team-x, but NOT on its roster + with self._patch(teams_by_id=teams, user_teams=["team-x"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert result == [] + + async def test_tool_resolution_fails_closed_on_db_error(self): + """Security regression: ANY error resolving the tool allowlist for a keyless admitted subject + must DENY the server's tools ([]) rather than collapse to allow-all (None). Patches an await + OUTSIDE the multi-team fan-out (the team-object lookup) to prove the whole function fails + closed, not just the one helper — mirroring the fail-closed server path.""" + auth = self._admitted_subject("sso-user") + with patch.object( + MCPRequestHandler, + "_get_team_object_permission", + new=AsyncMock(side_effect=RuntimeError("db blip")), + ): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert tools == [] + + async def test_admission_marker_cannot_be_set_from_validated_input(self): + """Defense-in-depth: the mcp_admitted_user_subject marker is server-only. Supplying it in any + validated input (constructor kwargs OR model_validate, e.g. a future JWT/key claim splat) is + stripped by the before-validator, so ONLY the admission path's post-construction assignment + can set it.""" + via_kwarg = UserAPIKeyAuth(user_id="u", api_key=None, mcp_admitted_user_subject=True) + via_validate = UserAPIKeyAuth.model_validate({"user_id": "u", "mcp_admitted_user_subject": True}) + assert via_kwarg.mcp_admitted_user_subject is False + assert via_validate.mcp_admitted_user_subject is False + assert _is_mcp_admitted_user_subject(via_kwarg) is False + assert _is_mcp_admitted_user_subject(via_validate) is False + + +@pytest.mark.asyncio +class TestAdmittedSubjectPerTeamOrgCap: + """A keyless admitted subject unions grants across teams that may span organizations. Each team's + grant (servers AND tools) is capped by that team's OWN org, and the user's direct grants by the + user's own org — never the caller's primary org applied over the whole cross-org union. Guards the + Veria 'team grants bypass their owning policies' finding.""" + + def _team(self, team_id, mcp_servers, *, org_id=None, tool_perms=None, members=("sso-user",)): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member + + return LiteLLM_TeamTable( + team_id=team_id, + organization_id=org_id, + members_with_roles=[Member(user_id=u, role="user") for u in members], + access_group_ids=[], + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id=f"op-{team_id}", + mcp_servers=mcp_servers, + mcp_tool_permissions=tool_perms, + ), + ) + + @staticmethod + def _admitted_subject(user_id, *, org_id=None, own_servers=None, own_tool_perms=None): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + op = None + if own_servers is not None or own_tool_perms is not None: + op = LiteLLM_ObjectPermissionTable( + object_permission_id=f"userop-{user_id}", + mcp_servers=own_servers or [], + mcp_tool_permissions=own_tool_perms, + ) + auth = UserAPIKeyAuth(user_id=user_id, api_key=None, org_id=org_id, object_permission=op) + auth.mcp_admitted_user_subject = True + return auth + + #: sentinel for org_perms: org has an object_permission_id but its load returns None (a swallowed + #: DB error / dangling id), which _object_permission_for_org must treat as fail-closed. + LOAD_FAILS = "__load_fails__" + + @contextlib.contextmanager + def _patch(self, *, teams_by_id, user_teams, org_perms=None, registry=None): + """org_perms: {org_id: LiteLLM_ObjectPermissionTable | None | LOAD_FAILS}. + - table → org exists, ceiling = that permission. + - None → org exists but carries no object_permission (no ceiling). + - LOAD_FAILS → org exists with an object_permission_id, but the permission load returns None. + - org_id ABSENT from the map → org row missing: get_org_object RAISES a bare Exception, exactly + as production does (it does NOT return None or raise HTTPException).""" + org_perms = org_perms or {} + + async def _get_team_object(team_id, **kw): + return teams_by_id.get(team_id) + + async def _get_user_object(user_id, **kw): + return MagicMock(user_id=user_id, teams=user_teams) + + async def _get_org_object(org_id, **kw): + if org_id not in org_perms: + from litellm.proxy.auth.auth_checks import OrganizationNotFoundError + + # matches production: a CONFIRMED-absent org raises this specific type, so callers + # can tell it apart from an outage (a bare Exception now means "lookup failed"). + raise OrganizationNotFoundError(f"Organization doesn't exist. Org={org_id}.") + op = org_perms[org_id] + has_permission_id = op is not None # a table OR LOAD_FAILS carries an id; None does not + return MagicMock( + organization_id=org_id, + object_permission_id=(f"orgop-{org_id}" if has_permission_id else None), + # Real typed values: the budget owners compare these, and a bare MagicMock attribute + # would explode the comparison and silently drop the source (bare-Mock rule). + litellm_budget_table=None, + spend=0.0, + ) + + async def _get_object_permission(object_permission_id, **kw): + for oid, op in org_perms.items(): + if op is not None and op != self.LOAD_FAILS and object_permission_id == f"orgop-{oid}": + return op + return None # LOAD_FAILS (or an unknown id) → None, simulating get_object_permission's swallow + + cms = [ + patch("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object), + patch("litellm.proxy.auth.auth_checks.get_user_object", _get_user_object), + patch("litellm.proxy.auth.auth_checks.get_org_object", _get_org_object), + patch("litellm.proxy.auth.auth_checks.get_object_permission", _get_object_permission), + patch( + "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + AsyncMock(return_value=[]), + ), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ] + if registry is not None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + # registry may be a list of bare server_ids (MagicMock servers) OR a dict of + # {server_id: server_obj} for tests that need real alias/name resolution (config servers). + reg = registry if isinstance(registry, dict) else {s: MagicMock() for s in registry} + cms.append(patch.object(global_mcp_server_manager, "get_registry", return_value=reg)) + with contextlib.ExitStack() as es: + for cm in cms: + es.enter_context(cm) + yield + + # ---- server axis ---- + + async def test_team_grant_capped_by_its_own_org(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a")} + org_perms = {"org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"])} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1"} # srv2 capped out by org-a's ceiling + + async def test_cross_org_teams_each_capped_by_own_org(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = { + "team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a"), + "team-b": self._team("team-b", ["srv3", "srv4"], org_id="org-b"), + } + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"]), + "org-b": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-b", mcp_servers=["srv3"]), + } + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1", "srv3"} # each team clipped by its OWN org, then unioned + + async def test_all_proxy_grant_capped_by_org(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, SpecialMCPServerName + + teams = {"team-a": self._team("team-a", [SpecialMCPServerName.all_proxy_servers.value], org_id="org-a")} + org_perms = {"org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"])} + auth = self._admitted_subject("sso-user") + with self._patch( + teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms, registry=["srv1", "srv2", "srv3"] + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # all_proxy expands to the whole registry, then org-a caps to {srv1} — the cell the old partial + # patch missed (it returned the full registry before capping). + assert set(result) == {"srv1"} + + async def test_org_row_without_object_permission_does_not_cap(self): + teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a")} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": None}): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1", "srv2"} # empty ceiling = no restriction + + async def test_direct_grants_unioned_with_team_and_capped_by_user_org(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"team-a": self._team("team-a", ["srv1"], org_id="org-a")} + org_perms = { + "org-a": None, # the team's org imposes no ceiling + "org-u": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-u", mcp_servers=["srvD", "srv1"]), + } + auth = self._admitted_subject("sso-user", org_id="org-u", own_servers=["srvD", "srvX"]) + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # direct {srvD,srvX} ∩ user-org {srvD,srv1} = {srvD}; UNIONed with team {srv1} (not intersected). + # srvX capped out by the user's org; team's srv1 NOT clipped by the user's primary org. + assert set(result) == {"srvD", "srv1"} + + async def test_single_team_key_uses_primary_org_cap_not_per_team(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + # A KEY (not admitted): the per-team org cap must NOT fire; the top-level primary-org cap applies, + # byte-identical to before. team-a (org-a) grants {srv1,srv2}; the key's primary org is org-k. + teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a")} + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv2"]), + "org-k": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-k", mcp_servers=["srv1"]), + } + key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", team_id="team-a", org_id="org-k") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(key_auth) + # If the per-team (org-a) cap wrongly fired, team-a would clip to {srv2} then org-k → {} (empty). + # Correct key behavior: no per-team cap; primary-org (org-k) cap → {srv1}. + assert set(result) == {"srv1"} + + # ---- tool axis ---- + + async def test_org_tool_ceiling_binds_when_team_places_no_tool_restriction(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + # team grants srv1 with NO tool restriction; org-a restricts srv1's tools to {tool_a}. + teams = {"team-a": self._team("team-a", ["srv1"], org_id="org-a")} + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable( + object_permission_id="orgop-org-a", + mcp_servers=["srv1"], + mcp_tool_permissions={"srv1": ["tool_a"]}, + ) + } + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + # Without the per-team org tool ceiling this would be None (all tools) — org-a's tool ceiling + # would be bypassed exactly like the server case. + assert tools == ["tool_a"] + + async def test_tool_union_across_cross_org_teams(self): + teams = { + "team-a": self._team("team-a", ["srv1"], org_id="org-a", tool_perms={"srv1": ["t1"]}), + "team-b": self._team("team-b", ["srv1"], org_id="org-b", tool_perms={"srv1": ["t2"]}), + } + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"], org_perms={"org-a": None, "org-b": None}): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert set(tools) == {"t1", "t2"} + + async def test_tool_deny_all_when_team_grant_and_org_tool_ceiling_disjoint(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"team-a": self._team("team-a", ["srv1"], org_id="org-a", tool_perms={"srv1": ["t1"]})} + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable( + object_permission_id="orgop-org-a", + mcp_servers=["srv1"], + mcp_tool_permissions={"srv1": ["t2"]}, + ) + } + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + # team {t1} ∩ org {t2} = {} → deny every tool ([]), NOT allow-all (None). + assert tools == [] + + # ---- error contract (adversarial-review findings) ---- + + async def test_missing_org_row_is_treated_as_no_ceiling_not_lockout(self): + """A team's organization_id may point to an org row that no longer exists (deleted / not yet + synced). get_org_object RAISES a bare Exception for that; it must be treated as 'no ceiling' and + must NOT lock the admitted subject out of the team's grants (parity with the key path, which + tolerates a deleted org).""" + teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-gone")} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={}): # org-gone absent → raises + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1", "srv2"} + + async def test_org_permission_load_failure_fails_closed(self): + """The org carries an object_permission_id but the permission load returns None (a swallowed DB + error / dangling id). The ceiling cannot be verified, so the admitted subject must fail CLOSED + for that team — NOT skip the ceiling, which would leak org-forbidden servers.""" + teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a")} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": self.LOAD_FAILS}): + # Asserted through the PUBLIC resolver: the per-source org ceiling is applied there now, + # so calling the single-team helper would return [] for an admitted subject either way + # and pin nothing. + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert result == [] # fail closed, not {srv1, srv2} + + # ---- open bot-thread findings (2026-07-21 re-review) ---- + + async def test_org_less_team_grant_capped_by_user_primary_org(self): + """HIGH (cursor): a team with NO organization_id must still be bounded by the user's PRIMARY + org — otherwise, since admitted subjects skip the top-level primary-org cap, an org-less team's + grant would bypass every org ceiling and reach servers the user's home org forbids.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"team-noorg": self._team("team-noorg", ["srv1", "srv2"], org_id=None)} + org_perms = {"org-U": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-U", mcp_servers=["srv1"])} + auth = self._admitted_subject("sso-user", org_id="org-U") + with self._patch(teams_by_id=teams, user_teams=["team-noorg"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # org-less team falls back to the user's primary org (org-U → {srv1}); srv2 capped out. + assert set(result) == {"srv1"} + + async def test_tool_empty_contributions_fails_closed(self): + """MEDIUM (greptile/cursor): when no source in the tool-resolution view grants the server (a + TOCTOU/cache-lag inconsistency on a server that passed the server gate), the admitted path must + fail CLOSED (deny all tools = []), NOT allow-all (None).""" + teams = {"team-a": self._team("team-a", ["srv1"], org_id="org-a")} + auth = self._admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": None}): + # 'srv-nobody' is granted by neither the team nor the user directly → empty contributions. + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-nobody", auth) + assert tools == [] + + async def test_tool_no_db_honors_in_memory_direct_restriction(self): + """MEDIUM (cursor): with no DB, the tool path must still honor the user's OWN in-memory + object_permission tool restriction (resolvable without a DB) rather than blanket-allow (None).""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + auth = self._admitted_subject( + "sso-user", own_servers=["srv1"], own_tool_perms={"srv1": ["t1"]} + ) # no org_id, direct grant of srv1 restricted to {t1} + with patch("litellm.proxy.proxy_server.prisma_client", None): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert tools == ["t1"] # in-memory restriction honored, not widened to all tools + + # ---- config.yaml-defined servers (incl. OAuth) ---- + + async def test_config_defined_oauth_server_by_alias_reached_and_org_capped(self): + """A config.yaml-defined MCP OAuth server flows through the SAME resolution as a DB server: + the team grant (and the org ceiling) reference it by ALIAS, expand_permission_list resolves it + via the config+DB registry union to its server_id, and the per-team org cap applies identically. + (The config server's OAuth *client* persistence is #33768 — an orthogonal egress concern; this + pins the grant/reachability side of the 10x flow for config-defined servers.)""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + cfg_server = MagicMock() + cfg_server.server_id = "cfg-oauth-1" + cfg_server.alias = "linear_cfg" + cfg_server.server_name = "linear_cfg" + cfg_server.name = "linear_cfg" + + # team grants the config server BY ALIAS alongside a DB-style bare id; org-a's ceiling lists + # ONLY the config server (also by alias). + teams = {"team-a": self._team("team-a", ["linear_cfg", "srv-db"], org_id="org-a")} + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["linear_cfg"]) + } + auth = self._admitted_subject("sso-user") + with self._patch( + teams_by_id=teams, + user_teams=["team-a"], + org_perms=org_perms, + registry={"cfg-oauth-1": cfg_server}, + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # 'linear_cfg' alias resolves to the config server_id and survives org-a's ceiling; 'srv-db' + # (not in org-a's allowlist) is capped out — same per-team org cap, config server included. + assert set(result) == {"cfg-oauth-1"} + + async def test_config_oauth_server_alias_resolution_feeds_the_org_cap(self): + """A config-defined OAuth server granted BY ALIAS whose OWN org forbids it is capped out — AND the + cap is proven to run on RESOLVED server_ids, not raw strings. A control config server, granted by + alias and allowed by the org via its RESOLVED id, must survive: that inclusion is impossible unless + expand_permission_list resolved the grant alias to the id the ceiling lists, so a broken alias path + yields {} and FAILS this test — whereas a bare `assert empty` would pass even if resolution never + ran (the weakness Cursor flagged).""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + forbidden = MagicMock() # granted by alias, but its org forbids it → must be capped out + forbidden.server_id = "cfg-oauth-1" + forbidden.alias = forbidden.server_name = forbidden.name = "linear_cfg" + control = MagicMock() # granted by alias, allowed by the org via its RESOLVED id → must survive + control.server_id = "control-id" + control.alias = control.server_name = control.name = "control_alias" + + teams = {"team-b": self._team("team-b", ["linear_cfg", "control_alias"], org_id="org-b")} + # org-b's ceiling allows ONLY the control server, referenced by its RESOLVED server_id. + org_perms = { + "org-b": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-b", mcp_servers=["control-id"]) + } + auth = self._admitted_subject("sso-user") + with self._patch( + teams_by_id=teams, + user_teams=["team-b"], + org_perms=org_perms, + registry={"cfg-oauth-1": forbidden, "control-id": control}, + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # control survives ('control_alias' resolved to 'control-id', matching the id-based ceiling); the + # forbidden config server ('cfg-oauth-1') is capped out. A broken alias path → {} → fails here. + assert set(result) == {"control-id"} + + +@pytest.mark.asyncio +class TestSessionBearerEgressScrub: + """The gateway session bearer / bridge envelope is an admission credential, never an upstream token. + The leak-defense scrub is anchored to the credential SHAPE, so a session-shaped Authorization is + stripped from every egress context even when it reaches a non-aggregate scope that never set the + admission marker (design-review finding: a session bearer misdirected to a per-server true_passthrough + path would otherwise be forwarded upstream verbatim and replayed against the aggregate endpoint).""" + + async def test_session_bearer_misdirected_to_passthrough_is_scrubbed(self): + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [(b"authorization", b"Bearer llm_session_synthetic-shaped-token")], + } + ttp_server = MagicMock() + ttp_server.auth_type = MCPAuth.true_passthrough + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ) as mock_auth, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ttp_server + (_auth, _mah, _srv, _sah, oauth2_headers, raw_headers) = await MCPRequestHandler.process_mcp_request(scope) + + mock_auth.assert_not_called() # true_passthrough → LiteLLM auth skipped (anonymous arm, no marker) + assert oauth2_headers is None # session-shaped bearer scrubbed from oauth2 egress + assert all(k.lower() != "authorization" for k in raw_headers) # ...and from raw egress headers + + async def test_legitimate_upstream_token_is_not_scrubbed(self): + """A genuine upstream/passthrough token is never session- or envelope-shaped, so the shape-anchored + scrub must leave it intact for forwarding (guards against over-stripping).""" + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [(b"authorization", b"Bearer real-upstream-opaque-token-xyz")], + } + ttp_server = MagicMock() + ttp_server.auth_type = MCPAuth.true_passthrough + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ttp_server + (_auth, _mah, _srv, _sah, oauth2_headers, _raw) = await MCPRequestHandler.process_mcp_request(scope) + + assert oauth2_headers.get("Authorization") == "Bearer real-upstream-opaque-token-xyz" + + async def test_scrub_removes_gateway_credential_from_every_egress_context(self): + """The scrub is anchored to the credential SHAPE and covers ALL egress contexts, not just + Authorization: a session bearer placed in x-mcp-auth OR a per-server x-mcp-{alias}-authorization + header is stripped too (the High-severity gap: those were forwarded upstream before).""" + sess = "Bearer llm_session_abc" + oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": sess}, + raw_headers={ + "authorization": sess, + "x-mcp-auth": "llm_session_xyz", + "x-mcp-github-authorization": "llm_session_ghi", + }, + mcp_auth_header="llm_session_xyz", + mcp_server_auth_headers={"github": {"Authorization": "llm_session_ghi"}}, + ) + assert oauth2 is None + assert "authorization" not in {k.lower() for k in raw} + assert all("llm_session_" not in v for v in raw.values()) # x-mcp-auth + per-server raw values gone + assert mcp_auth is None # deprecated x-mcp-auth value scrubbed + assert per_server == {} # per-server session bearer removed → now-empty server dict dropped + + async def test_scrub_keeps_real_upstream_tokens(self): + """A legitimate upstream token is never session-/envelope-shaped, so every context is forwarded + unchanged — guards against over-stripping a real credential the caller meant for the upstream.""" + oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": "Bearer real-upstream-xyz"}, + raw_headers={"authorization": "Bearer real-upstream-xyz", "x-mcp-github-authorization": "Bearer gh_real"}, + mcp_auth_header="some-api-key-123", + mcp_server_auth_headers={"github": {"Authorization": "Bearer gh_real"}}, + ) + assert oauth2 == {"Authorization": "Bearer real-upstream-xyz"} + assert raw["authorization"] == "Bearer real-upstream-xyz" + assert mcp_auth == "some-api-key-123" + assert per_server == {"github": {"Authorization": "Bearer gh_real"}} + + async def test_scrub_admitted_drops_authorization_but_keeps_injected_upstream_token(self): + """An admitted subject's top-level Authorization is dropped unconditionally, while the real + upstream token the bridge arm INJECTS into a per-server header (not gateway-shaped) survives.""" + oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + admitted=True, + oauth2_headers={"Authorization": "Bearer llm_session_abc"}, + raw_headers={"authorization": "Bearer llm_session_abc"}, + mcp_auth_header=None, + mcp_server_auth_headers={"github": {"Authorization": "Bearer gh_injected_upstream"}}, + ) + assert oauth2 is None + assert "authorization" not in {k.lower() for k in raw} + assert per_server == {"github": {"Authorization": "Bearer gh_injected_upstream"}} 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 692e5340f48..45d30244cef 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 @@ -3302,19 +3302,20 @@ async def test_token_root_does_not_resolve_private_server_for_external_client(): @pytest.mark.asyncio -async def test_register_root_resolves_single_oauth2_server(): - """When /register is hit without server name and exactly 1 OAuth2 server exists, resolve it.""" - try: - from fastapi import Request +async def test_register_root_does_aggregate_dcr_not_single_server_resolution(): + """Root /register is the aggregate DCR endpoint: it mints a stateless llm_dcrc_ client + from the request's redirect_uris and does NOT resolve a single configured oauth2 server + (a single-server deployment registers at /{server}/register instead).""" + import json - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - register_client, - ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - except ImportError: - pytest.skip("MCP discoverable endpoints not available") + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) global_mcp_server_manager.registry.clear() oauth2_server = _create_oauth2_server() @@ -3325,33 +3326,37 @@ async def test_register_root_resolves_single_oauth2_server(): mock_request.headers = {} try: - with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", - new=AsyncMock(return_value={}), + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}), + ), + patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"), ): - result = await register_client(request=mock_request, mcp_server_name=None) + response = await register_client(request=mock_request, mcp_server_name=None) - # Should resolve to the single server and return its name as client_id - assert result["client_id"] == "test_oauth" - assert "redirect_uris" in result + body = json.loads(response.body) + assert body["client_id"].startswith("llm_dcrc_") + assert body["client_id"] != "test_oauth" + assert body["token_endpoint_auth_method"] == "none" finally: global_mcp_server_manager.registry.clear() @pytest.mark.asyncio -async def test_register_root_does_not_resolve_private_server_for_external_client(): - """Root /register must not reveal or use a hidden MCP server.""" - try: - from fastapi import Request +async def test_register_root_does_not_leak_a_private_server(): + """Root /register never resolves or reveals a configured server, so a private one cannot + leak to an external caller: it always mints the aggregate DCR client instead.""" + import json - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - register_client, - ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - except ImportError: - pytest.skip("MCP discoverable endpoints not available") + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) global_mcp_server_manager.registry.clear() oauth2_server = _create_oauth2_server(available_on_public_internet=False) @@ -3365,17 +3370,19 @@ async def test_register_root_does_not_resolve_private_server_for_external_client with ( patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", - new=AsyncMock(return_value={}), + new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}), ), patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip", return_value="198.51.100.10", ), + patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"), ): - result = await register_client(request=mock_request, mcp_server_name=None) + response = await register_client(request=mock_request, mcp_server_name=None) - assert result["client_id"] == "dummy_client" - assert result["redirect_uris"] == ["https://llm.example.com/callback"] + body = json.loads(response.body) + assert body["client_id"].startswith("llm_dcrc_") + assert "test_oauth" not in body["client_id"] finally: global_mcp_server_manager.registry.clear() @@ -5155,7 +5162,10 @@ async def test_bridge_refresh_grant_with_non_envelope_is_invalid_grant_before_up def _mint_test_refresh_envelope( - server_id="bridge_srv", key_hash="hashed-litellm-key-77", upstream_refresh="UPSTREAM-REFRESH", identity=None, + server_id="bridge_srv", + key_hash="hashed-litellm-key-77", + upstream_refresh="UPSTREAM-REFRESH", + identity=None, scope=None, ): """Mint a refresh envelope the way the producer does, for driving the refresh_token grant in tests. @@ -5178,7 +5188,9 @@ def _mint_test_refresh_envelope( keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY) identity = identity if identity is not None else key_hash_identity(server_id=server_id, key_hash=key_hash) sealed = build_bridge_refresh_token_response( - identity, RefreshCredential(refresh_token=SecretStr(upstream_refresh), scope=scope), keys, + identity, + RefreshCredential(refresh_token=SecretStr(upstream_refresh), scope=scope), + keys, datetime.now(timezone.utc), ) assert isinstance(sealed, SealedEnvelope) @@ -5413,7 +5425,10 @@ async def test_bridge_refresh_re_requests_the_sealed_scope_when_client_omits_it( ) captured: dict = {} response = await _refresh_for_bridge_server( - server, refresh_env, {"access_token": "NEW-ACCESS", "token_type": "Bearer", "expires_in": 3600}, None, + server, + refresh_env, + {"access_token": "NEW-ACCESS", "token_type": "Bearer", "expires_in": 3600}, + None, fake_client_out=captured, ) @@ -5553,7 +5568,9 @@ async def test_bridge_refresh_upstream_invalid_grant_maps_to_invalid_grant(): error_response = MagicMock() error_response.status_code = 400 error_response.text = '{"error": "invalid_grant", "error_description": "refresh token expired"}' - error_response.json = MagicMock(return_value={"error": "invalid_grant", "error_description": "refresh token expired"}) + error_response.json = MagicMock( + return_value={"error": "invalid_grant", "error_description": "refresh token expired"} + ) error_response.raise_for_status = MagicMock( side_effect=httpx.HTTPStatusError("bad", request=MagicMock(), response=error_response) ) @@ -7183,7 +7200,9 @@ def _upstream_token_response(status_code: int, *, json_body: object = None, text return httpx.Response(status_code, text=text_body, request=request) -async def _exchange_with_upstream_response(upstream_response, *, server_client_id="web-client.apps.googleusercontent.com"): +async def _exchange_with_upstream_response( + upstream_response, *, server_client_id="web-client.apps.googleusercontent.com" +): """Run the raw (non-bridge) authorization_code exchange against a canned upstream token-endpoint response and return what the gateway would hand the client. ``server_client_id=None`` models the caller-supplied-credentials flow (no stored client on the server).""" @@ -7334,9 +7353,7 @@ async def test_token_exchange_bounds_relayed_error_fields(): async def test_token_exchange_200_without_access_token_is_502_not_keyerror(): """A 200 whose body has no usable access_token used to KeyError into a 500; the raw arm now answers 502 with the same wording as the bridge arm's no_upstream_token rejection.""" - response = await _exchange_with_upstream_response( - _upstream_token_response(200, json_body={"token_type": "Bearer"}) - ) + response = await _exchange_with_upstream_response(_upstream_token_response(200, json_body={"token_type": "Bearer"})) assert response.status_code == 502 body = json.loads(response.body) @@ -7357,7 +7374,9 @@ async def test_token_exchange_relays_rejection_when_http_client_raises(): ) raising_client = MagicMock() raising_client.post = AsyncMock( - side_effect=httpx.HTTPStatusError("Client error '401 Unauthorized'", request=rejection.request, response=rejection) + side_effect=httpx.HTTPStatusError( + "Client error '401 Unauthorized'", request=rejection.request, response=rejection + ) ) from fastapi import Request @@ -7422,7 +7441,9 @@ async def test_register_relays_rejection_when_http_client_raises(): ) raising_client = MagicMock() raising_client.post = AsyncMock( - side_effect=httpx.HTTPStatusError("Client error '400 Bad Request'", request=rejection.request, response=rejection) + side_effect=httpx.HTTPStatusError( + "Client error '400 Bad Request'", request=rejection.request, response=rejection + ) ) oauth2_server = _bridge_server(auth_type=MCPAuth.oauth2, dcr_bridge=None) @@ -7800,9 +7821,7 @@ async def test_hydrate_does_not_overwrite_explicit_config_client_id(): auth_type=MCPAuth.oauth2, client_id="explicit-from-config", ) - store_read = AsyncMock( - return_value={"client_id": "stale-store-client", "client_secret": "x", "redirect_uris": []} - ) + store_read = AsyncMock(return_value={"client_id": "stale-store-client", "client_secret": "x", "redirect_uris": []}) with ( patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), patch( @@ -8019,9 +8038,7 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate(): mock_request.headers = {} try: - authorization_response = _build_oauth_authorization_server_response( - request=mock_request, mcp_server_name=None - ) + authorization_response = _build_oauth_authorization_server_response(request=mock_request, mcp_server_name=None) resource_response = await _build_oauth_protected_resource_response( request=mock_request, mcp_server_name=None, use_standard_pattern=True ) @@ -8033,6 +8050,65 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate(): global_mcp_server_manager.registry.clear() +def test_gateway_dcr_flow_routing_engages_only_for_llm_dcrc_clients(monkeypatch): + """The aggregate DCR arms engage for llm_dcrc_ client_ids (register always mints one, + authorize/token route into the aggregate flow); a non-gateway client_id keeps the + per-server behavior, and /authorize/complete exists but 400s without a valid flow.""" + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-for-lit3637") + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637", raising=False) + global_mcp_server_manager.registry.clear() + app = FastAPI() + app.include_router(router) + client = TestClient(app) + + registered = client.post("/register", json={"redirect_uris": ["https://claude.ai/cb"]}) + assert registered.status_code == 201 + assert registered.json()["client_id"].startswith("llm_dcrc_") + assert registered.json()["token_endpoint_auth_method"] == "none" + + authorize_params = { + "client_id": "llm_dcrc_bogus", + "redirect_uri": "https://claude.ai/cb", + "response_type": "code", + "code_challenge": "c" * 43, + "code_challenge_method": "S256", + } + bogus_client = client.get("/authorize", params=authorize_params) + assert bogus_client.status_code == 400 + assert bogus_client.json()["error"] == "invalid_client" + + no_cookie = client.post("/authorize/complete", data={"flow": "h"}) + assert no_cookie.status_code == 400 + assert no_cookie.json()["error"] == "invalid_request" + + token_response = client.post( + "/token", + data={ + "grant_type": "authorization_code", + "client_id": "llm_dcrc_bogus", + "code": "x", + "redirect_uri": "https://claude.ai/cb", + "code_verifier": "v" * 43, + }, + ) + assert token_response.status_code == 400 + assert token_response.json()["error"] == "invalid_grant" + + upstream_shaped = client.post( + "/token", + data={"grant_type": "authorization_code", "client_id": "regular-upstream-client", "code": "x"}, + ) + assert upstream_shaped.status_code == 404 + + @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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py new file mode 100644 index 00000000000..375ec022115 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -0,0 +1,590 @@ +"""Tests for the aggregate gateway DCR flow (register, authorize, complete, token).""" + +import hashlib +import json +from base64 import urlsafe_b64encode +from datetime import datetime, timedelta, timezone +from http.cookies import SimpleCookie +from urllib.parse import parse_qs, urlparse + +import pytest +from starlette.requests import Request + +from litellm.caching.caching import DualCache +from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( + CONNECT_FLOW_COOKIE_PREFIX, + GATEWAY_AUTH_CODE_PREFIX, + GATEWAY_AUTH_CODE_TTL_SECONDS, + GATEWAY_DCR_CLIENT_ID_PREFIX, + _GatewayAuthCode, + _seal, + aggregate_authorize, + aggregate_token, + complete_connect_flow, + is_gateway_dcr_client_id, + open_gateway_dcr_client, + register_aggregate_client, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + resolve_session_bearer, + session_keys_from_master_key, + SessionBearerAdmitted, +) + +MASTER_KEY = "sk-gateway-dcr-flow-tests" +REDIRECT_URI = "https://claude.ai/api/mcp/auth_callback" +CODE_VERIFIER = "verifier-" + "v" * 43 +CODE_CHALLENGE = urlsafe_b64encode(hashlib.sha256(CODE_VERIFIER.encode("ascii")).digest()).rstrip(b"=").decode("ascii") + + +@pytest.fixture(autouse=True) +def _salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", MASTER_KEY) + + +def _request(path="/authorize", query="", cookies=None, method="GET"): + cookie_header = [] + if cookies: + cookie = SimpleCookie() + for name, value in cookies.items(): + cookie[name] = value + cookie_header = [(b"cookie", cookie.output(header="", sep="; ").strip().encode())] + return Request( + { + "type": "http", + "method": method, + "scheme": "https", + "path": path, + "query_string": query.encode(), + "headers": [(b"host", b"llm.example.com"), *cookie_header], + } + ) + + +async def _register(redirect_uris) -> dict: + response = await register_aggregate_client( + request=_request(path="/register", method="POST"), request_body={"redirect_uris": redirect_uris} + ) + return json.loads(response.body) + + +async def _reload_user_active(user_id: str): + return None + + +@pytest.mark.asyncio +async def test_register_mints_stateless_public_client(): + body = await _register([REDIRECT_URI]) + assert body["token_endpoint_auth_method"] == "none" + assert "client_secret" not in body + assert body["redirect_uris"] == [REDIRECT_URI] + assert is_gateway_dcr_client_id(body["client_id"]) + record = open_gateway_dcr_client(body["client_id"]) + assert record is not None + assert record.redirect_uris == (REDIRECT_URI,) + + +@pytest.mark.asyncio +async def test_register_allows_loopback_http_for_dev_clients(): + body = await _register(["http://localhost:6274/oauth/callback"]) + assert is_gateway_dcr_client_id(body["client_id"]) + + +@pytest.mark.parametrize( + "code_challenge", + ["short", "", "p" * 300, "ünïcode-challenge", "AAAA" * 20], +) +def test_pkce_mismatched_challenge_returns_false_never_raises(code_challenge): + """A wrong-length or non-ASCII code_challenge must VERIFY FALSE, not raise. + + Pins the reason this compares bytes rather than str: hmac.compare_digest raises TypeError on + two str with non-ASCII content, but on bytes of unequal length it simply returns False. A + review flagged this as an unhandled 500 on length mismatch; encoding both sides to bytes is + exactly what makes that impossible, so the claim is pinned here rather than in a comment.""" + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _pkce_verifier_matches + + assert _pkce_verifier_matches("a" * 43, code_challenge) is False + + +@pytest.mark.asyncio +async def test_register_allows_allowlisted_native_callback(): + """Native MCP clients register a private-use scheme, not https. Registration shares + the one redirect-URI shape owner with /authorize, so the callback the allowlist + already trusts there is registrable here rather than rejected as non-https.""" + body = await _register(["cursor://anysphere.cursor-mcp/oauth/callback"]) + assert is_gateway_dcr_client_id(body["client_id"]) + record = open_gateway_dcr_client(body["client_id"]) + assert record is not None + assert record.redirect_uris == ("cursor://anysphere.cursor-mcp/oauth/callback",) + + +@pytest.mark.asyncio +async def test_register_rejects_userinfo_spoofed_origin(): + """``https://claude.ai@attacker.example/cb`` parses with netloc + ``claude.ai@attacker.example``, so a naive origin display on the consent screen reads + as claude.ai while the code would be delivered to attacker.example. Rejected at + registration, which is the only way such a URI could enter a sealed client.""" + response = await register_aggregate_client( + request=_request(path="/register", method="POST"), + request_body={"redirect_uris": ["https://claude.ai@attacker.example/callback"]}, + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_redirect_uri" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "redirect_uris", + [ + [], + "not-a-list", + ["http://evil.example.com/callback"], + ["https://claude.ai/cb#fragment"], + ["ftp://claude.ai/cb"], + ["https://a.example.com/" + "p" * 300], + ["https://a.example.com/1", "https://a.example.com/2", "https://a.example.com/3", "https://a.example.com/4"], + [12345], + ], +) +async def test_register_rejects_bad_redirect_uris(redirect_uris): + response = await register_aggregate_client( + request=_request(path="/register", method="POST"), request_body={"redirect_uris": redirect_uris} + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] in ("invalid_redirect_uri", "invalid_client_metadata") + + +@pytest.mark.asyncio +async def test_tampered_client_id_does_not_open(): + body = await _register([REDIRECT_URI]) + tampered = body["client_id"][:-4] + "AAAA" + assert open_gateway_dcr_client(tampered) is None + assert open_gateway_dcr_client("llm_dcrc_garbage") is None + assert open_gateway_dcr_client("other_prefix") is None + + +def _authorize( + client_id, session_user_id, redirect_uri=REDIRECT_URI, challenge=CODE_CHALLENGE, method="S256", response_type="code" +): + return aggregate_authorize( + request=_request(query=f"client_id={client_id}"), + client_id=client_id, + redirect_uri=redirect_uri, + state="client-state-123", + code_challenge=challenge, + code_challenge_method=method, + response_type=response_type, + session_user_id=session_user_id, + ) + + +@pytest.mark.asyncio +async def test_authorize_validation_failures_never_redirect_to_client(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + for response, expected_error in ( + (_authorize("llm_dcrc_bogus", "u1"), "invalid_client"), + (_authorize(client_id, "u1", redirect_uri="https://attacker.example.com/cb"), "invalid_request"), + (_authorize(client_id, "u1", response_type="token"), "unsupported_response_type"), + (_authorize(client_id, "u1", challenge=None), "invalid_request"), + (_authorize(client_id, "u1", method="plain"), "invalid_request"), + ): + assert response.status_code == 400 + assert json.loads(response.body)["error"] == expected_error + + +@pytest.mark.asyncio +async def test_authorize_without_session_redirects_to_login_with_return_to(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = _authorize(client_id, session_user_id=None) + assert response.status_code == 303 + location = response.headers["location"] + assert location.startswith("https://llm.example.com/sso/key/generate?return_to=") + assert "return_to=%2Fauthorize" in location + + +@pytest.mark.asyncio +async def test_authorize_with_session_hands_browser_to_connect_page_with_flow_cookie(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = _authorize(client_id, session_user_id="u1") + assert response.status_code == 303 + location = urlparse(response.headers["location"]) + assert location.path == "/ui/chat/integrations" + params = parse_qs(location.query) + handle = params["connect_flow"][0] + assert params["connect_client"] == ["https://claude.ai"] + set_cookie = response.headers["set-cookie"] + assert f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}" in set_cookie + assert "HttpOnly" in set_cookie + return handle, set_cookie + + +def _flow_cookie_from(response) -> tuple: + location = urlparse(response.headers["location"]) + handle = parse_qs(location.query)["connect_flow"][0] + cookie = SimpleCookie() + cookie.load(response.headers["set-cookie"]) + name = f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}" + return handle, {name: cookie[name].value} + + +@pytest.mark.asyncio +async def test_full_walk_register_authorize_complete_token_and_replay(): + """The whole front door on one deterministic walk: register -> authorize -> + complete -> token, then the security edges on the same artifacts (user mismatch, + PKCE mismatch, single-use replay, refresh rotation, cross-client refresh).""" + client_id = (await _register([REDIRECT_URI]))["client_id"] + authorize_response = _authorize(client_id, session_user_id="u1") + handle, cookies = _flow_cookie_from(authorize_response) + + denied = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="attacker", + cache=DualCache(), + ) + assert denied.status_code == 403 + + anonymous = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id=None, + cache=DualCache(), + ) + assert anonymous.status_code == 401 + + completed = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=DualCache(), + ) + assert completed.status_code == 303 + redirect = urlparse(completed.headers["location"]) + assert f"{redirect.scheme}://{redirect.netloc}{redirect.path}" == REDIRECT_URI + params = parse_qs(redirect.query) + assert params["state"] == ["client-state-123"] + code = params["code"][0] + assert code.startswith(GATEWAY_AUTH_CODE_PREFIX) + + cache = DualCache() + + async def _token(**overrides): + arguments = { + "request": _request("/token", method="POST"), + "grant_type": "authorization_code", + "code": code, + "redirect_uri": REDIRECT_URI, + "client_id": client_id, + "code_verifier": CODE_VERIFIER, + "refresh_token": None, + "master_key": MASTER_KEY, + "reload_user": _reload_user_active, + "cache": cache, + } + return await aggregate_token(**{**arguments, **overrides}) + + wrong_verifier = await _token(code_verifier="wrong-" + "w" * 43) + assert json.loads(wrong_verifier.body)["error"] == "invalid_grant" + + wrong_client = await _token(client_id=(await _register([REDIRECT_URI]))["client_id"]) + assert json.loads(wrong_client.body)["error"] == "invalid_grant" + + token_response = await _token() + assert token_response.status_code == 200 + payload = json.loads(token_response.body) + assert payload["token_type"] == "Bearer" + assert 0 < payload["expires_in"] <= 3600 + + keys = session_keys_from_master_key(MASTER_KEY) + admitted = resolve_session_bearer(f"Bearer {payload['access_token']}", keys, datetime.now(timezone.utc)) + assert isinstance(admitted, SessionBearerAdmitted) + assert admitted.principal.user_id == "u1" + assert admitted.principal.client_id == client_id + + replay = await _token() + assert json.loads(replay.body)["error"] == "invalid_grant" + + refreshed = await _token(grant_type="refresh_token", code=None, refresh_token=payload["refresh_token"]) + assert refreshed.status_code == 200 + rotated = json.loads(refreshed.body) + assert rotated["refresh_token"] != payload["refresh_token"] + + # Rotation is single-use: replaying the now-consumed refresh token cannot mint a second pair + # (a captured token is dead once the legitimate holder has rotated). + replayed = await _token(grant_type="refresh_token", code=None, refresh_token=payload["refresh_token"]) + assert json.loads(replayed.body)["error"] == "invalid_grant" + assert "already used" in json.loads(replayed.body).get("error_description", "") + + cross_client = await _token( + grant_type="refresh_token", + code=None, + refresh_token=payload["refresh_token"], + client_id=(await _register([REDIRECT_URI]))["client_id"], + ) + assert json.loads(cross_client.body)["error"] == "invalid_grant" + + +@pytest.mark.asyncio +async def test_complete_rejects_missing_tampered_and_expired_flows(): + missing = await complete_connect_flow( + request=_request("/authorize/complete", method="POST"), + flow_handle="nope", + session_user_id="u1", + cache=DualCache(), + ) + assert missing.status_code == 400 + + tampered = await complete_connect_flow( + request=_request("/authorize/complete", cookies={f"{CONNECT_FLOW_COOKIE_PREFIX}h1": "garbage"}, method="POST"), + flow_handle="h1", + session_user_id="u1", + cache=DualCache(), + ) + assert tampered.status_code == 400 + + +@pytest.mark.asyncio +async def test_token_rejects_expired_code_and_missing_configuration(): + expired_code = _seal( + GATEWAY_AUTH_CODE_PREFIX, + _GatewayAuthCode( + user_id="u1", + client_id="llm_dcrc_x", + redirect_uri=REDIRECT_URI, + code_challenge=CODE_CHALLENGE, + jti="jti-1", + iat=int((datetime.now(timezone.utc) - timedelta(seconds=500)).timestamp()), + exp=int((datetime.now(timezone.utc) - timedelta(seconds=500 - GATEWAY_AUTH_CODE_TTL_SECONDS)).timestamp()), + ), + ) + response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=expired_code, + redirect_uri=REDIRECT_URI, + client_id="llm_dcrc_x", + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert json.loads(response.body)["error"] == "invalid_grant" + + no_master_key = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code="llm_gcode_x", + redirect_uri=REDIRECT_URI, + client_id="llm_dcrc_x", + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=None, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert no_master_key.status_code == 500 + assert json.loads(no_master_key.body)["error"] == "server_error" + + unsupported = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="password", + code=None, + redirect_uri=None, + client_id="llm_dcrc_x", + code_verifier=None, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert json.loads(unsupported.body)["error"] == "unsupported_grant_type" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "failure,expected_status,expected_error", + [ + ("no_active_key", 400, "invalid_grant"), + ("unavailable", 503, "temporarily_unavailable"), + ("unresolvable", 500, "server_error"), + ], +) +async def test_token_gates_on_live_user_revalidation(failure, expected_status, expected_error): + client_id = (await _register([REDIRECT_URI]))["client_id"] + authorize_response = _authorize(client_id, session_user_id="deactivated-user") + handle, cookies = _flow_cookie_from(authorize_response) + completed = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="deactivated-user", + cache=DualCache(), + ) + code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0] + + async def _reload_user_failing(user_id: str): + return failure + + response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=code, + redirect_uri=REDIRECT_URI, + client_id=client_id, + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_failing, + cache=DualCache(), + ) + assert response.status_code == expected_status + assert json.loads(response.body)["error"] == expected_error + + +@pytest.mark.asyncio +async def test_flow_is_single_use_shared_cache_rejects_second_complete(): + """A double-submit of the finish step mints only ONE code: the second complete over the + same cache fails invalid_request (atomic flow claim), so one sign-in cannot yield two codes.""" + cache = DualCache() + client_id = (await _register([REDIRECT_URI]))["client_id"] + handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1")) + + first = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=cache, + ) + assert first.status_code == 303 + second = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=cache, + ) + assert second.status_code == 400 + assert json.loads(second.body)["error"] == "invalid_request" + + +@pytest.mark.asyncio +async def test_token_rejects_out_of_range_code_verifier(): + """RFC 7636: a code_verifier outside 43-128 chars is invalid_request, not a confusing + invalid_grant PKCE-mismatch.""" + for bad in ["short", "x" * 200]: + response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code="llm_gcode_whatever", + redirect_uri=REDIRECT_URI, + client_id="llm_dcrc_x", + code_verifier=bad, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_request" + + +@pytest.mark.asyncio +async def test_authorize_rejects_over_long_state(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = aggregate_authorize( + request=_request(query=f"client_id={client_id}"), + client_id=client_id, + redirect_uri=REDIRECT_URI, + state="s" * 2000, + code_challenge=CODE_CHALLENGE, + code_challenge_method="S256", + response_type="code", + session_user_id="u1", + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_request" + + +@pytest.mark.asyncio +async def test_non_ascii_code_challenge_fails_grant_not_500(): + """A non-ASCII code_challenge (unvalidated from the client) must yield a clean + invalid_grant, never a TypeError-driven 500 (bytes comparison, not str).""" + client_id = (await _register([REDIRECT_URI]))["client_id"] + # Seal a code carrying a non-ASCII challenge directly (authorize requires S256 shape, + # but the challenge charset is not validated there, so this state is reachable). + from datetime import datetime, timezone + + code = _seal( + GATEWAY_AUTH_CODE_PREFIX, + _GatewayAuthCode( + user_id="u1", + client_id=client_id, + redirect_uri=REDIRECT_URI, + code_challenge="challenge-with-€-non-ascii", + jti="jti-x", + iat=int(datetime.now(timezone.utc).timestamp()), + exp=int(datetime.now(timezone.utc).timestamp()) + 120, + ), + ) + response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=code, + redirect_uri=REDIRECT_URI, + client_id=client_id, + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_grant" + + +@pytest.mark.asyncio +async def test_single_use_guard_in_memory_is_single_use_within_process(): + """No Redis configured (single-replica): the in-memory increment is authoritative — the first claim + wins, a replay of the same id loses.""" + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard + + guard = _SingleUseGuard(DualCache()) # redis_cache is None + assert await guard.claim("jti-inmem", 60) is True + assert await guard.claim("jti-inmem", 60) is False # replay of the same id + + +@pytest.mark.asyncio +async def test_single_use_guard_uses_redis_as_sole_authority_when_configured(): + """With Redis configured it is the SOLE authority: the shared INCR result decides the claim (1 → + first caller, >1 → replay), and the per-worker in-memory count is never consulted.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard + + cache = DualCache() + cache.redis_cache = MagicMock() + cache.redis_cache.async_increment = AsyncMock(return_value=1) + # in-memory must NOT be consulted when Redis is configured — poison it so any fallback is visible. + cache.async_increment_cache = AsyncMock(side_effect=AssertionError("must not fall back to in-memory")) + + guard = _SingleUseGuard(cache) + assert await guard.claim("jti-redis", 60) is True + cache.redis_cache.async_increment = AsyncMock(return_value=2) + assert await guard.claim("jti-redis", 60) is False # Redis says 2 → replay + + +@pytest.mark.asyncio +async def test_single_use_guard_fails_closed_when_redis_errors(): + """A Redis fault must fail the claim CLOSED (refuse the id) rather than fall back to the per-worker + in-memory count — which would let each replica observe count==1 and replay the one-time id (the + Cursor/Veria replay-across-workers finding).""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard + + cache = DualCache() + cache.redis_cache = MagicMock() + cache.redis_cache.async_increment = AsyncMock(side_effect=ConnectionError("redis down")) + cache.async_increment_cache = AsyncMock(return_value=1) # would fail OPEN if the guard fell back + + guard = _SingleUseGuard(cache) + assert await guard.claim("jti-fault", 60) is False # fail closed, not a fallback count of 1 diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 288e2533b72..c589014f276 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -176,9 +176,7 @@ async def test_authenticate_user_invalid_credentials(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) - with patch.dict( - os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": "correct-password"} - ): + with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": "correct-password"}): with pytest.raises(ProxyException) as exc_info: await authenticate_user( username=ui_username, @@ -227,9 +225,7 @@ async def test_authenticate_user_wrong_password(): ) mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - return_value=mock_user - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=mock_user) with patch.dict( os.environ, @@ -279,9 +275,7 @@ async def test_authenticate_user_email_case_insensitive_login(): return None mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - side_effect=mock_find_first - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(side_effect=mock_find_first) with patch.dict( os.environ, @@ -334,9 +328,7 @@ async def test_authenticate_user_database_required_for_admin(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) - with patch.dict( - os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": ui_password} - ): + with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": ui_password}): with patch( "litellm.proxy.auth.login_utils.user_update", new_callable=AsyncMock, @@ -429,9 +421,7 @@ def test_authenticate_user_non_ascii_direct_comparison(): assert result is True # And correctly returns False for different passwords - result = secrets.compare_digest( - password.encode("utf-8"), "different£pass".encode("utf-8") - ) + result = secrets.compare_digest(password.encode("utf-8"), "different£pass".encode("utf-8")) assert result is False @@ -531,9 +521,7 @@ async def test_authenticate_user_database_login_with_non_ascii_password(): return None mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - side_effect=mock_find_first - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(side_effect=mock_find_first) with patch.dict( os.environ, @@ -559,3 +547,58 @@ async def test_authenticate_user_database_login_with_non_ascii_password(): assert isinstance(result, LoginResult) assert result.user_id == "test-user-123" assert result.user_email == user_email + + +class TestEncodeUiSessionJwt: + """The UI session cookie must carry a bounded exp so it does not stay + signature-valid until the master key rotates, and so the session-cookie readers + that require a bounded lifetime (the MCP interactive sign-in) accept it.""" + + def _decode(self, token: str) -> dict: + import jwt + + return jwt.decode(token, "sk-master-for-tests", algorithms=["HS256"]) + + def test_encoded_cookie_carries_bounded_exp(self): + import time + + from litellm.proxy.auth.login_utils import encode_ui_session_jwt + + token_object = {"user_id": "u1", "key": "sk-abc", "login_method": "username_password"} + with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "24h"): + token = encode_ui_session_jwt(token_object, "sk-master-for-tests") + claims = self._decode(token) + assert claims["user_id"] == "u1" + assert claims["login_method"] == "username_password" + remaining = claims["exp"] - int(time.time()) + assert 23 * 3600 < remaining <= 24 * 3600 + + def test_duration_is_honored_from_env(self): + import time + + from litellm.proxy.auth.login_utils import encode_ui_session_jwt + + with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "1h"): + token = encode_ui_session_jwt({"user_id": "u1"}, "sk-master-for-tests") + remaining = self._decode(token)["exp"] - int(time.time()) + assert 0 < remaining <= 3600 + + def test_cookie_is_accepted_by_the_exp_requiring_session_reader(self): + """The regression this change exists for: before it, the UI cookie carried no + exp and _user_id_from_session_cookie (require=["exp"]) rejected every real login, + so the MCP interactive sign-in could never capture identity. A cookie minted by + this helper must now be accepted.""" + from unittest.mock import MagicMock + + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( + _user_id_from_session_cookie, + ) + from litellm.proxy.auth.login_utils import encode_ui_session_jwt + + token_object = {"user_id": "cornell-user", "key": "sk-abc", "login_method": "sso"} + with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "24h"): + token = encode_ui_session_jwt(token_object, "sk-master-for-tests") + request = MagicMock() + request.cookies = {"token": token} + with patch("litellm.proxy.proxy_server.master_key", "sk-master-for-tests"): + assert _user_id_from_session_cookie(request) == "cornell-user" 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 c693017e134..63a47428780 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -7763,3 +7763,80 @@ async def test_cli_completion_persists_assertion_under_db_user_id(): retain_mock.assert_awaited_once_with(user_id="cli-user-id", assertion=assertion) assert response.status_code == 200 + + +class TestSameOriginReturnPath: + """The same-origin relative return_to arm added for the MCP gateway DCR authorize + round-trip: only strictly relative paths qualify, so login can never redirect the + browser off the gateway origin.""" + + def test_accepts_relative_paths(self): + from litellm.proxy.management_endpoints.ui_sso import _is_same_origin_return_path + + assert _is_same_origin_return_path("/authorize?client_id=llm_dcrc_x&state=s") is True + assert _is_same_origin_return_path("/some_server/authorize") is True + + def test_rejects_absolute_protocol_relative_and_backslash_paths(self): + from litellm.proxy.management_endpoints.ui_sso import _is_same_origin_return_path + + assert _is_same_origin_return_path("https://evil.example.com/authorize") is False + assert _is_same_origin_return_path("//evil.example.com/authorize") is False + assert _is_same_origin_return_path("/\\evil.example.com") is False + assert _is_same_origin_return_path("javascript:alert(1)") is False + assert _is_same_origin_return_path("") is False + + +class TestPersistReturnToCookieSharedHelper: + """The single shared return_to helper used by EVERY sign-in branch (SSO / Okta / generic AND the + username/password form). It must be best-effort and NEVER raise — a bad return_to can never block + sign-in. Regression: the password form previously 400'd because it called _validate_return_to + directly (which raises for a non-matching absolute return_to when control_plane_url is set).""" + + @staticmethod + def _cookie(resp) -> str: + return resp.headers.get("set-cookie", "") + + def test_sets_cookie_for_same_origin_relative_path(self, monkeypatch): + from fastapi import Response + + from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie + + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + resp = Response() + _persist_return_to_cookie(resp, "/mcp/authorize?client_id=llm_dcrc_abc") + assert "litellm_cp_return_to=" in self._cookie(resp) + + def test_bad_absolute_with_control_plane_configured_does_not_raise_and_is_not_stored(self, monkeypatch): + """THE regression: a non-matching absolute return_to with control_plane_url set must NOT raise + (it did, blocking the login form) and must NOT be stored — sign-in proceeds.""" + from fastapi import Response + + from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie + + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", {"control_plane_url": "https://cp.example.com"} + ) + resp = Response() + _persist_return_to_cookie(resp, "https://evil.example.com/steal") # must not raise + assert "litellm_cp_return_to=" not in self._cookie(resp) + + def test_none_return_to_is_a_noop(self): + from fastapi import Response + + from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie + + resp = Response() + _persist_return_to_cookie(resp, None) + assert "litellm_cp_return_to=" not in self._cookie(resp) + + def test_control_plane_matching_absolute_is_stored(self, monkeypatch): + from fastapi import Response + + from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie + + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", {"control_plane_url": "https://cp.example.com"} + ) + resp = Response() + _persist_return_to_cookie(resp, "https://cp.example.com/ui?page=models") + assert "litellm_cp_return_to=" in self._cookie(resp) diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index f0250bbe1a6..a75d5bd5730 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -49,9 +49,7 @@ def _install_login_mocks(monkeypatch, raise_on_auth: bool = False) -> None: } monkeypatch.setattr("litellm.proxy.auth.login_utils.authenticate_user", _fake_auth) - monkeypatch.setattr( - "litellm.proxy.auth.login_utils.create_ui_token_object", _fake_token_object - ) + monkeypatch.setattr("litellm.proxy.auth.login_utils.create_ui_token_object", _fake_token_object) monkeypatch.setattr(ps, "master_key", "sk-test-master") monkeypatch.setattr(ps, "general_settings", {}) monkeypatch.setattr(ps, "premium_user", False) @@ -69,9 +67,7 @@ def test_fallback_login_returns_html_form(client, monkeypatch): body_lower = response.text.lower() shape = { "status": response.status_code, - "content_type_html": response.headers.get("content-type", "").startswith( - "text/html" - ), + "content_type_html": response.headers.get("content-type", "").startswith("text/html"), "has_form": "", "token": ""} + assert normalize(response.json(), volatile=frozenset({"token", "redirect_url"})) == { + "redirect_url": "", + "token": "", + } body = response.json() set_cookie = response.headers.get("set-cookie", "") shape = { "redirect_url_has_ui": "/ui/" in body.get("redirect_url", ""), - "redirect_url_has_login_success": "login=success" - in body.get("redirect_url", ""), + "redirect_url_has_login_success": "login=success" in body.get("redirect_url", ""), "token_in_body": bool(body.get("token")), "token_cookie_set": "token=" in set_cookie, } @@ -264,9 +254,7 @@ def test_v3_login_success_returns_code(client, monkeypatch): from litellm.proxy import proxy_server as ps _install_login_mocks(monkeypatch) - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) # Force the local (non-redis) cache path monkeypatch.setattr(ps, "redis_usage_cache", None) fake_cache = MagicMock() @@ -301,9 +289,7 @@ def test_v3_login_authenticate_failure_500(client, monkeypatch): from litellm.proxy import proxy_server as ps _install_login_mocks(monkeypatch, raise_on_auth=True) - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) response = client.post( "/v3/login", @@ -337,9 +323,7 @@ def test_v3_login_exchange_missing_code_400(client, monkeypatch): """Error path: missing 'code' in body -> 400 with 'Missing' message.""" from litellm.proxy import proxy_server as ps - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) response = client.post("/v3/login/exchange", json={}) assert response.status_code == 400 @@ -352,9 +336,7 @@ def test_v3_login_exchange_invalid_code_401(client, monkeypatch): """Error path: code that isn't in cache -> 401 'Invalid or expired'.""" from litellm.proxy import proxy_server as ps - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) monkeypatch.setattr(ps, "redis_usage_cache", None) fake_cache = MagicMock() fake_cache.async_get_cache = AsyncMock(return_value=None) @@ -372,9 +354,7 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc """Pin: valid code -> JSON {token, redirect_url} + token cookie + cache deleted (single-use).""" from litellm.proxy import proxy_server as ps - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) monkeypatch.setattr(ps, "redis_usage_cache", None) cached_payload = { @@ -388,9 +368,10 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc response = client.post("/v3/login/exchange", json={"code": "valid-code"}) assert response.status_code == 200 - assert normalize( - response.json(), volatile=frozenset({"token", "redirect_url"}) - ) == {"token": "", "redirect_url": ""} + assert normalize(response.json(), volatile=frozenset({"token", "redirect_url"})) == { + "token": "", + "redirect_url": "", + } body = response.json() set_cookie = response.headers.get("set-cookie", "") shape = { @@ -405,3 +386,77 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc "token_cookie_set": True, "cache_deleted_once": True, } + + +def test_login_form_honors_same_origin_return_to_cookie(client, monkeypatch): + """The aggregate DCR connect flow preserves a same-origin return_to in the litellm_cp_return_to + cookie; /login must RESUME there after password sign-in instead of dead-ending at the dashboard.""" + _install_login_mocks(monkeypatch) + return_to = "/mcp/authorize?client_id=llm_dcrc_abc&response_type=code" + response = client.post( + "/login", + data={"username": "admin", "password": "password"}, + cookies={"litellm_cp_return_to": return_to}, + follow_redirects=False, + ) + assert response.status_code == 303 + assert response.headers.get("location", "") == return_to # resumed the connect flow, not the dashboard + assert "token=" in response.headers.get("set-cookie", "") + + +def test_login_form_honors_control_plane_return_to_cookie(client, monkeypatch): + """/login resumes through the SAME resumer the SSO callback uses, so it honors BOTH shapes + _persist_return_to_cookie is willing to store. Honoring only the relative one silently dropped + a control-plane return_to and landed the user on the dashboard.""" + import litellm.proxy.proxy_server as ps + + _install_login_mocks(monkeypatch) + monkeypatch.setitem(ps.general_settings, "control_plane_url", "https://cp.example.com") + response = client.post( + "/login", + data={"username": "admin", "password": "password"}, + cookies={"litellm_cp_return_to": "https://cp.example.com/console"}, + follow_redirects=False, + ) + location = response.headers.get("location", "") + assert response.status_code == 303 + assert location.startswith("https://cp.example.com/console") + # Cross-origin arm hands the JWT off via a one-time code rather than a cookie. + assert "code=" in location and "login=success" in location + assert "token=" not in response.headers.get("set-cookie", "") + + +def test_login_form_survives_stale_control_plane_return_to(client, monkeypatch): + """A stale one-shot cookie must NEVER fail a completed sign-in. The resumer rejects a return_to + that no longer matches control_plane_url (a config change between the cookie's write and this + read); the user has already authenticated, so land on the dashboard instead of erroring.""" + import litellm.proxy.proxy_server as ps + + _install_login_mocks(monkeypatch) + monkeypatch.setitem(ps.general_settings, "control_plane_url", "https://new-cp.example.com") + response = client.post( + "/login", + data={"username": "admin", "password": "password"}, + cookies={"litellm_cp_return_to": "https://old-cp.example.com/console"}, + follow_redirects=False, + ) + assert response.status_code == 303, "login must not break on a stale return_to cookie" + location = response.headers.get("location", "") + assert "old-cp.example.com" not in location + assert "/ui/" in location + + +def test_login_form_ignores_open_redirect_return_to(client, monkeypatch): + """A non-same-origin return_to (open-redirect attempt) is rejected — /login falls back to the + dashboard rather than honoring an absolute/foreign URL.""" + _install_login_mocks(monkeypatch) + response = client.post( + "/login", + data={"username": "admin", "password": "password"}, + cookies={"litellm_cp_return_to": "https://evil.example.com/steal"}, + follow_redirects=False, + ) + assert response.status_code == 303 + location = response.headers.get("location", "") + assert "evil.example.com" not in location + assert "/ui/" in location # dashboard fallback diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 5b67780dc58..bad76864ca7 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -127,11 +127,15 @@ def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch): general_settings={}, premium_user=False, ) - mock_jwt_encode.assert_called_once_with( - {"user_id": "test-user"}, - "test-master-key", - algorithm="HS256", - ) + mock_jwt_encode.assert_called_once() + payload, secret = mock_jwt_encode.call_args.args + # The UI session token carries a bounded-lifetime `exp` claim (dynamic timestamp), alongside + # the user_id; assert its presence rather than an exact expiry value. + assert payload["user_id"] == "test-user" + assert isinstance(payload.get("exp"), int) and payload["exp"] > 0 + assert set(payload.keys()) == {"user_id", "exp"} + assert secret == "test-master-key" + assert mock_jwt_encode.call_args.kwargs == {"algorithm": "HS256"} def test_login_v2_returns_json_on_proxy_exception(monkeypatch): diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 12f19eaa1fe..6335de32147 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -2972,7 +2972,7 @@ "count": 1 }, "no-nested-ternary": { - "count": 7 + "count": 6 } }, "src/components/chat/MCPConnectPicker.tsx": { diff --git a/ui/litellm-dashboard/src/app/chat/integrations/page.tsx b/ui/litellm-dashboard/src/app/chat/integrations/page.tsx index f85dd591199..30ce62d8081 100644 --- a/ui/litellm-dashboard/src/app/chat/integrations/page.tsx +++ b/ui/litellm-dashboard/src/app/chat/integrations/page.tsx @@ -4,6 +4,7 @@ import { Suspense, useEffect } from "react"; import { useRouter, useSearchParams } from "next/navigation"; import { useChatShell } from "@/contexts/ChatShellContext"; import MCPAppsPanel from "@/components/chat/MCPAppsPanel"; +import ConnectFlowBanner from "@/components/chat/ConnectFlowBanner"; // useSearchParams() requires a Suspense boundary for static export. function IntegrationsPageContent() { @@ -11,6 +12,13 @@ function IntegrationsPageContent() { const router = useRouter(); const searchParams = useSearchParams(); const oauthReturn = searchParams.get("mcpOauthReturn"); + // Set by the gateway DCR authorize when a DCR client sends the user here to + // authorize servers before finishing sign-in (see gateway_dcr_flow.py). The + // handle keys the sealed per-flow cookie; connect_client is the client origin + // for display only. connect_flow is NOT cleaned from the URL: the finish form + // needs it, and the sealed cookie (not the URL) is the security boundary. + const connectFlow = searchParams.get("connect_flow"); + const connectClient = searchParams.get("connect_client"); // Clean up the OAuth return param after it's been consumed — real routing means // we no longer need it to pick a tab, but it should not linger in the address bar. @@ -24,7 +32,13 @@ function IntegrationsPageContent() { return (
- + {connectFlow && } +
); } diff --git a/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.test.tsx b/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.test.tsx new file mode 100644 index 00000000000..a565ae5db08 --- /dev/null +++ b/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.test.tsx @@ -0,0 +1,51 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; +import { render, screen } from "@testing-library/react"; +import ConnectFlowBanner from "./ConnectFlowBanner"; + +vi.mock("@/components/networking", () => ({ + getProxyBaseUrl: () => "https://gateway.example.com", +})); + +afterEach(() => { + vi.restoreAllMocks(); + sessionStorage.clear(); +}); + +describe("ConnectFlowBanner", () => { + it("posts the flow handle to the proxy /authorize/complete as a full-page form", () => { + const { container } = render(); + + const form = container.querySelector("form")!; + expect(form.getAttribute("method")).toBe("POST"); + expect(form.getAttribute("action")).toBe("https://gateway.example.com/authorize/complete"); + + const hidden = form.querySelector('input[name="flow"]') as HTMLInputElement; + expect(hidden.value).toBe("flow-handle-123"); + // No token, code, or secret is ever placed in the form; the sealed cookie carries them. + expect(form.innerHTML).not.toContain("token"); + }); + + it("shows the client origin so the user knows what they are connecting to", () => { + render(); + expect(screen.getAllByText(/claude\.ai/).length).toBeGreaterThan(0); + expect(screen.getByRole("button", { name: /finish connecting/i })).toBeInTheDocument(); + }); + + it("falls back to a generic label when the client origin is unknown", () => { + render(); + expect(screen.getAllByText(/the application/).length).toBeGreaterThan(0); + }); + + it("does NOT complete the flow on pagehide (completion requires the explicit button)", () => { + // Security regression: an attacker could lure a signed-in victim to their own client's + // authorize URL; the victim merely closing the tab must NOT deliver a victim-bound code. + // Completion is a deliberate button press, never a side effect of leaving the page. + const beaconMock = vi.fn(() => true); + vi.stubGlobal("navigator", { ...navigator, sendBeacon: beaconMock }); + render(); + + window.dispatchEvent(new Event("pagehide")); + + expect(beaconMock).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.tsx b/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.tsx new file mode 100644 index 00000000000..ac2c508e815 --- /dev/null +++ b/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.tsx @@ -0,0 +1,59 @@ +"use client"; + +import React from "react"; +import { CheckCircle } from "lucide-react"; +import { getProxyBaseUrl } from "@/components/networking"; + +interface Props { + flowHandle: string; + clientOrigin: string | null; +} + +/** + * The interlude shown when a DCR client (Claude Desktop, MCP Inspector) sends the user + * through the gateway sign-in and lands them on the apps grid to authorize servers. The + * grid below authorizes individual servers into the per-user vault; this banner is the + * finish step that returns the user to the client. + * + * Finishing requires the explicit "Finish connecting" button: a native form POST to the proxy's + * /authorize/complete, which mints the gateway authorization code and 303-redirects to the DCR + * client's own redirect URI (the full-page navigation carries the HttpOnly per-flow cookie and + * follows the cross-origin redirect to the client's loopback). + * + * The button press IS the consent gate and must not be bypassed. An earlier version auto-finished + * on tab close via navigator.sendBeacon; that let an attacker who lured a signed-in victim to their + * own client's authorize URL harvest a victim-bound code the moment the victim closed the tab + * (no click). Merely visiting the authorize URL is attacker-inducible, so completion has to be a + * deliberate user action, not a side effect of leaving the page. + */ +const ConnectFlowBanner: React.FC = ({ flowHandle, clientOrigin }) => { + const action = `${getProxyBaseUrl()}/authorize/complete`; + const clientLabel = clientOrigin ?? "the application"; + + return ( +
+
+
+ +
+

Connect your MCP servers to {clientLabel}

+

+ Authorize the servers you want to use below, then click Finish connecting to return to {clientLabel}. +

+
+
+
+ + +
+
+
+ ); +}; + +export default ConnectFlowBanner; diff --git a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx index 25ced2d62c3..329e4ec99fa 100644 --- a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx +++ b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx @@ -13,7 +13,7 @@ import { getMCPOAuthUserCredentialStatus, listMCPTools, } from "../networking"; -import { AUTH_TYPE, MCPServer, MCPTool, handleTransport } from "../mcp_tools/types"; +import { AUTH_TYPE, MCPServer, MCPTool, handleTransport, isUnsupportedOnGatewayConnect } from "../mcp_tools/types"; import { Logo } from "@/components/molecules/logo/Logo"; import MessageManager from "@/components/molecules/message_manager"; import { useUserMcpOAuthFlow } from "@/hooks/useUserMcpOAuthFlow"; @@ -71,6 +71,7 @@ interface Props { accessToken: string; selectedServers: string[]; onChange: (servers: string[]) => void; + connectMode?: boolean; } const AVATAR_COLORS = [ @@ -96,7 +97,7 @@ type TabKey = "all" | "connected"; const TOOLS_FETCH_CONCURRENCY = 5; -const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange }) => { +const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange, connectMode }) => { const [servers, setServers] = useState([]); const [loading, setLoading] = useState(true); const [query, setQuery] = useState(""); @@ -106,6 +107,7 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange const [toolCounts, setToolCounts] = useState>({}); const [loadingCounts, setLoadingCounts] = useState(false); const [oauthConnected, setOauthConnected] = useState>(new Set()); + const [oauthChecking, setOauthChecking] = useState>(new Set()); const serversRef = useRef([]); useEffect(() => { @@ -148,6 +150,14 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange } } catch { // ignore + } finally { + if (!fetchLoadCancelledRef.current) { + setOauthChecking((prev) => { + const next = new Set(prev); + next.delete(server.server_id); + return next; + }); + } } }, [accessToken], @@ -160,9 +170,13 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange .then(async (serverData) => { if (fetchLoadCancelledRef.current) return; const list: MCPServer[] = Array.isArray(serverData) ? serverData : serverData?.data ?? []; + const oauthServers = list.filter((s) => s.auth_type === AUTH_TYPE.OAUTH2); setServers(list); + setOauthChecking(new Set(oauthServers.map((s) => s.server_id))); setLoading(false); + oauthServers.forEach((s) => checkOauthCredential(s)); + setLoadingCounts(true); const chunks = Array.from({ length: Math.ceil(list.length / TOOLS_FETCH_CONCURRENCY) }, (_, i) => list.slice(i * TOOLS_FETCH_CONCURRENCY, (i + 1) * TOOLS_FETCH_CONCURRENCY), @@ -172,9 +186,6 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange await Promise.allSettled(chunk.map((s) => fetchToolCount(s))); } if (!fetchLoadCancelledRef.current) setLoadingCounts(false); - - const oauthServers = list.filter((s) => s.auth_type === AUTH_TYPE.OAUTH2); - oauthServers.forEach((s) => checkOauthCredential(s)); }) .catch(() => { if (!fetchLoadCancelledRef.current) { @@ -231,6 +242,36 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange } }; + const renderConnectionIndicator = (server: MCPServer) => { + if (connectMode && isUnsupportedOnGatewayConnect(server.auth_type)) { + return ( + + Not supported on this connection + + ); + } + if (server.auth_type === AUTH_TYPE.OAUTH2) { + if (oauthConnected.has(server.server_id)) { + return ; + } + if (oauthChecking.has(server.server_id)) { + return ; + } + return ( + setOauthConnected((prev) => new Set(prev).add(id))} + variant="badge" + /> + ); + } + if (selectedServers.includes(nameOf(server))) { + return ; + } + return null; + }; + const { data: detailToolsResult, isLoading: loadingTools } = useQuery({ queryKey: ["mcp-apps-panel-detail-tools", detailServer?.server_id], queryFn: () => listMCPTools(accessToken, detailServer!.server_id), @@ -390,24 +431,30 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange

MCP Servers

- - Beta - -
-
-

Browse tools, authenticate once, use in chat

- {loadingCounts ? ( - - - Loading tools... + {!connectMode && ( + + Beta - ) : totalTools > 0 ? ( - - - {totalTools} tool{totalTools !== 1 ? "s" : ""} available - - ) : null} + )}
+ {connectMode ? ( +

Click a server to see its tools and connect

+ ) : ( +
+

Browse tools, authenticate once, use in chat

+ {loadingCounts ? ( + + + Loading tools... + + ) : totalTools > 0 ? ( + + + {totalTools} tool{totalTools !== 1 ? "s" : ""} available + + ) : null} +
+ )}
@@ -458,10 +505,10 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange
{filtered.map((server, idx) => { const name = nameOf(server); - const isConnected = selectedServers.includes(name); const color = getAvatarColor(name); const isLeftCol = idx % 2 === 0; const count = toolCounts[name]; + const unsupported = !!connectMode && isUnsupportedOnGatewayConnect(server.auth_type); return (
= ({ accessToken, selectedServers, onChange onClick={() => setDetailServer(server)} className={`flex items-center gap-3 p-4 bg-card cursor-pointer transition-colors hover:bg-accent/30 min-w-0 ${ isLeftCol ? "border-r" : "" - } ${Math.floor(idx / 2) < Math.floor((filtered.length - 1) / 2) ? "border-b" : ""}`} + } ${Math.floor(idx / 2) < Math.floor((filtered.length - 1) / 2) ? "border-b" : ""} ${ + unsupported ? "opacity-50" : "" + }`} > {server.mcp_info?.logo_url ? ( = ({ accessToken, selectedServers, onChange ) : null}
- {server.auth_type === AUTH_TYPE.OAUTH2 ? ( - oauthConnected.has(server.server_id) ? ( - - ) : ( - { - setOauthConnected((prev) => new Set(prev).add(id)); - }} - variant="badge" - /> - ) - ) : isConnected ? ( - - ) : null} + {renderConnectionIndicator(server)}
); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.test.tsx index 1b2feb3c731..fc987ee7230 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.test.tsx @@ -14,6 +14,7 @@ import { preservedDeclaredAppCredentials, withoutMintedTokenCredentials, credentialAuthClass, + isUnsupportedOnGatewayConnect, } from "./types"; describe("getOAuthAuthorizationIdentity", () => { @@ -267,3 +268,23 @@ describe("credentialAuthClass", () => { expect(credentialAuthClass(null)).toBeNull(); }); }); + +describe("isUnsupportedOnGatewayConnect", () => { + it("flags the modes that need a caller-supplied upstream token or subject", () => { + // client-forwarded: caller presents the upstream Authorization per call + expect(isUnsupportedOnGatewayConnect(AUTH_TYPE.TRUE_PASSTHROUGH)).toBe(true); + expect(isUnsupportedOnGatewayConnect(AUTH_TYPE.OAUTH_DELEGATE)).toBe(true); + // OBO: caller's own IdP token is the exchange subject, which the session bearer is not + expect(isUnsupportedOnGatewayConnect(AUTH_TYPE.OAUTH2_TOKEN_EXCHANGE)).toBe(true); + }); + + it("does not flag modes the gateway can serve from server-side state or interactive vaulting", () => { + // interactive authorization_code is the one mode the connect grid vaults per user + expect(isUnsupportedOnGatewayConnect(AUTH_TYPE.OAUTH2)).toBe(false); + // server-configured credentials need no per-user connect + expect(isUnsupportedOnGatewayConnect(AUTH_TYPE.API_KEY)).toBe(false); + expect(isUnsupportedOnGatewayConnect(AUTH_TYPE.NONE)).toBe(false); + expect(isUnsupportedOnGatewayConnect(null)).toBe(false); + expect(isUnsupportedOnGatewayConnect(undefined)).toBe(false); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 049779e9fe2..038dc5cb2ca 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -65,6 +65,15 @@ export const gatewayMintsClientFor = (server: { auth_type?: string | null; dcr_b server.auth_type === AUTH_TYPE.TRUE_PASSTHROUGH || (server.auth_type === AUTH_TYPE.OAUTH_DELEGATE && !server.dcr_bridge); +// Auth modes that cannot be used through the gateway aggregate connect flow, where the client holds +// only an identity-only session bearer and upstream credentials are resolved server-side from the +// per-user vault. The vault is only populated by interactive authorization_code (oauth2). The +// client-forwarded modes need the caller to present the upstream Authorization per call, and +// oauth2_token_exchange (OBO) needs the caller's own IdP token as the subject to exchange; the +// session bearer is neither, so none of these can complete a tool call on this connection. +export const isUnsupportedOnGatewayConnect = (authType?: string | null): boolean => + isClientForwardedTokenMode(authType) || authType === AUTH_TYPE.OAUTH2_TOKEN_EXCHANGE; + export const OAUTH_FLOW = { INTERACTIVE: "interactive", M2M: "m2m", From ffa0dffbc6aa7918a9c2c3a1088aef18d0c2dc63 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Wed, 22 Jul 2026 15:39:03 -0700 Subject: [PATCH 05/25] refactor(mcp): trim redundant comments and dedupe admission-arm tests Compress the security rationale in the gateway-session admission path of user_api_key_auth_mcp.py, keeping the load-bearing "why" and dropping the restatement, and remove a garbled dead comment in get_allowed_tools_for_server In the tests, hoist the duplicated _team / _admitted_subject fixtures to module-level factories and parametrize the four fail-closed session-bearer variants into one case. No behavior change; the 294 tests in the file still pass --- .../mcp_server/auth/user_api_key_auth_mcp.py | 367 ++++++------------ .../auth/test_user_api_key_auth_mcp.py | 323 +++++++-------- 2 files changed, 269 insertions(+), 421 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 01e8b1490b2..a27d6b92843 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -131,11 +131,9 @@ def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> b """True when this auth is a keyless subject admitted by the gateway session / bridge user path, as opposed to a JWT or other keyless auth that merely lacks a ``team_id``. - Reads the server-only ``UserAPIKeyAuth.mcp_admitted_user_subject`` field, set exclusively by - ``_reload_admitted_user`` at admission. It is deliberately NOT a ``metadata`` key: virtual-key - metadata is caller-controlled at key creation, so a metadata marker could be forged on a - personal key to gain the team-inherited grant union or to dodge the caller-Authorization - egress scrub. This field cannot be set from caller input.""" + Reads the server-only ``mcp_admitted_user_subject`` field, set only by ``_reload_admitted_user``. It + is deliberately NOT a ``metadata`` key, which is caller-controlled at key creation and so forgeable + on a personal key to gain the team grant union or dodge the egress scrub; this field cannot be.""" return user_api_key_auth is not None and user_api_key_auth.mcp_admitted_user_subject is True @@ -391,11 +389,9 @@ class MCPRequestHandler: and oauth2_headers and is_session_bearer_shaped(oauth2_headers["Authorization"]) ): - # A gateway DCR session bearer at the aggregate /mcp scope: open the - # identity-only session token and admit under the live litellm user it - # references. A session-shaped bearer that does not open fails closed with - # the aggregate invalid_token challenge; a non-session bearer never reaches - # here (is_session_bearer_shaped is false) and falls through to the oauth2 arm. + # A gateway DCR session bearer at the aggregate /mcp scope: open the identity-only session + # token and admit under the live litellm user. One that does not open fails closed with the + # aggregate invalid_token challenge; a non-session bearer falls through to the oauth2 arm. validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session( authorization_value=oauth2_headers["Authorization"], request=request, @@ -431,13 +427,10 @@ class MCPRequestHandler: bearer_presented=False, ) - # Leak-defense (single chokepoint): a gateway admission credential — the session bearer or the - # bridge envelope — is NEVER a valid upstream MCP token. Scrub it from EVERY egress header context - # (top-level Authorization, the deprecated `x-mcp-auth`, and per-server `x-mcp-{alias}-authorization`) - # so no client-forwarded, OBO-subject, or passthrough path can send it upstream, where a hostile - # server could capture and replay it against the aggregate endpoint as this user. Anchored to the - # credential SHAPE, so a legitimate upstream/passthrough token (never session- or envelope-shaped) - # is forwarded unchanged; per-server vaulted credentials (resolved at egress) are unaffected. + # Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge + # envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no + # client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the + # credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged. raw_headers = dict(headers) ( oauth2_headers, @@ -463,10 +456,9 @@ class MCPRequestHandler: @staticmethod def _is_gateway_admission_credential(value: str | None) -> bool: - """True when a header value is a gateway admission credential — a session bearer (``llm_session_`` / - ``llm_srefresh_``) or a bridge envelope. Such a value proves who signed in to the GATEWAY; it is - never a valid credential for an UPSTREAM MCP server, so it must never be forwarded, where a hostile - upstream could capture and replay it against the aggregate ``/mcp`` endpoint as this user.""" + """True when a header value is a gateway admission credential — a session bearer or bridge + envelope. It proves who signed in to the GATEWAY, never a valid UPSTREAM token, so it must never + be forwarded (a hostile upstream could capture and replay it against the aggregate ``/mcp`` scope).""" return value is not None and (is_session_bearer_shaped(value) or is_bridge_envelope_shaped(value)) @staticmethod @@ -478,13 +470,10 @@ class MCPRequestHandler: mcp_server_auth_headers: dict[str, dict[str, str]] | None, ) -> tuple[dict[str, str] | None, dict[str, str], str | None, dict[str, dict[str, str]] | None]: """Remove any gateway admission credential from EVERY egress header context, keyed on the credential - SHAPE: the top-level ``Authorization`` (``oauth2_headers`` + ``raw_headers``), the deprecated - ``x-mcp-auth`` (``mcp_auth_header``), and per-server ``x-mcp-{alias}-authorization`` - (``mcp_server_auth_headers``). A legitimate upstream/passthrough token is never session- or - envelope-shaped, so it is forwarded unchanged; the per-server token the bridge arm injects is the - real upstream credential (also not gateway-shaped), so it survives. An admitted subject's top-level - Authorization IS the admission bearer, so it is dropped unconditionally as defense-in-depth even - though it is already gateway-shaped.""" + SHAPE: top-level ``Authorization`` (oauth2 + raw), the deprecated ``x-mcp-auth``, and per-server + ``x-mcp-{alias}-authorization``. A legitimate upstream/passthrough token is never gateway-shaped so + it survives (including the real upstream token the bridge arm injects per-server); an admitted + subject's top-level Authorization is dropped unconditionally as defense-in-depth.""" cred = MCPRequestHandler._is_gateway_admission_credential # 1. Top-level Authorization → oauth2_headers. @@ -745,21 +734,12 @@ class MCPRequestHandler: ) -> UserAPIKeyAuth: """Open a gateway DCR session bearer and admit the live litellm user it references. - The custody sibling of :meth:`_admit_dcr_bridge_delegate`: the session token seals - no upstream credential (those are vaulted per user and resolved at egress), so this - admits identity only and injects no per-server header. The token's signature proves - the user signed in when it was minted, but authorization is resolved fresh here, the - sealed ``user_id`` reloads the current user record through the SAME - :meth:`_reload_admitted_user` the bridge user-subject path uses, and the admitted - identity runs through the centralized policy gate, so the user's present team, org, - budget, and SCIM state gate the request rather than a snapshot frozen at mint time. - - Fails closed with the aggregate ``invalid_token`` challenge on an expired, tampered, - or foreign token, on a refresh token presented at the tool edge, and when the - referenced user is missing, deactivated, or rejected by the policy gate. The - pre-DB gates (size, IP, route allowlist) run first, mirroring the bridge arm and the - standard pipeline, so a caller blocked by IP or route is turned away before any - crypto or DB read.""" + Identity-only sibling of :meth:`_admit_dcr_bridge_delegate`: the session token seals no + upstream credential (those are vaulted per user, resolved at egress), so authorization is + resolved fresh via :meth:`_reload_admitted_user` + the centralized policy gate rather than a + mint-time snapshot. Pre-DB gates (size, IP, route allowlist) run first, mirroring the standard + pipeline. Fails closed with the aggregate ``invalid_token`` challenge on an expired, tampered, + foreign, or refresh token, or a missing/deactivated/policy-rejected user.""" from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( NotSessionBearer, SessionBearerAdmitted, @@ -842,34 +822,18 @@ class MCPRequestHandler: @staticmethod async def _reload_admitted_user(user_id: str) -> UserAPIKeyAuth: - """Reload the live user an interactively-minted envelope references and admit them as - themselves. + """Reload the live user an interactively-minted envelope references and admit them as themselves. - The DCR client authenticates via SSO at the bridged authorize, which yields a user - subject rather than a virtual key, so the envelope admits under the user's own - identity: the reloaded ``user_id``, the user's own MCP object permission, and the user's - ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the SAME ``get_allowed_mcp_servers`` - the key path uses then computes which servers the user may reach, so the user's litellm MCP - grants and access groups gate the request exactly as a key's do. Because the returned auth is - stamped ``mcp_admitted_user_subject`` (below), ``get_allowed_mcp_servers`` unions the servers the - user reaches through ANY of their teams on top of these direct grants — a ``UserAPIKeyAuth`` - pins one ``team_id`` but a user belongs to many, so the team fan-out happens off the marker, not - the single ``team_id``. Each source is bounded by ITS OWN org: the user's direct grants by the - bound ``org_id`` (their primary org), and each team's grant by that team's owning org inside - ``_allowed_mcp_servers_for_single_team`` — so a user who spans organizations does not leak one - org's servers past another org's ceiling. The caller's centralized policy gate enforces the - user's live budget and org state, and a SCIM-deactivated owner fails closed. + The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the + SAME ``get_allowed_mcp_servers`` the key path uses gates the request. The ``mcp_admitted_user_subject`` + marker (set below) makes that resolver union the servers the user reaches through ANY of their teams + on top of these direct grants, each source bounded by ITS OWN org, so a user spanning organizations + cannot leak one org's servers past another's ceiling. - Error handling mirrors the key path's retryable-503 contract, but ``get_user_object`` defeats a - type-based check: where ``get_key_object`` raises a typed ``ProxyException`` for a missing key - and lets a DB outage propagate raw, ``get_user_object`` catches every DB failure and re-raises a - bare ``ValueError``, so a missing user and a real outage look identical and the original error - survives only as ``__context__``. ``_raise_503_if_db_unavailable`` therefore walks the cause - chain: a transient DB outage still surfaces as a retryable 503, while a missing user, or any - other non-outage resolution failure, fails closed as a 401 rather than an opaque 500. The - object-permission load shares this one boundary, so an outage there is classified the same - way (``get_object_permission`` itself swallows a failed load to ``None``, matching how - ``get_key_object`` best-effort-loads a key's object permission).""" + Error handling: ``get_user_object`` catches every DB failure and re-raises a bare ``ValueError``, so a + missing user and a real outage look identical (the cause survives only as ``__context__``). + ``_raise_503_if_db_unavailable`` walks the cause chain so an outage stays a retryable 503 while any + other failure fails closed as 401, not an opaque 500; the object-permission load shares that boundary.""" from litellm.proxy.auth.auth_checks import get_object_permission, get_user_object from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -907,49 +871,35 @@ class MCPRequestHandler: org_id=user_object.organization_id, object_permission=object_permission, object_permission_id=user_object.object_permission_id, - # Copy the live user's rate limits, exactly as the standard user-subject auth path does - # (user_api_key_auth.py). The parallel limiter reads these off the auth object rather than - # re-fetching, and treats None as sys.maxsize (unlimited), so a keyless admitted user with - # them unset would invoke tools past their configured user RPM/TPM. - # - # Rate-limit model for the keyless admitted subject: bounded by their USER rpm/tpm - # (copied here; those descriptors key off user_id, which is set) AND by the per-server - # mcp_rpm_limit of EVERY team it reaches servers through, stamped below. Per-KEY MCP - # limits genuinely do not apply, because there is no key. + # Copy the live user's rate limits, as the standard user-subject path does: the parallel + # limiter reads these off the auth object and treats None as unlimited, so a keyless subject + # with them unset would outrun its user RPM/TPM. (Per-team mcp_rpm_limit is stamped below; + # per-KEY limits do not apply, there being no key.) user_tpm_limit=user_object.tpm_limit, user_rpm_limit=user_object.rpm_limit, ) - # Set the server-only admission marker AFTER construction: the before-validator strips it - # from any validated input, so a post-construction assignment is the only way to set it, and - # caller-supplied data (key metadata, JWT claims) can never forge it. + # Server-only marker, set AFTER construction: the before-validator strips it from any validated + # input, so caller-supplied data (key metadata, JWT claims) can never forge it. admitted.mcp_admitted_user_subject = True - # Carry each granting team's per-server MCP rpm limit. A key is pinned to one team so the - # limiter reads team_metadata directly; this subject reaches servers through several teams - # under its own identity, so without this the team ceiling silently does not apply to it and - # a cross-team user outruns every team's mcp_rpm_limit. Resolved from the same roster-checked - # sources the grant union uses, so a team can only throttle what it actually granted. + # Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through + # several teams under its own identity, so without this a cross-team user outruns every team's + # limit. Resolved from the same roster-checked sources as the grant union, so a team throttles + # only what it granted. admitted.mcp_source_team_rpm_limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(admitted) return admitted @staticmethod async def _admitted_subject_team_rpm_limits(auth: UserAPIKeyAuth) -> dict[str, dict[str, int]] | None: - """``team_id -> mcp_rpm_limit`` for every team this subject reaches servers through, with each - team's map filtered to the servers THAT team's grant actually reaches. + """``team_id -> mcp_rpm_limit`` for every team this subject reaches servers through, each map + filtered to the servers THAT team's grant actually reaches. - A limit rides the same scope as the access it bounds: a team's throttle exists to cap usage of - the access the team granted, so a roster team whose grant does not reach a server (not granted, - blocked, org-forbidden, opted out) must not be charged when the user reaches that server - through a DIFFERENT team — otherwise this user's calls drain a bucket shared by that team's own - keys for access the team never provided. The grant scope comes from the SAME - ``get_allowed_mcp_servers(source)`` call authorization uses, so the throttle scope cannot - diverge from the access scope. Limit maps are keyed by server name/alias (the limiter matches - on the called server's name) while grants are ids, so each key is resolved through - ``expand_permission_list`` — the one existing name->id owner — before the membership check. - - Returns None when no team contributes an applicable limit, so the limiter adds no descriptors - rather than empty ones. A lookup failure narrows to None rather than raising: rate limiting - must not be able to deny a request that authorization already allowed, and the user's own - rpm/tpm still bounds them.""" + A limit rides the same scope as the access it bounds, so a roster team is charged only for a + server its OWN grant reaches (never one the user reaches through a different team, which would + drain a bucket shared by that team's keys for access it never provided). Grant scope comes from + the SAME ``get_allowed_mcp_servers(source)`` authorization uses; limit-map keys are names/aliases + so each is resolved to an id via ``expand_permission_list`` before the membership check. Returns + None (no descriptors) when nothing applies; a lookup failure narrows to None rather than raising, + since rate limiting must not deny a request authorization already allowed.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) @@ -969,12 +919,9 @@ class MCPRequestHandler: for server_id in global_mcp_server_manager.expand_permission_list([server_name]): if server_id not in granted_ids: continue - # Charge ONLY the source the call is attributed to — the same single source - # billing picks, from the same owner. Adding a descriptor for every granting - # team let one cross-team user drain several teams' SHARED buckets at once, - # blocking their other members for access those teams did not provide on - # this call; and when the user's OWN grant reaches the server, no team - # provided it, so no team bucket is charged at all. + # Charge ONLY the source billing attributes the call to (same owner), so one + # cross-team user cannot drain several teams' shared buckets on a single call, + # and a server the user's OWN grant reaches charges no team bucket. attributed = await MCPRequestHandler.attributing_source_for_server( auth, server_id, source_grants=source_grants ) @@ -1358,11 +1305,10 @@ class MCPRequestHandler: from litellm.proxy.proxy_server import general_settings try: - # A keyless admitted subject is resolved entirely per source, BEFORE any single-source - # rule runs here. Ordering matters: the no_mcp_servers opt-out below reads the caller's - # own object_permission, so leaving it above this branch let a user's own opt-out zero - # their TEAMS' grants too — the sources are independent, and an opt-out on one of them - # must silence only that one (it is applied per source, inside the recursive call). + # A keyless admitted subject resolves per source BEFORE any single-source rule here. Ordering + # matters: the no_mcp_servers opt-out below reads the caller's own object_permission, so above + # this branch a user's own opt-out would wrongly zero their TEAMS' grants too (each source is + # independent; an opt-out silences only its own source, inside the recursive call). if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None: return await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth) @@ -1484,22 +1430,14 @@ class MCPRequestHandler: has_lower_level_mcp_restrictions: bool, keyless_source: bool = False, ) -> list[str]: - """Cap the resolved server list by this caller's org ceiling. If the org names an explicit MCP - list, lower-level restrictions are intersected with it, else the org list becomes the ceiling. - No org, or an empty org list, leaves the result unchanged. + """Cap the resolved server list by this caller's org ceiling: an explicit org list intersects + lower-level restrictions (else becomes the ceiling); no org or an empty list leaves it unchanged. - ``keyless_source`` marks one grant source of a keyless admitted subject and governs BOTH - org divergences, because they are the same fact about that caller shape. - - First, what an UNRESOLVABLE ceiling means. A virtual key keeps the - long-standing fail-open behavior (a DB blip must not lock working keys out mid-incident). A - keyless admitted subject fails CLOSED, because its only org bound is this ceiling: silently - dropping it on a transient fault would widen a cross-org user to servers their team's org - forbids, which is a privilege escalation rather than an availability blip. - - Second, whether the org list may SUBSTITUTE for absent lower-level grants. For a key it may - (that is the key ceiling model). For a source it may only ever intersect, because the - admitted model is a union of grants and a ceiling that grants is not a ceiling.""" + ``keyless_source`` governs both divergences for a keyless admitted source. An UNRESOLVABLE ceiling + fails CLOSED for it (its only org bound is this ceiling, so dropping it on a fault would escalate a + cross-org user) while a key stays fail-open. And an org list may only ever INTERSECT a source (the + admitted model unions grants, so a ceiling must not become one), whereas for a key it may + substitute, that being the key ceiling model.""" if not (user_api_key_auth and user_api_key_auth.org_id): return allowed_mcp_servers allowed_mcp_servers_for_org = await MCPRequestHandler._get_allowed_mcp_servers_for_org(user_api_key_auth) @@ -1512,12 +1450,8 @@ class MCPRequestHandler: if len(allowed_mcp_servers_for_org) == 0: return allowed_mcp_servers if has_lower_level_mcp_restrictions or keyless_source: - # Lower-level restrictions exist, so org can only cap them. - # - # A keyless admitted source ALWAYS takes this arm: its model is a union of GRANTS, so an - # org list may only narrow what a source already grants, never become one. Letting it - # substitute would hand every admitted user with an org_id that org's whole server list - # without any direct or team grant — a ceiling silently acting as a grant. + # Org can only cap lower-level restrictions. A keyless admitted source ALWAYS takes this + # arm: its model unions GRANTS, so an org list may only narrow a source, never become one. capped = [s for s in allowed_mcp_servers if s in allowed_mcp_servers_for_org] else: # No lower-level restrictions → org list becomes the ceiling. @@ -1535,15 +1469,12 @@ class MCPRequestHandler: ) -> UserAPIKeyAuth: """A plain, UNMARKED auth describing ONE grant source of an admitted subject. - Only the fields the resolver actually consults are carried. Everything else is left at its - default on purpose: ``api_key``/``token`` stay unset (this is not a key), budget, spend and - rate-limit fields stay unset because the admitted subject's own user-level limits are what - the request is metered against and cloning them per source would show the limiter N copies of - the same descriptor, and ``user_role`` stays unset because an admin role would grant every - server if this auth ever reached the server-manager wrapper. The admission marker cannot be - set through the constructor at all (a before-validator pops it), so each source is resolved - as an ordinary caller and cannot re-enter the admitted path. - """ + Only the fields the resolver consults are carried; everything else is left at its default on + purpose: no ``api_key``/``token`` (not a key), no budget/spend/rate-limit (the subject's own + user-level limits meter the request, and per-source copies would double descriptors), no + ``user_role`` (an admin role would grant every server at the server-manager wrapper). The + admission marker cannot be set via the constructor (a before-validator pops it), so each source + resolves as an ordinary caller and cannot re-enter the admitted path.""" scoped = UserAPIKeyAuth( user_id=auth.user_id, team_id=team_id, @@ -1551,9 +1482,8 @@ class MCPRequestHandler: parent_otel_span=auth.parent_otel_span, ) if carry_user_grants: - # The user's OWN grants. A team source deliberately carries none of these: the resolver - # loads that team's object_permission and access groups from team_id itself, and mixing - # the user's in would widen the team source with grants the team never made. + # The user's OWN grants. A team source carries none of these (the resolver loads the team's + # own object_permission from team_id); mixing them in would widen the team with grants it never made. scoped.object_permission = auth.object_permission scoped.object_permission_id = auth.object_permission_id scoped.access_group_ids = auth.access_group_ids @@ -1564,17 +1494,11 @@ class MCPRequestHandler: """The independent sources a keyless admitted subject reaches MCP servers through: their own direct grants, plus every team they are a live roster member of. - Each team source carries that TEAM's org as its ``org_id``, which is what makes the canonical - resolver apply the team's OWN owning-org ceiling to it — a cross-org user's teams are each - bounded by their own org rather than by the caller's home org. A team with no organization - falls back to the user's org so it is bounded rather than unbounded. - - Roster membership is checked HERE because it is a property of the source list, not of any one - resolution: a key is structurally pinned to a team it belongs to, while a user's cached - ``teams`` array can name a team whose ``members_with_roles`` no longer contains them (SCIM - group sync, or cache lag after a team_member_delete), and JWT auth can rewrite that array - outright. The roster is the source of truth for revocation. - """ + Each team source carries that TEAM's org (falling back to the user's), so the canonical resolver + applies the team's OWN owning-org ceiling — a cross-org user's teams are each bounded by their + own org, not the caller's home org. Roster membership is checked HERE (not per resolution) + because a user's cached ``teams`` array can name a team whose ``members_with_roles`` no longer + lists them; the roster is the source of truth for revocation.""" from litellm.proxy.proxy_server import prisma_client sources = [ @@ -1622,11 +1546,9 @@ class MCPRequestHandler: proxy_logging_obj=proxy_logging_obj, ) except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others - # The unit of fault isolation is the SOURCE: a team that cannot be resolved contributes - # nothing this request (fail closed for that team alone — access only ever narrows), - # while the user's own grants and every other resolvable team stand. Raising here - # instead would collapse the whole union to deny-all because one team's row was - # momentarily unreadable, on the servers, tools and throttle axes alike. + # Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for + # it alone, access only narrows) while every other source stands. Raising would collapse the + # whole union to deny-all over one momentarily-unreadable row. verbose_logger.warning(f"MCP admitted-subject source team {team_id!r} unresolvable, skipping: {str(e)}") return None if team_obj is None: @@ -1634,14 +1556,11 @@ class MCPRequestHandler: member_user_ids = {getattr(m, "user_id", None) for m in (team_obj.members_with_roles or [])} - {None} if auth.user_id not in member_user_ids: return None - # A team over its own max budget — or owned by an org over ITS budget — is not a live - # grantor, exactly as it is not for a virtual key pinned to it (common_checks rejects that - # key outright). Enforced with the SAME owners the key path uses (_team_max_budget_check / - # _organization_max_budget_check, cross-pod Redis-first spend), targeted at the TEAM's org - # via the scoped source view, so a cross-org team is judged by its own org's budget. This is - # budget ENFORCEMENT of an already-exceeded state; ATTRIBUTION of new spend stays with the - # user (documented deferral) — the two are different questions. Sitting here, no consumer of - # the source list (servers, tools, throttle stamping) can ever see an over-budget team. + # A team (or its owning org) over budget is not a live grantor, exactly as it is not for a key + # pinned to it. Enforced via the SAME owners the key path uses (_team_max_budget_check / + # _organization_max_budget_check), targeted at the TEAM's org through the scoped source view, so + # no consumer of the source list ever sees an over-budget team. This is ENFORCEMENT of an + # already-exceeded state; ATTRIBUTION of new spend stays with the user (documented deferral). from litellm.exceptions import BudgetExceededError from litellm.proxy.auth.auth_checks import ( _organization_max_budget_check, @@ -1696,16 +1615,11 @@ class MCPRequestHandler: async def billing_auth_for_tool_call(auth: UserAPIKeyAuth, tool_name: str) -> UserAPIKeyAuth: """The auth object a tool call's SPEND should be recorded against. - Returns ``auth`` unchanged for every caller that is not a keyless admitted subject, so key - and JWT billing is byte-identical. For an admitted subject whose call is reached through a - team's grant, returns a copy carrying that team's ``team_id`` and its owning ``org_id`` so - the team's budget accumulates and the correct organization is charged. - - Inert rather than wrong when the target server cannot be resolved from the tool name (a - display-name override, or a REST caller passing server_id with an unprefixed name): billing - then falls back to today's user-level attribution instead of guessing a team. Resolution - reuses the manager's own tool-name lookup rather than re-deriving prefix rules that live - there.""" + ``auth`` unchanged for any non-admitted caller (key/JWT billing byte-identical). For an admitted + subject whose call is reached through a team's grant, a copy carrying that team's ``team_id`` and + owning ``org_id`` so the team's budget accumulates and the right org is charged. Falls back to + user-level attribution (rather than guessing a team) when the tool name does not resolve to a + server, reusing the manager's own tool-name lookup.""" if not _is_mcp_admitted_user_subject(auth): return auth try: @@ -1733,19 +1647,14 @@ class MCPRequestHandler: server_id: str, source_grants: list[tuple[UserAPIKeyAuth, set[str]]] | None = None, ) -> UserAPIKeyAuth | None: - """The source a billable call to ``server_id`` is attributed to, or None to bill the caller - as themselves (their own grant reaches it, or nothing does). + """The source a billable call to ``server_id`` is attributed to, or None to bill the caller as + themselves (their own grant reaches it, or nothing does). - A keyless admitted subject carries no ``team_id``, so downstream spend skipped team updates - entirely and charged the user's PRIMARY org — a team-derived call neither accumulated its - team's budget (so that budget could never begin to block) nor charged the org that owns the - granting team. Attribution restores both. - - The rule: a user's OWN grant is not "through a team", so it bills the user. Otherwise the - call is billed to a granting team — deterministically the lowest ``team_id`` when several - grant the same server, so the choice is stable, reproducible and auditable rather than - dependent on dict ordering. Reads the one grant owner, so the team that gets billed is - always a team that actually granted the server.""" + The rule: a user's OWN grant is not "through a team", so it bills the user; otherwise the call + bills a granting team, deterministically the lowest ``team_id`` when several grant the server so + the pick is stable rather than dict-ordering-dependent. Reads the one grant owner, so the billed + team is always one that actually granted the server (restoring the team budget accrual and + owning-org charge that a keyless, team_id-less subject otherwise skipped).""" source_grants = source_grants or await MCPRequestHandler.admitted_source_grants(auth) granting = [(source, granted) for source, granted in source_grants if server_id in granted] if not granting: @@ -1768,14 +1677,10 @@ class MCPRequestHandler: global_mcp_server_manager, ) - # An OPEN channel (operator-opened allow_all_keys, the user's own BYOM submission) makes the - # server REACHABLE through the user themselves — no grant source names it, so without this - # the union below would return [] and leave it listable but uninvokable. Reachability is ALL - # it confers: it is not a waiver of the ceilings that bound the server. The user's own - # mcp_tool_permissions and their org's tool ceiling still bind, which is what a virtual key - # on the same allow_all server gets (its key_tools and _apply_agent_and_org_tool_ceilings - # both run). Returning None here instead skipped both and let a session holder invoke tools - # their own or their org's policy excludes. + # An OPEN channel (allow_all_keys, the user's own BYOM) makes the server REACHABLE through the + # user, though no grant source names it — without this the union returns [], listable but + # uninvokable. Reachability is ALL it confers, NOT a ceiling waiver: the user's own + # mcp_tool_permissions and org tool ceiling still bind, exactly as a key's do on an allow_all server. reachable_via_open_channel = server_id in await global_mcp_server_manager.operator_open_server_ids(auth) allowed: set[str] = set() @@ -1867,12 +1772,9 @@ class MCPRequestHandler: return None try: - # FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject is - # resolved per grant source and shares NOTHING with the single-credential prelude below. - # Ordering is the invariant, not a nicety — when this branch sat after the prelude, a - # fault in a lookup the subject never uses (its own mcp_toolsets, its team_obj_perm) hit - # the fail-closed handler and denied tools its teams did grant. Nothing that resolves a - # single credential's scope may run before this line. + # FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per + # source and shares nothing with the single-credential prelude below. Ordering is the invariant: + # sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant. if _is_mcp_admitted_user_subject(user_api_key_auth): return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth) @@ -1917,9 +1819,6 @@ class MCPRequestHandler: else None ) - # A keyless gateway/bridge-admitted user has no single team_id, so team_obj_perm above is - # None and the single-team lookup yields allow-all — silently dropping every team's - # per-server tool exclusions. Resolve it as the union over the sources that grant the # Apply same inheritance logic as get_allowed_mcp_servers if team_tools: if key_tools: @@ -1938,15 +1837,10 @@ class MCPRequestHandler: except Exception as e: verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}") - # Fail CLOSED for a keyless admitted subject: ANY error resolving the tool allowlist - # (multi-team fan-out, org/agent lookups) must deny the server's tools ([]) for this - # request rather than collapse to allow-all (None), mirroring the fail-closed server - # path. Key/JWT auth keeps its prior allow-all-on-error behavior. - # - # keyless_source matters as much as the marker: each source of an admitted subject is - # resolved through an UNMARKED auth, so without it a fault under a source returned None, - # and None wins the union as allow-all — dropping every team and org tool ceiling on a - # blip. The marker alone only covers a fault raised before the fan-out. + # Fail CLOSED for a keyless admitted subject: ANY error must deny the server's tools ([]), + # not collapse to allow-all (None); key/JWT auth keeps its prior allow-all-on-error. Both + # keyless_source AND the marker are needed: each source resolves through an UNMARKED auth, so + # without keyless_source a fault under a source returns None and wins the union as allow-all. return [] if (keyless_source or _is_mcp_admitted_user_subject(user_api_key_auth)) else None @staticmethod @@ -1957,16 +1851,12 @@ class MCPRequestHandler: keyless_source: bool = False, ) -> list[str] | None: """Narrow a key/team tool allowlist by the agent's tool permissions and the caller's org tool - ceiling. Each level only ever intersects, and None at a level means "no restriction from this - level". + ceiling. Each level only intersects; None at a level means no restriction from it. - An UNRESOLVABLE org ceiling (``_get_org_object_permission`` raises: the org names a permission - that cannot be loaded) is decided here, per caller shape, mirroring the servers axis: a - virtual key keeps its long-standing fail-open — the org step is skipped and the key/team/agent - restrictions already computed STAND (letting the raise escape would collapse them to - allow-all, which is fail-open WIDER than before the fault). A keyless admitted source - re-raises, and the outer handler denies tools for that one source while the subject's other - sources stand — its only org bound is this ceiling, so skipping it would widen access.""" + An UNRESOLVABLE org ceiling is decided per caller shape, mirroring the servers axis: a key stays + fail-open (skip the org step, keep the key/team/agent restrictions; letting the raise escape + would collapse them to allow-all, WIDER than before the fault), while a keyless source re-raises + so the outer handler denies that one source (its only org bound is this ceiling).""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) @@ -2346,9 +2236,8 @@ class MCPRequestHandler: verbose_logger.debug("prisma_client is None") return None - # An ABSENT org is a determinate fact, not a failure: a team's organization_id can point at a - # row that was deleted or has not synced yet, and get_org_object raises for that. It places no - # ceiling, exactly as a key with a dangling org_id is not locked out. + # A team's organization_id can point at a deleted or not-yet-synced row; get_org_object raises + # OrganizationNotFoundError for that. That is a determinate ABSENCE (no ceiling), handled below. try: org_obj = await get_org_object( org_id=user_api_key_auth.org_id, @@ -2358,20 +2247,17 @@ class MCPRequestHandler: proxy_logging_obj=proxy_logging_obj, ) except OrganizationNotFoundError as e: - # CONFIRMED absent (deleted org, not-yet-synced organization_id): a determinate fact, so - # it places no ceiling. Every OTHER exception is an operational failure and propagates — - # caught upstream as an unresolvable ceiling, which denies for a keyless source and stays - # fail-open for a key. Catching bare Exception here treated a DB outage as "no org", which - # silently dropped a real org's ceiling for exactly as long as the outage lasted. + # CONFIRMED absent: places no ceiling. Every OTHER exception propagates as an unresolvable + # ceiling (denies for a keyless source, fail-open for a key); catching bare Exception here + # would treat a DB outage as "no org" and silently drop a real ceiling for its duration. verbose_logger.debug(f"MCP org ceiling: org {user_api_key_auth.org_id!r} does not exist: {e}") return None if org_obj is None or not org_obj.object_permission_id: return None - # From here the org NAMES a permission. Failing to read it is INDETERMINATE, so it must not - # collapse into the same None that means "no ceiling" -- that is what would silently drop a - # real ceiling on a transient fault. Raise and let each caller pick fail-open or fail-closed. + # The org NAMES a permission; failing to read it is INDETERMINATE and must not collapse into the + # None that means "no ceiling". Raise and let each caller pick fail-open or fail-closed. object_permission = await get_object_permission( object_permission_id=org_obj.object_permission_id, prisma_client=prisma_client, @@ -2420,9 +2306,8 @@ class MCPRequestHandler: all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers return list(set(all_servers)) except Exception as e: - # None = the org ceiling could NOT be resolved, which is not the same fact as [] = the - # org places no restriction. Collapsing the two is what let a transient DB fault silently - # remove an org's ceiling; the caller picks fail-open or fail-closed from this signal. + # None = ceiling UNRESOLVED, distinct from [] = org places no restriction. Collapsing them + # let a DB fault silently drop a ceiling; the caller picks fail-open/closed from this signal. verbose_logger.warning(f"Failed to get allowed MCP servers for org: {str(e)}") return None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 5d32a9c740d..b3c0dcd1681 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -6405,28 +6405,33 @@ class TestGatewaySessionAdmission: # bridge envelope arm); the headers dict is whatever the request carried, here empty. assert not mcp_server_auth_headers - async def test_expired_session_fails_closed_with_invalid_token_challenge(self): + @pytest.mark.parametrize( + "scenario, expect_challenge", + [("expired", True), ("tampered", False), ("refresh_at_tool_edge", False), ("foreign_key", False)], + ) + async def test_bad_session_bearer_fails_closed(self, scenario, expect_challenge): + # Every non-admissible session-shaped bearer fails closed with 401; a valid-but-unusable one + # (expired) additionally carries the invalid_token challenge so the DCR client re-authorizes. from datetime import datetime, timezone - mint, _refresh, principal, keys = self._session_bearer() - token = mint(principal, keys, datetime(2020, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() - with ( - patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - ): + if scenario == "expired": + mint, _refresh, principal, keys = self._session_bearer() + bearer = mint(principal, keys, datetime(2020, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + elif scenario == "tampered": + token = self._access_token() + bearer = token[:-3] + ("aaa" if not token.endswith("aaa") else "bbb") + elif scenario == "refresh_at_tool_edge": + _mint, refresh, principal, keys = self._session_bearer() + bearer = refresh(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + else: # foreign_key: minted under the real master key, presented while the proxy uses another + bearer = self._access_token() + master_key = "sk-a-totally-different-master-key" if scenario == "foreign_key" else self._MASTER_KEY + with patch("litellm.proxy.proxy_server.master_key", master_key): with pytest.raises(HTTPException) as exc_info: - await MCPRequestHandler.process_mcp_request(self._scope(token)) - assert exc_info.value.status_code == 401 - assert 'error="invalid_token"' in (exc_info.value.headers or {})["WWW-Authenticate"] - - async def test_tampered_session_fails_closed(self): - token = self._access_token() - tampered = token[:-3] + ("aaa" if not token.endswith("aaa") else "bbb") - with ( - patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - ): - with pytest.raises(HTTPException) as exc_info: - await MCPRequestHandler.process_mcp_request(self._scope(tampered)) + await MCPRequestHandler.process_mcp_request(self._scope(bearer)) assert exc_info.value.status_code == 401 + if expect_challenge: + assert 'error="invalid_token"' in (exc_info.value.headers or {})["WWW-Authenticate"] async def test_deactivated_user_fails_with_invalid_token_challenge(self): """A cryptographically valid bearer whose referenced user is SCIM-deactivated must fail with @@ -6466,27 +6471,6 @@ class TestGatewaySessionAdmission: assert oauth2_headers is None assert not any(k.lower() == "authorization" for k in (raw_headers or {})) - async def test_refresh_token_is_not_admitted_at_the_tool_edge(self): - from datetime import datetime, timezone - - _mint, refresh, principal, keys = self._session_bearer() - refresh_token = refresh(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() - with ( - patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - ): - with pytest.raises(HTTPException) as exc_info: - await MCPRequestHandler.process_mcp_request(self._scope(refresh_token)) - assert exc_info.value.status_code == 401 - - async def test_foreign_key_session_fails_closed(self): - token = self._access_token() - with ( - patch("litellm.proxy.proxy_server.master_key", "sk-a-totally-different-master-key"), - ): - with pytest.raises(HTTPException) as exc_info: - await MCPRequestHandler.process_mcp_request(self._scope(token)) - assert exc_info.value.status_code == 401 - async def test_arm_does_not_fire_for_named_server(self): """A session-shaped bearer aimed at a named server (path scope) does not enter the aggregate arm; it is treated as an ordinary bearer on that server.""" @@ -6504,24 +6488,41 @@ class TestGatewaySessionAdmission: mock_auth.assert_called_once() +def _make_team(team_id, mcp_servers, *, org_id=None, tool_perms=None, members=("sso-user",)): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member + + return LiteLLM_TeamTable( + team_id=team_id, + organization_id=org_id, + members_with_roles=[Member(user_id=u, role="user") for u in members], + access_group_ids=[], + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id=f"op-{team_id}", mcp_servers=mcp_servers, mcp_tool_permissions=tool_perms + ), + ) + + +def _make_admitted_subject(user_id, *, org_id=None, own_servers=None, own_tool_perms=None): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + op = None + if own_servers is not None or own_tool_perms is not None: + op = LiteLLM_ObjectPermissionTable( + object_permission_id=f"userop-{user_id}", + mcp_servers=own_servers or [], + mcp_tool_permissions=own_tool_perms, + ) + auth = UserAPIKeyAuth(user_id=user_id, api_key=None, org_id=org_id, object_permission=op) + auth.mcp_admitted_user_subject = True + return auth + + @pytest.mark.asyncio class TestUserSubjectTeamUnion: """_get_allowed_mcp_servers_for_team unions across ALL a user's teams for a keyless user-subject caller (the gateway DCR session bearer and bridge user-envelope), while a key-based caller keeps its single-team behavior byte-identically.""" - def _team(self, team_id, mcp_servers, members=("sso-user",), tool_perms=None): - from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member - - return LiteLLM_TeamTable( - team_id=team_id, - members_with_roles=[Member(user_id=u, role="user") for u in members], - access_group_ids=[], - object_permission=LiteLLM_ObjectPermissionTable( - object_permission_id=f"op-{team_id}", mcp_servers=mcp_servers, mcp_tool_permissions=tool_perms - ), - ) - @contextlib.contextmanager def _patch(self, *, teams_by_id, user_teams=None, orgs_by_id=None): async def _get_team_object(team_id, **kw): @@ -6550,15 +6551,9 @@ class TestUserSubjectTeamUnion: ): yield - @staticmethod - def _admitted_subject(user_id): - auth = UserAPIKeyAuth(user_id=user_id, api_key=None) - auth.mcp_admitted_user_subject = True - return auth - async def test_keyless_user_unions_servers_across_all_their_teams(self): - teams = {"team-a": self._team("team-a", ["srv1", "srv2"]), "team-b": self._team("team-b", ["srv2", "srv3"])} - auth = self._admitted_subject("sso-user") + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"]), "team-b": _make_team("team-b", ["srv2", "srv3"])} + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv1", "srv2", "srv3"} @@ -6566,7 +6561,7 @@ class TestUserSubjectTeamUnion: async def test_key_based_caller_uses_single_team_only(self): """A key-based caller (api_key set) with a team_id sees ONLY that team, even though the same user belongs to other teams: key auth must be byte-identical to before.""" - teams = {"team-a": self._team("team-a", ["srv1"]), "team-b": self._team("team-b", ["srv2", "srv3"])} + teams = {"team-a": _make_team("team-a", ["srv1"]), "team-b": _make_team("team-b", ["srv2", "srv3"])} auth = UserAPIKeyAuth(user_id="sso-user", api_key="sk-hash", team_id="team-a") with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) @@ -6575,14 +6570,14 @@ class TestUserSubjectTeamUnion: async def test_keyless_user_with_explicit_team_id_uses_that_team_only(self): """A keyless caller that already pins a team_id (not the user-subject fan-out shape) resolves only that team; the union is strictly for the no-team-id user-subject case.""" - teams = {"team-a": self._team("team-a", ["srv1"]), "team-b": self._team("team-b", ["srv2"])} + teams = {"team-a": _make_team("team-a", ["srv1"]), "team-b": _make_team("team-b", ["srv2"])} auth = UserAPIKeyAuth(user_id="sso-user", api_key=None, team_id="team-a") with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) assert set(result) == {"srv1"} async def test_keyless_user_with_no_teams_gets_nothing_from_teams(self): - auth = self._admitted_subject("lonely-user") + auth = _make_admitted_subject("lonely-user") with self._patch(teams_by_id={}, user_teams=[]): result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) assert result == [] @@ -6606,7 +6601,7 @@ class TestUserSubjectTeamUnion: # those pins a team_id, so this helper only ever answers the single-team question. The fan-out # itself is _admitted_subject_sources' job, asserted below. with self._patch(teams_by_id={}, user_teams=["t2", "t3"]): - assert await MCPRequestHandler._team_ids_for_mcp_grant(self._admitted_subject("u")) == [] + assert await MCPRequestHandler._team_ids_for_mcp_grant(_make_admitted_subject("u")) == [] # keyless, no user_id -> nothing assert await MCPRequestHandler._team_ids_for_mcp_grant(UserAPIKeyAuth(api_key=None)) == [] # keyless with a user_id but NOT admission-marked (JWT auth) -> nothing (unchanged behavior) @@ -6629,9 +6624,9 @@ class TestUserSubjectTeamUnion: outage -> the keyless source denies.""" from litellm.proxy.auth.auth_checks import OrganizationNotFoundError - teams = {"t1": self._team("t1", ["srv1"])} + teams = {"t1": _make_team("t1", ["srv1"])} teams["t1"].organization_id = "org-a" - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") absent = AsyncMock(side_effect=OrganizationNotFoundError("Organization doesn't exist in db.")) with self._patch(teams_by_id=teams, user_teams=["t1"]): @@ -6657,7 +6652,7 @@ class TestUserSubjectTeamUnion: boom = AsyncMock(side_effect=RuntimeError("org lookup exploded")) # The subject must actually REACH something, or the assertion passes either way and pins # nothing (a fail-open mutant survived an earlier version of this test for exactly that). - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") auth.org_id = "org-a" auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) with self._patch(teams_by_id={}, user_teams=[]): @@ -6667,7 +6662,7 @@ class TestUserSubjectTeamUnion: assert admitted == [], "admitted subject must fail CLOSED when its org ceiling cannot resolve" key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", team_id="t1", org_id="org-a") - with self._patch(teams_by_id={"t1": self._team("t1", ["srv1"])}, user_teams=[]): + with self._patch(teams_by_id={"t1": _make_team("t1", ["srv1"])}, user_teams=[]): with patch.object(MCPRequestHandler, "_get_org_object_permission", boom): keyed = await MCPRequestHandler.get_allowed_mcp_servers(key_auth) assert set(keyed) == {"srv1"}, "key auth must keep its long-standing fail-open behavior" @@ -6679,11 +6674,11 @@ class TestUserSubjectTeamUnion: SAME source billing picks — one owner for both, so they cannot disagree.""" from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 - t1 = self._team("t1", ["srv1"]) + t1 = _make_team("t1", ["srv1"]) t1.metadata = {"mcp_rpm_limit": {"srv1": 5}} - t2 = self._team("t2", ["srv1"]) + t2 = _make_team("t2", ["srv1"]) t2.metadata = {"mcp_rpm_limit": {"srv1": 9}} - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id={"t1": t1, "t2": t2}, user_teams=["t1", "t2"]): auth.mcp_source_team_rpm_limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) billed = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") @@ -6700,9 +6695,9 @@ class TestUserSubjectTeamUnion: """When the user's OWN grant reaches the server, no team provided the access, so no team bucket may be charged — the user's own rpm/tpm is what bounds them. Mirrors billing, which bills the user and their own org for exactly this case.""" - t1 = self._team("t1", ["srv1"]) + t1 = _make_team("t1", ["srv1"]) t1.metadata = {"mcp_rpm_limit": {"srv1": 5}} - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) with self._patch(teams_by_id={"t1": t1}, user_teams=["t1"]): limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) @@ -6732,9 +6727,9 @@ class TestUserSubjectTeamUnion: """ACCOUNTING half of team budgets. Without attribution the admitted auth kept team_id=None, so spend skipped team updates (the team's budget never accumulated, so it could never begin to block) and charged the user's PRIMARY org rather than the org owning the granting team.""" - t_grant = self._team("t-grant", ["srv1"]) + t_grant = _make_team("t-grant", ["srv1"]) t_grant.organization_id = "org-team" - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") auth.org_id = "org-user-primary" with self._patch(teams_by_id={"t-grant": t_grant}, user_teams=["t-grant"]): source = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") @@ -6746,9 +6741,9 @@ class TestUserSubjectTeamUnion: """Asserted on billing_auth_for_tool_call itself, not on the source it picks: the source already carries the team's org by construction, so asserting there leaves the copy step unpinned (a mutant dropping org_id survived exactly that). This is the object spend reads.""" - t_grant = self._team("t-grant", ["srv1"]) + t_grant = _make_team("t-grant", ["srv1"]) t_grant.organization_id = "org-team" - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") auth.org_id = "org-user-primary" server = MagicMock(server_id="srv1") with self._patch(teams_by_id={"t-grant": t_grant}, user_teams=["t-grant"]): @@ -6766,8 +6761,8 @@ class TestUserSubjectTeamUnion: would charge that team for access it never provided.""" from litellm.proxy._types import LiteLLM_ObjectPermissionTable - t_other = self._team("t-other", ["srv1"]) - auth = self._admitted_subject("sso-user") + t_other = _make_team("t-other", ["srv1"]) + auth = _make_admitted_subject("sso-user") auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) with self._patch(teams_by_id={"t-other": t_other}, user_teams=["t-other"]): assert await MCPRequestHandler.attributing_source_for_server(auth, "srv1") is None @@ -6775,8 +6770,8 @@ class TestUserSubjectTeamUnion: async def test_billing_attribution_is_deterministic_across_several_granting_teams(self): """When several teams grant the same server the pick must be stable and reproducible rather than dependent on dict/roster ordering, or the same call bills different teams run to run.""" - teams = {"t-b": self._team("t-b", ["srv1"]), "t-a": self._team("t-a", ["srv1"])} - auth = self._admitted_subject("sso-user") + teams = {"t-b": _make_team("t-b", ["srv1"]), "t-a": _make_team("t-a", ["srv1"])} + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["t-b", "t-a"]): first = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") with self._patch(teams_by_id=teams, user_teams=["t-a", "t-b"]): @@ -6797,8 +6792,8 @@ class TestUserSubjectTeamUnion: such a fault hit the fail-closed handler and denied tools its teams did grant.""" from litellm.proxy._types import LiteLLM_ObjectPermissionTable - teams = {"t1": self._team("t1", ["srv1"], tool_perms={"srv1": ["read"]})} - auth = self._admitted_subject("sso-user") + teams = {"t1": _make_team("t1", ["srv1"], tool_perms={"srv1": ["read"]})} + auth = _make_admitted_subject("sso-user") # The subject must carry a toolset, or the prelude never resolves one and the fault below is # unreachable — the branch could sit anywhere and the test would still pass (it did). auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_toolsets=["ts-1"]) @@ -6825,7 +6820,7 @@ class TestUserSubjectTeamUnion: manager._get_active_submitted_mcp_server_ids_for_user = AsyncMock(return_value=["srv-byom"]) db_default_perm = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=[]) - admitted = self._admitted_subject("sso-user") + admitted = _make_admitted_subject("sso-user") admitted.object_permission = db_default_perm scoped_key = UserAPIKeyAuth(user_id="u", api_key="sk-hash", object_permission=db_default_perm) @@ -6840,7 +6835,7 @@ class TestUserSubjectTeamUnion: from litellm.proxy._types import LitellmUserRoles manager = self._manager_with(["srv-granted", "srv-secret"]) - admitted = self._admitted_subject("admin-user") + admitted = _make_admitted_subject("admin-user") admitted.user_role = LitellmUserRoles.PROXY_ADMIN with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-granted"])): admitted_view = set(await manager.get_allowed_mcp_servers(admitted)) @@ -6863,7 +6858,7 @@ class TestUserSubjectTeamUnion: opt_out = LiteLLM_ObjectPermissionTable( object_permission_id="op-u", mcp_servers=[SpecialMCPServerNames.no_mcp_servers.value] ) - admitted = self._admitted_subject("sso-user") + admitted = _make_admitted_subject("sso-user") admitted.object_permission = opt_out with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-team"])): admitted_view = set(await manager.get_allowed_mcp_servers(admitted)) @@ -6880,7 +6875,7 @@ class TestUserSubjectTeamUnion: session holder invoke tools their own policy excludes.""" from litellm.proxy._types import LiteLLM_ObjectPermissionTable - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") # The user is restricted to `read` on srv-open, and NO grant source names that server — # it is reachable only through the open channel, which is exactly the bypass path. auth.object_permission = LiteLLM_ObjectPermissionTable( @@ -6900,7 +6895,7 @@ class TestUserSubjectTeamUnion: source, so the source union alone returns [] — listable but uninvokable. The tools axis asks the same open-channel owner the server union uses, so the server is default-open for tools exactly as a virtual key experiences it.""" - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") open_ids = AsyncMock(return_value={"srv-open"}) with self._patch(teams_by_id={}, user_teams=[]): with patch( @@ -6918,13 +6913,13 @@ class TestUserSubjectTeamUnion: not keep granting servers, tools or throttle scope to a keyless union subject either. Enforced through the SAME owner the key path uses (_team_max_budget_check). Distinct from budget ATTRIBUTION of new spend, which stays with the user (documented deferral).""" - t_over = self._team("t-over", ["srv1"]) + t_over = _make_team("t-over", ["srv1"]) t_over.max_budget = 10.0 t_over.spend = 11.0 - t_ok = self._team("t-ok", ["srv2"]) + t_ok = _make_team("t-ok", ["srv2"]) t_ok.max_budget = 10.0 t_ok.spend = 1.0 - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id={"t-over": t_over, "t-ok": t_ok}, user_teams=["t-over", "t-ok"]): servers = set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) @@ -6935,13 +6930,13 @@ class TestUserSubjectTeamUnion: """The org axis of the same rule, judged against the TEAM's own org (not the caller's primary): a team owned by an org over its budget grants nothing, exactly as a key in that org is rejected by _organization_max_budget_check.""" - t_in_broke_org = self._team("t-b", ["srv1"]) + t_in_broke_org = _make_team("t-b", ["srv1"]) t_in_broke_org.organization_id = "org-broke" # object_permission_id=None: the org has NO MCP ceiling, so the source is denied by the # budget gate alone. A truthy auto-Mock id here made an earlier version of this test pass # through the org-CEILING fault path with the budget gate deleted — vacuous. org = MagicMock(object_permission_id=None, litellm_budget_table=MagicMock(max_budget=5.0), spend=9.0) - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id={"t-b": t_in_broke_org}, user_teams=["t-b"], orgs_by_id={"org-broke": org}): servers = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert servers == [], "a team in an over-budget org must not grant through the union" @@ -6953,8 +6948,8 @@ class TestUserSubjectTeamUnion: axis.""" from litellm.proxy._types import LiteLLM_ObjectPermissionTable - t_ok = self._team("t-ok", ["srv1"], tool_perms={"srv1": ["read"]}) - auth = self._admitted_subject("sso-user") + t_ok = _make_team("t-ok", ["srv1"], tool_perms={"srv1": ["read"]}) + auth = _make_admitted_subject("sso-user") auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv-own"]) teams = {"t-ok": t_ok} # t-boom absent from the map -> our patched get_team_object RAISES for it @@ -6993,17 +6988,17 @@ class TestUserSubjectTeamUnion: (else this user's calls drain a bucket shared by that team's keys for access the team never provided), not for map entries beyond its grant, and never when the team is blocked.""" # t-granting grants srv1 and limits it; also names srv9 in its map, which it does NOT grant. - t_granting = self._team("t-granting", ["srv1"]) + t_granting = _make_team("t-granting", ["srv1"]) t_granting.metadata = {"mcp_rpm_limit": {"srv1": 5, "srv9": 7}} # t-other grants only srv2 but retains limit metadata for srv1 -> must not be charged for it. - t_other = self._team("t-other", ["srv2"]) + t_other = _make_team("t-other", ["srv2"]) t_other.metadata = {"mcp_rpm_limit": {"srv1": 3}} # t-blocked grants srv1 and limits it, but is blocked -> grants nothing, charges nothing. - t_blocked = self._team("t-blocked", ["srv1"]) + t_blocked = _make_team("t-blocked", ["srv1"]) t_blocked.metadata = {"mcp_rpm_limit": {"srv1": 2}} t_blocked.blocked = True - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") teams = {"t-granting": t_granting, "t-other": t_other, "t-blocked": t_blocked} with self._patch(teams_by_id=teams, user_teams=["t-granting", "t-other", "t-blocked"]): limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) @@ -7015,9 +7010,9 @@ class TestUserSubjectTeamUnion: async def test_non_roster_team_rpm_limit_does_not_apply(self): """The roster gates grants and throttles through one owner, so a team the user was removed from neither grants servers nor gets charged for their calls.""" - stale = self._team("t-stale", ["srv1"], members=("someone-else",)) + stale = _make_team("t-stale", ["srv1"], members=("someone-else",)) stale.metadata = {"mcp_rpm_limit": {"srv1": 1}} - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id={"t-stale": stale}, user_teams=["t-stale"]): limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) assert limits is None @@ -7029,7 +7024,7 @@ class TestUserSubjectTeamUnion: org_id their whole org's server list with no direct or team grant behind it.""" from litellm.proxy._types import LiteLLM_ObjectPermissionTable - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") auth.org_id = "org-a" # org allows srv1+srv2; the user and their teams grant NOTHING org_perm = AsyncMock( return_value=LiteLLM_ObjectPermissionTable(object_permission_id="op-org-a", mcp_servers=["srv1", "srv2"]) @@ -7043,8 +7038,8 @@ class TestUserSubjectTeamUnion: """Each source is resolved through an UNMARKED auth, so a fault under a source must still deny. Returning None there would win the union as allow-all and drop every team/org tool ceiling on a DB blip -- the marker alone only covers faults raised before the fan-out.""" - auth = self._admitted_subject("sso-user") - teams = {"t1": self._team("t1", ["srv1"])} + auth = _make_admitted_subject("sso-user") + teams = {"t1": _make_team("t1", ["srv1"])} # Fault INSIDE the tool resolution only. Faulting something the server path also uses would # make the source grant nothing, so the union would return [] without the tool path ever # running -- the test would pass while pinning nothing (an earlier version did exactly that). @@ -7061,11 +7056,11 @@ class TestUserSubjectTeamUnion: virtual KEY still overrides team inheritance -- that is the key ceiling model, unchanged.)""" from litellm.proxy._types import LiteLLM_ObjectPermissionTable, SpecialMCPServerNames - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") auth.object_permission = LiteLLM_ObjectPermissionTable( object_permission_id="op-user", mcp_servers=[SpecialMCPServerNames.no_mcp_servers.value] ) - with self._patch(teams_by_id={"t1": self._team("t1", ["srv1"])}, user_teams=["t1"]): + with self._patch(teams_by_id={"t1": _make_team("t1", ["srv1"])}, user_teams=["t1"]): result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv1"}, "the user's own opt-out must not zero their team's grants" @@ -7077,11 +7072,11 @@ class TestUserSubjectTeamUnion: on. Each team source carries that team's own org, which is what makes the shared resolver apply the team's owning-org ceiling rather than the caller's home org.""" teams = { - "t-member": self._team("t-member", ["srv1"], members=("sso-user",)), - "t-stale": self._team("t-stale", ["srv2"], members=("someone-else",)), + "t-member": _make_team("t-member", ["srv1"], members=("sso-user",)), + "t-stale": _make_team("t-stale", ["srv2"], members=("someone-else",)), } teams["t-member"].organization_id = "org-a" - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["t-member", "t-stale"]): sources = await MCPRequestHandler._admitted_subject_sources(auth) @@ -7100,7 +7095,7 @@ class TestUserSubjectTeamUnion: user_id and (with no team claim) no team_id, but it is NOT admission-marked, so it must keep its prior behavior of inheriting no team grants rather than silently gaining the union across every team the user belongs to.""" - teams = {"team-a": self._team("team-a", ["srv1"]), "team-b": self._team("team-b", ["srv2"])} + teams = {"team-a": _make_team("team-a", ["srv1"]), "team-b": _make_team("team-b", ["srv2"])} jwt_auth = UserAPIKeyAuth(user_id="jwt-user", api_key=None) # no admission marker with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(jwt_auth) @@ -7112,7 +7107,7 @@ class TestUserSubjectTeamUnion: metadata is caller-controlled at key creation. A user who sets ``mcp_admitted_user_subject: true`` in their own key's metadata (api_key present, no team_id) must NOT be treated as an admitted subject and must gain no cross-team union.""" - teams = {"team-a": self._team("team-a", ["srv1"]), "team-b": self._team("team-b", ["srv2"])} + teams = {"team-a": _make_team("team-a", ["srv1"]), "team-b": _make_team("team-b", ["srv2"])} forged = UserAPIKeyAuth( user_id="attacker", api_key="sk-real-key", @@ -7140,7 +7135,7 @@ class TestUserSubjectTeamUnion: mcp_tool_permissions={"srv1": ["tool_a"]}, ), ) - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id={"team-a": team}, user_teams=["team-a"]): tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) assert tools == ["tool_a"] @@ -7158,8 +7153,8 @@ class TestUserSubjectTeamUnion: access_group_ids=[], object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-blk", mcp_servers=["srv-secret"]), ) - teams = {"team-ok": self._team("team-ok", ["srv-ok"]), "team-blocked": blocked} - auth = self._admitted_subject("sso-user") + teams = {"team-ok": _make_team("team-ok", ["srv-ok"]), "team-blocked": blocked} + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-ok", "team-blocked"]): result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv-ok"} @@ -7169,8 +7164,8 @@ class TestUserSubjectTeamUnion: team's roster inherits nothing from it, even when the team id lingers in the user's (stale or cached) teams array. The team roster is the source of truth, so a removed or foreign membership revokes access at the union rather than granting it.""" - teams = {"team-x": self._team("team-x", ["srv-x"], members=("someone-else",))} - auth = self._admitted_subject("sso-user") # in user.teams for team-x, but NOT on its roster + teams = {"team-x": _make_team("team-x", ["srv-x"], members=("someone-else",))} + auth = _make_admitted_subject("sso-user") # in user.teams for team-x, but NOT on its roster with self._patch(teams_by_id=teams, user_teams=["team-x"]): result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) assert result == [] @@ -7180,7 +7175,7 @@ class TestUserSubjectTeamUnion: must DENY the server's tools ([]) rather than collapse to allow-all (None). Patches an await OUTSIDE the multi-team fan-out (the team-object lookup) to prove the whole function fails closed, not just the one helper — mirroring the fail-closed server path.""" - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with patch.object( MCPRequestHandler, "_get_team_object_permission", @@ -7209,36 +7204,6 @@ class TestAdmittedSubjectPerTeamOrgCap: user's own org — never the caller's primary org applied over the whole cross-org union. Guards the Veria 'team grants bypass their owning policies' finding.""" - def _team(self, team_id, mcp_servers, *, org_id=None, tool_perms=None, members=("sso-user",)): - from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member - - return LiteLLM_TeamTable( - team_id=team_id, - organization_id=org_id, - members_with_roles=[Member(user_id=u, role="user") for u in members], - access_group_ids=[], - object_permission=LiteLLM_ObjectPermissionTable( - object_permission_id=f"op-{team_id}", - mcp_servers=mcp_servers, - mcp_tool_permissions=tool_perms, - ), - ) - - @staticmethod - def _admitted_subject(user_id, *, org_id=None, own_servers=None, own_tool_perms=None): - from litellm.proxy._types import LiteLLM_ObjectPermissionTable - - op = None - if own_servers is not None or own_tool_perms is not None: - op = LiteLLM_ObjectPermissionTable( - object_permission_id=f"userop-{user_id}", - mcp_servers=own_servers or [], - mcp_tool_permissions=own_tool_perms, - ) - auth = UserAPIKeyAuth(user_id=user_id, api_key=None, org_id=org_id, object_permission=op) - auth.mcp_admitted_user_subject = True - return auth - #: sentinel for org_perms: org has an object_permission_id but its load returns None (a swallowed #: DB error / dangling id), which _object_permission_for_org must treat as fail-closed. LOAD_FAILS = "__load_fails__" @@ -7315,9 +7280,9 @@ class TestAdmittedSubjectPerTeamOrgCap: async def test_team_grant_capped_by_its_own_org(self): from litellm.proxy._types import LiteLLM_ObjectPermissionTable - teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a")} + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a")} org_perms = {"org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"])} - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv1"} # srv2 capped out by org-a's ceiling @@ -7326,14 +7291,14 @@ class TestAdmittedSubjectPerTeamOrgCap: from litellm.proxy._types import LiteLLM_ObjectPermissionTable teams = { - "team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a"), - "team-b": self._team("team-b", ["srv3", "srv4"], org_id="org-b"), + "team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a"), + "team-b": _make_team("team-b", ["srv3", "srv4"], org_id="org-b"), } org_perms = { "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"]), "org-b": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-b", mcp_servers=["srv3"]), } - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"], org_perms=org_perms): result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv1", "srv3"} # each team clipped by its OWN org, then unioned @@ -7341,9 +7306,9 @@ class TestAdmittedSubjectPerTeamOrgCap: async def test_all_proxy_grant_capped_by_org(self): from litellm.proxy._types import LiteLLM_ObjectPermissionTable, SpecialMCPServerName - teams = {"team-a": self._team("team-a", [SpecialMCPServerName.all_proxy_servers.value], org_id="org-a")} + teams = {"team-a": _make_team("team-a", [SpecialMCPServerName.all_proxy_servers.value], org_id="org-a")} org_perms = {"org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"])} - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch( teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms, registry=["srv1", "srv2", "srv3"] ): @@ -7353,8 +7318,8 @@ class TestAdmittedSubjectPerTeamOrgCap: assert set(result) == {"srv1"} async def test_org_row_without_object_permission_does_not_cap(self): - teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a")} - auth = self._admitted_subject("sso-user") + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a")} + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": None}): result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv1", "srv2"} # empty ceiling = no restriction @@ -7362,12 +7327,12 @@ class TestAdmittedSubjectPerTeamOrgCap: async def test_direct_grants_unioned_with_team_and_capped_by_user_org(self): from litellm.proxy._types import LiteLLM_ObjectPermissionTable - teams = {"team-a": self._team("team-a", ["srv1"], org_id="org-a")} + teams = {"team-a": _make_team("team-a", ["srv1"], org_id="org-a")} org_perms = { "org-a": None, # the team's org imposes no ceiling "org-u": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-u", mcp_servers=["srvD", "srv1"]), } - auth = self._admitted_subject("sso-user", org_id="org-u", own_servers=["srvD", "srvX"]) + auth = _make_admitted_subject("sso-user", org_id="org-u", own_servers=["srvD", "srvX"]) with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): result = await MCPRequestHandler.get_allowed_mcp_servers(auth) # direct {srvD,srvX} ∩ user-org {srvD,srv1} = {srvD}; UNIONed with team {srv1} (not intersected). @@ -7379,7 +7344,7 @@ class TestAdmittedSubjectPerTeamOrgCap: # A KEY (not admitted): the per-team org cap must NOT fire; the top-level primary-org cap applies, # byte-identical to before. team-a (org-a) grants {srv1,srv2}; the key's primary org is org-k. - teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a")} + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a")} org_perms = { "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv2"]), "org-k": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-k", mcp_servers=["srv1"]), @@ -7397,7 +7362,7 @@ class TestAdmittedSubjectPerTeamOrgCap: from litellm.proxy._types import LiteLLM_ObjectPermissionTable # team grants srv1 with NO tool restriction; org-a restricts srv1's tools to {tool_a}. - teams = {"team-a": self._team("team-a", ["srv1"], org_id="org-a")} + teams = {"team-a": _make_team("team-a", ["srv1"], org_id="org-a")} org_perms = { "org-a": LiteLLM_ObjectPermissionTable( object_permission_id="orgop-org-a", @@ -7405,7 +7370,7 @@ class TestAdmittedSubjectPerTeamOrgCap: mcp_tool_permissions={"srv1": ["tool_a"]}, ) } - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) # Without the per-team org tool ceiling this would be None (all tools) — org-a's tool ceiling @@ -7414,10 +7379,10 @@ class TestAdmittedSubjectPerTeamOrgCap: async def test_tool_union_across_cross_org_teams(self): teams = { - "team-a": self._team("team-a", ["srv1"], org_id="org-a", tool_perms={"srv1": ["t1"]}), - "team-b": self._team("team-b", ["srv1"], org_id="org-b", tool_perms={"srv1": ["t2"]}), + "team-a": _make_team("team-a", ["srv1"], org_id="org-a", tool_perms={"srv1": ["t1"]}), + "team-b": _make_team("team-b", ["srv1"], org_id="org-b", tool_perms={"srv1": ["t2"]}), } - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"], org_perms={"org-a": None, "org-b": None}): tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) assert set(tools) == {"t1", "t2"} @@ -7425,7 +7390,7 @@ class TestAdmittedSubjectPerTeamOrgCap: async def test_tool_deny_all_when_team_grant_and_org_tool_ceiling_disjoint(self): from litellm.proxy._types import LiteLLM_ObjectPermissionTable - teams = {"team-a": self._team("team-a", ["srv1"], org_id="org-a", tool_perms={"srv1": ["t1"]})} + teams = {"team-a": _make_team("team-a", ["srv1"], org_id="org-a", tool_perms={"srv1": ["t1"]})} org_perms = { "org-a": LiteLLM_ObjectPermissionTable( object_permission_id="orgop-org-a", @@ -7433,7 +7398,7 @@ class TestAdmittedSubjectPerTeamOrgCap: mcp_tool_permissions={"srv1": ["t2"]}, ) } - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) # team {t1} ∩ org {t2} = {} → deny every tool ([]), NOT allow-all (None). @@ -7446,8 +7411,8 @@ class TestAdmittedSubjectPerTeamOrgCap: synced). get_org_object RAISES a bare Exception for that; it must be treated as 'no ceiling' and must NOT lock the admitted subject out of the team's grants (parity with the key path, which tolerates a deleted org).""" - teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-gone")} - auth = self._admitted_subject("sso-user") + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-gone")} + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={}): # org-gone absent → raises result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv1", "srv2"} @@ -7456,8 +7421,8 @@ class TestAdmittedSubjectPerTeamOrgCap: """The org carries an object_permission_id but the permission load returns None (a swallowed DB error / dangling id). The ceiling cannot be verified, so the admitted subject must fail CLOSED for that team — NOT skip the ceiling, which would leak org-forbidden servers.""" - teams = {"team-a": self._team("team-a", ["srv1", "srv2"], org_id="org-a")} - auth = self._admitted_subject("sso-user") + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a")} + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": self.LOAD_FAILS}): # Asserted through the PUBLIC resolver: the per-source org ceiling is applied there now, # so calling the single-team helper would return [] for an admitted subject either way @@ -7473,9 +7438,9 @@ class TestAdmittedSubjectPerTeamOrgCap: grant would bypass every org ceiling and reach servers the user's home org forbids.""" from litellm.proxy._types import LiteLLM_ObjectPermissionTable - teams = {"team-noorg": self._team("team-noorg", ["srv1", "srv2"], org_id=None)} + teams = {"team-noorg": _make_team("team-noorg", ["srv1", "srv2"], org_id=None)} org_perms = {"org-U": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-U", mcp_servers=["srv1"])} - auth = self._admitted_subject("sso-user", org_id="org-U") + auth = _make_admitted_subject("sso-user", org_id="org-U") with self._patch(teams_by_id=teams, user_teams=["team-noorg"], org_perms=org_perms): result = await MCPRequestHandler.get_allowed_mcp_servers(auth) # org-less team falls back to the user's primary org (org-U → {srv1}); srv2 capped out. @@ -7485,8 +7450,8 @@ class TestAdmittedSubjectPerTeamOrgCap: """MEDIUM (greptile/cursor): when no source in the tool-resolution view grants the server (a TOCTOU/cache-lag inconsistency on a server that passed the server gate), the admitted path must fail CLOSED (deny all tools = []), NOT allow-all (None).""" - teams = {"team-a": self._team("team-a", ["srv1"], org_id="org-a")} - auth = self._admitted_subject("sso-user") + teams = {"team-a": _make_team("team-a", ["srv1"], org_id="org-a")} + auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": None}): # 'srv-nobody' is granted by neither the team nor the user directly → empty contributions. tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-nobody", auth) @@ -7495,9 +7460,7 @@ class TestAdmittedSubjectPerTeamOrgCap: async def test_tool_no_db_honors_in_memory_direct_restriction(self): """MEDIUM (cursor): with no DB, the tool path must still honor the user's OWN in-memory object_permission tool restriction (resolvable without a DB) rather than blanket-allow (None).""" - from litellm.proxy._types import LiteLLM_ObjectPermissionTable - - auth = self._admitted_subject( + auth = _make_admitted_subject( "sso-user", own_servers=["srv1"], own_tool_perms={"srv1": ["t1"]} ) # no org_id, direct grant of srv1 restricted to {t1} with patch("litellm.proxy.proxy_server.prisma_client", None): @@ -7522,11 +7485,11 @@ class TestAdmittedSubjectPerTeamOrgCap: # team grants the config server BY ALIAS alongside a DB-style bare id; org-a's ceiling lists # ONLY the config server (also by alias). - teams = {"team-a": self._team("team-a", ["linear_cfg", "srv-db"], org_id="org-a")} + teams = {"team-a": _make_team("team-a", ["linear_cfg", "srv-db"], org_id="org-a")} org_perms = { "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["linear_cfg"]) } - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch( teams_by_id=teams, user_teams=["team-a"], @@ -7554,12 +7517,12 @@ class TestAdmittedSubjectPerTeamOrgCap: control.server_id = "control-id" control.alias = control.server_name = control.name = "control_alias" - teams = {"team-b": self._team("team-b", ["linear_cfg", "control_alias"], org_id="org-b")} + teams = {"team-b": _make_team("team-b", ["linear_cfg", "control_alias"], org_id="org-b")} # org-b's ceiling allows ONLY the control server, referenced by its RESOLVED server_id. org_perms = { "org-b": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-b", mcp_servers=["control-id"]) } - auth = self._admitted_subject("sso-user") + auth = _make_admitted_subject("sso-user") with self._patch( teams_by_id=teams, user_teams=["team-b"], From bb388b25664c6f23379b2efeca02ce16eb145184 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 23 Jul 2026 10:19:04 -0700 Subject: [PATCH 06/25] refactor(ui): migrate workflow runs to shadcn (#34370) * test(ui): characterise the workflow runs detail drawer before migrating it Pins the drawer's behaviour against the current antd implementation: the metadata fields it surfaces, the timeline ordered by sequence number, the empty-events copy, the messages section staying collapsed until opened, the in-drawer refresh refetching, and the close control dismissing it. Every assertion is role/text based so the same file can stay green once the component moves off antd, without being edited. * refactor(ui): migrate workflow runs to shadcn Replaces the antd Drawer, Collapse, Button, Spin, Tooltip and Empty on the Workflow Runs page with the installed Base UI primitives (Sheet, Collapsible, Button, UiLoadingSpinner, Tooltip) and lucide icons, and moves the page's hardcoded hex colours, fonts and geometry onto design tokens and utility classes so the page can be themed. The only inline styles left are the gantt bars' computed left/width, which are runtime values. Behaviour is unchanged: the drawer's characterisation tests were written against the antd version in the previous commit and pass here without being edited. Retires the file's now-unused antd no-restricted-imports suppression. --- ui/litellm-dashboard/eslint-suppressions.json | 3 - .../workflows/WorkflowRuns.test.tsx | 150 ++++- .../(dashboard)/workflows/WorkflowRuns.tsx | 630 +++++++----------- 3 files changed, 378 insertions(+), 405 deletions(-) diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 6335de32147..1bd3fceb6ff 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -2134,9 +2134,6 @@ "no-nested-ternary": { "count": 1 }, - "no-restricted-imports": { - "count": 1 - }, "no-restricted-syntax": { "count": 3 }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.test.tsx index 1aee3fcc8ab..6707330a54c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.test.tsx @@ -1,4 +1,4 @@ -import { render, screen, waitFor } from "@testing-library/react"; +import { render, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { afterEach, describe, expect, it, vi } from "vitest"; @@ -96,3 +96,151 @@ describe("WorkflowRuns (migrated onto shared DataTable)", () => { } }); }); + +interface FakeEvent { + event_id: string; + event_type: string; + step_name: string; + sequence_number: number; + created_at: string; + data: Record | null; +} + +interface FakeMessage { + message_id: string; + role: string; + content: string; + sequence_number: number; + created_at: string; +} + +const DETAIL_RUN = { + run_id: "run-aaaaaaaa-1111", + status: "completed", + workflow_type: "grill", + created_at: "2026-01-01T00:00:00Z", + metadata: { title: "First run", state: "done", pr_url: "https://example.com/pr/1", worktree_path: "/tmp/wt" }, +}; + +const DETAIL_EVENTS: FakeEvent[] = [ + { + event_id: "ev-2", + event_type: "hook.waiting", + step_name: "review", + sequence_number: 2, + created_at: "2026-01-01T00:00:05Z", + data: null, + }, + { + event_id: "ev-1", + event_type: "step.started", + step_name: "plan", + sequence_number: 1, + created_at: "2026-01-01T00:00:01Z", + data: { attempt: 1 }, + }, +]; + +const DETAIL_MESSAGES: FakeMessage[] = [ + { + message_id: "msg-1", + role: "user", + content: "kick off the run", + sequence_number: 1, + created_at: "2026-01-01T00:00:02Z", + }, +]; + +function mockDetailFetch(events: FakeEvent[], messages: FakeMessage[]) { + return vi.fn((url: string) => { + if (url.includes("/runs?limit")) { + return Promise.resolve({ ok: true, json: () => Promise.resolve({ runs: [DETAIL_RUN] }) }); + } + if (url.includes("/events")) { + return Promise.resolve({ ok: true, json: () => Promise.resolve({ events }) }); + } + if (url.includes("/messages")) { + return Promise.resolve({ ok: true, json: () => Promise.resolve({ messages }) }); + } + return Promise.resolve({ ok: false, status: 404, json: () => Promise.resolve({}) }); + }); +} + +async function openDetailDrawer(events = DETAIL_EVENTS, messages = DETAIL_MESSAGES) { + const user = userEvent.setup(); + const fetchSpy = mockDetailFetch(events, messages); + vi.stubGlobal("fetch", fetchSpy); + render(); + + await user.click(await screen.findByText("First run")); + const drawer = await screen.findByRole("dialog"); + await waitFor(() => expect(within(drawer).getByText("Timeline")).toBeInTheDocument()); + return { user, fetchSpy, drawer }; +} + +describe("WorkflowRuns detail drawer", () => { + it("shows the run's identity and metadata fields", async () => { + const { drawer } = await openDetailDrawer(); + + expect(within(drawer).getAllByText("First run")).toHaveLength(2); + expect(within(drawer).getByText("run-aaaa")).toBeInTheDocument(); + expect(within(drawer).getByText("grill")).toBeInTheDocument(); + expect(within(drawer).getByText("completed")).toBeInTheDocument(); + expect(within(drawer).getByText("done")).toBeInTheDocument(); + expect(within(drawer).getByText("/tmp/wt")).toBeInTheDocument(); + expect(within(drawer).getByRole("link", { name: "https://example.com/pr/1" })).toHaveAttribute( + "href", + "https://example.com/pr/1", + ); + }); + + it("renders every event in the timeline, ordered by sequence number", async () => { + const { drawer } = await openDetailDrawer(); + + expect(within(drawer).getByText("2 events")).toBeInTheDocument(); + + const stepLabels = within(drawer) + .getAllByText(/^(plan|review)$/) + .map((el) => el.textContent); + expect(stepLabels).toEqual(["plan", "review"]); + + expect(within(drawer).getByText("step.started")).toBeInTheDocument(); + expect(within(drawer).getByText("hook.waiting")).toBeInTheDocument(); + }); + + it("says no events were recorded when the run has none", async () => { + const { drawer } = await openDetailDrawer([], DETAIL_MESSAGES); + + expect(within(drawer).getByText("No events recorded")).toBeInTheDocument(); + }); + + it("keeps the messages section collapsed until it is opened", async () => { + const { user, drawer } = await openDetailDrawer(); + + expect(within(drawer).queryByText("kick off the run")).not.toBeInTheDocument(); + + await user.click(within(drawer).getByRole("button", { name: /Messages/ })); + + expect(await within(drawer).findByText("kick off the run")).toBeInTheDocument(); + expect(within(drawer).getByText("[user]")).toBeInTheDocument(); + }); + + it("refetches events and messages when the drawer's refresh button is clicked", async () => { + const { user, fetchSpy, drawer } = await openDetailDrawer(); + + const eventFetches = () => fetchSpy.mock.calls.filter(([url]) => String(url).includes("/events")).length; + expect(eventFetches()).toBe(1); + + await user.click(within(drawer).getByRole("button", { name: /refresh/i })); + + await waitFor(() => expect(eventFetches()).toBe(2)); + }); + + it("dismisses the drawer when its close control is clicked", async () => { + const { user, drawer } = await openDetailDrawer(); + + await user.click(within(drawer).getByRole("button", { name: /close/i })); + + await waitFor(() => expect(screen.queryAllByRole("dialog")).toHaveLength(0)); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx index 7354b7479f4..5efd56257e2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx @@ -1,6 +1,5 @@ import React, { useState, useEffect, useCallback, useMemo } from "react"; -import { Button, Collapse, Drawer, Empty, Spin, Tooltip, Typography } from "antd"; -import { ReloadOutlined } from "@ant-design/icons"; +import { ArrowLeft, ChevronDown, RefreshCw } from "lucide-react"; import type { ColumnDef, ColumnFiltersState } from "@tanstack/react-table"; import { getGlobalLitellmHeaderName, proxyBaseUrl } from "@/components/networking"; import { @@ -9,10 +8,14 @@ import { DataTableFilterField, DataTableToolbar, } from "@/components/shared/DataTable"; +import { Button } from "@/components/ui/button"; +import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; import { Input } from "@/components/ui/input"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; - -const { Text } = Typography; +import { Sheet, SheetContent, SheetDescription, SheetTitle } from "@/components/ui/sheet"; +import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; +import { cn } from "@/lib/cva.config"; interface WorkflowRunsProps { accessToken: string | null; @@ -59,11 +62,11 @@ interface WorkflowRunMessage { // ── design tokens ───────────────────────────────────────────────────────────── const STATUS_DOT: Record = { - pending: "#a1a1aa", - running: "#3b82f6", - paused: "#f59e0b", - completed: "#22c55e", - failed: "#ef4444", + pending: "bg-gray-400", + running: "bg-blue-500", + paused: "bg-amber-500", + completed: "bg-green-500", + failed: "bg-red-500", }; const RUN_STATUS_OPTIONS: RunStatus[] = ["pending", "running", "paused", "completed", "failed"]; @@ -75,15 +78,15 @@ const STATUS_LABELS: Record = { failed: "Failed", }; -const EVENT_COLOR: Record = { - "step.started": { bar: "#f0fdf4", border: "#86efac", text: "#16a34a" }, - "step.failed": { bar: "#fef2f2", border: "#fca5a5", text: "#dc2626" }, - "hook.waiting": { bar: "#fffbeb", border: "#fcd34d", text: "#d97706" }, - "hook.received": { bar: "#eff6ff", border: "#93c5fd", text: "#2563eb" }, +const EVENT_COLOR: Record = { + "step.started": { bar: "border-green-300 bg-green-50", text: "text-green-600" }, + "step.failed": { bar: "border-red-300 bg-red-50", text: "text-red-600" }, + "hook.waiting": { bar: "border-amber-300 bg-amber-50", text: "text-amber-600" }, + "hook.received": { bar: "border-blue-300 bg-blue-50", text: "text-blue-600" }, }; function eventStyle(type: string) { - return EVENT_COLOR[type] ?? { bar: "#f4f4f5", border: "#d4d4d8", text: "#52525b" }; + return EVENT_COLOR[type] ?? { bar: "border-border bg-muted", text: "text-muted-foreground" }; } // ── helpers ─────────────────────────────────────────────────────────────────── @@ -118,17 +121,8 @@ function shortId(id: string): string { // ── status dot ──────────────────────────────────────────────────────────────── -const StatusDot: React.FC<{ status: RunStatus; size?: number }> = ({ status, size = 8 }) => ( - +const StatusDot: React.FC<{ status: RunStatus; className?: string }> = ({ status, className }) => ( + ); // ── truncated text value ────────────────────────────────────────────────────── @@ -138,25 +132,14 @@ const TRUNCATE_AT = 120; const TruncatedValue: React.FC<{ value: string }> = ({ value }) => { const [expanded, setExpanded] = useState(false); if (value.length <= TRUNCATE_AT) { - return {value}; + return {value}; } return ( - + {expanded ? value : value.slice(0, TRUNCATE_AT) + "…"} - + ); }; @@ -179,67 +162,24 @@ const MetadataCard: React.FC<{ run: WorkflowRun }> = ({ run }) => { ); return ( -
+
{/* title bar */} -
- - {runTitle(run)} - +
+ + {runTitle(run)} + {shortId(run.run_id)} - - {run.workflow_type} - + {run.workflow_type}
{/* key fields grid */} -
+
- {run.status} + {run.status} - {timeAgo(run.created_at)} + {timeAgo(run.created_at)} {meta.pr_url && ( @@ -248,7 +188,7 @@ const MetadataCard: React.FC<{ run: WorkflowRun }> = ({ run }) => { href={String(meta.pr_url)} target="_blank" rel="noopener noreferrer" - style={{ color: "#2563eb", textDecoration: "none", wordBreak: "break-all" }} + className="break-all text-primary underline-offset-4 hover:underline" > {String(meta.pr_url)} @@ -280,9 +220,9 @@ const MetadataCard: React.FC<{ run: WorkflowRun }> = ({ run }) => { }; const FieldPair: React.FC<{ label: string; children: React.ReactNode }> = ({ label, children }) => ( -
- {label} - {children} +
+ {label} + {children}
); @@ -293,11 +233,7 @@ const GanttTimeline: React.FC<{ events: WorkflowRunEvent[]; }> = ({ run, events }) => { if (events.length === 0) { - return ( -
- No events recorded -
- ); + return
No events recorded
; } const runStart = new Date(run.created_at).getTime(); @@ -307,187 +243,139 @@ const GanttTimeline: React.FC<{ const totalDur = fmtDuration(lastTime - runStart); return ( -
- {/* ruler */} -
-
-
- {[0, 100].map((pct) => ( - - {pct === 0 ? "0" : totalDur} - - ))} + +
+ {/* ruler */} +
+
+
+ 0 + {totalDur} +
-
- {/* outer run bar */} -
-
- {runTitle(run)} + {/* outer run bar */} +
+
{runTitle(run)}
+
+ {totalDur} +
-
- {totalDur} -
-
- {/* event rows */} -
- {events.map((ev) => { - const evTime = new Date(ev.created_at).getTime(); - const leftPct = ((evTime - runStart) / totalSpan) * 100; + {/* event rows */} +
+ {events.map((ev) => { + const evTime = new Date(ev.created_at).getTime(); + const leftPct = ((evTime - runStart) / totalSpan) * 100; - const nextIdx = events.findIndex((e) => e.sequence_number > ev.sequence_number); - const nextTime = - nextIdx >= 0 ? new Date(events[nextIdx].created_at).getTime() : lastTime + Math.max(totalSpan * 0.12, 500); - const widthPct = Math.max(8, ((nextTime - evTime) / totalSpan) * 100); - const style = eventStyle(ev.event_type); - const dur = fmtDuration(nextTime - evTime); + const nextIdx = events.findIndex((e) => e.sequence_number > ev.sequence_number); + const nextTime = + nextIdx >= 0 + ? new Date(events[nextIdx].created_at).getTime() + : lastTime + Math.max(totalSpan * 0.12, 500); + const widthPct = Math.max(8, ((nextTime - evTime) / totalSpan) * 100); + const style = eventStyle(ev.event_type); + const dur = fmtDuration(nextTime - evTime); - return ( - -
- {ev.step_name || ev.event_type} -
-
- -
- type: - {ev.event_type} -
-
- step: - {ev.step_name} -
-
- seq: - {ev.sequence_number} -
-
- time: - {timeAgo(ev.created_at)} -
- {ev.data && Object.keys(ev.data).length > 0 && ( + return ( + +
{ev.step_name || ev.event_type}
+
+ + + } + > + {ev.event_type} + {dur && {dur}} + + +
- data: - {JSON.stringify(ev.data)} + type: + {ev.event_type}
- )} -
- } - > -
- {ev.event_type} - {dur && {dur}} -
-
-
-
- ); - })} +
+ step: + {ev.step_name} +
+
+ seq: + {ev.sequence_number} +
+
+ time: + {timeAgo(ev.created_at)} +
+ {ev.data && Object.keys(ev.data).length > 0 && ( +
+ data: + {JSON.stringify(ev.data)} +
+ )} +
+ + +
+ + ); + })} +
-
+
); }; // ── message row ─────────────────────────────────────────────────────────────── -const MessageRow: React.FC<{ msg: WorkflowRunMessage }> = ({ msg }) => { - const roleColor: Record = { - user: "#2563eb", - assistant: "#16a34a", - system: "#7c3aed", - tool_result: "#d97706", - }; - const color = roleColor[msg.role] ?? "#52525b"; - - return ( -
- [{msg.role}] -
- - {msg.content} - - - {timeAgo(msg.created_at)} - -
-
- ); +const ROLE_COLOR: Record = { + user: "text-blue-600", + assistant: "text-green-600", + system: "text-violet-600", + tool_result: "text-amber-600", }; +const MessageRow: React.FC<{ msg: WorkflowRunMessage }> = ({ msg }) => ( +
+ [{msg.role}] +
+ {msg.content} + {timeAgo(msg.created_at)} +
+
+); + +// ── collapsible section ─────────────────────────────────────────────────────── + +const DetailSection: React.FC<{ + title: string; + meta: React.ReactNode; + defaultOpen?: boolean; + children: React.ReactNode; +}> = ({ title, meta, defaultOpen = false, children }) => ( + + + + + {title} + {meta} + + + {children} + +); + // ── main component ──────────────────────────────────────────────────────────── const WorkflowRuns: React.FC = ({ accessToken }) => { @@ -572,11 +460,11 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { cell: ({ row }) => { const run = row.original; return ( -
- +
+
-
{runTitle(run)}
-
{shortId(run.run_id)}
+
{runTitle(run)}
+
{shortId(run.run_id)}
); @@ -588,7 +476,7 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { meta: { title: "Type" }, filterFn: "includesString", cell: ({ row }) => ( - {row.original.workflow_type} + {row.original.workflow_type} ), }, { @@ -600,11 +488,9 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { cell: ({ row }) => { const run = row.original; return ( -
- - - {run.metadata?.state ?? run.status} - +
+ + {run.metadata?.state ?? run.status}
); }, @@ -613,26 +499,18 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { accessorKey: "created_at", header: "Created", meta: { title: "Created" }, - cell: ({ row }) => {timeAgo(row.original.created_at)}, + cell: ({ row }) => {timeAgo(row.original.created_at)}, }, ], [], ); return ( -
+
{/* page header */} -
-
Workflow Runs
-
+
+
Workflow Runs
+
Durable state tracking for agents and automated workflows
@@ -643,12 +521,7 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { getRowId={(run) => run.run_id} isLoading={loadingRuns} loadingMessage="Loading workflow runs…" - noDataMessage={ - No workflow runs yet} - image={Empty.PRESENTED_IMAGE_SIMPLE} - /> - } + noDataMessage={
No workflow runs yet
} paginationMode="client" pageSizeOptions={[50, 100]} filterMode="client" @@ -712,115 +585,70 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { /> {/* detail drawer */} - setDrawerOpen(false)} - width={680} - title={null} - closable={false} - bodyStyle={{ padding: 0 }} - styles={{ body: { padding: 0 } }} - > - {!selectedRun ? null : loadingDetail ? ( -
- -
- ) : ( -
- {/* drawer close + refresh */} -
- - + + + Workflow run details + + Metadata, timeline and messages for the selected workflow run + + {!selectedRun ? null : loadingDetail ? ( +
+
+ ) : ( +
+ {/* drawer close + refresh */} +
+ + +
- {/* metadata card — top */} - + {/* metadata card — top */} + - {/* collapsible sections */} - - Timeline - - {events.length} {events.length === 1 ? "event" : "events"} - - - ), - children: ( -
- + {/* collapsible sections */} +
+ + {events.length} {events.length === 1 ? "event" : "events"} + + } + defaultOpen + > + + + + {messages.length === 0 ? ( +
No messages
+ ) : ( +
+ {messages.map((msg) => ( + + ))}
- ), - }, - { - key: "messages", - label: ( - - Messages - - {messages.length} - - - ), - children: - messages.length === 0 ? ( -
- No messages -
- ) : ( -
- {messages.map((msg) => ( - - ))} -
- ), - }, - ]} - /> -
- )} - + )} + +
+
+ )} +
+
); }; From fd494d2fb236783ff7f004de177f4f7adff6dc41 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 23 Jul 2026 10:20:16 -0700 Subject: [PATCH 07/25] refactor(ui): migrate models and endpoints table onto the shared DataTable (#34363) * refactor(ui): migrate models and endpoints table onto the shared DataTable Rebuilds the All Models table on the shared DataTable, following the 2a treatment from the Models + Endpoints design: one card holding search, the Team and View selectors, refresh, columns and filters, with the active filters on a chip row and the pagination footer at the bottom. Retires the last hand-rolled tremor renderer (all_models_table.tsx) and the antd/tremor column defs in molecules/models/columns.tsx, replacing them with a thin AllModelsTable consumer plus AllModelsTableColumns built from the shared cell library. Behavior is preserved end to end. The server sort field mapping now lives next to the column ids so the two cannot drift. Status keeps its column and its sort, hidden by default behind the Columns menu because the design shows nine columns. Access groups collapse into a "+N more" tooltip instead of a per-row expand toggle, and the full reset moves into the filter drawer footer where the design puts it. Adds the shadcn hover-card primitive (Base UI PreviewCard in the base-vega style) for the model information hover, which needs an interactive surface a tooltip cannot provide. * fix(ui): stop the models tab re-querying on mount The mount-time effect fires the debounced search with the initial empty value, and its callback rebuilt the pagination object unconditionally. That produced a second render (and a second query) roughly 300ms after mount with no user input, which on a slow CI machine swapped the table's row nodes mid-interaction and made a click land on a detached node. resetToFirstPage now returns the existing state when already on the first page, so React bails out instead of re-rendering. Pinned with a test that asserts no additional query after the debounce settles; it fails without the fix. --- ui/litellm-dashboard/eslint-suppressions.json | 57 +- .../components/AllModelsTab.test.tsx | 826 ++++--------- .../components/AllModelsTab.tsx | 631 +++------- .../components/AllModelsTable.test.tsx | 403 ++++++ .../components/AllModelsTable.tsx | 298 +++++ .../components/ModelsTableColumns.tsx | 488 ++++++++ .../model_dashboard/all_models_table.tsx | 221 ---- .../src/components/model_dashboard/types.ts | 1 + .../molecules/models/columns.test.tsx | 1099 ----------------- .../components/molecules/models/columns.tsx | 423 ------- .../DataTable/DataTableFilterDrawer.tsx | 7 + .../src/components/ui/hover-card.tsx | 46 + 12 files changed, 1704 insertions(+), 2796 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx delete mode 100644 ui/litellm-dashboard/src/components/model_dashboard/all_models_table.tsx delete mode 100644 ui/litellm-dashboard/src/components/molecules/models/columns.test.tsx delete mode 100644 ui/litellm-dashboard/src/components/molecules/models/columns.tsx create mode 100644 ui/litellm-dashboard/src/components/ui/hover-card.tsx diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 1bd3fceb6ff..6a50f4aa99e 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1143,25 +1143,6 @@ "count": 1 } }, - "src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx": { - "max-params": { - "count": 1 - }, - "unused-imports/no-unused-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx": { - "local/no-complex-jsx-arrow": { - "count": 2 - }, - "no-restricted-imports": { - "count": 2 - }, - "react-hooks/set-state-in-effect": { - "count": 3 - } - }, "src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.test.tsx": { "react/display-name": { "count": 1 @@ -3483,17 +3464,6 @@ "count": 1 } }, - "src/components/model_dashboard/all_models_table.tsx": { - "local/filename-pascal-case": { - "count": 1 - }, - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/model_filters.tsx": { "local/filename-pascal-case": { "count": 1 @@ -3560,28 +3530,6 @@ "count": 2 } }, - "src/components/molecules/models/columns.test.tsx": { - "no-restricted-imports": { - "count": 1 - }, - "react/display-name": { - "count": 1 - } - }, - "src/components/molecules/models/columns.tsx": { - "local/filename-pascal-case": { - "count": 1 - }, - "max-params": { - "count": 1 - }, - "no-nested-ternary": { - "count": 2 - }, - "no-restricted-imports": { - "count": 2 - } - }, "src/components/molecules/notifications_manager.test.tsx": { "no-restricted-imports": { "count": 1 @@ -4179,6 +4127,11 @@ "count": 1 } }, + "src/components/ui/hover-card.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/ui/input-group.tsx": { "local/filename-pascal-case": { "count": 1 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx index 045bf0a5f44..659a618c8f6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx @@ -1,662 +1,364 @@ import * as useAuthorizedModule from "@/app/(dashboard)/hooks/useAuthorized"; -import { fireEvent, render, screen, waitFor } from "@testing-library/react"; -import { renderWithProviders } from "../../../../../tests/test-utils"; +import { render, screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import AllModelsTab from "./AllModelsTab"; -// Mock modelDeleteCall +import AllModelsTab from "./AllModelsTab"; +import { STATUS_COLUMN_ID, toServerSortField } from "./ModelsTableColumns"; + const mockModelDeleteCall = vi.fn().mockResolvedValue({}); +const mockModelPatchUpdateCall = vi.fn().mockResolvedValue({}); vi.mock("@/components/networking", () => ({ - modelDeleteCall: (...args: any[]) => mockModelDeleteCall(...args), + modelDeleteCall: (...args: unknown[]) => mockModelDeleteCall(...args), + modelPatchUpdateCall: (...args: unknown[]) => mockModelPatchUpdateCall(...args), })); -// Mock NotificationsManager vi.mock("@/components/molecules/notifications_manager", () => ({ - default: { - success: vi.fn(), - fromBackend: vi.fn(), + default: { success: vi.fn(), fromBackend: vi.fn() }, +})); + +vi.mock("@/components/model_dashboard/ModelSettingsModal/ModelSettingsModal", () => ({ + default: function ModelSettingsModalMock({ isVisible }: { isVisible: boolean }) { + return isVisible ?
: null; }, })); -// Mock react-query const mockInvalidateQueries = vi.fn(); vi.mock("@tanstack/react-query", async (importOriginal) => { - const actual = (await importOriginal()) as any; - return { - ...actual, - useQueryClient: () => ({ - invalidateQueries: mockInvalidateQueries, - }), - }; + const actual = await importOriginal(); + return { ...actual, useQueryClient: () => ({ invalidateQueries: mockInvalidateQueries }) }; }); -// Mock the useModelsInfo hook -const mockUseModelsInfo = vi.fn(() => ({ - data: { data: [], total_count: 0, current_page: 1, total_pages: 1, size: 50 }, - isLoading: false, - error: null, -})) as any; +interface ModelsInfoArgs { + page?: number; + size?: number; + search?: string; + teamId?: string; + sortBy?: string; + sortOrder?: string; +} + +const modelsInfoCalls: ModelsInfoArgs[] = []; +const mockRefetch = vi.fn(); +let modelsInfoResult: Record = {}; + +type UseModelsInfoArgs = [ + page?: number, + size?: number, + search?: string, + modelId?: string, + teamId?: string, + sortBy?: string, + sortOrder?: string, +]; vi.mock("../../hooks/models/useModels", () => ({ - useModelsInfo: (page?: number, size?: number, search?: string) => mockUseModelsInfo(page, size, search), -})); - -// Mock the useModelCostMap hook -const mockUseModelCostMap = vi.fn(() => ({ - data: { - "gpt-4": { litellm_provider: "openai" }, - "gpt-3.5-turbo": { litellm_provider: "openai" }, - "gpt-4-accessible": { litellm_provider: "openai" }, - "gpt-3.5-turbo-blocked": { litellm_provider: "openai" }, - "gpt-4-sales": { litellm_provider: "openai" }, - "gpt-4-engineering": { litellm_provider: "openai" }, - "gpt-4-personal": { litellm_provider: "openai" }, - "gpt-4-team-only": { litellm_provider: "openai" }, - "gpt-4-config": { litellm_provider: "openai" }, - "gpt-4-db": { litellm_provider: "openai" }, + useModelsInfo: (...args: UseModelsInfoArgs) => { + const [page, size, search, , teamId, sortBy, sortOrder] = args; + const call: ModelsInfoArgs = { page, size, search, teamId, sortBy, sortOrder }; + modelsInfoCalls.push(call); + return { ...modelsInfoResult, refetch: mockRefetch }; }, - isLoading: false, - error: null, -})) as any; +})); vi.mock("../../hooks/models/useModelCostMap", () => ({ - useModelCostMap: () => mockUseModelCostMap(), + useModelCostMap: () => ({ data: { "gpt-4": { litellm_provider: "openai" } }, isLoading: false, error: null }), })); -// Mock the useTeams hook (react-query implementation) -const mockUseTeams = vi.fn(() => ({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), -})) as any; - +const mockTeams = [{ team_id: "team-1", team_alias: "Engineering" }]; vi.mock("../../hooks/teams/useTeams", () => ({ - useTeams: () => mockUseTeams(), + useTeams: () => ({ data: mockTeams, isLoading: false, error: null, refetch: vi.fn() }), })); -// Helper function to create model cost map mock return value -const createModelCostMapMock = (data: Record) => ({ - data, - isLoading: false, - error: null, +const BASE_MODEL_INFO = { + id: "model-1", + db_model: true, + created_by: "user-123", + created_at: "2024-01-01T00:00:00Z", + updated_at: "2024-01-02T00:00:00Z", + team_id: "team-1", + access_groups: [], +}; + +const makeRow = (overrides: Record = {}) => ({ + model_name: "gpt-4", + litellm_params: { model: "openai/gpt-4", custom_llm_provider: "openai" }, + model_info: { ...BASE_MODEL_INFO, ...((overrides.model_info as Record) ?? {}) }, }); -// Helper function to create paginated model data mock -const createPaginatedModelData = ( - models: any[], - totalCount: number = models.length, - currentPage: number = 1, - totalPages: number = 1, - size: number = 50, -) => ({ - data: models, - total_count: totalCount, - current_page: currentPage, - total_pages: totalPages, - size: size, -}); +const setModelsInfo = (rows: Record[], totalCount = rows.length, isLoading = false) => { + modelsInfoResult = { + data: { data: rows, total_count: totalCount, current_page: 1, total_pages: 1, size: 50 }, + isLoading, + isFetching: false, + error: null, + }; +}; + +const lastModelsInfoCall = (): ModelsInfoArgs => modelsInfoCalls[modelsInfoCalls.length - 1]; + +const SEARCH_SETTLE_MS = 400; + +const MOCK_AUTHORIZED = { + isLoading: false, + isAuthorized: true, + token: "mock-token", + accessToken: "mock-access-token", + userId: "user-123", + userEmail: "test@example.com", + userRole: "Admin", + premiumUser: true, + disabledPersonalKeyCreation: false, + showSSOBanner: false, +}; + +const mockSetSelectedModelGroup = vi.fn(); +const mockSetSelectedModelId = vi.fn(); +const mockSetSelectedTeamId = vi.fn(); + +const defaultProps = { + selectedModelGroup: "all", + setSelectedModelGroup: mockSetSelectedModelGroup, + availableModelGroups: ["gpt-4", "gpt-3.5-turbo"], + availableModelAccessGroups: ["sales-team"], + setSelectedModelId: mockSetSelectedModelId, + setSelectedTeamId: mockSetSelectedTeamId, +}; describe("AllModelsTab", () => { - const mockSetSelectedModelGroup = vi.fn(); - const mockSetSelectedModelId = vi.fn(); - const mockSetSelectedTeamId = vi.fn(); - - const defaultProps = { - selectedModelGroup: "all", - setSelectedModelGroup: mockSetSelectedModelGroup, - availableModelGroups: ["gpt-4", "gpt-3.5-turbo"], - availableModelAccessGroups: ["sales-team", "engineering-team"], - setSelectedModelId: mockSetSelectedModelId, - setSelectedTeamId: mockSetSelectedTeamId, - }; - - const mockUseAuthorized = { - token: "mock-token", - accessToken: "mock-access-token", - userId: "user-123", - userEmail: "test@example.com", - userRole: "Admin", - premiumUser: true, - disabledPersonalKeyCreation: false, - showSSOBanner: false, - }; - beforeEach(() => { vi.clearAllMocks(); - vi.spyOn(useAuthorizedModule, "default").mockReturnValue(mockUseAuthorized); + modelsInfoCalls.length = 0; + setModelsInfo([makeRow()]); + vi.spyOn(useAuthorizedModule, "default").mockReturnValue(MOCK_AUTHORIZED); }); - it("should render with empty data", () => { - mockUseModelsInfo.mockReturnValueOnce({ - data: createPaginatedModelData([], 0, 1, 1, 50), - isLoading: false, - error: null, - }); + it("renders the fetched models and the server row count", async () => { + setModelsInfo([makeRow()], 137); + render(); - mockUseTeams.mockReturnValueOnce({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), - }); - - mockUseModelCostMap.mockReturnValueOnce(createModelCostMapMock({})); - - renderWithProviders(); - expect(screen.getByText("Current Team:")).toBeInTheDocument(); + expect(await screen.findByText("gpt-4")).toBeInTheDocument(); + expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-50 of 137"); }); - it("should filter models by direct team access when current team is selected", async () => { - const mockTeams = [ - { - team_id: "team-456", - team_alias: "Engineering Team", - models: ["gpt-4"], - max_budget: null, - budget_duration: null, - tpm_limit: null, - rpm_limit: null, - organization_id: "org-123", - created_at: "2024-01-01", - keys: [], - members_with_roles: [], - }, + it("does not re-query after the mount-time debounced search settles unchanged", async () => { + render(); + const callsAfterMount = modelsInfoCalls.length; + + await new Promise((resolve) => setTimeout(resolve, SEARCH_SETTLE_MS)); + + expect(modelsInfoCalls.length).toBe(callsAfterMount); + }); + + it("shows the empty state when the proxy returns no models", () => { + setModelsInfo([], 0); + render(); + + expect(screen.getByText("No models found")).toBeInTheDocument(); + }); + + it("shows the loading skeleton while the first page is in flight", () => { + setModelsInfo([], 0, true); + render(); + + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByText("No models found")).not.toBeInTheDocument(); + }); + + describe("server sort contract", () => { + const sortHeader = (columnId: string): HTMLElement => screen.getByTestId(`sort-header-${columnId}`); + + const expectIndicator = async (columnId: string, state: "asc" | "desc" | "none") => { + await waitFor(() => { + expect(sortHeader(columnId).querySelector(`[data-sort-indicator="${state}"]`)).not.toBeNull(); + }); + }; + + const cases: [string, string, string, "asc" | "desc"][] = [ + ["Model Information", "model_name", "model_name", "asc"], + ["Created By", "model_info_created_by", "created_at", "asc"], + ["Updated At", "model_info_updated_at", "updated_at", "asc"], + ["Costs", "input_cost", "costs", "desc"], ]; - mockUseTeams.mockReturnValueOnce({ - data: mockTeams, - isLoading: false, - error: null, - refetch: vi.fn(), + it.each(cases)("sorts %s using the server field %s", async (_label, columnId, serverField, firstDirection) => { + const user = userEvent.setup(); + render(); + + await user.click(sortHeader(columnId)); + await expectIndicator(columnId, firstDirection); + + expect(lastModelsInfoCall().sortBy).toBe(serverField); + expect(lastModelsInfoCall().sortOrder).toBe(firstDirection); }); - mockUseModelCostMap.mockReturnValueOnce( - createModelCostMapMock({ - "gpt-4-accessible": { litellm_provider: "openai" }, - "gpt-3.5-turbo-blocked": { litellm_provider: "openai" }, - }), - ); + it("maps the hidden Status column to the server field status", () => { + expect(toServerSortField(STATUS_COLUMN_ID)).toBe("status"); + }); - const modelData = createPaginatedModelData( - [ - { - model_name: "gpt-4-accessible", - model_info: { - id: "model-1", - access_via_team_ids: ["team-456"], - access_groups: [], - }, - }, - { - model_name: "gpt-3.5-turbo-blocked", - model_info: { - id: "model-2", - access_via_team_ids: ["team-789"], - access_groups: [], - }, - }, - ], - 2, - 1, - 1, - 50, - ); + it("cycles a sorted column back to unsorted", async () => { + const user = userEvent.setup(); + render(); - mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); + await user.click(sortHeader("model_info_updated_at")); + await expectIndicator("model_info_updated_at", "asc"); + expect(lastModelsInfoCall().sortOrder).toBe("asc"); - renderWithProviders(); + await user.click(sortHeader("model_info_updated_at")); + await expectIndicator("model_info_updated_at", "desc"); + expect(lastModelsInfoCall().sortOrder).toBe("desc"); - // Component shows API total_count (2), not filtered count - // Since default is "personal" team and models don't have direct_access, they're filtered out - await waitFor(() => { - expect(screen.getByText("Showing 1 - 2 of 2 results")).toBeInTheDocument(); + await user.click(sortHeader("model_info_updated_at")); + await expectIndicator("model_info_updated_at", "none"); + expect(lastModelsInfoCall().sortBy).toBeUndefined(); }); }); - it("should filter models by access group matching when team models match model access groups", async () => { - const mockTeams = [ - { - team_id: "team-sales", - team_alias: "Sales Team", - models: ["sales-model-group"], - max_budget: null, - budget_duration: null, - tpm_limit: null, - rpm_limit: null, - organization_id: "org-123", - created_at: "2024-01-01", - keys: [], - members_with_roles: [], - }, - ]; + it("queries the selected team and resets to the first page", async () => { + const user = userEvent.setup(); + render(); - mockUseTeams.mockReturnValue({ - data: mockTeams, - isLoading: false, - error: null, - refetch: vi.fn(), - }); + expect(lastModelsInfoCall().teamId).toBeUndefined(); - mockUseModelCostMap.mockReturnValueOnce( - createModelCostMapMock({ - "gpt-4-sales": { litellm_provider: "openai" }, - "gpt-4-engineering": { litellm_provider: "openai" }, - }), - ); + await user.click(screen.getByTestId("models-team-select")); + await user.click(await screen.findByRole("option", { name: "Engineering" })); - const modelData = createPaginatedModelData( - [ - { - model_name: "gpt-4-sales", - model_info: { - id: "model-sales-1", - access_via_team_ids: [], - access_groups: ["sales-model-group"], - }, - }, - { - model_name: "gpt-4-engineering", - model_info: { - id: "model-eng-1", - access_via_team_ids: [], - access_groups: ["engineering-model-group"], - }, - }, - ], - 2, - 1, - 1, - 50, - ); - - mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); - - renderWithProviders(); - - // Component shows API total_count (2), not filtered count - // Since default is "personal" team and models don't have direct_access, they're filtered out await waitFor(() => { - expect(screen.getByText("Showing 1 - 2 of 2 results")).toBeInTheDocument(); + expect(lastModelsInfoCall().teamId).toBe("team-1"); + }); + expect(lastModelsInfoCall().page).toBe(1); + }); + + it("debounces the model name search into the server query", async () => { + const user = userEvent.setup(); + render(); + + await user.type(screen.getByTestId("datatable-search"), "claude"); + + await waitFor(() => { + expect(lastModelsInfoCall().search).toBe("claude"); }); }); - it("should filter models by direct_access for personal team", async () => { - mockUseTeams.mockReturnValue({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), - }); + it("applies a public model name filter through the drawer", async () => { + const user = userEvent.setup(); + render(); - mockUseModelCostMap.mockReturnValueOnce( - createModelCostMapMock({ - "gpt-4-personal": { litellm_provider: "openai" }, - "gpt-4-team-only": { litellm_provider: "openai" }, - }), - ); + await user.click(screen.getByTestId("datatable-filters-trigger")); + await user.click(await screen.findByPlaceholderText("Filter by Public Model Name")); + await user.click(await screen.findByRole("option", { name: "gpt-3.5-turbo" })); + await user.click(screen.getByTestId("filter-drawer-apply")); - const modelData = createPaginatedModelData( - [ - { - model_name: "gpt-4-personal", - model_info: { - id: "model-personal-1", - direct_access: true, - access_via_team_ids: [], - access_groups: [], - }, - }, - { - model_name: "gpt-4-team-only", - model_info: { - id: "model-team-1", - direct_access: false, - access_via_team_ids: ["team-123"], - access_groups: [], - }, - }, - ], - 2, - 1, - 1, - 50, - ); - - mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); - - renderWithProviders(); - - // Component shows API total_count (2), but only 1 model has direct_access await waitFor(() => { - expect(screen.getByText("Showing 1 - 2 of 2 results")).toBeInTheDocument(); + expect(mockSetSelectedModelGroup).toHaveBeenCalledWith("gpt-3.5-turbo"); }); }); - it("should show config model status for models defined in configs", async () => { - mockUseTeams.mockReturnValue({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), - }); + it("filters the fetched page down to the selected model group", () => { + setModelsInfo([makeRow(), { ...makeRow(), model_name: "claude-opus" }], 2); + render(); - mockUseModelCostMap.mockReturnValueOnce( - createModelCostMapMock({ - "gpt-4-config": { litellm_provider: "openai" }, - "gpt-4-db": { litellm_provider: "openai" }, - }), - ); + const table = screen.getByRole("table"); + expect(within(table).getByText("claude-opus")).toBeInTheDocument(); + expect(within(table).queryByText("gpt-4")).not.toBeInTheDocument(); + }); - const modelData = createPaginatedModelData( - [ - { - model_name: "gpt-4-config", - litellm_model_name: "gpt-4-config", - provider: "openai", - model_info: { - id: "model-config-1", - db_model: false, - direct_access: true, - access_via_team_ids: [], - access_groups: [], - created_by: "user-123", - created_at: "2024-01-01", - updated_at: "2024-01-01", - }, - }, - { - model_name: "gpt-4-db", - litellm_model_name: "gpt-4-db", - provider: "openai", - model_info: { - id: "model-db-1", - db_model: true, - direct_access: true, - access_via_team_ids: [], - access_groups: [], - created_by: "user-123", - created_at: "2024-01-01", - updated_at: "2024-01-01", - }, - }, - ], - 2, - 1, - 1, - 50, - ); + it("resets search, filters, team and sorting from the drawer reset button", async () => { + const user = userEvent.setup(); + render(); - mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); + await user.click(screen.getByTestId("models-team-select")); + await user.click(await screen.findByRole("option", { name: "Engineering" })); + await waitFor(() => expect(lastModelsInfoCall().teamId).toBe("team-1")); - renderWithProviders(); + await user.click(screen.getByTestId("datatable-filters-trigger")); + await user.click(await screen.findByTestId("filter-drawer-reset")); + expect(mockSetSelectedModelGroup).toHaveBeenCalledWith("all"); await waitFor(() => { - expect(screen.getByText("Config Model")).toBeInTheDocument(); - expect(screen.getByText("DB Model")).toBeInTheDocument(); + expect(lastModelsInfoCall().teamId).toBeUndefined(); }); }); - it("should show 'Defined in config' for models defined in configs", async () => { - mockUseTeams.mockReturnValue({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), - }); + it("opens the delete modal from the row and deletes the model", async () => { + const user = userEvent.setup(); + render(); - mockUseModelCostMap.mockReturnValueOnce( - createModelCostMapMock({ - "gpt-4-config": { litellm_provider: "openai" }, - }), - ); + await user.click(await screen.findByTestId("model-delete-model-1")); + expect(await screen.findByText("Delete Model")).toBeInTheDocument(); - const modelData = createPaginatedModelData( - [ - { - model_name: "gpt-4-config", - litellm_model_name: "gpt-4-config", - provider: "openai", - model_info: { - id: "model-config-1", - db_model: false, - direct_access: true, - access_via_team_ids: [], - access_groups: [], - created_by: "user-123", - created_at: "2024-01-01", - updated_at: "2024-01-01", - }, - }, - ], - 1, - 1, - 1, - 50, - ); - - mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null }); - - renderWithProviders(); + await user.click(screen.getByRole("button", { name: /^delete$/i })); await waitFor(() => { - expect(screen.getByText("Defined in config")).toBeInTheDocument(); + expect(mockModelDeleteCall).toHaveBeenCalledWith("mock-access-token", "model-1"); }); }); - it("should handle pagination: Previous button is disabled on first page and Next button works", async () => { - mockUseTeams.mockReturnValue({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), - }); + it("pauses a model through the row toggle", async () => { + const user = userEvent.setup(); + render(); - mockUseModelCostMap.mockReturnValue( - createModelCostMapMock({ - "gpt-4-page1": { litellm_provider: "openai" }, - "gpt-4-page2": { litellm_provider: "openai" }, - }), - ); - - // Mock first page response (page 1 of 2) - const page1Data = createPaginatedModelData( - [ - { - model_name: "gpt-4-page1", - model_info: { - id: "model-page1-1", - direct_access: true, - access_via_team_ids: [], - access_groups: [], - }, - }, - ], - 2, // total_count - 1, // current_page - 2, // total_pages - 50, // size - ); - - // Set up mock to return page1Data for page 1 - mockUseModelsInfo.mockImplementation((page: number = 1, size?: number, search?: string) => { - return { data: page1Data, isLoading: false, error: null }; - }); - - renderWithProviders(); + await user.click(await screen.findByTestId("model-pause-toggle-model-1")); await waitFor(() => { - // Component calculates: ((1-1)*50)+1 = 1, Math.min(1*50, 2) = 2 - expect(screen.getByText("Showing 1 - 2 of 2 results")).toBeInTheDocument(); + expect(mockModelPatchUpdateCall).toHaveBeenCalledWith("mock-access-token", { blocked: true }, "model-1"); }); - - // Check that Previous button is disabled on first page - const previousButton = screen.getByRole("button", { name: /previous/i }); - expect(previousButton).toBeDisabled(); - - // Check that Next button is enabled (since we're on page 1 of 2) - const nextButton = screen.getByRole("button", { name: /next/i }); - expect(nextButton).not.toBeDisabled(); }); - it("should handle pagination: Next button is disabled on last page", async () => { - mockUseTeams.mockReturnValue({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), - }); + it("opens the model settings modal from the toolbar", async () => { + const user = userEvent.setup(); + render(); - mockUseModelCostMap.mockReturnValue( - createModelCostMapMock({ - "gpt-4-page2": { litellm_provider: "openai" }, - }), - ); - - // Mock single page response (page 1 of 1 - last page) - const singlePageData = createPaginatedModelData( - [ - { - model_name: "gpt-4-page2", - model_info: { - id: "model-page2-1", - direct_access: true, - access_via_team_ids: [], - access_groups: [], - }, - }, - ], - 1, // total_count - 1, // current_page - 1, // total_pages (only 1 page, so this is the last page) - 50, // size - ); - - mockUseModelsInfo.mockImplementation((page?: number, size?: number, search?: string) => { - return { data: singlePageData, isLoading: false, error: null }; - }); - - renderWithProviders(); - - await waitFor(() => { - expect(screen.getByText("Showing 1 - 1 of 1 results")).toBeInTheDocument(); - }); - - // When there's only 1 page (last page), Next should be disabled - const nextButton = screen.getByRole("button", { name: /next/i }); - expect(nextButton).toBeDisabled(); - - // Previous should also be disabled on the first (and only) page - const previousButton = screen.getByRole("button", { name: /previous/i }); - expect(previousButton).toBeDisabled(); + expect(screen.queryByTestId("model-settings-modal")).not.toBeInTheDocument(); + await user.click(screen.getByTestId("models-settings-trigger")); + expect(screen.getByTestId("model-settings-modal")).toBeInTheDocument(); }); - it("should pass setDeleteModalModelId to columns for delete functionality", async () => { - // This test verifies that the delete modal setter is passed to columns - // The actual modal rendering is handled by DeleteResourceModal component - mockUseTeams.mockReturnValue({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), - }); + it("opens the model detail view from the model ID cell", async () => { + const user = userEvent.setup(); + render(); - mockUseModelCostMap.mockReturnValue( - createModelCostMapMock({ - "gpt-4-delete-test": { litellm_provider: "openai" }, - }), - ); + await user.click(await screen.findByTestId("model-id-model-1")); - const modelData = createPaginatedModelData( - [ - { - model_name: "gpt-4-delete-test", - litellm_model_name: "gpt-4-delete-test", - provider: "openai", - model_info: { - id: "model-to-delete", - db_model: true, - direct_access: true, - access_via_team_ids: [], - access_groups: [], - created_by: "user-123", - created_at: "2024-01-01", - updated_at: "2024-01-01", - }, - }, - ], - 1, - 1, - 1, - 50, - ); - - mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null, refetch: vi.fn() }); - - renderWithProviders(); - - await waitFor(() => { - expect(screen.getByText("gpt-4-delete-test")).toBeInTheDocument(); - }); - - // Verify the DB Model badge is shown (indicating it can be deleted) - expect(screen.getByText("DB Model")).toBeInTheDocument(); + expect(mockSetSelectedModelId).toHaveBeenCalledWith("model-1"); }); - it("should render clickable model ID that calls setSelectedModelId", async () => { - mockUseTeams.mockReturnValue({ - data: [], - isLoading: false, - error: null, - refetch: vi.fn(), + it("opens the team detail view from the team ID cell", async () => { + const user = userEvent.setup(); + render(); + + await user.click(await screen.findByTestId("model-team-id-model-1")); + + expect(mockSetSelectedTeamId).toHaveBeenCalledWith("team-1"); + }); + + describe("virtual key hint", () => { + it("explains personal key creation while viewing current team models", () => { + render(); + + expect(screen.getByText(/create a Virtual Key without selecting a team/i)).toBeInTheDocument(); }); - mockUseModelCostMap.mockReturnValue( - createModelCostMapMock({ - "gpt-4-clickable": { litellm_provider: "openai" }, - }), - ); + it("names the selected team in the hint", async () => { + const user = userEvent.setup(); + render(); - const modelData = createPaginatedModelData( - [ - { - model_name: "gpt-4-clickable", - litellm_model_name: "gpt-4-clickable", - provider: "openai", - model_info: { - id: "clickable-model-id", - db_model: true, - direct_access: true, - access_via_team_ids: [], - access_groups: [], - created_by: "user-123", - created_at: "2024-01-01", - updated_at: "2024-01-01", - }, - }, - ], - 1, - 1, - 1, - 50, - ); + await user.click(screen.getByTestId("models-team-select")); + await user.click(await screen.findByRole("option", { name: "Engineering" })); - mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null, refetch: vi.fn() }); - - renderWithProviders(); - - await waitFor(() => { - expect(screen.getByText("gpt-4-clickable")).toBeInTheDocument(); + expect(await screen.findByText(/select Team as "Engineering"/i)).toBeInTheDocument(); }); - // Click on the Model ID cell which should call setSelectedModelId - const modelIdCell = screen.getByText("clickable-model-id"); - expect(modelIdCell).toBeInTheDocument(); + it("hides the hint when viewing all available models", async () => { + const user = userEvent.setup(); + render(); - fireEvent.click(modelIdCell); + await user.click(screen.getByTestId("models-view-select")); + await user.click(await screen.findByRole("option", { name: "All Available Models" })); - await waitFor(() => { - expect(mockSetSelectedModelId).toHaveBeenCalledWith("clickable-model-id"); + await waitFor(() => { + expect(screen.queryByText(/create a Virtual Key/i)).not.toBeInTheDocument(); + }); }); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx index 6efc85c019b..1dc7736d5ac 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx @@ -1,27 +1,33 @@ +"use client"; + import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { Team } from "@/components/key_team_helpers/key_list"; -import { AllModelsDataTable } from "@/components/model_dashboard/all_models_table"; -import { columns } from "@/components/molecules/models/columns"; -import { getDisplayModelName } from "@/components/view_model/model_name_display"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; +import ModelSettingsModal from "@/components/model_dashboard/ModelSettingsModal/ModelSettingsModal"; +import { ModelData } from "@/components/model_dashboard/types"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { modelDeleteCall, modelPatchUpdateCall } from "@/components/networking"; -import { InfoCircleOutlined, SettingOutlined } from "@ant-design/icons"; -import { PaginationState, SortingState } from "@tanstack/react-table"; import { useQueryClient } from "@tanstack/react-query"; -import { Grid } from "@tremor/react"; -import { Badge, Button, Select, Skeleton, Space, Typography } from "antd"; -import ModelSettingsModal from "@/components/model_dashboard/ModelSettingsModal/ModelSettingsModal"; import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer"; -import { useEffect, useMemo, useState } from "react"; +import { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; +import { Info } from "lucide-react"; +import { useCallback, useEffect, useMemo, useState } from "react"; + import { useModelsInfo } from "../../hooks/models/useModels"; import { transformModelData } from "../utils/modelDataTransformer"; -type ModelViewMode = "all" | "current_team"; +import { + ALL_MODEL_GROUPS_VALUE, + AllModelsTable, + ModelViewMode, + PERSONAL_TEAM_VALUE, + WILDCARD_MODEL_GROUP_VALUE, +} from "./AllModelsTable"; +import { ACCESS_GROUPS_COLUMN_ID, MODEL_NAME_COLUMN_ID, toServerSortField } from "./ModelsTableColumns"; const SEARCH_DEBOUNCE_WAIT_MS = 200; -const { Text } = Typography; +const DEFAULT_PAGE_SIZE = 50; +const DEFAULT_PAGINATION: PaginationState = { pageIndex: 0, pageSize: DEFAULT_PAGE_SIZE }; interface AllModelsTabProps { selectedModelGroup: string | null; @@ -41,31 +47,30 @@ const AllModelsTab = ({ setSelectedTeamId, }: AllModelsTabProps) => { const { data: modelCostMapData, isLoading: isLoadingModelCostMap } = useModelCostMap(); - const { accessToken, userId, userRole, premiumUser } = useAuthorized(); + const { accessToken, userId, userRole } = useAuthorized(); const { data: teams, isLoading: isLoadingTeams } = useTeams(); const queryClient = useQueryClient(); const [modelNameSearch, setModelNameSearch] = useState(""); const [debouncedSearch, setDebouncedSearch] = useState(""); const [modelViewMode, setModelViewMode] = useState("current_team"); - const [currentTeam, setCurrentTeam] = useState("personal"); - const [showFilters, setShowFilters] = useState(false); + const [selectedTeamValue, setSelectedTeamValue] = useState(PERSONAL_TEAM_VALUE); const [selectedModelAccessGroupFilter, setSelectedModelAccessGroupFilter] = useState(null); - const [expandedRows, setExpandedRows] = useState>(new Set()); - const [currentPage, setCurrentPage] = useState(1); - const [pageSize] = useState(50); - const [pagination, setPagination] = useState({ - pageIndex: 0, - pageSize: 50, - }); + const [pagination, setPagination] = useState(DEFAULT_PAGINATION); const [sorting, setSorting] = useState([]); const [isModelSettingsModalVisible, setIsModelSettingsModalVisible] = useState(false); + const [deleteModalModelId, setDeleteModalModelId] = useState(null); + const [deleteLoading, setDeleteLoading] = useState(false); + const [pausingModelId, setPausingModelId] = useState(null); + + const resetToFirstPage = useCallback(() => { + setPagination((previous) => (previous.pageIndex === 0 ? previous : { ...previous, pageIndex: 0 })); + }, []); const debouncedUpdateSearch = useDebouncedCallback( (value: string) => { setDebouncedSearch(value); - setCurrentPage(1); - setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 })); + resetToFirstPage(); }, { wait: SEARCH_DEBOUNCE_WAIT_MS }, ); @@ -74,125 +79,130 @@ const AllModelsTab = ({ debouncedUpdateSearch(modelNameSearch); }, [modelNameSearch, debouncedUpdateSearch]); - // Determine teamId to pass to the query - only pass if not "personal" - const teamIdForQuery = currentTeam === "personal" ? undefined : currentTeam.team_id; + const teamIdForQuery = selectedTeamValue === PERSONAL_TEAM_VALUE ? undefined : selectedTeamValue; - // Convert sorting state to sortBy and sortOrder for API const sortBy = useMemo(() => { if (sorting.length === 0) return undefined; - const sort = sorting[0]; - const columnIdToServerField: Record = { - input_cost: "costs", // Map input_cost column to "costs" for server-side sorting - model_info_db_model: "status", // Map model_info.db_model column to "status" for server-side sorting - model_info_created_by: "created_at", // Map model_info.created_by column to "created_at" for server-side sorting - model_info_updated_at: "updated_at", // Map model_info.updated_at column to "updated_at" for server-side sorting - }; - return columnIdToServerField[sort.id] || sort.id; + return toServerSortField(sorting[0].id); }, [sorting]); const sortOrder = useMemo(() => { if (sorting.length === 0) return undefined; - const sort = sorting[0]; - return sort.desc ? "desc" : "asc"; + return sorting[0].desc ? "desc" : "asc"; }, [sorting]); const { data: rawModelData, isLoading: isLoadingModelsInfo, + isFetching: isFetchingModelsInfo, refetch: refetchModels, - } = useModelsInfo(currentPage, pageSize, debouncedSearch || undefined, undefined, teamIdForQuery, sortBy, sortOrder); + } = useModelsInfo( + pagination.pageIndex + 1, + pagination.pageSize, + debouncedSearch || undefined, + undefined, + teamIdForQuery, + sortBy, + sortOrder, + ); const isLoading = isLoadingModelsInfo || isLoadingModelCostMap; - const getProviderFromModel = (model: string) => { - if (modelCostMapData !== null && modelCostMapData !== undefined) { - if (typeof modelCostMapData == "object" && model in modelCostMapData) { - return modelCostMapData[model]["litellm_provider"]; + const getProviderFromModel = useCallback( + (model: string) => { + if (modelCostMapData !== null && modelCostMapData !== undefined) { + if (typeof modelCostMapData == "object" && model in modelCostMapData) { + return modelCostMapData[model]["litellm_provider"]; + } } - } - return "openai"; - }; + return "openai"; + }, + [modelCostMapData], + ); const modelData = useMemo(() => { if (!rawModelData) return { data: [] }; return transformModelData(rawModelData, getProviderFromModel); - }, [rawModelData, modelCostMapData]); + }, [rawModelData, getProviderFromModel]); - const [deleteModalModelId, setDeleteModalModelId] = useState(null); - const [deleteLoading, setDeleteLoading] = useState(false); - - // Get pagination metadata from the response - const paginationMeta = useMemo(() => { - if (!rawModelData) { - return { - total_count: 0, - current_page: 1, - total_pages: 1, - size: pageSize, - }; - } - return { - total_count: rawModelData.total_count ?? 0, - current_page: rawModelData.current_page ?? 1, - total_pages: rawModelData.total_pages ?? 1, - size: rawModelData.size ?? pageSize, - }; - }, [rawModelData, pageSize]); - - const filteredData = useMemo(() => { + const filteredData = useMemo(() => { if (!modelData || !modelData.data || modelData.data.length === 0) { return []; } - // Server-side search is now handled by the API, so we only filter by other criteria - return modelData.data.filter((model: any) => { + return modelData.data.filter((model: ModelData) => { const modelNameMatch = - selectedModelGroup === "all" || + selectedModelGroup === ALL_MODEL_GROUPS_VALUE || model.model_name === selectedModelGroup || !selectedModelGroup || - (selectedModelGroup === "wildcard" && model.model_name?.includes("*")); + (selectedModelGroup === WILDCARD_MODEL_GROUP_VALUE && model.model_name?.includes("*")); const accessGroupMatch = - selectedModelAccessGroupFilter === "all" || - model.model_info["access_groups"]?.includes(selectedModelAccessGroupFilter) || + selectedModelAccessGroupFilter === ALL_MODEL_GROUPS_VALUE || + model.model_info["access_groups"]?.includes(selectedModelAccessGroupFilter ?? "") || !selectedModelAccessGroupFilter; - // Team filtering is now handled server-side via teamId query parameter - // Only apply client-side filtering for model groups and access groups return modelNameMatch && accessGroupMatch; }); }, [modelData, selectedModelGroup, selectedModelAccessGroupFilter]); - useEffect(() => { - setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 })); - setCurrentPage(1); - }, [selectedModelGroup, selectedModelAccessGroupFilter]); + const columnFilters = useMemo( + () => + [ + selectedModelGroup && selectedModelGroup !== ALL_MODEL_GROUPS_VALUE + ? { id: MODEL_NAME_COLUMN_ID, value: selectedModelGroup } + : null, + selectedModelAccessGroupFilter ? { id: ACCESS_GROUPS_COLUMN_ID, value: selectedModelAccessGroupFilter } : null, + ].filter((entry) => entry !== null), + [selectedModelGroup, selectedModelAccessGroupFilter], + ); - // Reset pagination when team changes - useEffect(() => { - setCurrentPage(1); - setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 })); - }, [teamIdForQuery]); + const handleColumnFiltersChange: OnChangeFn = (updater) => { + const next = typeof updater === "function" ? updater(columnFilters) : updater; + const modelGroup = next.find((entry) => entry.id === MODEL_NAME_COLUMN_ID)?.value; + const accessGroup = next.find((entry) => entry.id === ACCESS_GROUPS_COLUMN_ID)?.value; + setSelectedModelGroup(typeof modelGroup === "string" ? modelGroup : ALL_MODEL_GROUPS_VALUE); + setSelectedModelAccessGroupFilter(typeof accessGroup === "string" ? accessGroup : null); + resetToFirstPage(); + }; - // Reset pagination when sorting changes - useEffect(() => { - setCurrentPage(1); - setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 })); - }, [sorting]); + const handleSortingChange: OnChangeFn = (updater) => { + setSorting(typeof updater === "function" ? updater(sorting) : updater); + resetToFirstPage(); + }; + + const handleTeamChange = (value: string) => { + setSelectedTeamValue(value); + resetToFirstPage(); + }; const resetFilters = () => { setModelNameSearch(""); - setSelectedModelGroup("all"); + setSelectedModelGroup(ALL_MODEL_GROUPS_VALUE); setSelectedModelAccessGroupFilter(null); - setCurrentTeam("personal"); + setSelectedTeamValue(PERSONAL_TEAM_VALUE); setModelViewMode("current_team"); - setCurrentPage(1); - setPagination({ pageIndex: 0, pageSize: 50 }); + setPagination(DEFAULT_PAGINATION); setSorting([]); }; + const teamOptions = useMemo( + () => [ + { value: PERSONAL_TEAM_VALUE, label: "Personal" }, + ...(teams ?? []) + .filter((team) => team.team_id) + .map((team) => ({ value: team.team_id, label: team.team_alias ? team.team_alias : team.team_id })), + ], + [teams], + ); + + const selectedTeam = useMemo( + () => (teams ?? []).find((team) => team.team_id === selectedTeamValue) ?? null, + [teams, selectedTeamValue], + ); + const modelToDelete = useMemo(() => { if (!deleteModalModelId || !modelData?.data) return null; - return modelData.data.find((model: any) => model.model_info.id === deleteModalModelId); + return modelData.data.find((model: ModelData) => model.model_info.id === deleteModalModelId); }, [deleteModalModelId, modelData]); const handleDeleteModel = async () => { @@ -212,356 +222,99 @@ const AllModelsTab = ({ } }; - const [pausingModelId, setPausingModelId] = useState(null); + const handleTogglePause = useCallback( + async (modelId: string, blocked: boolean) => { + if (!accessToken) return; + try { + setPausingModelId(modelId); + await modelPatchUpdateCall(accessToken, { blocked }, modelId); + NotificationsManager.success(blocked ? "Model paused" : "Model resumed"); + // invalidateQueries already schedules a refetch for active observers + // on this key — no need to also call refetchModels() (would double-fetch). + queryClient.invalidateQueries({ queryKey: ["models", "list"] }); + } catch (error) { + console.error("Error toggling model pause state:", error); + NotificationsManager.fromBackend(error); + } finally { + setPausingModelId(null); + } + }, + [accessToken, queryClient], + ); - const handleTogglePause = async (modelId: string, blocked: boolean) => { - if (!accessToken) return; - try { - setPausingModelId(modelId); - await modelPatchUpdateCall(accessToken, { blocked }, modelId); - NotificationsManager.success(blocked ? "Model paused" : "Model resumed"); - // invalidateQueries already schedules a refetch for active observers - // on this key — no need to also call refetchModels() (would double-fetch). - queryClient.invalidateQueries({ queryKey: ["models", "list"] }); - } catch (error) { - console.error("Error toggling model pause state:", error); - NotificationsManager.fromBackend(error); - } finally { - setPausingModelId(null); - } - }; + const handleRefresh = useCallback(() => { + void refetchModels(); + }, [refetchModels]); + + const handleDeleteClick = useCallback((modelId: string) => { + setDeleteModalModelId(modelId); + }, []); + + const handleOpenModelSettings = useCallback(() => { + setIsModelSettingsModalVisible(true); + }, []); + + const teamAccessLabel = selectedTeam?.team_alias || selectedTeam?.team_id || ""; return (
- -
-
- {/* Current Team and View Mode Selector - Prominent Section */} -
-
-
- Current Team: -
- {isLoading ? ( - - ) : ( - setModelViewMode(value as "current_team" | "all")} - options={[ - { - value: "current_team", - label: ( - - - Current Team Models - - ), - }, - { - value: "all", - label: ( - - - All Available Models - - ), - }, - ]} - /> - )} -
-
-
+
+ - {modelViewMode === "current_team" && ( -
- -
- {currentTeam === "personal" ? ( - - To access these models: Create a Virtual Key without selecting a team on the{" "} - - Virtual Keys page - - - ) : ( - - To access these models: Create a Virtual Key and select Team as " - {typeof currentTeam !== "string" ? currentTeam.team_alias || currentTeam.team_id : ""}" on - the{" "} - - Virtual Keys page - - - )} -
-
- )} -
- - {/* Search and Filter Controls */} -
-
- {/* Search and Filter Controls */} -
-
- {/* Model Name Search */} -
- setModelNameSearch(e.target.value)} - /> - - - -
- - {/* Filter Button */} - - - {/* Reset Filters Button */} - -
- - {/* Model Settings Button */} -
- - {/* Additional Filters */} - {showFilters && ( -
- {/* Model Name Filter */} -
- setSelectedModelAccessGroupFilter(value === "all" ? null : value)} - placeholder="Filter by Model Access Group" - showSearch - options={[ - { value: "all", label: "All Model Access Groups" }, - ...availableModelAccessGroups.map((accessGroup, idx) => ({ - value: accessGroup, - label: accessGroup, - })), - ]} - /> -
-
- )} - - {/* Results Count and Pagination Controls */} -
- {isLoading ? ( - - ) : ( - - {paginationMeta.total_count > 0 - ? `Showing ${(currentPage - 1) * pageSize + 1} - ${Math.min(currentPage * pageSize, paginationMeta.total_count)} of ${paginationMeta.total_count} results` - : "Showing 0 results"} - - )} - -
- {isLoading ? ( - - ) : ( - - )} - - {isLoading ? ( - - ) : ( - - )} -
-
-
-
- - {}, - () => {}, - expandedRows, - setExpandedRows, - setDeleteModalModelId, - handleTogglePause, - pausingModelId, - )} - data={filteredData} - isLoading={isLoadingModelsInfo} - sorting={sorting} - onSortingChange={setSorting} - pagination={pagination} - onPaginationChange={setPagination} - enablePagination={true} - onRowClick={(model: any) => setSelectedModelId(model.model_info.id)} - /> + {modelViewMode === "current_team" && ( +
+ + {selectedTeamValue === PERSONAL_TEAM_VALUE ? ( + + To access these models, create a Virtual Key without selecting a team on the{" "} + + Virtual Keys page + + . + + ) : ( + + To access these models, create a Virtual Key and select Team as "{teamAccessLabel}" on the{" "} + + Virtual Keys page + + . + + )}
-
- + )} +
({ + default: { success: vi.fn(), fromBackend: vi.fn() }, +})); + +const makeModel = (overrides: Partial = {}): ModelData => + ({ + model_name: "gpt-4-public", + litellm_model_name: "openai/gpt-4", + provider: "openai", + input_cost: 30 as unknown as number, + output_cost: 60 as unknown as number, + max_tokens: 8192, + max_input_tokens: 8192, + litellm_params: { model: "openai/gpt-4" }, + cleanedLitellmParams: {}, + ...overrides, + model_info: { + id: "model-1", + created_at: "2024-01-02T00:00:00Z", + updated_at: "2024-03-04T00:00:00Z", + created_by: "alice", + team_id: "team-1", + db_model: true, + access_groups: null, + ...(overrides.model_info ?? {}), + }, + }) as ModelData; + +const baseProps = { + data: [makeModel()], + rowCount: 1, + isLoading: false, + isRefreshing: false, + onRefresh: vi.fn(), + sorting: [], + onSortingChange: vi.fn(), + pagination: { pageIndex: 0, pageSize: 50 }, + onPaginationChange: vi.fn(), + columnFilters: [], + onColumnFiltersChange: vi.fn(), + onResetFilters: vi.fn(), + searchValue: "", + onSearchChange: vi.fn(), + teamOptions: [ + { value: "personal", label: "Personal" }, + { value: "team-1", label: "Engineering" }, + ], + selectedTeamValue: "personal", + onTeamChange: vi.fn(), + isLoadingTeams: false, + viewMode: "current_team" as const, + onViewModeChange: vi.fn(), + onOpenModelSettings: vi.fn(), + availableModelGroups: ["gpt-4", "gpt-3.5-turbo"], + availableModelAccessGroups: ["sales-team"], + userRole: "Admin", + userID: "alice", + onModelIdClick: vi.fn(), + onTeamIdClick: vi.fn(), + onDeleteClick: vi.fn(), + onTogglePauseClick: vi.fn(), + pausingModelId: null, +}; + +const row = (modelId: string): HTMLElement => { + const element = document.querySelector(`[data-row-id="${modelId}"]`); + if (!(element instanceof HTMLElement)) { + throw new Error(`row ${modelId} not rendered`); + } + return element; +}; + +describe("AllModelsTable", () => { + it("renders the nine design columns and hides Status behind the Columns menu", async () => { + const user = userEvent.setup(); + render(); + + for (const header of [ + "Model ID", + "Model Information", + "Credentials", + "Created By", + "Updated At", + "Costs", + "Team ID", + "Model Access Group", + "Actions", + ]) { + expect(screen.getByRole("columnheader", { name: new RegExp(header, "i") })).toBeInTheDocument(); + } + + expect(screen.queryByRole("columnheader", { name: /^status$/i })).not.toBeInTheDocument(); + expect(screen.queryByText("DB Model")).not.toBeInTheDocument(); + + await user.click(screen.getByRole("button", { name: /columns/i })); + await user.click(await screen.findByRole("menuitemcheckbox", { name: /status/i })); + + expect(await screen.findByText("DB Model")).toBeInTheDocument(); + }); + + it("opens the model detail from the model ID cell", async () => { + const user = userEvent.setup(); + const onModelIdClick = vi.fn(); + render(); + + await user.click(screen.getByTestId("model-id-model-1")); + + expect(onModelIdClick).toHaveBeenCalledWith("model-1"); + }); + + it("opens the team detail from the team ID cell", async () => { + const user = userEvent.setup(); + const onTeamIdClick = vi.fn(); + render(); + + await user.click(screen.getByTestId("model-team-id-model-1")); + + expect(onTeamIdClick).toHaveBeenCalledWith("team-1"); + }); + + it("shows a dash when the model has no team", () => { + render( + , + ); + + expect(within(row("model-1")).getAllByText("-").length).toBeGreaterThan(0); + expect(screen.queryByTestId("model-team-id-model-1")).not.toBeInTheDocument(); + }); + + it("renders the model name over the litellm model name", () => { + render(); + + const cell = screen.getByTestId("model-information-model-1"); + expect(within(cell).getByText("gpt-4-public")).toBeInTheDocument(); + expect(within(cell).getByText("openai/gpt-4")).toBeInTheDocument(); + }); + + it("renders a reusable credential by name and falls back to Manual", () => { + const { rerender } = render( + , + ); + expect(screen.getByText("openai-prod")).toBeInTheDocument(); + expect(screen.queryByText("Manual")).not.toBeInTheDocument(); + + rerender(); + expect(screen.getByText("Manual")).toBeInTheDocument(); + }); + + it("shows 'Defined in config' for a config model and the creator for a DB model", () => { + const { rerender } = render(); + expect(screen.getByText("alice")).toBeInTheDocument(); + + rerender( + , + ); + expect(screen.getByText("Defined in config")).toBeInTheDocument(); + }); + + it("renders input and output costs and a dash when both are missing", () => { + const { rerender } = render(); + expect(screen.getByText("$30")).toBeInTheDocument(); + expect(screen.getByText("$60")).toBeInTheDocument(); + + rerender( + , + ); + expect(screen.queryByText(/^\$/)).not.toBeInTheDocument(); + }); + + it("collapses extra access groups behind a +N more badge", () => { + render( + , + ); + + expect(screen.getByText("sales-team")).toBeInTheDocument(); + expect(screen.getByText("+2 more")).toBeInTheDocument(); + }); + + describe("pause / resume", () => { + it("renders the toggle on for an active DB model and off for a blocked one", () => { + const { rerender } = render(); + expect(screen.getByTestId("model-pause-toggle-model-1")).toBeChecked(); + + rerender( + , + ); + expect(screen.getByTestId("model-pause-toggle-model-1")).not.toBeChecked(); + }); + + it("pauses an active model and resumes a blocked one", async () => { + const user = userEvent.setup(); + const onTogglePauseClick = vi.fn(); + const { rerender } = render(); + + await user.click(screen.getByTestId("model-pause-toggle-model-1")); + expect(onTogglePauseClick).toHaveBeenCalledWith("model-1", true); + + onTogglePauseClick.mockClear(); + rerender( + , + ); + + await user.click(screen.getByTestId("model-pause-toggle-model-1")); + expect(onTogglePauseClick).toHaveBeenCalledWith("model-1", false); + }); + + it("does not let a non-admin toggle a model", async () => { + const user = userEvent.setup(); + const onTogglePauseClick = vi.fn(); + render(); + + const toggle = screen.getByTestId("model-pause-toggle-model-1"); + expect(toggle).toHaveAttribute("data-disabled"); + await user.click(toggle); + expect(onTogglePauseClick).not.toHaveBeenCalled(); + }); + + it("does not let anyone toggle a config model", async () => { + const user = userEvent.setup(); + const onTogglePauseClick = vi.fn(); + render( + , + ); + + const toggle = screen.getByTestId("model-pause-toggle-model-1"); + expect(toggle).toHaveAttribute("data-disabled"); + await user.click(toggle); + expect(onTogglePauseClick).not.toHaveBeenCalled(); + }); + + it("replaces the toggle with a pending indicator while a PATCH is in flight", () => { + render(); + + expect(screen.getByTestId("model-pause-pending-model-1")).toBeInTheDocument(); + expect(screen.queryByTestId("model-pause-toggle-model-1")).not.toBeInTheDocument(); + }); + }); + + describe("delete", () => { + it("lets an admin delete a DB model", async () => { + const user = userEvent.setup(); + const onDeleteClick = vi.fn(); + render(); + + await user.click(screen.getByTestId("model-delete-model-1")); + expect(onDeleteClick).toHaveBeenCalledWith("model-1"); + }); + + it("lets the creator delete their own DB model", async () => { + const user = userEvent.setup(); + const onDeleteClick = vi.fn(); + render(); + + await user.click(screen.getByTestId("model-delete-model-1")); + expect(onDeleteClick).toHaveBeenCalledWith("model-1"); + }); + + it("blocks deleting a model the user did not create", async () => { + const user = userEvent.setup(); + const onDeleteClick = vi.fn(); + render(); + + const deleteButton = screen.getByTestId("model-delete-model-1"); + expect(deleteButton).toBeDisabled(); + await user.click(deleteButton); + expect(onDeleteClick).not.toHaveBeenCalled(); + }); + + it("blocks deleting a config model", async () => { + const user = userEvent.setup(); + const onDeleteClick = vi.fn(); + render( + , + ); + + const deleteButton = screen.getByTestId("model-delete-model-1"); + expect(deleteButton).toBeDisabled(); + await user.click(deleteButton); + expect(onDeleteClick).not.toHaveBeenCalled(); + }); + }); + + describe("toolbar", () => { + it("wires search, refresh, team, view and model settings", async () => { + const user = userEvent.setup(); + const onSearchChange = vi.fn(); + const onRefresh = vi.fn(); + const onOpenModelSettings = vi.fn(); + render( + , + ); + + await user.type(screen.getByTestId("datatable-search"), "gpt"); + expect(onSearchChange).toHaveBeenCalled(); + + await user.click(screen.getByTestId("datatable-refresh")); + expect(onRefresh).toHaveBeenCalled(); + + await user.click(screen.getByTestId("models-settings-trigger")); + expect(onOpenModelSettings).toHaveBeenCalled(); + + expect(screen.getByTestId("models-team-select")).toHaveTextContent("Personal"); + expect(screen.getByTestId("models-view-select")).toHaveTextContent("Current Team Models"); + }); + + it("switches the current team", async () => { + const user = userEvent.setup(); + const onTeamChange = vi.fn(); + render(); + + await user.click(screen.getByTestId("models-team-select")); + await user.click(await screen.findByRole("option", { name: "Engineering" })); + + expect(onTeamChange).toHaveBeenCalledWith("team-1"); + }); + + it("runs the full reset from the filter drawer", async () => { + const user = userEvent.setup(); + const onResetFilters = vi.fn(); + render(); + + await user.click(screen.getByTestId("datatable-filters-trigger")); + await user.click(await screen.findByTestId("filter-drawer-reset")); + + expect(onResetFilters).toHaveBeenCalled(); + }); + + it("renders active filters as removable chips", async () => { + const user = userEvent.setup(); + const onColumnFiltersChange = vi.fn(); + render( + , + ); + + const chip = screen.getByTestId("filter-chip-model_name"); + expect(chip).toHaveTextContent("Public Model Name"); + expect(chip).toHaveTextContent("Wildcard Models (*)"); + + await user.click(screen.getByTestId("filter-chip-remove-model_name")); + expect(onColumnFiltersChange).toHaveBeenCalled(); + }); + }); + + it("shows the server row count in the pagination footer", () => { + render(); + + expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-50 of 137"); + }); + + it("shows the empty state when there are no models", () => { + render(); + + expect(screen.getByText("No models found")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.tsx new file mode 100644 index 00000000000..d073519d162 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.tsx @@ -0,0 +1,298 @@ +"use client"; + +import { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; +import { Search, Settings } from "lucide-react"; +import { useMemo, useState } from "react"; + +import { ModelData } from "@/components/model_dashboard/types"; +import { + DataTable, + DataTableFilterDrawer, + DataTableFilterField, + DataTableToolbar, +} from "@/components/shared/DataTable"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Button } from "@/components/ui/button"; +import { Select, SelectContent, SelectItem, SelectTrigger } from "@/components/ui/select"; +import { Separator } from "@/components/ui/separator"; +import { cn } from "@/lib/cva.config"; + +import { + ACCESS_GROUPS_COLUMN_ID, + getModelsTableColumns, + MODEL_NAME_COLUMN_ID, + STATUS_COLUMN_ID, +} from "./ModelsTableColumns"; + +export type ModelViewMode = "all" | "current_team"; + +export const PERSONAL_TEAM_VALUE = "personal"; +export const ALL_MODEL_GROUPS_VALUE = "all"; +export const WILDCARD_MODEL_GROUP_VALUE = "wildcard"; + +const MODEL_TABLE_BODY_HEIGHT = 600; + +const FILTER_LABELS: Record = { + [MODEL_NAME_COLUMN_ID]: "Public Model Name", + [ACCESS_GROUPS_COLUMN_ID]: "Model Access Group", +}; + +const VIEW_MODE_LABELS: Record = { + current_team: "Current Team Models", + all: "All Available Models", +}; + +export interface ModelsTableTeamOption { + value: string; + label: string; +} + +interface AllModelsTableProps { + data: ModelData[]; + rowCount: number; + isLoading: boolean; + isRefreshing: boolean; + onRefresh: () => void; + sorting: SortingState; + onSortingChange: OnChangeFn; + pagination: PaginationState; + onPaginationChange: OnChangeFn; + columnFilters: ColumnFiltersState; + onColumnFiltersChange: OnChangeFn; + onResetFilters: () => void; + searchValue: string; + onSearchChange: (value: string) => void; + teamOptions: ModelsTableTeamOption[]; + selectedTeamValue: string; + onTeamChange: (value: string) => void; + isLoadingTeams: boolean; + viewMode: ModelViewMode; + onViewModeChange: (viewMode: ModelViewMode) => void; + onOpenModelSettings: () => void; + availableModelGroups: string[]; + availableModelAccessGroups: string[]; + userRole: string; + userID: string; + onModelIdClick: (modelId: string) => void; + onTeamIdClick: (teamId: string) => void; + onDeleteClick: (modelId: string) => void; + onTogglePauseClick: (modelId: string, blocked: boolean) => void | Promise; + pausingModelId: string | null; +} + +function EmptyState() { + return ( +
+
+ +
+
No models found
+
+ No models match your search or filters. Try resetting them. +
+
+ ); +} + +export function AllModelsTable({ + data, + rowCount, + isLoading, + isRefreshing, + onRefresh, + sorting, + onSortingChange, + pagination, + onPaginationChange, + columnFilters, + onColumnFiltersChange, + onResetFilters, + searchValue, + onSearchChange, + teamOptions, + selectedTeamValue, + onTeamChange, + isLoadingTeams, + viewMode, + onViewModeChange, + onOpenModelSettings, + availableModelGroups, + availableModelAccessGroups, + userRole, + userID, + onModelIdClick, + onTeamIdClick, + onDeleteClick, + onTogglePauseClick, + pausingModelId, +}: AllModelsTableProps) { + const [filtersOpen, setFiltersOpen] = useState(false); + + const columns = useMemo(() => { + const columnDeps = { + userRole, + userID, + onModelIdClick, + onTeamIdClick, + onDeleteClick, + onTogglePauseClick, + pausingModelId, + }; + return getModelsTableColumns(columnDeps); + }, [userRole, userID, onModelIdClick, onTeamIdClick, onDeleteClick, onTogglePauseClick, pausingModelId]); + + const modelGroupOptions = useMemo( + () => [ + { label: "All Models", value: ALL_MODEL_GROUPS_VALUE }, + { label: "Wildcard Models (*)", value: WILDCARD_MODEL_GROUP_VALUE }, + ...availableModelGroups.map((group) => ({ label: group, value: group })), + ], + [availableModelGroups], + ); + + const accessGroupOptions = useMemo( + () => [ + { label: "All Model Access Groups", value: ALL_MODEL_GROUPS_VALUE }, + ...availableModelAccessGroups.map((accessGroup) => ({ label: accessGroup, value: accessGroup })), + ], + [availableModelAccessGroups], + ); + + const formatFilterValue = (columnId: string, value: unknown): string => { + const raw = String(value); + if (columnId === MODEL_NAME_COLUMN_ID && raw === WILDCARD_MODEL_GROUP_VALUE) { + return "Wildcard Models (*)"; + } + return raw; + }; + + const selectedTeamLabel = + teamOptions.find((option) => option.value === selectedTeamValue)?.label ?? teamOptions[0]?.label ?? ""; + + return ( + row.model_info?.id ?? String(index)} + sortingMode="server" + sorting={sorting} + onSortingChange={onSortingChange} + enableSortingRemoval + paginationMode="server" + pagination={pagination} + onPaginationChange={onPaginationChange} + rowCount={rowCount} + pageSizeOptions={[10, 25, 50]} + filterMode="server" + columnFilters={columnFilters} + onColumnFiltersChange={onColumnFiltersChange} + defaultColumnVisibility={{ [STATUS_COLUMN_ID]: false }} + enableColumnResizing + maxBodyHeight={MODEL_TABLE_BODY_HEIGHT} + isLoading={isLoading} + loadingMessage="Loading models…" + noDataMessage={} + size="compact" + toolbar={(table) => ( + <> + setFiltersOpen(true)} + onRefresh={onRefresh} + isRefreshing={isRefreshing} + filterLabels={FILTER_LABELS} + formatFilterValue={formatFilterValue} + > + + + + + + + + + + {({ get, set }) => ( + <> + + + set(MODEL_NAME_COLUMN_ID, value === ALL_MODEL_GROUPS_VALUE ? undefined : value) + } + placeholder="Filter by Public Model Name" + emptyText="No models found" + /> + + + + set(ACCESS_GROUPS_COLUMN_ID, value === ALL_MODEL_GROUPS_VALUE ? undefined : value) + } + placeholder="Filter by Model Access Group" + emptyText="No model access groups found" + /> + + + )} + + + )} + /> + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx new file mode 100644 index 00000000000..93ee6d9f0ab --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx @@ -0,0 +1,488 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { Copy, Info, Loader2, Pencil, RefreshCw, Trash2 } from "lucide-react"; + +import { ProviderLogo } from "@/components/molecules/models/ProviderLogo"; +import { ModelData } from "@/components/model_dashboard/types"; +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { CellTooltip, DateCell, formatCellDate, IdCell, StatusBadge } from "@/components/shared/table_cells"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { HoverCard, HoverCardContent, HoverCardTrigger } from "@/components/ui/hover-card"; +import { Switch } from "@/components/ui/switch"; +import { getDisplayModelName } from "@/components/view_model/model_name_display"; +import { copyToClipboard } from "@/utils/dataUtils"; + +export const MODEL_ID_COLUMN_ID = "model_info_id"; +export const MODEL_NAME_COLUMN_ID = "model_name"; +export const CREDENTIALS_COLUMN_ID = "litellm_credential_name"; +export const CREATED_BY_COLUMN_ID = "model_info_created_by"; +export const UPDATED_AT_COLUMN_ID = "model_info_updated_at"; +export const COSTS_COLUMN_ID = "input_cost"; +export const TEAM_ID_COLUMN_ID = "model_info_team_id"; +export const ACCESS_GROUPS_COLUMN_ID = "model_info_access_groups"; +export const STATUS_COLUMN_ID = "model_info_db_model"; + +const COLUMN_ID_TO_SERVER_SORT_FIELD: Record = { + [COSTS_COLUMN_ID]: "costs", + [STATUS_COLUMN_ID]: "status", + [CREATED_BY_COLUMN_ID]: "created_at", + [UPDATED_AT_COLUMN_ID]: "updated_at", +}; + +export const toServerSortField = (columnId: string): string => COLUMN_ID_TO_SERVER_SORT_FIELD[columnId] ?? columnId; + +const formatShortDate = (value: string | null | undefined): string | null => { + if (!value) { + return null; + } + const date = new Date(value); + return Number.isNaN(date.getTime()) ? null : formatCellDate(date, "date"); +}; + +function ModelInformationCell({ model, displayName }: { model: ModelData; displayName: string }) { + const litellmModelName = model.litellm_model_name || "-"; + + return ( + + + } + > + {model.provider ? ( + + ) : ( + + - + + )} + + + {displayName} + + + {litellmModelName} + + + + +
+
+ {model.provider ? : null} + {model.provider || "Unknown provider"} +
+
+ Public Model Name + + {displayName} + +
+
+ LiteLLM Model Name + + + {litellmModelName} + + + +
+
+
+
+ ); +} + +function CredentialsHeader() { + return ( + + Credentials + + + } + > + + + +
+ Credential types +
+ + + Reusable + + + Credentials saved in LiteLLM that can be added to models repeatedly. + +
+
+ + + Manual + + + Credentials added directly during model creation or defined in the config file. + +
+
+
+
+
+ ); +} + +function CredentialsCell({ credentialName }: { credentialName: string | undefined }) { + if (!credentialName) { + return ( + + + Manual + + ); + } + + return ( + + + {credentialName} + + ); +} + +function CreatedByCell({ model }: { model: ModelData }) { + const isConfigModel = !model.model_info?.db_model; + const createdAt = formatShortDate(model.model_info.created_at); + const primary = isConfigModel ? "Defined in config" : model.model_info.created_by || "Unknown"; + const secondaryForDbModel = createdAt ?? "Unknown date"; + + return ( +
+ + {primary} + + {isConfigModel ? "-" : secondaryForDbModel} +
+ ); +} + +function CostsCell({ model }: { model: ModelData }) { + const { input_cost: inputCost, output_cost: outputCost } = model; + + if (inputCost == null && outputCost == null) { + return -; + } + + return ( + + {inputCost != null && ( + + IN + ${inputCost} + + )} + {outputCost != null && ( + + OUT + ${outputCost} + + )} +
+ } + /> + ); +} + +function AccessGroupsCell({ accessGroups }: { accessGroups: string[] | null }) { + if (!accessGroups || accessGroups.length === 0) { + return -; + } + + const [first, ...overflow] = accessGroups; + + return ( +
+ + {first} + + {overflow.length > 0 && ( + + {overflow.map((group) => ( + {group} + ))} +
+ } + trigger={ + + +{overflow.length} more + + } + /> + )} +
+ ); +} + +interface ModelRowActionsProps { + model: ModelData; + userRole: string; + userID: string; + isPausing: boolean; + onDeleteClick?: (modelId: string) => void; + onTogglePauseClick?: (modelId: string, blocked: boolean) => void | Promise; +} + +function ModelRowActions({ + model, + userRole, + userID, + isPausing, + onDeleteClick, + onTogglePauseClick, +}: ModelRowActionsProps) { + const modelId = model.model_info?.id; + const isConfigModel = !model.model_info?.db_model; + const isAdmin = userRole === "Admin"; + const canEditModel = isAdmin || model.model_info?.created_by === userID; + const isBlocked = model.model_info?.blocked === true; + const isPauseToggleable = !isConfigModel && isAdmin && Boolean(onTogglePauseClick); + + const resolvePauseTooltip = (): string => { + if (isConfigModel) { + return "Config models cannot be paused from the dashboard. Pause is DB-backed."; + } + if (!isAdmin) { + return "Only proxy admins can pause or resume a model."; + } + return isBlocked ? "Resume model — restore normal routing." : "Pause model — stop routing requests until resumed."; + }; + + const deleteTooltip = isConfigModel + ? "Config model cannot be deleted on the dashboard. Please delete it from the config file." + : "Delete model"; + + return ( +
+ + {isPausing ? ( + + ) : ( + + { + if (isPauseToggleable && onTogglePauseClick && modelId) { + void onTogglePauseClick(modelId, !nextChecked); + } + }} + /> + + } + /> + )} + + + + + } + /> +
+ ); +} + +export interface ModelsTableColumnDeps { + userRole: string; + userID: string; + onModelIdClick: (modelId: string) => void; + onTeamIdClick: (teamId: string) => void; + onDeleteClick?: (modelId: string) => void; + onTogglePauseClick?: (modelId: string, blocked: boolean) => void | Promise; + pausingModelId?: string | null; +} + +export const getModelsTableColumns = ({ + userRole, + userID, + onModelIdClick, + onTeamIdClick, + onDeleteClick, + onTogglePauseClick, + pausingModelId, +}: ModelsTableColumnDeps): ColumnDef[] => [ + { + id: MODEL_ID_COLUMN_ID, + accessorFn: (row) => row.model_info.id, + meta: { title: "Model ID" }, + header: "Model ID", + enableSorting: false, + size: 140, + minSize: 90, + cell: ({ row }) => ( + + ), + }, + { + id: MODEL_NAME_COLUMN_ID, + accessorFn: (row) => row.model_name ?? "", + meta: { title: "Model Information", skeleton: "twoLine" }, + header: ({ column }) => , + enableSorting: true, + size: 280, + minSize: 160, + cell: ({ row }) => ( + + ), + }, + { + id: CREDENTIALS_COLUMN_ID, + accessorFn: (row) => row.litellm_params?.litellm_credential_name ?? "", + meta: { title: "Credentials" }, + header: () => , + enableSorting: false, + size: 180, + minSize: 110, + cell: ({ row }) => , + }, + { + id: CREATED_BY_COLUMN_ID, + accessorFn: (row) => row.model_info.created_by ?? "", + meta: { title: "Created By", skeleton: "twoLine" }, + header: ({ column }) => , + enableSorting: true, + size: 180, + minSize: 110, + cell: ({ row }) => , + }, + { + id: UPDATED_AT_COLUMN_ID, + accessorFn: (row) => row.model_info.updated_at ?? "", + meta: { title: "Updated At" }, + header: ({ column }) => , + enableSorting: true, + size: 140, + minSize: 100, + cell: ({ row }) => , + }, + { + id: COSTS_COLUMN_ID, + accessorFn: (row) => row.input_cost, + meta: { title: "Costs" }, + header: ({ column }) => , + enableSorting: true, + size: 130, + minSize: 90, + cell: ({ row }) => , + }, + { + id: TEAM_ID_COLUMN_ID, + accessorFn: (row) => row.model_info.team_id ?? "", + meta: { title: "Team ID" }, + header: "Team ID", + enableSorting: false, + size: 140, + minSize: 90, + cell: ({ row }) => ( + + ), + }, + { + id: ACCESS_GROUPS_COLUMN_ID, + accessorFn: (row) => row.model_info.access_groups ?? [], + meta: { title: "Model Access Group", skeleton: "chips" }, + header: "Model Access Group", + enableSorting: false, + size: 200, + minSize: 120, + cell: ({ row }) => , + }, + { + id: STATUS_COLUMN_ID, + accessorFn: (row) => row.model_info.db_model, + meta: { title: "Status", skeleton: "badge" }, + header: ({ column }) => , + enableSorting: true, + size: 140, + minSize: 100, + cell: ({ row }) => + row.original.model_info.db_model ? ( + + ) : ( + + ), + }, + { + id: "actions", + meta: { title: "Actions", className: "text-right", headerClassName: "text-right" }, + header: "Actions", + enableSorting: false, + enableHiding: false, + enableResizing: false, + size: 110, + minSize: 110, + cell: ({ row }) => ( + + ), + }, +]; diff --git a/ui/litellm-dashboard/src/components/model_dashboard/all_models_table.tsx b/ui/litellm-dashboard/src/components/model_dashboard/all_models_table.tsx deleted file mode 100644 index 4372c6e0efc..00000000000 --- a/ui/litellm-dashboard/src/components/model_dashboard/all_models_table.tsx +++ /dev/null @@ -1,221 +0,0 @@ -import { - ColumnDef, - flexRender, - getCoreRowModel, - getPaginationRowModel, - SortingState, - useReactTable, - ColumnResizeMode, - VisibilityState, - PaginationState, - OnChangeFn, -} from "@tanstack/react-table"; -import React from "react"; -import { Table, TableHead, TableHeaderCell, TableBody, TableRow, TableCell } from "@tremor/react"; -import { - TableHeaderSortDropdown, - SortState, -} from "../common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; - -// Extend the column meta type to include className -declare module "@tanstack/react-table" { - interface ColumnMeta { - className?: string; - } -} - -interface AllModelsDataTableProps { - data: TData[]; - columns: ColumnDef[]; - isLoading?: boolean; - sorting?: SortingState; - onSortingChange?: OnChangeFn; - pagination?: PaginationState; - onPaginationChange?: OnChangeFn; - enablePagination?: boolean; - onRowClick?: (row: TData) => void; -} - -export function AllModelsDataTable({ - data = [], - columns, - isLoading = false, - sorting = [], - onSortingChange, - pagination, - onPaginationChange, - enablePagination = false, - onRowClick, -}: AllModelsDataTableProps) { - const [columnResizeMode] = React.useState("onChange"); - const [columnSizing, setColumnSizing] = React.useState({}); - const [columnVisibility, setColumnVisibility] = React.useState({}); - - const tableInstance = useReactTable({ - data, - columns, - state: { - sorting, - columnSizing, - columnVisibility, - ...(enablePagination && pagination ? { pagination } : {}), - }, - columnResizeMode, - onSortingChange: onSortingChange, - onColumnSizingChange: setColumnSizing, - onColumnVisibilityChange: setColumnVisibility, - ...(enablePagination && onPaginationChange ? { onPaginationChange } : {}), - getCoreRowModel: getCoreRowModel(), - // NO getSortedRowModel - sorting is handled server-side - ...(enablePagination ? { getPaginationRowModel: getPaginationRowModel() } : {}), - enableSorting: true, - enableColumnResizing: true, - manualSorting: true, // Enable manual sorting for server-side sorting - defaultColumn: { - minSize: 40, - maxSize: 500, - }, - }); - - const getHeaderText = (header: any): string => { - if (typeof header === "string") { - return header; - } - if (typeof header === "function") { - const headerElement = header(); - if (headerElement && headerElement.props && headerElement.props.children) { - const children = headerElement.props.children; - if (typeof children === "string") { - return children; - } - if (children.props && children.props.children) { - return children.props.children; - } - } - } - return ""; - }; - - return ( -
-
-
- - - {tableInstance.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - -
-
- {header.isPlaceholder - ? null - : flexRender(header.column.columnDef.header, header.getContext())} -
- {header.id !== "actions" && header.column.getCanSort() && onSortingChange && ( - { - // Convert SortState to TanStack SortingState - // Only allow one column to be sorted at a time - if (newState === false) { - onSortingChange([]); - } else { - onSortingChange([ - { - id: header.column.id, - desc: newState === "desc", - }, - ]); - } - }} - columnId={header.column.id} - /> - )} -
- {header.column.getCanResize() && ( -
- )} - - ))} - - ))} - - - {isLoading ? ( - - -
-

🚅 Loading models...

-
-
-
- ) : tableInstance.getRowModel().rows.length > 0 ? ( - tableInstance.getRowModel().rows.map((row) => ( - onRowClick?.(row.original)} - > - {row.getVisibleCells().map((cell) => ( - - {flexRender(cell.column.columnDef.cell, cell.getContext())} - - ))} - - )) - ) : ( - - -
-

No models found

-
-
-
- )} -
-
-
-
-
- ); -} diff --git a/ui/litellm-dashboard/src/components/model_dashboard/types.ts b/ui/litellm-dashboard/src/components/model_dashboard/types.ts index 77a03d2c039..e58204995dd 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard/types.ts +++ b/ui/litellm-dashboard/src/components/model_dashboard/types.ts @@ -7,6 +7,7 @@ export interface ModelInfo { db_model: boolean; access_groups: string[] | null; blocked?: boolean; + team_public_model_name?: string; } export interface LiteLLMParams { diff --git a/ui/litellm-dashboard/src/components/molecules/models/columns.test.tsx b/ui/litellm-dashboard/src/components/molecules/models/columns.test.tsx deleted file mode 100644 index 9628c706c83..00000000000 --- a/ui/litellm-dashboard/src/components/molecules/models/columns.test.tsx +++ /dev/null @@ -1,1099 +0,0 @@ -import { render, screen } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { describe, expect, it, vi, beforeEach } from "vitest"; -import { useReactTable, getCoreRowModel, flexRender } from "@tanstack/react-table"; -import { columns } from "./columns"; -import { ModelData } from "../../model_dashboard/types"; -import { Table, TableHead, TableHeaderCell, TableBody, TableRow, TableCell } from "@tremor/react"; -import * as providerInfoHelpers from "../../provider_info_helpers"; - -vi.mock("../../provider_info_helpers"); - -vi.mock("@tremor/react", async (importOriginal) => { - const React = await import("react"); - const actual = await importOriginal(); - const IconComponent = React.forwardRef( - ({ icon: IconComp, onClick, className, ...props }, ref) => { - const ariaLabel = className?.includes("cursor-not-allowed") - ? "Config model cannot be deleted on the dashboard. Please delete it from the config file." - : "Delete model"; - return React.createElement( - "button", - { ...props, onClick, className, ref, "aria-label": ariaLabel }, - IconComp && React.createElement(IconComp, { className: "w-4 h-4" }), - ); - }, - ); - IconComponent.displayName = "Icon"; - // Re-apply the global Button/Tooltip overrides from tests/setupTests.ts. A file-level - // vi.mock fully replaces the setup-level mock, so without this the real Tremor Button - // leaks through and its useTooltip(300) schedules a native setTimeout that can fire - // post-teardown -> "window is not defined". - const Button = React.forwardRef(({ children, ...props }, ref) => - React.createElement("button", { ...props, ref }, children), - ); - const Tooltip = ({ children }: any) => React.createElement(React.Fragment, null, children); - return { - ...actual, - Icon: IconComponent, - Button, - Tooltip, - }; -}); - -const createMockModel = (overrides: Partial = {}): ModelData => ({ - model_info: { - id: "test-model-id", - created_at: "2024-01-01T00:00:00Z", - updated_at: "2024-01-02T00:00:00Z", - created_by: "test-user", - team_id: "test-team-id", - db_model: true, - access_groups: ["group1"], - }, - model_name: "test-model", - provider: "openai", - litellm_model_name: "gpt-4", - input_cost: 0.01, - output_cost: 0.03, - max_tokens: 4096, - max_input_tokens: 8192, - litellm_params: { - model: "gpt-4", - litellm_credential_name: "test-credential", - }, - cleanedLitellmParams: {}, - ...overrides, -}); - -const TestTable = ({ data, columns: cols }: { data: ModelData[]; columns: ReturnType }) => { - const table = useReactTable({ - data, - columns: cols, - getCoreRowModel: getCoreRowModel(), - }); - - return ( - - - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - - {header.isPlaceholder ? null : flexRender(header.column.columnDef.header, header.getContext())} - - ))} - - ))} - - - {table.getRowModel().rows.map((row) => ( - - {row.getVisibleCells().map((cell) => ( - {flexRender(cell.column.columnDef.cell, cell.getContext())} - ))} - - ))} - -
- ); -}; - -describe("columns", () => { - beforeEach(() => { - vi.mocked(providerInfoHelpers.getProviderLogoAndName).mockImplementation((provider: string) => { - const providerMap: Record = { - openai: { displayName: "OpenAI", logo: "/openai-logo.png" }, - anthropic: { displayName: "Anthropic", logo: "/anthropic-logo.png" }, - }; - return providerMap[provider] || { displayName: provider || "Unknown provider", logo: "" }; - }); - }); - - const defaultProps = { - userRole: "Admin", - userID: "test-user", - premiumUser: false, - setSelectedModelId: vi.fn(), - setSelectedTeamId: vi.fn(), - getDisplayModelName: vi.fn((model: ModelData) => model.model_name || "-"), - handleEditClick: vi.fn(), - handleRefreshClick: vi.fn(), - expandedRows: new Set(), - setExpandedRows: vi.fn(), - }; - - it("should render columns with table structure", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel(); - render(); - - expect(screen.getByText("Model ID")).toBeInTheDocument(); - expect(screen.getByText("Model Information")).toBeInTheDocument(); - expect(screen.getByText("Credentials")).toBeInTheDocument(); - expect(screen.getByText("Created By")).toBeInTheDocument(); - expect(screen.getByText("Updated At")).toBeInTheDocument(); - expect(screen.getByText("Costs")).toBeInTheDocument(); - expect(screen.getByText("Team ID")).toBeInTheDocument(); - expect(screen.getByText("Model Access Group")).toBeInTheDocument(); - expect(screen.getByText("Status")).toBeInTheDocument(); - expect(screen.getByText("Actions")).toBeInTheDocument(); - }); - - it("should display model information with provider logo", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_name: "GPT-4", - provider: "openai", - litellm_model_name: "gpt-4", - }); - render(); - - expect(screen.getByText("GPT-4")).toBeInTheDocument(); - expect(screen.getByText("gpt-4")).toBeInTheDocument(); - }); - - it("should display credential name when available", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - litellm_params: { - model: "gpt-4", - litellm_credential_name: "my-credential", - }, - }); - render(); - - expect(screen.getByText("my-credential")).toBeInTheDocument(); - }); - - it("should display 'Manual' when credential name is missing", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - litellm_params: { - model: "gpt-4", - }, - }); - render(); - - expect(screen.getByText("Manual")).toBeInTheDocument(); - }); - - describe("credentials column", () => { - it("should display Credentials header with info icon", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel(); - render(); - - expect(screen.getByText("Credentials")).toBeInTheDocument(); - // Info icon is in a flex container with Credentials - ant icons render as span with role="img" - const credentialsHeader = screen.getByText("Credentials").closest("span"); - expect(credentialsHeader?.parentElement?.querySelector('[role="img"]')).toBeInTheDocument(); - }); - - it("should display reusable credential with SyncOutlined icon and credential name", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - litellm_params: { - model: "gpt-4", - litellm_credential_name: "my-reusable-credential", - }, - }); - render(); - - expect(screen.getByText("my-reusable-credential")).toBeInTheDocument(); - const credentialCell = screen.getByText("my-reusable-credential").closest("div"); - expect(credentialCell).toHaveClass("flex"); - expect(screen.getByText("my-reusable-credential")).toHaveClass("text-blue-600"); - }); - - it("should display Manual with EditOutlined when no credential name", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - litellm_params: { - model: "gpt-4", - }, - }); - render(); - - expect(screen.getByText("Manual")).toBeInTheDocument(); - expect(screen.getByText("Manual")).toHaveClass("text-gray-500"); - }); - - it("should display Manual when litellm_params is undefined", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - litellm_params: undefined as any, - }); - render(); - - expect(screen.getByText("Manual")).toBeInTheDocument(); - }); - - it("should display Manual when litellm_credential_name is empty string", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - litellm_params: { - model: "gpt-4", - litellm_credential_name: "", - }, - }); - render(); - - expect(screen.getByText("Manual")).toBeInTheDocument(); - }); - }); - - it("should display created by information for DB models", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: true, - created_by: "admin-user", - created_at: "2024-01-15T10:30:00Z", - }, - }); - render(); - - expect(screen.getByText("admin-user")).toBeInTheDocument(); - }); - - it("should display 'Defined in config' for config models", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: false, - }, - }); - render(); - - expect(screen.getByText("Defined in config")).toBeInTheDocument(); - }); - - it("should display costs when available", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - input_cost: 0.01, - output_cost: 0.03, - }); - render(); - - expect(screen.getByText("In: $0.01")).toBeInTheDocument(); - expect(screen.getByText("Out: $0.03")).toBeInTheDocument(); - }); - - it("should display '-' when costs are missing", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - input_cost: undefined as any, - output_cost: undefined as any, - }); - render(); - - const costCells = screen.getAllByText("-"); - expect(costCells.length).toBeGreaterThan(0); - }); - - it("should call setSelectedModelId without triggering the row click when the model ID pill is clicked", async () => { - const user = userEvent.setup(); - const setSelectedModelId = vi.fn(); - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel(); - const rowClick = vi.fn(); - render( -
- -
, - ); - - await user.click(screen.getByText("test-model-id")); - expect(setSelectedModelId).toHaveBeenCalledWith("test-model-id"); - expect(rowClick).not.toHaveBeenCalled(); - }); - - it("should call setSelectedTeamId without triggering the row click when the team ID pill is clicked", async () => { - const user = userEvent.setup(); - const setSelectedTeamId = vi.fn(); - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel(); - const rowClick = vi.fn(); - render( -
- -
, - ); - - await user.click(screen.getByText("test-team-id")); - expect(setSelectedTeamId).toHaveBeenCalledWith("test-team-id"); - expect(rowClick).not.toHaveBeenCalled(); - }); - - it("should display '-' when team ID is missing", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - team_id: "", - }, - }); - render(); - - const teamIdCells = screen.getAllByText("-"); - expect(teamIdCells.length).toBeGreaterThan(0); - }); - - it("should display access groups", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - access_groups: ["group1", "group2"], - }, - }); - render(); - - expect(screen.getByText("group1")).toBeInTheDocument(); - expect(screen.getByText("+1")).toBeInTheDocument(); - }); - - it("should expand access groups when expand button is clicked", async () => { - const user = userEvent.setup(); - const setExpandedRows = vi.fn(); - const expandedRows = new Set(); - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - expandedRows, - setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - id: "model-with-groups", - access_groups: ["group1", "group2", "group3"], - }, - }); - render(); - - const expandButton = screen.getByText("+2"); - expect(expandButton).toBeInTheDocument(); - - await user.click(expandButton); - expect(setExpandedRows).toHaveBeenCalled(); - }); - - it("should display '-' when access groups are empty", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - access_groups: null, - }, - }); - render(); - - const emptyCells = screen.getAllByText("-"); - expect(emptyCells.length).toBeGreaterThan(0); - }); - - it("should display 'DB Model' status for DB models", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: true, - }, - }); - render(); - - expect(screen.getByText("DB Model")).toBeInTheDocument(); - }); - - it("should display 'Config Model' status for config models", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: false, - }, - }); - render(); - - expect(screen.getByText("Config Model")).toBeInTheDocument(); - }); - - it("should allow Admin to delete DB models", async () => { - const user = userEvent.setup(); - const onDeleteClick = vi.fn(); - const cols = columns( - "Admin", - "admin-user", - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - onDeleteClick, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: true, - id: "deletable-model", - }, - }); - render(); - - const deleteButton = screen.getByRole("button", { name: "Delete model" }); - expect(deleteButton).toBeInTheDocument(); - - await user.click(deleteButton); - expect(onDeleteClick).toHaveBeenCalledWith("deletable-model"); - }); - - it("should allow model creator to delete their own DB models", async () => { - const user = userEvent.setup(); - const onDeleteClick = vi.fn(); - const cols = columns( - "User", - "model-creator", - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - onDeleteClick, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: true, - created_by: "model-creator", - id: "user-model", - }, - }); - render(); - - const deleteButton = screen.getByRole("button", { name: "Delete model" }); - expect(deleteButton).toBeInTheDocument(); - - await user.click(deleteButton); - expect(onDeleteClick).toHaveBeenCalledWith("user-model"); - }); - - it("should disable delete for config models", () => { - const cols = columns( - "Admin", - "admin-user", - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: false, - }, - }); - render(); - - const deleteButton = screen.getByRole("button", { name: /config model cannot be deleted/i }); - expect(deleteButton).toBeInTheDocument(); - expect(deleteButton).toHaveClass("cursor-not-allowed"); - }); - - it("should display collapsed access groups with expand button", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - new Set(), - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - access_groups: ["group1", "group2", "group3"], - }, - }); - render(); - - expect(screen.getByText("group1")).toBeInTheDocument(); - expect(screen.getByText("+2")).toBeInTheDocument(); - expect(screen.queryByText("group2")).not.toBeInTheDocument(); - expect(screen.queryByText("group3")).not.toBeInTheDocument(); - }); - - it("should display expanded access groups when expanded", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - new Set(["test-model-id"]), - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - id: "test-model-id", - access_groups: ["group1", "group2", "group3"], - }, - }); - render(); - - expect(screen.getByText("group1")).toBeInTheDocument(); - expect(screen.getByText("group2")).toBeInTheDocument(); - expect(screen.getByText("group3")).toBeInTheDocument(); - expect(screen.getByText("−")).toBeInTheDocument(); - }); - - it("should display single access group without expand button", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - access_groups: ["group1"], - }, - }); - render(); - - expect(screen.getByText("group1")).toBeInTheDocument(); - expect(screen.queryByText(/\+/)).not.toBeInTheDocument(); - }); - - it("should handle missing display name gracefully", () => { - const getDisplayModelName = vi.fn(() => ""); - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel(); - render(); - - expect(screen.getByText("-")).toBeInTheDocument(); - }); - - it("should handle missing created_at date", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - created_at: "", - }, - }); - render(); - - expect(screen.getByText("Unknown date")).toBeInTheDocument(); - }); - - it("should handle missing updated_at date", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - updated_at: "", - }, - }); - render(); - - const updatedAtCells = screen.getAllByText("-"); - expect(updatedAtCells.length).toBeGreaterThan(0); - }); - - it("should handle missing created_by for DB models", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: true, - created_by: "", - }, - }); - render(); - - expect(screen.getByText("Unknown")).toBeInTheDocument(); - }); - - it("should display only input cost when output cost is missing", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - input_cost: 0.01, - output_cost: undefined as any, - }); - render(); - - expect(screen.getByText("In: $0.01")).toBeInTheDocument(); - expect(screen.queryByText(/Out:/)).not.toBeInTheDocument(); - }); - - it("should display only output cost when input cost is missing", () => { - const cols = columns( - defaultProps.userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - ); - - const model = createMockModel({ - input_cost: undefined as any, - output_cost: 0.03, - }); - render(); - - expect(screen.getByText("Out: $0.03")).toBeInTheDocument(); - expect(screen.queryByText(/In:/)).not.toBeInTheDocument(); - }); - - describe("pause/resume toggle", () => { - const renderWithToggle = ( - overrides: Partial["model_info"]> = {}, - togglePauseHandler?: ReturnType, - userRole: string = "Admin", - ) => { - const handler = togglePauseHandler ?? vi.fn(); - const cols = columns( - userRole, - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - vi.fn(), - handler, - ); - const model = createMockModel({ - model_info: { ...createMockModel().model_info, ...overrides }, - }); - render(); - return { handler }; - }; - - it("renders the toggle ON for a db_model that is not blocked", () => { - renderWithToggle({ db_model: true, blocked: false }); - const toggle = screen.getByRole("switch", { name: /pause model/i }); - expect(toggle).toBeEnabled(); - expect(toggle).toHaveAttribute("aria-checked", "true"); - }); - - it("renders the toggle OFF for a db_model that is blocked", () => { - renderWithToggle({ db_model: true, blocked: true }); - const toggle = screen.getByRole("switch", { name: /resume model/i }); - expect(toggle).toBeEnabled(); - expect(toggle).toHaveAttribute("aria-checked", "false"); - }); - - it("calls the handler with blocked=true when an admin flips an active toggle off", async () => { - const handler = vi.fn(); - renderWithToggle({ db_model: true, blocked: false }, handler); - await userEvent.click(screen.getByRole("switch", { name: /pause model/i })); - expect(handler).toHaveBeenCalledWith("test-model-id", true); - }); - - it("calls the handler with blocked=false when an admin flips a paused toggle on", async () => { - const handler = vi.fn(); - renderWithToggle({ db_model: true, blocked: true }, handler); - await userEvent.click(screen.getByRole("switch", { name: /resume model/i })); - expect(handler).toHaveBeenCalledWith("test-model-id", false); - }); - - it("disables the toggle for non-admin users", () => { - const handler = vi.fn(); - renderWithToggle({ db_model: true, blocked: false }, handler, "User"); - const toggle = screen.getByRole("switch", { name: /pause model/i }); - expect(toggle).toBeDisabled(); - }); - - it("disables the toggle for config models", () => { - const handler = vi.fn(); - renderWithToggle({ db_model: false, blocked: false }, handler, "Admin"); - const toggle = screen.getByRole("switch", { name: /pause model/i }); - expect(toggle).toBeDisabled(); - }); - - it("disables the toggle while a PATCH for the same row is in-flight", () => { - // Regression for Greptile P1 on PR #28151 — antd's `loading` prop is - // visual only and does not prevent click events, so the row needs to - // be explicitly disabled while its PATCH is pending to avoid - // racing/conflicting PATCH calls on double-click. - const handler = vi.fn(); - const model = createMockModel({ - model_info: { - ...createMockModel().model_info, - db_model: true, - blocked: false, - }, - }); - const cols = columns( - "Admin", - defaultProps.userID, - defaultProps.premiumUser, - defaultProps.setSelectedModelId, - defaultProps.setSelectedTeamId, - defaultProps.getDisplayModelName, - defaultProps.handleEditClick, - defaultProps.handleRefreshClick, - defaultProps.expandedRows, - defaultProps.setExpandedRows, - vi.fn(), - handler, - model.model_info.id, // pausingModelId matches this row - ); - render(); - const toggle = screen.getByRole("switch", { name: /pause model/i }); - expect(toggle).toBeDisabled(); - }); - }); -}); diff --git a/ui/litellm-dashboard/src/components/molecules/models/columns.tsx b/ui/litellm-dashboard/src/components/molecules/models/columns.tsx deleted file mode 100644 index d043d0820f8..00000000000 --- a/ui/litellm-dashboard/src/components/molecules/models/columns.tsx +++ /dev/null @@ -1,423 +0,0 @@ -import { EditOutlined, InfoCircleOutlined, SyncOutlined } from "@ant-design/icons"; -import { TrashIcon } from "@heroicons/react/outline"; -import { ColumnDef } from "@tanstack/react-table"; -import { Badge, Icon } from "@tremor/react"; -import { Divider, Flex, Popover, Space, Switch, Tooltip, Typography } from "antd"; -import { DateCell, IdCell, StatusBadge } from "@/components/shared/table_cells"; -import { ModelData } from "../../model_dashboard/types"; -import { ProviderLogo } from "./ProviderLogo"; - -const { Text, Title } = Typography; - -const credentialsInfoPopoverContent = ( - - - Credential types - - - - - - - - Reusable - - - Credentials saved in LiteLLM that can be added to models repeatedly. - - - - - - - - - Manual - - - Credentials added directly during model creation or defined in the config file. - - - - -); - -export const columns = ( - userRole: string, - userID: string, - premiumUser: boolean, - setSelectedModelId: (id: string) => void, - setSelectedTeamId: (id: string) => void, - getDisplayModelName: (model: any) => string, - handleEditClick: (model: any) => void, - handleRefreshClick: () => void, - expandedRows: Set, - setExpandedRows: (expandedRows: Set) => void, - onDeleteClick?: (modelId: string) => void, - onTogglePauseClick?: (modelId: string, blocked: boolean) => void | Promise, - pausingModelId?: string | null, -): ColumnDef[] => [ - { - header: () => Model ID, - accessorKey: "model_info.id", - enableSorting: false, - size: 130, - minSize: 80, - cell: ({ row }) => { - const model = row.original; - return ( -
e.stopPropagation()}> - -
- ); - }, - }, - { - header: () => Model Information, - accessorKey: "model_name", - size: 250, - minSize: 120, - cell: ({ row }) => { - const model = row.original; - const displayName = getDisplayModelName(row.original) || "-"; - const popoverContent = ( - - - - - {model.provider || "Unknown provider"} - - - - - - - Public Model Name - - - {displayName} - - - - - - LiteLLM Model Name - - - {model.litellm_model_name || "-"} - - - - - ); - - return ( - -
-
- {model.provider ? ( - - ) : ( -
-
- )} -
- -
- - {displayName} - - - {model.litellm_model_name || "-"} - -
-
-
- ); - }, - }, - { - header: () => ( - - Credentials - - - - - ), - accessorKey: "litellm_credential_name", - enableSorting: false, - size: 180, - minSize: 100, - cell: ({ row }) => { - const model = row.original; - const credentialName = model.litellm_params?.litellm_credential_name; - const isReusable = !!credentialName; - - return ( -
- {isReusable ? ( - <> - - - {credentialName} - - - ) : ( - <> - - Manual - - )} -
- ); - }, - }, - { - header: () => Created By, - accessorKey: "model_info.created_by", - sortingFn: "datetime", - size: 160, - minSize: 100, - cell: ({ row }) => { - const model = row.original; - const isConfigModel = !model.model_info?.db_model; - const createdBy = model.model_info.created_by; - const createdAt = model.model_info.created_at ? new Date(model.model_info.created_at).toLocaleDateString() : null; - - return ( -
- {/* Created By - Primary */} -
- {isConfigModel ? "Defined in config" : createdBy || "Unknown"} -
- {/* Created At - Secondary */} -
- {isConfigModel ? "-" : createdAt || "Unknown date"} -
-
- ); - }, - }, - { - header: () => Updated At, - accessorKey: "model_info.updated_at", - sortingFn: "datetime", - size: 120, - minSize: 80, - cell: ({ row }) => { - const model = row.original; - return ; - }, - }, - { - header: () => Costs, - accessorKey: "input_cost", - size: 120, - minSize: 80, - cell: ({ row }) => { - const model = row.original; - const inputCost = model.input_cost; - const outputCost = model.output_cost; - - // If both costs are missing or undefined, show "-" - if (inputCost == null && outputCost == null) { - return ( -
- - -
- ); - } - - return ( - -
- {/* Input Cost - Primary */} - {inputCost != null &&
In: ${inputCost}
} - {/* Output Cost - Secondary */} - {outputCost != null &&
Out: ${outputCost}
} -
-
- ); - }, - }, - { - header: () => Team ID, - accessorKey: "model_info.team_id", - enableSorting: false, - size: 130, - minSize: 80, - cell: ({ row }) => { - const model = row.original; - return model.model_info.team_id ? ( -
e.stopPropagation()}> - -
- ) : ( - "-" - ); - }, - }, - { - header: () => Model Access Group, - accessorKey: "model_info.model_access_group", - enableSorting: false, - size: 180, - minSize: 100, - cell: ({ row }) => { - const model = row.original; - const accessGroups = model.model_info.access_groups; - - if (!accessGroups || accessGroups.length === 0) { - return "-"; - } - - const modelId = model.model_info.id; - const isExpanded = expandedRows.has(modelId); - const shouldShowExpandButton = accessGroups.length > 1; - - const toggleExpanded = () => { - const newExpanded = new Set(expandedRows); - if (isExpanded) { - newExpanded.delete(modelId); - } else { - newExpanded.add(modelId); - } - setExpandedRows(newExpanded); - }; - - return ( -
- - {accessGroups[0]} - - - {(isExpanded || (!shouldShowExpandButton && accessGroups.length === 2)) && - accessGroups.slice(1).map((group: string, index: number) => ( - - {group} - - ))} - - {shouldShowExpandButton && ( - - )} -
- ); - }, - }, - { - header: () => Status, - accessorKey: "model_info.db_model", - size: 120, - minSize: 80, - cell: ({ row }) => { - const model = row.original; - return model.model_info.db_model ? ( - - ) : ( - - ); - }, - }, - { - id: "actions", - header: () => Actions, - size: 100, - minSize: 80, - enableResizing: false, - cell: ({ row }) => { - const model = row.original; - const canEditModel = userRole === "Admin" || model.model_info?.created_by === userID; - const isConfigModel = !model.model_info?.db_model; - const isAdmin = userRole === "Admin"; - const isBlocked = model.model_info?.blocked === true; - const isPauseToggleable = !isConfigModel && isAdmin && Boolean(onTogglePauseClick); - const pauseTooltip = isConfigModel - ? "Config models cannot be paused from the dashboard. Pause is DB-backed." - : !isAdmin - ? "Only proxy admins can pause or resume a model." - : isBlocked - ? "Resume model — restore normal routing." - : "Pause model — stop routing requests until resumed."; - // antd's `loading` prop on Switch is purely cosmetic — it does not block - // clicks. Pair `loading` with `disabled` derived from the same condition - // so a double-click during a pending PATCH cannot send a second, - // conflicting `blocked` value. - const isPausing = pausingModelId === model.model_info?.id; - return ( -
- - { - e.stopPropagation(); - }} - onChange={(nextChecked) => { - const modelId = model.model_info?.id; - if (isPauseToggleable && onTogglePauseClick && modelId) { - void onTogglePauseClick(modelId, !nextChecked); - } - }} - /> - - {isConfigModel ? ( - - - - ) : ( - - { - e.stopPropagation(); - if (canEditModel && onDeleteClick) { - onDeleteClick(model.model_info.id); - } - }} - className={!canEditModel ? "opacity-50 cursor-not-allowed" : "cursor-pointer hover:text-red-600"} - /> - - )} -
- ); - }, - }, -]; diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.tsx index 8aaec3f13a0..ea986a84cc3 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.tsx @@ -20,6 +20,8 @@ interface DataTableFilterDrawerProps { description?: React.ReactNode; applyLabel?: string; resetLabel?: string; + /** Runs instead of the default "clear this table's column filters" when the reset button is pressed. */ + onReset?: () => void; children: (draft: FilterDraft) => React.ReactNode; } @@ -48,6 +50,7 @@ export function DataTableFilterDrawer({ description, applyLabel = "Apply Filters", resetLabel = "Reset", + onReset, children, }: DataTableFilterDrawerProps) { const [draft, setDraft] = React.useState>(() => toDraft(table.getState().columnFilters)); @@ -72,6 +75,10 @@ export function DataTableFilterDrawer({ const reset = () => { setDraft({}); + if (onReset !== undefined) { + onReset(); + return; + } table.setColumnFilters([]); }; diff --git a/ui/litellm-dashboard/src/components/ui/hover-card.tsx b/ui/litellm-dashboard/src/components/ui/hover-card.tsx new file mode 100644 index 00000000000..586a77b7cb6 --- /dev/null +++ b/ui/litellm-dashboard/src/components/ui/hover-card.tsx @@ -0,0 +1,46 @@ +"use client"; + +import { PreviewCard as PreviewCardPrimitive } from "@base-ui/react/preview-card"; + +import { cn } from "@/lib/cva.config"; + +function HoverCard({ ...props }: PreviewCardPrimitive.Root.Props) { + return ; +} + +function HoverCardTrigger({ ...props }: PreviewCardPrimitive.Trigger.Props) { + return ; +} + +function HoverCardContent({ + className, + side = "bottom", + sideOffset = 4, + align = "center", + alignOffset = 4, + ...props +}: PreviewCardPrimitive.Popup.Props & + Pick) { + return ( + + + + + + ); +} + +export { HoverCard, HoverCardTrigger, HoverCardContent }; From 0a4333580faaa70a66ca378f7f5f47cc1da5efa4 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 23 Jul 2026 10:32:08 -0700 Subject: [PATCH 08/25] refactor(ui): migrate request logs table onto the shared DataTable (#34343) * refactor(ui): migrate request logs table onto the shared DataTable Moves the Request Logs tab off the local view_logs/table.tsx clone and onto the shared DataTable in server sort, pagination, and filter mode. The container is split into RequestLogsPanel (data owner: the spend-logs query, the session dedup and composition pipeline, and the detail drawer), a thin RequestLogsTable, and RequestLogsTableColumns. The clone itself stays for now because TopModelView and TopKeyView still consume it The advanced filter bar moves into the shared DataTableFilterDrawer, so filters commit on Apply and render as removable chips. That makes the per-keystroke debounce in the query hook redundant, and the hook now takes ColumnFiltersState, PaginationState, and SortingState directly instead of carrying its own filter shape. Reset still restores the default 24 hour window alongside the filters Adds shared/PaginatedSearchSelect, a Base UI combobox with server-side search and infinite scroll, and uses it for the Key Alias and Model filters. That retires the three logs-only antd pickers (PaginatedKeyAliasSelect, PaginatedModelSelect, FilterTeamDropdown) and the FilterComponent molecule they plugged into. The shared TeamDropdown is deliberately untouched: six other surfaces still render it, five of them as a bare child of an antd Form.Item that injects value/onChange implicitly * test(ui): pin team-scoped key alias filtering in the logs filter drawer The Key Alias filter narrows its options to the team selected in the same drawer, a cross-filter dependency carried over from the antd picker it replaced. Nothing covered it: the live QA pass explicitly did not exercise it either, so it was the one behaviour in this migration that could regress silently Asserts the selected team id reaches useInfiniteKeyAliases, that the lookup stays unscoped when no team is picked, and that the scope does not leak into the Model lookup, which shares the same combobox but takes no team --- ui/litellm-dashboard/eslint-suppressions.json | 41 -- .../PaginatedKeyAliasSelect.test.tsx | 249 ------- .../PaginatedKeyAliasSelect.tsx | 107 --- .../PaginatedModelSelect.test.tsx | 293 -------- .../PaginatedModelSelect.tsx | 140 ---- .../common_components/FilterTeamDropdown.tsx | 9 - .../src/components/molecules/filter.test.tsx | 615 ---------------- .../src/components/molecules/filter.tsx | 221 ------ .../shared/PaginatedSearchSelect.test.tsx | 165 +++++ .../shared/PaginatedSearchSelect.tsx | 119 ++++ .../components/view_logs/LogsTableToolbar.tsx | 295 +++----- .../view_logs/RequestLogsFilters.test.tsx | 91 +++ .../view_logs/RequestLogsFilters.tsx | 319 +++++++++ .../view_logs/RequestLogsPanel.test.tsx | 202 ++++++ .../components/view_logs/RequestLogsPanel.tsx | 264 +++++++ .../components/view_logs/RequestLogsTable.tsx | 131 ++++ .../RequestLogsTableColumns.test.tsx | 136 ++++ .../view_logs/RequestLogsTableColumns.tsx | 317 +++++++++ .../src/components/view_logs/columns.test.tsx | 70 -- .../src/components/view_logs/columns.tsx | 492 ------------- .../components/view_logs/filter_options.ts | 82 --- .../src/components/view_logs/index.test.tsx | 195 ++--- .../src/components/view_logs/index.tsx | 284 +------- .../view_logs/log_filter_logic.test.tsx | 673 ++++-------------- .../components/view_logs/log_filter_logic.tsx | 196 ++--- 25 files changed, 2092 insertions(+), 3614 deletions(-) delete mode 100644 ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.test.tsx delete mode 100644 ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx delete mode 100644 ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.test.tsx delete mode 100644 ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx delete mode 100644 ui/litellm-dashboard/src/components/common_components/FilterTeamDropdown.tsx delete mode 100644 ui/litellm-dashboard/src/components/molecules/filter.test.tsx delete mode 100644 ui/litellm-dashboard/src/components/molecules/filter.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/PaginatedSearchSelect.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/PaginatedSearchSelect.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx delete mode 100644 ui/litellm-dashboard/src/components/view_logs/columns.test.tsx delete mode 100644 ui/litellm-dashboard/src/components/view_logs/filter_options.ts diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 6a50f4aa99e..8e8c1447d22 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -2349,11 +2349,6 @@ "count": 1 } }, - "src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/LicenseExpiryBanner.tsx": { "no-restricted-imports": { "count": 1 @@ -2364,14 +2359,6 @@ "count": 1 } }, - "src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/Navbar/BlogDropdown/BlogDropdown.test.tsx": { "max-nested-callbacks": { "count": 12 @@ -3511,17 +3498,6 @@ "count": 1 } }, - "src/components/molecules/filter.tsx": { - "local/filename-pascal-case": { - "count": 1 - }, - "no-nested-ternary": { - "count": 2 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/molecules/message_manager.tsx": { "local/filename-pascal-case": { "count": 1 @@ -4404,17 +4380,6 @@ "count": 2 } }, - "src/components/view_logs/LogsTableToolbar.tsx": { - "local/no-complex-jsx-arrow": { - "count": 1 - }, - "no-nested-ternary": { - "count": 4 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/view_logs/ToolsSection/FormattedToolView.tsx": { "no-restricted-imports": { "count": 1 @@ -4443,9 +4408,6 @@ "src/components/view_logs/columns.tsx": { "local/filename-pascal-case": { "count": 1 - }, - "no-restricted-imports": { - "count": 1 } }, "src/components/view_logs/index.tsx": { @@ -4454,9 +4416,6 @@ }, "no-restricted-imports": { "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 } }, "src/components/view_logs/log_filter_logic.tsx": { diff --git a/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.test.tsx b/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.test.tsx deleted file mode 100644 index 7ef2e9d0def..00000000000 --- a/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.test.tsx +++ /dev/null @@ -1,249 +0,0 @@ -import { screen, waitFor } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { beforeEach, describe, expect, it, vi } from "vitest"; -import { renderWithProviders } from "../../../../tests/test-utils"; -import { PaginatedKeyAliasSelect } from "./PaginatedKeyAliasSelect"; - -const mockFetchNextPage = vi.fn(); - -vi.mock("@/app/(dashboard)/hooks/keys/useKeyAliases", () => ({ - useInfiniteKeyAliases: vi.fn(), -})); - -vi.mock("@tanstack/react-pacer/debouncer", async () => { - const React = await vi.importActual("react"); - return { - useDebouncedState: (initial: string) => { - const [value, setValue] = React.useState(initial); - return [value, setValue]; - }, - }; -}); - -import { useInfiniteKeyAliases } from "@/app/(dashboard)/hooks/keys/useKeyAliases"; - -const mockUseInfiniteKeyAliases = vi.mocked(useInfiniteKeyAliases); - -const mockPagesWithAliases = { - pages: [ - { - aliases: ["alias-1", "alias-2"], - total_count: 2, - current_page: 1, - total_pages: 1, - size: 50, - }, - ], -}; - -const mockEmptyPages = { - pages: [{ aliases: [], total_count: 0, current_page: 1, total_pages: 1, size: 50 }], -}; - -describe("PaginatedKeyAliasSelect", () => { - const mockOnChange = vi.fn(); - - const defaultHookReturn = { - data: mockPagesWithAliases, - fetchNextPage: mockFetchNextPage, - hasNextPage: false, - isFetchingNextPage: false, - isLoading: false, - }; - - beforeEach(() => { - vi.clearAllMocks(); - mockUseInfiniteKeyAliases.mockReturnValue(defaultHookReturn as any); - }); - - it("should render", () => { - renderWithProviders(); - - expect(screen.getByRole("combobox")).toBeInTheDocument(); - expect(screen.getByText("Select a key alias")).toBeInTheDocument(); - }); - - it("should display custom placeholder when provided", () => { - renderWithProviders(); - - expect(screen.getByText("Choose alias")).toBeInTheDocument(); - }); - - it("should display alias options when data is loaded", async () => { - renderWithProviders(); - - const combobox = screen.getByRole("combobox"); - await userEvent.click(combobox); - - await waitFor(() => { - expect(screen.getByRole("option", { name: "alias-1" })).toBeInTheDocument(); - expect(screen.getByRole("option", { name: "alias-2" })).toBeInTheDocument(); - }); - }); - - it("should call onChange when user selects an alias", async () => { - const user = userEvent.setup({ delay: null }); - renderWithProviders(); - - const combobox = screen.getByRole("combobox"); - await user.click(combobox); - - const option = await screen.findByTitle("alias-1"); - await user.click(option); - - await waitFor(() => { - expect(mockOnChange).toHaveBeenCalledWith("alias-1"); - }); - }); - - it("should show loading state when isLoading is true", () => { - mockUseInfiniteKeyAliases.mockReturnValue({ - ...defaultHookReturn, - isLoading: true, - } as any); - - renderWithProviders(); - - expect(screen.getByRole("combobox")).toHaveAttribute("aria-expanded", "false"); - }); - - it("should pass pageSize to useInfiniteKeyAliases", () => { - renderWithProviders(); - - expect(mockUseInfiniteKeyAliases).toHaveBeenCalledWith(25, undefined, undefined); - }); - - it("should pass search to useInfiniteKeyAliases when user types", async () => { - const user = userEvent.setup(); - renderWithProviders(); - - const combobox = screen.getByRole("combobox"); - await user.click(combobox); - await user.keyboard("my-alias"); - - await waitFor(() => { - expect(mockUseInfiniteKeyAliases).toHaveBeenCalledWith(50, "my-alias", undefined); - }); - }); - - it("should have scroll container for infinite loading when hasNextPage is true", async () => { - mockUseInfiniteKeyAliases.mockReturnValue({ - ...defaultHookReturn, - hasNextPage: true, - isFetchingNextPage: false, - } as any); - - renderWithProviders(); - - const combobox = screen.getByRole("combobox"); - await userEvent.click(combobox); - - await waitFor(() => { - expect(screen.getByRole("option", { name: "alias-1" })).toBeInTheDocument(); - }); - - const scrollableContainer = document.querySelector(".ant-select-dropdown .rc-virtual-list-holder"); - expect(scrollableContainer).toBeInTheDocument(); - }); - - it("should deduplicate aliases with the same value across pages", async () => { - mockUseInfiniteKeyAliases.mockReturnValue({ - ...defaultHookReturn, - data: { - pages: [ - { - aliases: ["alias-1", "alias-1"], - total_count: 2, - current_page: 1, - total_pages: 1, - size: 50, - }, - ], - }, - } as any); - - renderWithProviders(); - - const combobox = screen.getByRole("combobox"); - await userEvent.click(combobox); - - await waitFor(() => { - const options = screen.queryAllByRole("option", { name: "alias-1" }); - expect(options.length).toBe(1); - }); - }); - - it("should skip empty aliases", async () => { - mockUseInfiniteKeyAliases.mockReturnValue({ - ...defaultHookReturn, - data: { - pages: [ - { - aliases: ["valid-alias", "", null], - total_count: 3, - current_page: 1, - total_pages: 1, - size: 50, - }, - ], - }, - } as any); - - renderWithProviders(); - - const combobox = screen.getByRole("combobox"); - await userEvent.click(combobox); - - await waitFor(() => { - expect(screen.getByRole("option", { name: "valid-alias" })).toBeInTheDocument(); - const allOptions = screen.queryAllByRole("option"); - expect(allOptions.length).toBe(1); - }); - }); - - it("should respect allowClear prop", () => { - renderWithProviders(); - - expect(screen.getByRole("combobox")).toBeInTheDocument(); - }); - - it("should respect disabled prop", () => { - renderWithProviders(); - - const combobox = screen.getByRole("combobox"); - expect(combobox.closest(".ant-select")).toHaveClass("ant-select-disabled"); - }); - - it("should not call fetchNextPage when hasNextPage is false", async () => { - mockUseInfiniteKeyAliases.mockReturnValue({ - ...defaultHookReturn, - hasNextPage: false, - } as any); - - renderWithProviders(); - - await userEvent.click(screen.getByRole("combobox")); - - await waitFor(() => { - expect(screen.getByRole("option", { name: "alias-1" })).toBeInTheDocument(); - }); - - expect(mockFetchNextPage).not.toHaveBeenCalled(); - }); - - it("should show no aliases found when data is empty", async () => { - mockUseInfiniteKeyAliases.mockReturnValue({ - ...defaultHookReturn, - data: mockEmptyPages, - } as any); - - renderWithProviders(); - - const combobox = screen.getByRole("combobox"); - await userEvent.click(combobox); - - await waitFor(() => { - expect(screen.getByText("No key aliases found")).toBeInTheDocument(); - }); - }); -}); diff --git a/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx b/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx deleted file mode 100644 index 1d19ba3255d..00000000000 --- a/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx +++ /dev/null @@ -1,107 +0,0 @@ -import { useInfiniteKeyAliases } from "@/app/(dashboard)/hooks/keys/useKeyAliases"; -import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; -import { LoadingOutlined } from "@ant-design/icons"; -import { useDebouncedState } from "@tanstack/react-pacer/debouncer"; -import { Select } from "antd"; -import { useMemo, useState, type UIEvent } from "react"; - -export interface PaginatedKeyAliasSelectProps { - value?: string; - onChange?: (value: string) => void; - placeholder?: string; - style?: React.CSSProperties; - pageSize?: number; - allowClear?: boolean; - disabled?: boolean; - allFilters?: { [key: string]: string }; -} - -const SCROLL_THRESHOLD = 0.8; - -export const PaginatedKeyAliasSelect = ({ - value, - onChange, - placeholder = "Select a key alias", - style, - pageSize = 50, - allowClear = true, - disabled = false, - allFilters, -}: PaginatedKeyAliasSelectProps) => { - const [searchInput, setSearchInput] = useState(""); - const [debouncedSearch, setDebouncedSearch] = useDebouncedState("", { - wait: DEBOUNCE_WAIT_MS, - }); - - const teamId = allFilters?.["Team ID"] || undefined; - - const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteKeyAliases( - pageSize, - debouncedSearch || undefined, - teamId, - ); - - const options = useMemo(() => { - if (!data?.pages) return []; - - const seen = new Set(); - const result: { label: string; value: string }[] = []; - - for (const page of data.pages) { - for (const alias of page.aliases) { - if (!alias || seen.has(alias)) continue; - seen.add(alias); - result.push({ label: alias, value: alias }); - } - } - - return result; - }, [data]); - - const handlePopupScroll = (e: UIEvent) => { - const target = e.currentTarget; - const scrollRatio = (target.scrollTop + target.clientHeight) / target.scrollHeight; - - if (scrollRatio >= SCROLL_THRESHOLD && hasNextPage && !isFetchingNextPage) { - fetchNextPage(); - } - }; - - const handleSearch = (value: string) => { - setSearchInput(value); - setDebouncedSearch(value); - }; - - const handleChange = (v: string | null) => { - onChange?.(v ?? ""); - }; - - return ( - : "No models found"} - options={options} - optionRender={optionRender} - popupRender={(menu) => ( - <> - {menu} - {isFetchingNextPage && ( -
- -
- )} - - )} - /> - ); -}; diff --git a/ui/litellm-dashboard/src/components/common_components/FilterTeamDropdown.tsx b/ui/litellm-dashboard/src/components/common_components/FilterTeamDropdown.tsx deleted file mode 100644 index 3b756b94f12..00000000000 --- a/ui/litellm-dashboard/src/components/common_components/FilterTeamDropdown.tsx +++ /dev/null @@ -1,9 +0,0 @@ -import React from "react"; -import TeamDropdown from "./team_dropdown"; -import type { FilterOptionCustomComponentProps } from "../molecules/filter"; - -const FilterTeamDropdown: React.FC = ({ value, onChange }) => ( - -); - -export default FilterTeamDropdown; diff --git a/ui/litellm-dashboard/src/components/molecules/filter.test.tsx b/ui/litellm-dashboard/src/components/molecules/filter.test.tsx deleted file mode 100644 index 66c671b12d4..00000000000 --- a/ui/litellm-dashboard/src/components/molecules/filter.test.tsx +++ /dev/null @@ -1,615 +0,0 @@ -import { screen, waitFor, within } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { beforeEach, describe, expect, it, vi } from "vitest"; -import { renderWithProviders } from "../../../tests/test-utils"; -import FilterComponent, { FilterOption } from "./filter"; - -describe("FilterComponent", () => { - const mockOnApplyFilters = vi.fn(); - const mockOnResetFilters = vi.fn(); - - const defaultOptions: FilterOption[] = [ - { - name: "teamId", - label: "Team ID", - options: [ - { label: "Team 1", value: "team1" }, - { label: "Team 2", value: "team2" }, - ], - }, - { - name: "status", - label: "Status", - options: [ - { label: "Active", value: "active" }, - { label: "Inactive", value: "inactive" }, - ], - }, - { - name: "userId", - label: "User ID", - }, - ]; - - beforeEach(() => { - vi.clearAllMocks(); - }); - - it("should render", () => { - renderWithProviders( - , - ); - expect(screen.getByRole("button", { name: "Filters" })).toBeInTheDocument(); - }); - - it("should display custom button label", () => { - renderWithProviders( - , - ); - expect(screen.getByRole("button", { name: "Custom Filters" })).toBeInTheDocument(); - }); - - it("should toggle filters visibility when filter button is clicked", async () => { - const user = userEvent.setup({ delay: null }); - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - expect(screen.queryByPlaceholderText("Enter User ID...")).not.toBeInTheDocument(); - - await user.click(filterButton); - - await waitFor(() => { - expect(screen.getByPlaceholderText("Enter User ID...")).toBeInTheDocument(); - }); - - await user.click(filterButton); - - await waitFor(() => { - expect(screen.queryByPlaceholderText("Enter User ID...")).not.toBeInTheDocument(); - }); - }); - - it("should call onResetFilters when reset button is clicked", async () => { - const user = userEvent.setup({ delay: null }); - renderWithProviders( - , - ); - - const resetButton = screen.getByRole("button", { name: "Reset Filters" }); - await user.click(resetButton); - - await waitFor(() => { - expect(mockOnResetFilters).toHaveBeenCalledTimes(1); - }); - }); - - it("renders filters in the caller-supplied order", async () => { - const user = userEvent.setup({ delay: null }); - const options: FilterOption[] = [ - { name: "model", label: "Model" }, - { name: "teamId", label: "Team ID" }, - { name: "status", label: "Status" }, - { name: "userId", label: "User ID" }, - ]; - - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - const labels = screen.getAllByText(/^(Team ID|Status|User ID|Model)$/); - expect(labels.map((l) => l.textContent)).toEqual(["Model", "Team ID", "Status", "User ID"]); - }); - }); - - it("should handle input filter changes", async () => { - const user = userEvent.setup({ delay: null }); - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - const userIdInput = screen.getByPlaceholderText("Enter User ID..."); - await user.type(userIdInput, "user123"); - - await waitFor(() => { - expect(mockOnApplyFilters).toHaveBeenCalledWith({ userId: "user123" }); - }); - }); - - it("should display initial values in filters", async () => { - const user = userEvent.setup({ delay: null }); - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - const userIdInput = screen.getByPlaceholderText("Enter User ID...") as HTMLInputElement; - expect(userIdInput.value).toBe("user123"); - }); - }); - - it("should handle select dropdown filter changes", async () => { - const user = userEvent.setup({ delay: null }); - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - const teamIdLabel = screen.getByText("Team ID"); - const teamIdSection = teamIdLabel.closest("div"); - const teamIdSelect = within(teamIdSection!).getByRole("combobox"); - - await user.click(teamIdSelect); - - await waitFor(() => { - expect(screen.getByText("Team 1")).toBeInTheDocument(); - }); - - await user.click(screen.getByText("Team 1")); - - await waitFor(() => { - expect(mockOnApplyFilters).toHaveBeenCalledWith({ teamId: "team1" }); - }); - }); - - it("should handle searchable filter with search function", async () => { - const user = userEvent.setup({ delay: null }); - const mockSearchFn = vi.fn().mockResolvedValue([ - { label: "Result 1", value: "result1" }, - { label: "Result 2", value: "result2" }, - ]); - - const options: FilterOption[] = [ - { - name: "model", - label: "Model", - isSearchable: true, - searchFn: mockSearchFn, - }, - ]; - - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - expect(mockSearchFn).toHaveBeenCalledWith(""); - }); - - const modelLabel = screen.getByText("Model"); - const modelSection = modelLabel.closest("div"); - const modelSelect = within(modelSection!).getByRole("combobox"); - await user.click(modelSelect); - - await waitFor(() => { - expect(screen.getByText("Result 1")).toBeInTheDocument(); - expect(screen.getByText("Result 2")).toBeInTheDocument(); - }); - }); - - it("should debounce search input for searchable filters", async () => { - const user = userEvent.setup({ delay: null }); - const mockSearchFn = vi.fn().mockResolvedValue([{ label: "Result", value: "result" }]); - - const options: FilterOption[] = [ - { - name: "model", - label: "Model", - isSearchable: true, - searchFn: mockSearchFn, - }, - ]; - - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - expect(mockSearchFn).toHaveBeenCalledWith(""); - }); - - vi.clearAllMocks(); - - const modelLabel = screen.getByText("Model"); - const modelSection = modelLabel.closest("div"); - const modelSelect = within(modelSection!).getByRole("combobox"); - await user.click(modelSelect); - await user.type(modelSelect, "test"); - - expect(mockSearchFn).not.toHaveBeenCalled(); - - await waitFor( - () => { - expect(mockSearchFn).toHaveBeenCalledWith("test"); - }, - { timeout: 500 }, - ); - }); - - it("should show loading state when searching", async () => { - const user = userEvent.setup({ delay: null }); - let resolveSearch: (value: Array<{ label: string; value: string }>) => void; - const mockSearchFn = vi.fn().mockImplementation( - () => - new Promise>((resolve) => { - resolveSearch = resolve; - }), - ); - - const options: FilterOption[] = [ - { - name: "model", - label: "Model", - isSearchable: true, - searchFn: mockSearchFn, - }, - ]; - - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - expect(mockSearchFn).toHaveBeenCalledWith(""); - }); - - const modelLabel = screen.getByText("Model"); - const modelSection = modelLabel.closest("div"); - const modelSelect = within(modelSection!).getByRole("combobox"); - await user.click(modelSelect); - await user.type(modelSelect, "test"); - - await waitFor( - () => { - expect(screen.getByText("Loading...")).toBeInTheDocument(); - }, - { timeout: 500 }, - ); - - resolveSearch!([{ label: "Result", value: "result" }]); - - await waitFor(() => { - expect(screen.queryByText("Loading...")).not.toBeInTheDocument(); - }); - }); - - it("shows a loading state (not an empty list) while a searchable filter's data is still loading", async () => { - const user = userEvent.setup({ delay: null }); - const mockSearchFn = vi.fn().mockResolvedValue([]); - - const options: FilterOption[] = [ - { - name: "model", - label: "Model", - isSearchable: true, - loading: true, - searchFn: mockSearchFn, - }, - ]; - - renderWithProviders( - , - ); - - await user.click(screen.getByRole("button", { name: "Filters" })); - - const modelLabel = screen.getByText("Model"); - const modelSelect = within(modelLabel.closest("div")!).getByRole("combobox"); - await user.click(modelSelect); - - await waitFor(() => { - expect(screen.getByText("Loading...")).toBeInTheDocument(); - }); - expect(screen.queryByText("No results found")).not.toBeInTheDocument(); - // It must not cache an empty initial-options list while the source is still loading. - expect(mockSearchFn).not.toHaveBeenCalled(); - }); - - it("loads initial options once a searchable filter's data finishes loading", async () => { - const user = userEvent.setup({ delay: null }); - const mockSearchFn = vi.fn().mockResolvedValue([{ label: "Team A", value: "team-a" }]); - const baseOption: FilterOption = { name: "model", label: "Model", isSearchable: true, searchFn: mockSearchFn }; - - const { rerender } = renderWithProviders( - , - ); - - await user.click(screen.getByRole("button", { name: "Filters" })); - expect(mockSearchFn).not.toHaveBeenCalled(); - - rerender( - , - ); - - await waitFor(() => { - expect(mockSearchFn).toHaveBeenCalledWith(""); - }); - }); - - it("should handle search errors gracefully", async () => { - const user = userEvent.setup({ delay: null }); - const consoleErrorSpy = vi.spyOn(console, "error").mockImplementation(() => {}); - const mockSearchFn = vi.fn().mockRejectedValue(new Error("Search failed")); - - const options: FilterOption[] = [ - { - name: "model", - label: "Model", - isSearchable: true, - searchFn: mockSearchFn, - }, - ]; - - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - expect(mockSearchFn).toHaveBeenCalledWith(""); - }); - - const modelLabel = screen.getByText("Model"); - const modelSection = modelLabel.closest("div"); - const modelSelect = within(modelSection!).getByRole("combobox"); - await user.click(modelSelect); - await user.type(modelSelect, "test"); - - await waitFor( - () => { - expect(consoleErrorSpy).toHaveBeenCalledWith("Error searching:", expect.any(Error)); - expect(screen.getByText("No results found")).toBeInTheDocument(); - }, - { timeout: 500 }, - ); - - consoleErrorSpy.mockRestore(); - }); - - it("should load initial options when dropdown opens for searchable filter", async () => { - const user = userEvent.setup({ delay: null }); - const mockSearchFn = vi.fn().mockResolvedValue([{ label: "Initial Result", value: "initial" }]); - - const options: FilterOption[] = [ - { - name: "model", - label: "Model", - isSearchable: true, - searchFn: mockSearchFn, - }, - ]; - - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - expect(mockSearchFn).toHaveBeenCalledWith(""); - }); - - vi.clearAllMocks(); - - const modelLabel = screen.getByText("Model"); - const modelSection = modelLabel.closest("div"); - const modelSelect = within(modelSection!).getByRole("combobox"); - await user.click(modelSelect); - - await waitFor(() => { - expect(screen.getByText("Initial Result")).toBeInTheDocument(); - }); - }); - - it("renders caller-supplied options that match no predefined filter name (LIT-3151)", async () => { - const user = userEvent.setup({ delay: null }); - const options: FilterOption[] = [ - { - name: "unknownFilter", - label: "Unknown Filter", - }, - ]; - - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - expect(screen.getByText("Unknown Filter")).toBeInTheDocument(); - expect(screen.getByPlaceholderText("Enter Unknown Filter...")).toBeInTheDocument(); - }); - }); - - it("renders every Tool Policies filter when none match a predefined name (LIT-3151)", async () => { - const user = userEvent.setup({ delay: null }); - const options: FilterOption[] = [ - { name: "Input Policy", label: "Input Policy", options: [{ label: "Trusted", value: "trusted" }] }, - { name: "Output Policy", label: "Output Policy", options: [{ label: "Blocked", value: "blocked" }] }, - { name: "Team Name", label: "Team Name", options: [] }, - { name: "Key Name", label: "Key Name", options: [] }, - ]; - - renderWithProviders( - , - ); - - await user.click(screen.getByRole("button", { name: "Filters" })); - - await waitFor(() => { - const labels = screen.getAllByText(/^(Input Policy|Output Policy|Team Name|Key Name)$/); - expect(labels.map((l) => l.textContent)).toEqual(["Input Policy", "Output Policy", "Team Name", "Key Name"]); - }); - }); - - it("should call onApplyFilters with updated values when multiple filters change", async () => { - const user = userEvent.setup({ delay: null }); - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - const userIdInput = screen.getByPlaceholderText("Enter User ID..."); - await user.type(userIdInput, "user123"); - - await waitFor(() => { - expect(mockOnApplyFilters).toHaveBeenCalledWith({ userId: "user123" }); - }); - - const teamIdLabel = screen.getByText("Team ID"); - const teamIdSection = teamIdLabel.closest("div"); - const teamIdSelect = within(teamIdSection!).getByRole("combobox"); - await user.click(teamIdSelect); - - await waitFor(() => { - expect(screen.getByText("Team 1")).toBeInTheDocument(); - }); - - await user.click(screen.getByText("Team 1")); - - await waitFor(() => { - expect(mockOnApplyFilters).toHaveBeenCalledWith({ - userId: "user123", - teamId: "team1", - }); - }); - }); - - it("cancels a pending debounced search when the component unmounts mid-type", async () => { - const user = userEvent.setup({ delay: null }); - const mockSearchFn = vi.fn().mockResolvedValue([{ label: "Result", value: "result" }]); - - const options: FilterOption[] = [ - { - name: "model", - label: "Model", - isSearchable: true, - searchFn: mockSearchFn, - }, - ]; - - const { unmount } = renderWithProviders( - , - ); - - await user.click(screen.getByRole("button", { name: "Filters" })); - - await waitFor(() => { - expect(mockSearchFn).toHaveBeenCalledWith(""); - }); - - vi.clearAllMocks(); - - const modelLabel = screen.getByText("Model"); - const modelSelect = within(modelLabel.closest("div")!).getByRole("combobox"); - await user.click(modelSelect); - await user.type(modelSelect, "test"); - - expect(mockSearchFn).not.toHaveBeenCalled(); - - unmount(); - - await new Promise((resolve) => setTimeout(resolve, 400)); - expect(mockSearchFn).not.toHaveBeenCalled(); - }); - - it("should reset all filter values when reset button is clicked", async () => { - const user = userEvent.setup({ delay: null }); - renderWithProviders( - , - ); - - const filterButton = screen.getByRole("button", { name: "Filters" }); - await user.click(filterButton); - - await waitFor(() => { - const userIdInput = screen.getByPlaceholderText("Enter User ID...") as HTMLInputElement; - expect(userIdInput.value).toBe("user123"); - }); - - const resetButton = screen.getByRole("button", { name: "Reset Filters" }); - await user.click(resetButton); - - await waitFor(() => { - const userIdInput = screen.getByPlaceholderText("Enter User ID...") as HTMLInputElement; - expect(userIdInput.value).toBe(""); - }); - }); -}); diff --git a/ui/litellm-dashboard/src/components/molecules/filter.tsx b/ui/litellm-dashboard/src/components/molecules/filter.tsx deleted file mode 100644 index 8218de41a19..00000000000 --- a/ui/litellm-dashboard/src/components/molecules/filter.tsx +++ /dev/null @@ -1,221 +0,0 @@ -import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; -import { FilterIcon } from "@heroicons/react/outline"; -import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer"; -import { Button, Input, Select } from "antd"; -import React, { useCallback, useEffect, useState } from "react"; - -export interface FilterOptionCustomComponentProps { - value?: string; - onChange: (value: string) => void; - placeholder?: string; - allFilters?: { [key: string]: string }; -} - -export interface FilterOption { - name: string; - label?: string; - isSearchable?: boolean; - searchFn?: (searchText: string) => Promise>; - options?: Array<{ label: string; value: string }>; - customComponent?: React.ComponentType; - loading?: boolean; -} - -interface FilterValues { - [key: string]: string; -} - -interface FilterComponentProps { - options: FilterOption[]; - onApplyFilters: (filters: FilterValues) => void; - initialValues?: FilterValues; - buttonLabel?: string; - onResetFilters: () => void; -} - -const FilterComponent: React.FC = ({ - options, - onApplyFilters, - onResetFilters, - initialValues = {}, - buttonLabel = "Filters", -}) => { - const [showFilters, setShowFilters] = useState(false); - const [tempValues, setTempValues] = useState(initialValues); - const [searchOptionsMap, setSearchOptionsMap] = useState<{ - [key: string]: Array<{ label: string; value: string }>; - }>({}); - const [searchLoadingMap, setSearchLoadingMap] = useState<{ - [key: string]: boolean; - }>({}); - const [searchInputValueMap, setSearchInputValueMap] = useState<{ - [key: string]: string; - }>({}); - const [initialOptionsLoaded, setInitialOptionsLoaded] = useState<{ - [key: string]: boolean; - }>({}); - - const debouncedSearch = useDebouncedCallback( - async (value: string, option: FilterOption) => { - if (!option.isSearchable || !option.searchFn) return; - - setSearchLoadingMap((prev) => ({ ...prev, [option.name]: true })); - try { - const results = await option.searchFn(value); - setSearchOptionsMap((prev) => ({ ...prev, [option.name]: results })); - } catch (error) { - console.error("Error searching:", error); - setSearchOptionsMap((prev) => ({ ...prev, [option.name]: [] })); - } finally { - setSearchLoadingMap((prev) => ({ ...prev, [option.name]: false })); - } - }, - { wait: DEBOUNCE_WAIT_MS }, - ); - - // Load initial options for searchable filters - const loadInitialOptions = useCallback( - async (option: FilterOption) => { - if (!option.isSearchable || !option.searchFn || option.loading || initialOptionsLoaded[option.name]) return; - - setSearchLoadingMap((prev) => ({ ...prev, [option.name]: true })); - setInitialOptionsLoaded((prev) => ({ ...prev, [option.name]: true })); - - try { - // Load initial options with empty search to get some default results - const results = await option.searchFn(""); - setSearchOptionsMap((prev) => ({ ...prev, [option.name]: results })); - } catch (error) { - console.error("Error loading initial options:", error); - setSearchOptionsMap((prev) => ({ ...prev, [option.name]: [] })); - } finally { - setSearchLoadingMap((prev) => ({ ...prev, [option.name]: false })); - } - }, - [initialOptionsLoaded], - ); - - // Load initial options when filters are shown - useEffect(() => { - if (showFilters) { - options.forEach((option) => { - if (option.isSearchable && !initialOptionsLoaded[option.name]) { - loadInitialOptions(option); - } - }); - } - }, [showFilters, options, loadInitialOptions, initialOptionsLoaded]); - - const handleFilterChange = (name: string, value: string) => { - const newValues = { - ...tempValues, - [name]: value, - }; - setTempValues(newValues); - onApplyFilters(newValues); - }; - - const resetFilters = () => { - const emptyValues: FilterValues = {}; - options.forEach((option) => { - emptyValues[option.name] = ""; - }); - setTempValues(emptyValues); - onResetFilters(); - }; - - // Handle dropdown open to load initial options - const handleDropdownVisibleChange = (open: boolean, option: FilterOption) => { - if (open && option.isSearchable && !initialOptionsLoaded[option.name]) { - loadInitialOptions(option); - } - }; - - return ( -
-
- - -
- - {showFilters && ( -
- {options.map((option) => { - const isOptionLoading = searchLoadingMap[option.name] || option.loading; - return ( -
- - {option.isSearchable ? ( - handleFilterChange(option.name, value)} - allowClear - > - {option.options.map((opt) => ( - - {opt.label} - - ))} - - ) : option.customComponent ? ( - (() => { - const CustomComponent = option.customComponent; - return ( - handleFilterChange(option.name, value ?? "")} - placeholder={`Select ${option.label || option.name}...`} - allFilters={tempValues} - /> - ); - })() - ) : ( - handleFilterChange(option.name, e.target.value)} - allowClear - /> - )} -
- ); - })} -
- )} -
- ); -}; - -export default FilterComponent; diff --git a/ui/litellm-dashboard/src/components/shared/PaginatedSearchSelect.test.tsx b/ui/litellm-dashboard/src/components/shared/PaginatedSearchSelect.test.tsx new file mode 100644 index 00000000000..cfabeeab362 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/PaginatedSearchSelect.test.tsx @@ -0,0 +1,165 @@ +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { useState } from "react"; +import { describe, expect, it, vi } from "vitest"; + +import { PaginatedSearchSelect } from "./PaginatedSearchSelect"; +import type { SearchSelectOption } from "./SearchSelect"; + +const OPTIONS: SearchSelectOption[] = [ + { label: "alias-alpha", value: "alias-alpha" }, + { label: "alias-beta", value: "alias-beta" }, + { label: "gamma-key", value: "gamma-key" }, +]; + +function renderSelect(overrides: Partial> = {}) { + const props: React.ComponentProps = { + options: OPTIONS, + onValueChange: vi.fn(), + onSearchChange: vi.fn(), + onLoadMore: vi.fn(), + ...overrides, + }; + render(); + return props; +} + +function setListMetrics(list: HTMLElement, metrics: { scrollTop: number; clientHeight: number; scrollHeight: number }) { + Object.defineProperty(list, "scrollTop", { value: metrics.scrollTop, configurable: true }); + Object.defineProperty(list, "clientHeight", { value: metrics.clientHeight, configurable: true }); + Object.defineProperty(list, "scrollHeight", { value: metrics.scrollHeight, configurable: true }); +} + +describe("PaginatedSearchSelect", () => { + it("reports the typed query to the server instead of filtering locally", async () => { + const user = userEvent.setup(); + const onSearchChange = vi.fn(); + renderSelect({ onSearchChange }); + + const input = screen.getByRole("combobox"); + await user.click(input); + await user.type(input, "gamma"); + + await waitFor(() => expect(onSearchChange).toHaveBeenCalledWith("gamma")); + + expect(await screen.findByText("alias-alpha")).toBeInTheDocument(); + }); + + it("does not re-query the server when an item is selected", async () => { + const user = userEvent.setup(); + const onSearchChange = vi.fn(); + + function Controlled() { + const [value, setValue] = useState(""); + return ( + + ); + } + render(); + + await user.click(screen.getByRole("combobox")); + await user.click(await screen.findByText("alias-beta")); + + expect(screen.getByRole("combobox")).toHaveValue("alias-beta"); + await new Promise((resolve) => setTimeout(resolve, 400)); + expect(onSearchChange).not.toHaveBeenCalled(); + }); + + it("still reports a cleared input so the unfiltered page comes back", async () => { + const user = userEvent.setup(); + const onSearchChange = vi.fn(); + renderSelect({ onSearchChange, value: "alias-alpha" }); + + await user.click(document.querySelector('[data-slot="combobox-clear"]') as HTMLElement); + + await waitFor(() => expect(onSearchChange).toHaveBeenCalledWith("")); + }); + + it("requests the next page once the list is scrolled near the bottom", async () => { + const user = userEvent.setup(); + const onLoadMore = vi.fn(); + renderSelect({ onLoadMore, hasNextPage: true }); + + await user.click(screen.getByRole("combobox")); + const list = await screen.findByTestId("paginated-search-select-list"); + + setListMetrics(list, { scrollTop: 0, clientHeight: 100, scrollHeight: 1000 }); + fireEvent.scroll(list); + expect(onLoadMore).not.toHaveBeenCalled(); + + setListMetrics(list, { scrollTop: 850, clientHeight: 100, scrollHeight: 1000 }); + fireEvent.scroll(list); + expect(onLoadMore).toHaveBeenCalledTimes(1); + }); + + it("does not request more pages when there is no next page or one is already in flight", async () => { + const user = userEvent.setup(); + const onLoadMore = vi.fn(); + const { unmount } = render( + , + ); + await user.click(screen.getByRole("combobox")); + let list = await screen.findByTestId("paginated-search-select-list"); + setListMetrics(list, { scrollTop: 900, clientHeight: 100, scrollHeight: 1000 }); + fireEvent.scroll(list); + expect(onLoadMore).not.toHaveBeenCalled(); + unmount(); + + renderSelect({ onLoadMore, hasNextPage: true, isFetchingNextPage: true }); + await user.click(screen.getByRole("combobox")); + list = await screen.findByTestId("paginated-search-select-list"); + setListMetrics(list, { scrollTop: 900, clientHeight: 100, scrollHeight: 1000 }); + fireEvent.scroll(list); + expect(onLoadMore).not.toHaveBeenCalled(); + }); + + it("keeps showing a selected value that is absent from the current page of options", () => { + renderSelect({ options: [], value: "alias-from-an-earlier-page" }); + + expect(screen.getByRole("combobox")).toHaveValue("alias-from-an-earlier-page"); + }); + + it("reports the selected option's value", async () => { + const user = userEvent.setup(); + const onValueChange = vi.fn(); + renderSelect({ onValueChange }); + + await user.click(screen.getByRole("combobox")); + await user.click(await screen.findByText("alias-beta")); + + expect(onValueChange).toHaveBeenCalledWith("alias-beta"); + }); + + it("surfaces loading and fetching-more affordances", async () => { + const user = userEvent.setup(); + const { unmount } = render( + , + ); + await user.click(screen.getByRole("combobox")); + expect(await screen.findByText("Loading key aliases…")).toBeInTheDocument(); + unmount(); + + renderSelect({ isFetchingNextPage: true }); + await user.click(screen.getByRole("combobox")); + expect(await screen.findByTestId("paginated-search-select-loading-more")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/PaginatedSearchSelect.tsx b/ui/litellm-dashboard/src/components/shared/PaginatedSearchSelect.tsx new file mode 100644 index 00000000000..b59cd1263ea --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/PaginatedSearchSelect.tsx @@ -0,0 +1,119 @@ +"use client"; + +import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer"; +import { Loader2 } from "lucide-react"; +import { useMemo, type UIEvent } from "react"; + +import { + Combobox, + ComboboxContent, + ComboboxEmpty, + ComboboxInput, + ComboboxItem, + ComboboxList, +} from "@/components/ui/combobox"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; + +import type { SearchSelectOption } from "./SearchSelect"; + +const SCROLL_THRESHOLD = 0.8; + +const SEARCH_REASONS: ReadonlySet = new Set(["input-change", "input-clear", "clear-press"]); + +interface PaginatedSearchSelectProps { + options: SearchSelectOption[]; + value?: string; + onValueChange: (value: string) => void; + onSearchChange: (query: string) => void; + onLoadMore: () => void; + hasNextPage?: boolean; + isLoading?: boolean; + isFetchingNextPage?: boolean; + placeholder?: string; + emptyText?: string; + loadingText?: string; + disabled?: boolean; + className?: string; +} + +export function PaginatedSearchSelect({ + options, + value, + onValueChange, + onSearchChange, + onLoadMore, + hasNextPage = false, + isLoading = false, + isFetchingNextPage = false, + placeholder = "Search…", + emptyText = "No results", + loadingText = "Loading…", + disabled = false, + className, +}: PaginatedSearchSelectProps) { + const selected = useMemo(() => { + if (value === undefined || value === "") return null; + return options.find((option) => option.value === value) ?? { label: value, value }; + }, [options, value]); + + const items = useMemo(() => { + if (selected === null) return options; + if (options.some((option) => option.value === selected.value)) return options; + return [selected, ...options]; + }, [options, selected]); + + const debouncedSearch = useDebouncedCallback(onSearchChange, { wait: DEBOUNCE_WAIT_MS }); + + const handleInputValueChange = (next: string, reason: string) => { + if (!SEARCH_REASONS.has(reason)) return; + debouncedSearch(next); + }; + + const handleScroll = (event: UIEvent) => { + const target = event.currentTarget; + if (target.scrollHeight === 0) return; + const ratio = (target.scrollTop + target.clientHeight) / target.scrollHeight; + if (ratio >= SCROLL_THRESHOLD && hasNextPage && !isFetchingNextPage) { + onLoadMore(); + } + }; + + return ( + onValueChange(item?.value ?? "")} + onInputValueChange={(next, eventDetails) => handleInputValueChange(next, eventDetails.reason)} + isItemEqualToValue={(a: SearchSelectOption, b: SearchSelectOption) => a.value === b.value} + itemToStringLabel={(item: SearchSelectOption) => item.label} + filter={null} + disabled={disabled} + > + + + {isLoading ? loadingText : emptyText} + + {(item: SearchSelectOption) => ( + + + {item.label} + {item.sublabel != null && item.sublabel !== "" && ( + {item.sublabel} + )} + + + )} + + {isFetchingNextPage && ( +
+ +
+ )} +
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/view_logs/LogsTableToolbar.tsx b/ui/litellm-dashboard/src/components/view_logs/LogsTableToolbar.tsx index f65ff2cc6ca..a522e75a4af 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogsTableToolbar.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogsTableToolbar.tsx @@ -1,14 +1,18 @@ +"use client"; + import moment from "moment"; -import { useEffect, useRef, useState } from "react"; -import { SyncOutlined } from "@ant-design/icons"; -import { Button, Switch } from "antd"; +import { CalendarDays } from "lucide-react"; +import { useState } from "react"; + +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover"; +import { Switch } from "@/components/ui/switch"; + import { QUICK_SELECT_OPTIONS } from "./constants"; import { getTimeRangeDisplay } from "./logs_utils"; -import type { PaginatedResponse } from "./log_filter_logic"; interface LogsTableToolbarProps { - searchTerm: string; - onSearchChange: (value: string) => void; startTime: string; onStartTimeChange: (value: string) => void; endTime: string; @@ -19,18 +23,11 @@ interface LogsTableToolbarProps { onSelectedTimeIntervalChange: (value: { value: number; unit: string }) => void; isLiveTail: boolean; onIsLiveTailChange: (value: boolean) => void; - currentPage: number; - onCurrentPageChange: (updater: number | ((prev: number) => number)) => void; - pageSize: number; - isLoading: boolean; - isButtonLoading: boolean; - onRefetch: () => void; - filteredLogs: PaginatedResponse; + onResetToFirstPage: () => void; + onResetFilters: () => void; } export function LogsTableToolbar({ - searchTerm, - onSearchChange, startTime, onStartTimeChange, endTime, @@ -41,26 +38,23 @@ export function LogsTableToolbar({ onSelectedTimeIntervalChange, isLiveTail, onIsLiveTailChange, - currentPage, - onCurrentPageChange, - pageSize, - isLoading, - isButtonLoading, - onRefetch, - filteredLogs, + onResetToFirstPage, + onResetFilters, }: LogsTableToolbarProps) { const [quickSelectOpen, setQuickSelectOpen] = useState(false); - const quickSelectRef = useRef(null); - useEffect(() => { - function handleClickOutside(event: MouseEvent) { - if (quickSelectRef.current && !quickSelectRef.current.contains(event.target as Node)) { - setQuickSelectOpen(false); - } - } - document.addEventListener("mousedown", handleClickOutside); - return () => document.removeEventListener("mousedown", handleClickOutside); - }, []); + const applyQuickSelect = (option: { label: string; value: number; unit: string }) => { + onResetToFirstPage(); + onEndTimeChange(moment().format("YYYY-MM-DDTHH:mm")); + onStartTimeChange( + moment() + .subtract(option.value, option.unit as moment.unitOfTime.DurationConstructor) + .format("YYYY-MM-DDTHH:mm"), + ); + onSelectedTimeIntervalChange({ value: option.value, unit: option.unit }); + onIsCustomDateChange(false); + setQuickSelectOpen(false); + }; const selectedOption = QUICK_SELECT_OPTIONS.find( (option) => option.value === selectedTimeInterval.value && option.unit === selectedTimeInterval.unit, @@ -68,178 +62,83 @@ export function LogsTableToolbar({ const displayLabel = isCustomDate ? getTimeRangeDisplay(isCustomDate, startTime, endTime) : selectedOption?.label; return ( - <> -
-
-
-
- onSearchChange(e.target.value)} - /> - - - -
- -
-
- - - {quickSelectOpen && ( -
-
- {QUICK_SELECT_OPTIONS.map((option) => ( - - ))} -
- -
-
- )} -
- -
- Live Tail - -
- +
+ + + + {displayLabel} + + } + /> + +
+ {QUICK_SELECT_OPTIONS.map((option) => ( -
- - {isCustomDate && ( -
-
- { - onStartTimeChange(e.target.value); - onCurrentPageChange(1); - }} - className="px-3 py-2 border rounded-md text-sm focus:outline-hidden focus:ring-2 focus:ring-blue-500 focus:border-blue-500" - /> -
- to -
- { - onEndTimeChange(e.target.value); - onCurrentPageChange(1); - }} - className="px-3 py-2 border rounded-md text-sm focus:outline-hidden focus:ring-2 focus:ring-blue-500 focus:border-blue-500" - /> -
-
- )} -
- -
- + - -
+ Custom Range +
-
-
- {isLiveTail && currentPage === 1 && ( -
-
- Auto-refreshing every 15 seconds -
- + + + + {isCustomDate && ( +
+ { + onStartTimeChange(event.target.value); + onResetToFirstPage(); + }} + /> + to + { + onEndTimeChange(event.target.value); + onResetToFirstPage(); + }} + />
)} - + +
+ Live Tail + +
+ + +
+ ); +} + +export function LiveTailBanner({ onStop }: { onStop: () => void }) { + return ( +
+ Auto-refreshing every 15 seconds + +
); } diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx new file mode 100644 index 00000000000..0e6e60ad05d --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx @@ -0,0 +1,91 @@ +import { screen, waitFor } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; +import { LOG_FILTER_IDS } from "./log_filter_logic"; +import { RequestLogsFilters } from "./RequestLogsFilters"; + +vi.mock("@/app/(dashboard)/hooks/keys/useKeyAliases", () => ({ + useInfiniteKeyAliases: vi.fn(), +})); + +vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ + useInfiniteModelInfo: vi.fn(), +})); + +vi.mock("../networking", async (importOriginal) => { + const actual = await importOriginal(); + return { ...actual, allEndUsersCall: vi.fn().mockResolvedValue([]) }; +}); + +import { useInfiniteKeyAliases } from "@/app/(dashboard)/hooks/keys/useKeyAliases"; +import { useInfiniteModelInfo } from "@/app/(dashboard)/hooks/models/useModels"; + +const emptyInfiniteQuery = { + data: { pages: [], pageParams: [] }, + fetchNextPage: vi.fn(), + hasNextPage: false, + isFetchingNextPage: false, + isLoading: false, +}; + +function renderFilters(filters: Record = {}) { + const set = vi.fn(); + renderWithProviders( + filters[id]} set={set} teams={[]} accessToken="test-token" />, + ); + return { set }; +} + +describe("RequestLogsFilters", () => { + beforeEach(() => { + vi.clearAllMocks(); + testQueryClient.clear(); + vi.mocked(useInfiniteKeyAliases).mockReturnValue( + emptyInfiniteQuery as unknown as ReturnType, + ); + vi.mocked(useInfiniteModelInfo).mockReturnValue( + emptyInfiniteQuery as unknown as ReturnType, + ); + }); + + it("renders every backend-supported filter field", async () => { + renderFilters(); + + for (const label of [ + "Team ID", + "Status", + "Key Alias", + "End User", + "Error Code", + "Error Message", + "Key Hash", + "Session ID", + "Model", + "Public model / search tool", + ]) { + expect(await screen.findByText(label)).toBeInTheDocument(); + } + }); + + it("scopes the Key Alias lookup to the selected team", async () => { + renderFilters({ [LOG_FILTER_IDS.TEAM_ID]: "team-42" }); + + await waitFor(() => expect(useInfiniteKeyAliases).toHaveBeenCalled()); + expect(useInfiniteKeyAliases).toHaveBeenCalledWith(50, undefined, "team-42"); + }); + + it("leaves the Key Alias lookup unscoped when no team is selected", async () => { + renderFilters(); + + await waitFor(() => expect(useInfiniteKeyAliases).toHaveBeenCalled()); + expect(useInfiniteKeyAliases).toHaveBeenCalledWith(50, undefined, undefined); + }); + + it("does not leak the team scope into the Model lookup", async () => { + renderFilters({ [LOG_FILTER_IDS.TEAM_ID]: "team-42" }); + + await waitFor(() => expect(useInfiniteModelInfo).toHaveBeenCalled()); + expect(useInfiniteModelInfo).toHaveBeenCalledWith(50, undefined); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.tsx new file mode 100644 index 00000000000..ec20f4d0e47 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.tsx @@ -0,0 +1,319 @@ +"use client"; + +import { useQuery } from "@tanstack/react-query"; +import { useMemo, useState } from "react"; + +import { useInfiniteKeyAliases } from "@/app/(dashboard)/hooks/keys/useKeyAliases"; +import { useInfiniteModelInfo } from "@/app/(dashboard)/hooks/models/useModels"; +import { DataTableFilterField } from "@/components/shared/DataTable"; +import { PaginatedSearchSelect } from "@/components/shared/PaginatedSearchSelect"; +import { SearchSelect, type SearchSelectOption } from "@/components/shared/SearchSelect"; +import { + Combobox, + ComboboxContent, + ComboboxEmpty, + ComboboxInput, + ComboboxItem, + ComboboxList, +} from "@/components/ui/combobox"; +import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; + +import type { Team } from "../key_team_helpers/key_list"; +import { allEndUsersCall } from "../networking"; +import { ERROR_CODE_OPTIONS } from "./constants"; +import { LOG_FILTER_IDS } from "./log_filter_logic"; + +const ALL_VALUE = "all"; +const PAGE_SIZE = 50; + +const asString = (value: unknown): string => (typeof value === "string" ? value : ""); +const emptyToUndefined = (value: string): string | undefined => (value === "" ? undefined : value); + +function TeamFilterField({ + value, + onChange, + teams, +}: { + value: string; + onChange: (value: string | undefined) => void; + teams: Team[]; +}) { + const options = useMemo( + () => + teams.map((team) => ({ + label: team.team_alias || team.team_id, + value: team.team_id, + sublabel: team.team_id, + })), + [teams], + ); + + return ( + + onChange(emptyToUndefined(next))} + placeholder="Search or select a team" + emptyText="No teams found" + /> + + ); +} + +function KeyAliasFilterField({ + value, + onChange, + teamId, +}: { + value: string; + onChange: (value: string | undefined) => void; + teamId: string; +}) { + const [search, setSearch] = useState(""); + const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteKeyAliases( + PAGE_SIZE, + emptyToUndefined(search), + emptyToUndefined(teamId), + ); + + const options = useMemo(() => { + const seen = new Set(); + return (data?.pages ?? []).flatMap((page) => + page.aliases.flatMap((alias) => { + if (!alias || seen.has(alias)) return []; + seen.add(alias); + return [{ label: alias, value: alias }]; + }), + ); + }, [data]); + + return ( + + onChange(emptyToUndefined(next))} + onSearchChange={setSearch} + onLoadMore={() => void fetchNextPage()} + hasNextPage={hasNextPage} + isLoading={isLoading} + isFetchingNextPage={isFetchingNextPage} + placeholder="Search a key alias" + emptyText="No key aliases found" + /> + + ); +} + +function ModelFilterField({ value, onChange }: { value: string; onChange: (value: string | undefined) => void }) { + const [search, setSearch] = useState(""); + const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteModelInfo( + PAGE_SIZE, + emptyToUndefined(search), + ); + + const options = useMemo(() => { + const seen = new Set(); + return (data?.pages ?? []).flatMap((page) => + page.data.flatMap((model) => { + const modelId = model.model_info?.id ?? ""; + const modelName = model.model_name ?? ""; + if (!modelId || seen.has(modelId)) return []; + seen.add(modelId); + return [{ label: modelName || modelId, value: modelId, sublabel: `Model ID: ${modelId}` }]; + }), + ); + }, [data]); + + return ( + + onChange(emptyToUndefined(next))} + onSearchChange={setSearch} + onLoadMore={() => void fetchNextPage()} + hasNextPage={hasNextPage} + isLoading={isLoading} + isFetchingNextPage={isFetchingNextPage} + placeholder="Search a model" + emptyText="No models found" + /> + + ); +} + +function EndUserFilterField({ + value, + onChange, + accessToken, +}: { + value: string; + onChange: (value: string | undefined) => void; + accessToken: string; +}) { + const { data } = useQuery({ + queryKey: ["logFilterEndUsers", accessToken], + queryFn: async () => { + const endUsers = await allEndUsersCall(accessToken); + return (endUsers ?? []).flatMap((endUser: { user_id?: string }) => + typeof endUser.user_id === "string" ? [endUser.user_id] : [], + ); + }, + enabled: accessToken !== "", + }); + + const options = useMemo( + () => (data ?? []).map((userId) => ({ label: userId, value: userId })), + [data], + ); + + return ( + + onChange(emptyToUndefined(next))} + placeholder="Search an end user" + emptyText="No end users found" + /> + + ); +} + +function ErrorCodeFilterField({ value, onChange }: { value: string; onChange: (value: string | undefined) => void }) { + const [query, setQuery] = useState(""); + + const options = useMemo(() => { + const trimmed = query.trim(); + const lowered = trimmed.toLowerCase(); + const matches = ERROR_CODE_OPTIONS.filter((option) => option.label.toLowerCase().includes(lowered)); + if (trimmed === "" || ERROR_CODE_OPTIONS.some((option) => option.value === trimmed)) return matches; + return [...matches, { label: `Use custom code: ${trimmed}`, value: trimmed }]; + }, [query]); + + const selected = useMemo(() => { + if (value === "") return null; + return ERROR_CODE_OPTIONS.find((option) => option.value === value) ?? { label: value, value }; + }, [value]); + + const items = useMemo(() => { + if (selected === null) return options; + if (options.some((option) => option.value === selected.value)) return options; + return [selected, ...options]; + }, [options, selected]); + + return ( + + onChange(emptyToUndefined(item?.value ?? ""))} + onInputValueChange={setQuery} + isItemEqualToValue={(a: SearchSelectOption, b: SearchSelectOption) => a.value === b.value} + itemToStringLabel={(item: SearchSelectOption) => item.label} + filter={null} + > + + + No error codes found + + {(item: SearchSelectOption) => ( + + {item.label} + + )} + + + + + ); +} + +interface RequestLogsFiltersProps { + get: (columnId: string) => unknown; + set: (columnId: string, value: unknown) => void; + teams: Team[]; + accessToken: string; +} + +export function RequestLogsFilters({ get, set, teams, accessToken }: RequestLogsFiltersProps) { + const valueOf = (id: string): string => asString(get(id)); + const setter = (id: string) => (next: string | undefined) => set(id, next); + + return ( + <> + + + + + + + + + + + + + + set(LOG_FILTER_IDS.ERROR_MESSAGE, emptyToUndefined(event.target.value))} + placeholder="Enter error message…" + /> + + + + set(LOG_FILTER_IDS.KEY_HASH, emptyToUndefined(event.target.value))} + placeholder="Enter key hash…" + /> + + + + set(LOG_FILTER_IDS.SESSION_ID, emptyToUndefined(event.target.value))} + placeholder="Enter session ID…" + /> + + + + + + set(LOG_FILTER_IDS.PUBLIC_MODEL_OR_SEARCH_TOOL, emptyToUndefined(event.target.value))} + placeholder="Enter public model or search tool…" + /> + + + ); +} diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx new file mode 100644 index 00000000000..9469c273d20 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx @@ -0,0 +1,202 @@ +import { screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import moment from "moment"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; +import type { LogEntry } from "./columns"; +import RequestLogsPanel from "./RequestLogsPanel"; + +vi.mock("../networking", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + uiSpendLogsCall: vi.fn(), + keyInfoV1Call: vi.fn().mockResolvedValue({ info: {} }), + allEndUsersCall: vi.fn().mockResolvedValue([]), + }; +}); + +vi.mock("@/components/key_team_helpers/filter_helpers", () => ({ + fetchAllTeams: vi.fn().mockResolvedValue([]), +})); + +vi.mock("./LogDetailsDrawer", () => ({ + LogDetailsDrawer: function LogDetailsDrawerMock({ open }: { open: boolean }) { + return
{open ? "open" : "closed"}
; + }, +})); + +import { uiSpendLogsCall } from "../networking"; + +const logEntry = (overrides: Partial): LogEntry => ({ + request_id: "req-1", + api_key: "key-1", + team_id: "team-1", + model: "gpt-4o", + model_id: "model-1", + call_type: "acompletion", + spend: 0.01, + total_tokens: 10, + prompt_tokens: 5, + completion_tokens: 5, + startTime: "2026-07-07T09:50:13Z", + endTime: "2026-07-07T09:50:14Z", + cache_hit: "false", + messages: [], + response: {}, + ...overrides, +}); + +const respondWith = (data: LogEntry[]) => + vi.mocked(uiSpendLogsCall).mockResolvedValue({ + data, + total: data.length, + page: 1, + page_size: 50, + total_pages: 1, + }); + +const defaultProps = { + accessToken: "test-token", + token: "test-token", + userRole: "Admin", + userID: "user-1", + isActive: true, +}; + +const row = (requestId: string) => document.querySelector(`[data-row-id="${requestId}"]`); +const lastCall = () => vi.mocked(uiSpendLogsCall).mock.calls.at(-1)?.[0]; + +describe("RequestLogsPanel", () => { + beforeEach(() => { + vi.clearAllMocks(); + sessionStorage.clear(); + testQueryClient.clear(); + respondWith([]); + }); + + describe("multi-call session collapsing", () => { + const sessionRows = [ + logEntry({ request_id: "req-mcp", call_type: "call_mcp_tool", session_id: "sess-1", session_total_count: 3 }), + logEntry({ request_id: "req-llm", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }), + logEntry({ request_id: "req-llm-2", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }), + ]; + + it("collapses a multi-call session to a single representative row", async () => { + respondWith(sessionRows); + renderWithProviders(); + + await waitFor(() => expect(row("req-mcp") ?? row("req-llm") ?? row("req-llm-2")).not.toBeNull()); + + const rendered = ["req-mcp", "req-llm", "req-llm-2"].filter((id) => row(id) !== null); + expect(rendered).toHaveLength(1); + }); + + it("prefers an LLM call over an MCP call as the session's representative", async () => { + respondWith(sessionRows); + renderWithProviders(); + + await waitFor(() => expect(row("req-llm")).not.toBeNull()); + expect(row("req-mcp")).toBeNull(); + }); + + it("shows the session's call count and composition on the representative row", async () => { + respondWith(sessionRows); + renderWithProviders(); + + await waitFor(() => expect(row("req-llm")).not.toBeNull()); + expect(within(row("req-llm") as HTMLElement).getByText("3")).toBeInTheDocument(); + }); + + it("leaves single-call rows untouched", async () => { + respondWith([ + logEntry({ request_id: "req-solo-a", session_id: "sess-a", session_total_count: 1 }), + logEntry({ request_id: "req-solo-b" }), + ]); + renderWithProviders(); + + await waitFor(() => expect(row("req-solo-a")).not.toBeNull()); + expect(row("req-solo-b")).not.toBeNull(); + }); + }); + + describe("client-side search", () => { + it("narrows the visible rows without refetching", async () => { + const user = userEvent.setup(); + respondWith([ + logEntry({ request_id: "req-alpha", model: "gpt-4o" }), + logEntry({ request_id: "req-beta", model: "claude-opus" }), + ]); + renderWithProviders(); + + await waitFor(() => expect(row("req-alpha")).not.toBeNull()); + const callsBefore = vi.mocked(uiSpendLogsCall).mock.calls.length; + + await user.type(screen.getByTestId("datatable-search"), "alpha"); + + await waitFor(() => expect(row("req-beta")).toBeNull()); + expect(row("req-alpha")).not.toBeNull(); + expect(vi.mocked(uiSpendLogsCall).mock.calls.length).toBe(callsBefore); + }); + }); + + describe("time range", () => { + it("requests a ~15 minute window when Last 15 Minutes is picked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); + await user.click(await screen.findByRole("button", { name: "Last 15 Minutes" })); + + await waitFor(() => { + const call = lastCall(); + if (!call) throw new Error("no call"); + const diff = moment + .utc(call.end_date, "YYYY-MM-DD HH:mm:ss") + .diff(moment.utc(call.start_date, "YYYY-MM-DD HH:mm:ss"), "seconds"); + expect(diff).toBeGreaterThanOrEqual(15 * 60); + expect(diff).toBeLessThanOrEqual(16 * 60); + }); + }); + + it("restores the default 24 hour window when filters are reset", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); + await user.click(await screen.findByRole("button", { name: "Last 15 Minutes" })); + await waitFor(() => expect(screen.getByRole("button", { name: /Last 15 Minutes/i })).toBeInTheDocument()); + + await user.click(screen.getByRole("button", { name: "Reset Filters" })); + + expect(await screen.findByRole("button", { name: /Last 24 Hours/i })).toBeInTheDocument(); + + await user.click(screen.getByTestId("datatable-refresh")); + + await waitFor(() => { + const call = lastCall(); + if (!call) throw new Error("no call"); + const diff = moment + .utc(call.end_date, "YYYY-MM-DD HH:mm:ss") + .diff(moment.utc(call.start_date, "YYYY-MM-DD HH:mm:ss"), "seconds"); + expect(diff).toBeGreaterThanOrEqual(24 * 60 * 60); + }); + }); + }); + + describe("live tail", () => { + it("shows the auto-refresh banner on the first page and hides it once stopped", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + expect(await screen.findByText("Auto-refreshing every 15 seconds")).toBeInTheDocument(); + + await user.click(screen.getByRole("button", { name: "Stop" })); + + expect(screen.queryByText("Auto-refreshing every 15 seconds")).toBeNull(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx new file mode 100644 index 00000000000..05044dda791 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx @@ -0,0 +1,264 @@ +"use client"; + +import { useQuery, type UseQueryOptions } from "@tanstack/react-query"; +import type { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; +import moment from "moment"; +import { useCallback, useEffect, useMemo, useState } from "react"; + +import { internalUserRoles } from "../../utils/roles"; +import type { KeyResponse } from "../key_team_helpers/key_list"; +import { keyInfoV1Call } from "../networking"; +import KeyInfoView from "../templates/key_info_view"; +import type { LogEntry } from "./columns"; +import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants"; +import { DEFAULT_LOGS_SORTING, useLogFilterLogic } from "./log_filter_logic"; +import { LogDetailsDrawer } from "./LogDetailsDrawer"; +import { LiveTailBanner, LogsTableToolbar } from "./LogsTableToolbar"; +import { RequestLogsTable } from "./RequestLogsTable"; + +const PAGE_SIZE = 50; +const DEFAULT_INTERVAL = { value: 24, unit: "hours" }; + +interface RequestLogsPanelProps { + accessToken: string; + token: string; + userRole: string; + userID: string; + isActive: boolean; +} + +interface SessionComposition { + llm: number; + agent: number; + mcp: number; +} + +export default function RequestLogsPanel({ accessToken, token, userRole, userID, isActive }: RequestLogsPanelProps) { + const [searchTerm, setSearchTerm] = useState(""); + const [pagination, setPagination] = useState({ pageIndex: 0, pageSize: PAGE_SIZE }); + const [sorting, setSorting] = useState(DEFAULT_LOGS_SORTING); + const [columnFilters, setColumnFilters] = useState([]); + + const [startTime, setStartTime] = useState(moment().subtract(24, "hours").format("YYYY-MM-DDTHH:mm")); + const [endTime, setEndTime] = useState(moment().format("YYYY-MM-DDTHH:mm")); + const [isCustomDate, setIsCustomDate] = useState(false); + const [selectedTimeInterval, setSelectedTimeInterval] = useState<{ value: number; unit: string }>(DEFAULT_INTERVAL); + + const [selectedKeyIdInfoView, setSelectedKeyIdInfoView] = useState(null); + const [selectedLog, setSelectedLog] = useState(null); + const [isDrawerOpen, setIsDrawerOpen] = useState(false); + const [selectedSessionId, setSelectedSessionId] = useState(null); + + const [isLiveTail, setIsLiveTail] = useState(() => { + const storedValue = sessionStorage.getItem("isLiveTail"); + return storedValue !== null ? JSON.parse(storedValue) : true; + }); + + useEffect(() => { + sessionStorage.setItem("isLiveTail", JSON.stringify(isLiveTail)); + }, [isLiveTail]); + + const filterByCurrentUser = internalUserRoles.includes(userRole); + + const { logsQuery, filteredLogs, allTeams } = useLogFilterLogic({ + accessToken, + token, + userRole, + userID, + columnFilters, + filterByCurrentUser, + activeTab: isActive ? "request logs" : "inactive", + isLiveTail, + startTime, + endTime, + pagination, + isCustomDate, + sorting, + }); + + const keyInfoQueryOptions: UseQueryOptions = { + queryKey: ["requestLogsKeyInfo", selectedKeyIdInfoView, accessToken], + queryFn: async () => { + if (selectedKeyIdInfoView === null) return null; + const keyData = await keyInfoV1Call(accessToken, selectedKeyIdInfoView); + return { + ...keyData["info"], + token: selectedKeyIdInfoView, + api_key: selectedKeyIdInfoView, + }; + }, + enabled: selectedKeyIdInfoView !== null, + }; + + const { data: selectedKeyInfo } = useQuery(keyInfoQueryOptions); + + const rows = useMemo(() => { + const searchedLogs = filteredLogs.data.filter((log) => { + if (!searchTerm) return true; + return ( + log.request_id.includes(searchTerm) || + log.model.includes(searchTerm) || + (log.user !== undefined && log.user.includes(searchTerm)) + ); + }); + + const sessionCompositionById = searchedLogs.reduce>((acc, log) => { + if (!log.session_id) return acc; + if (!acc[log.session_id]) { + acc[log.session_id] = { llm: 0, agent: 0, mcp: 0 }; + } + if (MCP_CALL_TYPES.includes(log.call_type)) { + acc[log.session_id].mcp += 1; + } else if (AGENT_CALL_TYPES.includes(log.call_type)) { + acc[log.session_id].agent += 1; + } else { + acc[log.session_id].llm += 1; + } + return acc; + }, {}); + + const sessionRepresentativeMap = new Map(); + for (const log of searchedLogs) { + if (!log.session_id || (log.session_total_count || 1) <= 1) continue; + const isMcp = MCP_CALL_TYPES.includes(log.call_type); + const existing = sessionRepresentativeMap.get(log.session_id); + if (!existing || (existing.isMcp && !isMcp)) { + sessionRepresentativeMap.set(log.session_id, { requestId: log.request_id, isMcp }); + } + } + + return searchedLogs + .map((log) => { + const sessionComposition = log.session_id ? sessionCompositionById[log.session_id] : undefined; + return { + ...log, + session_llm_count: sessionComposition?.llm ?? undefined, + session_mcp_count: sessionComposition?.mcp ?? undefined, + session_agent_count: sessionComposition?.agent ?? undefined, + }; + }) + .filter((log) => { + if (!log.session_id || (log.session_total_count || 1) <= 1) return true; + return sessionRepresentativeMap.get(log.session_id)?.requestId === log.request_id; + }); + }, [filteredLogs.data, searchTerm]); + + const handleSortingChange = useCallback>((updaterOrValue) => { + setSorting(updaterOrValue); + setPagination((previous) => ({ ...previous, pageIndex: 0 })); + }, []); + + const handleColumnFiltersChange = useCallback>((updaterOrValue) => { + setColumnFilters(updaterOrValue); + setPagination((previous) => ({ ...previous, pageIndex: 0 })); + }, []); + + const resetToFirstPage = useCallback(() => { + setPagination((previous) => ({ ...previous, pageIndex: 0 })); + }, []); + + const handleResetFilters = useCallback(() => { + setColumnFilters([]); + setSearchTerm(""); + setStartTime(moment().subtract(24, "hours").format("YYYY-MM-DDTHH:mm")); + setEndTime(moment().format("YYYY-MM-DDTHH:mm")); + setIsCustomDate(false); + setSelectedTimeInterval(DEFAULT_INTERVAL); + resetToFirstPage(); + }, [resetToFirstPage]); + + const handleRowClick = useCallback((log: LogEntry) => { + const isMultiCallSession = log.session_id !== undefined && (log.session_total_count || 1) > 1; + setSelectedSessionId(isMultiCallSession ? log.session_id ?? null : null); + setSelectedLog(log); + setIsDrawerOpen(true); + }, []); + + const handleSessionClick = useCallback( + (sessionId: string) => { + if (!sessionId) return; + const log = rows.find((candidate) => candidate.session_id === sessionId) ?? null; + setSelectedSessionId(sessionId); + setSelectedLog(log); + setIsDrawerOpen(true); + }, + [rows], + ); + + const handleKeyHashClick = useCallback((keyHash: string) => { + setSelectedKeyIdInfoView(keyHash); + }, []); + + if (selectedKeyInfo && selectedKeyIdInfoView && selectedKeyInfo.api_key === selectedKeyIdInfoView) { + return ( + setSelectedKeyIdInfoView(null)} + backButtonText="Back to Logs" + /> + ); + } + + return ( + <> +
+

Request Logs

+
+ + {isLiveTail && pagination.pageIndex === 0 && setIsLiveTail(false)} />} + + void logsQuery.refetch()} + onRowClick={handleRowClick} + onKeyHashClick={handleKeyHashClick} + onSessionClick={handleSessionClick} + teams={allTeams ?? []} + accessToken={accessToken} + toolbarChildren={ + + } + /> + + { + setIsDrawerOpen(false); + setSelectedSessionId(null); + }} + logEntry={selectedLog} + sessionId={selectedSessionId} + accessToken={accessToken} + allLogs={rows} + onSelectLog={setSelectedLog} + startTime={moment(startTime).utc().format("YYYY-MM-DD HH:mm:ss")} + /> + + ); +} diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx new file mode 100644 index 00000000000..2b2b8448604 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx @@ -0,0 +1,131 @@ +"use client"; + +import type { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; +import { ScrollText } from "lucide-react"; +import { useMemo, useState, type ReactNode } from "react"; + +import { DataTable, DataTableFilterDrawer, DataTableToolbar } from "@/components/shared/DataTable"; + +import type { Team } from "../key_team_helpers/key_list"; +import type { LogEntry } from "./columns"; +import { LOG_FILTER_LABELS } from "./log_filter_logic"; +import { RequestLogsFilters } from "./RequestLogsFilters"; +import { getRequestLogsTableColumns } from "./RequestLogsTableColumns"; + +interface RequestLogsTableProps { + data: LogEntry[]; + rowCount: number; + isLoading: boolean; + isRefreshing: boolean; + pagination: PaginationState; + onPaginationChange: OnChangeFn; + sorting: SortingState; + onSortingChange: OnChangeFn; + columnFilters: ColumnFiltersState; + onColumnFiltersChange: OnChangeFn; + searchValue: string; + onSearchChange: (value: string) => void; + onRefresh: () => void; + onRowClick: (log: LogEntry) => void; + onKeyHashClick: (keyHash: string) => void; + onSessionClick: (sessionId: string) => void; + teams: Team[]; + accessToken: string; + toolbarChildren?: ReactNode; +} + +function RequestLogsEmptyState({ filtered }: { filtered: boolean }) { + return ( +
+
+ +
+
{filtered ? "No matching requests" : "No requests yet"}
+
+ {filtered + ? "No requests match your filters for this time range." + : "Requests proxied through LiteLLM will appear here."} +
+
+ ); +} + +export function RequestLogsTable({ + data, + rowCount, + isLoading, + isRefreshing, + pagination, + onPaginationChange, + sorting, + onSortingChange, + columnFilters, + onColumnFiltersChange, + searchValue, + onSearchChange, + onRefresh, + onRowClick, + onKeyHashClick, + onSessionClick, + teams, + accessToken, + toolbarChildren, +}: RequestLogsTableProps) { + const [filtersOpen, setFiltersOpen] = useState(false); + + const columns = useMemo(() => { + const deps = { onKeyHashClick, onSessionClick }; + return getRequestLogsTableColumns(deps); + }, [onKeyHashClick, onSessionClick]); + + const isFiltered = columnFilters.length > 0 || searchValue !== ""; + + return ( + row.request_id} + sortingMode="server" + sorting={sorting} + onSortingChange={onSortingChange} + paginationMode="server" + pagination={pagination} + onPaginationChange={onPaginationChange} + rowCount={rowCount} + filterMode="server" + columnFilters={columnFilters} + onColumnFiltersChange={onColumnFiltersChange} + isLoading={isLoading} + loadingMessage="Loading request logs…" + noDataMessage={} + size="compact" + onRowClick={onRowClick} + toolbar={(table) => ( + <> + setFiltersOpen(true)} + filterLabels={LOG_FILTER_LABELS} + showViewOptions={false} + > + {toolbarChildren} + + + {({ get, set }) => } + + + )} + /> + ); +} diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx new file mode 100644 index 00000000000..2fdc8455ca1 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx @@ -0,0 +1,136 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; + +import { DataTable } from "@/components/shared/DataTable"; + +import type { LogEntry } from "./columns"; +import { getRequestLogsTableColumns } from "./RequestLogsTableColumns"; + +const logEntry = (overrides: Partial): LogEntry => ({ + request_id: "req-1", + api_key: "key-1", + team_id: "team-1", + model: "gpt-4o", + model_id: "model-1", + call_type: "acompletion", + spend: 0, + total_tokens: 10, + prompt_tokens: 5, + completion_tokens: 5, + startTime: "2026-07-07T09:50:13Z", + endTime: "2026-07-07T09:50:14Z", + cache_hit: "false", + messages: [], + response: {}, + ...overrides, +}); + +const noopDeps = { onKeyHashClick: vi.fn(), onSessionClick: vi.fn() }; + +function renderRows(rows: LogEntry[], deps = noopDeps) { + render( + row.request_id} + size="compact" + />, + ); +} + +describe("Cost column", () => { + it("renders '-' for zero spend with no tooltip, so hovering never shows a contradictory $0", async () => { + const user = userEvent.setup(); + renderRows([logEntry({ request_id: "req-zero", spend: 0 })]); + + for (const dash of screen.getAllByText("-")) { + await user.hover(dash); + } + expect(screen.queryByText("$0")).not.toBeInTheDocument(); + }); + + it("shows the full-precision raw value in the tooltip for a real spend", async () => { + const user = userEvent.setup(); + renderRows([logEntry({ request_id: "req-spend", spend: 0.00012345678 })]); + + await user.hover(screen.getByText("$0.000123")); + expect(await screen.findByText("$0.00012345678")).toBeInTheDocument(); + }); + + it("shows the summed session total, not the representative call's spend, for a multi-round session", () => { + renderRows([ + logEntry({ + request_id: "req-session", + spend: 0.01, + session_id: "sess-1", + session_total_count: 3, + session_total_spend: 0.06, + }), + ]); + + expect(screen.getByText("$0.060000")).toBeInTheDocument(); + expect(screen.queryByText("$0.010000")).not.toBeInTheDocument(); + expect(screen.getByText("session total")).toBeInTheDocument(); + }); +}); + +describe("row action cells", () => { + it("reports the key hash through the injected dependency rather than a row field", async () => { + const user = userEvent.setup(); + const deps = { onKeyHashClick: vi.fn(), onSessionClick: vi.fn() }; + renderRows([logEntry({ request_id: "req-key", metadata: { user_api_key: "sk-hash-9" } })], deps); + + await user.click(screen.getByText("sk-hash-9")); + expect(deps.onKeyHashClick).toHaveBeenCalledWith("sk-hash-9"); + }); + + it("reports the session id from the session cell", async () => { + const user = userEvent.setup(); + const deps = { onKeyHashClick: vi.fn(), onSessionClick: vi.fn() }; + renderRows([logEntry({ request_id: "req-sess", session_id: "sess-42" })], deps); + + await user.click(screen.getByText("sess-42")); + expect(deps.onSessionClick).toHaveBeenCalledWith("sess-42"); + }); +}); + +describe("sortable headers", () => { + it("exposes sort controls only for the backend-sortable fields", () => { + renderRows([logEntry({})]); + + for (const field of ["startTime", "spend", "request_duration_ms", "ttft_ms", "model", "total_tokens"]) { + expect(screen.getByTestId(`sort-trigger-${field}`)).toBeInTheDocument(); + } + for (const field of ["request_id", "session_id", "status", "type", "end_user"]) { + expect(screen.queryByTestId(`sort-trigger-${field}`)).toBeNull(); + } + }); +}); + +describe("TTFT column", () => { + it("renders '-' when the completion start equals the end time, since TTFT is meaningless there", () => { + renderRows([ + logEntry({ + request_id: "req-nonstream", + endTime: "2026-07-07T09:50:14Z", + completionStartTime: "2026-07-07T09:50:14Z", + }), + ]); + + expect(screen.queryByText("1.00")).not.toBeInTheDocument(); + }); + + it("renders seconds when streaming produced a real first token", () => { + renderRows([ + logEntry({ + request_id: "req-stream", + startTime: "2026-07-07T09:50:13Z", + endTime: "2026-07-07T09:50:16Z", + completionStartTime: "2026-07-07T09:50:14Z", + }), + ]); + + expect(screen.getByText("1.00")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx new file mode 100644 index 00000000000..4e5a83ac7dd --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx @@ -0,0 +1,317 @@ +"use client"; + +import type { ColumnDef } from "@tanstack/react-table"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { CellTooltip, DateCell, IdCell, MoneyCell, StatusBadge } from "@/components/shared/table_cells"; +import { getSpendString } from "@/utils/dataUtils"; + +import { getProviderLogoAndName } from "../provider_info_helpers"; +import type { LogEntry } from "./columns"; +import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants"; +import { AgentBadge, AgentIcon, LlmBadge, McpBadge, SparkleIcon, WrenchIcon } from "./TypeBadges"; + +export interface RequestLogsTableColumnsDeps { + onKeyHashClick: (keyHash: string) => void; + onSessionClick: (sessionId: string) => void; +} + +const readMetaString = (metadata: Record | undefined, key: string): string | undefined => { + const value = metadata?.[key]; + return typeof value === "string" && value !== "" ? value : undefined; +}; + +const readMcpLogoUrl = (metadata: Record | undefined): string | undefined => { + const mcpMetadata = metadata?.["mcp_tool_call_metadata"]; + if (typeof mcpMetadata !== "object" || mcpMetadata === null) return undefined; + const url = (mcpMetadata as Record)["mcp_server_logo_url"]; + return typeof url === "string" && url !== "" ? url : undefined; +}; + +const getLogoUrl = (row: LogEntry, provider: string): string => + readMcpLogoUrl(row.metadata) ?? (provider ? getProviderLogoAndName(provider).logo : ""); + +function TruncatedText({ value }: { value: string | undefined }) { + const display = value ?? "-"; + return {display}} />; +} + +export const getRequestLogsTableColumns = ({ + onKeyHashClick, + onSessionClick, +}: RequestLogsTableColumnsDeps): ColumnDef[] => [ + { + id: "startTime", + accessorKey: "startTime", + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "type", + header: "Type", + size: 90, + enableSorting: false, + meta: { skeleton: "badge" }, + cell: ({ row }) => { + const log = row.original; + const sessionCount = log.session_total_count || 1; + const isMcp = MCP_CALL_TYPES.includes(log.call_type); + const isAgent = AGENT_CALL_TYPES.includes(log.call_type); + const sessionLlmCount = log.session_llm_count ?? (isMcp || isAgent ? 0 : sessionCount); + const sessionAgentCount = log.session_agent_count ?? (isAgent ? sessionCount : 0); + const sessionMcpCount = log.session_mcp_count ?? (isMcp ? sessionCount : 0); + + if (isMcp) return ; + if (isAgent && sessionCount <= 1) return ; + if (sessionCount <= 1) return ; + + const sessionTypeBadge = ( + + + {sessionCount} + {sessionAgentCount > 0 && ( + <> + · + + + )} + {sessionMcpCount > 0 && ( + <> + · + + + )} + + ); + + const tooltipParts = [ + sessionLlmCount > 0 && `${sessionLlmCount} LLM`, + sessionAgentCount > 0 && `${sessionAgentCount} Agent`, + sessionMcpCount > 0 && `${sessionMcpCount} MCP`, + ].filter(Boolean); + return ; + }, + }, + { + id: "status", + header: "Status", + size: 100, + enableSorting: false, + meta: { skeleton: "badge" }, + cell: ({ row }) => { + const status = readMetaString(row.original.metadata, "status") ?? "Success"; + const isSuccess = status.toLowerCase() !== "failure"; + return ; + }, + }, + { + id: "session_id", + accessorKey: "session_id", + header: "Session ID", + size: 120, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "request_id", + accessorKey: "request_id", + header: "Request ID", + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "spend", + accessorKey: "spend", + header: ({ column }) => , + size: 110, + enableSorting: true, + meta: { numeric: true, skeleton: "twoLine" }, + cell: ({ row }) => { + const log = row.original; + const mcpCount = log.mcp_tool_call_count || 0; + const mcpSpend = log.mcp_tool_call_spend || 0; + const isMultiCallSession = (log.session_total_count || 1) > 1; + const spend = isMultiCallSession && log.session_total_spend != null ? log.session_total_spend : log.spend; + const money = ( + + + + ); + + return ( +
+ {spend ? : money} + {isMultiCallSession && session total} + {mcpCount > 0 && mcpSpend > 0 && ( + + incl. {getSpendString(mcpSpend)} from {mcpCount} MCP + + )} +
+ ); + }, + }, + { + id: "request_duration_ms", + accessorKey: "request_duration_ms", + header: ({ column }) => , + enableSorting: true, + meta: { numeric: true }, + cell: ({ row }) => { + const ms = row.original.request_duration_ms; + if (ms == null) return -; + return ( + {(ms / 1000).toFixed(2)}} + /> + ); + }, + }, + { + id: "ttft_ms", + accessorKey: "completionStartTime", + header: ({ column }) => , + enableSorting: true, + meta: { numeric: true }, + cell: ({ row }) => { + const log = row.original; + const completionStartTime = log.completionStartTime; + if (!completionStartTime) return -; + if (completionStartTime === log.endTime) return -; + const ttftMs = new Date(completionStartTime).getTime() - new Date(log.startTime).getTime(); + if (ttftMs <= 0) return -; + return ( + {(ttftMs / 1000).toFixed(2)}} + /> + ); + }, + }, + { + id: "team_alias", + header: "Team Name", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "key_hash", + header: "Key Hash", + size: 110, + enableSorting: false, + cell: ({ row }) => ( + + ), + }, + { + id: "key_alias", + header: "Key Alias", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "model", + accessorKey: "model", + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => { + const log = row.original; + const provider = log.custom_llm_provider; + const modelName = log.model ?? ""; + return ( +
+ {provider && ( + { + event.currentTarget.style.display = "none"; + }} + /> + )} + {modelName}} /> +
+ ); + }, + }, + { + id: "total_tokens", + accessorKey: "total_tokens", + header: ({ column }) => , + size: 140, + enableSorting: true, + meta: { numeric: true }, + cell: ({ row }) => { + const log = row.original; + return ( + + {String(log.total_tokens || "0")} + + ({String(log.prompt_tokens || "0")}+{String(log.completion_tokens || "0")}) + + + ); + }, + }, + { + id: "user", + accessorKey: "user", + header: "Internal User", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "end_user", + accessorKey: "end_user", + header: "End User", + size: 140, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "request_tags", + accessorKey: "request_tags", + header: "Tags", + size: 150, + enableSorting: false, + meta: { skeleton: "chips" }, + cell: ({ row }) => { + const tags = row.original.request_tags; + if (!tags || Object.keys(tags).length === 0) return "-"; + + const tagEntries = Object.entries(tags); + const [firstTagKey, firstTagValue] = tagEntries[0]; + const remainingCount = tagEntries.length - 1; + + return ( +
+ + {tagEntries.map(([key, value]) => ( + + {key}: {String(value)} + + ))} +
+ } + trigger={ + + {firstTagKey}: {String(firstTagValue)} + {remainingCount > 0 && ` +${remainingCount}`} + + } + /> +
+ ); + }, + }, +]; diff --git a/ui/litellm-dashboard/src/components/view_logs/columns.test.tsx b/ui/litellm-dashboard/src/components/view_logs/columns.test.tsx deleted file mode 100644 index c3d9bfdd0a5..00000000000 --- a/ui/litellm-dashboard/src/components/view_logs/columns.test.tsx +++ /dev/null @@ -1,70 +0,0 @@ -import { render, screen } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { describe, expect, it } from "vitest"; - -import { createColumns, type LogEntry } from "./columns"; -import { DataTable } from "./table"; - -const logEntry = (overrides: Partial): LogEntry => ({ - request_id: "req-1", - api_key: "key-1", - team_id: "team-1", - model: "gpt-4o", - model_id: "model-1", - call_type: "acompletion", - spend: 0, - total_tokens: 10, - prompt_tokens: 5, - completion_tokens: 5, - startTime: "2026-07-07T09:50:13Z", - endTime: "2026-07-07T09:50:14Z", - cache_hit: "false", - messages: [], - response: {}, - ...overrides, -}); - -describe("Cost column", () => { - it("renders '-' for zero spend with no tooltip, so hovering never shows a contradictory $0", async () => { - const user = userEvent.setup(); - render( - r.request_id} - />, - ); - for (const dash of screen.getAllByText("-")) { - await user.hover(dash); - } - expect(screen.queryByText("$0")).not.toBeInTheDocument(); - }); - - it("shows the full-precision raw value in the tooltip for a real spend", async () => { - const user = userEvent.setup(); - render( - r.request_id} - />, - ); - const formatted = screen.getByText("$0.000123"); - await user.hover(formatted); - expect(await screen.findByText("$0.00012345678")).toBeInTheDocument(); - }); - - it("shows the summed session total, not the representative call's spend, for a multi-round session", () => { - const overrides: Partial = { - request_id: "req-session", - spend: 0.01, - session_id: "sess-1", - session_total_count: 3, - session_total_spend: 0.06, - }; - render( r.request_id} />); - expect(screen.getByText("$0.060000")).toBeInTheDocument(); - expect(screen.queryByText("$0.010000")).not.toBeInTheDocument(); - expect(screen.getByText("session total")).toBeInTheDocument(); - }); -}); diff --git a/ui/litellm-dashboard/src/components/view_logs/columns.tsx b/ui/litellm-dashboard/src/components/view_logs/columns.tsx index d3ba90d4e2d..a4b8015892a 100644 --- a/ui/litellm-dashboard/src/components/view_logs/columns.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/columns.tsx @@ -1,13 +1,3 @@ -import { DateCell, IdCell, MoneyCell, StatusBadge } from "@/components/shared/table_cells"; -import { getSpendString } from "@/utils/dataUtils"; -import type { ColumnDef } from "@tanstack/react-table"; -import { Tooltip } from "antd"; -import React from "react"; -import { getProviderLogoAndName } from "../provider_info_helpers"; -import { TableHeaderSortDropdown } from "../common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; -import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants"; -import { AgentBadge, AgentIcon, LlmBadge, McpBadge, SparkleIcon, WrenchIcon } from "./TypeBadges"; - /** API sort field mapping for /spend/logs/ui endpoint */ export const LOGS_SORT_FIELD_MAP = { startTime: "startTime", @@ -20,22 +10,6 @@ export const LOGS_SORT_FIELD_MAP = { export type LogsSortField = keyof typeof LOGS_SORT_FIELD_MAP; -export interface LogsSortProps { - sortBy: LogsSortField; - sortOrder: "asc" | "desc"; - onSortChange: (sortBy: LogsSortField, sortOrder: "asc" | "desc") => void; -} - -// Helper to get the appropriate logo URL -const getLogoUrl = (row: LogEntry, provider: string) => { - // Check if mcp_tool_call_metadata exists and contains mcp_server_logo_url - if (row.metadata?.mcp_tool_call_metadata?.mcp_server_logo_url) { - return row.metadata.mcp_tool_call_metadata.mcp_server_logo_url; - } - // Fall back to default provider logo - return provider ? getProviderLogoAndName(provider).logo : ""; -}; - export type LogEntry = { request_id: string; api_key: string; @@ -72,470 +46,4 @@ export type LogEntry = { session_llm_count?: number; session_mcp_count?: number; session_agent_count?: number; - onKeyHashClick?: (keyHash: string) => void; - onSessionClick?: (sessionId: string) => void; -}; - -const SortableHeader = ({ - label, - field, - sortBy, - sortOrder, - onSortChange, -}: { - label: string; - field: LogsSortField; - sortBy: LogsSortField; - sortOrder: "asc" | "desc"; - onSortChange: (sortBy: LogsSortField, sortOrder: "asc" | "desc") => void; -}) => ( -
- {label} - { - if (newState === false) { - onSortChange("startTime", "desc"); - } else { - onSortChange(field, newState); - } - }} - /> -
-); - -export const createColumns = (sortProps?: LogsSortProps): ColumnDef[] => [ - { - header: sortProps - ? () => ( - - ) - : "Time", - accessorKey: "startTime", - size: 200, - cell: (info: any) => , - }, - { - header: "Type", - id: "type", - size: 90, - cell: (info: any) => { - const row = info.row.original; - const sessionCount = row.session_total_count || 1; - const isMcp = MCP_CALL_TYPES.includes(row.call_type); - const isAgent = AGENT_CALL_TYPES.includes(row.call_type); - const sessionLlmCount = row.session_llm_count ?? (isMcp || isAgent ? 0 : sessionCount); - const sessionAgentCount = row.session_agent_count ?? (isAgent ? sessionCount : 0); - const sessionMcpCount = row.session_mcp_count ?? (isMcp ? sessionCount : 0); - - if (isMcp) return ; - if (isAgent && sessionCount <= 1) return ; - if (sessionCount <= 1) return ; - - // Multi-call session — show total count, plus Agent/MCP indicators when mixed. - const sessionTypeBadge = ( - - - {sessionCount} - {sessionAgentCount > 0 && ( - <> - · - - - )} - {sessionMcpCount > 0 && ( - <> - · - - - )} - - ); - - const tooltipParts = [ - sessionLlmCount > 0 && `${sessionLlmCount} LLM`, - sessionAgentCount > 0 && `${sessionAgentCount} Agent`, - sessionMcpCount > 0 && `${sessionMcpCount} MCP`, - ].filter(Boolean); - return {sessionTypeBadge}; - }, - }, - { - header: "Status", - accessorKey: "metadata.status", - size: 100, - cell: (info: any) => { - const status = info.getValue() || "Success"; - const isSuccess = status.toLowerCase() !== "failure"; - return ; - }, - }, - { - header: "Session ID", - accessorKey: "session_id", - size: 120, - cell: (info: any) => , - }, - - { - header: "Request ID", - accessorKey: "request_id", - cell: (info: any) => , - }, - { - header: sortProps - ? () => ( - - ) - : "Cost", - accessorKey: "spend", - size: 110, - meta: { numeric: true }, - cell: (info: any) => { - const row = info.row.original; - const mcpCount = row.mcp_tool_call_count || 0; - const mcpSpend = row.mcp_tool_call_spend || 0; - const isMultiCallSession = (row.session_total_count || 1) > 1; - const spend = isMultiCallSession && row.session_total_spend != null ? row.session_total_spend : info.getValue(); - - return ( -
- - - - - - {isMultiCallSession && session total} - {mcpCount > 0 && mcpSpend > 0 && ( - - incl. {getSpendString(mcpSpend)} from {mcpCount} MCP - - )} -
- ); - }, - }, - { - header: sortProps - ? () => ( - - ) - : "Duration (s)", - accessorKey: "request_duration_ms", - meta: { numeric: true }, - cell: (info: any) => { - const ms = info.getValue(); - if (ms == null) return -; - const seconds = (ms / 1000).toFixed(2); - return ( - - {seconds} - - ); - }, - }, - { - header: sortProps - ? () => ( - - ) - : "TTFT (s)", - accessorKey: "completionStartTime", - meta: { numeric: true }, - cell: (info: any) => { - const row = info.row.original; - const completionStartTime = info.getValue(); - if (!completionStartTime) return -; - // For non-streaming, completionStartTime == endTime so TTFT is not meaningful - if (completionStartTime === row.endTime) return -; - const ttftMs = new Date(completionStartTime).getTime() - new Date(row.startTime).getTime(); - if (ttftMs <= 0) return -; - const ttftSeconds = (ttftMs / 1000).toFixed(2); - return ( - - {ttftSeconds} - - ); - }, - }, - { - header: "Team Name", - accessorKey: "metadata.user_api_key_team_alias", - size: 150, - cell: (info: any) => ( - - {String(info.getValue() || "-")} - - ), - }, - { - header: "Key Hash", - accessorKey: "metadata.user_api_key", - size: 110, - cell: (info: any) => , - }, - { - header: "Key Alias", - accessorKey: "metadata.user_api_key_alias", - size: 150, - cell: (info: any) => ( - - {String(info.getValue() || "-")} - - ), - }, - { - header: sortProps - ? () => ( - - ) - : "Model", - accessorKey: "model", - size: 200, - cell: (info: any) => { - const row = info.row.original; - const provider = row.custom_llm_provider; - const modelName = String(info.getValue() || ""); - return ( -
- {provider && ( - { - const target = e.target as HTMLImageElement; - target.style.display = "none"; - }} - /> - )} - - {modelName} - -
- ); - }, - }, - { - header: sortProps - ? () => ( - - ) - : "Tokens", - accessorKey: "total_tokens", - size: 140, - meta: { numeric: true }, - cell: (info: any) => { - const row = info.row.original; - return ( - - {String(row.total_tokens || "0")} - - ({String(row.prompt_tokens || "0")}+{String(row.completion_tokens || "0")}) - - - ); - }, - }, - { - header: "Internal User", - accessorKey: "user", - size: 150, - cell: (info: any) => ( - - {String(info.getValue() || "-")} - - ), - }, - { - header: "End User", - accessorKey: "end_user", - size: 140, - cell: (info: any) => ( - - {String(info.getValue() || "-")} - - ), - }, - - { - header: "Tags", - accessorKey: "request_tags", - size: 150, - cell: (info: any) => { - const tags = info.getValue(); - if (!tags || Object.keys(tags).length === 0) return "-"; - - const tagEntries = Object.entries(tags); - const firstTag = tagEntries[0]; - const remainingTags = tagEntries.slice(1); - - return ( -
- - {tagEntries.map(([key, value]) => ( - - {key}: {String(value)} - - ))} -
- } - > - - {firstTag[0]}: {String(firstTag[1])} - {remainingTags.length > 0 && ` +${remainingTags.length}`} - - -
- ); - }, - }, -]; - -/** Default columns without sort (for backward compatibility) */ -export const columns = createColumns(); - -const formatMessage = (message: any): string => { - if (!message) return "N/A"; - if (typeof message === "string") return message; - if (typeof message === "object") { - // Handle the {text, type} object specifically - if (message.text) return message.text; - if (message.content) return message.content; - return JSON.stringify(message); - } - return String(message); -}; - -// Add this new component for displaying request/response with copy buttons -export const RequestResponsePanel = ({ request, response }: { request: any; response: any }) => { - const requestStr = typeof request === "object" ? JSON.stringify(request, null, 2) : String(request || "{}"); - const responseStr = typeof response === "object" ? JSON.stringify(response, null, 2) : String(response || "{}"); - - const copyToClipboard = async (text: string) => { - try { - await navigator.clipboard.writeText(text); - } catch (err) { - console.error("Failed to copy text: ", err); - } - }; - - return ( -
-
-
-

Request

- -
-
{requestStr}
-
- -
-
-

Response

- -
-
-          {responseStr}
-        
-
-
- ); -}; - -// New component for collapsible JSON display -const CollapsibleJsonCell = ({ jsonData }: { jsonData: any }) => { - const [isExpanded, setIsExpanded] = React.useState(false); - const jsonString = JSON.stringify(jsonData, null, 2); - - if (!jsonData || Object.keys(jsonData).length === 0) { - return -; - } - - return ( -
- - {isExpanded && ( -
{jsonString}
- )} -
- ); }; diff --git a/ui/litellm-dashboard/src/components/view_logs/filter_options.ts b/ui/litellm-dashboard/src/components/view_logs/filter_options.ts deleted file mode 100644 index 90e27b6144d..00000000000 --- a/ui/litellm-dashboard/src/components/view_logs/filter_options.ts +++ /dev/null @@ -1,82 +0,0 @@ -import FilterTeamDropdown from "../common_components/FilterTeamDropdown"; -import { PaginatedKeyAliasSelect } from "../KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect"; -import { PaginatedModelSelect } from "../ModelSelect/PaginatedModelSelect/PaginatedModelSelect"; -import { FilterOption } from "../molecules/filter"; -import { allEndUsersCall } from "../networking"; -import { ERROR_CODE_OPTIONS } from "./constants"; -import { FILTER_KEYS } from "./log_filter_logic"; - -export function getLogFilterOptions(accessToken: string): FilterOption[] { - return [ - { - name: "Team ID", - label: "Team ID", - customComponent: FilterTeamDropdown, - }, - { - name: "Status", - label: "Status", - isSearchable: false, - options: [ - { label: "Success", value: "success" }, - { label: "Failure", value: "failure" }, - ], - }, - { - name: "Key Alias", - label: "Key Alias", - customComponent: PaginatedKeyAliasSelect, - }, - { - name: "End User", - label: "End User", - isSearchable: true, - searchFn: async (searchText: string) => { - const data = await allEndUsersCall(accessToken); - const users = data?.map((u: any) => u.user_id) || []; - const filtered = users.filter((u: string) => u.toLowerCase().includes(searchText.toLowerCase())); - return filtered.map((u: string) => ({ label: u, value: u })); - }, - }, - { - name: "Error Code", - label: "Error Code", - isSearchable: true, - searchFn: async (searchText: string) => { - if (!searchText) return ERROR_CODE_OPTIONS; - const lower = searchText.toLowerCase(); - const filtered = ERROR_CODE_OPTIONS.filter((opt) => opt.label.toLowerCase().includes(lower)); - const isExactValue = ERROR_CODE_OPTIONS.some((opt) => opt.value === searchText.trim()); - if (!isExactValue && searchText.trim()) { - filtered.push({ label: `Use custom code: ${searchText.trim()}`, value: searchText.trim() }); - } - return filtered; - }, - }, - { - name: "Error Message", - label: "Error Message", - isSearchable: false, - }, - { - name: "Key Hash", - label: "Key Hash", - isSearchable: false, - }, - { - name: FILTER_KEYS.SESSION_ID, - label: "Session ID", - isSearchable: false, - }, - { - name: "Model", - label: "Model", - customComponent: PaginatedModelSelect, - }, - { - name: FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL, - label: "Public model / search tool", - isSearchable: false, - }, - ]; -} diff --git a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx index 6f1c4126282..b2e77ec7fd5 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx @@ -1,106 +1,60 @@ -import { screen, waitFor } from "@testing-library/react"; +import { screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; -import moment from "moment"; -import { beforeEach, describe, expect, it, vi } from "vitest"; +import { describe, expect, it, vi } from "vitest"; import SpendLogsTable from "./index"; import { renderWithProviders } from "../../../tests/test-utils"; -import { uiSpendLogsCall } from "../networking"; -import { useLogFilterLogic } from "./log_filter_logic"; -const mockHandleFilterResetFromHook = vi.fn(); -vi.mock("./log_filter_logic", async (importOriginal) => { - const actual = await importOriginal(); - return { - ...actual, - useLogFilterLogic: vi.fn(() => ({ - logsQuery: { isLoading: false, isFetching: false, refetch: vi.fn() }, - filteredLogs: { data: [], total: 0, page: 1, page_size: 50, total_pages: 1 }, - allTeams: [], - handleFilterChange: vi.fn(), - handleFilterReset: mockHandleFilterResetFromHook, - })), - }; -}); - -vi.mock("../networking", async (importOriginal) => { - const actual = await importOriginal(); - return { - ...actual, - uiSpendLogsCall: vi.fn().mockResolvedValue({ - data: [], - total: 0, - page: 1, - page_size: 50, - total_pages: 0, - }), - keyListCall: vi.fn().mockResolvedValue({ keys: [] }), - keyInfoV1Call: vi.fn().mockResolvedValue({ info: {} }), - allEndUsersCall: vi.fn().mockResolvedValue([]), - }; -}); - -vi.mock("../key_team_helpers/filter_helpers", () => ({ - fetchAllTeams: vi.fn().mockResolvedValue([]), +vi.mock("./RequestLogsPanel", () => ({ + default: function RequestLogsPanelMock({ isActive }: { isActive: boolean }) { + return
{isActive ? "active" : "inactive"}
; + }, })); +vi.mock("./AuditLogsPanel", () => ({ + default: function AuditLogsPanelMock({ isActive }: { isActive: boolean }) { + return
{isActive ? "active" : "inactive"}
; + }, +})); + +vi.mock("../DeletedKeysPage/DeletedKeysPage", () => ({ + default: function DeletedKeysPageMock() { + return
; + }, +})); + +vi.mock("../DeletedTeamsPage/DeletedTeamsPage", () => ({ + default: function DeletedTeamsPageMock() { + return
; + }, +})); + +const defaultProps = { + accessToken: "test-token", + token: "test-token", + userRole: "Admin", + userID: "user-1", + premiumUser: false, +}; + describe("SpendLogsTable", () => { - const defaultProps = { - accessToken: "test-token", - token: "test-token", - userRole: "Admin", - userID: "user-1", - premiumUser: false, - }; + it("renders the four log tabs", () => { + renderWithProviders(); - beforeEach(() => { - vi.clearAllMocks(); - // Clear sessionStorage to avoid isLiveTail state from previous tests - sessionStorage.clear(); + for (const label of ["Request Logs", "Audit Logs", "Deleted Keys", "Deleted Teams"]) { + expect(screen.getByRole("tab", { name: label })).toBeInTheDocument(); + } }); - it("should call handleFilterResetFromHook when Reset Filters is clicked", async () => { + it("marks only the visible tab's panel active so background tabs do not query", async () => { const user = userEvent.setup(); renderWithProviders(); - const resetButton = screen.getByRole("button", { name: "Reset Filters" }); - await user.click(resetButton); + expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("active"); - await waitFor(() => { - expect(mockHandleFilterResetFromHook).toHaveBeenCalledTimes(1); - }); - }); + await user.click(screen.getByRole("tab", { name: "Audit Logs" })); - it("should reset custom date range to default when Reset Filters is clicked", async () => { - const user = userEvent.setup(); - renderWithProviders(); - - // Open the time range quick select dropdown (button shows current range like "Last 24 Hours") - const quickSelectButton = screen.getByRole("button", { - name: /Last 24 Hours|Last 15 Minutes|Last Hour|Last 4 Hours|Last 7 Days/i, - }); - await user.click(quickSelectButton); - - // Click "Custom Range" to enable custom date selection - const customRangeButton = await screen.findByRole("button", { name: "Custom Range" }); - await user.click(customRangeButton); - - // Custom date inputs should now be visible (start and end datetime-local inputs) - const datetimeInputs = document.querySelectorAll('input[type="datetime-local"]'); - expect(datetimeInputs.length).toBeGreaterThanOrEqual(2); - - // Click Reset Filters - this should reset the custom date range and hide custom inputs - const resetButton = screen.getByRole("button", { name: "Reset Filters" }); - await user.click(resetButton); - - await waitFor(() => { - expect(mockHandleFilterResetFromHook).toHaveBeenCalled(); - }); - - // After reset, custom date inputs should be hidden (isCustomDate reset to false) - await waitFor(() => { - const inputsAfterReset = document.querySelectorAll('input[type="datetime-local"]'); - expect(inputsAfterReset.length).toBe(0); - }); + expect(await screen.findByTestId("audit-logs-panel")).toHaveTextContent("active"); + expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("inactive"); }); describe("auth-not-ready guard", () => { @@ -108,73 +62,14 @@ describe("SpendLogsTable", () => { renderWithProviders(); expect(document.querySelector(".ant-spin")).toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Reset Filters" })).not.toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Request Logs" })).not.toBeInTheDocument(); }); - it("renders the table (no spinner) once all credentials are present", () => { + it("renders the tabs (no spinner) once all credentials are present", () => { renderWithProviders(); expect(document.querySelector(".ant-spin")).not.toBeInTheDocument(); - expect(screen.getByRole("button", { name: "Reset Filters" })).toBeInTheDocument(); - }); - }); - - describe("Quick Select time range", () => { - // uiSpendLogsCall fires from the real useLogFilterLogic query, so restore it here. - beforeEach(async () => { - const actual = await vi.importActual("./log_filter_logic"); - vi.mocked(useLogFilterLogic).mockImplementation(actual.useLogFilterLogic); - }); - - const waitForWindowSeconds = async (minMinutes: number) => { - let diff = -1; - await waitFor(() => { - const lastCall = vi.mocked(uiSpendLogsCall).mock.calls.at(-1)?.[0]; - if (!lastCall) throw new Error("uiSpendLogsCall was not called"); - diff = moment - .utc(lastCall.end_date, "YYYY-MM-DD HH:mm:ss") - .diff(moment.utc(lastCall.start_date, "YYYY-MM-DD HH:mm:ss"), "seconds"); - // start_date is rounded down to the minute boundary, end_date is the - // current wall-clock at queryFn time. The dropped sub-minute fraction - // on start_date can push the diff up to (minMinutes+1)*60 seconds - // exactly (e.g. click at HH:MM:59.9 → start floors to HH:MM:00 and - // queryFn fires just past HH:(MM+1):00), so allow equality on the - // upper bound. - expect(diff).toBeGreaterThanOrEqual(minMinutes * 60); - expect(diff).toBeLessThanOrEqual((minMinutes + 1) * 60); - }); - return diff; - }; - - it("should pass a ~1-minute window to uiSpendLogsCall when 'Last Minute' is selected", async () => { - const user = userEvent.setup(); - renderWithProviders(); - - await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); - await user.click(await screen.findByRole("button", { name: "Last Minute" })); - - await waitForWindowSeconds(1); - }); - - it("should pass a ~15-minute window to uiSpendLogsCall when 'Last 15 Minutes' is selected", async () => { - const user = userEvent.setup(); - renderWithProviders(); - - await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); - await user.click(await screen.findByRole("button", { name: "Last 15 Minutes" })); - - await waitForWindowSeconds(15); - }); - - it("should update the time-range button label to 'Last Minute' after selecting it", async () => { - const user = userEvent.setup(); - renderWithProviders(); - - await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); - await user.click(await screen.findByRole("button", { name: "Last Minute" })); - - expect(screen.getByRole("button", { name: "Last Minute" })).toBeInTheDocument(); - expect(screen.queryByRole("button", { name: /Last 24 Hours/i })).not.toBeInTheDocument(); + expect(screen.getByRole("tab", { name: "Request Logs" })).toBeInTheDocument(); }); }); }); diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx index aa8077a02e9..8e7423e3fae 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx @@ -1,21 +1,9 @@ -import moment from "moment"; -import { useCallback, useDeferredValue, useEffect, useMemo, useState } from "react"; +import { useState } from "react"; import { Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; -import { internalUserRoles } from "../../utils/roles"; import DeletedKeysPage from "../DeletedKeysPage/DeletedKeysPage"; import DeletedTeamsPage from "../DeletedTeamsPage/DeletedTeamsPage"; -import { KeyResponse } from "../key_team_helpers/key_list"; -import FilterComponent from "../molecules/filter"; -import { keyInfoV1Call } from "../networking"; -import KeyInfoView from "../templates/key_info_view"; import AuditLogsPanel from "./AuditLogsPanel"; -import { createColumns, LogEntry, type LogsSortField } from "./columns"; -import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants"; -import { getLogFilterOptions } from "./filter_options"; -import { useLogFilterLogic, defaultFilters, type LogFilterState } from "./log_filter_logic"; -import { LogDetailsDrawer } from "./LogDetailsDrawer"; -import { LogsTableToolbar } from "./LogsTableToolbar"; -import { DataTable } from "./table"; +import RequestLogsPanel from "./RequestLogsPanel"; import { AntDLoadingSpinner } from "../ui/AntDLoadingSpinner"; interface SpendLogsTableProps { @@ -27,190 +15,8 @@ interface SpendLogsTableProps { } export default function SpendLogsTable({ accessToken, token, userRole, userID, premiumUser }: SpendLogsTableProps) { - const [searchTerm, setSearchTerm] = useState(""); - const [currentPage, setCurrentPage] = useState(1); - const [pageSize] = useState(50); - - // New state variables for Start and End Time - const [startTime, setStartTime] = useState(moment().subtract(24, "hours").format("YYYY-MM-DDTHH:mm")); - const [endTime, setEndTime] = useState(moment().format("YYYY-MM-DDTHH:mm")); - - const [isCustomDate, setIsCustomDate] = useState(false); - const [filters, setFilters] = useState(defaultFilters); - const [selectedKeyInfo, setSelectedKeyInfo] = useState(null); - const [selectedKeyIdInfoView, setSelectedKeyIdInfoView] = useState(null); - const [filterByCurrentUser, setFilterByCurrentUser] = useState(userRole && internalUserRoles.includes(userRole)); const [activeTab, setActiveTab] = useState("request logs"); - const [selectedLog, setSelectedLog] = useState(null); - const [isDrawerOpen, setIsDrawerOpen] = useState(false); - const [selectedSessionId, setSelectedSessionId] = useState(null); - - const [sortBy, setSortBy] = useState("startTime"); - const [sortOrder, setSortOrder] = useState<"asc" | "desc">("desc"); - - const [selectedTimeInterval, setSelectedTimeInterval] = useState<{ value: number; unit: string }>({ - value: 24, - unit: "hours", - }); - - const [isLiveTail, setIsLiveTail] = useState(() => { - const storedValue = sessionStorage.getItem("isLiveTail"); - // default to true if nothing is stored - return storedValue !== null ? JSON.parse(storedValue) : true; - }); - - useEffect(() => { - sessionStorage.setItem("isLiveTail", JSON.stringify(isLiveTail)); - }, [isLiveTail]); - - useEffect(() => { - const fetchKeyInfo = async () => { - if (selectedKeyIdInfoView && accessToken) { - const keyData = await keyInfoV1Call(accessToken, selectedKeyIdInfoView); - - const keyResponse: KeyResponse = { - ...keyData["info"], - token: selectedKeyIdInfoView, - api_key: selectedKeyIdInfoView, - }; - setSelectedKeyInfo(keyResponse); - } - }; - fetchKeyInfo(); - }, [selectedKeyIdInfoView, accessToken]); - - useEffect(() => { - if (userRole && internalUserRoles.includes(userRole)) { - setFilterByCurrentUser(true); - } - }, [userRole]); - - const { - logsQuery, - filteredLogs, - allTeams, - handleFilterChange, - handleFilterReset: handleFilterResetFromHook, - } = useLogFilterLogic({ - accessToken, - token, - userRole, - userID, - filters, - setFilters, - filterByCurrentUser: !!filterByCurrentUser, - activeTab, - isLiveTail, - startTime, - endTime, - pageSize, - isCustomDate, - setCurrentPage, - sortBy, - sortOrder, - currentPage, - }); - - const handleFilterReset = useCallback(() => { - handleFilterResetFromHook(); - setStartTime(moment().subtract(24, "hours").format("YYYY-MM-DDTHH:mm")); - setEndTime(moment().format("YYYY-MM-DDTHH:mm")); - setIsCustomDate(false); - setSelectedTimeInterval({ value: 24, unit: "hours" }); - setCurrentPage(1); - }, [handleFilterResetFromHook]); - - const handleSortChange = useCallback((newSortBy: LogsSortField, newSortOrder: "asc" | "desc") => { - setSortBy(newSortBy); - setSortOrder(newSortOrder); - setCurrentPage(1); - }, []); - - const columns = useMemo( - () => createColumns({ sortBy, sortOrder, onSortChange: handleSortChange }), - [sortBy, sortOrder, handleSortChange], - ); - - const filteredData = useMemo(() => { - const searchedLogs = filteredLogs.data.filter((log) => { - const matchesSearch = - !searchTerm || - log.request_id.includes(searchTerm) || - log.model.includes(searchTerm) || - (log.user && log.user.includes(searchTerm)); - - // No need for additional filtering since we're now handling this in the API call - return matchesSearch; - }); - - const sessionCompositionById = searchedLogs.reduce>( - (acc, log) => { - if (!log.session_id) return acc; - if (!acc[log.session_id]) { - acc[log.session_id] = { llm: 0, agent: 0, mcp: 0 }; - } - if (MCP_CALL_TYPES.includes(log.call_type)) { - acc[log.session_id].mcp += 1; - } else if (AGENT_CALL_TYPES.includes(log.call_type)) { - acc[log.session_id].agent += 1; - } else { - acc[log.session_id].llm += 1; - } - return acc; - }, - {}, - ); - - // Build a single-pass map of session_id → representative request_id. - // Prefers an LLM row over an MCP row as the representative. - const sessionRepresentativeMap = new Map(); - for (const log of searchedLogs) { - if (!log.session_id || (log.session_total_count || 1) <= 1) continue; - const isMcp = MCP_CALL_TYPES.includes(log.call_type); - const existing = sessionRepresentativeMap.get(log.session_id); - if (!existing || (existing.isMcp && !isMcp)) { - sessionRepresentativeMap.set(log.session_id, { requestId: log.request_id, isMcp }); - } - } - - return ( - searchedLogs - .map((log) => { - const sessionComposition = log.session_id ? sessionCompositionById[log.session_id] : undefined; - return { - ...log, - request_duration_ms: log.request_duration_ms, - session_llm_count: sessionComposition?.llm ?? undefined, - session_mcp_count: sessionComposition?.mcp ?? undefined, - session_agent_count: sessionComposition?.agent ?? undefined, - onKeyHashClick: (keyHash: string) => setSelectedKeyIdInfoView(keyHash), - onSessionClick: (sessionId: string) => { - if (sessionId) { - setSelectedSessionId(sessionId); - setSelectedLog(log); - setIsDrawerOpen(true); - } - }, - }; - }) - // Deduplicate multi-call sessions using the pre-built map (O(1) per row). - .filter((log) => { - if (!log.session_id || (log.session_total_count || 1) <= 1) return true; - return sessionRepresentativeMap.get(log.session_id)?.requestId === log.request_id; - }) - ); - }, [filteredLogs.data, searchTerm]); - - // Keep the Fetch button busy until the table has actually committed the new - // rows. `keepPreviousData` leaves logsQuery.isLoading false on refetch, so - // without this the button clears while stale rows are still on screen. - const deferredData = useDeferredValue(filteredData); - const isStale = deferredData !== filteredData; - const isButtonLoading = logsQuery.isFetching || isStale; - const isRefiltering = logsQuery.isPlaceholderData; - const isLogsLoading = logsQuery.isLoading || isRefiltering; - if (!accessToken || !token || !userRole || !userID) { return (
@@ -219,20 +25,6 @@ export default function SpendLogsTable({ accessToken, token, userRole, userID, p ); } - const handleRowClick = (log: LogEntry) => { - // Multi-call session row: open in the same right-side drawer (session mode) - if (log.session_id && (log.session_total_count || 1) > 1) { - setSelectedSessionId(log.session_id); - setSelectedLog(log); - setIsDrawerOpen(true); - return; - } - // Single-call row: open the detail drawer - setSelectedSessionId(null); - setSelectedLog(log); - setIsDrawerOpen(true); - }; - return (
setActiveTab(index === 0 ? "request logs" : "audit logs")}> @@ -244,56 +36,13 @@ export default function SpendLogsTable({ accessToken, token, userRole, userID, p -
-

Request Logs

-
- {selectedKeyInfo && selectedKeyIdInfoView && selectedKeyInfo.api_key === selectedKeyIdInfoView ? ( - setSelectedKeyIdInfoView(null)} - backButtonText="Back to Logs" - /> - ) : ( - <> - -
- logsQuery.refetch()} - filteredLogs={filteredLogs} - /> - row.request_id} - onRowClick={handleRowClick} - isLoading={isLogsLoading} - /> -
- - )} +
- - {/* Log Details Drawer */} - { - setIsDrawerOpen(false); - setSelectedSessionId(null); - }} - logEntry={selectedLog} - sessionId={selectedSessionId} - accessToken={accessToken} - allLogs={filteredData} - onSelectLog={setSelectedLog} - startTime={moment(startTime).utc().format("YYYY-MM-DD HH:mm:ss")} - />
); } diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx index ef550baea91..9f8f2fbe542 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx @@ -1,14 +1,16 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { act, renderHook, waitFor } from "@testing-library/react"; -import React, { ReactNode, useState } from "react"; +import type { ColumnFiltersState, PaginationState, SortingState } from "@tanstack/react-table"; +import { renderHook, waitFor } from "@testing-library/react"; +import moment from "moment"; +import React, { ReactNode } from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import type { LogsSortField } from "./columns"; import { - defaultFilters, + DEFAULT_LOGS_SORTING, + getFilterValue, getLiveTailRefetchInterval, LIVE_TAIL_INTERVAL_MS, + LOG_FILTER_IDS, useLogFilterLogic, - type LogFilterState, type PaginatedResponse, } from "./log_filter_logic"; @@ -30,31 +32,33 @@ const emptyResponse: PaginatedResponse = { total_pages: 0, }; +const FIRST_PAGE: PaginationState = { pageIndex: 0, pageSize: 50 }; + const defaultProps = { accessToken: "test-token" as string | null, token: "test-token" as string | null, userRole: "Admin" as string | null, userID: "user-1" as string | null, + columnFilters: [] as ColumnFiltersState, filterByCurrentUser: false, activeTab: "request logs", isLiveTail: false, startTime: "2025-01-01T00:00:00", endTime: "2025-01-01T23:59:59", + pagination: FIRST_PAGE, isCustomDate: true, - sortBy: "startTime" as LogsSortField, - sortOrder: "desc" as "asc" | "desc", - currentPage: 1, + sorting: DEFAULT_LOGS_SORTING, }; -type HookOverrides = Partial[0], "filters" | "setFilters">>; +type HookOverrides = Partial[0]>; + +const lastCallParams = () => vi.mocked(uiSpendLogsCall).mock.calls.at(-1)?.[0]; describe("useLogFilterLogic", () => { let queryClient: QueryClient; beforeEach(() => { - queryClient = new QueryClient({ - defaultOptions: { queries: { retry: false } }, - }); + queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); vi.clearAllMocks(); vi.mocked(uiSpendLogsCall).mockResolvedValue(emptyResponse); }); @@ -63,602 +67,169 @@ describe("useLogFilterLogic", () => { React.createElement(QueryClientProvider, { client: queryClient }, children); function renderFilterHook(overrides: HookOverrides = {}) { - const setCurrentPage = overrides.setCurrentPage ?? vi.fn(); - const rendered = renderHook( - () => { - const [filters, setFilters] = useState(defaultFilters); - const hook = useLogFilterLogic({ - ...defaultProps, - ...overrides, - filters, - setFilters, - setCurrentPage, - }); - return { ...hook, filters, setFilters }; - }, - { wrapper }, - ); - return { ...rendered, setCurrentPage }; + return renderHook(() => useLogFilterLogic({ ...defaultProps, ...overrides }), { wrapper }); } - describe("return shape", () => { - it("exposes filteredLogs, allTeams, handleFilterChange, handleFilterReset", () => { - const { result } = renderFilterHook(); - - expect(result.current.filteredLogs).toBeDefined(); - expect(result.current).toHaveProperty("allTeams"); - expect(result.current.handleFilterChange).toBeInstanceOf(Function); - expect(result.current.handleFilterReset).toBeInstanceOf(Function); - }); - }); - - describe("handleFilterReset", () => { - it("restores filters to defaults after changes", () => { - const { result } = renderFilterHook(); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-1", Status: "success" }); - }); - - expect(result.current.filters["Team ID"]).toBe("team-1"); - expect(result.current.filters["Status"]).toBe("success"); - - act(() => { - result.current.handleFilterReset(); - }); - - expect(result.current.filters["Team ID"]).toBe(""); - expect(result.current.filters["Status"]).toBe(""); - }); - - it("calls setCurrentPage(1)", () => { - const setCurrentPage = vi.fn(); - const { result } = renderFilterHook({ setCurrentPage }); - - act(() => { - result.current.handleFilterReset(); - }); - - expect(setCurrentPage).toHaveBeenCalledWith(1); - }); - - it("triggers a fetch with all filter params undefined", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue(emptyResponse); - const { result } = renderFilterHook(); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); - - act(() => { - result.current.handleFilterReset(); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenLastCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ - team_id: undefined, - api_key: undefined, - request_id: undefined, - user_id: undefined, - end_user: undefined, - status_filter: undefined, - model_id: undefined, - key_alias: undefined, - error_code: undefined, - error_message: undefined, - }), - }), - ); - }, - { timeout: 500 }, - ); - }); - }); - - describe("handleFilterChange", () => { - it("calls setCurrentPage(1) when filters change", () => { - const setCurrentPage = vi.fn(); - const { result } = renderFilterHook({ setCurrentPage }); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-1" }); - }); - - expect(setCurrentPage).toHaveBeenCalledWith(1); - }); - - it("merges partial updates without clobbering other filter keys", () => { - const { result } = renderFilterHook(); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-a" }); - }); - expect(result.current.filters["Team ID"]).toBe("team-a"); - - act(() => { - result.current.handleFilterChange({ Model: "gpt-4" }); - }); - - expect(result.current.filters["Team ID"]).toBe("team-a"); - expect(result.current.filters["Model"]).toBe("gpt-4"); - }); - - it("does not call setCurrentPage when filters are identical", async () => { - const setCurrentPage = vi.fn(); - const { result } = renderFilterHook({ setCurrentPage }); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-1" }); - }); - - await waitFor(() => expect(setCurrentPage).toHaveBeenCalledTimes(1), { timeout: 500 }); - - setCurrentPage.mockClear(); - - await act(async () => { - result.current.handleFilterChange({ "Team ID": "team-1" }); - await new Promise((resolve) => setTimeout(resolve, 350)); - }); - - expect(setCurrentPage).not.toHaveBeenCalled(); - }); - }); - - describe("query params — filter keys", () => { - const filterCases: Array<{ - filterKey: keyof LogFilterState; - paramName: string; - value: string; - }> = [ - { filterKey: "Team ID", paramName: "team_id", value: "team-a" }, - { filterKey: "Key Hash", paramName: "api_key", value: "key-x" }, - { filterKey: "Request ID", paramName: "request_id", value: "req-xyz" }, - { filterKey: "Session ID", paramName: "session_id", value: "sess-42" }, - { filterKey: "User ID", paramName: "user_id", value: "user-123" }, - { filterKey: "End User", paramName: "end_user", value: "user-a" }, - { filterKey: "Status", paramName: "status_filter", value: "error" }, - { filterKey: "Model", paramName: "model_id", value: "gpt-4" }, - { filterKey: "Public model / search tool", paramName: "model", value: "tavily-marketing" }, - { filterKey: "Error Code", paramName: "error_code", value: "429" }, - { filterKey: "Error Message", paramName: "error_message", value: "rate limit exceeded" }, + describe("column filters map onto backend query params", () => { + const cases: ReadonlyArray<{ id: string; value: string; param: string }> = [ + { id: LOG_FILTER_IDS.KEY_HASH, value: "sk-hash-1", param: "api_key" }, + { id: LOG_FILTER_IDS.TEAM_ID, value: "team-1", param: "team_id" }, + { id: LOG_FILTER_IDS.REQUEST_ID, value: "req-1", param: "request_id" }, + { id: LOG_FILTER_IDS.SESSION_ID, value: "sess-1", param: "session_id" }, + { id: LOG_FILTER_IDS.END_USER, value: "end-user-1", param: "end_user" }, + { id: LOG_FILTER_IDS.STATUS, value: "failure", param: "status_filter" }, + { id: LOG_FILTER_IDS.MODEL_ID, value: "model-uuid-1", param: "model_id" }, + { id: LOG_FILTER_IDS.PUBLIC_MODEL_OR_SEARCH_TOOL, value: "gpt-4o", param: "model" }, + { id: LOG_FILTER_IDS.KEY_ALIAS, value: "alias-1", param: "key_alias" }, + { id: LOG_FILTER_IDS.ERROR_CODE, value: "429", param: "error_code" }, + { id: LOG_FILTER_IDS.ERROR_MESSAGE, value: "rate limited", param: "error_message" }, + { id: LOG_FILTER_IDS.USER_ID, value: "user-9", param: "user_id" }, ]; - it.each(filterCases)( - "forwards $filterKey as params.$paramName to uiSpendLogsCall", - async ({ filterKey, paramName, value }) => { - const { result } = renderFilterHook(); + it.each(cases)("sends $id as $param", async ({ id, value, param }) => { + renderFilterHook({ columnFilters: [{ id, value }] }); - act(() => { - result.current.handleFilterChange({ [filterKey]: value } as Partial); - }); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + expect(lastCallParams()?.params).toMatchObject({ [param]: value }); + }); - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ [paramName]: value }), - }), - ); - }, - { timeout: 500 }, - ); - }, - ); - }); - - describe("query params — date & sort", () => { - it("passes start_date, end_date, sort_by, and sort_order to uiSpendLogsCall", async () => { - const { result } = renderFilterHook({ - startTime: "2025-01-15T00:00:00Z", - endTime: "2025-01-15T23:59:59Z", - isCustomDate: true, - sortBy: "spend" as LogsSortField, - sortOrder: "asc", + it("omits params for filters that are absent, blank, or whitespace-only", async () => { + renderFilterHook({ + columnFilters: [ + { id: LOG_FILTER_IDS.TEAM_ID, value: " " }, + { id: LOG_FILTER_IDS.KEY_HASH, value: "" }, + ], }); - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - start_date: "2025-01-15 00:00:00", - end_date: "2025-01-15 23:59:59", - params: expect.objectContaining({ - sort_by: "spend", - sort_order: "asc", - }), - }), - ); - }, - { timeout: 500 }, - ); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + const params = lastCallParams()?.params; + expect(params?.team_id).toBeUndefined(); + expect(params?.api_key).toBeUndefined(); + expect(params?.error_code).toBeUndefined(); }); }); - describe("debounce", () => { - it("calls uiSpendLogsCall after the debounce elapses for text filters", async () => { - const { result } = renderFilterHook(); + describe("paging, dates, and sort", () => { + it("sends a 1-based page derived from pageIndex", async () => { + renderFilterHook({ pagination: { pageIndex: 2, pageSize: 25 } }); - act(() => { - result.current.handleFilterChange({ "Key Hash": "hash-1" }); - }); - - await waitFor( - () => - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ api_key: "hash-1" }), - }), - ), - { timeout: 500 }, - ); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + expect(lastCallParams()).toMatchObject({ page: 3, page_size: 25 }); }); - it("does not call uiSpendLogsCall with a text filter before the debounce elapses", async () => { - const { result } = renderFilterHook(); + it("passes start_date, end_date, sort_by, and sort_order", async () => { + renderFilterHook({ sorting: [{ id: "spend", desc: false }] }); - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); - vi.mocked(uiSpendLogsCall).mockClear(); - - act(() => { - result.current.handleFilterChange({ "Key Hash": "hash-1" }); - }); - - await new Promise((resolve) => setTimeout(resolve, 100)); - expect(uiSpendLogsCall).not.toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ api_key: "hash-1" }), - }), - ); - - await waitFor( - () => - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ api_key: "hash-1" }), - }), - ), - { timeout: 500 }, - ); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + const call = lastCallParams(); + expect(call?.start_date).toBe(moment(defaultProps.startTime).utc().format("YYYY-MM-DD HH:mm:ss")); + expect(call?.end_date).toBe(moment(defaultProps.endTime).utc().format("YYYY-MM-DD HH:mm:ss")); + expect(call?.params).toMatchObject({ sort_by: "spend", sort_order: "asc" }); }); - it("applies dropdown filter changes without waiting for the debounce", async () => { - const { result } = renderFilterHook(); + it("falls back to the default sort when the sorting state is empty", async () => { + renderFilterHook({ sorting: [] }); - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); - vi.mocked(uiSpendLogsCall).mockClear(); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-instant" }); - }); - - await waitFor( - () => - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ team_id: "team-instant" }), - }), - ), - { timeout: 100 }, - ); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + expect(lastCallParams()?.params).toMatchObject({ sort_by: "startTime", sort_order: "desc" }); }); - // Guards the TEXT_FILTER_KEYS fix: this free-text filter must debounce, not fire per keystroke. - it("debounces the 'Public model / search tool' text filter", async () => { - const { result } = renderFilterHook(); + it("ignores a sort id the backend does not support", async () => { + renderFilterHook({ sorting: [{ id: "request_id", desc: false }] as SortingState }); - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); - vi.mocked(uiSpendLogsCall).mockClear(); - - act(() => { - result.current.handleFilterChange({ "Public model / search tool": "tavily-marketing" }); - }); - - await new Promise((resolve) => setTimeout(resolve, 100)); - expect(uiSpendLogsCall).not.toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ model: "tavily-marketing" }), - }), - ); - - await waitFor( - () => - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ model: "tavily-marketing" }), - }), - ), - { timeout: 500 }, - ); - }); - }); - - describe("handleFilterReset", () => { - it("flushes the text-filter debounce so a pending typed value is not sent", async () => { - const { result } = renderFilterHook(); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); - vi.mocked(uiSpendLogsCall).mockClear(); - - act(() => { - result.current.handleFilterChange({ "Key Hash": "pending-hash" }); - }); - - act(() => { - result.current.handleFilterReset(); - }); - - await new Promise((resolve) => setTimeout(resolve, 400)); - - for (const call of vi.mocked(uiSpendLogsCall).mock.calls) { - expect(call[0].params?.api_key).toBeUndefined(); - } - }); - }); - - describe("backend filtered logs", () => { - it("returns the query payload as filteredLogs when backend filters are active", async () => { - const backendLog = { request_id: "backend-req" }; - vi.mocked(uiSpendLogsCall).mockResolvedValue({ - data: [backendLog], - total: 1, - page: 1, - page_size: 50, - total_pages: 1, - } as PaginatedResponse); - - const { result } = renderFilterHook(); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor( - () => { - expect(result.current.filteredLogs.data).toHaveLength(1); - expect(result.current.filteredLogs.data[0].request_id).toBe("backend-req"); - }, - { timeout: 500 }, - ); - }); - - it("returns empty data when the API returns an empty payload", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue(emptyResponse); - const { result } = renderFilterHook(); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); - - expect(result.current.filteredLogs.data).toHaveLength(0); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + expect(lastCallParams()?.params).toMatchObject({ sort_by: "startTime" }); }); }); describe("refetch triggers", () => { - it("refetches when sortBy changes", async () => { - const { rerender } = renderHook( - (props: { sortBy: LogsSortField }) => { - const [filters, setFilters] = useState(defaultFilters); - return useLogFilterLogic({ - ...defaultProps, - filters, - setFilters, - setCurrentPage: vi.fn(), - sortBy: props.sortBy, - }); - }, - { wrapper, initialProps: { sortBy: "startTime" } }, - ); + it.each([ + ["sorting", { sorting: [{ id: "spend", desc: true }] as SortingState }], + ["pagination", { pagination: { pageIndex: 1, pageSize: 50 } }], + ["startTime", { startTime: "2025-02-02T00:00:00" }], + ["columnFilters", { columnFilters: [{ id: LOG_FILTER_IDS.TEAM_ID, value: "team-2" }] }], + ])("refetches when %s changes", async (_label, nextProps) => { + const { rerender } = renderHook((props: HookOverrides) => useLogFilterLogic({ ...defaultProps, ...props }), { + wrapper, + initialProps: {}, + }); - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); - - rerender({ sortBy: "spend" }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); - expect(uiSpendLogsCall).toHaveBeenLastCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ sort_by: "spend" }), - }), - ); - }); - - it("refetches when sortOrder changes", async () => { - const { rerender } = renderHook( - (props: { sortOrder: "asc" | "desc" }) => { - const [filters, setFilters] = useState(defaultFilters); - return useLogFilterLogic({ - ...defaultProps, - filters, - setFilters, - setCurrentPage: vi.fn(), - sortOrder: props.sortOrder, - }); - }, - { wrapper, initialProps: { sortOrder: "desc" } }, - ); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); - - rerender({ sortOrder: "asc" }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); - expect(uiSpendLogsCall).toHaveBeenLastCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ sort_order: "asc" }), - }), - ); - }); - - it("refetches when currentPage changes", async () => { - const { rerender } = renderHook( - (props: { currentPage: number }) => { - const [filters, setFilters] = useState(defaultFilters); - return useLogFilterLogic({ - ...defaultProps, - filters, - setFilters, - setCurrentPage: vi.fn(), - currentPage: props.currentPage, - }); - }, - { wrapper, initialProps: { currentPage: 1 } }, - ); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); - - rerender({ currentPage: 2 }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); - expect(uiSpendLogsCall).toHaveBeenLastCalledWith(expect.objectContaining({ page: 2 })); - }); - - it("refetches when startTime changes", async () => { - const { rerender } = renderHook( - (props: { startTime: string }) => { - const [filters, setFilters] = useState(defaultFilters); - return useLogFilterLogic({ - ...defaultProps, - filters, - setFilters, - setCurrentPage: vi.fn(), - startTime: props.startTime, - }); - }, - { wrapper, initialProps: { startTime: "2025-01-01T00:00:00Z" } }, - ); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); - - rerender({ startTime: "2025-01-02T00:00:00Z" }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); - expect(uiSpendLogsCall).toHaveBeenLastCalledWith(expect.objectContaining({ start_date: "2025-01-02 00:00:00" })); - }); - - it("refetches with a different end_date when isCustomDate toggles", async () => { - const customEndTime = "2025-01-15T23:59:59Z"; - const customEndFormatted = "2025-01-15 23:59:59"; - - const { rerender } = renderHook( - (props: { isCustomDate: boolean }) => { - const [filters, setFilters] = useState(defaultFilters); - return useLogFilterLogic({ - ...defaultProps, - endTime: customEndTime, - filters, - setFilters, - setCurrentPage: vi.fn(), - isCustomDate: props.isCustomDate, - }); - }, - { wrapper, initialProps: { isCustomDate: false } }, - ); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); - const firstEndDate = vi.mocked(uiSpendLogsCall).mock.calls[0][0].end_date; - expect(firstEndDate).not.toBe(customEndFormatted); - - rerender({ isCustomDate: true }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); - expect(vi.mocked(uiSpendLogsCall).mock.calls[1][0].end_date).toBe(customEndFormatted); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1)); + rerender(nextProps); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2)); }); }); describe("query enablement", () => { - const nullCredentialCases: Array<{ name: string; override: HookOverrides }> = [ - { name: "accessToken", override: { accessToken: null } }, - { name: "token", override: { token: null } }, - { name: "userRole", override: { userRole: null } }, - { name: "userID", override: { userID: null } }, - ]; - - it.each(nullCredentialCases)("does not call uiSpendLogsCall when $name is null", async ({ override }) => { - const { result } = renderFilterHook(override); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await new Promise((resolve) => setTimeout(resolve, 350)); + it("does not query when the request logs tab is inactive", async () => { + renderFilterHook({ activeTab: "audit logs" }); + await new Promise((resolve) => setTimeout(resolve, 50)); expect(uiSpendLogsCall).not.toHaveBeenCalled(); }); - it("does not call uiSpendLogsCall when activeTab is not 'request logs'", async () => { - const { result } = renderFilterHook({ activeTab: "audit logs" }); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await new Promise((resolve) => setTimeout(resolve, 350)); + it("does not query when credentials are missing", async () => { + renderFilterHook({ accessToken: null }); + await new Promise((resolve) => setTimeout(resolve, 50)); expect(uiSpendLogsCall).not.toHaveBeenCalled(); }); }); describe("filterByCurrentUser", () => { - it("sends user_id: userID when the User ID filter is blank", async () => { - const { result } = renderFilterHook({ + it("scopes to the current user when no explicit user filter is set", async () => { + renderFilterHook({ filterByCurrentUser: true }); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + expect(lastCallParams()?.params).toMatchObject({ user_id: "user-1" }); + }); + + it("lets an explicit user filter win over the current-user scope", async () => { + renderFilterHook({ filterByCurrentUser: true, - userID: "me-123", + columnFilters: [{ id: LOG_FILTER_IDS.USER_ID, value: "someone-else" }], }); - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ user_id: "me-123" }), - }), - ); - }, - { timeout: 500 }, - ); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + expect(lastCallParams()?.params).toMatchObject({ user_id: "someone-else" }); }); }); - describe("error handling", () => { - it("does not crash when uiSpendLogsCall throws", async () => { - vi.mocked(uiSpendLogsCall).mockRejectedValue(new Error("Network error")); - const { result } = renderFilterHook(); + it("returns an empty payload and does not crash when the call fails", async () => { + vi.mocked(uiSpendLogsCall).mockRejectedValue(new Error("boom")); + const { result } = renderFilterHook(); - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled()); + expect(result.current.filteredLogs.data).toEqual([]); + expect(result.current.filteredLogs.total).toBe(0); + }); +}); - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); +describe("getFilterValue", () => { + it("trims values and treats blank ones as absent", () => { + const filters: ColumnFiltersState = [ + { id: "team_id", value: " team-1 " }, + { id: "key_hash", value: " " }, + { id: "status", value: 42 }, + ]; - expect(result.current.filteredLogs).toBeDefined(); - expect(result.current.filteredLogs.data).toEqual([]); - }); + expect(getFilterValue(filters, "team_id")).toBe("team-1"); + expect(getFilterValue(filters, "key_hash")).toBeUndefined(); + expect(getFilterValue(filters, "status")).toBeUndefined(); + expect(getFilterValue(filters, "missing")).toBeUndefined(); }); }); describe("getLiveTailRefetchInterval", () => { - it("polls every 15s when live tail is on and on page 1", () => { - expect(getLiveTailRefetchInterval(true, 1)).toBe(LIVE_TAIL_INTERVAL_MS); + it("polls every 15s when live tail is on and on the first page", () => { + expect(getLiveTailRefetchInterval(true, 0)).toBe(LIVE_TAIL_INTERVAL_MS); }); it("does not poll when live tail is off", () => { - expect(getLiveTailRefetchInterval(false, 1)).toBe(false); + expect(getLiveTailRefetchInterval(false, 0)).toBe(false); }); - it("does not poll when not on page 1, even with live tail on", () => { - expect(getLiveTailRefetchInterval(true, 2)).toBe(false); + it("does not poll past the first page, even with live tail on", () => { + expect(getLiveTailRefetchInterval(true, 1)).toBe(false); }); }); diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx index caa9d8d9361..ae229056fa3 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx @@ -1,13 +1,11 @@ import moment from "moment"; -import { useEffect, useMemo, useState } from "react"; -import { useDebouncer } from "@tanstack/react-pacer/debouncer"; -import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; +import { keepPreviousData, useQuery, type UseQueryOptions } from "@tanstack/react-query"; +import type { ColumnFiltersState, PaginationState, SortingState } from "@tanstack/react-table"; import { uiSpendLogsCall } from "../networking"; import { Team } from "../key_team_helpers/key_list"; -import { keepPreviousData, useQuery } from "@tanstack/react-query"; import { fetchAllTeams } from "../../components/key_team_helpers/filter_helpers"; import { defaultPageSize } from "../constants"; -import type { LogEntry, LogsSortField } from "./columns"; +import { LOGS_SORT_FIELD_MAP, type LogEntry, type LogsSortField } from "./columns"; export interface PaginatedResponse { data: LogEntry[]; @@ -18,54 +16,48 @@ export interface PaginatedResponse { total_is_capped?: boolean; } -/** Spend log `model` column (LLM public model name or `search_tool_name` for /search). */ -export const FILTER_KEYS = { - TEAM_ID: "Team ID", - KEY_HASH: "Key Hash", - REQUEST_ID: "Request ID", - SESSION_ID: "Session ID", - MODEL: "Model", - /** Exact match on LiteLLM_SpendLogs.model — use for search tools and public model names. */ - PUBLIC_MODEL_OR_SEARCH_TOOL: "Public model / search tool", - USER_ID: "User ID", - END_USER: "End User", - STATUS: "Status", - KEY_ALIAS: "Key Alias", - ERROR_CODE: "Error Code", - ERROR_MESSAGE: "Error Message", +export const LOG_FILTER_IDS = { + TEAM_ID: "team_id", + STATUS: "status", + KEY_ALIAS: "key_alias", + END_USER: "end_user", + ERROR_CODE: "error_code", + ERROR_MESSAGE: "error_message", + KEY_HASH: "key_hash", + SESSION_ID: "session_id", + MODEL_ID: "model_id", + PUBLIC_MODEL_OR_SEARCH_TOOL: "model", + REQUEST_ID: "request_id", + USER_ID: "user_id", } as const; -export type FilterKey = keyof typeof FILTER_KEYS; -export type LogFilterState = Record<(typeof FILTER_KEYS)[FilterKey], string>; +export const LOG_FILTER_LABELS: Record = { + [LOG_FILTER_IDS.TEAM_ID]: "Team ID", + [LOG_FILTER_IDS.STATUS]: "Status", + [LOG_FILTER_IDS.KEY_ALIAS]: "Key Alias", + [LOG_FILTER_IDS.END_USER]: "End User", + [LOG_FILTER_IDS.ERROR_CODE]: "Error Code", + [LOG_FILTER_IDS.ERROR_MESSAGE]: "Error Message", + [LOG_FILTER_IDS.KEY_HASH]: "Key Hash", + [LOG_FILTER_IDS.SESSION_ID]: "Session ID", + [LOG_FILTER_IDS.MODEL_ID]: "Model", + [LOG_FILTER_IDS.PUBLIC_MODEL_OR_SEARCH_TOOL]: "Public model / search tool", +}; -// Keys whose UI is a free-form text input; only these need debouncing. -const TEXT_FILTER_KEYS: readonly (keyof LogFilterState)[] = [ - FILTER_KEYS.KEY_HASH, - FILTER_KEYS.ERROR_MESSAGE, - FILTER_KEYS.REQUEST_ID, - FILTER_KEYS.SESSION_ID, - FILTER_KEYS.USER_ID, - FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL, -]; - -// Live-tail polls every 15s, but only on page 1 (newest) while live tail is on. export const LIVE_TAIL_INTERVAL_MS = 15000; -export const getLiveTailRefetchInterval = (isLiveTail: boolean, currentPage: number): number | false => - isLiveTail && currentPage === 1 ? LIVE_TAIL_INTERVAL_MS : false; -export const defaultFilters: LogFilterState = { - [FILTER_KEYS.TEAM_ID]: "", - [FILTER_KEYS.KEY_HASH]: "", - [FILTER_KEYS.REQUEST_ID]: "", - [FILTER_KEYS.SESSION_ID]: "", - [FILTER_KEYS.MODEL]: "", - [FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL]: "", - [FILTER_KEYS.USER_ID]: "", - [FILTER_KEYS.END_USER]: "", - [FILTER_KEYS.STATUS]: "", - [FILTER_KEYS.KEY_ALIAS]: "", - [FILTER_KEYS.ERROR_CODE]: "", - [FILTER_KEYS.ERROR_MESSAGE]: "", +export const getLiveTailRefetchInterval = (isLiveTail: boolean, pageIndex: number): number | false => + isLiveTail && pageIndex === 0 ? LIVE_TAIL_INTERVAL_MS : false; + +export const DEFAULT_LOGS_SORTING: SortingState = [{ id: "startTime", desc: true }]; + +const isSortField = (id: string): id is LogsSortField => Object.hasOwn(LOGS_SORT_FIELD_MAP, id); + +export const getFilterValue = (columnFilters: ColumnFiltersState, columnId: string): string | undefined => { + const entry = columnFilters.find((filter) => filter.id === columnId); + if (typeof entry?.value !== "string") return undefined; + const trimmed = entry.value.trim(); + return trimmed === "" ? undefined : trimmed; }; export function useLogFilterLogic({ @@ -73,63 +65,45 @@ export function useLogFilterLogic({ token, userRole, userID, - filters, - setFilters, + columnFilters, filterByCurrentUser, activeTab, isLiveTail, startTime, endTime, - pageSize = defaultPageSize, + pagination, isCustomDate, - setCurrentPage, - sortBy = "startTime", - sortOrder = "desc", - currentPage = 1, + sorting, }: { accessToken: string | null; token: string | null; userRole: string | null; userID: string | null; - filters: LogFilterState; - setFilters: React.Dispatch>; + columnFilters: ColumnFiltersState; filterByCurrentUser: boolean | null; activeTab: string; isLiveTail: boolean; startTime: string; endTime: string; - pageSize?: number; + pagination: PaginationState; isCustomDate: boolean; - setCurrentPage: (page: number) => void; - sortBy?: LogsSortField; - sortOrder?: "asc" | "desc"; - currentPage?: number; + sorting: SortingState; }) { - const [debouncedFilters, setDebouncedFilters] = useState(filters); - const debouncer = useDebouncer(setDebouncedFilters, { wait: DEBOUNCE_WAIT_MS }); - useEffect(() => { - debouncer.maybeExecute(filters); - }, [filters, debouncer]); + const pageSize = pagination.pageSize || defaultPageSize; + const activeSort = sorting[0] ?? DEFAULT_LOGS_SORTING[0]; + const sortBy: LogsSortField = isSortField(activeSort.id) ? activeSort.id : "startTime"; + const sortOrder: "asc" | "desc" = activeSort.desc ? "desc" : "asc"; - // Live values for dropdown keys, debounced for text keys. - const effectiveFilters = useMemo(() => { - const merged = { ...filters }; - for (const k of TEXT_FILTER_KEYS) { - merged[k] = debouncedFilters[k]; - } - return merged; - }, [filters, debouncedFilters]); - - const logsQuery = useQuery({ + const logsQueryOptions: UseQueryOptions = { queryKey: [ "logs", "table", - currentPage, + pagination.pageIndex, pageSize, startTime, endTime, isCustomDate, - effectiveFilters, + columnFilters, filterByCurrentUser ? userID : null, sortBy, sortOrder, @@ -150,38 +124,39 @@ export function useLogFilterLogic({ ? moment(endTime).utc().format("YYYY-MM-DD HH:mm:ss") : moment().utc().format("YYYY-MM-DD HH:mm:ss"); - const response = await uiSpendLogsCall({ + const userIdFilter = getFilterValue(columnFilters, LOG_FILTER_IDS.USER_ID); + + return await uiSpendLogsCall({ accessToken, start_date: formattedStartTime, end_date: formattedEndTime, - page: currentPage, + page: pagination.pageIndex + 1, page_size: pageSize, params: { - api_key: effectiveFilters[FILTER_KEYS.KEY_HASH] || undefined, - team_id: effectiveFilters[FILTER_KEYS.TEAM_ID] || undefined, - request_id: effectiveFilters[FILTER_KEYS.REQUEST_ID] || undefined, - session_id: effectiveFilters[FILTER_KEYS.SESSION_ID] || undefined, - user_id: effectiveFilters[FILTER_KEYS.USER_ID] || (filterByCurrentUser ? userID ?? undefined : undefined), - end_user: effectiveFilters[FILTER_KEYS.END_USER] || undefined, - status_filter: effectiveFilters[FILTER_KEYS.STATUS] || undefined, - model_id: effectiveFilters[FILTER_KEYS.MODEL] || undefined, - model: effectiveFilters[FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL] || undefined, - key_alias: effectiveFilters[FILTER_KEYS.KEY_ALIAS] || undefined, - error_code: effectiveFilters[FILTER_KEYS.ERROR_CODE] || undefined, - error_message: effectiveFilters[FILTER_KEYS.ERROR_MESSAGE] || undefined, + api_key: getFilterValue(columnFilters, LOG_FILTER_IDS.KEY_HASH), + team_id: getFilterValue(columnFilters, LOG_FILTER_IDS.TEAM_ID), + request_id: getFilterValue(columnFilters, LOG_FILTER_IDS.REQUEST_ID), + session_id: getFilterValue(columnFilters, LOG_FILTER_IDS.SESSION_ID), + user_id: userIdFilter ?? (filterByCurrentUser ? userID ?? undefined : undefined), + end_user: getFilterValue(columnFilters, LOG_FILTER_IDS.END_USER), + status_filter: getFilterValue(columnFilters, LOG_FILTER_IDS.STATUS), + model_id: getFilterValue(columnFilters, LOG_FILTER_IDS.MODEL_ID), + model: getFilterValue(columnFilters, LOG_FILTER_IDS.PUBLIC_MODEL_OR_SEARCH_TOOL), + key_alias: getFilterValue(columnFilters, LOG_FILTER_IDS.KEY_ALIAS), + error_code: getFilterValue(columnFilters, LOG_FILTER_IDS.ERROR_CODE), + error_message: getFilterValue(columnFilters, LOG_FILTER_IDS.ERROR_MESSAGE), sort_by: sortBy, sort_order: sortOrder, }, }); - - return response; }, enabled: !!accessToken && !!token && !!userRole && !!userID && activeTab === "request logs", - refetchInterval: getLiveTailRefetchInterval(isLiveTail, currentPage), + refetchInterval: getLiveTailRefetchInterval(isLiveTail, pagination.pageIndex), placeholderData: keepPreviousData, - // Only live-tail-poll while the tab is visible. refetchIntervalInBackground: false, - }); + }; + + const logsQuery = useQuery(logsQueryOptions); const filteredLogs: PaginatedResponse = logsQuery.data ?? { data: [], @@ -191,7 +166,7 @@ export function useLogFilterLogic({ total_pages: 0, }; - const { data: allTeams } = useQuery({ + const allTeamsQueryOptions: UseQueryOptions = { queryKey: ["allTeamsForLogFilters", accessToken], queryFn: async () => { if (!accessToken) return []; @@ -199,34 +174,13 @@ export function useLogFilterLogic({ return teamsData || []; }, enabled: !!accessToken, - }); - - const handleFilterChange = (newFilters: Partial) => { - setFilters((prev) => { - const updatedFilters = { ...prev, ...newFilters }; - for (const key of Object.keys(defaultFilters) as Array) { - if (!(key in updatedFilters)) { - updatedFilters[key] = defaultFilters[key]; - } - } - if (JSON.stringify(updatedFilters) !== JSON.stringify(prev)) { - setCurrentPage(1); - } - return updatedFilters as LogFilterState; - }); }; - const handleFilterReset = () => { - setFilters(defaultFilters); - setDebouncedFilters(defaultFilters); - setCurrentPage(1); - }; + const { data: allTeams } = useQuery(allTeamsQueryOptions); return { logsQuery, filteredLogs, allTeams, - handleFilterChange, - handleFilterReset, }; } From 6ca40e7dcd5e0e5f337067526ac82aef305e2ca9 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Thu, 23 Jul 2026 10:41:15 -0700 Subject: [PATCH 09/25] refactor(mcp): delete unreachable v1 OBO handler and gate REST OAuth on v2 resolver The v2 credential resolver owns oauth2_token_exchange end to end: any server with a token-exchange config maps to a non-None TokenExchangeConfig spec, and that config is in _create_mcp_client's override-exclusion set, so a caller x-mcp-* override cannot force it back to v1 either. The v1 handler resolve_mcp_auth reached at spec is None was therefore dead for OBO, including its warn-then-proceed-unauthenticated fall-through. Delete auth/token_exchange.py and the exchange branch, dropping the subject_token parameter that only fed it. Separately, the REST listing and call paths still ran the v1 per-user OAuth lookup for servers the v2 resolver owns. Unlike the two protocol-path call sites they gated on auth_type == oauth2 only, with no to_server_spec check, so a migrated authorization_code server did a DB round-trip whose Authorization header _resolve_v2_auth then discards. Add the same guard via _is_v1_resolved_oauth2_server, shared by the per-server lookup and the prefetch preflight. Also collapses MCPOAuth2TokenCache.async_get_token's now single-caller require_client_credentials_flow kwarg and removes the dead _get_bulk_user_oauth_headers helper (zero callers). --- .../mcp_server/auth/token_exchange.py | 192 ------- .../mcp_server/mcp_server_manager.py | 4 +- .../mcp_server/oauth2_token_cache.py | 34 +- .../mcp_server/rest_endpoints.py | 65 +-- .../types/mcp_server/mcp_server_manager.py | 9 - .../mcp_server/auth/test_token_exchange.py | 539 ------------------ .../mcp_server/test_mcp_server_manager.py | 53 ++ .../mcp_server/test_mcp_tool_search.py | 4 +- .../mcp_server/test_rest_endpoints.py | 65 +++ 9 files changed, 151 insertions(+), 814 deletions(-) delete mode 100644 litellm/proxy/_experimental/mcp_server/auth/token_exchange.py delete mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/auth/test_token_exchange.py diff --git a/litellm/proxy/_experimental/mcp_server/auth/token_exchange.py b/litellm/proxy/_experimental/mcp_server/auth/token_exchange.py deleted file mode 100644 index cd41dd648ee..00000000000 --- a/litellm/proxy/_experimental/mcp_server/auth/token_exchange.py +++ /dev/null @@ -1,192 +0,0 @@ -""" -OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers. - -Exchanges a user's incoming JWT (subject_token) for a scoped access token -at an IDP's token exchange endpoint. The exchanged token is then used to -authenticate requests to the upstream MCP server. - -See: https://datatracker.ietf.org/doc/html/rfc8693 -""" - -import asyncio -import hashlib -import weakref -from typing import TYPE_CHECKING, Dict, Tuple - -import httpx - -from litellm._logging import verbose_logger -from litellm.caching.in_memory_cache import InMemoryCache -from litellm.constants import ( - MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, - MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, - MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE, -) -from litellm.llms.custom_httpx.http_handler import get_async_httpx_client -from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( - build_token_endpoint_client_auth, -) -from litellm.types.llms.custom_http import httpxSpecialProvider -from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE - -if TYPE_CHECKING: - from litellm.types.mcp_server.mcp_server_manager import MCPServer - -# RFC 8693 grant type constant -TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange" - - -class TokenExchangeHandler: - """Handles OAuth 2.0 Token Exchange (RFC 8693) for MCP servers. - - Caches exchanged tokens keyed by ``hash(subject_token + server_id)`` so - repeated calls with the same user token skip the IDP round-trip. - """ - - def __init__(self) -> None: - self._cache = InMemoryCache( - max_size_in_memory=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE, - default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, - ) - # WeakValueDictionary so locks are GC'd once no coroutine holds a reference, - # preventing unbounded growth with many rotating user tokens. - self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = weakref.WeakValueDictionary() - - def _get_lock(self, cache_key: str) -> asyncio.Lock: - lock = self._locks.get(cache_key) - if lock is None: - lock = asyncio.Lock() - self._locks[cache_key] = lock - return lock - - @staticmethod - def _cache_key(subject_token: str, server_id: str) -> str: - raw = f"{subject_token}:{server_id}" - return hashlib.sha256(raw.encode()).hexdigest() - - async def exchange_token( - self, - subject_token: str, - server: "MCPServer", - ) -> str: - """Exchange *subject_token* for a scoped access token. - - Returns the exchanged ``access_token`` string (suitable for a - ``Bearer`` header). - - Raises ``ValueError`` on configuration or IDP errors. - """ - cache_key = self._cache_key(subject_token, server.server_id) - - # Fast path - cached = self._cache.get_cache(cache_key) - if cached is not None: - return cached - - # Slow path — one exchange at a time per (user, server) pair - async with self._get_lock(cache_key): - cached = self._cache.get_cache(cache_key) - if cached is not None: - return cached - - token, ttl = await self._do_exchange(subject_token, server) - self._cache.set_cache(cache_key, token, ttl=ttl) - return token - - async def _do_exchange( - self, - subject_token: str, - server: "MCPServer", - ) -> Tuple[str, int]: - """POST to the token exchange endpoint with RFC 8693 parameters. - - Returns ``(access_token, ttl_seconds)``. - """ - endpoint = server.token_exchange_endpoint or server.token_url - if not endpoint: - raise ValueError( - f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange " - f"but no token_exchange_endpoint or token_url configured" - ) - if not server.client_id or not server.client_secret: - raise ValueError( - f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange " - f"but missing client_id or client_secret" - ) - - client_auth = build_token_endpoint_client_auth( - auth_method=server.token_endpoint_auth_method, - client_id=server.client_id, - client_secret=server.client_secret, - ) - data: Dict[str, str] = { - "grant_type": TOKEN_EXCHANGE_GRANT_TYPE, - "subject_token": subject_token, - "subject_token_type": server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE, - **client_auth.body, - } - if server.audience: - data["audience"] = server.audience - if server.scopes: - data["scope"] = " ".join(server.scopes) - - verbose_logger.debug( - "Exchanging token for MCP server %s at %s (audience=%s)", - server.server_id, - endpoint, - server.audience, - ) - - client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) - post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})} - try: - response = await client.post(endpoint, **post_kwargs) - response.raise_for_status() - except httpx.HTTPStatusError as exc: - verbose_logger.debug( - "Token exchange IDP error for MCP server %s (status %d)", - server.server_id, - exc.response.status_code, - ) - raise ValueError( - f"Token exchange for MCP server '{server.server_id}' failed with status {exc.response.status_code}" - ) from exc - - body = response.json() - if not isinstance(body, dict): - raise ValueError( - f"Token exchange response for MCP server '{server.server_id}' " - f"returned non-object JSON (got {type(body).__name__})" - ) - - access_token = body.get("access_token") - if not access_token: - raise ValueError(f"Token exchange response for MCP server '{server.server_id}' missing 'access_token'") - - raw_expires_in = body.get("expires_in") - try: - expires_in = int(raw_expires_in) if raw_expires_in is not None else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL - except (TypeError, ValueError): - expires_in = MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL - - ttl = max( - expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, - MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, - ) - - verbose_logger.info( - "Token exchange succeeded for MCP server %s (expires in %ds)", - server.server_id, - expires_in, - ) - return access_token, ttl - - def invalidate(self, subject_token: str, server_id: str) -> None: - """Remove a cached exchanged token (e.g. after a 401).""" - cache_key = self._cache_key(subject_token, server_id) - self._cache.delete_cache(cache_key) - - -# Module-level singleton -mcp_token_exchange_handler = TokenExchangeHandler() diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b442ea5de70..0ee74960293 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3086,9 +3086,7 @@ class MCPServerManager: ) ): spec = None - auth_value = ( - await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token) if spec is None else None - ) + auth_value = await resolve_mcp_auth(server, mcp_auth_header) if spec is None else None # Create sampling and elicitation callbacks for this client sampling_cb = _create_sampling_callback(user_api_key_auth=user_api_key_auth) if server.allow_sampling else None diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index 43fe3999291..a6acaf8e1d6 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -26,7 +26,6 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy._experimental.mcp_server.auth import token_exchange from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( build_token_endpoint_client_auth, ) @@ -58,17 +57,12 @@ class MCPOAuth2TokenCache(InMemoryCache): def _has_client_credentials_config(server: "MCPServer") -> bool: return bool(server.client_id and server.client_secret and server.token_url) - async def async_get_token( - self, - server: "MCPServer", - *, - require_client_credentials_flow: bool = True, - ) -> Optional[str]: + async def async_get_token(self, server: "MCPServer") -> Optional[str]: """Return a valid access token, fetching or refreshing as needed. Returns ``None`` when the server lacks client credentials config. """ - if require_client_credentials_flow and not server.has_client_credentials: + if not server.has_client_credentials: return None if not self._has_client_credentials_config(server): return None @@ -278,36 +272,16 @@ mcp_per_user_token_cache = MCPPerUserTokenCache() async def resolve_mcp_auth( server: "MCPServer", mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, - subject_token: Optional[str] = None, ) -> Optional[Union[str, Dict[str, str]]]: """Resolve the auth value for an MCP server. Priority: 1. ``mcp_auth_header`` — per-request/per-user override - 2. OAuth2 Token Exchange (OBO / RFC 8693) — exchange user token for scoped token - 3. OAuth2 client_credentials token — auto-fetched and cached - 4. ``server.authentication_token`` — static token from config/DB + 2. OAuth2 client_credentials token — auto-fetched and cached + 3. ``server.authentication_token`` — static token from config/DB """ if mcp_auth_header: return mcp_auth_header - if server.has_token_exchange_config: - if subject_token: - return await token_exchange.mcp_token_exchange_handler.exchange_token(subject_token, server) - # No subject_token — fall back to client_credentials using the same client - # credentials and token_url so M2M scenarios still work. - if server.client_id and server.client_secret and server.token_url: - return await mcp_oauth2_token_cache.async_get_token( - server, - require_client_credentials_flow=False, - ) - # OBO configured but no subject_token and missing client credentials — warn - # rather than silently proceeding unauthenticated. - verbose_logger.warning( - "MCP server '%s' is configured for token exchange (OBO) but no subject_token " - "was provided and client credentials (client_id/client_secret/token_url) are " - "incomplete. The request will proceed without authentication.", - server.server_id, - ) if server.has_client_credentials: return await mcp_oauth2_token_cache.async_get_token(server) return server.authentication_token diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 94271c54f4b..26e4176e09b 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -230,16 +230,33 @@ if MCP_AVAILABLE: return server_auth return mcp_auth_header - def _get_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]: - """Return the subset of *allowed_server_ids* whose servers use OAuth2 auth. + def _is_v1_resolved_oauth2_server(server: Optional[MCPServer]) -> bool: + """Whether this server's per-user OAuth2 token is still resolved by v1. - Used as a cheap pre-flight check to skip bulk credential fetching when no - OAuth2 servers are involved in the current request. + A server the v2 resolver owns reads its stored token from the resolver at connect + time and drops any Authorization built for it here, so the v1 lookup would be a DB + round-trip whose result is discarded. Mirrors the same guard on the protocol listing + path and in ``_resolve_oauth2_headers_for_tool_call``. + """ + from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( + to_server_spec, + ) + + if getattr(server, "auth_type", None) != MCPAuth.oauth2: + return False + return to_server_spec(server) is None + + def _v1_resolved_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]: + """Return the subset of *allowed_server_ids* whose per-user OAuth2 token is still + resolved by v1. + + Used as a cheap pre-flight check to skip bulk credential fetching when no such + server is involved in the current request. """ return { sid for sid in allowed_server_ids - if getattr(global_mcp_server_manager.get_mcp_server_by_id(sid), "auth_type", None) == MCPAuth.oauth2 + if _is_v1_resolved_oauth2_server(global_mcp_server_manager.get_mcp_server_by_id(sid)) } async def _get_user_oauth_extra_headers( @@ -253,11 +270,13 @@ if MCP_AVAILABLE: the MCP server the same way the admin "Add MCP / Authorize and Fetch" flow does. Returns None for non-OAuth2 servers or when no credential is stored. + A server the v2 resolver owns is skipped; see ``_is_v1_resolved_oauth2_server``. + Args: prefetched_creds: Optional dict keyed by server_id with credential payloads. When provided, avoids a per-server DB round-trip. """ - if getattr(server, "auth_type", None) != MCPAuth.oauth2: + if not _is_v1_resolved_oauth2_server(server): return None user_id = getattr(user_api_key_dict, "user_id", None) server_id = getattr(server, "server_id", None) @@ -320,38 +339,6 @@ if MCP_AVAILABLE: verbose_logger.warning(f"_prefetch_user_oauth_creds: failed to prefetch for user={user_id}: {e}") return {} - async def _get_bulk_user_oauth_headers( - user_api_key_dict: UserAPIKeyAuth, - ) -> Dict[str, Dict[str, str]]: - """ - Fetch ALL OAuth2 credentials for the current user in a single DB query and - return a mapping of server_id → {"Authorization": "Bearer "}. - - This is the batch alternative to calling _get_user_oauth_extra_headers - per-server inside a loop (N+1 DB queries). - """ - user_id = getattr(user_api_key_dict, "user_id", None) - if not user_id: - return {} - try: - from litellm.proxy._experimental.mcp_server.db import ( - list_user_oauth_credentials, - ) - from litellm.proxy.utils import get_prisma_client_or_throw - - prisma_client = get_prisma_client_or_throw( - "Database not connected. Connect a database to use OAuth2 MCP tools." - ) - creds = await list_user_oauth_credentials(prisma_client, user_id) - return { - c["server_id"]: {"Authorization": f"Bearer {c['access_token']}"} - for c in creds - if c.get("access_token") and c.get("server_id") - } - except Exception: - verbose_logger.debug("Failed to bulk-fetch OAuth credentials", exc_info=True) - return {} - def _create_tool_response_objects(tools, server: MCPServer): """Helper function to create tool response objects. @@ -825,7 +812,7 @@ if MCP_AVAILABLE: # to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers. prefetched_oauth_creds = ( await _prefetch_user_oauth_creds(user_api_key_dict) - if _get_oauth2_server_ids(allowed_server_ids) + if _v1_resolved_oauth2_server_ids(allowed_server_ids) else {} ) diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 8ae974b19a6..b0af22e7c3f 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -261,12 +261,3 @@ class MCPServer(BaseModel): if self.oauth_passthrough is not True: return False return any(h.lower() == "authorization" for h in self.extra_headers) - - @property - def has_token_exchange_config(self) -> bool: - """True if this server is configured for OAuth2 token exchange (OBO / RFC 8693).""" - return ( - self.auth_type == MCPAuth.oauth2_token_exchange - and bool(self.client_id and self.client_secret) - and bool(self.token_exchange_endpoint or self.token_url) - ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_token_exchange.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_token_exchange.py deleted file mode 100644 index d2aa58e29ea..00000000000 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_token_exchange.py +++ /dev/null @@ -1,539 +0,0 @@ -""" -Tests for OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers. - -Covers: exchange flow, caching, error handling, resolve_mcp_auth integration, -bearer token extraction, and config loading. -""" - -from unittest.mock import AsyncMock, MagicMock, patch - -import httpx -import pytest - -from litellm.proxy._experimental.mcp_server.auth.token_exchange import ( - TOKEN_EXCHANGE_GRANT_TYPE, - TokenExchangeHandler, -) -from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - MCPServerManager, -) -from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - resolve_mcp_auth, -) -from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport -from litellm.types.mcp import MCPAuth -from litellm.types.mcp_server.mcp_server_manager import MCPServer - - -def _obo_server(**overrides) -> MCPServer: - defaults = dict( - server_id="srv-obo-1", - name="test-obo", - url="https://mcp.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - client_id="litellm-client-id", - client_secret="litellm-client-secret", - token_exchange_endpoint="https://idp.example.com/oauth2/token", - audience="api://mcp-server", - scopes=["mcp.tools.read", "mcp.tools.execute"], - ) - defaults.update(overrides) - return MCPServer(**defaults) - - -def _exchange_response(token="exchanged-tok-abc", expires_in=3600): - resp = MagicMock() - resp.json.return_value = { - "access_token": token, - "token_type": "Bearer", - "expires_in": expires_in, - } - resp.raise_for_status = MagicMock() - resp.text = "" - return resp - - -# ── Exchange Flow ── - - -@pytest.mark.asyncio -async def test_exchange_token_success(): - """Token exchange sends correct RFC 8693 parameters and returns access_token.""" - handler = TokenExchangeHandler() - server = _obo_server() - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("scoped-token-1") - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - result = await handler.exchange_token("user-jwt-xyz", server) - - assert result == "scoped-token-1" - mock_client.post.assert_called_once() - - _, kwargs = mock_client.post.call_args - data = kwargs["data"] - assert data["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE - assert data["subject_token"] == "user-jwt-xyz" - assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:access_token" - assert data["audience"] == "api://mcp-server" - assert data["scope"] == "mcp.tools.read mcp.tools.execute" - assert data["client_id"] == "litellm-client-id" - assert data["client_secret"] == "litellm-client-secret" - - -@pytest.mark.asyncio -async def test_exchange_token_no_audience(): - """When audience is None, it is omitted from the request.""" - handler = TokenExchangeHandler() - server = _obo_server(audience=None) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response() - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - await handler.exchange_token("user-jwt", server) - - _, kwargs = mock_client.post.call_args - assert "audience" not in kwargs["data"] - - -@pytest.mark.asyncio -async def test_exchange_token_no_scopes(): - """When scopes is None, scope param is omitted from the request.""" - handler = TokenExchangeHandler() - server = _obo_server(scopes=None) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response() - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - await handler.exchange_token("user-jwt", server) - - _, kwargs = mock_client.post.call_args - assert "scope" not in kwargs["data"] - - -# ── Caching ── - - -@pytest.mark.asyncio -async def test_exchange_token_cached(): - """Second call with same user token uses cache — only 1 HTTP POST.""" - handler = TokenExchangeHandler() - server = _obo_server() - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("cached-exchange-tok") - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - t1 = await handler.exchange_token("same-jwt", server) - t2 = await handler.exchange_token("same-jwt", server) - - assert t1 == t2 == "cached-exchange-tok" - assert mock_client.post.call_count == 1 - - -@pytest.mark.asyncio -async def test_different_user_tokens_not_shared(): - """Different user JWTs get different exchanged tokens.""" - handler = TokenExchangeHandler() - server = _obo_server() - call_count = 0 - - async def mock_post(url, data=None): - nonlocal call_count - call_count += 1 - resp = MagicMock() - resp.json.return_value = { - "access_token": f"exchanged-{call_count}", - "expires_in": 3600, - } - resp.raise_for_status = MagicMock() - return resp - - mock_client = AsyncMock() - mock_client.post = mock_post - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - t1 = await handler.exchange_token("user-a-jwt", server) - t2 = await handler.exchange_token("user-b-jwt", server) - - assert t1 == "exchanged-1" - assert t2 == "exchanged-2" - assert call_count == 2 - - -# ── Error Handling ── - - -@pytest.mark.asyncio -async def test_exchange_token_http_error(): - """HTTP errors from the IDP are wrapped in a ValueError.""" - handler = TokenExchangeHandler() - server = _obo_server() - mock_response = MagicMock() - mock_response.status_code = 400 - mock_response.text = "invalid_grant" - mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( - "Bad Request", - request=MagicMock(), - response=mock_response, - ) - mock_client = AsyncMock() - mock_client.post.return_value = mock_response - - with ( - patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ), - pytest.raises(ValueError, match="failed with status 400"), - ): - await handler.exchange_token("bad-jwt", server) - - -@pytest.mark.asyncio -async def test_exchange_token_http_error_does_not_log_response_body(): - """Raw IDP error bodies are not logged because they can contain credentials.""" - handler = TokenExchangeHandler() - server = _obo_server() - raw_response_body = "client_secret=do-not-log" - mock_response = MagicMock() - mock_response.status_code = 401 - mock_response.text = raw_response_body - mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( - "Unauthorized", - request=MagicMock(), - response=mock_response, - ) - mock_client = AsyncMock() - mock_client.post.return_value = mock_response - - with ( - patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ), - patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.verbose_logger.debug" - ) as mock_debug, - pytest.raises(ValueError, match="failed with status 401"), - ): - await handler.exchange_token("bad-jwt", server) - - logged_values = " ".join( - str(value) - for call in mock_debug.call_args_list - for value in [*call.args, *call.kwargs.values()] - ) - assert raw_response_body not in logged_values - - -@pytest.mark.asyncio -async def test_exchange_token_missing_access_token(): - """Response without access_token raises ValueError.""" - handler = TokenExchangeHandler() - server = _obo_server() - resp = MagicMock() - resp.json.return_value = {"token_type": "Bearer"} - resp.raise_for_status = MagicMock() - mock_client = AsyncMock() - mock_client.post.return_value = resp - - with ( - patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ), - pytest.raises(ValueError, match="missing 'access_token'"), - ): - await handler.exchange_token("jwt", server) - - -@pytest.mark.asyncio -async def test_exchange_token_missing_endpoint(): - """Missing token_exchange_endpoint and token_url raises ValueError.""" - handler = TokenExchangeHandler() - server = _obo_server(token_exchange_endpoint=None, token_url=None) - - with pytest.raises(ValueError, match="no token_exchange_endpoint or token_url"): - await handler.exchange_token("jwt", server) - - -@pytest.mark.asyncio -async def test_exchange_token_missing_credentials(): - """Missing client_id or client_secret raises ValueError.""" - handler = TokenExchangeHandler() - server = _obo_server(client_id=None, client_secret=None) - # has_token_exchange_config will be False, so we call _do_exchange directly - with pytest.raises(ValueError, match="missing client_id or client_secret"): - await handler._do_exchange("jwt", server) - - -# ── resolve_mcp_auth Integration ── - - -@pytest.mark.asyncio -async def test_resolve_mcp_auth_with_token_exchange(): - """resolve_mcp_auth delegates to token exchange when server has OBO config and subject_token provided.""" - server = _obo_server() - mock_handler = AsyncMock() - mock_handler.exchange_token.return_value = "obo-scoped-token" - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.mcp_token_exchange_handler", - mock_handler, - ): - result = await resolve_mcp_auth(server, subject_token="user-jwt") - - assert result == "obo-scoped-token" - mock_handler.exchange_token.assert_called_once_with("user-jwt", server) - - -@pytest.mark.asyncio -async def test_resolve_mcp_auth_obo_without_subject_token_falls_through(): - """Without a subject_token, resolve_mcp_auth falls through to client_credentials.""" - server = _obo_server( - token_url="https://auth.example.com/token", - ) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("cc-token") - - with patch( - "litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client", - return_value=mock_client, - ): - result = await resolve_mcp_auth(server, subject_token=None) - - # Falls through to client_credentials since subject_token is None - # The server has client_id/client_secret/token_url so has_client_credentials is True - assert result == "cc-token" - - -@pytest.mark.asyncio -async def test_resolve_mcp_auth_obo_without_subject_token_uses_cached_client_credentials(): - """The M2M fallback for OBO servers reuses the client_credentials cache.""" - server = _obo_server( - server_id="srv-obo-m2m-cache", - token_url="https://auth.example.com/token", - ) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("cached-cc-token") - - with patch( - "litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client", - return_value=mock_client, - ): - first = await resolve_mcp_auth(server, subject_token=None) - second = await resolve_mcp_auth(server, subject_token=None) - - assert first == second == "cached-cc-token" - mock_client.post.assert_called_once() - - -@pytest.mark.asyncio -async def test_resolve_mcp_auth_header_beats_obo(): - """An explicit mcp_auth_header takes priority over OBO token exchange.""" - server = _obo_server() - result = await resolve_mcp_auth( - server, mcp_auth_header="Bearer override", subject_token="user-jwt" - ) - assert result == "Bearer override" - - -# ── Bearer Token Extraction ── - - -def test_extract_bearer_token_from_oauth2_headers(): - """Extracts token from oauth2_headers Authorization header.""" - result = MCPServerManager._extract_bearer_token( - oauth2_headers={"Authorization": "Bearer my-jwt-token"}, - raw_headers=None, - ) - assert result == "my-jwt-token" - - -def test_extract_bearer_token_from_raw_headers(): - """Falls back to raw_headers when oauth2_headers missing.""" - result = MCPServerManager._extract_bearer_token( - oauth2_headers=None, - raw_headers={"authorization": "Bearer raw-jwt"}, - ) - assert result == "raw-jwt" - - -def test_extract_bearer_token_no_bearer_prefix(): - """Returns token as-is when no Bearer prefix.""" - result = MCPServerManager._extract_bearer_token( - oauth2_headers={"Authorization": "some-opaque-token"}, - raw_headers=None, - ) - assert result == "some-opaque-token" - - -def test_extract_bearer_token_none(): - """Returns None when no auth headers present.""" - result = MCPServerManager._extract_bearer_token( - oauth2_headers=None, - raw_headers=None, - ) - assert result is None - - -# ── MCPServer Properties ── - - -def test_has_token_exchange_config_true(): - """has_token_exchange_config is True for a fully configured OBO server.""" - server = _obo_server() - assert server.has_token_exchange_config is True - - -def test_has_token_exchange_config_false_wrong_auth_type(): - """has_token_exchange_config is False when auth_type is not oauth2_token_exchange.""" - server = _obo_server(auth_type=MCPAuth.oauth2) - assert server.has_token_exchange_config is False - - -def test_has_token_exchange_config_false_missing_creds(): - """has_token_exchange_config is False when client_id/client_secret missing.""" - server = _obo_server(client_id=None) - assert server.has_token_exchange_config is False - - -def test_has_token_exchange_config_uses_token_url_fallback(): - """has_token_exchange_config is True when token_url is set instead of token_exchange_endpoint.""" - server = _obo_server( - token_exchange_endpoint=None, - token_url="https://idp.example.com/token", - ) - assert server.has_token_exchange_config is True - - -# ── Config Loading ── - - -@pytest.mark.asyncio -async def test_config_loading_token_exchange_fields(): - """load_servers_from_config correctly maps OBO config fields to MCPServer.""" - manager = MCPServerManager() - config = { - "my_obo_server": { - "url": "https://mcp.example.com/mcp", - "transport": "http", - "auth_type": "oauth2_token_exchange", - "client_id": "my-client", - "client_secret": "my-secret", - "token_exchange_endpoint": "https://idp.example.com/oauth2/token", - "audience": "api://my-mcp", - "scopes": ["read", "write"], - "subject_token_type": "urn:ietf:params:oauth:token-type:jwt", - } - } - await manager.load_servers_from_config(config) - - servers = list(manager.config_mcp_servers.values()) - assert len(servers) == 1 - - server = servers[0] - assert server.auth_type == MCPAuth.oauth2_token_exchange - assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token" - assert server.audience == "api://my-mcp" - assert server.subject_token_type == "urn:ietf:params:oauth:token-type:jwt" - assert server.client_id == "my-client" - assert server.client_secret == "my-secret" - assert server.scopes == ["read", "write"] - assert server.has_token_exchange_config is True - - -@pytest.mark.asyncio -async def test_config_loading_default_subject_token_type(): - """subject_token_type defaults to access_token when not specified in config.""" - manager = MCPServerManager() - config = { - "obo_defaults": { - "url": "https://mcp.example.com/mcp", - "transport": "http", - "auth_type": "oauth2_token_exchange", - "client_id": "cid", - "client_secret": "csec", - "token_exchange_endpoint": "https://idp.example.com/token", - } - } - await manager.load_servers_from_config(config) - - server = list(manager.config_mcp_servers.values())[0] - assert server.subject_token_type == "urn:ietf:params:oauth:token-type:access_token" - - -@pytest.mark.asyncio -async def test_database_loading_token_exchange_scopes_from_credentials(): - """DB-loaded OBO server credentials retain configured scopes.""" - manager = MCPServerManager() - db_server = LiteLLM_MCPServerTable( - server_id="srv-obo-db", - server_name="obo_db_server", - url="https://mcp.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - credentials={ - "client_id": "db-client", - "client_secret": "db-secret", - "token_exchange_endpoint": "https://idp.example.com/oauth2/token", - "audience": "api://db-mcp", - "scopes": ["db.read", "db.write"], - }, - ) - - server = await manager.build_mcp_server_from_table( - db_server, - credentials_are_encrypted=False, - ) - - assert server.auth_type == MCPAuth.oauth2_token_exchange - assert server.client_id == "db-client" - assert server.client_secret == "db-secret" - assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token" - assert server.audience == "api://db-mcp" - assert server.scopes == ["db.read", "db.write"] - - -@pytest.mark.asyncio -async def test_exchange_token_uses_client_secret_basic_when_configured(): - """LIT-4091: token exchange with token_endpoint_auth_method=client_secret_basic sends the - client credentials as HTTP Basic and omits client_secret from the body.""" - import base64 - - handler = TokenExchangeHandler() - server = _obo_server( - server_id="srv-obo-basic", token_endpoint_auth_method="client_secret_basic" - ) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("scoped-basic") - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - result = await handler.exchange_token("user-jwt-basic", server) - - assert result == "scoped-basic" - _, kwargs = mock_client.post.call_args - expected = "Basic " + base64.b64encode(b"litellm-client-id:litellm-client-secret").decode() - assert kwargs["headers"]["Authorization"] == expected - assert "client_secret" not in kwargs["data"] - assert "client_id" not in kwargs["data"] - assert kwargs["data"]["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE 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 77f072b81d5..42b6cbee1c4 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 @@ -2333,6 +2333,59 @@ class TestMCPServerManager: assert emitted.headers["Authorization"] == "Bearer upstream-token" assert not kwargs["extra_headers"] or "authorization" not in {k.lower() for k in kwargs["extra_headers"]} + @pytest.mark.asyncio + async def test_create_mcp_client_token_exchange_never_falls_back_to_v1(self): + """A configured OBO server is owned end to end by the v2 token_exchange arm, even when the + caller supplies an x-mcp-* override. This is what makes the v1 OBO handler unreachable, so if + it ever defers to v1 again the deleted handler is silently needed back.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + OAuthToken, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import ( + UpstreamCredentialProvider, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + + class _StubExchanger: + def __init__(self): + self.subject_tokens = [] + + async def exchange(self, subject_token, server, config, *, tenant_id=""): + self.subject_tokens.append(subject_token) + return Ok(OAuthToken(access_token="exchanged-token")) + + async def invalidate(self, subject_token, server, config, *, tenant_id=""): + return None + + exchanger = _StubExchanger() + manager = MCPServerManager() + server = MCPServer( + server_id="obo-egress", + name="obo", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + client_id="gateway-client", + client_secret="gateway-secret", + token_exchange_endpoint="https://idp.example.com/oauth2/token", + ) + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth", + new_callable=AsyncMock, + ) as mock_resolve, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls, + ): + await manager._create_mcp_client( + server=server, + mcp_auth_header="Bearer caller-override", + subject_token="eyJ-subject-token", + cred_provider=UpstreamCredentialProvider(token_exchanger=exchanger), + ) + mock_resolve.assert_not_awaited() + assert exchanger.subject_tokens == ["eyJ-subject-token"] + assert self._emitted_authorization(mock_client_cls) == "Bearer exchanged-token" + @staticmethod def _emitted_authorization(mock_client_cls) -> str: kwargs = mock_client_cls.call_args.kwargs diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 1da44029b5c..0e442102e53 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -237,7 +237,7 @@ class TestListToolRestApiWithToolSearch: return_value={}, ), patch( - "litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids", + "litellm.proxy._experimental.mcp_server.rest_endpoints._v1_resolved_oauth2_server_ids", return_value=[], ), patch( @@ -316,7 +316,7 @@ class TestListToolRestApiWithToolSearch: return_value={}, ), patch( - "litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids", + "litellm.proxy._experimental.mcp_server.rest_endpoints._v1_resolved_oauth2_server_ids", return_value=[], ), patch( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index d4ba66c4381..5c9612a055e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -2783,3 +2783,68 @@ class TestRestListToolsetFiltering: ) assert [tool.name for tool in result] == ["lookup_status"] + + +class TestV1ResolvedOauth2Gate: + """The REST surface must stop resolving per-user OAuth2 tokens for servers the v2 resolver owns. + + ``_resolve_v2_auth`` drops any Authorization built here for an ``authorization_code`` server and + injects the resolver's own token, so the v1 lookup was a DB round-trip whose result was discarded. + A server that still defers to v1 (upstream-delegated oauth2) must keep resolving, which is what + makes these assertions non-vacuous. + """ + + @staticmethod + def _oauth2_server(*, delegate_auth_to_upstream: bool) -> Any: + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + return MCPServer( + server_id="oauth2-srv", + name="oauth2-srv", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=delegate_auth_to_upstream, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "delegate_auth_to_upstream, expected_headers, expected_lookups", + [ + (False, None, 0), + (True, {"Authorization": "Bearer stored-token"}, 1), + ], + ) + async def test_user_oauth_headers_skip_v2_owned_servers( + self, delegate_auth_to_upstream, expected_headers, expected_lookups, monkeypatch + ): + from litellm.proxy._experimental.mcp_server import db as mcp_db + + server = self._oauth2_server(delegate_auth_to_upstream=delegate_auth_to_upstream) + resolve_token = AsyncMock(return_value={"access_token": "stored-token"}) + monkeypatch.setattr(mcp_db, "resolve_valid_user_oauth_token", resolve_token) + + headers = await rest_endpoints._get_user_oauth_extra_headers( + server, + UserAPIKeyAuth(user_id="alice", api_key="sk-1234"), + prefetched_creds={"oauth2-srv": {"access_token": "stored-token"}}, + ) + + assert headers == expected_headers + assert resolve_token.await_count == expected_lookups + + def test_prefetch_preflight_only_counts_v1_resolved_servers(self, monkeypatch): + v2_owned = self._oauth2_server(delegate_auth_to_upstream=False) + v1_resolved = self._oauth2_server(delegate_auth_to_upstream=True) + v1_resolved.server_id = "delegate-srv" + registry = {"oauth2-srv": v2_owned, "delegate-srv": v1_resolved} + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: registry.get(server_id), + ) + + assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv"]) == set() + assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv", "delegate-srv"]) == {"delegate-srv"} From 1b2a7ce5185b76433b2df7d0a8de6c2067b02942 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 23 Jul 2026 10:45:04 -0700 Subject: [PATCH 10/25] feat(ui): rebuild Organization Settings on react-hook-form + zod with a dirty-field PATCH (#34324) * feat(ui): rebuild Organization Settings on react-hook-form + zod with a dirty-field PATCH Replaces the antd Settings form in organization_view.tsx with OrgSettingsForm, the first consumer of the shared RHF + zod form kit. The form derives a minimal payload from RHF dirty tracking via pickDirty and sends it to the typed PATCH /v2/organization/{organization_id}, so untouched fields are omitted, emptied widgets clear with null ([] for lists), and the old full-send builder with its length > 0 clear-dropping guards is deleted. Adds src/lib/forms/useZodForm.ts so every form gets the z.input/z.output generics and zodResolver wiring from one place, and forwardRefs ui/textarea so RHF can register it under React 18 * fix(ui): forwardRef InputGroupTextarea to match the forwardRef'd Textarea * test(ui): pin that an mcp server edit preserves existing org toolsets * docs(ui): explain the useZodForm generics * chore(ui): re-prune eslint suppressions after rebase onto staging --- ui/litellm-dashboard/eslint-suppressions.json | 3 - .../src/components/networking.tsx | 1 + .../org-settings/OrgSettingsForm.test.tsx | 230 ++++++++++++++++++ .../org-settings/OrgSettingsForm.tsx | 169 +++++++++++++ .../organization/org-settings/mapper.test.ts | 108 ++++++++ .../organization/org-settings/mapper.ts | 73 ++++++ .../organization/org-settings/schema.ts | 43 ++++ .../organization/organization_view.tsx | 172 +------------ .../src/components/ui/input-group.tsx | 28 ++- .../src/components/ui/textarea.tsx | 10 +- .../src/lib/forms/useZodForm.test.tsx | 46 ++++ .../src/lib/forms/useZodForm.ts | 18 ++ 12 files changed, 720 insertions(+), 181 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/organization/org-settings/OrgSettingsForm.test.tsx create mode 100644 ui/litellm-dashboard/src/components/organization/org-settings/OrgSettingsForm.tsx create mode 100644 ui/litellm-dashboard/src/components/organization/org-settings/mapper.test.ts create mode 100644 ui/litellm-dashboard/src/components/organization/org-settings/mapper.ts create mode 100644 ui/litellm-dashboard/src/components/organization/org-settings/schema.ts create mode 100644 ui/litellm-dashboard/src/lib/forms/useZodForm.test.tsx create mode 100644 ui/litellm-dashboard/src/lib/forms/useZodForm.ts diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 8e8c1447d22..4fa5528aab8 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -3610,9 +3610,6 @@ }, "no-restricted-imports": { "count": 3 - }, - "unused-imports/no-unused-imports": { - "count": 1 } }, "src/components/page_utils.test.ts": { diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index f5fe3990ba7..5f79b7f95a9 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -212,6 +212,7 @@ export interface Organization { object_permission_id: string; mcp_servers: string[]; mcp_access_groups?: string[]; + mcp_toolsets?: string[]; vector_stores: string[]; }; } diff --git a/ui/litellm-dashboard/src/components/organization/org-settings/OrgSettingsForm.test.tsx b/ui/litellm-dashboard/src/components/organization/org-settings/OrgSettingsForm.test.tsx new file mode 100644 index 00000000000..5bd809bcfd5 --- /dev/null +++ b/ui/litellm-dashboard/src/components/organization/org-settings/OrgSettingsForm.test.tsx @@ -0,0 +1,230 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import React from "react"; +import { describe, expect, it, vi } from "vitest"; + +vi.mock("@/components/molecules/notifications_manager", () => ({ + __esModule: true, + default: { success: vi.fn(), fromBackend: vi.fn() }, +})); +vi.mock("@/components/ModelSelect/ModelSelect", () => ({ + ModelSelect: ({ onChange }: { onChange: (values: string[]) => void }) => ( + + ), +})); +vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({ + __esModule: true, + default: ({ onChange }: { onChange: (values: string[]) => void }) => ( + + ), +})); +vi.mock("@/components/mcp_server_management/MCPServerSelector", () => ({ + __esModule: true, + default: ({ + value, + onChange, + }: { + value?: { servers: string[]; accessGroups: string[]; toolsets?: string[] }; + onChange: (values: { servers: string[]; accessGroups: string[]; toolsets: string[] }) => void; + }) => ( + + ), +})); + +import type { Organization } from "@/components/networking"; + +import { OrgSettingsForm } from "./OrgSettingsForm"; + +const org: Organization = { + organization_id: "org-1", + organization_alias: "acme", + budget_id: "budget-1", + metadata: {}, + models: ["gpt-5.2"], + spend: 0, + model_spend: {}, + created_at: "2026-01-01T00:00:00Z", + created_by: "admin", + updated_at: "2026-01-01T00:00:00Z", + updated_by: "admin", + litellm_budget_table: { max_budget: 100, budget_duration: "30d", tpm_limit: 1000, rpm_limit: 50 }, + teams: null, + users: null, + members: null, + object_permission: { + object_permission_id: "op-1", + mcp_servers: ["srv-1"], + mcp_access_groups: [], + vector_stores: ["vs-1"], + }, +}; + +const renderForm = (overrides?: { + patchOrganization?: ReturnType; + onSaved?: () => void; + org?: Organization; +}) => { + const patchOrganization = overrides?.patchOrganization ?? vi.fn().mockResolvedValue({}); + const onSaved = overrides?.onSaved ?? vi.fn(); + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + render( + + + , + ); + return { patchOrganization, onSaved }; +}; + +describe("OrgSettingsForm", () => { + it("disables Save while the form is pristine", () => { + renderForm(); + + expect(screen.getByRole("button", { name: "Save Changes" })).toBeDisabled(); + }); + + it("sends only the edited field", async () => { + const user = userEvent.setup(); + const { patchOrganization } = renderForm(); + + await user.clear(screen.getByLabelText("Organization Name")); + await user.type(screen.getByLabelText("Organization Name"), "acme-2"); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + await waitFor(() => expect(patchOrganization).toHaveBeenCalledTimes(1)); + expect(patchOrganization).toHaveBeenCalledWith("org-1", { organization_alias: "acme-2" }); + }); + + it("sends null when a limit is cleared", async () => { + const user = userEvent.setup(); + const { patchOrganization } = renderForm(); + + await user.clear(screen.getByLabelText("Tokens per minute Limit (TPM)")); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + await waitFor(() => expect(patchOrganization).toHaveBeenCalledTimes(1)); + expect(patchOrganization).toHaveBeenCalledWith("org-1", { tpm_limit: null }); + }); + + it("sends models as [] when the selector is cleared", async () => { + const user = userEvent.setup(); + const { patchOrganization } = renderForm(); + + await user.click(screen.getByRole("button", { name: "clear-models" })); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + await waitFor(() => expect(patchOrganization).toHaveBeenCalledTimes(1)); + expect(patchOrganization).toHaveBeenCalledWith("org-1", { models: [] }); + }); + + it("wraps a vector store change in object_permission without mcp keys", async () => { + const user = userEvent.setup(); + const { patchOrganization } = renderForm(); + + await user.click(screen.getByRole("button", { name: "set-vector-stores" })); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + await waitFor(() => expect(patchOrganization).toHaveBeenCalledTimes(1)); + expect(patchOrganization).toHaveBeenCalledWith("org-1", { + object_permission: { vector_stores: ["vs-2"] }, + }); + }); + + it("wraps an mcp change in object_permission with all three mcp keys", async () => { + const user = userEvent.setup(); + const { patchOrganization } = renderForm(); + + await user.click(screen.getByRole("button", { name: "set-mcp" })); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + await waitFor(() => expect(patchOrganization).toHaveBeenCalledTimes(1)); + expect(patchOrganization).toHaveBeenCalledWith("org-1", { + object_permission: { mcp_servers: ["srv-2"], mcp_access_groups: [], mcp_toolsets: [] }, + }); + }); + + it("preserves existing toolsets when only the servers change", async () => { + const user = userEvent.setup(); + const orgWithToolsets: Organization = { + ...org, + object_permission: { ...org.object_permission!, mcp_toolsets: ["ts-1"] }, + }; + const { patchOrganization } = renderForm({ org: orgWithToolsets }); + + await user.click(screen.getByRole("button", { name: "set-mcp" })); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + await waitFor(() => expect(patchOrganization).toHaveBeenCalledTimes(1)); + expect(patchOrganization).toHaveBeenCalledWith("org-1", { + object_permission: { mcp_servers: ["srv-2"], mcp_access_groups: [], mcp_toolsets: ["ts-1"] }, + }); + }); + + it("does not send a patch when an edit is reverted to the original value", async () => { + const user = userEvent.setup(); + renderForm(); + + const alias = screen.getByLabelText("Organization Name"); + await user.clear(alias); + await user.type(alias, "acme"); + + await waitFor(() => expect(screen.getByRole("button", { name: "Save Changes" })).toBeDisabled()); + }); + + it("blocks submit and shows an error for invalid metadata JSON", async () => { + const user = userEvent.setup(); + const { patchOrganization } = renderForm(); + + await user.type(screen.getByLabelText("Metadata"), "not json"); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + expect(await screen.findByRole("alert")).toHaveTextContent("Metadata must be a valid JSON object"); + expect(patchOrganization).not.toHaveBeenCalled(); + }); + + it("keeps the view open when the patch fails", async () => { + const user = userEvent.setup(); + const onSaved = vi.fn(); + const { patchOrganization } = renderForm({ + patchOrganization: vi.fn().mockRejectedValue(new Error("boom")), + onSaved, + }); + + await user.clear(screen.getByLabelText("Organization Name")); + await user.type(screen.getByLabelText("Organization Name"), "acme-2"); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + await waitFor(() => expect(patchOrganization).toHaveBeenCalledTimes(1)); + expect(onSaved).not.toHaveBeenCalled(); + }); + + it("calls onSaved after a successful patch", async () => { + const user = userEvent.setup(); + const onSaved = vi.fn(); + renderForm({ onSaved }); + + await user.clear(screen.getByLabelText("Requests per minute Limit (RPM)")); + await user.type(screen.getByLabelText("Requests per minute Limit (RPM)"), "75"); + await user.click(screen.getByRole("button", { name: "Save Changes" })); + + await waitFor(() => expect(onSaved).toHaveBeenCalledTimes(1)); + }); +}); diff --git a/ui/litellm-dashboard/src/components/organization/org-settings/OrgSettingsForm.tsx b/ui/litellm-dashboard/src/components/organization/org-settings/OrgSettingsForm.tsx new file mode 100644 index 00000000000..8e970b88cd2 --- /dev/null +++ b/ui/litellm-dashboard/src/components/organization/org-settings/OrgSettingsForm.tsx @@ -0,0 +1,169 @@ +"use client"; + +import { useMutation, useQueryClient } from "@tanstack/react-query"; +import * as React from "react"; + +import { organizationKeys } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; +import { ModelSelect } from "@/components/ModelSelect/ModelSelect"; +import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector"; +import NotificationsManager from "@/components/molecules/notifications_manager"; +import type { Organization } from "@/components/networking"; +import { FieldGroup } from "@/components/shared/form/field"; +import { FormField } from "@/components/shared/form/FormField"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Textarea } from "@/components/ui/textarea"; +import VectorStoreSelector from "@/components/vector_store_management/VectorStoreSelector"; +import { pickDirty } from "@/lib/forms/pickDirty"; +import { useZodForm } from "@/lib/forms/useZodForm"; +import { fetchClient } from "@/lib/http/api"; + +import { buildOrgPatch, orgToForm, type OrgPatchBody } from "./mapper"; +import { orgSettingsSchema } from "./schema"; + +const NO_RESET = "never"; + +const BUDGET_DURATION_OPTIONS = [ + { value: NO_RESET, label: "No reset" }, + { value: "24h", label: "daily" }, + { value: "7d", label: "weekly" }, + { value: "30d", label: "monthly" }, +] as const; + +const defaultPatchOrganization = async (organizationId: string, body: OrgPatchBody): Promise => { + const { data } = await fetchClient.PATCH("/v2/organization/{organization_id}", { + params: { path: { organization_id: organizationId } }, + body, + }); + return data; +}; + +interface OrgSettingsFormProps { + organizationId: string; + org: Organization; + accessToken: string; + onCancel: () => void; + onSaved: () => void; + patchOrganization?: (organizationId: string, body: OrgPatchBody) => Promise; +} + +export const OrgSettingsForm = ({ + organizationId, + org, + accessToken, + onCancel, + onSaved, + patchOrganization = defaultPatchOrganization, +}: OrgSettingsFormProps) => { + const queryClient = useQueryClient(); + const form = useZodForm(orgSettingsSchema, { defaultValues: orgToForm(org) }); + const { isDirty } = form.formState; + + const mutation = useMutation({ + mutationFn: (body: OrgPatchBody) => patchOrganization(organizationId, body), + onSuccess: () => { + NotificationsManager.success("Organization settings updated successfully"); + queryClient.invalidateQueries({ queryKey: organizationKeys.all }); + onSaved(); + }, + onError: (error: unknown) => + NotificationsManager.fromBackend( + error instanceof Error ? error.message : "Failed to update organization settings", + ), + }); + + const onSubmit = form.handleSubmit((values) => { + mutation.mutate(buildOrgPatch(pickDirty(values, form.formState.dirtyFields))); + }); + + return ( +
+ + + {({ ref, ...field }) => } + + + + {(field) => ( + + )} + + + + {({ ref, ...field }) => } + + + + {({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => ( + + )} + + + + {({ ref, ...field }) => } + + + + {({ ref, ...field }) => } + + + + {(field) => ( + + )} + + + + {(field) => ( + + )} + + + + {({ ref, ...field }) =>