fix(mcp): count agent toolsets in the agent MCP ceiling and tool grants

An agent's object_permission.mcp_toolsets could be persisted through the
new edit form and PATCH /v1/agents but never reached the request-time
checks: _get_allowed_mcp_servers_for_agent read only mcp_servers and
mcp_access_groups, so an agent granted a toolset alone resolved to [] and
placed no ceiling on its keys, and _get_agent_tool_permissions_for_server
ignored the tools those toolsets grant. Both helpers now resolve toolsets
through the shared _toolset_tool_permissions / _toolset_tools_for_server
helpers the key, team, and org levels already use, and a declared toolset
that resolves to nothing raises UnloadableEntitlementError so the
resolver denies instead of reading the agent as unrestricted
This commit is contained in:
mateo-berri 2026-09-01 19:12:24 -07:00
parent e4c6badca2
commit 3c2138f037
2 changed files with 186 additions and 35 deletions

View file

@ -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

View file

@ -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