fix(mcp): reject configuration changes during scoped dispatch

This commit is contained in:
Joshua Valluru 2026-09-22 12:11:14 -07:00
parent c6b5f6da7f
commit 44a5992cf8
4 changed files with 141 additions and 4 deletions

View file

@ -37,7 +37,31 @@ class TargetCatalog:
async def refresh(self) -> None:
async with self._reload_lock:
await self._manager._reload_servers_from_database() # pyright: ignore[reportPrivateUsage] # existing staged loader
await self._reload()
async def _reload(self, *, reuse_unchanged: bool = False) -> None:
token: Final = self._scope.set(None)
try:
await self._manager._reload_servers_from_database(reuse_unchanged=reuse_unchanged) # pyright: ignore[reportPrivateUsage] # existing staged loader
finally:
self._scope.reset(token)
def assert_current(self, server: MCPServer) -> None:
snapshot: Final = self.current
if snapshot is None or server.server_id not in snapshot:
return
expected: Final = snapshot[server.server_id]
registered: Final = self._manager.registry.get(server.server_id) or self._manager.config_mcp_servers.get(
server.server_id
)
if (
registered is None
or registered.updated_at != expected.updated_at
or server.updated_at != expected.updated_at
):
from fastapi import HTTPException # noqa: PLC0415 # optional proxy dependency
raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation")
async def list(self) -> Mapping[str, MCPServer]:
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime proxy dependency
@ -48,7 +72,7 @@ class TargetCatalog:
async with self._reload_lock:
if prisma_client is not None:
try:
await self._manager._reload_servers_from_database(reuse_unchanged=True) # pyright: ignore[reportPrivateUsage] # existing staged loader
await self._reload(reuse_unchanged=True)
except Exception as exc: # noqa: BLE001 # never serve an unverified database snapshot
from fastapi import HTTPException # noqa: PLC0415 # optional proxy dependency

View file

@ -2181,6 +2181,7 @@ class MCPServerManager:
incomplete metadata for a server whose OAuth flow the gateway
runs itself.
"""
self.catalog.assert_current(server)
acquisition: Final = self._get_or_start_oauth_discovery_task(server)
if acquisition is None:
return self._registered_server(server)
@ -2193,11 +2194,13 @@ class MCPServerManager:
raise
match outcome:
case _OAuthDiscoveryResolved(resolved_server):
self.catalog.assert_current(resolved_server)
return resolved_server
case _OAuthDiscoveryStale():
return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale)
case _OAuthDiscoveryFailed(timed_out=timed_out):
current: Final = self._registered_server(server)
self.catalog.assert_current(current)
if current.is_client_forwarded_token:
return current
server_ref: Final = current.alias or current.server_name or current.name or current.server_id
@ -5544,6 +5547,7 @@ class MCPServerManager:
# Registration used add_server_prefix_to_name(base, get_server_prefix(server)),
# and tool_name is the bare base name by the time call_tool reaches here, so
# rebuilding the key the same way reproduces it exactly
self.catalog.assert_current(server)
registry_key: Final = add_server_prefix_to_name(tool_name, get_server_prefix(server))
tool: Final = global_mcp_tool_registry.get_tool(registry_key)
if tool is None:
@ -6599,9 +6603,9 @@ class MCPServerManager:
# prefix that lookups will use.
if not reuse_unchanged or new_server is not previous_registry.get(server_id):
await self._maybe_register_openapi_tools(new_server, initialize_mapping=False)
if new_server.spec_path:
registered_openapi_tools = True
registered_registry[server_id] = new_server
if new_server.spec_path:
registered_openapi_tools = True
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",

View file

@ -2552,6 +2552,9 @@ async def _handle_local_mcp_tool(name: str, arguments: dict[str, object]) -> Cal
"""
import inspect
server: Final = global_mcp_server_manager.server_owning_tool_name_prefix(name)
if server is not None:
global_mcp_server_manager.catalog.assert_current(server)
tool: Final = global_mcp_tool_registry.get_tool(name)
if not tool:
raise HTTPException(status_code=404, detail=f"Tool '{name}' not found")

View file

@ -14722,3 +14722,109 @@ async def test_catalog_reuses_openapi_tools_until_configuration_or_background_re
read_rows.return_value = []
await manager.catalog.list()
assert registry.list_tools() == []
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["edit", "delete"])
async def test_catalog_rejects_configuration_switch_before_client_creation(monkeypatch, change):
read_rows = AsyncMock(return_value=[_catalog_row()])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
async with manager.catalog.operation():
admitted = manager.get_mcp_server_by_id("catalog-server")
read_rows.return_value = [_catalog_row("updated")] if change == "edit" else []
await asyncio.create_task(manager.reload_servers_from_database())
with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory:
with pytest.raises(HTTPException) as exc:
await manager._create_mcp_client(admitted)
assert exc.value.status_code == 503
assert exc.value.detail == "MCP server configuration changed; retry the operation"
factory.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("dispatch", ["managed", "local"])
async def test_catalog_changed_openapi_handler_never_dispatches(monkeypatch, tmp_path, respx_mock, dispatch):
from litellm.proxy._experimental.mcp_server import operations, tool_registry
from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix
registry = tool_registry.MCPToolRegistry()
monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry)
monkeypatch.setattr(operations, "global_mcp_tool_registry", registry)
spec_path = tmp_path / "catalog-race.json"
spec_path.write_text(json.dumps({
"openapi": "3.0.0", "info": {"title": "Catalog", "version": "1"},
"paths": {"/echo": {"get": {"operationId": "echo"}}},
}))
row = _catalog_row().model_copy(update={"spec_path": str(spec_path)})
read_rows = AsyncMock(return_value=[row])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
destination = respx_mock.get("https://changed.example/echo").respond(200, text="must not execute")
async with manager.catalog.operation():
admitted = manager.get_mcp_server_by_id("catalog-server")
read_rows.return_value = [row.model_copy(update={
"url": "https://changed.example", "updated_at": datetime(2026, 1, 2),
})]
await asyncio.create_task(manager.reload_servers_from_database())
with pytest.raises(HTTPException) as exc:
if dispatch == "managed":
await manager._call_openapi_tool_handler(admitted, "echo", {})
else:
await operations._handle_local_mcp_tool(
add_server_prefix_to_name("echo", get_server_prefix(admitted)), {}
)
assert exc.value.status_code == 503
assert destination.call_count == 0
@pytest.mark.asyncio
async def test_catalog_rejects_old_admission_in_a_new_operation(monkeypatch):
read_rows = AsyncMock(return_value=[_catalog_row()])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
original = (await manager.catalog.list())["catalog-server"]
read_rows.return_value = [_catalog_row("updated")]
async with manager.catalog.operation():
with pytest.raises(HTTPException) as exc:
await manager._create_mcp_client(original)
assert exc.value.detail == "MCP server configuration changed; retry the operation"
@pytest.mark.asyncio
@pytest.mark.parametrize("change", [False, True])
async def test_catalog_deferred_discovery_preserves_admitted_configuration(monkeypatch, change):
row = _catalog_row().model_copy(update={"auth_type": MCPAuth.true_passthrough})
read_rows = AsyncMock(return_value=[row])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
entered = asyncio.Event()
release = asyncio.Event()
metadata = MCPOAuthMetadata(
authorization_url="https://issuer.example/authorize", token_url="https://issuer.example/token",
)
async def discover(server):
entered.set()
await release.wait()
return metadata
with patch.object(manager, "_discover_oauth_metadata_for_server", side_effect=discover):
async with manager.catalog.operation():
admitted = manager.get_mcp_server_by_id("catalog-server")
resolving = asyncio.create_task(manager.ensure_oauth_metadata_discovered(admitted))
await asyncio.wait_for(entered.wait(), timeout=1)
if change:
read_rows.return_value = [row.model_copy(update={"updated_at": datetime(2026, 1, 2)})]
await manager.reload_servers_from_database()
release.set()
if change:
with pytest.raises(HTTPException) as exc:
await asyncio.wait_for(resolving, timeout=1)
assert exc.value.detail == "MCP server configuration changed; retry the operation"
else:
resolved = await asyncio.wait_for(resolving, timeout=1)
assert resolved.authorization_url == metadata.authorization_url
assert resolved.token_url == metadata.token_url
assert resolved.updated_at == admitted.updated_at