mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(mcp): retain valid live routes and handlers during refresh
This commit is contained in:
parent
42821c102e
commit
d4c1afdebc
2 changed files with 143 additions and 9 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue