mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(mcp): reject configuration changes during scoped dispatch
This commit is contained in:
parent
c6b5f6da7f
commit
44a5992cf8
4 changed files with 141 additions and 4 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue