mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
fix(proxy): show MCP server names and team alias in key MCP allowlist 403
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
9496f16f12
commit
e71c051ccf
2 changed files with 136 additions and 9 deletions
|
|
@ -300,6 +300,41 @@ async def _resolve_mcp_server_identifiers_to_ids(
|
|||
return resolved
|
||||
|
||||
|
||||
async def _mcp_server_display_names(
|
||||
server_ids: AbstractSet[str],
|
||||
prisma_client: PrismaClient | None,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Map MCP server IDs to human-readable names for error messages.
|
||||
|
||||
For each id, prefer alias, then server_name, then name, falling back to the
|
||||
raw id when the server is unknown or has no name. DB rows win over the
|
||||
in-memory registry when both know the server.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
registry_names: Final = MappingProxyType(
|
||||
{
|
||||
server_id: server.alias or server.server_name or server.name or server_id
|
||||
for registry_key, server in global_mcp_server_manager.get_registry().items()
|
||||
if (server_id := server.server_id or registry_key)
|
||||
}
|
||||
)
|
||||
db_names: Final = MappingProxyType(
|
||||
{
|
||||
server.server_id: server.alias or server.server_name or server.server_id
|
||||
for server in await _get_db_mcp_servers_by_identifiers(
|
||||
identifiers=server_ids,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
}
|
||||
)
|
||||
id_to_name: Final = MappingProxyType({**registry_names, **db_names})
|
||||
return sorted(id_to_name.get(server_id, server_id) for server_id in server_ids)
|
||||
|
||||
|
||||
_MCP_TOOL_PERMISSIONS_ADAPTER: Final = TypeAdapter(dict[str, list[str] | None])
|
||||
|
||||
|
||||
|
|
@ -694,19 +729,31 @@ async def validate_key_mcp_servers_against_team(
|
|||
)
|
||||
disallowed_servers: Final = active_requested_servers - allowed_servers - grandfathered_servers
|
||||
if disallowed_servers:
|
||||
disallowed_names: Final = await _mcp_server_display_names(
|
||||
server_ids=disallowed_servers,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
allow_all_names: Final = await _mcp_server_display_names(
|
||||
server_ids=allow_all_keys_servers,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if team_obj is not None:
|
||||
team_id = team_obj.team_id
|
||||
team_allowed_names: Final = await _mcp_server_display_names(
|
||||
server_ids=team_allowed_servers,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
team_display: Final = team_obj.team_alias or team_obj.team_id
|
||||
detail = (
|
||||
f"Key requests MCP servers not allowed by team '{team_id}': "
|
||||
f"{sorted(disallowed_servers)}. "
|
||||
f"Team allows: {sorted(team_allowed_servers)}. "
|
||||
f"Global (allow_all_keys) servers: {sorted(allow_all_keys_servers)}."
|
||||
f"Key requests MCP servers not allowed by team '{team_display}': "
|
||||
f"{disallowed_names}. "
|
||||
f"Team allows: {team_allowed_names}. "
|
||||
f"Global (allow_all_keys) servers: {allow_all_names}."
|
||||
)
|
||||
else:
|
||||
detail = (
|
||||
f"Key is not in a team. Only globally available (allow_all_keys) MCP servers "
|
||||
f"can be assigned: {sorted(allow_all_keys_servers)}. "
|
||||
f"Disallowed servers: {sorted(disallowed_servers)}."
|
||||
f"can be assigned: {allow_all_names}. "
|
||||
f"Disallowed servers: {disallowed_names}."
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
|
|
|
|||
|
|
@ -214,6 +214,7 @@ def test_extract_requested_mcp_access_groups_none():
|
|||
|
||||
def _make_team_obj(
|
||||
team_id="team-1",
|
||||
team_alias=None,
|
||||
mcp_servers=None,
|
||||
mcp_access_groups=None,
|
||||
mcp_tool_permissions=None,
|
||||
|
|
@ -221,6 +222,7 @@ def _make_team_obj(
|
|||
"""Create a mock team object with the given MCP permissions."""
|
||||
mock_team = MagicMock()
|
||||
mock_team.team_id = team_id
|
||||
mock_team.team_alias = team_alias
|
||||
|
||||
if (
|
||||
mcp_servers is not None
|
||||
|
|
@ -336,6 +338,84 @@ async def test_validate_key_servers_outside_team_scope_raises(
|
|||
assert "server-outside" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
new=_make_mock_mcp_manager(
|
||||
servers=[
|
||||
_make_mock_mcp_server("server-1", alias="github_mcp"),
|
||||
_make_mock_mcp_server("server-outside", alias="jira_mcp"),
|
||||
]
|
||||
),
|
||||
)
|
||||
@patch(
|
||||
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
|
||||
return_value=set(),
|
||||
)
|
||||
@patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[],
|
||||
)
|
||||
async def test_validate_key_servers_outside_team_scope_error_uses_names(
|
||||
mock_access_groups, mock_allow_all
|
||||
):
|
||||
"""The 403 detail should show server aliases and the team alias, not raw IDs."""
|
||||
team_obj = _make_team_obj(
|
||||
team_id="team-uuid",
|
||||
team_alias="mcp-test-team",
|
||||
mcp_servers=["server-1"],
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await validate_key_mcp_servers_against_team(
|
||||
object_permission={"mcp_servers": ["server-1", "server-outside"]},
|
||||
team_obj=team_obj,
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
detail = str(exc_info.value.detail)
|
||||
assert "mcp-test-team" in detail
|
||||
assert "jira_mcp" in detail
|
||||
assert "github_mcp" in detail
|
||||
assert "team-uuid" not in detail
|
||||
assert "server-outside" not in detail
|
||||
assert "server-1" not in detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
new=_make_mock_mcp_manager(
|
||||
servers=[
|
||||
_make_mock_mcp_server("server-1", alias="github_mcp"),
|
||||
_make_mock_mcp_server("server-outside", alias="jira_mcp"),
|
||||
]
|
||||
),
|
||||
)
|
||||
@patch(
|
||||
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
|
||||
return_value=set(),
|
||||
)
|
||||
@patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[],
|
||||
)
|
||||
async def test_validate_key_servers_no_team_error_uses_names(
|
||||
mock_access_groups, mock_allow_all
|
||||
):
|
||||
"""The teamless 403 detail should show server aliases, not raw IDs."""
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await validate_key_mcp_servers_against_team(
|
||||
object_permission={"mcp_servers": ["server-outside"]},
|
||||
team_obj=None,
|
||||
is_proxy_admin=False,
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
detail = str(exc_info.value.detail)
|
||||
assert "jira_mcp" in detail
|
||||
assert "server-outside" not in detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
|
|
@ -690,7 +770,7 @@ async def test_validate_mcp_server_alias_outside_team_scope_raises(
|
|||
team_obj=team_obj,
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "private-server-id" in str(exc_info.value.detail)
|
||||
assert "private-alias" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -788,7 +868,7 @@ async def test_validate_db_mcp_server_alias_outside_team_scope_raises_when_regis
|
|||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "private-server-id" in str(exc_info.value.detail)
|
||||
assert "private-alias" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue