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 2c63d957e75..9c7778e2b77 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 @@ -1947,10 +1947,41 @@ class MCPRequestHandler: 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 [ + grants: Final = [ (source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True))) for source in await MCPRequestHandler.admitted_subject_sources(auth, allowed_team_ids=allowed_team_ids) ] + scope: Final = await MCPRequestHandler._toolset_scope(auth) + if scope is None: + return grants + return [(source, granted & frozenset(scope)) for source, granted in grants] + + @staticmethod + async def _toolset_scope(auth: UserAPIKeyAuth) -> dict[str, list[str]] | None: + """The ``server_id -> tools`` a namespaced toolset route pinned this subject to via + ``mcp_toolset_id``, or None on the aggregate scope. Every source's servers and tools are + intersected with it, so the route narrows a team grant exactly as it narrows the user's own.""" + if auth.mcp_toolset_id is None: + return None + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + return await global_mcp_server_manager.resolve_toolset_tool_permissions( + toolset_ids=[auth.mcp_toolset_id], requires_fresh_policy=auth.requires_fresh_policy + ) + + @staticmethod + async def _narrow_tools_to_toolset( + tools: list[str] | None, + server_id: str, + auth: UserAPIKeyAuth, + ) -> list[str] | None: + scope: Final = await MCPRequestHandler._toolset_scope(auth) + if scope is None: + return tools + scoped: Final = frozenset(scope.get(server_id, ())) + return sorted(scoped if tools is None else scoped & frozenset(tools)) @staticmethod async def resolve_admitted_subject_servers( @@ -2048,9 +2079,9 @@ class MCPRequestHandler: continue tools = await MCPRequestHandler.get_allowed_tools_for_server(server_id, source, keyless_source=True) if tools is None: - return None + return await MCPRequestHandler._narrow_tools_to_toolset(None, server_id, auth) allowed.update(tools) - return sorted(allowed) + return await MCPRequestHandler._narrow_tools_to_toolset(sorted(allowed), server_id, auth) @staticmethod def _get_key_object_permission( @@ -2067,6 +2098,16 @@ class MCPRequestHandler: return user_api_key_auth.object_permission + @staticmethod + async def team_object_permission(user_api_key_auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return await MCPRequestHandler._get_team_object_permission(user_api_key_auth) + + @staticmethod + async def key_object_permission_hydrated( + user_api_key_auth: UserAPIKeyAuth, + ) -> LiteLLM_ObjectPermissionTable | None: + return await MCPRequestHandler._key_object_permission_hydrated(user_api_key_auth) + @staticmethod async def _get_team_object_permission( user_api_key_auth: UserAPIKeyAuth | None = None, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 44523f5cf45..d17b849750b 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -915,11 +915,9 @@ if MCP_AVAILABLE: ) -> UserAPIKeyAuth: """The one credential this tools request acts as. - A toolset name narrows the caller's own credential to that toolset; otherwise a dashboard - session is swapped for its admitted subject. The two are mutually exclusive by construction, - which is why they share an owner: the admitted subject resolves per grant source and a team - source deliberately carries none of the caller's ``object_permission``, so a toolset - narrowing layered on top would evaporate on every team-granted server.""" + A toolset name pins the acting principal to that toolset through ``_apply_toolset_scope``, + which itself swaps a dashboard session for its admitted subject; otherwise the swap happens + here so both shapes resolve as the same identity.""" if not toolset_name: return await acting_user_auth(user_api_key_dict) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 44979d12a00..2090a0c7421 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -66,7 +66,13 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import ( get_route_relative_request_path, well_known_root_suffix, ) -from litellm.proxy._experimental.mcp_server.ui_session_utils import is_ui_session_credential +from litellm.proxy._experimental.mcp_server.ui_session_utils import ( + ActingUser, + GrantedToolsetIds, + acting_user_auth, + granted_toolset_ids, + is_ui_session_credential, +) from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_DESCRIPTION, LITELLM_MCP_SERVER_NAME, @@ -1515,16 +1521,21 @@ if MCP_AVAILABLE: async def _apply_toolset_scope( user_api_key_auth: UserAPIKeyAuth, toolset_id: str, + acting_user: ActingUser = acting_user_auth, + granted: GrantedToolsetIds = granted_toolset_ids, ) -> UserAPIKeyAuth: """ - Restrict a key's MCP permissions to a single toolset. + Pin a principal's MCP permissions to a single toolset for /toolset/{name}/mcp. - When a request arrives via /toolset/{name}/mcp we override the key's - object_permission so that only the toolset's tools are visible. + A virtual key (and an admin session) has its object_permission rewritten to + the toolset's servers and tools. A keyless subject resolves per grant source, + so a non-admin dashboard session first becomes its admitted user and the + toolset rides along as ``mcp_toolset_id``, which every source's grant is + intersected with; a team-granted toolset is served without the user's own + row capping it. - Raises HTTPException(403) if the key has an explicit toolset grant list - that does not include toolset_id (i.e. mcp_toolsets is set but empty, - or set to a list that omits this toolset). Admin keys always pass. + Raises HTTPException(403) unless the principal holds toolset_id through one + of its grant sources. Admins always pass. """ from litellm.proxy._types import LiteLLM_ObjectPermissionTable from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view @@ -1540,26 +1551,31 @@ if MCP_AVAILABLE: detail="API key is scoped to no MCP servers; toolset access is denied.", ) - # Access control: non-admin keys must have this toolset in their grant list. - # Use _user_has_admin_view so that PROXY_ADMIN_VIEW_ONLY is also treated as admin. - is_admin: Final = _user_has_admin_view(user_api_key_auth) - if not is_admin: - op: Final = user_api_key_auth.object_permission - granted: Final = getattr(op, "mcp_toolsets", None) if op else None - # granted=None → key has no explicit toolset grants → deny (same semantics as - # fetch_mcp_toolsets which returns [] for non-admin keys with no grants configured). - # granted=[] or list without toolset_id → also deny. - if granted is None or toolset_id not in granted: + acting: Final = await acting_user(user_api_key_auth) + is_admin: Final = _user_has_admin_view(acting) + if not is_admin and toolset_id not in await granted(acting): + raise HTTPException( + status_code=403, + detail=f"API key does not have access to toolset '{toolset_id}'.", + ) + if _is_mcp_admitted_user_subject(acting): + resource_server_id: Final = acting.mcp_session_resource_server_id + if resource_server_id is not None and resource_server_id not in ( + await operations.global_mcp_server_manager.resolve_toolset_tool_permissions( + toolset_ids=[toolset_id], requires_fresh_policy=acting.requires_fresh_policy + ) + ): raise HTTPException( status_code=403, detail=f"API key does not have access to toolset '{toolset_id}'.", ) + return acting.model_copy(update={"mcp_toolset_id": toolset_id}) tool_permissions = await operations.global_mcp_server_manager.resolve_toolset_tool_permissions( toolset_ids=[toolset_id] ) server_ids: Final = list(tool_permissions.keys()) - existing_op: Final = user_api_key_auth.object_permission + existing_op: Final = acting.object_permission if existing_op is not None: updated_op = existing_op.model_copy( update={ @@ -1576,7 +1592,12 @@ if MCP_AVAILABLE: mcp_servers=server_ids, mcp_tool_permissions=tool_permissions, ) - return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) + return acting.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) + + async def _toolset_server_ids(toolset_id: str) -> set[str]: + return set( + await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[toolset_id]) + ) async def _raise_preemptive_401_for_unauthenticated_servers( scope: Scope, @@ -2007,8 +2028,7 @@ if MCP_AVAILABLE: toolset_allowed_server_ids: set[str] | None = None if active_toolset_id and user_api_key_auth is not None: user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) - op: Final = user_api_key_auth.object_permission - toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set() + toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # Must run after toolset scoping so the challenge set is derived @@ -2355,8 +2375,7 @@ if MCP_AVAILABLE: toolset_allowed_server_ids: set[str] | None = None if active_toolset_id and user_api_key_auth is not None: user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) - op: Final = user_api_key_auth.object_permission - toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set() + toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # Must run after toolset scoping so the challenge set is derived diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py index 901259c18ad..b9e25259868 100644 --- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py +++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py @@ -2,14 +2,23 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable -from typing import Final +import asyncio +from collections.abc import Awaitable, Callable, Sequence +from itertools import chain +from typing import Final, TypeAlias from fastapi import HTTPException from litellm._logging import verbose_logger from litellm.constants import UI_SESSION_TOKEN_TEAM_ID -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth + +EffectiveAuthContexts: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[Sequence[UserAPIKeyAuth]]] +TeamObjectPermission: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[LiteLLM_ObjectPermissionTable | None]] +OwnObjectPermission: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[LiteLLM_ObjectPermissionTable | None]] +AdmittedContext: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[UserAPIKeyAuth | None]] +ActingUser: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[UserAPIKeyAuth]] +GrantedToolsetIds: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[frozenset[str]]] def clone_user_api_key_auth_with_team( @@ -109,10 +118,10 @@ async def acting_user_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth: both surfaces. An admin session keeps its operator view and any caller-passed credential is returned unchanged, never widened. - Do not combine this with a narrowing that rewrites a single credential's ``object_permission`` - (toolset scope): the admitted subject resolves per grant source and a team source deliberately - carries none of the caller's own grants, so the narrowing would silently evaporate on every - team-granted server. A request carrying such a scope keeps the caller's own credential.""" + A toolset narrowing is never applied to the admitted subject by rewriting its ``object_permission``: + it resolves per grant source and a team source deliberately carries none of the caller's own grants, + so the rewrite would evaporate on every team-granted server. The route pins ``mcp_toolset_id`` + instead, which every source's grant is intersected with.""" if not is_ui_session_credential(user_api_key_auth): return user_api_key_auth @@ -153,3 +162,135 @@ async def can_access_mcp_server( if server_id in await allowed_servers(context): return True return False + + +def _restricts_mcp(permission: LiteLLM_ObjectPermissionTable | None) -> bool: + return permission is not None and bool( + permission.mcp_servers + or permission.mcp_toolsets + or permission.mcp_tool_permissions + or permission.mcp_access_groups + ) + + +def is_keyless_mcp_subject(user_api_key_auth: UserAPIKeyAuth) -> bool: + """A principal with no virtual key to declare MCP access on: the dashboard's own session token or a + gateway-admitted user. Its grants are resolved per source, never through a key row.""" + + return is_ui_session_credential(user_api_key_auth) or user_api_key_auth.mcp_admitted_user_subject is True + + +async def toolset_grant_contexts( + user_api_key_auth: UserAPIKeyAuth, + admitted_context: AdmittedContext = admitted_user_context, + admitted_sources: EffectiveAuthContexts | None = None, +) -> Sequence[UserAPIKeyAuth]: + """The grant sources a toolset is looked up through. A virtual key is its own single source. A keyless + subject fans out exactly as the aggregate /mcp resolution does: its own user row plus every team whose + live roster still lists it, so a membership that only survives in the user's cached team list grants + nothing.""" + + if not is_keyless_mcp_subject(user_api_key_auth): + return (user_api_key_auth,) + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + + load_sources: Final = admitted_sources or MCPRequestHandler.admitted_subject_sources + acting: Final = await admitted_context(user_api_key_auth) + return tuple(await load_sources(acting if acting is not None else user_api_key_auth)) + + +async def _own_toolset_ids( + context: UserAPIKeyAuth, + load_own_permission: OwnObjectPermission, +) -> Sequence[str] | None: + """The source's own toolsets, or None when it declares no MCP grant of its own. A source that names an + ``object_permission_id`` is a known restriction even when the row is unhydrated, unreadable or gone, + so it is loaded rather than read as unrestricted, and grants nothing when it cannot be read.""" + if context.object_permission is None and not context.object_permission_id: + return None + try: + own: Final = await load_own_permission(context) + except Exception as exc: # noqa: BLE001 # a named but unreadable own grant must deny, not widen to the team + verbose_logger.warning( + "MCP toolset grants: object permission %s unreadable, granting nothing through it: %s", + context.object_permission_id, + exc, + ) + return () + if own is None: + return () + if not _restricts_mcp(own): + return None + return own.mcp_toolsets or () + + +async def _inherited_toolset_ids( + context: UserAPIKeyAuth, + load_team_permission: TeamObjectPermission, +) -> Sequence[str]: + try: + team: Final = await load_team_permission(context) + except Exception as exc: # noqa: BLE001 # an unreadable team grants nothing through this source and must not fail the caller's other sources + verbose_logger.warning( + "MCP toolset grants: team %s unreadable, inheriting nothing from it: %s", + context.team_id, + exc, + ) + return () + return () if team is None else (team.mcp_toolsets or ()) + + +async def _context_toolset_ids( + context: UserAPIKeyAuth, + inherits_team: bool, + load_team_permission: TeamObjectPermission, + load_own_permission: OwnObjectPermission, +) -> Sequence[str]: + own: Final = await _own_toolset_ids(context, load_own_permission) + if own is not None: + return own + if not inherits_team or not context.team_id: + return () + return await _inherited_toolset_ids(context, load_team_permission) + + +async def granted_toolset_ids( + user_api_key_auth: UserAPIKeyAuth, + effective_contexts: EffectiveAuthContexts = toolset_grant_contexts, + team_object_permission: TeamObjectPermission | None = None, + require_key_access: bool | None = None, + own_object_permission: OwnObjectPermission | None = None, +) -> frozenset[str]: + """Toolset ids the principal holds, resolved per grant source with the key/team rule the aggregate + /mcp listing applies: a source that declares any MCP grant of its own is scoped to its own toolsets and + never reads its team, one that declares none inherits its team's, except a virtual key under + ``require_key_mcp_access_defined``, which inherits nothing. A keyless subject's team sources always + inherit. A team that cannot be read contributes nothing while every other source still counts, and an + own grant that is named but cannot be read grants nothing. No grant anywhere yields the empty set.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy.proxy_server import general_settings + + require: Final = ( + require_key_access + if require_key_access is not None + else bool( + general_settings.get( # pyright: ignore[reportUnknownArgumentType] # general_settings is an untyped dict; truthiness must match the /mcp path's read of this flag + "require_key_mcp_access_defined", False + ) + ) + ) + inherits_team: Final = is_keyless_mcp_subject(user_api_key_auth) or not require + load_team_permission: Final = team_object_permission or MCPRequestHandler.team_object_permission + load_own_permission: Final = own_object_permission or MCPRequestHandler.key_object_permission_hydrated + contexts: Final = await effective_contexts(user_api_key_auth) + per_context: Final = await asyncio.gather( + *( + _context_toolset_ids(context, inherits_team, load_team_permission, load_own_permission) + for context in contexts + ) + ) + return frozenset(chain.from_iterable(per_context)) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 0974a2a952d..8ed8ec8752e 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -184,6 +184,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.ui_session_utils import ( admitted_user_context, build_effective_auth_contexts, + granted_toolset_ids, is_ui_session_credential, ) from litellm.proxy._types import ( @@ -3332,18 +3333,15 @@ if MCP_AVAILABLE: ): """Return toolsets the calling key is allowed to access.""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - is_admin: Final = _user_has_admin_view(user_api_key_dict) - op: Final = user_api_key_dict.object_permission - # mcp_toolsets=None or [] both mean "not restricted by toolsets". - # For admins: either value → no restriction → return all. - # For non-admins: either value → no toolsets explicitly granted → return nothing. - # (An admin whose DB row has mcp_toolsets=[] should still see all toolsets.) - raw_toolsets: Final = getattr(op, "mcp_toolsets", None) if op else None - if not raw_toolsets: - if is_admin: + if _user_has_admin_view(user_api_key_dict): + op: Final = user_api_key_dict.object_permission + if op is None or not op.mcp_toolsets: return await list_mcp_toolsets(prisma_client) + return await list_mcp_toolsets(prisma_client, toolset_ids=op.mcp_toolsets) + granted: Final = await granted_toolset_ids(user_api_key_dict) + if not granted: return [] - return await list_mcp_toolsets(prisma_client, toolset_ids=raw_toolsets) + return await list_mcp_toolsets(prisma_client, toolset_ids=sorted(granted)) @router.get( "/toolset/{toolset_id}", @@ -3355,15 +3353,13 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - # Non-admin keys may only fetch toolsets they've been explicitly granted. - if not _user_has_admin_view(user_api_key_dict): - op: Final = user_api_key_dict.object_permission - granted: Final = getattr(op, "mcp_toolsets", None) if op else None - if granted is None or toolset_id not in granted: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={"error": "API key does not have access to this toolset."}, - ) + if not _user_has_admin_view(user_api_key_dict) and toolset_id not in await granted_toolset_ids( + user_api_key_dict + ): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={"error": "API key does not have access to this toolset."}, + ) toolset: Final = await get_mcp_toolset(prisma_client, toolset_id) if toolset is None: raise HTTPException( diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 86b37f31090..68f635f7799 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -3466,6 +3466,12 @@ async def team_member_add( litellm_proxy_admin_name=litellm_proxy_admin_name, ) + await delete_cache_team_object( + team_id=data.team_id, + team_alias=complete_team_data.team_alias, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) await evict_and_broadcast( cache_keys=tuple(sorted(user.user_id for user in updated_users)), user_api_key_cache=user_api_key_cache, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 71f61079154..c9e0861935d 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -40,6 +40,7 @@ if TYPE_CHECKING: from mcp.types import CallToolResult from mcp.types import Tool as MCPTool + from litellm.proxy._experimental.mcp_server.ui_session_utils import GrantedToolsetIds from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging else: @@ -223,7 +224,9 @@ class LiteLLM_Proxy_MCP_Handler: mcp_servers=all_server_ids, mcp_tool_permissions=tool_permissions, ) - return user_api_key_auth.model_copy(update={"object_permission": updated_op}) + return user_api_key_auth.model_copy( + update={"object_permission": updated_op, "mcp_explicit_grants_only": True} + ) except Exception as _e: verbose_logger.debug("Could not apply toolset permissions: %s", _e) return user_api_key_auth @@ -237,6 +240,7 @@ class LiteLLM_Proxy_MCP_Handler: mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, request_tags: list[str] | None = None, raw_headers: dict[str, str] | None = None, + granted_toolsets: "GrantedToolsetIds | None" = None, ) -> tuple[list[MCPTool], list[str]]: """ Get available tools from the MCP server manager. @@ -279,23 +283,19 @@ class LiteLLM_Proxy_MCP_Handler: if prisma_client is not None: toolset = await global_mcp_server_manager.get_toolset_by_name_cached(prisma_client, name) if toolset is not None: - # Access control: only allow if the key explicitly grants this toolset. if user_api_key_auth is not None: + from litellm.proxy._experimental.mcp_server.ui_session_utils import ( + granted_toolset_ids, + ) from litellm.proxy.management_endpoints.common_utils import ( _user_has_admin_view, ) - is_admin = _user_has_admin_view(user_api_key_auth) - if not is_admin: - op = user_api_key_auth.object_permission - granted = getattr(op, "mcp_toolsets", None) if op else None - # None means no grants configured → deny (consistent with - # fetch_mcp_toolsets which returns [] for unconfigured keys) - if granted is None or toolset.toolset_id not in granted: - verbose_logger.debug( - "Key does not have access to toolset '%s', skipping.", name - ) - continue + if not _user_has_admin_view(user_api_key_auth) and toolset.toolset_id not in ( + await (granted_toolsets or granted_toolset_ids)(user_api_key_auth) + ): + verbose_logger.debug("Key does not have access to toolset '%s', skipping.", name) + continue resolved_toolset_ids.append(toolset.toolset_id) # Don't add to resolved_mcp_servers — toolset scope # restricts via object_permission, not server name filter. diff --git a/tests/integration/_support/mcp_grants.py b/tests/integration/_support/mcp_grants.py index 5fa9eeaa0b6..82e79b2c3bd 100644 --- a/tests/integration/_support/mcp_grants.py +++ b/tests/integration/_support/mcp_grants.py @@ -51,12 +51,12 @@ def delete_toolset(gateway: Gateway, identity: str) -> None: assert response.status_code in (200, 202, 204), response.text -def create_toolset(scenario: Scenario, tools: tuple[tuple[str, str], ...]) -> str: +def create_toolset(scenario: Scenario, tools: tuple[tuple[str, str], ...], toolset_name: str | None = None) -> str: response: Final = scenario.gateway.request( "POST", "/v1/mcp/toolset", { - "toolset_name": f"integration-{uuid.uuid4().hex[:10]}", + "toolset_name": toolset_name or f"integration-{uuid.uuid4().hex[:10]}", "tools": [{"server_id": server_id, "tool_name": tool} for server_id, tool in tools], }, ) diff --git a/tests/integration/mcp/test_mcp_llm_endpoints.py b/tests/integration/mcp/test_mcp_llm_endpoints.py index 26f5ced4de7..01bf006a03e 100644 --- a/tests/integration/mcp/test_mcp_llm_endpoints.py +++ b/tests/integration/mcp/test_mcp_llm_endpoints.py @@ -9,6 +9,7 @@ import httpx import pytest from integration._support.client import Gateway, Scenario from integration._support.mcp import McpPeer, mcp_peer, register_mcp, tool_calls +from integration._support.mcp_grants import create_toolset from integration._support.wire import Reply, Request, Wire, wire_server Surface = Literal["chat", "responses", "messages", "messages_bridge"] @@ -194,9 +195,7 @@ class Rig: ) def upstream_tools(self) -> tuple[tuple[str, ...], ...]: - return tuple( - _tool_names(json.loads(request.body)) for request in self.wire.drain() if request.method == "POST" - ) + return tuple(_tool_names(json.loads(request.body)) for request in self.wire.drain() if request.method == "POST") def final_text(self, body: Mapping[str, object]) -> str: if self.surface == "chat": @@ -314,6 +313,25 @@ def test_allowed_tools_narrows_the_tool_list_handed_to_the_model(gateway: Gatewa assert [call["body"]["params"]["name"] for call in _peer_add_calls(rig.peer)] == ["add"] +@pytest.mark.parametrize("surface", ("chat", "responses", "messages")) +def test_toolset_gateway_url_serves_a_team_granted_toolset_to_a_key_without_its_own_grant( + gateway: Gateway, surface: Surface +) -> None: + with _rig(gateway, surface) as rig: + register_mcp(rig.scenario, rig.peer, "open" + uuid.uuid4().hex[:8], allow_all_keys=True) + toolset_name: Final = "ts" + uuid.uuid4().hex[:8] + toolset_id: Final = create_toolset(rig.scenario, ((rig.server_id, "add"),), toolset_name=toolset_name) + sibling_id: Final = create_toolset(rig.scenario, ((rig.server_id, "multiply"),)) + team_id: Final = rig.scenario.team(object_permission={"mcp_toolsets": [toolset_id, sibling_id]}) + key: Final = rig.scenario.key(team_id=team_id) + response: Final = rig.send(key, [{**AUTO, "server_url": f"litellm_proxy/mcp/{toolset_name}"}]) + assert response.status_code == 200, response.text + requests: Final = rig.upstream_tools() + assert requests, "model was never called" + assert all(names == (rig.tool,) for names in requests), requests + assert [call["body"]["params"]["name"] for call in _peer_add_calls(rig.peer)] == ["add"] + + @pytest.mark.parametrize("surface", ("chat", "responses", "messages")) def test_server_scoped_gateway_url_exposes_only_that_servers_tools(gateway: Gateway, surface: Surface) -> None: with _rig(gateway, surface) as rig, mcp_peer() as other_peer: @@ -347,3 +365,41 @@ def test_streaming_chat_executes_the_tool_once_and_streams_the_follow_up(gateway assert text == ANSWER, response.text assert [call["body"]["params"]["name"] for call in _peer_add_calls(rig.peer)] == ["add"] assert len(rig.upstream_tools()) == 2 + + +def test_streaming_chat_through_a_toolset_gateway_url_serves_a_team_key_without_its_own_grant( + gateway: Gateway, +) -> None: + with _rig(gateway, "chat") as rig: + register_mcp(rig.scenario, rig.peer, "open" + uuid.uuid4().hex[:8], allow_all_keys=True) + toolset_name: Final = "ts" + uuid.uuid4().hex[:8] + toolset_id: Final = create_toolset(rig.scenario, ((rig.server_id, "add"),), toolset_name=toolset_name) + key: Final = rig.scenario.key(team_id=rig.scenario.team(object_permission={"mcp_toolsets": [toolset_id]})) + response: Final = rig.send(key, [{**AUTO, "server_url": f"litellm_proxy/mcp/{toolset_name}"}], stream=True) + assert response.status_code == 200, response.text + chunks: Final = tuple( + json.loads(line.removeprefix("data: ")) + for line in response.text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + text: Final = "".join( + str(chunk["choices"][0]["delta"].get("content") or "") for chunk in chunks if chunk.get("choices") + ) + assert text == ANSWER, response.text + assert [call["body"]["params"]["name"] for call in _peer_add_calls(rig.peer)] == ["add"] + requests: Final = rig.upstream_tools() + assert len(requests) == 2 and all(names == (rig.tool,) for names in requests), requests + + +@pytest.mark.parametrize("surface", ("chat", "responses", "messages")) +def test_toolset_gateway_url_gives_a_key_of_an_ungranted_team_no_tools_and_never_reaches_the_peer( + gateway: Gateway, surface: Surface +) -> None: + with _rig(gateway, surface) as rig: + toolset_name: Final = "ts" + uuid.uuid4().hex[:8] + create_toolset(rig.scenario, ((rig.server_id, "add"),), toolset_name=toolset_name) + key: Final = rig.scenario.key(team_id=rig.scenario.team()) + response: Final = rig.send(key, [{**AUTO, "server_url": f"litellm_proxy/mcp/{toolset_name}"}]) + assert _peer_add_calls(rig.peer) == (), "denied caller reached the peer" + assert all(rig.tool not in names for names in rig.upstream_tools()), rig.upstream_tools() + assert response.status_code in (200, 400, 401, 403), response.text diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 7a60c8ede30..60563e7aacd 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -1,25 +1,31 @@ import base64 import hashlib import secrets +import time import uuid from dataclasses import dataclass from typing import Final from urllib.parse import parse_qs, urlsplit import httpx +import jwt import pytest from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, + INITIALIZE, EntryPoint, McpCaller, McpPeer, + Outcome, + _outcome_from_rpc, call_tool, mcp_peer, register_mcp, tool_calls, ) +from integration._support.mcp_grants import create_toolset from integration._support.oauth_server import AuthorizationServer, oauth_server ADD: Final = {"a": 2, "b": 3} @@ -383,3 +389,117 @@ def test_dcr_bridge_relays_client_registration_and_advertises_gateway_endpoints( assert issuer.json()["authorization_endpoint"] == f"{_base(gateway)}/{alias}/authorize" assert issuer.json()["token_endpoint"] == f"{_base(gateway)}/{alias}/token" assert "S256" in issuer.json()["code_challenge_methods_supported"] + + +def _ui_session_cookie(gateway: Gateway, user_id: str) -> dict[str, str]: + claims: Final = {"user_id": user_id, "login_method": "username_password", "exp": int(time.time()) + 600} + return {"token": jwt.encode(claims, gateway.key, algorithm="HS256")} + + +def _gateway_session_bearer(gateway: Gateway, user_id: str, resource: str | None = None) -> str: + registered: Final = gateway.client.post( + "/register", json={"redirect_uris": [CLIENT_REDIRECT], "client_name": "integration"} + ) + assert registered.status_code in (200, 201), registered.text + client_id: Final = registered.json()["client_id"] + pkce: Final = _Pkce(secrets.token_urlsafe(48)) + cookies: Final = _ui_session_cookie(gateway, user_id) + started: Final = gateway.client.get( + "/authorize", + params={ + "client_id": client_id, + "redirect_uri": CLIENT_REDIRECT, + "response_type": "code", + "state": "lit6029", + "code_challenge": pkce.challenge, + "code_challenge_method": "S256", + **({} if resource is None else {"resource": resource}), + }, + cookies=cookies, + ) + assert started.status_code == 303, started.text + handle: Final = parse_qs(urlsplit(started.headers["location"]).query)["connect_flow"][0] + completed: Final = gateway.client.post( + "/authorize/complete", data={"flow": handle}, cookies={**cookies, **dict(started.cookies)} + ) + assert completed.status_code == 303, completed.text + callback: Final = parse_qs(urlsplit(completed.headers["location"]).query) + assert "code" in callback, completed.headers["location"] + issued: Final = gateway.client.post( + "/token", + data={ + "grant_type": "authorization_code", + "code": callback["code"][0], + "redirect_uri": CLIENT_REDIRECT, + "client_id": client_id, + "code_verifier": pkce.verifier, + }, + ) + assert issued.status_code == 200, issued.text + return _issued_token(issued.json()) + + +def _toolset_rpc(gateway: Gateway, bearer: str, name: str, method: str, params: dict[str, object]) -> Outcome: + def post(rpc_method: str, rpc_params: dict[str, object]) -> httpx.Response: + return gateway.client.post( + f"/toolset/{name}/mcp", + json={"jsonrpc": "2.0", "id": 1, "method": rpc_method, "params": rpc_params}, + headers={"Authorization": f"Bearer {bearer}", "Accept": "application/json, text/event-stream"}, + ) + + initialized: Final = _outcome_from_rpc(post("initialize", dict(INITIALIZE))) + if not initialized.ok: + return initialized + return _outcome_from_rpc(post(method, params)) + + +def test_gateway_session_bearer_of_a_team_member_is_served_the_team_toolset_on_its_route(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029sess" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_name: Final = "lit6029g" + uuid.uuid4().hex[:8] + withheld_name: Final = "lit6029w" + uuid.uuid4().hex[:8] + granted_id: Final = create_toolset(scenario, ((server_id, "add"),), toolset_name=granted_name) + create_toolset(scenario, ((server_id, "multiply"),), toolset_name=withheld_name) + member: Final = scenario.member(scenario.team(object_permission={"mcp_toolsets": [granted_id]})) + bearer: Final = _gateway_session_bearer(gateway, member) + assert bearer.startswith("llm_session_"), bearer[:16] + listed: Final = _toolset_rpc(gateway, bearer, granted_name, "tools/list", {}) + assert listed.tools == (f"{alias}-add",), listed.raw + peer.drain() + called: Final = _toolset_rpc( + gateway, bearer, granted_name, "tools/call", {"name": f"{alias}-add", "arguments": {"a": 4, "b": 5}} + ) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + denied: Final = _toolset_rpc(gateway, bearer, withheld_name, "tools/list", {}) + assert denied.status == 403, denied.raw + assert tool_calls(peer.drain()) == () + + +def test_resource_scoped_session_bearer_opens_a_team_toolset_inside_its_server_and_none_outside( + gateway: Gateway, +) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + inside: Final = "lit6029in" + uuid.uuid4().hex[:6] + outside: Final = "lit6029out" + uuid.uuid4().hex[:6] + inside_server: Final = register_mcp(scenario, peer, inside) + outside_server: Final = register_mcp(scenario, peer, outside) + inside_name: Final = "lit6029i" + uuid.uuid4().hex[:8] + outside_name: Final = "lit6029o" + uuid.uuid4().hex[:8] + inside_id: Final = create_toolset(scenario, ((inside_server, "add"),), toolset_name=inside_name) + outside_id: Final = create_toolset(scenario, ((outside_server, "add"),), toolset_name=outside_name) + member: Final = scenario.member(scenario.team(object_permission={"mcp_toolsets": [inside_id, outside_id]})) + bearer: Final = _gateway_session_bearer(gateway, member, resource=f"{_base(gateway)}/{inside}/mcp") + assert bearer.startswith("llm_session_"), bearer[:16] + listed: Final = _toolset_rpc(gateway, bearer, inside_name, "tools/list", {}) + assert listed.tools == (f"{inside}-add",), listed.raw + peer.drain() + called: Final = _toolset_rpc( + gateway, bearer, inside_name, "tools/call", {"name": f"{inside}-add", "arguments": {"a": 4, "b": 5}} + ) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + refused: Final = _toolset_rpc(gateway, bearer, outside_name, "tools/list", {}) + assert refused.status == 403, refused.raw + assert tool_calls(peer.drain()) == () diff --git a/tests/integration/mcp/test_mcp_toolsets.py b/tests/integration/mcp/test_mcp_toolsets.py new file mode 100644 index 00000000000..3dd309db665 --- /dev/null +++ b/tests/integration/mcp/test_mcp_toolsets.py @@ -0,0 +1,509 @@ +import secrets +import uuid +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, Scenario, object_value +from integration._support.mcp import ( + INITIALIZE, + Outcome, + _outcome_from_rest, + _outcome_from_rpc, + mcp_peer, + register_mcp, + tool_calls, +) +from integration._support.mcp_grants import create_toolset +from integration._support.process import owned_proxy + +from litellm.models.user import LiteLLM_UserTable +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import LITELLM_SESSION_TOKEN_PREFIX, ExperimentalUIJWTToken +from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_bearer_token + +ADD: Final = {"a": 4, "b": 5} + + +def _dashboard_ui_session_token(user_id: str) -> str: + user: Final = LiteLLM_UserTable(user_id=user_id, user_role="internal_user", models=[]) + return ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(user) + + +def _toolset(scenario: Scenario, server_id: str, tool: str) -> tuple[str, str]: + name: Final = "lit6029_" + uuid.uuid4().hex[:10] + return create_toolset(scenario, ((server_id, tool),), toolset_name=name), name + + +def _toolset_rpc( + gateway: Gateway, headers: dict[str, str], name: str, method: str, params: dict[str, object] +) -> Outcome: + def post(rpc_method: str, rpc_params: dict[str, object]) -> httpx.Response: + return gateway.client.post( + f"/toolset/{name}/mcp", + json={"jsonrpc": "2.0", "id": 1, "method": rpc_method, "params": rpc_params}, + headers={**headers, "Accept": "application/json, text/event-stream"}, + ) + + initialized: Final = _outcome_from_rpc(post("initialize", dict(INITIALIZE))) + if not initialized.ok: + return initialized + return _outcome_from_rpc(post(method, params)) + + +def _listed_toolset_ids(gateway: Gateway, headers: dict[str, str]) -> tuple[str, ...]: + response: Final = gateway.client.get("/v1/mcp/toolset", headers=headers) + assert response.status_code == 200, response.text + return tuple(toolset["toolset_id"] for toolset in response.json()) + + +def _assert_team_grants_only(gateway: Gateway, team_id: str, key: str, toolset_id: str) -> None: + team: Final = object_value(gateway.get("/team/info", {"team_id": team_id})["team_info"]) + assert object_value(team["object_permission"])["mcp_toolsets"] == [toolset_id], team + key_info: Final = object_value(gateway.get("/key/info", {"key": key})["info"]) + assert key_info.get("object_permission") is None, f"key must carry no grant of its own: {key_info}" + + +def test_team_granted_toolset_is_listed_and_served_to_a_team_key(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + withheld_id, withheld_name = _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + key: Final = scenario.key(team_id=team_id) + _assert_team_grants_only(gateway, team_id, key, granted_id) + headers: Final = {"Authorization": f"Bearer {key}"} + + assert _listed_toolset_ids(gateway, headers) == (granted_id,) + detail: Final = gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers) + assert detail.status_code == 200, detail.text + assert detail.json()["toolset_name"] == granted_name, detail.text + withheld_detail: Final = gateway.client.get(f"/v1/mcp/toolset/{withheld_id}", headers=headers) + assert withheld_detail.status_code == 403, withheld_detail.text + + listed: Final = _toolset_rpc(gateway, headers, granted_name, "tools/list", {}) + assert listed.ok, listed.raw + assert listed.tools == (f"{alias}-add",), listed.raw + peer.drain() + called: Final = _toolset_rpc( + gateway, headers, granted_name, "tools/call", {"name": f"{alias}-add", "arguments": ADD} + ) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + denied: Final = _toolset_rpc(gateway, headers, withheld_name, "tools/list", {}) + assert denied.status == 403, denied.raw + + +def test_dashboard_session_of_a_team_member_lists_the_team_granted_toolset( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + user_id: Final = scenario.user(user_role="internal_user", teams=[team_id]) + user: Final = object_value(gateway.get("/user/info", {"user_id": user_id})["user_info"]) + assert user["teams"] == [team_id], user + headers: Final = {"Authorization": f"Bearer {_dashboard_ui_session_token(user_id)}"} + + assert _listed_toolset_ids(gateway, headers) == (granted_id,) + detail: Final = gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers) + assert detail.status_code == 200, detail.text + assert detail.json()["toolset_name"] == granted_name, detail.text + + +def test_direct_grants_no_grants_and_admin_listing_are_unchanged_by_team_resolution(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + withheld_id, withheld_name = _toolset(scenario, server_id, "multiply") + direct: Final = {"Authorization": f"Bearer {scenario.key(object_permission={'mcp_toolsets': [granted_id]})}"} + ungranted_team: Final = scenario.team() + no_grant: Final = {"Authorization": f"Bearer {scenario.key(team_id=ungranted_team)}"} + admin: Final = {"Authorization": f"Bearer {gateway.key}"} + + assert _listed_toolset_ids(gateway, direct) == (granted_id,) + assert _toolset_rpc(gateway, direct, granted_name, "tools/list", {}).tools == (f"{alias}-add",) + assert _toolset_rpc(gateway, direct, withheld_name, "tools/list", {}).status == 403 + assert gateway.client.get(f"/v1/mcp/toolset/{withheld_id}", headers=direct).status_code == 403 + + assert _listed_toolset_ids(gateway, no_grant) == () + assert gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=no_grant).status_code == 403 + assert _toolset_rpc(gateway, no_grant, granted_name, "tools/list", {}).status == 403 + + assert {granted_id, withheld_id} <= set(_listed_toolset_ids(gateway, admin)) + assert gateway.client.get(f"/v1/mcp/toolset/{withheld_id}", headers=admin).status_code == 200 + assert _toolset_rpc(gateway, admin, withheld_name, "tools/list", {}).tools == (f"{alias}-multiply",) + + +def test_a_key_with_its_own_toolset_grant_does_not_inherit_the_team_toolset(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + own_id, own_name = _toolset(scenario, server_id, "add") + team_only_id, team_only_name = _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [own_id, team_only_id]}) + key: Final = scenario.key(team_id=team_id, object_permission={"mcp_toolsets": [own_id]}) + headers: Final = {"Authorization": f"Bearer {key}"} + + assert _listed_toolset_ids(gateway, headers) == (own_id,) + assert gateway.client.get(f"/v1/mcp/toolset/{team_only_id}", headers=headers).status_code == 403 + assert _toolset_rpc(gateway, headers, team_only_name, "tools/list", {}).status == 403 + assert _toolset_rpc(gateway, headers, own_name, "tools/list", {}).tools == (f"{alias}-add",) + + +def _team_member_with_own_grant(scenario: Scenario, team_id: str, own_server_id: str) -> str: + user_id: Final = scenario.user(user_role="internal_user", object_permission={"mcp_servers": [own_server_id]}) + scenario.gateway.post("/team/member_add", {"team_id": team_id, "member": {"role": "user", "user_id": user_id}}) + return user_id + + +def test_dashboard_session_serves_the_team_toolset_despite_a_disjoint_grant_on_the_user_row( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + """The member's own row grants a different server outright. That grant must not cap the team's toolset + to nothing, and the team's sibling toolset must not leak onto the granted toolset's route.""" + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + own_server_id: Final = register_mcp(scenario, peer, "lit6029_own_" + uuid.uuid4().hex[:8]) + granted_id, granted_name = _toolset(scenario, server_id, "add") + sibling_id, sibling_name = _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id, sibling_id]}) + user_id: Final = _team_member_with_own_grant(scenario, team_id, own_server_id) + headers: Final = {"Authorization": f"Bearer {_dashboard_ui_session_token(user_id)}"} + + assert set(_listed_toolset_ids(gateway, headers)) == {granted_id, sibling_id} + listed: Final = _toolset_rpc(gateway, headers, granted_name, "tools/list", {}) + assert listed.ok, listed.raw + assert listed.tools == (f"{alias}-add",), listed.raw + assert _toolset_rpc(gateway, headers, sibling_name, "tools/list", {}).tools == (f"{alias}-multiply",) + peer.drain() + called: Final = _toolset_rpc( + gateway, headers, granted_name, "tools/call", {"name": f"{alias}-add", "arguments": ADD} + ) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + stranger: Final = { + "Authorization": f"Bearer {_dashboard_ui_session_token(scenario.user(user_role='internal_user'))}" + } + assert _toolset_rpc(gateway, stranger, granted_name, "tools/list", {}).status == 403 + + +def test_a_member_removed_from_the_team_loses_its_toolset_on_the_dashboard_session( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + user_id: Final = scenario.member(team_id) + headers: Final = {"Authorization": f"Bearer {_dashboard_ui_session_token(user_id)}"} + assert _listed_toolset_ids(gateway, headers) == (granted_id,) + assert _toolset_rpc(gateway, headers, granted_name, "tools/list", {}).tools == (f"{alias}-add",) + + gateway.post("/team/member_delete", {"team_id": team_id, "user_id": user_id}) + + assert _listed_toolset_ids(gateway, headers) == () + assert gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers).status_code == 403 + assert _toolset_rpc(gateway, headers, granted_name, "tools/list", {}).status == 403 + + +def _bearer(token: str) -> dict[str, str]: + return {"Authorization": f"Bearer {token}"} + + +def _expired_dashboard_token(user_id: str) -> str: + expired: Final = (datetime.now(timezone.utc) - timedelta(minutes=5)).strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + stale: Final = UserAPIKeyAuth( + token="ui-token", + key_name="ui-token", + key_alias="ui-token", + expires=expired + "+00:00", + user_id=user_id, + team_id="litellm-dashboard", + models=[], + user_role=LitellmUserRoles.INTERNAL_USER, + ) + return encrypt_bearer_token(stale.model_dump_json(exclude_none=True), prefix=LITELLM_SESSION_TOKEN_PREFIX) + + +def _rest_list(gateway: Gateway, headers: dict[str, str], params: object) -> httpx.Response: + return gateway.client.get("/mcp-rest/tools/list", headers=headers, params=params) + + +def _rest_call(gateway: Gateway, headers: dict[str, str], name: str, server_id: str) -> Outcome: + return _outcome_from_rest( + gateway.client.post( + "/mcp-rest/tools/call", + headers=headers, + json={"name": name, "arguments": dict(ADD), "server_id": server_id}, + ) + ) + + +def _route_tools(gateway: Gateway, headers: dict[str, str], name: str) -> Outcome: + return _toolset_rpc(gateway, headers, name, "tools/list", {}) + + +def _route_call(gateway: Gateway, headers: dict[str, str], name: str, tool: str) -> Outcome: + return _toolset_rpc(gateway, headers, name, "tools/call", {"name": tool, "arguments": dict(ADD)}) + + +def _strict_config(directory: Path) -> Path: + base: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text()) + strict: Final = { + **base, + "general_settings": {**base.get("general_settings", {}), "require_key_mcp_access_defined": True}, + } + path: Final = directory / "require_key_mcp_access.yaml" + path.write_text(yaml.safe_dump(strict)) + return path + + +def test_mcp_rest_toolset_name_narrows_the_list_and_serves_the_call_for_a_team_key_and_a_dashboard_member( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029rest" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + _, withheld_name = _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + key: Final = scenario.key(team_id=team_id) + member: Final = scenario.member(team_id) + for headers in (_bearer(key), _bearer(_dashboard_ui_session_token(member))): + listed: Final = _rest_list(gateway, headers, {"toolset_name": granted_name}) + assert listed.status_code == 200, listed.text + listed_names: Final = tuple(tool["name"] for tool in listed.json()["tools"]) + assert len(listed_names) == 1 and listed_names[0].endswith("add"), listed.text + denied: Final = _rest_list(gateway, headers, {"toolset_name": withheld_name}) + assert denied.status_code == 200 and denied.json()["tools"] == [], denied.text + assert "does not have access to toolset" in denied.json()["message"], denied.text + peer.drain() + called: Final = _rest_call(gateway, headers, listed_names[0], server_id) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + + +def test_a_key_restricted_to_its_own_servers_does_not_inherit_the_team_toolset(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029own" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + other_id: Final = register_mcp(scenario, peer, "lit6029other" + uuid.uuid4().hex[:6]) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id], "mcp_servers": [other_id]}) + key: Final = scenario.key(team_id=team_id, object_permission={"mcp_servers": [other_id]}) + headers: Final = _bearer(key) + assert _listed_toolset_ids(gateway, headers) == () + assert gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers).status_code == 403 + assert _route_tools(gateway, headers, granted_name).status == 403 + assert tool_calls(peer.drain()) == () + + +def test_a_team_with_an_empty_or_absent_toolset_grant_gives_its_keys_nothing(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + server_id: Final = register_mcp(scenario, peer, "lit6029empty" + uuid.uuid4().hex[:6]) + granted_id, granted_name = _toolset(scenario, server_id, "add") + teams: Final = ( + scenario.team(object_permission={"mcp_toolsets": []}), + scenario.team(object_permission={"mcp_toolsets": None}), + scenario.team(), + ) + for team_id in teams: + headers: Final = _bearer(scenario.key(team_id=team_id)) + assert _listed_toolset_ids(gateway, headers) == (), team_id + assert gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers).status_code == 403 + assert _route_tools(gateway, headers, granted_name).status == 403, team_id + assert tool_calls(peer.drain()) == () + + +def test_an_unknown_or_malformed_toolset_name_is_refused_without_peer_traffic_and_the_route_keeps_serving( + gateway: Gateway, +) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029bad" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + _, withheld_name = _toolset(scenario, server_id, "multiply") + headers: Final = _bearer(scenario.key(team_id=scenario.team(object_permission={"mcp_toolsets": [granted_id]}))) + unknown: Final = "missing" + uuid.uuid4().hex[:8] + assert _route_tools(gateway, headers, unknown).status == 404 + assert _rest_list(gateway, headers, {"toolset_name": unknown}).status_code == 404 + malformed: Final = ( + [("toolset_name", granted_name), ("toolset_name", withheld_name)], + {"toolset_name": "x" * 5000}, + {"toolset_name": ""}, + {"toolset_name": granted_name + "\x00"}, + ) + statuses: Final = tuple(_rest_list(gateway, headers, params).status_code for params in malformed) + assert all(status < 500 for status in statuses), statuses + assert tool_calls(peer.drain()) == () + assert gateway.client.get("/health/liveliness").status_code == 200 + served: Final = _route_tools(gateway, headers, granted_name) + assert served.tools == (f"{alias}-add",), served.raw + + +def test_garbage_expired_and_tampered_credentials_are_refused_on_every_toolset_surface( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + server_id: Final = register_mcp(scenario, peer, "lit6029cred" + uuid.uuid4().hex[:6]) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + member: Final = scenario.member(team_id) + forged: Final = ( + "sk-" + secrets.token_urlsafe(24), + _expired_dashboard_token(member), + "llm_session_" + secrets.token_urlsafe(32), + ) + for bearer in forged: + headers: Final = _bearer(bearer) + listed: Final = gateway.client.get("/v1/mcp/toolset", headers=headers) + assert listed.status_code == 401, (bearer[:12], listed.text) + detail: Final = gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers) + assert detail.status_code == 401, (bearer[:12], detail.text) + routed: Final = _route_tools(gateway, headers, granted_name) + assert routed.status == 401, (bearer[:12], routed.raw) + rest: Final = _rest_list(gateway, headers, {"toolset_name": granted_name}) + assert rest.status_code == 401, (bearer[:12], rest.text) + assert peer.drain() == () + + +def test_a_dashboard_member_of_a_deleted_team_loses_the_toolset_while_a_direct_grant_survives( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029gone" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + team_id_, team_name = _toolset(scenario, server_id, "add") + own_id, own_name = _toolset(scenario, server_id, "multiply") + doomed: Final = scenario.gateway.post( + "/team/new", + {"team_alias": f"integration-{uuid.uuid4().hex}", "object_permission": {"mcp_toolsets": [team_id_]}}, + ) + doomed_team: Final = str(doomed["team_id"]) + member: Final = scenario.user(user_role="internal_user", teams=[doomed_team]) + granted: Final = scenario.user( + user_role="internal_user", teams=[doomed_team], object_permission={"mcp_toolsets": [own_id]} + ) + member_headers: Final = _bearer(_dashboard_ui_session_token(member)) + granted_headers: Final = _bearer(_dashboard_ui_session_token(granted)) + assert _listed_toolset_ids(gateway, member_headers) == (team_id_,) + assert set(_listed_toolset_ids(gateway, granted_headers)) == {team_id_, own_id} + scenario.delete_team(doomed_team) + assert _listed_toolset_ids(gateway, member_headers) == () + assert gateway.client.get(f"/v1/mcp/toolset/{team_id_}", headers=member_headers).status_code == 403 + assert _route_tools(gateway, member_headers, team_name).status == 403 + assert _listed_toolset_ids(gateway, granted_headers) == (own_id,) + assert _route_tools(gateway, granted_headers, team_name).status == 403 + assert _route_tools(gateway, granted_headers, own_name).tools == (f"{alias}-multiply",) + peer.drain() + kept: Final = _route_call(gateway, granted_headers, own_name, f"{alias}-multiply") + assert kept.ok and kept.text == "20", kept.raw + assert [call["body"]["params"]["name"] for call in tool_calls(peer.drain())] == ["multiply"] + + +def test_require_key_mcp_access_defined_stops_key_inheritance_but_not_the_dashboard_member( + gateway: Gateway, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with ( + owned_proxy(gateway, tmp_path, {}, config=_strict_config(tmp_path), workers=2) as strict, + mcp_peer() as peer, + strict.scenario() as scenario, + ): + alias: Final = "lit6029strict" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + inheriting: Final = _bearer(scenario.key(team_id=team_id)) + own: Final = _bearer(scenario.key(team_id=team_id, object_permission={"mcp_toolsets": [granted_id]})) + member: Final = _bearer(_dashboard_ui_session_token(scenario.member(team_id))) + assert _listed_toolset_ids(strict, inheriting) == () + assert _route_tools(strict, inheriting, granted_name).status == 403 + assert _listed_toolset_ids(strict, own) == (granted_id,) + assert _route_tools(strict, own, granted_name).tools == (f"{alias}-add",) + assert _listed_toolset_ids(strict, member) == (granted_id,) + assert _route_tools(strict, member, granted_name).tools == (f"{alias}-add",) + peer.drain() + called: Final = _route_call(strict, member, granted_name, f"{alias}-add") + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + + +def test_a_key_with_only_a_vector_store_grant_still_inherits_the_team_toolset(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029vs" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + headers: Final = _bearer( + scenario.key(team_id=team_id, object_permission={"vector_stores": ["vs-" + uuid.uuid4().hex[:8]]}) + ) + assert _listed_toolset_ids(gateway, headers) == (granted_id,) + assert _route_tools(gateway, headers, granted_name).tools == (f"{alias}-add",) + peer.drain() + called: Final = _route_call(gateway, headers, granted_name, f"{alias}-add") + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + + +def test_a_member_added_after_the_team_was_cached_sees_the_toolset_on_every_following_request( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + server_id: Final = register_mcp(scenario, peer, "lit6029cache" + uuid.uuid4().hex[:6]) + granted_id, _ = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + user_id: Final = scenario.user(user_role="internal_user") + headers: Final = _bearer(_dashboard_ui_session_token(user_id)) + warm: Final = _bearer(scenario.key(team_id=team_id)) + assert tuple(_listed_toolset_ids(gateway, warm) for _ in range(4)) == ((granted_id,),) * 4 + assert tuple(_listed_toolset_ids(gateway, headers) for _ in range(4)) == ((),) * 4 + added: Final = gateway.request( + "POST", "/team/member_add", {"team_id": team_id, "member": {"user_id": user_id, "role": "user"}} + ) + assert added.status_code == 200, added.text + listings: Final = tuple(_listed_toolset_ids(gateway, headers) for _ in range(8)) + assert listings == ((granted_id,),) * 8, listings + + +def test_a_member_of_two_teams_sees_the_union_and_each_route_stays_narrowed_to_its_own_toolset( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029two" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + first_id, first_name = _toolset(scenario, server_id, "add") + second_id, second_name = _toolset(scenario, server_id, "multiply") + teams: Final = ( + scenario.team(object_permission={"mcp_toolsets": [first_id]}), + scenario.team(object_permission={"mcp_toolsets": [second_id]}), + ) + headers: Final = _bearer( + _dashboard_ui_session_token(scenario.user(user_role="internal_user", teams=list(teams))) + ) + assert set(_listed_toolset_ids(gateway, headers)) == {first_id, second_id} + assert _route_tools(gateway, headers, first_name).tools == (f"{alias}-add",) + assert _route_tools(gateway, headers, second_name).tools == (f"{alias}-multiply",) + peer.drain() + crossed: Final = _route_call(gateway, headers, first_name, f"{alias}-multiply") + assert not crossed.ok, crossed.raw + assert tool_calls(peer.drain()) == () diff --git a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 0d0c65e3650..02ac1540071 100644 --- a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -8201,7 +8201,7 @@ class TestGatewaySessionAdmission: assert not any(k.lower() == "authorization" for k in (raw_headers or {})) -def _make_team(team_id, mcp_servers, *, org_id=None, tool_perms=None, members=("sso-user",)): +def _make_team(team_id, mcp_servers, *, org_id=None, tool_perms=None, members=("sso-user",), toolsets=None): from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member return LiteLLM_TeamTable( @@ -8210,7 +8210,10 @@ def _make_team(team_id, mcp_servers, *, org_id=None, tool_perms=None, members=(" 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 + object_permission_id=f"op-{team_id}", + mcp_servers=mcp_servers, + mcp_tool_permissions=tool_perms, + mcp_toolsets=toolsets, ), ) @@ -8271,6 +8274,64 @@ class TestUserSubjectTeamUnion: result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv1", "srv2", "srv3"} + async def test_toolsets_of_a_team_that_dropped_the_user_from_its_roster_are_not_granted(self): + """The user's cached team list still names team-revoked, but its live roster no longer lists + the user, so its toolset is withheld exactly as its servers are on the aggregate /mcp.""" + from litellm.proxy._experimental.mcp_server.ui_session_utils import granted_toolset_ids + + teams = { + "team-kept": _make_team("team-kept", [], toolsets=["ts-kept"]), + "team-revoked": _make_team("team-revoked", [], toolsets=["ts-revoked"], members=("someone-else",)), + } + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-kept", "team-revoked"]): + granted = await granted_toolset_ids(auth) + assert granted == {"ts-kept"} + + async def test_a_pinned_toolset_narrows_every_source_to_the_toolset_servers_and_tools(self): + """On /toolset/{name}/mcp the admitted subject carries mcp_toolset_id; team-a's grant on srv1 and + srv2 with every tool collapses to the toolset's srv1 and its one tool, and team-b's srv3 drops.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"]), "team-b": _make_team("team-b", ["srv3"])} + auth = _make_admitted_subject("sso-user") + pinned = auth.model_copy(update={"mcp_toolset_id": "ts-1"}) + resolve = AsyncMock(return_value={"srv1": ["add"]}) + with ( + self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]), + patch.object(global_mcp_server_manager, "resolve_toolset_tool_permissions", resolve), + ): + servers = await MCPRequestHandler.resolve_admitted_subject_servers(pinned) + tools = await MCPRequestHandler.resolve_admitted_subject_tools("srv1", pinned) + unpinned_servers = await MCPRequestHandler.resolve_admitted_subject_servers(auth) + unpinned_tools = await MCPRequestHandler.resolve_admitted_subject_tools("srv1", auth) + assert servers == ["srv1"] + assert tools == ["add"] + assert set(unpinned_servers) == {"srv1", "srv2", "srv3"} + assert unpinned_tools is None + assert {call.kwargs["toolset_ids"][0] for call in resolve.await_args_list} == {"ts-1"} + + async def test_a_fresh_policy_pinned_toolset_bypasses_the_toolset_permission_cache(self): + """A session admitted under requires_fresh_policy reads the pinned toolset from the writer, so a + tool revoked from the toolset is gone on the very next request (Devin Review 4150024092).""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"])} + auth = _make_admitted_subject("sso-user") + auth.requires_fresh_policy = True + pinned = auth.model_copy(update={"mcp_toolset_id": "ts-1"}) + resolve = AsyncMock(return_value={"srv1": ["add"]}) + with ( + self._patch(teams_by_id=teams, user_teams=["team-a"]), + patch.object(global_mcp_server_manager, "resolve_toolset_tool_permissions", resolve), + ): + servers = await MCPRequestHandler.resolve_admitted_subject_servers(pinned) + tools = await MCPRequestHandler.resolve_admitted_subject_tools("srv1", pinned) + assert servers == ["srv1"] + assert tools == ["add"] + assert resolve.await_args_list + assert all(call.kwargs == {"toolset_ids": ["ts-1"], "requires_fresh_policy": True} for call in resolve.await_args_list) + 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.""" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py index 1398884783e..95be2b8b12b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py @@ -1,10 +1,12 @@ """Tests for MCP toolset scope enforcement.""" import asyncio +from collections.abc import Awaitable, Callable from typing import Dict, List, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import HTTPException from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -30,6 +32,19 @@ def _make_auth( ) +def _granted_through_team(*team_toolset_ids: str) -> Callable[[UserAPIKeyAuth], Awaitable[frozenset[str]]]: + """The real grant resolver over a team that holds ``team_toolset_ids``, with no key access rule.""" + from litellm.proxy._experimental.mcp_server.ui_session_utils import granted_toolset_ids + + async def team_permission(context: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable(object_permission_id="team-op", mcp_toolsets=list(team_toolset_ids)) + + async def granted(context: UserAPIKeyAuth) -> frozenset[str]: + return await granted_toolset_ids(context, team_object_permission=team_permission, require_key_access=False) + + return granted + + class TestApplyToolsetScope: """Tests for _apply_toolset_scope helper.""" @@ -97,6 +112,122 @@ class TestApplyToolsetScope: assert op.mcp_servers == ["server-a"] assert op.mcp_tool_permissions == toolset_perms + @pytest.mark.asyncio + async def test_team_granted_toolset_is_served_to_a_key_without_its_own_grant(self): + """A team key whose own row carries no toolset grant is admitted to the toolset its team + holds (LIT-6029), scoped to that toolset's servers and tools.""" + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + toolset_perms = {"server-a": ["tool1"]} + auth = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=None) + with patch( + "litellm.proxy._experimental.mcp_server.server." + "global_mcp_server_manager.resolve_toolset_tool_permissions", + new=AsyncMock(return_value=toolset_perms), + ): + result = await _apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-123")) + + assert result.mcp_toolset_id == "toolset-123" + assert result.object_permission is not None + assert result.object_permission.mcp_servers == ["server-a"] + assert result.object_permission.mcp_tool_permissions == toolset_perms + + @pytest.mark.asyncio + async def test_a_non_admin_dashboard_session_is_pinned_as_its_admitted_user_instead_of_rewritten(self): + """The dashboard session acts as its admitted user, whose team grants resolve per source, so a + team-granted toolset is not capped by the user's own row: the row stays intact and the toolset + rides along as mcp_toolset_id (LIT-6029).""" + from litellm.constants import UI_SESSION_TOKEN_TEAM_ID + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") + own_row = LiteLLM_ObjectPermissionTable(object_permission_id="user-op", mcp_servers=["server-own"]) + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=own_row) + admitted.mcp_admitted_user_subject = True + granted = AsyncMock(return_value=frozenset({"toolset-123"})) + resolve = AsyncMock(return_value={"server-team": ["tool1"]}) + with patch( + "litellm.proxy._experimental.mcp_server.server." + "global_mcp_server_manager.resolve_toolset_tool_permissions", + new=resolve, + ): + result = await _apply_toolset_scope( + session, "toolset-123", acting_user=AsyncMock(return_value=admitted), granted=granted + ) + + assert granted.await_args is not None and granted.await_args.args[0].mcp_admitted_user_subject is True + assert result.mcp_admitted_user_subject is True + assert result.mcp_toolset_id == "toolset-123" + assert result.object_permission == own_row + resolve.assert_not_awaited() + + @pytest.mark.asyncio + async def test_a_gateway_admitted_user_without_the_toolset_in_any_source_is_denied(self): + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) + admitted.mcp_admitted_user_subject = True + granted = AsyncMock(return_value=frozenset({"toolset-other"})) + with pytest.raises(HTTPException) as exc_info: + await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + + assert exc_info.value.status_code == 403 + granted.assert_awaited_once_with(admitted) + + @pytest.mark.asyncio + async def test_a_resource_scoped_admitted_user_is_denied_a_team_toolset_on_another_server(self): + """A gateway bearer scoped to server-own (RFC 8707 resource) cannot open a team toolset whose + servers lie outside that resource, even though the team grants it (Devin Review 4150024267).""" + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) + admitted.mcp_admitted_user_subject = True + admitted.mcp_session_resource_server_id = "server-own" + admitted.requires_fresh_policy = True + granted = AsyncMock(return_value=frozenset({"toolset-123"})) + resolve = AsyncMock(return_value={"server-team": ["tool1"]}) + with patch( + "litellm.proxy._experimental.mcp_server.server." + "global_mcp_server_manager.resolve_toolset_tool_permissions", + new=resolve, + ): + with pytest.raises(HTTPException) as exc_info: + await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + + assert exc_info.value.status_code == 403 + resolve.assert_awaited_once_with(toolset_ids=["toolset-123"], requires_fresh_policy=True) + + @pytest.mark.asyncio + async def test_a_resource_scoped_admitted_user_opens_a_toolset_inside_its_resource(self): + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) + admitted.mcp_admitted_user_subject = True + admitted.mcp_session_resource_server_id = "server-team" + granted = AsyncMock(return_value=frozenset({"toolset-123"})) + resolve = AsyncMock(return_value={"server-team": ["tool1"], "server-other": ["tool2"]}) + with patch( + "litellm.proxy._experimental.mcp_server.server." + "global_mcp_server_manager.resolve_toolset_tool_permissions", + new=resolve, + ): + result = await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + + assert result.mcp_toolset_id == "toolset-123" + assert result.mcp_session_resource_server_id == "server-team" + resolve.assert_awaited_once_with(toolset_ids=["toolset-123"], requires_fresh_policy=False) + + @pytest.mark.asyncio + async def test_team_grant_for_another_toolset_does_not_admit_a_key_to_this_one(self): + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + auth = _make_auth(mcp_toolsets=[]) + auth.team_id = "team-a" + with pytest.raises(HTTPException) as exc_info: + await _apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-other")) + + assert exc_info.value.status_code == 403 + @pytest.mark.asyncio async def test_non_admin_no_object_permission_raises_403(self): """Non-admin key with object_permission=None is denied (no grants configured).""" @@ -250,6 +381,131 @@ class TestFetchMCPToolsetsAccess: assert len(result) == 2 mock_list.assert_called_once_with(mock_client, toolset_ids=["ts-1", "ts-2"]) + @pytest.mark.asyncio + async def test_team_granted_toolsets_are_listed_for_a_key_without_its_own_grant(self): + """GET /v1/mcp/toolset for a team key lists the team's toolsets (LIT-6029).""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_toolsets, + ) + + auth = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=None) + team_permission = LiteLLM_ObjectPermissionTable(object_permission_id="team-op", mcp_toolsets=["ts-team"]) + fake_toolsets = [MagicMock(toolset_id="ts-team")] + mock_client = MagicMock() + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=mock_client, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.list_mcp_toolsets", + new=AsyncMock(return_value=fake_toolsets), + ) as mock_list, + patch.object( + MCPRequestHandler, + "_get_team_object_permission", + new=AsyncMock(return_value=team_permission), + ), + ): + result = await fetch_mcp_toolsets(user_api_key_dict=auth) + + assert result == fake_toolsets + mock_list.assert_called_once_with(mock_client, toolset_ids=["ts-team"]) + + @pytest.mark.asyncio + async def test_admin_with_own_grants_is_not_narrowed_by_a_team_lookup(self): + """An admin's own grant list is the only filter; no team lookup runs for admins.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_toolsets, + ) + + auth = _make_auth(mcp_toolsets=["ts-1"]) + auth.user_role = LitellmUserRoles.PROXY_ADMIN + mock_client = MagicMock() + own_toolsets = [{"toolset_id": "ts-1", "toolset_name": "own"}] + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=mock_client, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.list_mcp_toolsets", + new=AsyncMock(return_value=own_toolsets), + ) as mock_list, + patch.object( + MCPRequestHandler, "_get_team_object_permission", new=AsyncMock(return_value=None) + ) as team_lookup, + ): + result = await fetch_mcp_toolsets(user_api_key_dict=auth) + + assert result == own_toolsets + mock_list.assert_called_once_with(mock_client, toolset_ids=["ts-1"]) + team_lookup.assert_not_awaited() + + +class TestFetchMCPToolsetAccess: + """Tests for GET /v1/mcp/toolset/{toolset_id} access control.""" + + @staticmethod + async def _fetch(auth: UserAPIKeyAuth, toolset_id: str, team_toolsets: list[str] | None): + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_toolset, + ) + + team_permission = ( + LiteLLM_ObjectPermissionTable(object_permission_id="team-op", mcp_toolsets=team_toolsets) + if team_toolsets is not None + else None + ) + toolset = MagicMock(toolset_id=toolset_id) + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_toolset", + new=AsyncMock(return_value=toolset), + ), + patch.object( + MCPRequestHandler, + "_get_team_object_permission", + new=AsyncMock(return_value=team_permission), + ), + ): + return await fetch_mcp_toolset(toolset_id=toolset_id, user_api_key_dict=auth) + + @pytest.mark.asyncio + async def test_team_granted_toolset_detail_is_served_to_a_key_without_its_own_grant(self): + auth = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=None) + + toolset = await self._fetch(auth, "ts-team", team_toolsets=["ts-team"]) + + assert toolset.toolset_id == "ts-team" + + @pytest.mark.asyncio + async def test_toolset_detail_stays_forbidden_when_neither_key_nor_team_holds_it(self): + from fastapi import HTTPException + + auth = _make_auth(mcp_toolsets=["ts-own"]) + auth.team_id = "team-a" + + with pytest.raises(HTTPException) as exc_info: + await self._fetch(auth, "ts-withheld", team_toolsets=["ts-team"]) + + assert exc_info.value.status_code == 403 + class TestToolsetPrefixResolution: """Regression for LIT-3419. diff --git a/tests/unit/proxy/_experimental/mcp_server/test_ui_session_utils.py b/tests/unit/proxy/_experimental/mcp_server/test_ui_session_utils.py index 816ccc5e7e6..293d9443ced 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_ui_session_utils.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_ui_session_utils.py @@ -6,13 +6,14 @@ import pytest from fastapi import HTTPException from litellm.constants import UI_SESSION_TOKEN_TEAM_ID -from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy._experimental.mcp_server.ui_session_utils import ( build_effective_auth_contexts, clone_user_api_key_auth_with_team, + granted_toolset_ids, + toolset_grant_contexts, resolve_ui_session_team_ids, ) +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth def test_clone_user_api_key_auth_with_team_creates_independent_copy(): @@ -258,3 +259,234 @@ async def test_admitted_user_context_carries_the_request_span(monkeypatch): assert (await acting_user_auth(user_auth)).parent_otel_span is parent_span assert (await build_effective_auth_contexts(user_auth))[-1].parent_otel_span is parent_span + + +def _toolset_permission(*toolset_ids: str) -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable( + object_permission_id=f"op-{'-'.join(toolset_ids)}", mcp_toolsets=list(toolset_ids) + ) + + +@pytest.mark.asyncio +async def test_granted_toolset_ids_unions_own_and_team_grants_over_every_effective_context(): + """A dashboard session of a user in two teams holds the toolsets of both teams plus the ones on + the user row itself, exactly the grant sources the aggregate /mcp listing expands.""" + session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") + team_a = UserAPIKeyAuth(team_id="team-a", user_id="user-1") + team_b = UserAPIKeyAuth(team_id="team-b", user_id="user-1", object_permission=_toolset_permission()) + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=_toolset_permission("ts-user")) + team_grants = {"team-a": _toolset_permission("ts-a", "ts-shared"), "team-b": _toolset_permission("ts-b")} + + async def effective_contexts(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + assert auth is session + return [team_a, team_b, admitted] + + async def team_permission(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return team_grants.get(auth.team_id or "") + + granted = await granted_toolset_ids(session, effective_contexts, team_permission) + + assert granted == frozenset({"ts-a", "ts-shared", "ts-b", "ts-user"}) + + +@pytest.mark.asyncio +async def test_granted_toolset_ids_is_empty_when_neither_key_nor_team_grants_a_toolset(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=_toolset_permission()) + + async def effective_contexts(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + return [auth] + + async def no_team_permission(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return None + + assert await granted_toolset_ids(key, effective_contexts, no_team_permission) == frozenset() + + +async def _same_context(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + return [auth] + + +async def _team_grants_ts_team(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return _toolset_permission("ts-team") + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "own", + [ + LiteLLM_ObjectPermissionTable(object_permission_id="op", mcp_toolsets=["ts-own"]), + LiteLLM_ObjectPermissionTable(object_permission_id="op", mcp_servers=["srv-own"]), + LiteLLM_ObjectPermissionTable(object_permission_id="op", mcp_tool_permissions={"srv-own": ["add"]}), + LiteLLM_ObjectPermissionTable(object_permission_id="op", mcp_access_groups=["group-own"]), + ], +) +async def test_a_key_declaring_its_own_mcp_grant_does_not_inherit_the_team_toolsets(own): + """The key/team rule of the aggregate listing: a key's own MCP grant is a ceiling the team cannot widen.""" + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=own) + + granted = await granted_toolset_ids(key, _same_context, _team_grants_ts_team, require_key_access=False) + + assert granted == frozenset(own.mcp_toolsets or ()) + + +@pytest.mark.asyncio +async def test_require_key_mcp_access_defined_stops_a_key_inheriting_team_toolsets_but_not_a_session(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a") + session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") + + assert await granted_toolset_ids(key, _same_context, _team_grants_ts_team, require_key_access=False) == {"ts-team"} + assert await granted_toolset_ids(key, _same_context, _team_grants_ts_team, require_key_access=True) == frozenset() + assert await granted_toolset_ids(session, _same_context, _team_grants_ts_team, require_key_access=True) == { + "ts-team" + } + + +@pytest.mark.asyncio +async def test_toolset_grant_contexts_of_a_virtual_key_is_the_key_alone(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a") + + async def never(auth: UserAPIKeyAuth) -> None: + raise AssertionError("a virtual key has no admitted sources") + + assert await toolset_grant_contexts(key, admitted_context=never, admitted_sources=never) == (key,) + + +def _admitted(user_id: str, own: LiteLLM_ObjectPermissionTable | None = None) -> UserAPIKeyAuth: + subject = UserAPIKeyAuth(user_id=user_id, object_permission=own) + subject.mcp_admitted_user_subject = True + return subject + + +@pytest.mark.asyncio +async def test_toolset_grant_contexts_of_a_dashboard_session_are_its_admitted_users_grant_sources(): + """The dashboard session fans out through the same roster-checked source builder as the aggregate + /mcp resolution, applied to the admitted user it acts as, so a cached membership a team has since + revoked never reaches the toolset check.""" + session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") + admitted = _admitted("user-1") + own_source = UserAPIKeyAuth(user_id="user-1") + team_source = UserAPIKeyAuth(user_id="user-1", team_id="team-a") + + async def admitted_context(auth: UserAPIKeyAuth) -> UserAPIKeyAuth: + assert auth is session + return admitted + + async def admitted_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + assert auth is admitted + return [own_source, team_source] + + assert await toolset_grant_contexts(session, admitted_context, admitted_sources) == (own_source, team_source) + + +@pytest.mark.asyncio +async def test_toolset_grant_contexts_of_a_gateway_admitted_user_are_its_own_grant_sources(): + admitted = _admitted("user-1") + team_source = UserAPIKeyAuth(user_id="user-1", team_id="team-a") + + async def no_dashboard_context(auth: UserAPIKeyAuth) -> None: + return None + + async def admitted_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + assert auth is admitted + return [team_source] + + assert await toolset_grant_contexts(admitted, no_dashboard_context, admitted_sources) == (team_source,) + + +@pytest.mark.asyncio +async def test_a_source_declaring_its_own_mcp_grant_never_reads_its_team(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=_toolset_permission("ts-own")) + team_reads: list[str | None] = [] # mutable-ok: records the lookups the code under test performs + + async def team_permission(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + team_reads.append(auth.team_id) + return _toolset_permission("ts-team") + + granted = await granted_toolset_ids(key, _same_context, team_permission, require_key_access=False) + + assert granted == {"ts-own"} + assert team_reads == [] + + +async def _team_a_unreadable(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + if auth.team_id == "team-a": + raise RuntimeError("team row unreadable") + return _toolset_permission("ts-b") + + +@pytest.mark.asyncio +async def test_an_unreadable_team_grants_nothing_while_the_direct_and_other_team_grants_still_count(): + """A dashboard user whose own row grants ts-user and who sits on team-a and team-b keeps ts-user and + ts-b when team-a cannot be read; team-a itself contributes nothing rather than failing the lookup.""" + admitted = _admitted("user-1", _toolset_permission("ts-user")) + team_a = UserAPIKeyAuth(user_id="user-1", team_id="team-a") + team_b = UserAPIKeyAuth(user_id="user-1", team_id="team-b") + + async def sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + return [admitted, team_a, team_b] + + assert await granted_toolset_ids(admitted, sources, _team_a_unreadable) == {"ts-user", "ts-b"} + + +@pytest.mark.asyncio +async def test_a_key_whose_only_grant_source_is_an_unreadable_team_is_granted_nothing(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a") + + assert await granted_toolset_ids(key, _same_context, _team_a_unreadable, require_key_access=False) == frozenset() + + +async def _hydrates_op_key_to_srv_own(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + if auth.object_permission is not None: + return auth.object_permission + if auth.object_permission_id == "op-key": + return LiteLLM_ObjectPermissionTable(object_permission_id="op-key", mcp_servers=["srv-own"]) + return None + + +@pytest.mark.asyncio +async def test_a_key_cached_with_its_own_grant_unhydrated_is_scoped_to_that_grant_not_its_team(): + """The main auth flow can cache a key with object_permission_id set and object_permission None. The + row it names is the key's ceiling, so it is loaded and read as the key's own grant instead of letting the + key inherit its team's toolsets.""" + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission_id="op-key") + + granted = await granted_toolset_ids( + key, + _same_context, + _team_grants_ts_team, + require_key_access=False, + own_object_permission=_hydrates_op_key_to_srv_own, + ) + + assert granted == frozenset() + + +async def _own_row_unreadable(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + raise RuntimeError("object permission row unreadable") + + +async def _own_row_gone(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("load_own", [_own_row_unreadable, _own_row_gone]) +async def test_a_key_naming_an_own_grant_that_cannot_be_read_is_granted_nothing_rather_than_its_team(load_own): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission_id="op-key") + + granted = await granted_toolset_ids( + key, _same_context, _team_grants_ts_team, require_key_access=False, own_object_permission=load_own + ) + + assert granted == frozenset() + + +@pytest.mark.asyncio +async def test_a_key_naming_no_own_grant_is_not_hydrated_before_inheriting_its_team(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a") + + granted = await granted_toolset_ids( + key, _same_context, _team_grants_ts_team, require_key_access=False, own_object_permission=_own_row_unreadable + ) + + assert granted == {"ts-team"} diff --git a/tests/unit/proxy/management_endpoints/test_team_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_endpoints.py index 32b3c919b07..6d56d325dc3 100644 --- a/tests/unit/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_endpoints.py @@ -13596,6 +13596,63 @@ async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeyp assert mock_audit.call_args.kwargs["team_alias"] == "list-audit" +@pytest.mark.asyncio +async def test_team_member_add_evicts_the_cached_team_roster(monkeypatch): + """Roster checks read the team through get_team_object, so a cached pre-add roster must be dropped.""" + from litellm.proxy._types import TeamMemberAddRequest + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.management_endpoints.team_endpoints import team_member_add + + team_id = "team-roster-evict" + team_row = LiteLLM_TeamTable(team_id=team_id, team_alias="roster-evict", members_with_roles=[]) + cache = UserApiKeyCache() + cache.set_cache(key=f"team_id:{team_id}", value=team_row) + cache.set_cache(key="team_alias:roster-evict", value=team_row) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id") + + joined_user = LiteLLM_UserTable(user_id="joiner", max_budget=None, spend=0.0, models=[]) + updated_team = MagicMock() + updated_team.model_dump.return_value = {"team_id": team_id, "members_with_roles": []} + + with ( + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + new_callable=AsyncMock, + return_value=team_row, + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._validate_team_member_add_permissions", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._validate_and_populate_member_user_info", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._resolve_existing_member_user_ids", + new_callable=AsyncMock, + return_value=frozenset(), + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team", + new_callable=AsyncMock, + return_value=(updated_team, [joined_user], []), + ), + patch("litellm.proxy.management_endpoints.team_endpoints._schedule_team_member_add_audit_logs"), + ): + await team_member_add( + data=TeamMemberAddRequest(team_id=team_id, member=Member(user_id="joiner", role="user")), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1"), + ) + + assert cache.get_cache(key=f"team_id:{team_id}") is None + assert cache.get_cache(key="team_alias:roster-evict") is None + + class _RecordingAuditLogger(CustomLogger): def __init__(self) -> None: super().__init__() diff --git a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index 643944673a2..2699f9445c9 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -8,7 +8,8 @@ from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import HTTPException -from mcp.types import CallToolResult, TextContent, Tool as MCPTool +from mcp.types import CallToolResult, TextContent +from mcp.types import Tool as MCPTool from openai.types.responses.tool_param import Mcp from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing @@ -110,9 +111,7 @@ def test_extract_tool_calls_from_chat_response_handles_tool_calls(): object="chat.completion", ) - tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response( - response - ) + tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(response) assert len(tool_calls) == 1 assert tool_calls[0]["function"]["name"] == "foo" @@ -182,9 +181,7 @@ def test_transform_mcp_tools_to_openai_uses_chat_format(monkeypatch): fake_transform_responses, ) - chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( - ["tool"], target_format="chat" - ) + chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"], target_format="chat") resp_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"]) assert chat_tools == [{"chat": True}] @@ -304,9 +301,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n ) from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm - _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock( - return_value=fake_server - ) + _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=fake_server) tool_name = "my_deepwiki-read_wiki_structure" tool_calls = [ @@ -380,7 +375,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey fake_manager = types.SimpleNamespace( get_registry=MagicMock(return_value={}), - call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")) + call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")), ) monkeypatch.setattr( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -388,9 +383,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey ) tool_name = "deepwiki-read_wiki_structure" - tool_calls = [ - {"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}} - ] + tool_calls = [{"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}}] user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") @@ -408,10 +401,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey post_call_failure_hook.assert_awaited_once() assert post_call_failure_hook.await_args is not None - assert ( - post_call_failure_hook.await_args.kwargs.get("route") - == "/responses/mcp/call_tool" - ) + assert post_call_failure_hook.await_args.kwargs.get("route") == "/responses/mcp/call_tool" @pytest.mark.asyncio @@ -434,9 +424,7 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio # NOTE: Don't patch via dotted string path here because `litellm.responses` # is a function attribute on the `litellm` package (shadowing the submodule), # which breaks monkeypatch's importpath resolution. - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" @@ -516,7 +504,9 @@ async def test_execute_tool_calls_applies_post_call_hook_content(monkeypatch): logging_obj = MagicMock() logging_obj.model_call_details = {} - logging_obj.async_post_mcp_tool_call_hook = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="[REDACTED]")], is_error=True)) + logging_obj.async_post_mcp_tool_call_hook = AsyncMock( + return_value=CallToolResult(content=[TextContent(type="text", text="[REDACTED]")], is_error=True) + ) logging_obj.async_success_handler = AsyncMock() handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", lambda *_args, **_kwargs: (logging_obj, None)) @@ -679,9 +669,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( user_api_key_auth=user_auth, - mcp_tools_with_litellm_proxy=[ - {"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"} - ], + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}], ) forwarded: Final = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(tools) @@ -700,9 +688,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch def test_get_parent_request_tags_from_metadata(): - tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags( - {"metadata": {"tags": ["team-a", "prod"]}} - ) + tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags({"metadata": {"tags": ["team-a", "prod"]}}) assert tags == ["team-a", "prod"] @@ -739,9 +725,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch): await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( user_api_key_auth=types.SimpleNamespace(api_key="k", user_id="u"), - mcp_tools_with_litellm_proxy=[ - {"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"} - ], + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}], request_tags=["team-a"], ) @@ -761,9 +745,7 @@ async def test_execute_tool_calls_exposes_sanitized_client_headers_to_logging(mo captured.update(kwargs) return None, None - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" @@ -789,9 +771,7 @@ async def test_execute_tool_calls_propagates_request_tags_to_function_setup(monk captured.update(kwargs) return None, None - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" @@ -1171,7 +1151,9 @@ def test_create_follow_up_input_keeps_each_reasoning_item_before_its_function_ca "function_call_output", "function_call_output", ] - assert [cast(dict[str, Any], item).get("id") or cast(dict[str, Any], item).get("call_id") for item in follow_up[1:5]] == [ + assert [ + cast(dict[str, Any], item).get("id") or cast(dict[str, Any], item).get("call_id") for item in follow_up[1:5] + ] == [ "rs_1", "call-1", "rs_2", @@ -1229,16 +1211,20 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false( async def fake_aresponses(**kwargs: Any) -> ResponsesAPIResponse: captured_calls.append(kwargs) - return first_response if len(captured_calls) == 1 else ResponsesAPIResponse( - id="resp_follow_up", - created_at=1234567891, - model="gpt-5", - object="response", - status="completed", - output=[], - parallel_tool_calls=False, - tool_choice="auto", - tools=[], + return ( + first_response + if len(captured_calls) == 1 + else ResponsesAPIResponse( + id="resp_follow_up", + created_at=1234567891, + model="gpt-5", + object="response", + status="completed", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ) ) async def fake_process(**kwargs: Any) -> tuple[list[Any], dict[str, str]]: @@ -1279,12 +1265,14 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false( @pytest.mark.asyncio async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: pytest.MonkeyPatch): - from litellm.proxy._experimental.mcp_server import operations - from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._experimental.mcp_server import mcp_server_manager, operations headers: Final = { - "x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a", - "x-mcp-deepwiki-authorization": "upstream-sentinel", "authorization": "proxy-sentinel", + "x-app-id": "app-a", + "x-nuid": "user-a", + "x-user-id": "identity-a", + "x-mcp-deepwiki-authorization": "upstream-sentinel", + "authorization": "proxy-sentinel", } manager: Final = types.SimpleNamespace( get_registry=MagicMock(return_value={}), @@ -1298,12 +1286,21 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[])) monkeypatch.setattr(operations, "function_setup", setup) response: Final = ResponsesAPIResponse( - id="resp_test", created_at=1234567891, model="test-model", object="response", - status="completed", output=[], parallel_tool_calls=False, tool_choice="auto", tools=[], + id="resp_test", + created_at=1234567891, + model="test-model", + object="response", + status="completed", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], ) monkeypatch.setattr(responses_main, "aresponses", AsyncMock(return_value=response)) result: Final = await responses_main.aresponses_api_with_mcp( - input="hi", model="test-model", tools=[{"type": "mcp", "server_url": "litellm_proxy"}], + input="hi", + model="test-model", + tools=[{"type": "mcp", "server_url": "litellm_proxy"}], secret_fields={"raw_headers": headers}, ) assert result is response @@ -1311,3 +1308,84 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py logged: Final = setup.call_args.kwargs["metadata"]["headers"] assert logged == {"x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a"} assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel" + + +def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace: + return types.SimpleNamespace( + get_registry=MagicMock(return_value={}), + get_allowed_mcp_servers=AsyncMock(return_value=[]), + get_mcp_servers_from_ids=MagicMock(return_value=[]), + get_mcp_server_by_name=MagicMock(return_value=None), + get_toolset_by_name_cached=AsyncMock(return_value=types.SimpleNamespace(toolset_id=toolset_id)), + resolve_toolset_tool_permissions=AsyncMock(return_value={server_id: ["add"]}), + ) + + +async def _tools_listing_kwargs_for_toolset_url(monkeypatch, team_toolset_id: str) -> dict[str, object]: + from litellm.proxy._experimental.mcp_server.ui_session_utils import granted_toolset_ids + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LitellmUserRoles, UserAPIKeyAuth + + mock_get_tools = AsyncMock(return_value=AggregateToolListing(tools=[], outcomes={})) + monkeypatch.setattr("litellm.proxy._experimental.mcp_server.server._get_tools_from_mcp_servers", mock_get_tools) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + _toolset_gateway_manager("ts-granted", "srv-1"), + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock()) + + async def team_permission(context: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable(object_permission_id="op-team", mcp_toolsets=[team_toolset_id]) + + async def granted_through_team(context: UserAPIKeyAuth) -> frozenset[str]: + return await granted_toolset_ids(context, team_object_permission=team_permission, require_key_access=False) + + team_key: Final = UserAPIKeyAuth(api_key="sk-team", team_id="team-1", user_role=LitellmUserRoles.INTERNAL_USER) + await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + user_api_key_auth=team_key, + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/team-toolset"}], + granted_toolsets=granted_through_team, + ) + assert mock_get_tools.await_args is not None + return mock_get_tools.await_args.kwargs + + +@pytest.mark.asyncio +async def test_toolset_gateway_url_scopes_a_team_granted_toolset_for_a_key_without_its_own_grant(monkeypatch): + kwargs: Final = await _tools_listing_kwargs_for_toolset_url(monkeypatch, team_toolset_id="ts-granted") + scoped = kwargs["user_api_key_auth"].object_permission + assert scoped is not None + assert scoped.mcp_servers == ["srv-1"] + assert scoped.mcp_tool_permissions == {"srv-1": ["add"]} + assert kwargs["mcp_servers"] is None + + +@pytest.mark.asyncio +async def test_toolset_gateway_url_skips_a_toolset_the_team_does_not_grant(monkeypatch): + kwargs: Final = await _tools_listing_kwargs_for_toolset_url(monkeypatch, team_toolset_id="ts-other") + assert kwargs["user_api_key_auth"].object_permission is None + assert kwargs["mcp_servers"] is None + + +@pytest.mark.asyncio +async def test_apply_toolset_permissions_pins_the_auth_to_explicit_grants_only(monkeypatch: pytest.MonkeyPatch): + """A toolset gateway URL must not widen to operator-open (allow_all_keys) servers.""" + from litellm.proxy._types import UserAPIKeyAuth + + fake_manager = types.SimpleNamespace( + resolve_toolset_tool_permissions=AsyncMock(return_value={"srv-1": ["add"]}), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + fake_manager, + ) + + scoped = await LiteLLM_Proxy_MCP_Handler._apply_toolset_permissions( + resolved_toolset_ids=["ts-1"], + resolved_mcp_servers=[], + user_api_key_auth=UserAPIKeyAuth(api_key="sk-test", user_id="u1"), + ) + + assert scoped.mcp_explicit_grants_only is True + assert scoped.object_permission is not None + assert scoped.object_permission.mcp_servers == ["srv-1"] + assert scoped.object_permission.mcp_tool_permissions == {"srv-1": ["add"]}