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 425f82794e6..ce539258fd2 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 @@ -3074,15 +3074,17 @@ class MCPRequestHandler: @staticmethod async def _get_allowed_mcp_servers_for_agent( user_api_key_auth: UserAPIKeyAuth | None = None, - agent_object_permission=None, + agent_object_permission: LiteLLM_ObjectPermissionTable | None = None, ) -> list[str]: """ Get allowed MCP servers for an agent (from the agent's object_permission). - Returns the MCP servers from the agent's object_permission. - If agent has no object_permission, returns [] (no extra restriction). An entitlement the - agent LINKS but that cannot be read raises ``UnloadableEntitlementError`` out of here so the - resolver denies. + Returns the agent's direct servers, the servers in its access groups, and the servers reached + through its toolsets, exactly as the key, team, and org levels count theirs. If agent has no + object_permission, returns [] (no extra restriction). An entitlement the agent LINKS but that + cannot be read, or a declared toolset that resolves to no grants, raises + ``UnloadableEntitlementError`` out of here so the resolver denies instead of reading the + agent as unrestricted. Args: user_api_key_auth: User auth with agent_id @@ -3092,31 +3094,30 @@ class MCPRequestHandler: if not user_api_key_auth or not user_api_key_auth.agent_id: return [] - obj_perm = agent_object_permission - if obj_perm is None: - obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + obj_perm: Final = ( + agent_object_permission + if agent_object_permission is not None + else await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + ) if obj_perm is None: return [] try: - direct_mcp_servers = getattr(obj_perm, "mcp_servers", None) or [] - if isinstance(direct_mcp_servers, str): - direct_mcp_servers = [] - mcp_access_groups = getattr(obj_perm, "mcp_access_groups", None) or [] - if isinstance(mcp_access_groups, str): - mcp_access_groups = [] - - # Permission entries may be server_ids OR names/aliases — expand to ids. from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) - expanded_direct_servers: Final = global_mcp_server_manager.expand_permission_list(list(direct_mcp_servers)) - - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(mcp_access_groups) - all_servers: Final = expanded_direct_servers + access_group_servers - return list(set(all_servers)) + expanded_direct_servers: Final = global_mcp_server_manager.expand_permission_list( + obj_perm.mcp_servers or [] + ) + access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + obj_perm.mcp_access_groups or [] + ) + toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(obj_perm) + return list({*expanded_direct_servers, *access_group_servers, *toolset_grants}) except Exception as e: + if isinstance(e, UnloadableEntitlementError): + raise verbose_logger.warning("Failed to get allowed MCP servers for agent: %s", e) return [] @@ -3124,13 +3125,15 @@ class MCPRequestHandler: async def _get_agent_tool_permissions_for_server( server_id: str, user_api_key_auth: UserAPIKeyAuth | None = None, - agent_object_permission=None, + agent_object_permission: LiteLLM_ObjectPermissionTable | None = None, ) -> list[str] | None: """ - Get allowed tool names for a server from the agent's object_permission. - Returns None if agent has no tool restrictions for this server. An entitlement the agent - LINKS but that cannot be read raises ``UnloadableEntitlementError`` out of here, which the - tool resolver turns into deny-all for the server rather than an unrestricted tool list. + Get allowed tool names for a server from the agent's object_permission: the union of its + direct tool permissions and the tools its toolsets grant on that server, mirroring the key and + team levels. Returns None if agent has no tool restrictions for this server. An entitlement the + agent LINKS but that cannot be read, or a declared toolset that resolves to no grants, raises + ``UnloadableEntitlementError`` out of here, which the tool resolver turns into deny-all for the + server rather than an unrestricted tool list. Args: server_id: Server ID to check permissions for @@ -3141,24 +3144,30 @@ class MCPRequestHandler: if not user_api_key_auth or not user_api_key_auth.agent_id: return None - obj_perm = agent_object_permission - if obj_perm is None: - obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + obj_perm: Final = ( + agent_object_permission + if agent_object_permission is not None + else await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + ) if obj_perm is None: return None try: - mcp_tool_permissions: Final = getattr(obj_perm, "mcp_tool_permissions", None) - if not mcp_tool_permissions or not isinstance(mcp_tool_permissions, dict): - return None - # Dict keys may be server_ids OR names/aliases; normalize before lookup. from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) - tools: Final = global_mcp_server_manager.expand_tool_permissions(mcp_tool_permissions).get(server_id) - return list(tools) if tools else None + direct_tools: Final = ( + global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions).get(server_id) + if obj_perm.mcp_tool_permissions + else None + ) + toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id) + agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools) + return list(agent_tools) if agent_tools else None except Exception as e: + if isinstance(e, UnloadableEntitlementError): + raise verbose_logger.warning("Failed to get agent tool permissions for server: %s", e) return None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 0144fbb17dd..7caf31ef6f1 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 @@ -11,6 +11,7 @@ from starlette.datastructures import Headers from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, + UnloadableEntitlementError, _is_mcp_admitted_user_subject, ) from litellm.proxy._types import ( @@ -4358,6 +4359,147 @@ class TestAgentMCPPermissions: ) assert sorted(result) == ["tool_a", "tool_b"] + def _agent_object_permission(self, *, toolset_ids, servers=(), tool_permissions=None): + agent_object_permission = MagicMock() + agent_object_permission.mcp_servers = list(servers) + agent_object_permission.mcp_access_groups = [] + agent_object_permission.mcp_tool_permissions = tool_permissions + agent_object_permission.mcp_toolsets = list(toolset_ids) + return agent_object_permission + + def _mock_manager_with_toolsets(self, toolset_perms): + mock_manager = MagicMock() + mock_manager.expand_permission_list = MagicMock(side_effect=lambda servers: list(servers)) + mock_manager.expand_tool_permissions = MagicMock(side_effect=lambda perms: perms or {}) + mock_manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms) + return mock_manager + + def _agent_toolset_patches(self, agent_object_permission, mock_manager): + return ( + patch.object( # test-quality-ok: stub the agent perm loader; the resolver reads module globals with no injection seam + MCPRequestHandler, "_get_agent_object_permission", AsyncMock(return_value=agent_object_permission) + ), + patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling toolset tests + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + mock_manager, + ), + patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here + MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + ), + ) + + async def test_get_allowed_mcp_servers_for_agent_includes_toolset_servers(self): + """An agent granted only mcp_toolsets reaches the toolset's servers, exactly as a + key, team, or org granted only toolsets does""" + user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets") + agent_object_permission = self._agent_object_permission(toolset_ids=["toolset-1"], servers=["server-direct"]) + mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]}) + + with contextlib.ExitStack() as stack: + for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): + stack.enter_context(patcher) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth) + + assert sorted(result) == ["server-a", "server-direct"] + mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"]) + + async def test_get_allowed_mcp_servers_toolset_only_agent_caps_key_servers(self): + """Regression: an agent whose only grant is a toolset used to resolve to [] and place + no ceiling at all, so a key bound to it kept every server the key itself granted""" + user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets") + agent_object_permission = self._agent_object_permission(toolset_ids=["toolset-1"]) + mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]}) + + with contextlib.ExitStack() as stack: + for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): + stack.enter_context(patcher) + stack.enter_context( + patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here + MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"]) + ) + ) + stack.enter_context( + patch.object( # test-quality-ok: team resolution has its own tests; pin it empty here + MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[]) + ) + ) + result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) + + assert result == ["server-a"] + + async def test_get_allowed_mcp_servers_agent_dangling_toolset_denies(self): + """An agent toolset that resolves to nothing is a known restriction with unknown + contents: deny, never fall through to the key's own servers""" + user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets") + agent_object_permission = self._agent_object_permission(toolset_ids=["toolset-gone"]) + mock_manager = self._mock_manager_with_toolsets({}) + + with contextlib.ExitStack() as stack: + for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): + stack.enter_context(patcher) + with pytest.raises(UnloadableEntitlementError): + await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth) + stack.enter_context( + patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here + MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"]) + ) + ) + stack.enter_context( + patch.object( # test-quality-ok: team resolution has its own tests; pin it empty here + MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[]) + ) + ) + result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) + + assert result == [] + + async def test_get_agent_tool_permissions_for_server_unions_direct_and_toolset_tools(self): + """The agent's tool ceiling on a server is its direct tool grants plus the tools its + toolsets grant there, and None only when neither names the server""" + user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets") + agent_object_permission = self._agent_object_permission( + toolset_ids=["toolset-1"], tool_permissions={"server-a": ["tool_direct"]} + ) + mock_manager = self._mock_manager_with_toolsets({"server-a": ["tool_via_toolset"], "server-b": ["tool_b"]}) + + with contextlib.ExitStack() as stack: + for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): + stack.enter_context(patcher) + server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-a", user_api_key_auth) + server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-b", user_api_key_auth) + server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-c", user_api_key_auth) + + assert sorted(server_a_tools) == ["tool_direct", "tool_via_toolset"] + assert server_b_tools == ["tool_b"] + assert server_c_tools is None + + async def test_get_allowed_tools_for_server_toolset_only_agent_caps_key_tools(self): + """Regression: a key allowing [tool_a, tool_b] bound to an agent whose toolset grants + only tool_a on the server ends with [tool_a]; the toolset used to be ignored""" + user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets") + agent_object_permission = self._agent_object_permission(toolset_ids=["toolset-1"]) + mock_manager = self._mock_manager_with_toolsets({"server-a": ["tool_a"]}) + key_perm = MagicMock() + key_perm.mcp_tool_permissions = {"server-a": ["tool_a", "tool_b"]} + key_perm.mcp_toolsets = [] + + with contextlib.ExitStack() as stack: + for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): + stack.enter_context(patcher) + stack.enter_context( + patch.object( # test-quality-ok: stub the key perm loader; the resolver reads module globals with no injection seam + MCPRequestHandler, "_get_key_object_permission", return_value=key_perm + ) + ) + stack.enter_context( + patch.object( # test-quality-ok: team resolution has its own tests; pin it absent here + MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=None) + ) + ) + result = await MCPRequestHandler.get_allowed_tools_for_server("server-a", user_api_key_auth) + + assert result == ["tool_a"] + async def test_get_agent_object_permission_uses_shared_helper(self): """``_get_agent_object_permission`` must resolve the agent's ``object_permission_id`` and then defer to the shared