fix(mcp): rotate new encrypted env var columns on master-key rotation

Master-key rotation re-encrypted only the credentials column and the
litellm_mcpusercredentials table, leaving the new global env_vars values
and the litellm_mcpuserenvvars values_b64 column encrypted under the old
key. After a rotation those values fail to decrypt, so global ${VAR}
headers are forwarded as empty substitutions and every per-user value
reads back as missing (412). Re-encrypt both new columns alongside the
existing ones, skipping undecryptable entries so a corrupt row is
preserved rather than overwritten.
This commit is contained in:
mateo-berri 2026-06-04 13:41:17 +00:00
parent 77ddd83c06
commit 43394c198d
No known key found for this signature in database
4 changed files with 239 additions and 21 deletions

View file

@ -90,6 +90,47 @@ def decrypt_global_env_var_values(env_vars: Optional[Iterable[Any]]) -> None:
entry.value = decrypted
def _reencrypt_global_env_var_values(
env_vars: Optional[Iterable[Any]], new_encryption_key: str
) -> Optional[List[Dict[str, Any]]]:
"""Re-encrypt ``scope="global"`` env var values for master-key rotation.
Each global value is decrypted with the current salt key and re-encrypted
under ``new_encryption_key``. Returns the rebuilt list when at least one
value was rotated, else ``None`` so the caller can skip the DB write. A
value that fails to decrypt is left untouched (and logged) so a corrupt
entry is preserved for recovery rather than overwritten.
"""
if not env_vars:
return None
rebuilt = [dict(v) for v in env_vars]
rotated = False
for entry in rebuilt:
if not _is_global_env_var_scope(entry.get("scope")):
continue
value = entry.get("value")
if not value:
continue
decrypted = decrypt_value_helper(
value=value,
key="mcp_global_env_var",
exception_type="debug",
return_original_value=False,
)
if decrypted is None:
verbose_proxy_logger.warning(
"rotate_mcp_server_credentials_master_key: could not decrypt "
"global env var %s, skipping",
entry.get("name"),
)
continue
entry["value"] = encrypt_value_helper(
decrypted, new_encryption_key=new_encryption_key
)
rotated = True
return rebuilt if rotated else None
def _prepare_mcp_server_data(
data: Union[NewMCPServerRequest, UpdateMCPServerRequest],
exclude_unset: bool = False,
@ -600,33 +641,38 @@ async def update_mcp_server(
async def rotate_mcp_server_credentials_master_key(
prisma_client: PrismaClient, touched_by: str, new_master_key: str
):
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
mcp_servers = await prisma_client.db.litellm_mcpservertable.find_many()
for mcp_server in mcp_servers:
update_data: Dict[str, Any] = {}
credentials = mcp_server.credentials
if not credentials:
if credentials:
# Decrypt with current key first, then re-encrypt with new key
decrypted_credentials = decrypt_credentials(
credentials=cast(MCPCredentials, dict(credentials)),
)
encrypted_credentials = encrypt_credentials(
credentials=decrypted_credentials,
encryption_key=new_master_key,
)
update_data["credentials"] = safe_dumps(encrypted_credentials)
rotated_env_vars = _reencrypt_global_env_var_values(
mcp_server.env_vars, new_master_key
)
if rotated_env_vars is not None:
update_data["env_vars"] = safe_dumps(rotated_env_vars)
if not update_data:
continue
credentials_copy = dict(credentials)
# Decrypt with current key first, then re-encrypt with new key
decrypted_credentials = decrypt_credentials(
credentials=cast(MCPCredentials, credentials_copy),
)
encrypted_credentials = encrypt_credentials(
credentials=decrypted_credentials,
encryption_key=new_master_key,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
serialized_credentials = safe_dumps(encrypted_credentials)
update_data["updated_by"] = touched_by
await prisma_client.db.litellm_mcpservertable.update(
where={"server_id": mcp_server.server_id},
data={
"credentials": serialized_credentials,
"updated_by": touched_by,
},
data=update_data,
)
@ -706,6 +752,46 @@ async def rotate_mcp_user_credentials_master_key(
)
async def rotate_mcp_user_env_vars_master_key(
prisma_client: PrismaClient, new_master_key: str
):
"""Re-encrypt every ``LiteLLM_MCPUserEnvVars`` row with ``new_master_key``.
Reads each ``values_b64`` blob with the current salt key and writes it back
encrypted under the new master key. Rows that fail to decrypt are logged and
skipped so one corrupt row does not abort the rotation nor overwrite values
that may still be recoverable.
"""
rows = await prisma_client.db.litellm_mcpuserenvvars.find_many()
for row in rows:
plaintext = decrypt_value_helper(
value=row.values_b64,
key="mcp_user_env_vars",
exception_type="debug",
return_original_value=False,
)
if plaintext is None:
verbose_proxy_logger.warning(
"rotate_mcp_user_env_vars_master_key: could not decrypt env vars "
"for user_id=%s server_id=%s, skipping",
row.user_id,
row.server_id,
)
continue
re_encrypted = encrypt_value_helper(
plaintext, new_encryption_key=new_master_key
)
await prisma_client.db.litellm_mcpuserenvvars.update(
where={
"user_id_server_id": {
"user_id": row.user_id,
"server_id": row.server_id,
}
},
data={"values_b64": re_encrypted},
)
async def store_user_credential(
prisma_client: PrismaClient,
user_id: str,

View file

@ -39,6 +39,7 @@ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._experimental.mcp_server.db import (
rotate_mcp_server_credentials_master_key,
rotate_mcp_user_credentials_master_key,
rotate_mcp_user_env_vars_master_key,
)
from litellm.proxy._types import *
from litellm.proxy._types import LiteLLM_VerificationToken
@ -4136,6 +4137,15 @@ async def _rotate_master_key( # noqa: PLR0915
"Failed to rotate MCP user credentials: %s", str(e)
)
# 4c. process MCP per-user environment variables table
try:
await rotate_mcp_user_env_vars_master_key(
prisma_client=prisma_client,
new_master_key=new_master_key,
)
except Exception as e:
verbose_proxy_logger.warning("Failed to rotate MCP user env vars: %s", str(e))
# 5. process credentials table
try:
credentials = await prisma_client.db.litellm_credentialstable.find_many()

View file

@ -20,11 +20,14 @@ from litellm.proxy._experimental.mcp_server.db import (
get_user_oauth_credential,
list_user_oauth_credentials,
rotate_mcp_user_credentials_master_key,
rotate_mcp_user_env_vars_master_key,
store_user_credential,
store_user_oauth_credential,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
SALT_KEY = "test-salt-key-for-byok-credential-tests-1234"
@ -400,3 +403,69 @@ async def test_rotate_skips_undecodable_rows():
assert prisma.db.litellm_mcpusercredentials.update.call_count == 1
where = prisma.db.litellm_mcpusercredentials.update.call_args.kwargs["where"]
assert where["user_id_server_id"]["server_id"] == "srv-ok"
# ── per-user env-var rotation ─────────────────────────────────────────────────
def _env_var_row(values_b64: str, user_id="alice", server_id="srv-1"):
row = MagicMock()
row.values_b64 = values_b64
row.user_id = user_id
row.server_id = server_id
return row
@pytest.mark.asyncio
async def test_rotate_user_env_vars_re_encrypts_with_new_key(monkeypatch):
# Encrypt env vars under the current salt, rotate to a new key, then confirm
# the stored ciphertext round-trips under the NEW key.
values = {"API_KEY": "sk-secret", "REGION": "us-east-1"}
encrypted_old = encrypt_value_helper(json.dumps(values))
prisma = MagicMock()
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(
return_value=[_env_var_row(encrypted_old)]
)
prisma.db.litellm_mcpuserenvvars.update = AsyncMock()
new_master_key = "rotated-env-key-1111-2222-3333-4444"
await rotate_mcp_user_env_vars_master_key(
prisma_client=prisma, new_master_key=new_master_key
)
new_stored = prisma.db.litellm_mcpuserenvvars.update.call_args.kwargs["data"][
"values_b64"
]
assert new_stored != encrypted_old, "rotation must produce different ciphertext"
monkeypatch.setenv("LITELLM_SALT_KEY", new_master_key)
decrypted = decrypt_value_helper(
value=new_stored,
key="mcp_user_env_vars",
exception_type="debug",
return_original_value=False,
)
assert json.loads(decrypted) == values
@pytest.mark.asyncio
async def test_rotate_user_env_vars_skips_undecryptable_rows():
# A corrupt row must be skipped (not overwritten) so recoverable data is
# preserved and one bad row does not abort the rest of the rotation.
good = _env_var_row(
encrypt_value_helper(json.dumps({"A": "1"})), server_id="srv-ok"
)
bad = _env_var_row("!!! not encrypted !!!", server_id="srv-corrupt")
prisma = MagicMock()
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[bad, good])
prisma.db.litellm_mcpuserenvvars.update = AsyncMock()
await rotate_mcp_user_env_vars_master_key(
prisma_client=prisma, new_master_key="new-key-xxxx"
)
assert prisma.db.litellm_mcpuserenvvars.update.call_count == 1
where = prisma.db.litellm_mcpuserenvvars.update.call_args.kwargs["where"]
assert where["user_id_server_id"]["server_id"] == "srv-ok"

View file

@ -1007,6 +1007,7 @@ class TestRotateCredentials:
"aws_secret_access_key": "enc_old:SAK",
"aws_region_name": "us-east-1",
}
server.env_vars = None
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
@ -1043,6 +1044,58 @@ class TestRotateCredentials:
# Non-secret fields should pass through unchanged
assert stored_creds["aws_region_name"] == "us-east-1"
@pytest.mark.asyncio
async def test_rotation_reencrypts_global_env_vars(self):
"""Global env var values are re-encrypted under the new key; user-scope
placeholders are left untouched."""
from litellm.proxy._experimental.mcp_server.db import (
rotate_mcp_server_credentials_master_key,
)
server = MagicMock()
server.server_id = "srv-env"
server.credentials = None
server.env_vars = [
{"name": "API_KEY", "value": "enc_old:secret", "scope": "global"},
{"name": "USER_TOKEN", "value": "", "scope": "user"},
]
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[server]
)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
with (
patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value="old-key",
),
patch(
"litellm.proxy._experimental.mcp_server.db.decrypt_value_helper",
side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace(
"enc_old:", ""
),
),
patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: f"enc_new:{value}",
),
):
await rotate_mcp_server_credentials_master_key(
mock_prisma, "admin", "new-key"
)
update_call = mock_prisma.db.litellm_mcpservertable.update
assert update_call.called
stored_env = json.loads(update_call.call_args[1]["data"]["env_vars"])
# Global value decrypted from old, then re-encrypted with new key
assert stored_env[0]["value"] == "enc_new:secret"
# User-scope placeholder untouched
assert stored_env[1]["value"] == ""
# Credentials column not written when the server has none
assert "credentials" not in update_call.call_args[1]["data"]
class TestAuthTypeSwitchClearsCredentials:
"""Test that switching auth_type without credentials clears stale secrets."""