From 5736fd32d96916ce088eb231b81e84bc5c52f7de Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Wed, 11 Feb 2026 18:07:24 -0800 Subject: [PATCH] MCP fixes * fix(oldteams.tsx): show policies when creating * fix(proxy/_types.py): ensure mcp rest endpoints can be called by virtual key ensures UI works with virtual key testing mcp endpoints * refactor: migrate get object permissions table logic to happen in user api key auth - allows functions to trust user api key object they receive has what they need * fix(rest_endpoints.py): filter for allowed tools based on what key has access to * fix(mcp_server_manager.py): ensure only allowed MCP's are returned to the user, via rest endpoints --- .../mcp_server/auth/user_api_key_auth_mcp.py | 80 ++--- .../mcp_server/mcp_server_manager.py | 40 ++- .../mcp_server/rest_endpoints.py | 43 ++- litellm/proxy/_types.py | 3 +- .../auth/agent_permission_handler.py | 148 ++++----- litellm/proxy/auth/auth_checks.py | 46 +++ .../auth/test_user_api_key_auth_mcp.py | 129 ++++---- .../mcp_server/test_rest_endpoints.py | 298 +++++++++++++++++- .../auth/test_object_permission_loading.py | 151 +++++++++ .../test_mcp_management_endpoints.py | 73 ++++- .../src/components/OldTeams.tsx | 48 ++- 11 files changed, 818 insertions(+), 241 deletions(-) create mode 100644 tests/test_litellm/proxy/auth/test_object_permission_loading.py 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 786cfbfb008..548e3bc3dbf 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 @@ -6,7 +6,12 @@ from starlette.requests import Request from starlette.types import Scope from litellm._logging import verbose_logger -from litellm.proxy._types import LiteLLM_TeamTable, ProxyException, SpecialHeaders, UserAPIKeyAuth +from litellm.proxy._types import ( + LiteLLM_TeamTable, + ProxyException, + SpecialHeaders, + UserAPIKeyAuth, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -372,45 +377,31 @@ class MCPRequestHandler: return [] @staticmethod - async def _get_key_object_permission( + def _get_key_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): - """Helper to get key object_permission from cache or DB.""" - from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) + """ + Get key object_permission - already loaded by get_key_object() in main auth flow. + Note: object_permission is automatically populated when the key is fetched via + get_key_object() in litellm/proxy/auth/auth_checks.py + """ if not user_api_key_auth: return None - # Already loaded - if user_api_key_auth.object_permission: - return user_api_key_auth.object_permission - - # Need to fetch from DB - if user_api_key_auth.object_permission_id and prisma_client: - return await get_object_permission( - object_permission_id=user_api_key_auth.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - - return None + return user_api_key_auth.object_permission @staticmethod async def _get_team_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): - """Helper to get team object_permission from cache or DB.""" - from litellm.proxy.auth.auth_checks import ( - get_object_permission, - get_team_object, - ) + """ + Get team object_permission - automatically loaded by get_team_object() in main auth flow. + + Note: object_permission is automatically populated when the team is fetched via + get_team_object() in litellm/proxy/auth/auth_checks.py + """ + from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.proxy_server import ( prisma_client, proxy_logging_obj, @@ -423,7 +414,7 @@ class MCPRequestHandler: if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client: return None - # First get the team object (which may have object_permission already loaded) + # Get the team object (which has object_permission already loaded) team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( team_id=user_api_key_auth.team_id, prisma_client=prisma_client, @@ -435,21 +426,7 @@ class MCPRequestHandler: if not team_obj: return None - # Already loaded - if team_obj.object_permission: - return team_obj.object_permission - - # Need to fetch from DB using object_permission_id - if team_obj.object_permission_id: - return await get_object_permission( - object_permission_id=team_obj.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - - return None + return team_obj.object_permission @staticmethod async def get_allowed_tools_for_server( @@ -471,8 +448,8 @@ class MCPRequestHandler: return None try: - # Get key and team object permissions - key_obj_perm = await MCPRequestHandler._get_key_object_permission( + # Get key and team object permissions (already loaded in main auth flow) + key_obj_perm = MCPRequestHandler._get_key_object_permission( user_api_key_auth ) team_obj_perm = await MCPRequestHandler._get_team_object_permission( @@ -559,7 +536,8 @@ class MCPRequestHandler: user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> List[str]: try: - key_object_permission = await MCPRequestHandler._get_key_object_permission( + # Get key object permission (already loaded in main auth flow) + key_object_permission = MCPRequestHandler._get_key_object_permission( user_api_key_auth ) if key_object_permission is None: @@ -591,12 +569,10 @@ class MCPRequestHandler: """ Get allowed MCP servers for a team. - Uses the helper _get_team_object_permission which: - 1. First checks if object_permission is already loaded on the team - 2. If not, fetches from DB using object_permission_id if it exists + Note: object_permission is automatically loaded by get_team_object() in main auth flow. """ try: - # Use the helper method that properly handles fetching from DB if needed + # Get team object permission (already loaded in main auth flow) object_permissions = await MCPRequestHandler._get_team_object_permission( user_api_key_auth ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 24eae430d85..fd251488db4 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -70,9 +70,7 @@ try: from mcp.shared.tool_name_validation import ( validate_tool_name, # pyright: ignore[reportAssignmentType] ) - from mcp.shared.tool_name_validation import ( - SEP_986_URL, - ) + from mcp.shared.tool_name_validation import SEP_986_URL except ImportError: from pydantic import BaseModel @@ -673,24 +671,47 @@ class MCPServerManager: return [ server.server_id for server in self.get_registry().values() - if server.allow_all_keys + if server.allow_all_keys is True ] async def get_allowed_mcp_servers( self, user_api_key_auth: Optional[UserAPIKeyAuth] = None ) -> List[str]: """ - Get the allowed MCP Servers for the user + Get the allowed MCP Servers for the user. + + Priority: + 1. If object_permission.mcp_servers is explicitly set, use it (even for admins) + 2. If admin and no object_permission, return all servers + 3. Otherwise, use standard permission checks """ from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view - # If admin, get all servers - if user_api_key_auth and _user_has_admin_view(user_api_key_auth): - return list(self.get_registry().keys()) - allow_all_server_ids = self.get_allow_all_keys_server_ids() try: + # Check if object_permission.mcp_servers is explicitly set + has_explicit_object_permission = False + if user_api_key_auth and user_api_key_auth.object_permission: + # Check if mcp_servers is explicitly set (not None, empty list is valid) + if user_api_key_auth.object_permission.mcp_servers is not None: + has_explicit_object_permission = True + verbose_logger.debug( + f"Object permission mcp_servers explicitly set: {user_api_key_auth.object_permission.mcp_servers}" + ) + + # If admin but NO explicit object permission, get all servers + if ( + user_api_key_auth + and _user_has_admin_view(user_api_key_auth) + and not has_explicit_object_permission + ): + verbose_logger.debug( + "Admin user without explicit object_permission - returning all servers" + ) + return list(self.get_registry().keys()) + + # Get allowed servers from object permissions (respects object_permission even for admins) allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers( user_api_key_auth ) @@ -2246,6 +2267,7 @@ class MCPServerManager: from litellm.proxy.proxy_server import ( general_settings as proxy_general_settings, ) + return proxy_general_settings except ImportError: # Fallback if proxy_server not available diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 43b388993d9..aed81afd254 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -37,6 +37,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.server import ( ListMCPToolsRestAPIResponseObject, MCPServer, + _tool_name_matches, execute_mcp_tool, filter_tools_by_allowed_tools, ) @@ -159,6 +160,7 @@ if MCP_AVAILABLE: server, server_auth_header, raw_headers: Optional[Dict[str, str]] = None, + user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): """Helper function to get tools for a single server.""" tools = await global_mcp_server_manager._get_tools_from_server( @@ -173,6 +175,29 @@ if MCP_AVAILABLE: if server.allowed_tools is not None and len(server.allowed_tools) > 0: tools = filter_tools_by_allowed_tools(tools, server) + # Filter tools based on user_api_key_auth.object_permission.mcp_tool_permissions + # This provides per-key/team/org control over which tools can be accessed + if ( + user_api_key_auth + and user_api_key_auth.object_permission + and user_api_key_auth.object_permission.mcp_tool_permissions + ): + allowed_tools_for_server = ( + user_api_key_auth.object_permission.mcp_tool_permissions.get( + server.server_id + ) + ) + if ( + allowed_tools_for_server is not None + and len(allowed_tools_for_server) > 0 + ): + # Filter tools to only include those in the allowed list + tools = [ + tool + for tool in tools + if _tool_name_matches(tool.name, allowed_tools_for_server) + ] + return _create_tool_response_objects(tools, server.mcp_info) async def _resolve_allowed_mcp_servers_for_tool_call( @@ -197,9 +222,7 @@ if MCP_AVAILABLE: ) allowed_mcp_servers: List[MCPServer] = [] for allowed_server_id in allowed_server_ids_set: - server = global_mcp_server_manager.get_mcp_server_by_id( - allowed_server_id - ) + server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id) if server is not None: allowed_mcp_servers.append(server) return allowed_mcp_servers @@ -276,9 +299,7 @@ if MCP_AVAILABLE: "message": f"The key is not allowed to access server {server_id}", }, ) - server = global_mcp_server_manager.get_mcp_server_by_id( - server_id - ) + server = global_mcp_server_manager.get_mcp_server_by_id(server_id) if server is None: return { "tools": [], @@ -292,7 +313,10 @@ if MCP_AVAILABLE: try: list_tools_result = await _get_tools_for_single_server( - server, server_auth_header, raw_headers_from_request + server, + server_auth_header, + raw_headers_from_request, + user_api_key_dict, ) except Exception as e: verbose_logger.exception( @@ -328,7 +352,10 @@ if MCP_AVAILABLE: try: tools_result = await _get_tools_for_single_server( - server, server_auth_header, raw_headers_from_request + server, + server_auth_header, + raw_headers_from_request, + user_api_key_dict, ) list_tools_result.extend(tools_result) except Exception as e: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c81c1e7e83f..87ff4a66e08 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -420,7 +420,8 @@ class LiteLLMRoutes(enum.Enum): "/mcp/tools", "/mcp/tools/list", "/mcp/tools/call", - # Read-only MCP discovery endpoint (virtual keys may be allowed here) + "/mcp-rest/tools/list", + "/mcp-rest/tools/call", "/v1/mcp/server", ] diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 96e7a21cc33..bf3256cf47b 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -21,7 +21,7 @@ class AgentRequestHandler: 1. Key-level agent permissions 2. Team-level agent permissions 3. Agent access group resolution - + Follows the same inheritance logic as MCP: - If team has restrictions and key has restrictions: use intersection - If team has restrictions and key has none: inherit from team @@ -35,7 +35,7 @@ class AgentRequestHandler: ) -> List[str]: """ Get list of allowed agent IDs for the given user/key based on permissions. - + Returns: List[str]: List of allowed agent IDs. Empty list means no restrictions (allow all). """ @@ -45,7 +45,9 @@ class AgentRequestHandler: await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth) ) allowed_agents_for_team = ( - await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth) + await AgentRequestHandler._get_allowed_agents_for_team( + user_api_key_auth + ) ) # If team has agent restrictions, handle inheritance and intersection logic @@ -73,62 +75,48 @@ class AgentRequestHandler: ) -> bool: """ Check if a specific agent is allowed for the given user/key. - + Args: agent_id: The agent ID to check user_api_key_auth: User authentication info - + Returns: bool: True if agent is allowed, False otherwise """ allowed_agents = await AgentRequestHandler.get_allowed_agents(user_api_key_auth) - + # Empty list means no restrictions - allow all if len(allowed_agents) == 0: return True - + return agent_id in allowed_agents @staticmethod - async def _get_key_object_permission( + def _get_key_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> Optional[LiteLLM_ObjectPermissionTable]: - """Helper to get key object_permission from cache or DB.""" - from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) + """ + Get key object_permission - already loaded by get_key_object() in main auth flow. + Note: object_permission is automatically populated when the key is fetched via + get_key_object() in litellm/proxy/auth/auth_checks.py + """ if not user_api_key_auth: return None - # Already loaded - if user_api_key_auth.object_permission: - return user_api_key_auth.object_permission - - # Need to fetch from DB - if user_api_key_auth.object_permission_id and prisma_client: - return await get_object_permission( - object_permission_id=user_api_key_auth.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - - return None + return user_api_key_auth.object_permission @staticmethod async def _get_team_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> Optional[LiteLLM_ObjectPermissionTable]: - """Helper to get team object_permission from cache or DB.""" - from litellm.proxy.auth.auth_checks import ( - get_object_permission, - get_team_object, - ) + """ + Get team object_permission - automatically loaded by get_team_object() in main auth flow. + + Note: object_permission is automatically populated when the team is fetched via + get_team_object() in litellm/proxy/auth/auth_checks.py + """ + from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.proxy_server import ( prisma_client, proxy_logging_obj, @@ -138,7 +126,7 @@ class AgentRequestHandler: if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client: return None - # First get the team object (which may have object_permission already loaded) + # Get the team object (which has object_permission already loaded) team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( team_id=user_api_key_auth.team_id, prisma_client=prisma_client, @@ -150,21 +138,7 @@ class AgentRequestHandler: if not team_obj: return None - # Already loaded - if team_obj.object_permission: - return team_obj.object_permission - - # Need to fetch from DB using object_permission_id - if team_obj.object_permission_id: - return await get_object_permission( - object_permission_id=team_obj.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - - return None + return team_obj.object_permission @staticmethod async def _get_allowed_agents_for_key( @@ -172,31 +146,16 @@ class AgentRequestHandler: ) -> List[str]: """ Get allowed agents for a key from its object_permission. - """ - from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) + Note: object_permission is already loaded by get_key_object() in main auth flow. + """ if user_api_key_auth is None: return [] - if user_api_key_auth.object_permission_id is None: - return [] - - if prisma_client is None: - verbose_logger.debug("prisma_client is None") - return [] - try: - key_object_permission = await get_object_permission( - object_permission_id=user_api_key_auth.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, + # Get key object permission (already loaded in main auth flow) + key_object_permission = AgentRequestHandler._get_key_object_permission( + user_api_key_auth ) if key_object_permission is None: return [] @@ -205,8 +164,10 @@ class AgentRequestHandler: direct_agents = key_object_permission.agents or [] # Get agents from access groups - access_group_agents = await AgentRequestHandler._get_agents_from_access_groups( - key_object_permission.agent_access_groups or [] + access_group_agents = ( + await AgentRequestHandler._get_agents_from_access_groups( + key_object_permission.agent_access_groups or [] + ) ) # Combine both lists @@ -222,6 +183,8 @@ class AgentRequestHandler: ) -> List[str]: """ Get allowed agents for a team from its object_permission. + + Note: object_permission is already loaded by get_team_object() in main auth flow. """ if user_api_key_auth is None: return [] @@ -230,7 +193,7 @@ class AgentRequestHandler: return [] try: - # Use the helper method that properly handles fetching from DB if needed + # Get team object permission (already loaded in main auth flow) object_permissions = await AgentRequestHandler._get_team_object_permission( user_api_key_auth ) @@ -242,8 +205,10 @@ class AgentRequestHandler: direct_agents = object_permissions.agents or [] # Get agents from access groups - access_group_agents = await AgentRequestHandler._get_agents_from_access_groups( - object_permissions.agent_access_groups or [] + access_group_agents = ( + await AgentRequestHandler._get_agents_from_access_groups( + object_permissions.agent_access_groups or [] + ) ) # Combine both lists @@ -284,9 +249,7 @@ class AgentRequestHandler: for agent in agents: agent_ids.add(agent.agent_id) except Exception as e: - verbose_logger.debug( - f"Error getting agents from access groups: {e}" - ) + verbose_logger.debug(f"Error getting agents from access groups: {e}") return agent_ids @staticmethod @@ -306,16 +269,16 @@ class AgentRequestHandler: ) # Use the helper for DB agents - db_agent_ids = await AgentRequestHandler._get_db_agent_ids_for_access_groups( - prisma_client, access_groups + db_agent_ids = ( + await AgentRequestHandler._get_db_agent_ids_for_access_groups( + prisma_client, access_groups + ) ) agent_ids.update(db_agent_ids) return list(agent_ids) except Exception as e: - verbose_logger.warning( - f"Failed to get agents from access groups: {str(e)}" - ) + verbose_logger.warning(f"Failed to get agents from access groups: {str(e)}") return [] @staticmethod @@ -326,11 +289,15 @@ class AgentRequestHandler: Get list of agent access groups for the given user/key based on permissions. """ access_groups: List[str] = [] - access_groups_for_key = await AgentRequestHandler._get_agent_access_groups_for_key( - user_api_key_auth + access_groups_for_key = ( + await AgentRequestHandler._get_agent_access_groups_for_key( + user_api_key_auth + ) ) - access_groups_for_team = await AgentRequestHandler._get_agent_access_groups_for_team( - user_api_key_auth + access_groups_for_team = ( + await AgentRequestHandler._get_agent_access_groups_for_team( + user_api_key_auth + ) ) # If team has access groups, then key must have a subset of the team's access groups @@ -378,7 +345,9 @@ class AgentRequestHandler: return key_object_permission.agent_access_groups or [] except Exception as e: - verbose_logger.warning(f"Failed to get agent access groups for key: {str(e)}") + verbose_logger.warning( + f"Failed to get agent access groups for key: {str(e)}" + ) return [] @staticmethod @@ -425,4 +394,3 @@ class AgentRequestHandler: f"Failed to get agent access groups for team: {str(e)}" ) return [] - diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index c6093172932..76ec67ab10e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1368,6 +1368,22 @@ async def _get_team_object_from_user_api_key_cache( raise Exception _response = LiteLLM_TeamTableCachedObj(**response.dict()) + + # Load object_permission if object_permission_id exists but object_permission is not loaded + if _response.object_permission_id and not _response.object_permission: + try: + _response.object_permission = await get_object_permission( + object_permission_id=_response.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Failed to load object_permission for team {team_id} with object_permission_id={_response.object_permission_id}: {e}" + ) + # save the team object to cache await _cache_team_object( team_id=team_id, @@ -1550,6 +1566,21 @@ async def get_team_object_by_alias( team = teams[0] team_obj = LiteLLM_TeamTableCachedObj(**team.model_dump()) + # Load object_permission if object_permission_id exists but object_permission is not loaded + if team_obj.object_permission_id and not team_obj.object_permission: + try: + team_obj.object_permission = await get_object_permission( + object_permission_id=team_obj.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Failed to load object_permission for team {team_obj.team_id} with object_permission_id={team_obj.object_permission_id}: {e}" + ) + # Cache the result by both alias and team_id await user_api_key_cache.async_set_cache( key=cache_key, @@ -1838,6 +1869,21 @@ async def get_key_object( _response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True)) + # Load object_permission if object_permission_id exists but object_permission is not loaded + if _response.object_permission_id and not _response.object_permission: + try: + _response.object_permission = await get_object_permission( + object_permission_id=_response.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Failed to load object_permission for key with object_permission_id={_response.object_permission_id}: {e}" + ) + # save the key object to cache await _cache_key_object( hashed_token=hashed_token, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 1ed21b07bb7..c2dbc94f721 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -331,11 +331,7 @@ class TestMCPRequestHandler: # Create an async mock for user_api_key_auth async def mock_user_api_key_auth(api_key, request): return UserAPIKeyAuth( - token=( - "test-token-sha256-empty-hash" - if api_key - else None - ), + token=("test-token-sha256-empty-hash" if api_key else None), api_key=api_key, user_id="test-user-id" if api_key else None, team_id="test-team-id" if api_key else None, @@ -632,8 +628,7 @@ class TestMCPOAuth2AuthFlow: # OAuth2 headers should still contain the Authorization token assert ( - oauth2_headers.get("Authorization") - == "Bearer atlassian-oauth2-token" + oauth2_headers.get("Authorization") == "Bearer atlassian-oauth2-token" ) async def test_litellm_key_in_authorization_backward_compat(self): @@ -1291,21 +1286,21 @@ async def test_get_team_object_permission_with_already_loaded_permission(): mcp_access_groups=["group1"], vector_stores=["store1"], ) - + # Create mock team object with object_permission already loaded mock_team_obj = LiteLLM_TeamTable( team_id="team-123", object_permission=mock_object_permission, object_permission_id="perm-123", ) - + # Create mock user auth mock_user_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", team_id="team-123", ) - + # Mock get_team_object to return our team with loaded permission # Also need to mock prisma_client from proxy_server mock_prisma = MagicMock() @@ -1313,96 +1308,81 @@ async def test_get_team_object_permission_with_already_loaded_permission(): "litellm.proxy.proxy_server.prisma_client", mock_prisma, ): - with patch( - "litellm.proxy.auth.auth_checks.get_team_object" - ) as mock_get_team: + with patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team: with patch( "litellm.proxy.auth.auth_checks.get_object_permission" ) as mock_get_perm: mock_get_team.return_value = mock_team_obj - + # Call the method result = await MCPRequestHandler._get_team_object_permission( mock_user_auth ) - + # Assert we got the object permission assert result == mock_object_permission assert result.mcp_servers == ["server1", "server2"] - + # Verify get_team_object was called mock_get_team.assert_called_once() - + # Verify get_object_permission was NOT called (since it was already loaded) mock_get_perm.assert_not_called() @pytest.mark.asyncio -async def test_get_team_object_permission_fetches_from_db_when_not_loaded(): +async def test_get_team_object_permission_with_core_auth_auto_loading(): """ - Test that _get_team_object_permission fetches from DB when object_permission - is not loaded but object_permission_id exists. + Test that _get_team_object_permission returns the object_permission that was + automatically loaded by get_team_object() in the core auth flow. + + Note: After migrating permission loading to core auth (get_team_object in auth_checks.py), + the team object returned by get_team_object() should already have object_permission loaded + when an object_permission_id exists. """ from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable - # Create mock object permission (to be returned from DB) + # Create mock object permission mock_object_permission = LiteLLM_ObjectPermissionTable( object_permission_id="perm-456", mcp_servers=["server3", "server4"], mcp_access_groups=["group2"], vector_stores=["store2"], ) - - # Create mock team object WITHOUT object_permission loaded (but has ID) + + # Create mock team object WITH object_permission already loaded + # (This is what get_team_object() returns after the core auth migration) mock_team_obj = LiteLLM_TeamTable( team_id="team-456", - object_permission=None, + object_permission=mock_object_permission, # Already loaded by core auth object_permission_id="perm-456", ) - + # Create mock user auth mock_user_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", team_id="team-456", ) - + # Mock the methods - # Also need to mock prisma_client from proxy_server mock_prisma = MagicMock() with patch( "litellm.proxy.proxy_server.prisma_client", mock_prisma, ): - with patch( - "litellm.proxy.auth.auth_checks.get_team_object" - ) as mock_get_team: - with patch( - "litellm.proxy.auth.auth_checks.get_object_permission" - ) as mock_get_perm: - mock_get_team.return_value = mock_team_obj - mock_get_perm.return_value = mock_object_permission - - # Call the method - result = await MCPRequestHandler._get_team_object_permission( - mock_user_auth - ) - - # Assert we got the object permission - assert result == mock_object_permission - assert result.mcp_servers == ["server3", "server4"] - - # Verify get_team_object was called - mock_get_team.assert_called_once() - - # Verify get_object_permission WAS called (since it wasn't loaded) - mock_get_perm.assert_called_once_with( - object_permission_id="perm-456", - prisma_client=mock.ANY, - user_api_key_cache=mock.ANY, - parent_otel_span=mock_user_auth.parent_otel_span, - proxy_logging_obj=mock.ANY, - ) + with patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team: + mock_get_team.return_value = mock_team_obj + + # Call the method + result = await MCPRequestHandler._get_team_object_permission(mock_user_auth) + + # Assert we got the object permission (already loaded by core auth) + assert result == mock_object_permission + assert result.mcp_servers == ["server3", "server4"] + + # Verify get_team_object was called + mock_get_team.assert_called_once() @pytest.mark.asyncio @@ -1420,14 +1400,14 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): mcp_access_groups=["dev-group"], vector_stores=[], ) - + # Create mock user auth mock_user_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", team_id="team-789", ) - + # Mock the helper methods with patch.object( MCPRequestHandler, "_get_team_object_permission" @@ -1437,13 +1417,16 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): ) as mock_get_access_group_servers: # Configure mocks mock_get_team_perm.return_value = mock_object_permission - mock_get_access_group_servers.return_value = ["group-server1", "group-server2"] - + mock_get_access_group_servers.return_value = [ + "group-server1", + "group-server2", + ] + # Call the method result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( mock_user_auth ) - + # Assert the result contains both direct and access group servers assert set(result) == { "direct-server1", @@ -1451,10 +1434,10 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): "group-server1", "group-server2", } - + # Verify _get_team_object_permission was called (the helper we fixed) mock_get_team_perm.assert_called_once_with(mock_user_auth) - + # Verify access groups were resolved mock_get_access_group_servers.assert_called_once_with(["dev-group"]) @@ -1471,21 +1454,21 @@ async def test_get_allowed_mcp_servers_for_team_with_no_object_permission(): user_id="test-user", team_id="team-no-perm", ) - + # Mock the helper to return None (no object permission) with patch.object( MCPRequestHandler, "_get_team_object_permission" ) as mock_get_team_perm: mock_get_team_perm.return_value = None - + # Call the method result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( mock_user_auth ) - + # Assert empty list is returned assert result == [] - + # Verify the helper was called mock_get_team_perm.assert_called_once_with(mock_user_auth) @@ -1509,9 +1492,7 @@ async def test_get_allowed_mcp_servers_for_team_without_team_id_returns_empty(): team_id=None, ) - result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( - mock_user_auth - ) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(mock_user_auth) assert result == [] @@ -1546,9 +1527,7 @@ async def test_get_allowed_mcp_servers_for_key_guard_conditions( "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, ) as mock_get_perm: - with patch( - "litellm.proxy.proxy_server.prisma_client", prisma_client_value - ): + with patch("litellm.proxy.proxy_server.prisma_client", prisma_client_value): result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( user_api_key_auth ) @@ -1569,9 +1548,7 @@ async def test_get_allowed_mcp_servers_for_key_returns_empty_when_db_returns_non mock_prisma = object() - with patch( - "litellm.proxy.proxy_server.prisma_client", mock_prisma - ), patch( + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, ) as mock_get_perm: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 9c77edc6743..4f93270c162 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -70,7 +70,9 @@ def _route_has_dependency(route, dependency) -> bool: dependant = getattr(route, "dependant", None) if dependant is None: return False - return any(getattr(dep, "call", None) == dependency for dep in dependant.dependencies) + return any( + getattr(dep, "call", None) == dependency for dep in dependant.dependencies + ) class TestExecuteWithMcpClient: @@ -481,7 +483,9 @@ class TestListToolsRestAPI: captured = {"called": False} - async def fake_get_tools(server, server_auth_header, raw_headers=None): + async def fake_get_tools( + server, server_auth_header, raw_headers=None, user_api_key_auth=None + ): captured["called"] = True captured["server"] = server captured["auth_header"] = server_auth_header @@ -659,3 +663,293 @@ class TestCallToolRestAPI: assert captured["name"] == "demo-tool" assert captured["arguments"] == {"foo": "bar"} assert captured["allowed_mcp_servers"] == [stub_server] + + +class TestGetToolsForSingleServer: + """Test _get_tools_for_single_server with object_permission filtering""" + + pytestmark = pytest.mark.asyncio + + async def test_filters_tools_by_object_permission_mcp_tool_permissions( + self, monkeypatch + ): + """Test that tools are filtered by user_api_key_auth.object_permission.mcp_tool_permissions""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + # Create mock tools + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + MockTool("tool3", "Third tool"), + ] + + # Mock _get_tools_from_server to return all tools + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + # Create server + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=None, # No server-level filtering + ) + + # Create UserAPIKeyAuth with object_permission + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-permission-id", + mcp_tool_permissions={"test-server-id": ["tool1", "tool3"]}, + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=object_permission, + ) + + # Call the function + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # Verify only allowed tools are returned + assert len(result) == 2 + tool_names = [tool.name for tool in result] + assert "tool1" in tool_names + assert "tool3" in tool_names + assert "tool2" not in tool_names + + async def test_no_filtering_when_object_permission_is_none(self, monkeypatch): + """Test that all tools are returned when object_permission is None""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + ] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=None, + ) + + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # All tools should be returned + assert len(result) == 2 + + async def test_no_filtering_when_mcp_tool_permissions_is_none(self, monkeypatch): + """Test that all tools are returned when mcp_tool_permissions is None""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + ] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=None, + ) + + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-permission-id", + mcp_tool_permissions=None, # No tool permissions set + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=object_permission, + ) + + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # All tools should be returned + assert len(result) == 2 + + async def test_no_filtering_when_server_not_in_mcp_tool_permissions( + self, monkeypatch + ): + """Test that all tools are returned when server is not in mcp_tool_permissions""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + ] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=None, + ) + + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-permission-id", + mcp_tool_permissions={"other-server-id": ["tool1"]}, # Different server + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=object_permission, + ) + + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # All tools should be returned since server is not in permissions + assert len(result) == 2 + + async def test_combines_server_allowed_tools_and_object_permission_filters( + self, monkeypatch + ): + """Test that both server.allowed_tools and object_permission.mcp_tool_permissions filters are applied""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + MockTool("tool3", "Third tool"), + MockTool("tool4", "Fourth tool"), + ] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + # Server allows tool1, tool2, tool3 + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=["tool1", "tool2", "tool3"], + ) + + # Object permission allows tool2, tool3, tool4 + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-permission-id", + mcp_tool_permissions={"test-server-id": ["tool2", "tool3", "tool4"]}, + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=object_permission, + ) + + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # Only tools in both lists should be returned (intersection): tool2, tool3 + assert len(result) == 2 + tool_names = [tool.name for tool in result] + assert "tool2" in tool_names + assert "tool3" in tool_names + assert "tool1" not in tool_names + assert "tool4" not in tool_names diff --git a/tests/test_litellm/proxy/auth/test_object_permission_loading.py b/tests/test_litellm/proxy/auth/test_object_permission_loading.py new file mode 100644 index 00000000000..54e4c82471e --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_object_permission_loading.py @@ -0,0 +1,151 @@ +""" +Test that object_permission is automatically loaded when fetching keys and teams. +""" +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.proxy._types import ( + LiteLLM_ObjectPermissionTable, + LiteLLM_TeamTableCachedObj, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_checks import get_key_object, get_team_object + + +@pytest.mark.asyncio +async def test_get_key_object_loads_object_permission(): + """ + Test that get_key_object automatically loads object_permission when object_permission_id exists. + """ + # Mock prisma client + mock_prisma_client = MagicMock() + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache + + # Mock the DB response with object_permission_id but no object_permission + mock_token_data = MagicMock() + mock_token_data.model_dump.return_value = { + "token": "test_token_hash", + "user_id": "test_user", + "object_permission_id": "test_perm_id", + "object_permission": None, + } + mock_prisma_client.get_data = AsyncMock(return_value=mock_token_data) + + # Mock the object_permission that should be loaded + mock_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test_perm_id", + mcp_servers=["server1", "server2"], + vector_stores=["store1"], + ) + + # Mock get_object_permission to return the permission + with patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + AsyncMock(return_value=mock_object_permission) + ), patch( + "litellm.proxy.auth.auth_checks._cache_key_object", + AsyncMock() + ): + result = await get_key_object( + hashed_token="test_token_hash", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + # Verify that object_permission was loaded + assert result.object_permission is not None + assert result.object_permission.object_permission_id == "test_perm_id" + assert result.object_permission.mcp_servers == ["server1", "server2"] + + +@pytest.mark.asyncio +async def test_get_key_object_no_permission_id(): + """ + Test that get_key_object works correctly when no object_permission_id exists. + """ + # Mock prisma client + mock_prisma_client = MagicMock() + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache + + # Mock the DB response without object_permission_id + mock_token_data = MagicMock() + mock_token_data.model_dump.return_value = { + "token": "test_token_hash", + "user_id": "test_user", + "object_permission_id": None, + "object_permission": None, + } + mock_prisma_client.get_data = AsyncMock(return_value=mock_token_data) + + with patch( + "litellm.proxy.auth.auth_checks._cache_key_object", + AsyncMock() + ): + result = await get_key_object( + hashed_token="test_token_hash", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + # Verify that object_permission is None + assert result.object_permission is None + + +@pytest.mark.asyncio +async def test_get_team_object_loads_object_permission(): + """ + Test that get_team_object automatically loads object_permission when object_permission_id exists. + """ + # Mock prisma client + mock_prisma_client = MagicMock() + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache + + # Mock team data with object_permission_id + mock_team = MagicMock() + mock_team.dict.return_value = { + "team_id": "test_team", + "team_alias": "Test Team", + "object_permission_id": "test_perm_id", + "object_permission": None, + } + + # Mock the object_permission that should be loaded + mock_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test_perm_id", + mcp_servers=["team_server1"], + vector_stores=["team_store1"], + ) + + with patch( + "litellm.proxy.auth.auth_checks._get_team_db_check", + AsyncMock(return_value=mock_team) + ), patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + AsyncMock(return_value=mock_object_permission) + ), patch( + "litellm.proxy.auth.auth_checks._cache_team_object", + AsyncMock() + ), patch( + "litellm.proxy.auth.auth_checks._should_check_db", + return_value=True + ), patch( + "litellm.proxy.auth.auth_checks._update_last_db_access_time" + ): + result = await get_team_object( + team_id="test_team", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + # Verify that object_permission was loaded + assert result.object_permission is not None + assert result.object_permission.object_permission_id == "test_perm_id" + assert result.object_permission.mcp_servers == ["team_server1"] diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index b4c8d5f07cc..e81c6264f7b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1,14 +1,15 @@ import os import sys import types -from types import SimpleNamespace from datetime import datetime, timedelta +from types import SimpleNamespace from typing import List, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient + from litellm._uuid import uuid from litellm.proxy.management_endpoints import ( mcp_management_endpoints as mgmt_endpoints, @@ -624,6 +625,72 @@ class TestListMCPServers: assert server.alias == "Allowed Zapier MCP" assert server.url == "https://actions.zapier.com/mcp/sse" + @pytest.mark.asyncio + async def test_admin_user_with_object_permission_respects_mcp_servers(self): + """ + Test that admin users with explicit object_permission.mcp_servers + only see the servers specified in object_permission. + + Scenario: Admin user has object_permission.mcp_servers set to specific servers + Expected: Only those servers are returned, not all servers in the registry + """ + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + # Create mock object permission with specific servers + mock_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-obj-perm-id", + mcp_servers=["server-1", "server-2"], # Only these two servers + mcp_access_groups=[], + mcp_tool_permissions={}, + vector_stores=[], + agents=[], + agent_access_groups=[], + ) + + # Create admin user with object permission + mock_user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin_user_id", + api_key="admin_api_key", + object_permission=mock_object_permission, + object_permission_id="test-obj-perm-id", + ) + + # Mock servers that the user should see + server_1 = generate_mock_mcp_server_db_record( + server_id="server-1", alias="Server 1", url="https://server1.example.com" + ) + server_2 = generate_mock_mcp_server_db_record( + server_id="server-2", alias="Server 2", url="https://server2.example.com" + ) + + # Mock manager + mock_manager = MagicMock() + mock_manager.get_all_allowed_mcp_servers = AsyncMock( + return_value=[server_1, server_2] + ) + + with patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts", + AsyncMock(return_value=[mock_user_auth]), + ): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_all_mcp_servers, + ) + + result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth) + + # Verify results - should only return the 2 servers in object_permission + assert len(result) == 2 + server_ids = {server.server_id for server in result} + assert server_ids == {"server-1", "server-2"} + + # Verify credentials are redacted + assert all(server.credentials is None for server in result) + @pytest.mark.asyncio async def test_fetch_single_mcp_server_redacts_credentials(self): mock_server = generate_mock_mcp_server_db_record( @@ -1251,7 +1318,9 @@ class TestMCPRegistryEndpoint: mock_manager = MagicMock() mock_manager.get_registry.return_value = {mock_server.server_id: mock_server} # The registry endpoint uses get_filtered_registry (filters by client IP) - mock_manager.get_filtered_registry.return_value = {mock_server.server_id: mock_server} + mock_manager.get_filtered_registry.return_value = { + mock_server.server_id: mock_server + } with ( patch_proxy_general_settings({"enable_mcp_registry": True}), diff --git a/ui/litellm-dashboard/src/components/OldTeams.tsx b/ui/litellm-dashboard/src/components/OldTeams.tsx index ecc6a624be0..7b906505759 100644 --- a/ui/litellm-dashboard/src/components/OldTeams.tsx +++ b/ui/litellm-dashboard/src/components/OldTeams.tsx @@ -52,7 +52,7 @@ import type { KeyResponse, Team } from "./key_team_helpers/key_list"; import MCPServerSelector from "./mcp_server_management/MCPServerSelector"; import MCPToolPermissions from "./mcp_server_management/MCPToolPermissions"; import NotificationsManager from "./molecules/notifications_manager"; -import { Organization, fetchMCPAccessGroups, getGuardrailsList, teamDeleteCall } from "./networking"; +import { Organization, fetchMCPAccessGroups, getGuardrailsList, getPoliciesList, teamDeleteCall } from "./networking"; import NumericalInput from "./shared/numerical_input"; import VectorStoreSelector from "./vector_store_management/VectorStoreSelector"; @@ -223,6 +223,7 @@ const Teams: React.FC = ({ const [isTeamDeleting, setIsTeamDeleting] = useState(false); // Add this state near the other useState declarations const [guardrailsList, setGuardrailsList] = useState([]); + const [policiesList, setPoliciesList] = useState([]); const [expandedAccordions, setExpandedAccordions] = useState>({}); const [loggingSettings, setLoggingSettings] = useState([]); const [mcpAccessGroups, setMcpAccessGroups] = useState([]); @@ -273,7 +274,22 @@ const Teams: React.FC = ({ } }; + const fetchPolicies = async () => { + try { + if (accessToken == null) { + return; + } + + const response = await getPoliciesList(accessToken); + const policyNames = response.policies.map((p: { policy_name: string }) => p.policy_name); + setPoliciesList(policyNames); + } catch (error) { + console.error("Failed to fetch policies:", error); + } + }; + fetchGuardrails(); + fetchPolicies(); }, [accessToken]); const fetchMcpAccessGroups = async () => { @@ -1330,6 +1346,36 @@ const Teams: React.FC = ({ } /> + + Policies{" "} + + e.stopPropagation()} + > + + + + + } + name="policies" + className="mt-8" + help="Select existing policies or enter new ones" + > + ({ + value: name, + label: name, + }))} + /> +