diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index a7796bc7a7c..41072a63659 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -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, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index c8c590af97c..99d0bac88af 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index 078adf72d4c..8aa9f107c8f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -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" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index 0c2a8bb8087..c2164a9f19f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -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."""