mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(mcp): preserve managed agent scope through server resolution
This commit is contained in:
parent
057d751da6
commit
71594c29b2
3 changed files with 46 additions and 2 deletions
|
|
@ -3434,6 +3434,10 @@ class MCPServerManager:
|
|||
2. If admin and no object_permission, return all servers
|
||||
3. Otherwise, use standard permission checks
|
||||
"""
|
||||
if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
|
||||
managed: Final = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
|
||||
return managed if access is None else [server for server in managed if server in access.server_ids]
|
||||
|
||||
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
|
||||
|
||||
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
|
||||
|
|
|
|||
|
|
@ -259,3 +259,43 @@ async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatc
|
|||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("role", (None, "proxy_admin", "internal_user"))
|
||||
@pytest.mark.parametrize("scoped", (False, True))
|
||||
async def test_manager_preserves_managed_server_grants_across_open_channels(
|
||||
monkeypatch: pytest.MonkeyPatch, role: str | None, scoped: bool
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import db
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {
|
||||
"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True),
|
||||
"submitted": MCPServer(server_id="submitted", name="submitted", transport="http"),
|
||||
"passthrough": MCPServer(
|
||||
server_id="passthrough", name="passthrough", transport="http", auth_type="true_passthrough"
|
||||
),
|
||||
}
|
||||
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["submitted"]))
|
||||
auth: Final = actor(None)
|
||||
auth.user_role = role
|
||||
assert not auth.mcp_explicit_grants_only
|
||||
access: Final = MCPServerAccess(server_ids=("slack", "open")) if scoped else None
|
||||
assert set(await manager.get_allowed_mcp_servers(auth, access=access)) == ({"slack"} if scoped else {"slack", "linear"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_manager_does_not_replace_managed_policy_failure_with_open_servers(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True)}
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("writer unavailable")))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await manager.get_allowed_mcp_servers(actor(None, delegated=True))
|
||||
assert failure.value.status_code == 503
|
||||
|
|
|
|||
|
|
@ -4708,7 +4708,7 @@ class TestAgentMCPPermissions:
|
|||
),
|
||||
)
|
||||
|
||||
async def testget_allowed_mcp_servers_for_agent_includes_toolset_servers(self):
|
||||
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")
|
||||
|
|
@ -4777,7 +4777,7 @@ class TestAgentMCPPermissions:
|
|||
|
||||
assert result == []
|
||||
|
||||
async def testget_agent_tool_permissions_for_server_unions_direct_and_toolset_tools(self):
|
||||
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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue