mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
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:
parent
e4c6badca2
commit
3c2138f037
2 changed files with 186 additions and 35 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue