diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 2938f1362af..f49fa718139 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 768d985a095..d8511617699 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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): """