diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index fddccf202b8..67486efd5bb 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -692,6 +692,7 @@ async def rotate_mcp_server_credentials_master_key( mcp_servers = await prisma_client.db.litellm_mcpservertable.find_many() + updated = 0 for mcp_server in mcp_servers: update_data: Dict[str, Any] = {} @@ -721,6 +722,11 @@ async def rotate_mcp_server_credentials_master_key( where={"server_id": mcp_server.server_id}, data=update_data, ) + updated += 1 + verbose_proxy_logger.info( + "rotate_mcp_server_credentials_master_key: rotated %d MCP server row(s)", + updated, + ) def _decode_user_credential(stored: str) -> Optional[str]: @@ -775,6 +781,8 @@ async def rotate_mcp_user_credentials_master_key( are logged and skipped so one corrupt row does not abort the rotation. """ rows = await prisma_client.db.litellm_mcpusercredentials.find_many() + rotated = 0 + skipped = 0 for row in rows: plaintext = _decode_user_credential(row.credential_b64) if plaintext is None: @@ -784,6 +792,7 @@ async def rotate_mcp_user_credentials_master_key( row.user_id, row.server_id, ) + skipped += 1 continue re_encrypted = encrypt_value_helper( plaintext, new_encryption_key=new_master_key @@ -797,6 +806,12 @@ async def rotate_mcp_user_credentials_master_key( }, data={"credential_b64": re_encrypted}, ) + rotated += 1 + verbose_proxy_logger.info( + "rotate_mcp_user_credentials_master_key: rotated %d row(s), skipped %d", + rotated, + skipped, + ) async def rotate_mcp_user_env_vars_master_key( @@ -810,6 +825,8 @@ async def rotate_mcp_user_env_vars_master_key( that may still be recoverable. """ rows = await prisma_client.db.litellm_mcpuserenvvars.find_many() + rotated = 0 + skipped = 0 for row in rows: plaintext = decrypt_value_helper( value=row.values_b64, @@ -824,6 +841,7 @@ async def rotate_mcp_user_env_vars_master_key( row.user_id, row.server_id, ) + skipped += 1 continue re_encrypted = encrypt_value_helper( plaintext, new_encryption_key=new_master_key @@ -837,6 +855,12 @@ async def rotate_mcp_user_env_vars_master_key( }, data={"values_b64": re_encrypted}, ) + rotated += 1 + verbose_proxy_logger.info( + "rotate_mcp_user_env_vars_master_key: rotated %d row(s), skipped %d", + rotated, + skipped, + ) async def store_user_credential( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py index af74fa29c76..febfa7153f8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -1384,6 +1384,70 @@ def test_reencrypt_global_env_var_values_handles_json_string(env_vars_salt_key): assert rebuilt[0]["value"] != "s3cr3t-p@ss" +@pytest.mark.asyncio +async def test_rotate_mcp_user_env_vars_logs_rotated_and_skipped_counts( + env_vars_salt_key, monkeypatch +): + """Master-key rotation is a rare, high-stakes batch op, so it emits one + summary line. The counts must track real work: a decryptable row is + re-encrypted and counted as rotated, while a row that no longer decrypts is + left untouched and counted as skipped.""" + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy._experimental.mcp_server.db as mcp_db + from litellm.proxy._experimental.mcp_server.db import ( + rotate_mcp_user_env_vars_master_key, + ) + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + def _row(user_id, server_id, blob): + row = MagicMock() + row.user_id = user_id + row.server_id = server_id + row.values_b64 = blob + return row + + import json + + # Encrypted under an unrelated key, so it won't decrypt under the active salt + # key and must be skipped rather than re-encrypted. + undecryptable = encrypt_value_helper( + json.dumps({"X": "y"}), new_encryption_key="unrelated-key-9999" + ) + good_one = _row("alice", "srv-1", _encrypted_user_env_blob({"GH_TOKEN": "tok-1"})) + good_two = _row("bob", "srv-2", _encrypted_user_env_blob({"GH_TOKEN": "tok-2"})) + bad = _row("carol", "srv-3", undecryptable) + + prisma = MagicMock() + prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock( + return_value=[good_one, good_two, bad] + ) + prisma.db.litellm_mcpuserenvvars.update = AsyncMock() + + logger = MagicMock() + monkeypatch.setattr(mcp_db, "verbose_proxy_logger", logger) + + await rotate_mcp_user_env_vars_master_key(prisma, new_master_key="rotated-key-0000") + + update = prisma.db.litellm_mcpuserenvvars.update + assert update.await_count == 2 + updated_servers = { + call.kwargs["where"]["user_id_server_id"]["server_id"] + for call in update.call_args_list + } + assert updated_servers == {"srv-1", "srv-2"} # srv-3 was skipped, not rotated + for call in update.call_args_list: + assert call.kwargs["data"]["values_b64"] not in ( + good_one.values_b64, + good_two.values_b64, + ) + + logger.info.assert_called_once() + info_args = logger.info.call_args.args + assert info_args[1] == 2 # rotated + assert info_args[2] == 1 # skipped + + def test_decrypt_global_env_var_drops_undecryptable_value( env_vars_salt_key, monkeypatch ):