This commit is contained in:
Fernando 2026-08-27 20:11:27 -05:00 committed by GitHub
commit 064ff52302
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 62 additions and 11 deletions

View file

@ -2734,7 +2734,10 @@ class MCPServerManager:
base_url=server.url or "",
)
if initialize_mapping:
self.initialize_tool_name_to_mcp_server_name_mapping()
# OpenAPI tools are served from the local tool registry, so only this
# server needs re-mapping. Sweeping the whole registry would call
# list_tools upstream on every other (remote) MCP server.
self.initialize_tool_name_to_mcp_server_name_mapping(servers=(server,))
async def add_server(self, mcp_server: LiteLLM_MCPServerTable):
# The runtime registry is the allowlist for tool calls and health
@ -5821,24 +5824,28 @@ class MCPServerManager:
# End of Methods that call the upstream MCP servers
#########################################################
def initialize_tool_name_to_mcp_server_name_mapping(self):
def initialize_tool_name_to_mcp_server_name_mapping(self, servers: Sequence[MCPServer] | None = None):
"""
On startup, initialize the tool name to MCP server name mapping
On startup, initialize the tool name to MCP server name mapping.
``servers`` restricts the sweep to those servers. Callers that only registered
OpenAPI tools pass the affected servers so a re-registration does not re-list
tools upstream from every remote server in the registry.
"""
try:
if asyncio.get_running_loop():
asyncio.create_task(self._initialize_tool_name_to_mcp_server_name_mapping())
asyncio.create_task(self._initialize_tool_name_to_mcp_server_name_mapping(servers))
except RuntimeError as e: # no running event loop
verbose_logger.exception(
"No running event loop - skipping tool name to MCP server name mapping initialization: %s", e
)
async def _initialize_tool_name_to_mcp_server_name_mapping(self):
async def _initialize_tool_name_to_mcp_server_name_mapping(self, servers: Sequence[MCPServer] | None = None):
"""
Call list_tools for each server and update the tool name to MCP server name mapping
Note: This now handles prefixed tool names
"""
for server in self.get_registry().values():
for server in servers if servers is not None else tuple(self.get_registry().values()):
if self._oauth_discovery_slot(server.server_id) is not None:
continue
if server.needs_user_oauth_token:
@ -5988,7 +5995,6 @@ class MCPServerManager:
# Assign short prefixes against the full candidate set without
# publishing the staged registry to concurrent callers.
registered_registry: Final[dict[str, MCPServer]] = {}
registered_openapi_tools = False
for server_id, new_server in new_registry.items():
try:
self._assign_unique_short_prefix(new_server, registry=new_registry)
@ -5997,8 +6003,6 @@ class MCPServerManager:
# prefix that lookups will use.
await self._maybe_register_openapi_tools(new_server, initialize_mapping=False)
registered_registry[server_id] = new_server
if new_server.spec_path:
registered_openapi_tools = True
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
@ -6019,8 +6023,13 @@ class MCPServerManager:
registered_servers: Final = tuple(registered_registry.values())
self._reconcile_oauth_discovery_slots_for_servers(registered_servers)
self._prime_oauth_metadata_discovery_for_servers(registered_servers)
if registered_openapi_tools:
self.initialize_tool_name_to_mcp_server_name_mapping()
openapi_servers: Final = tuple(server for server in registered_servers if server.spec_path)
if openapi_servers:
# This reload runs on a timer (``proxy_config_reload_interval_seconds``,
# 30s by default). Map only the OpenAPI servers, whose tools come from the
# local registry: a full sweep would list tools upstream from every remote
# MCP server on every reload.
self.initialize_tool_name_to_mcp_server_name_mapping(servers=openapi_servers)
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))

View file

@ -926,6 +926,48 @@ class TestMCPServerManager:
get_tools.assert_not_awaited()
@pytest.mark.asyncio
async def test_database_reload_maps_only_openapi_servers(self):
"""The DB reload runs on a timer, so the tool mapping it triggers for freshly
registered OpenAPI servers must not re-list tools upstream from every remote
server in the registry."""
manager = MCPServerManager()
openapi_row = LiteLLM_MCPServerTable(
server_id="openapi-1",
server_name="openapi_server",
alias="openapi_server",
url="https://openapi.example.com",
transport=MCPTransport.http,
spec_path="https://openapi.example.com/openapi.json",
)
remote_row = LiteLLM_MCPServerTable(
server_id="remote-1",
server_name="remote_server",
alias="remote_server",
url="https://remote.example.com/mcp",
transport=MCPTransport.http,
)
repository = MagicMock()
repository.table.find_many = AsyncMock(return_value=[openapi_row, remote_row])
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
return_value=repository,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch.object(manager, "_register_openapi_tools", new=AsyncMock()),
patch.object(manager, "initialize_tool_name_to_mcp_server_name_mapping") as init_mapping,
):
await manager.reload_servers_from_database()
init_mapping.assert_called_once()
mapped = init_mapping.call_args.kwargs["servers"]
assert [server.server_id for server in mapped] == ["openapi-1"]
@pytest.mark.asyncio
async def test_load_servers_from_config_requires_oauth2_flow(self):
"""auth_type oauth2 without an explicit oauth2_flow is a config error: the