mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
fix(mcp): keep server-wide listed tool metadata across per-user OAuth token invalidation
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
485f4522e4
commit
3c536bfc9a
2 changed files with 23 additions and 7 deletions
|
|
@ -2643,7 +2643,7 @@ class MCPServerManager:
|
|||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
|
||||
_warn_config_id_jag_server_outruns_sso(new_server)
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._invalidate_server_definition_caches(server_id)
|
||||
self.config_mcp_servers[server_id] = new_server
|
||||
self._set_oauth_discovery_deferred(
|
||||
server_id,
|
||||
|
|
@ -2845,7 +2845,7 @@ class MCPServerManager:
|
|||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
self._invalidate_discovery_lists(server.server_id)
|
||||
self._invalidate_server_definition_caches(server.server_id)
|
||||
prefix_root: Final = normalize_server_name(get_server_prefix(server))
|
||||
if server.spec_path and prefix_root:
|
||||
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
|
||||
|
|
@ -3222,7 +3222,7 @@ class MCPServerManager:
|
|||
# env_vars_are_encrypted=False.
|
||||
new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self._invalidate_discovery_lists(mcp_server.server_id)
|
||||
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)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -3259,7 +3259,7 @@ class MCPServerManager:
|
|||
previous_server=self.registry[mcp_server.server_id],
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self._invalidate_discovery_lists(mcp_server.server_id)
|
||||
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)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -4501,6 +4501,9 @@ class MCPServerManager:
|
|||
self._prompt_discovery_cache.invalidate(server_id)
|
||||
self._resource_discovery_cache.invalidate(server_id)
|
||||
self._template_discovery_cache.invalidate(server_id)
|
||||
|
||||
def _invalidate_server_definition_caches(self, server_id: str) -> None:
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._listed_tools_by_server_id.pop(server_id, None)
|
||||
|
||||
def _discovery_key(
|
||||
|
|
@ -6634,7 +6637,7 @@ class MCPServerManager:
|
|||
|
||||
for server_id in previous_registry.keys() | registered_registry.keys():
|
||||
if previous_registry.get(server_id) != registered_registry.get(server_id):
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._invalidate_server_definition_caches(server_id)
|
||||
self.registry = registered_registry
|
||||
# A discovery task may have published into ``previous_registry`` while
|
||||
# this replacement was being staged. Reconcile every published entry
|
||||
|
|
|
|||
|
|
@ -6606,19 +6606,32 @@ class TestMCPServerManager:
|
|||
assert by_prefixed_name is not None and by_prefixed_name.description == "v2"
|
||||
assert manager.get_listed_tool(server, "missing") is None
|
||||
|
||||
def test_invalidate_discovery_lists_drops_listed_tools(self):
|
||||
def test_server_definition_change_drops_listed_tools(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other")
|
||||
manager._create_prefixed_tools([MCPTool(name="echo", description="old", inputSchema={})], server)
|
||||
manager._create_prefixed_tools([MCPTool(name="ping", description="kept", inputSchema={})], other)
|
||||
|
||||
manager._invalidate_discovery_lists(server.server_id)
|
||||
manager._invalidate_server_definition_caches(server.server_id)
|
||||
|
||||
assert manager.get_listed_tool(server, "echo") is None
|
||||
kept = manager.get_listed_tool(other, "ping")
|
||||
assert kept is not None and kept.description == "kept"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_oauth_refresh_keeps_listed_tools(self):
|
||||
"""Tool definitions are server-wide, so one user's re-auth must not blank the metadata other
|
||||
callers' tool calls hand to pre-call guardrails."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
manager._create_prefixed_tools([MCPTool(name="echo", description="shared", inputSchema={})], server)
|
||||
|
||||
await manager.invalidate_user_oauth_token_cache("alice", server.server_id)
|
||||
|
||||
listed = manager.get_listed_tool(server, "echo")
|
||||
assert listed is not None and listed.description == "shared"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_allowed_mcp_servers_with_user_api_key_auth(self):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue