diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 50c96ce45ea..ed61ceac45c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -18568,20 +18568,34 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( @pytest.mark.asyncio -async def test_rotate_master_key_reencrypts_vector_store_litellm_params(): - """Master-key rotation re-encrypts the managed vector store credentials (step 4e).""" +async def test_rotate_master_key_reencrypts_vector_store_litellm_params(monkeypatch): + """Master-key rotation re-encrypts the managed vector store credentials (step 4e) under the new key.""" + import json from unittest.mock import AsyncMock, MagicMock + from litellm.proxy import proxy_server from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with from litellm.proxy.management_endpoints.key_management_endpoints import ( _rotate_master_key, ) + from litellm.proxy.vector_store_endpoints.litellm_params_encryption import encrypt_vector_store_litellm_params + + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(proxy_server, "master_key", "sk-old-master-key") + stored = encrypt_vector_store_litellm_params({"api_key": "sk-vs-secret"}) mock_prisma_client = AsyncMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[]) + vector_stores = mock_prisma_client.db.litellm_managedvectorstorestable + vector_stores.find_many = AsyncMock(return_value=[{"vector_store_id": "vs_1", "litellm_params": stored}]) + vector_stores.update_many = AsyncMock() + mock_prisma_client.tx = MagicMock() + mock_prisma_client.tx.return_value.__aenter__.return_value = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.proxy_config", MagicMock()), @@ -18589,10 +18603,6 @@ async def test_rotate_master_key_reencrypts_vector_store_litellm_params(): patch("litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_user_credentials_master_key"), patch("litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_user_env_vars_master_key"), patch("litellm.proxy.management_endpoints.key_management_endpoints.rotate_sso_identity_assertions_master_key"), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.reencrypt_vector_store_litellm_params", - new_callable=AsyncMock, - ) as mock_reencrypt_vector_stores, ): await _rotate_master_key( prisma_client=mock_prisma_client, @@ -18601,10 +18611,10 @@ async def test_rotate_master_key_reencrypts_vector_store_litellm_params(): new_master_key="sk-new-master-key", ) - mock_reencrypt_vector_stores.assert_awaited_once_with( - prisma_client=mock_prisma_client, - new_master_key="sk-new-master-key", - ) + written = json.loads(vector_stores.update_many.await_args.kwargs["data"]["litellm_params"]) + ciphertext = written["api_key"].removeprefix("litellm_enc::") + assert decrypt_if_encrypted_with(ciphertext, "sk-new-master-key") == "sk-vs-secret" + assert decrypt_if_encrypted_with(ciphertext, "sk-old-master-key") is None @pytest.mark.asyncio