mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): drop the catalog again once a saved OpenAPI server's registry is rebuilt
add_server and update_server publish the saved definition before the OpenAPI registry entries are rebuilt from the spec, so a listing recorded during that fetch held the pre-save entries under the new generation. The generation is bumped a second time after the registry refresh. The during-hook task no longer accepts a listed entry, the one-line wrapper over get_listed_tool is inlined at its two call sites, and the per-server generation map is a plain dict.
This commit is contained in:
parent
9f46ec11b2
commit
71be87b893
3 changed files with 46 additions and 20 deletions
|
|
@ -2035,7 +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._listed_tools_generations: dict[str, int] = {} # mutable-ok: bumped per server save
|
||||
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
|
||||
|
|
@ -3351,6 +3351,8 @@ class MCPServerManager:
|
|||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Added MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -3388,6 +3390,8 @@ class MCPServerManager:
|
|||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Updated MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -4672,9 +4676,7 @@ 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}
|
||||
)
|
||||
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:
|
||||
|
|
@ -4727,8 +4729,7 @@ class MCPServerManager:
|
|||
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."""
|
||||
read before the listing's upstream fetch; the record is skipped when it no longer matches."""
|
||||
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)
|
||||
|
|
@ -6059,7 +6060,6 @@ class MCPServerManager:
|
|||
start_time: datetime.datetime,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
tool: MCPTool | None = None,
|
||||
):
|
||||
"""Create and return a during hook task for MCP tool calls.
|
||||
|
||||
|
|
@ -6074,8 +6074,6 @@ class MCPServerManager:
|
|||
tool_name=name,
|
||||
arguments=arguments,
|
||||
server_name=server_name_from_prefix,
|
||||
tool_description=tool.description if tool is not None else None,
|
||||
tool_input_schema=tool.input_schema if tool is not None else None,
|
||||
start_time=start_time.timestamp() if start_time else None,
|
||||
hidden_params=HiddenParams(),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -81,7 +81,6 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
|||
outcome_wire_value,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
ListedToolsCaller,
|
||||
MCPServerManager,
|
||||
_caller_authorization_fans_out,
|
||||
_client_forwarded_authorization_headers,
|
||||
|
|
@ -1610,12 +1609,6 @@ async def _list_mcp_resource_templates(
|
|||
return managed_resource_templates
|
||||
|
||||
|
||||
def _registered_tool_metadata(name: str, server: MCPServer, caller: ListedToolsCaller) -> MCPTool | None:
|
||||
"""The tool as ``tools/list`` served it to this caller (pinned, overridden, guardrail-masked), or None when
|
||||
no listing was recorded so the call hands the hooks name and arguments only."""
|
||||
return global_mcp_server_manager.get_listed_tool(server, name, caller)
|
||||
|
||||
|
||||
def _resolve_display_name_to_original(
|
||||
name: str,
|
||||
allowed_mcp_servers: list[MCPServer],
|
||||
|
|
@ -2086,9 +2079,9 @@ async def _execute_mcp_tool(
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=_registered_tool_metadata(
|
||||
original_tool_name,
|
||||
tool=global_mcp_server_manager.get_listed_tool(
|
||||
mcp_server,
|
||||
original_tool_name,
|
||||
listed_tools_caller_for(
|
||||
mcp_server,
|
||||
user_api_key_auth,
|
||||
|
|
@ -2210,9 +2203,9 @@ async def _execute_mcp_tool(
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=_registered_tool_metadata(
|
||||
original_tool_name,
|
||||
tool=global_mcp_server_manager.get_listed_tool(
|
||||
prefix_server,
|
||||
original_tool_name,
|
||||
listed_tools_caller_for(
|
||||
prefix_server,
|
||||
user_api_key_auth,
|
||||
|
|
|
|||
|
|
@ -7232,6 +7232,41 @@ class TestMCPServerManager:
|
|||
listed = manager.get_listed_tool(server, "turn", caller)
|
||||
assert listed is not None and listed.description == "after save"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_server_refreshing_openapi_tools_drops_a_listing_recorded_during_the_spec_fetch(self):
|
||||
"""An OpenAPI server's registry entries are rebuilt after the save is published, so a listing that
|
||||
records while the spec is fetched holds the pre-save entries; the catalog is dropped again once the
|
||||
registry is current."""
|
||||
manager = MCPServerManager()
|
||||
old = MCPServer(
|
||||
server_id="srv", name="srv", transport=MCPTransport.http, url="http://old", spec_path="/old.json"
|
||||
)
|
||||
manager.registry[old.server_id] = old
|
||||
new = MCPServer(
|
||||
server_id="srv", name="srv", transport=MCPTransport.http, url="http://new", spec_path="/new.json"
|
||||
)
|
||||
caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister"))
|
||||
|
||||
async def register_while_a_listing_records(server: MCPServer, *, initialize_mapping: bool = True) -> None:
|
||||
manager._record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="search", description="pre-save", inputSchema={})],
|
||||
caller,
|
||||
manager._listed_tools_generations.get(server.server_id, 0),
|
||||
)
|
||||
|
||||
manager.build_mcp_server_from_table = AsyncMock(return_value=new)
|
||||
manager._maybe_register_openapi_tools = register_while_a_listing_records
|
||||
manager.prime_oauth_metadata_discovery = MagicMock()
|
||||
record = LiteLLM_MCPServerTable(
|
||||
server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http
|
||||
)
|
||||
|
||||
await manager.update_server(record)
|
||||
|
||||
assert manager.registry["srv"] is new
|
||||
assert manager.get_listed_tool(new, "search", caller) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_oauth_refresh_keeps_listed_tools(self):
|
||||
manager = MCPServerManager()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue