mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
77ddd83c06
commit
43394c198d
4 changed files with 239 additions and 21 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue