fix(mcp): retain valid live routes and handlers during refresh

This commit is contained in:
Joshua Valluru 2026-09-22 14:11:07 -07:00
parent 42821c102e
commit d4c1afdebc
2 changed files with 143 additions and 9 deletions

View file

@ -209,8 +209,9 @@ class TargetCatalog:
previous_config: Final = MappingProxyType(
{key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items()}
)
live_registry: Final = self.manager.registry
previous_servers: Final = previous_config | MappingProxyType(
{key: value.model_copy(deep=True) for key, value in self.manager.registry.items()}
{key: value.model_copy(deep=True) for key, value in live_registry.items()}
)
config_identities: Final = MappingProxyType(
{key: _configuration_identity(value) for key, value in previous_config.items()}
@ -223,17 +224,21 @@ class TargetCatalog:
staged_routing: Final = dict(initial_routing)
closed: Final = asyncio.Event()
routing_token: Final = self._staged_routing.set((staged_routing, closed))
initial_tools: Final = MappingProxyType(dict(global_mcp_tool_registry.published_tools))
try:
with global_mcp_tool_registry.catalog_scope(global_mcp_tool_registry.published_tools) as staged_tools:
with global_mcp_tool_registry.catalog_scope(initial_tools) as staged_tools:
await self._reload()
live_routes: Final = self._unchanged_routing(previous_servers, self.manager.published_tool_routes)
concurrent_routes: Final = self._unchanged_routing(
previous_servers,
self.manager.config_mcp_servers | live_registry,
MappingProxyType(
{
name: owner
for name, owner in self.manager.published_tool_routes.items()
if initial_routing.get(name) != owner
and staged_routing.get(name) == initial_routing.get(name)
and (
name not in staged_routing or staged_routing.get(name) == initial_routing.get(name)
)
}
),
)
@ -248,6 +253,18 @@ class TargetCatalog:
}
),
)
concurrent_tools: Final = MappingProxyType(
{
name: tool
for name, tool in global_mcp_tool_registry.published_tools.items()
if (
name in live_routes
or name in concurrent_routes
or (name not in initial_routing and name not in self.manager.published_tool_routes)
)
and initial_tools.get(name) is staged_tools.get(name)
}
)
self.manager.config_mcp_servers = {
key: value.model_copy(
update=staged_config[key].model_dump(
@ -258,8 +275,18 @@ class TargetCatalog:
else value
for key, value in self.manager.config_mcp_servers.items()
}
global_mcp_tool_registry.tools = staged_tools
self.manager.published_tool_routes = retained_staged_routes | concurrent_routes
global_mcp_tool_registry.tools = (
MappingProxyType(
{
name: tool
for name, tool in staged_tools.items()
if name in global_mcp_tool_registry.published_tools
or tool is not initial_tools.get(name)
}
)
| concurrent_tools
)
self.manager.published_tool_routes = live_routes | retained_staged_routes | concurrent_routes
finally:
closed.set()
self._staged_routing.reset(routing_token)

View file

@ -14607,34 +14607,63 @@ async def test_catalog_operation_retains_routes_only_for_same_configured_target(
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["concurrent", "staged", "deleted", "staged_after_delete", "staged_with_concurrent"])
async def test_catalog_reload_keeps_new_route_owner_over_earlier_route(change):
@pytest.mark.parametrize("mapped", [False, True])
async def test_catalog_reload_keeps_new_route_owner_over_earlier_route(change, mapped, monkeypatch):
from litellm.proxy._experimental.mcp_server import tool_registry
manager: Final = MCPServerManager()
registry: Final = tool_registry.MCPToolRegistry()
monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry)
servers: Final = {name: MCPServer(server_id=name, name=name, transport=MCPTransport.http,
client_id="configured-client") for name in ("old_owner", "new_owner", "concurrent_owner")}
manager.config_mcp_servers = servers
manager.published_tool_routes = {"shared-search": "old_owner"}
manager.published_tool_routes = {"shared-search": "old_owner"} if mapped else {}
async def original_handler():
return "original response"
async def new_handler():
return "new response"
async def concurrent_handler():
return "concurrent response"
registry.register_tool("shared-search", "Search", {}, original_handler)
original_tool: Final = registry.get_tool("shared-search")
async def read_rows(**_kwargs):
if change == "concurrent":
manager.published_tool_routes["shared-search"] = "new_owner"
registry.published_tools["shared-search"] = original_tool.model_copy(update={"handler": new_handler})
elif change in ("staged", "staged_after_delete", "staged_with_concurrent"):
if change == "staged_after_delete":
manager.published_tool_routes.clear()
registry.published_tools.clear()
elif change == "staged_with_concurrent":
manager.published_tool_routes["shared-search"] = "concurrent_owner"
registry.published_tools["shared-search"] = original_tool.model_copy(update={"handler": concurrent_handler})
manager.tool_name_to_mcp_server_name_mapping["shared-search"] = "new_owner"
registry.register_tool("shared-search", "Search", {}, new_handler)
else:
manager.published_tool_routes.clear()
registry.published_tools.clear()
if not mapped:
manager.published_tool_routes.clear()
manager.tool_name_to_mcp_server_name_mapping.clear()
return []
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await manager.reload_servers_from_database()
if change == "deleted":
if change == "deleted" or not mapped:
assert "shared-search" not in manager.published_tool_routes
else:
assert manager.published_tool_routes["shared-search"] == "new_owner"
if change == "deleted":
assert registry.get_tool("shared-search") is None
else:
assert await registry.get_tool("shared-search").handler() == "new response"
@pytest.mark.asyncio
@ -14666,6 +14695,84 @@ async def test_catalog_observes_committed_update_and_delete_without_background_r
assert prisma.db.litellm_mcpservertable.find_many.await_count == 3
@pytest.mark.asyncio
async def test_catalog_rebuilt_unchanged_server_keeps_discovered_tool_routes():
from mcp.types import Tool
row: Final = LiteLLM_MCPServerTable(server_id="unchanged-routes", alias="unchanged_routes",
url="https://upstream.example/mcp", transport=MCPTransport.http)
manager: Final = MCPServerManager()
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await manager.reload_servers_from_database()
server: Final = manager.get_mcp_server_by_id(row.server_id)
tools: Final = manager._create_prefixed_tools([Tool(name="search", inputSchema={})], server)
await manager.reload_servers_from_database()
assert manager.server_exposes_tool(manager.get_mcp_server_by_id(row.server_id), tools[0].name)
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["same", "changed", "deleted"])
@pytest.mark.parametrize("openapi", [False, True])
async def test_catalog_reload_retains_routes_discovered_for_a_late_server(change, openapi, monkeypatch):
from datetime import timedelta
from mcp.types import Tool
from litellm.proxy._experimental.mcp_server import tool_registry
row: Final = LiteLLM_MCPServerTable(server_id="late-routes", alias="late_routes",
url="https://upstream.example/mcp", transport=MCPTransport.http, updated_at=datetime.now(),
spec_path="late-openapi.json" if openapi else None)
manager: Final = MCPServerManager()
registry: Final = tool_registry.MCPToolRegistry()
monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry)
server: Final = await manager.build_mcp_server_from_table(row)
start: Final = asyncio.Event()
published: Final = asyncio.Event()
async def handler():
return "late server response"
async def publish():
await start.wait()
manager.registry[server.server_id] = server
manager._create_prefixed_tools([Tool(name="search", inputSchema={})], server)
if openapi:
registry.register_tool("late_routes-search", "Search", {}, handler)
published.set()
async def read_rows(**_kwargs):
start.set()
await published.wait()
if change == "deleted":
return []
return [row if change == "same" else row.model_copy(update={
"url": "https://updated.example/mcp", "spec_path": None,
"updated_at": row.updated_at + timedelta(seconds=1)})]
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
task: Final = asyncio.create_task(publish())
try:
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await manager.reload_servers_from_database()
finally:
task.cancel()
await asyncio.gather(task, return_exceptions=True)
if change == "deleted":
assert manager.get_mcp_server_by_id(row.server_id) is None
assert "late_routes-search" not in manager.published_tool_routes
else:
assert manager.server_exposes_tool(manager.get_mcp_server_by_id(row.server_id), "late_routes-search") is (change == "same")
tool: Final = registry.get_tool("late_routes-search")
if openapi and change == "same":
assert tool is not None
assert await tool.handler() == "late server response"
else:
assert tool is None
@pytest.mark.asyncio
async def test_catalog_lookup_uses_one_snapshot_until_operation_finishes():
from datetime import timedelta