diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 0f5eca285e3..ed61b4360d9 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -133,21 +133,25 @@ class TargetCatalog: self._retain_discovered_routing(snapshot) def _retain_discovered_routing(self, snapshot: CatalogSnapshot) -> None: + self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing( + snapshot.servers, snapshot.routing + ) + + def _unchanged_routing( + self, servers: Mapping[str, MCPServer], routing: Mapping[str, str] + ) -> MappingProxyType[str, str]: from litellm.proxy._experimental.mcp_server.utils import normalize_server_name current: Final = self.manager.config_mcp_servers | self.manager.registry unchanged_owners: Final = frozenset( owner - for key, server in snapshot.servers.items() - if current.get(key) == server + for key, server in servers.items() + if (candidate := current.get(key)) is not None + and _configuration_identity(candidate) == _configuration_identity(server) for owner in self.manager.owned_mapping_values(server) ) - self.manager.published_tool_routes = self.manager.published_tool_routes | MappingProxyType( - { - name: owner - for name, owner in snapshot.routing.items() - if normalize_server_name(owner) in unchanged_owners - } + return MappingProxyType( + {name: owner for name, owner in routing.items() if normalize_server_name(owner) in unchanged_owners} ) async def resolve(self, identifier: str, client_ip: str | None = None) -> MCPServer | None: @@ -202,9 +206,18 @@ class TargetCatalog: from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry async with self._refresh_lock: - staged_config: Final = { - key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items() - } + previous_config: Final = MappingProxyType( + {key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items()} + ) + previous_servers: Final = previous_config | MappingProxyType( + {key: value.model_copy(deep=True) for key, value in self.manager.registry.items()} + ) + config_identities: Final = MappingProxyType( + {key: _configuration_identity(value) for key, value in previous_config.items()} + ) + staged_config: Final = MappingProxyType( + {key: value.model_copy(deep=True) for key, value in previous_config.items()} + ) await self.manager.hydrate_config_servers_dcr_clients(tuple(staged_config.values())) staged_routing: Final = dict(self.manager.published_tool_routes) closed: Final = asyncio.Event() @@ -212,9 +225,24 @@ class TargetCatalog: try: with global_mcp_tool_registry.catalog_scope(global_mcp_tool_registry.published_tools) as staged_tools: await self._reload() - self.manager.config_mcp_servers = staged_config + concurrent_routes: Final = self._unchanged_routing( + previous_servers, self.manager.published_tool_routes + ) + retained_staged_routes: Final = self._unchanged_routing( + previous_config | self.manager.registry, staged_routing + ) + self.manager.config_mcp_servers = { + key: value.model_copy( + update=staged_config[key].model_dump( + include=frozenset(("client_id", "client_secret", "token_endpoint_auth_method")) + ) + ) + if config_identities.get(key) == _configuration_identity(value) + else value + for key, value in self.manager.config_mcp_servers.items() + } global_mcp_tool_registry.tools = staged_tools - self.manager.published_tool_routes = staged_routing + self.manager.published_tool_routes = concurrent_routes | retained_staged_routes finally: closed.set() self._staged_routing.reset(routing_token) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 2b92367f186..cff3d5f78c6 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -2,7 +2,7 @@ import os import pytest from unittest.mock import AsyncMock, MagicMock, patch -from contextlib import asynccontextmanager +from contextlib import asynccontextmanager, nullcontext from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -909,6 +909,7 @@ async def test_get_tools_from_mcp_servers(): # Create a mock manager mock_manager = AsyncMock() + mock_manager.catalog.operation = nullcontext mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id"] ) @@ -937,6 +938,7 @@ async def test_get_tools_from_mcp_servers(): # Test Case 2: Without specific MCP servers # Create a different mock manager for the second test case mock_manager_2 = AsyncMock() + mock_manager_2.catalog.operation = nullcontext mock_manager_2.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id"] ) @@ -984,6 +986,7 @@ async def test_get_tools_from_mcp_servers(): # Test Case 3: With specific MCP servers and access groups # Create a mock manager mock_manager = AsyncMock() + mock_manager.catalog.operation = nullcontext mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id", "server3_id"] ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_block_recording.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_block_recording.py index b8aadef430f..746130fb892 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_block_recording.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_block_recording.py @@ -66,8 +66,7 @@ async def _call_block(logging_obj, order: list, *, user_api_key_auth=mock.sentin proxy_logging_obj = mock.MagicMock() proxy_logging_obj.post_call_failure_hook.side_effect = _record_post_call_failure_hook - fake_proxy_server = types.ModuleType("litellm.proxy.proxy_server") - fake_proxy_server.proxy_logging_obj = proxy_logging_obj # pyright: ignore[reportAttributeAccessIssue] + fake_proxy_server = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj, prisma_client=None) with mock.patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}): with contextlib.suppress(HTTPException): 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 3615d9681f0..fc70689b6c9 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 @@ -14519,6 +14519,92 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie auth_context_var.reset(token) +@pytest.mark.asyncio +@pytest.mark.parametrize("anchored", [False, True]) +async def test_catalog_reload_preserves_concurrent_config_discovery_and_routes(anchored): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="config-race", name="config_race", transport=MCPTransport.http, + url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2, + issuer="https://issuer.example" if anchored else None, issuer_is_anchored=anchored) + manager.config_mcp_servers = {server.server_id: server} + manager._set_oauth_discovery_deferred(server.server_id, True) + generation: Final = manager.oauth_discovery_slot(server.server_id).generation + resolved: Final = server.model_copy(update={"authorization_url": "https://issuer.example/authorize", + "token_url": "https://issuer.example/token", "scopes": ["read"], "issuer": "https://issuer.example"}) + + async def hydrate(target: MCPServer) -> bool: + target.client_id = "persisted-client" + return True + + async def read_rows(**_kwargs): + assert manager._publish_resolved_oauth_server(resolved, generation) is resolved + manager.published_tool_routes["config_race-search"] = "config_race" + return [] + + prisma: Final = MagicMock() + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.hydrate_config_server_dcr_client", side_effect=hydrate), + ): + await manager.reload_servers_from_database() + current: Final = manager.get_mcp_server_by_id(server.server_id) + assert current.authorization_url == resolved.authorization_url + assert current.token_url == resolved.token_url + assert current.scopes == ["read"] + assert current.issuer == "https://issuer.example" + assert current.client_id == "persisted-client" + assert manager.oauth_discovery_slot(server.server_id) is None + assert manager.published_tool_routes["config_race-search"] == "config_race" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["credentials", "delete"]) +async def test_catalog_reload_does_not_restore_replaced_config_credentials_or_routes(change): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="config-replaced", name="config_replaced", transport=MCPTransport.http, + client_id="previous-client") + manager.config_mcp_servers = {server.server_id: server} + manager.published_tool_routes["config_replaced-search"] = "config_replaced" + + async def read_rows(**_kwargs): + manager.config_mcp_servers = {} if change == "delete" else { + server.server_id: server.model_copy(update={"client_id": "new-client"})} + return [] + + prisma: Final = MagicMock() + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await manager.reload_servers_from_database() + current: Final = manager.get_mcp_server_by_id(server.server_id) + if change == "delete": + assert current is None + else: + assert current.client_id == "new-client" + assert "config_replaced-search" not in manager.published_tool_routes + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["discovery", "update", "delete"]) +async def test_catalog_operation_retains_routes_only_for_same_configured_target(change): + from mcp.types import Tool + + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="route-race", name="route_race", transport=MCPTransport.http, + url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2) + manager.registry = {server.server_id: server} + with patch("litellm.proxy.proxy_server.prisma_client", None): + async with manager.catalog.operation(): + tools: Final = manager._create_prefixed_tools([Tool(name="search", inputSchema={})], server) + if change == "delete": + manager.registry = {} + else: + manager.registry[server.server_id] = server.model_copy(update=( + {"authorization_url": "https://issuer.example/authorize", "token_url": "https://issuer.example/token"} + if change == "discovery" else {"url": "https://changed.example/mcp"})) + assert (tools[0].name in manager.published_tool_routes) is (change == "discovery") + + @pytest.mark.asyncio async def test_catalog_observes_committed_update_and_delete_without_background_reload(): from datetime import timedelta diff --git a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py index c2c204e7024..7ed5b5e36a9 100644 --- a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py @@ -1,5 +1,6 @@ import sys import types +from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock @@ -77,6 +78,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: """Patch the MCP tool-call plumbing so _execute_tool_calls can run in tests.""" call_tool = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)) fake_manager = types.SimpleNamespace( + catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=call_tool, _get_mcp_server_from_tool_name=MagicMock(return_value=None),