fix(mcp): make the pre-update identity snapshot advisory so a read failure cannot fail the edit

The snapshot read only feeds the stale-token purge decision; leaving it unguarded meant a failed
read would 500 an edit whose update would have succeeded, and it broke
test_edit_mcp_server_redacts_credentials, whose mocked prisma is not awaitable on the un-patched
get_mcp_server path. A failure now logs and skips the purge, consistent with the purge half already
being best-effort. Adds the first endpoint-level coverage of the edit purge wiring: purge on a
mint-relevant change, no purge when the identity is unchanged, and edit success with purge skipped
when the snapshot read raises
This commit is contained in:
Tin 2026-07-09 13:10:24 -07:00
parent 42388c3d68
commit c75184bec9
2 changed files with 88 additions and 2 deletions

View file

@ -2320,8 +2320,18 @@ if MCP_AVAILABLE:
},
)
# Snapshot the pre-update identity so we can detect a mint-relevant change below.
old_server_record = await get_mcp_server(prisma_client, payload.server_id)
# Snapshot the pre-update identity so we can detect a mint-relevant change below. The read is
# advisory (it only feeds the stale-token purge decision), so a failure skips the purge with a
# warning instead of failing the edit, whose primary job is the update itself.
try:
old_server_record = await get_mcp_server(prisma_client, payload.server_id)
except Exception as exc: # noqa: BLE001 - advisory read; invalidation is best-effort end-to-end
verbose_logger.warning(
"MCP server %s: could not snapshot the pre-update record; skipping the stale-token check: %s",
payload.server_id,
exc,
)
old_server_record = None
# try to update the mcp server
mcp_server_record_updated = await update_mcp_server(

View file

@ -5134,3 +5134,79 @@ def test_stamp_oauth2_flow_ignores_non_oauth2():
payload = _oauth2_create_payload(auth_type="none")
mgmt_endpoints.stamp_omitted_oauth2_flow(payload)
assert payload.oauth2_flow is None
async def _run_edit(old_record, updated_record):
from litellm.proxy.management_endpoints.mcp_management_endpoints import edit_mcp_server
server_id = updated_record.server_id
with (
patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(side_effect=old_record)
if isinstance(old_record, Exception)
else AsyncMock(return_value=old_record),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server",
AsyncMock(return_value=updated_record),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload",
autospec=True,
),
patch("litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager") as mock_manager,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.purge_user_oauth_credentials_for_server",
AsyncMock(return_value=1),
) as mock_purge,
):
mock_manager.update_server = AsyncMock()
mock_manager.reload_servers_from_database = AsyncMock()
payload = UpdateMCPServerRequest(server_id=server_id, alias=updated_record.alias, url=updated_record.url)
user_auth = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
result = await edit_mcp_server(payload=payload, user_api_key_dict=user_auth)
return result, mock_purge
@pytest.mark.asyncio
async def test_edit_mcp_server_purges_user_tokens_on_mint_relevant_change():
server_id = str(uuid.uuid4())
old = generate_mock_mcp_server_db_record(server_id=server_id, url="https://old.example.com/mcp")
updated = generate_mock_mcp_server_db_record(server_id=server_id, url="https://new.example.com/mcp")
result, mock_purge = await _run_edit(old, updated)
assert result.server_id == server_id
mock_purge.assert_awaited_once()
assert mock_purge.await_args.args[1] == server_id
@pytest.mark.asyncio
async def test_edit_mcp_server_skips_purge_when_identity_unchanged():
server_id = str(uuid.uuid4())
old = generate_mock_mcp_server_db_record(server_id=server_id, alias="Before")
updated = generate_mock_mcp_server_db_record(server_id=server_id, alias="After")
result, mock_purge = await _run_edit(old, updated)
assert result.server_id == server_id
mock_purge.assert_not_awaited()
@pytest.mark.asyncio
async def test_edit_mcp_server_snapshot_failure_skips_purge_but_edit_succeeds():
"""The pre-update snapshot read is advisory (it only feeds the purge decision); a read failure
must skip the stale-token check with a warning, never fail the edit itself."""
server_id = str(uuid.uuid4())
updated = generate_mock_mcp_server_db_record(server_id=server_id, url="https://new.example.com/mcp")
result, mock_purge = await _run_edit(RuntimeError("db read failed"), updated)
assert result.server_id == server_id
mock_purge.assert_not_awaited()