mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(mcp): retain rediscovered routes across concurrent updates
This commit is contained in:
parent
c72b7b8824
commit
47e47f62ae
3 changed files with 43 additions and 8 deletions
|
|
@ -5,7 +5,8 @@ from __future__ import annotations
|
|||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
|
||||
from collections import UserDict
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, MutableMapping, Sequence
|
||||
from contextlib import asynccontextmanager
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass, replace
|
||||
|
|
@ -27,12 +28,23 @@ _P = ParamSpec("_P")
|
|||
_R = TypeVar("_R")
|
||||
|
||||
|
||||
class _OperationRoutes(UserDict[str, str]):
|
||||
def __init__(self, initial: Mapping[str, str]) -> None:
|
||||
super().__init__()
|
||||
self.data = dict(initial)
|
||||
self.written_names: set[str] = set() # mutable-ok: record route writes without copying the journal per tool
|
||||
|
||||
def __setitem__(self, name: str, owner: str) -> None:
|
||||
self.data[name] = owner
|
||||
self.written_names.add(name)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CatalogSnapshot:
|
||||
servers: Mapping[str, MCPServer]
|
||||
identity: str
|
||||
tools: Mapping[str, MCPTool]
|
||||
routing: dict[str, str] # mutable-ok: tool routes are remapped in place through this snapshot's live dict
|
||||
routing: MutableMapping[str, str]
|
||||
|
||||
|
||||
def _configuration_identity(server: MCPServer) -> str:
|
||||
|
|
@ -96,7 +108,7 @@ class TargetCatalog:
|
|||
scoped: Final = self._operation.get()
|
||||
return scoped[0] if scoped is not None and not scoped[1].is_set() else None
|
||||
|
||||
def routing(self) -> dict[str, str]: # mutable-ok: returns the live routing dict that callers remap in place
|
||||
def routing(self) -> MutableMapping[str, str]:
|
||||
staged: Final = self._staged_routing.get()
|
||||
if staged is not None and not staged[1].is_set():
|
||||
return staged[0]
|
||||
|
|
@ -166,10 +178,11 @@ class TargetCatalog:
|
|||
|
||||
shared: Final = await self._fresh_snapshot()
|
||||
initial_routing: Final = self._unchanged_routing(shared.servers, self.manager.published_tool_routes)
|
||||
routing: Final = _OperationRoutes(initial_routing)
|
||||
snapshot: Final = replace(
|
||||
shared,
|
||||
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
|
||||
routing=dict(initial_routing),
|
||||
routing=routing,
|
||||
)
|
||||
closed: Final = asyncio.Event()
|
||||
token: Final = self._operation.set((snapshot, closed))
|
||||
|
|
@ -179,12 +192,12 @@ class TargetCatalog:
|
|||
finally:
|
||||
closed.set()
|
||||
self._operation.reset(token)
|
||||
self._retain_discovered_routing(snapshot, initial_routing)
|
||||
self._retain_discovered_routing(snapshot, frozenset(routing.written_names))
|
||||
|
||||
def _retain_discovered_routing(self, snapshot: CatalogSnapshot, initial_routing: Mapping[str, str]) -> None:
|
||||
def _retain_discovered_routing(self, snapshot: CatalogSnapshot, written_names: frozenset[str]) -> None:
|
||||
self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing(
|
||||
snapshot.servers,
|
||||
{name: owner for name, owner in snapshot.routing.items() if initial_routing.get(name) != owner},
|
||||
{name: owner for name, owner in snapshot.routing.items() if name in written_names},
|
||||
)
|
||||
|
||||
def _unchanged_routing(
|
||||
|
|
|
|||
|
|
@ -2417,7 +2417,7 @@ class MCPServerManager:
|
|||
@property
|
||||
def tool_name_to_mcp_server_name_mapping(
|
||||
self,
|
||||
) -> dict[str, str]: # mutable-ok: callers remap tool routes through this dict
|
||||
) -> MutableMapping[str, str]:
|
||||
return self.catalog.routing()
|
||||
|
||||
@tool_name_to_mcp_server_name_mapping.setter
|
||||
|
|
|
|||
|
|
@ -16716,6 +16716,28 @@ async def test_catalog_delete_drops_derived_tool_mapping(monkeypatch):
|
|||
assert not await manager.catalog.list()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("rediscover", [False, True])
|
||||
async def test_catalog_publishes_rediscovered_routes_despite_concurrent_owner_change(monkeypatch, rediscover):
|
||||
from mcp.types import Tool
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
manager = MCPServerManager()
|
||||
selected = MCPServer(server_id="selected", name="selected", transport=MCPTransport.http)
|
||||
competing = MCPServer(server_id="competing", name="competing", transport=MCPTransport.http)
|
||||
manager.config_mcp_servers = {server.server_id: server for server in (selected, competing)}
|
||||
manager.published_tool_routes = {"shared_tool": selected.name}
|
||||
async with manager.catalog.operation():
|
||||
manager.published_tool_routes["shared_tool"] = competing.name
|
||||
if rediscover:
|
||||
manager._create_prefixed_tools([Tool(name="shared_tool", input_schema={})], selected)
|
||||
assert manager._get_mcp_server_from_tool_name("shared_tool").server_id == selected.server_id
|
||||
async with manager.catalog.operation():
|
||||
expected = selected if rediscover else competing
|
||||
assert manager._get_mcp_server_from_tool_name("shared_tool").server_id == expected.server_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("anchored", [False, True])
|
||||
async def test_catalog_reload_preserves_concurrent_config_discovery_and_routes(anchored):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue