mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): drop a listed catalog recorded across a server save
_record_listed_tools ran after the awaited upstream fetch, so a PUT /v1/mcp/server that landed mid-fetch had its invalidation undone when the fetch completed: hooks then saw the pre-save description next to the post-save definition until the next listing, instead of name and arguments only. The manager now keeps a per-server listed-tools generation, bumped by _invalidate_server_definition_caches. _get_tools_from_server reads it before the fetch and _record_listed_tools skips the write when it moved; the next listing records normally.
This commit is contained in:
parent
4b0b326dc2
commit
9f46ec11b2
2 changed files with 57 additions and 4 deletions
|
|
@ -2035,6 +2035,7 @@ class MCPServerManager:
|
|||
}
|
||||
"""
|
||||
self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list
|
||||
self._listed_tools_generations: Mapping[str, int] = MappingProxyType({})
|
||||
self._upstream_initialize_instructions_by_server_id: dict[str, str] = {}
|
||||
# Per-server monotonic timestamp of last upstream prefetch attempt (success,
|
||||
# empty result, or failure). Used to throttle re-probes for servers that do
|
||||
|
|
@ -4504,6 +4505,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
listed_generation: Final = self._listed_tools_generations.get(server.server_id, 0)
|
||||
|
||||
try:
|
||||
# Tool *listing* must not be blocked by missing per-user env vars —
|
||||
|
|
@ -4602,7 +4604,7 @@ class MCPServerManager:
|
|||
# through _create_prefixed_tools — that would add the prefix a second
|
||||
# time producing "test_petstore-test_petstore-getinventory".
|
||||
unprefixed_tools: Final = guarded_openapi
|
||||
self._record_listed_tools(server, unprefixed_tools, listed_caller)
|
||||
self._record_listed_tools(server, unprefixed_tools, listed_caller, listed_generation)
|
||||
if not add_prefix:
|
||||
return unprefixed_tools
|
||||
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
|
||||
|
|
@ -4618,7 +4620,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
)
|
||||
prefixed_or_original_tools: Final = self._create_prefixed_tools(
|
||||
guarded_tools, server, add_prefix=add_prefix, caller=listed_caller
|
||||
guarded_tools, server, add_prefix=add_prefix, caller=listed_caller, generation=listed_generation
|
||||
)
|
||||
|
||||
return prefixed_or_original_tools
|
||||
|
|
@ -4670,6 +4672,9 @@ class MCPServerManager:
|
|||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._listed_tools_by_server_id.pop(server_id, None)
|
||||
self._listed_tools_generations = MappingProxyType(
|
||||
{**self._listed_tools_generations, server_id: self._listed_tools_generations.get(server_id, 0) + 1}
|
||||
)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
|
||||
|
|
@ -4715,8 +4720,17 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
def _record_listed_tools(
|
||||
self, server: MCPServer, tools: Sequence[MCPTool], caller: ListedToolsCaller | None
|
||||
self,
|
||||
server: MCPServer,
|
||||
tools: Sequence[MCPTool],
|
||||
caller: ListedToolsCaller | None,
|
||||
generation: int | None = None,
|
||||
) -> None:
|
||||
"""Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation
|
||||
read before the listing's upstream fetch; a server save that landed mid-fetch moved it, and the
|
||||
pre-save catalog is then dropped rather than written over the invalidation."""
|
||||
if generation is not None and generation != self._listed_tools_generations.get(server.server_id, 0):
|
||||
return
|
||||
identity: Final = self._listed_tools_identity(server, caller)
|
||||
listing: Final = MappingProxyType({tool.name: tool for tool in tools})
|
||||
existing: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({}))
|
||||
|
|
@ -5638,6 +5652,7 @@ class MCPServerManager:
|
|||
server: MCPServer,
|
||||
add_prefix: bool = True,
|
||||
caller: ListedToolsCaller | None = None,
|
||||
generation: int | None = None,
|
||||
) -> list[MCPTool]:
|
||||
"""
|
||||
Create prefixed tools and update tool mapping.
|
||||
|
|
@ -5663,7 +5678,7 @@ class MCPServerManager:
|
|||
for spelling in iter_known_tool_name_spellings(original_name, server):
|
||||
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
|
||||
|
||||
self._record_listed_tools(server, tools, caller)
|
||||
self._record_listed_tools(server, tools, caller, generation)
|
||||
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
|
||||
return prefixed_tools
|
||||
|
||||
|
|
|
|||
|
|
@ -7194,6 +7194,44 @@ class TestMCPServerManager:
|
|||
kept = manager.get_listed_tool(other, "ping")
|
||||
assert kept is not None and kept.description == "kept"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_server_save_during_an_in_flight_listing_is_not_undone_by_the_stale_record(self):
|
||||
"""A PUT /v1/mcp/server that lands while a listing awaits its upstream fetch drops the server's
|
||||
catalog; the fetch completing afterwards must not write the pre-save catalog back, or hooks see
|
||||
the old description next to the new definition until the next listing."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
user = UserAPIKeyAuth(api_key="sk-litellm", user_id="saver")
|
||||
fetch_started = asyncio.Event()
|
||||
release_fetch = asyncio.Event()
|
||||
|
||||
async def fetch(client, name):
|
||||
fetch_started.set()
|
||||
await release_fetch.wait()
|
||||
return [MCPTool(name="turn", description="before save", inputSchema={})]
|
||||
|
||||
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
|
||||
manager._fetch_tools_with_timeout = fetch
|
||||
caller = ListedToolsCaller(user_api_key_auth=user)
|
||||
|
||||
async def list_tools() -> None:
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user)
|
||||
|
||||
listing = asyncio.create_task(list_tools())
|
||||
await fetch_started.wait()
|
||||
manager._invalidate_server_definition_caches(server.server_id)
|
||||
release_fetch.set()
|
||||
await listing
|
||||
|
||||
assert manager.get_listed_tool(server, "turn", caller) is None
|
||||
|
||||
manager._fetch_tools_with_timeout = AsyncMock(
|
||||
return_value=[MCPTool(name="turn", description="after save", inputSchema={})]
|
||||
)
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user)
|
||||
listed = manager.get_listed_tool(server, "turn", caller)
|
||||
assert listed is not None and listed.description == "after save"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_oauth_refresh_keeps_listed_tools(self):
|
||||
manager = MCPServerManager()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue