mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(mcp): preserve discovery published during catalog refresh
This commit is contained in:
parent
74a470410d
commit
ed7cdd151f
5 changed files with 134 additions and 16 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue