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 421f1dcfbea..f98bf8b22da 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 @@ -1425,13 +1425,7 @@ class MCPRequestHandler: key_object_permission.mcp_access_groups or [] ) - # servers referenced in tool permissions should also be accessible - tool_perm_servers = list( - global_mcp_server_manager.expand_tool_permissions(key_object_permission.mcp_tool_permissions).keys() - ) - - # Combine all lists - all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers + all_servers = direct_mcp_servers + access_group_servers return list(set(all_servers)) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers for key: {str(e)}") @@ -1502,13 +1496,7 @@ class MCPRequestHandler: object_permissions.mcp_access_groups or [] ) - tool_perm_servers = list( - global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys() - ) - - all_servers = ( - direct_mcp_servers + legacy_access_group_servers + tool_perm_servers + team_access_group_servers - ) + all_servers = direct_mcp_servers + legacy_access_group_servers + team_access_group_servers return list(set(all_servers)) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers for team: {str(e)}") @@ -1587,11 +1575,7 @@ class MCPRequestHandler: object_permissions.mcp_access_groups or [] ) - tool_perm_servers = list( - global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys() - ) - - all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers + all_servers = direct_mcp_servers + access_group_servers return list(set(all_servers)) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers for org: {str(e)}") @@ -1648,15 +1632,7 @@ class MCPRequestHandler: end_user_obj.object_permission.mcp_access_groups or [] ) - # servers referenced in tool permissions should also be accessible - tool_perm_servers = list( - global_mcp_server_manager.expand_tool_permissions( - end_user_obj.object_permission.mcp_tool_permissions - ).keys() - ) - - # Combine all lists - all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers + all_servers = direct_mcp_servers + access_group_servers return list(set(all_servers)) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers for end_user: {str(e)}") 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 6f132aaae9c..da5d04e23ea 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 @@ -3742,12 +3742,14 @@ class TestAgentMCPPermissions: @pytest.mark.asyncio -async def test_tool_permission_servers_included_in_allowed_servers(): +async def test_tool_permission_only_server_not_granted_server_access(): """ - Servers listed only in mcp_tool_permissions (not in mcp_servers) - should still be accessible. + mcp_servers is the single source of truth for which servers a key can + reach; mcp_tool_permissions only narrows which tools are callable on + already-granted servers. A server referenced only in mcp_tool_permissions + (absent from mcp_servers/mcp_access_groups) must NOT be reachable. - Regression test for https://github.com/BerriAI/litellm/issues/21954 + Regression test for https://github.com/BerriAI/litellm/issues/33397 """ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, @@ -3755,8 +3757,6 @@ async def test_tool_permission_servers_included_in_allowed_servers(): from litellm.types.mcp import MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer - # Register the server id so expand_permission_list resolves it rather than - # dropping it as stale. global_mcp_server_manager.registry["server_id_123"] = MCPServer( server_id="server_id_123", name="server_id_123", @@ -3787,11 +3787,65 @@ async def test_tool_permission_servers_included_in_allowed_servers(): result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( user_api_key_auth=user_api_key_auth, ) - assert "server_id_123" in result + assert "server_id_123" not in result + assert result == [] finally: global_mcp_server_manager.registry.pop("server_id_123", None) +@pytest.mark.asyncio +async def test_removing_server_revokes_access_despite_stale_tool_permissions(): + """ + Removing a server from mcp_servers must revoke it even when a stale + mcp_tool_permissions entry for that server still exists. Only the server + still in mcp_servers stays reachable. + + Regression test for https://github.com/BerriAI/litellm/issues/33397 + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + for server_id in ("kept_server", "removed_server"): + global_mcp_server_manager.registry[server_id] = MCPServer( + server_id=server_id, + name=server_id, + server_name=server_id, + url=f"https://{server_id}.example.com", + transport=MCPTransport.http, + ) + try: + perm = MagicMock() + perm.mcp_servers = ["kept_server"] + perm.mcp_access_groups = [] + # Stale entry left behind by the dashboard when the server was removed. + perm.mcp_tool_permissions = { + "kept_server": ["tool_a"], + "removed_server": ["tool_b"], + } + + user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + with ( + patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=perm), + patch.object( + MCPRequestHandler, + "_get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ), + ): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( + user_api_key_auth=user_api_key_auth, + ) + assert result == ["kept_server"] + finally: + for server_id in ("kept_server", "removed_server"): + global_mcp_server_manager.registry.pop(server_id, None) + + # --------------------------------------------------------------------------- # Org-level MCP permission tests # --------------------------------------------------------------------------- @@ -3963,6 +4017,7 @@ class TestOrgMCPPermissions: assert "group_server_1" in result async def test_get_allowed_mcp_servers_for_org_tool_permissions_only(self): + """A server referenced only in mcp_tool_permissions is not granted org access.""" auth = self._make_auth(org_id="org-123") mock_perm = MagicMock() @@ -3985,7 +4040,7 @@ class TestOrgMCPPermissions: ), ): result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth) - assert "tool_only_server" in result + assert result == [] async def test_get_allowed_mcp_servers_for_org_no_object_permission(self): auth = self._make_auth(org_id="org-123")