mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
chore(mcp): log master-key rotation summary counts
Master-key rotation is a rare, high-stakes batch operation that previously ran silently on success; only per-row decrypt failures were logged. Each of the three MCP rotation steps (server credentials/global env vars, per-user credentials, per-user env vars) now emits one info line summarising how many rows were rotated and skipped, so an operator can confirm the step ran and sanity-check the counts after rotating the key. Regression test asserts the per-user env-var summary reports rotated and skipped counts tied to real work: a decryptable row is re-encrypted and counted as rotated while an undecryptable row is left untouched and counted as skipped
This commit is contained in:
parent
faca3973ce
commit
4d7f160d25
2 changed files with 88 additions and 0 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue