From 0c79d08cba33b40dbc2775425b1825128bd7847f Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 10 Mar 2026 17:51:06 -0700 Subject: [PATCH] fix: team-scoped MCP server listing now queries DB for servers not in registry `_get_team_scoped_mcp_servers` only checked the in-memory registry via `get_all_mcp_servers_unfiltered()`, missing DB-only servers. Now checks the registry first via `get_mcp_server_by_id()`, then falls back to a direct DB query via `get_mcp_servers()` for any missing IDs. Co-Authored-By: Claude Opus 4.6 --- .../mcp_management_endpoints.py | 21 ++++++++++++++++--- .../test_mcp_management_endpoints.py | 21 ++++++++++++++----- 2 files changed, 34 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 2b4cb6a4434..716aa59fa16 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -324,6 +324,7 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth, ) -> List[LiteLLM_MCPServerTable]: """Return MCP servers allowed by the team + allow_all_keys servers.""" + from litellm.proxy._experimental.mcp_server.db import get_mcp_servers from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.management_helpers.object_permission_utils import ( get_team_mcp_permissions, @@ -346,9 +347,23 @@ if MCP_AVAILABLE: # Team has no MCP config - only allow_all_keys servers allowed_ids = allow_all_ids - all_servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered() - filtered = [s for s in all_servers if s.server_id in allowed_ids] - return _redact_mcp_credentials_list(filtered) + # Collect servers from both in-memory registry and DB + servers_by_id: Dict[str, LiteLLM_MCPServerTable] = {} + + # 1. Check in-memory registry (config + previously loaded DB servers) + for server_id in allowed_ids: + server = global_mcp_server_manager.get_mcp_server_by_id(server_id) + if server: + servers_by_id[server_id] = global_mcp_server_manager._build_mcp_server_table(server) + + # 2. For any IDs not found in registry, query DB directly + missing_ids = allowed_ids - set(servers_by_id.keys()) + if missing_ids and prisma_client is not None: + db_servers = await get_mcp_servers(prisma_client, missing_ids) + for s in db_servers: + servers_by_id[s.server_id] = s + + return _redact_mcp_credentials_list(servers_by_id.values()) async def _get_team_scoped_access_groups(team_id: str) -> dict: """Return access groups available to the specified team.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 10e2fdb05d0..d94d65d7a7a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1534,14 +1534,26 @@ class TestTeamIdParam: ) server_1 = generate_mock_mcp_server_db_record(server_id="server-1") - server_2 = generate_mock_mcp_server_db_record(server_id="server-2") server_public = generate_mock_mcp_server_db_record(server_id="server-public") + # Build config-style MCPServer mocks for the registry + config_server_1 = generate_mock_mcp_server_config_record(server_id="server-1") + config_server_public = generate_mock_mcp_server_config_record(server_id="server-public") + mock_manager = MagicMock() mock_manager.get_allow_all_keys_server_ids.return_value = ["server-public"] - mock_manager.get_all_mcp_servers_unfiltered = AsyncMock( - return_value=[server_1, server_2, server_public] - ) + + def mock_get_by_id(sid): + lookup = {"server-1": config_server_1, "server-public": config_server_public} + return lookup.get(sid) + + mock_manager.get_mcp_server_by_id.side_effect = mock_get_by_id + + def mock_build_table(server): + lookup = {"server-1": server_1, "server-public": server_public} + return lookup.get(server.server_id, server_1) + + mock_manager._build_mcp_server_table.side_effect = mock_build_table with ( patch( @@ -1581,7 +1593,6 @@ class TestTeamIdParam: server_ids = {s.server_id for s in result} assert "server-1" in server_ids assert "server-public" in server_ids - assert "server-2" not in server_ids @pytest.mark.asyncio async def test_fetch_mcp_servers_team_id_non_member_rejected(self):