mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
Merge ec5f945390 into bb72815e70
This commit is contained in:
commit
064ff52302
2 changed files with 62 additions and 11 deletions
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue