mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(mcp): reconcile concurrent route changes when publishing catalog
This commit is contained in:
parent
ed7cdd151f
commit
42821c102e
2 changed files with 53 additions and 4 deletions
|
|
@ -219,17 +219,34 @@ class TargetCatalog:
|
|||
{key: value.model_copy(deep=True) for key, value in previous_config.items()}
|
||||
)
|
||||
await self.manager.hydrate_config_servers_dcr_clients(tuple(staged_config.values()))
|
||||
staged_routing: Final = dict(self.manager.published_tool_routes)
|
||||
initial_routing: Final = MappingProxyType(dict(self.manager.published_tool_routes))
|
||||
staged_routing: Final = dict(initial_routing)
|
||||
closed: Final = asyncio.Event()
|
||||
routing_token: Final = self._staged_routing.set((staged_routing, closed))
|
||||
try:
|
||||
with global_mcp_tool_registry.catalog_scope(global_mcp_tool_registry.published_tools) as staged_tools:
|
||||
await self._reload()
|
||||
concurrent_routes: Final = self._unchanged_routing(
|
||||
previous_servers, self.manager.published_tool_routes
|
||||
previous_servers,
|
||||
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)
|
||||
}
|
||||
),
|
||||
)
|
||||
removed_routes: Final = initial_routing.keys() - self.manager.published_tool_routes.keys()
|
||||
retained_staged_routes: Final = self._unchanged_routing(
|
||||
previous_config | self.manager.registry, staged_routing
|
||||
previous_config | self.manager.registry,
|
||||
MappingProxyType(
|
||||
{
|
||||
name: owner
|
||||
for name, owner in staged_routing.items()
|
||||
if name not in removed_routes or owner != initial_routing[name]
|
||||
}
|
||||
),
|
||||
)
|
||||
self.manager.config_mcp_servers = {
|
||||
key: value.model_copy(
|
||||
|
|
@ -242,7 +259,7 @@ class TargetCatalog:
|
|||
for key, value in self.manager.config_mcp_servers.items()
|
||||
}
|
||||
global_mcp_tool_registry.tools = staged_tools
|
||||
self.manager.published_tool_routes = concurrent_routes | retained_staged_routes
|
||||
self.manager.published_tool_routes = retained_staged_routes | concurrent_routes
|
||||
finally:
|
||||
closed.set()
|
||||
self._staged_routing.reset(routing_token)
|
||||
|
|
|
|||
|
|
@ -14605,6 +14605,38 @@ async def test_catalog_operation_retains_routes_only_for_same_configured_target(
|
|||
assert (tools[0].name in manager.published_tool_routes) is (change == "discovery")
|
||||
|
||||
|
||||
@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):
|
||||
manager: Final = MCPServerManager()
|
||||
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"}
|
||||
|
||||
async def read_rows(**_kwargs):
|
||||
if change == "concurrent":
|
||||
manager.published_tool_routes["shared-search"] = "new_owner"
|
||||
elif change in ("staged", "staged_after_delete", "staged_with_concurrent"):
|
||||
if change == "staged_after_delete":
|
||||
manager.published_tool_routes.clear()
|
||||
elif change == "staged_with_concurrent":
|
||||
manager.published_tool_routes["shared-search"] = "concurrent_owner"
|
||||
manager.tool_name_to_mcp_server_name_mapping["shared-search"] = "new_owner"
|
||||
else:
|
||||
manager.published_tool_routes.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":
|
||||
assert "shared-search" not in manager.published_tool_routes
|
||||
else:
|
||||
assert manager.published_tool_routes["shared-search"] == "new_owner"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_observes_committed_update_and_delete_without_background_reload():
|
||||
from datetime import timedelta
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue