mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(mcp): retain discovered routes across cached catalog operations
This commit is contained in:
parent
886aadefef
commit
57dfed6e39
3 changed files with 49 additions and 6 deletions
|
|
@ -167,10 +167,11 @@ class TargetCatalog:
|
|||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
shared: Final = await self._fresh_snapshot()
|
||||
initial_routing: Final = self._unchanged_routing(shared.servers, self.manager.published_tool_routes)
|
||||
snapshot: Final = replace(
|
||||
shared,
|
||||
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
|
||||
routing=dict(shared.routing),
|
||||
routing=dict(initial_routing),
|
||||
)
|
||||
closed: Final = asyncio.Event()
|
||||
token: Final = self._operation.set((snapshot, closed))
|
||||
|
|
@ -180,11 +181,12 @@ class TargetCatalog:
|
|||
finally:
|
||||
closed.set()
|
||||
self._operation.reset(token)
|
||||
self._retain_discovered_routing(snapshot)
|
||||
self._retain_discovered_routing(snapshot, initial_routing)
|
||||
|
||||
def _retain_discovered_routing(self, snapshot: CatalogSnapshot) -> None:
|
||||
def _retain_discovered_routing(self, snapshot: CatalogSnapshot, initial_routing: Mapping[str, str]) -> None:
|
||||
self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing(
|
||||
snapshot.servers, snapshot.routing
|
||||
snapshot.servers,
|
||||
{name: owner for name, owner in snapshot.routing.items() if initial_routing.get(name) != owner},
|
||||
)
|
||||
|
||||
def _unchanged_routing(
|
||||
|
|
|
|||
|
|
@ -16990,3 +16990,44 @@ async def test_upstream_preparation_honors_case_sensitive_extra_command(monkeypa
|
|||
client: Final = await MCPServerManager()._create_mcp_client(server)
|
||||
assert client.stdio_config is not None
|
||||
assert client.stdio_config["command"] == "/opt/tools/CustomRunner"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("overlap", [False, True])
|
||||
async def test_catalog_cached_revision_retains_newly_discovered_tool_routes(monkeypatch, overlap):
|
||||
from mcp.types import Tool
|
||||
|
||||
read_rows = AsyncMock(return_value=[])
|
||||
_catalog_database(monkeypatch, read_rows, AsyncMock(return_value=_revision_row(7)))
|
||||
manager = MCPServerManager()
|
||||
first = MCPServer(server_id="first", name="first", transport=MCPTransport.http)
|
||||
second = MCPServer(server_id="second", name="second", transport=MCPTransport.http)
|
||||
manager.config_mcp_servers = {server.server_id: server for server in (first, second)}
|
||||
manager.published_tool_routes = {"search": "first"}
|
||||
ready = asyncio.Event()
|
||||
release = asyncio.Event()
|
||||
|
||||
async def read_catalog():
|
||||
async with manager.catalog.operation():
|
||||
assert manager._get_mcp_server_from_tool_name("search").server_id == "first"
|
||||
ready.set()
|
||||
await release.wait()
|
||||
|
||||
reader = asyncio.create_task(read_catalog()) if overlap else None
|
||||
try:
|
||||
if reader is not None:
|
||||
await asyncio.wait_for(ready.wait(), 2)
|
||||
async with manager.catalog.operation():
|
||||
manager._create_prefixed_tools([Tool(name="search", input_schema={})], second)
|
||||
assert manager.published_tool_routes["search"] == "second"
|
||||
release.set()
|
||||
if reader is not None:
|
||||
await asyncio.wait_for(reader, 2)
|
||||
async with manager.catalog.operation():
|
||||
assert manager._get_mcp_server_from_tool_name("search").server_id == "second"
|
||||
assert manager.published_tool_routes["search"] == "second"
|
||||
read_rows.assert_awaited_once()
|
||||
finally:
|
||||
if reader is not None:
|
||||
reader.cancel()
|
||||
await asyncio.gather(reader, return_exceptions=True)
|
||||
|
|
|
|||
|
|
@ -7735,8 +7735,8 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti
|
|||
),
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_registry",
|
||||
return_value={
|
||||
"registry",
|
||||
{
|
||||
requested_server.server_id: requested_server,
|
||||
collision_server.server_id: collision_server,
|
||||
},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue