diff --git a/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py b/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py index 96d3956e158..64a35b1995e 100644 --- a/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py +++ b/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py @@ -2,8 +2,11 @@ import os from collections.abc import Callable, Mapping +from datetime import timedelta from typing import TYPE_CHECKING, Final +from typing_extensions import TypeIs + from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -19,6 +22,14 @@ if TYPE_CHECKING: VECTOR_STORE_SECRET_PARAM_MASKER: Final = SensitiveDataMasker(extra_sensitive_patterns=frozenset(("connection",))) +def _is_json_object(value: object) -> TypeIs[dict[str, object]]: # guard-ok: trivial isinstance; JSON keys are str + return isinstance(value, dict) + + +def _is_json_array(value: object) -> TypeIs[list[object]]: # guard-ok: trivial isinstance narrowing + return isinstance(value, list) + + def _map_litellm_param_strings( litellm_params: Mapping[str, object], transform: Callable[[str, str, bool], str], @@ -27,15 +38,20 @@ def _map_litellm_param_strings( ) -> dict[str, object]: """Apply ``transform(key, value, is_secret)`` to every string in a nested ``litellm_params`` dict. - ``is_secret`` is true for a secret-named key and for every string nested under one. + ``is_secret`` is true for a secret-named key and for every string nested under one. A list item is passed + with the key of the list that holds it. """ mapped: Final[dict[str, object]] = {} for key, value in litellm_params.items(): is_secret = parent_is_secret or VECTOR_STORE_SECRET_PARAM_MASKER.is_sensitive_key(key) if isinstance(value, str): mapped[key] = transform(key, value, is_secret) - elif isinstance(value, dict) and depth < DEFAULT_MAX_RECURSE_DEPTH: + elif _is_json_object(value) and depth < DEFAULT_MAX_RECURSE_DEPTH: mapped[key] = _map_litellm_param_strings(value, transform, is_secret, depth + 1) + elif _is_json_array(value) and depth < DEFAULT_MAX_RECURSE_DEPTH: + mapped[key] = [ + _map_litellm_param_strings({key: item}, transform, parent_is_secret, depth + 1)[key] for item in value + ] else: mapped[key] = value return mapped @@ -102,7 +118,8 @@ async def reencrypt_vector_store_litellm_params(prisma_client: "PrismaClient", n """Re-encrypt every stored vector store's encrypted ``litellm_params`` values for a master-key rotation. Values are re-encrypted under ``LITELLM_SALT_KEY`` when it is set, else under ``new_master_key``. Plaintext - values and values that do not decrypt are left as stored. Returns the number of rows rewritten. + values and values that do not decrypt are left as stored. All rows are rewritten in one transaction. Returns the + number of rows rewritten. """ salt_key: Final = os.getenv("LITELLM_SALT_KEY") new_encryption_key: Final = new_master_key if salt_key is None else salt_key @@ -117,7 +134,7 @@ async def reencrypt_vector_store_litellm_params(prisma_client: "PrismaClient", n return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(plaintext, new_encryption_key=new_encryption_key) table: Final = ManagedVectorStoresRepository(prisma_client).table - rewritten = 0 + rewrites: Final[dict[str, dict[str, object]]] = {} for row in await table.find_many(): stored = LiteLLM_ManagedVectorStore(**dict(row)) litellm_params = stored.get("litellm_params") @@ -131,11 +148,15 @@ async def reencrypt_vector_store_litellm_params(prisma_client: "PrismaClient", n stored.get("vector_store_id"), sorted(undecryptable_keys), ) - if reencrypted == litellm_params: + vector_store_id = stored.get("vector_store_id") + if reencrypted == litellm_params or vector_store_id is None: continue - await table.update( - where={"vector_store_id": stored.get("vector_store_id")}, - data={"litellm_params": safe_dumps(reencrypted)}, - ) - rewritten += 1 - return rewritten + rewrites[vector_store_id] = reencrypted + if rewrites: + async with prisma_client.tx(timeout=timedelta(minutes=2)) as tx: + for vector_store_id, reencrypted in rewrites.items(): + await tx.litellm_managedvectorstorestable.update_many( + where={"vector_store_id": vector_store_id}, + data={"litellm_params": safe_dumps(reencrypted)}, + ) + return len(rewrites) diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 9d8644fb195..d116add8f4f 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -62,6 +62,26 @@ def _row_to_vector_store(row: "_VectorStoreRow") -> LiteLLM_ManagedVectorStore: return decrypt_vector_store_litellm_params(LiteLLM_ManagedVectorStore(**row.model_dump())) +def _registry_copy(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_ManagedVectorStore | None: + """The copy of a decrypted DB ``vector_store`` to put in the in-memory registry. + + When a secret in it does not decrypt with this process's key (the master key was rotated and the proxy not + restarted yet), the registry's current ``litellm_params`` are kept, and None is returned if it has none. + """ + if not holds_undecrypted_secret(vector_store): + return vector_store + registry: Final = litellm.vector_store_registry + vector_store_id: Final = vector_store.get("vector_store_id") + current: Final = ( + registry.get_litellm_managed_vector_store_from_registry(vector_store_id) + if registry is not None and vector_store_id is not None + else None + ) + if current is None: + return None + return vector_store | LiteLLM_ManagedVectorStore(litellm_params=current.get("litellm_params", {})) + + class _ConfigOwnedDetail(TypedDict): error: ReadOnly[str] vector_store_id: ReadOnly[str] @@ -419,14 +439,13 @@ async def list_vector_stores( litellm.vector_store_registry.delete_vector_store_from_registry(vector_store_id=vs_id) verbose_proxy_logger.debug("Removed deleted vector store %s from in-memory registry", vs_id) - # 2. Update in-memory registry with database versions (for updates). A row whose secret this - # process cannot decrypt (the master key was rotated and the proxy not yet restarted) keeps - # the registry's working copy. + # 2. Update in-memory registry with database versions (for updates) for vector_store in vector_stores_from_db: vector_store_id = vector_store.get("vector_store_id", None) - if vector_store_id and not holds_undecrypted_secret(vector_store): + registry_copy = _registry_copy(vector_store) + if vector_store_id and registry_copy is not None: litellm.vector_store_registry.update_vector_store_in_registry( - vector_store_id=vector_store_id, updated_data=vector_store + vector_store_id=vector_store_id, updated_data=registry_copy ) # Filter vector stores based on access control @@ -671,10 +690,11 @@ async def update_vector_store( updated_vs: Final = _row_to_vector_store(updated) # Immediately update in-memory registry to keep it in sync - if litellm.vector_store_registry is not None and not holds_undecrypted_secret(updated_vs): + registry_copy: Final = _registry_copy(updated_vs) + if litellm.vector_store_registry is not None and registry_copy is not None: litellm.vector_store_registry.update_vector_store_in_registry( vector_store_id=vector_store_id, - updated_data=updated_vs, + updated_data=registry_copy, ) verbose_proxy_logger.debug( "Updated vector store %s in both database and in-memory registry", vector_store_id diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_litellm_params_encryption.py b/tests/test_litellm/proxy/vector_store_endpoints/test_litellm_params_encryption.py index eb4278f183e..aaeba1000d2 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_litellm_params_encryption.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_litellm_params_encryption.py @@ -142,7 +142,8 @@ def _rotation_rows(encrypted_params): ] prisma_client = MagicMock() prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=rows) - prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock() + prisma_client.db.litellm_managedvectorstorestable.update_many = AsyncMock() + prisma_client.tx.return_value.__aenter__.return_value = prisma_client.db return prisma_client @@ -159,8 +160,8 @@ async def test_reencrypt_moves_encrypted_rows_to_the_new_master_key_and_leaves_p rewritten = await reencrypt_vector_store_litellm_params(prisma_client=prisma_client, new_master_key=new_key) assert rewritten == 1 - prisma_client.db.litellm_managedvectorstorestable.update.assert_awaited_once() - call = prisma_client.db.litellm_managedvectorstorestable.update.await_args + prisma_client.db.litellm_managedvectorstorestable.update_many.assert_awaited_once() + call = prisma_client.db.litellm_managedvectorstorestable.update_many.await_args assert call.kwargs["where"] == {"vector_store_id": "vs_new"} stored = json.loads(call.kwargs["data"]["litellm_params"]) assert _decrypted_under(stored["api_key"], new_key) == "sk-new-secret" @@ -175,7 +176,7 @@ async def test_reencrypt_keeps_values_under_the_salt_key_when_one_is_set(salt_ke await reencrypt_vector_store_litellm_params(prisma_client=prisma_client, new_master_key="sk-rotated-master-key") stored = json.loads( - prisma_client.db.litellm_managedvectorstorestable.update.await_args.kwargs["data"]["litellm_params"] + prisma_client.db.litellm_managedvectorstorestable.update_many.await_args.kwargs["data"]["litellm_params"] ) assert _decrypted_under(stored["api_key"], _SALT_KEY) == "sk-new-secret" @@ -191,7 +192,7 @@ async def test_reencrypt_warns_about_values_it_cannot_decrypt(salt_key, monkeypa rewritten = await reencrypt_vector_store_litellm_params(prisma_client=prisma_client, new_master_key="sk-new") assert rewritten == 0 - prisma_client.db.litellm_managedvectorstorestable.update.assert_not_awaited() + prisma_client.db.litellm_managedvectorstorestable.update_many.assert_not_awaited() warning.assert_called_once() assert "vs_new" in warning.call_args.args assert "sk-lost-secret" not in str(warning.call_args) @@ -206,7 +207,7 @@ async def test_reencrypt_treats_an_empty_salt_key_as_set(monkeypatch): await reencrypt_vector_store_litellm_params(prisma_client=prisma_client, new_master_key="sk-rotated-master-key") stored = json.loads( - prisma_client.db.litellm_managedvectorstorestable.update.await_args.kwargs["data"]["litellm_params"] + prisma_client.db.litellm_managedvectorstorestable.update_many.await_args.kwargs["data"]["litellm_params"] ) assert decrypt_vector_store_litellm_params(LiteLLM_ManagedVectorStore(vector_store_id="vs", litellm_params=stored))[ "litellm_params" @@ -227,3 +228,19 @@ def test_holds_undecrypted_secret(salt_key, monkeypatch): assert holds(readable) is False assert holds({"api_key": "sk-plain", "api_base": "litellm_enc::not-a-secret-key"}) is False assert holds(unreadable_nested) is True + + +def test_secrets_inside_lists_are_encrypted_and_decrypted(salt_key): + params = { + "api_key": ["sk-first", "sk-second"], + "extra": [{"api_key": "sk-in-list"}, "not-a-secret"], + } + + encrypted = encrypt_vector_store_litellm_params(params) + + assert all(_decrypted_under(value, _SALT_KEY) for value in encrypted["api_key"]) + assert _decrypted_under(encrypted["extra"][0]["api_key"], _SALT_KEY) == "sk-in-list" + assert encrypted["extra"][1] == "not-a-secret" + assert "sk-first" not in json.dumps(encrypted) + store = LiteLLM_ManagedVectorStore(vector_store_id="vs", litellm_params=encrypted) + assert decrypt_vector_store_litellm_params(store)["litellm_params"] == params diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index f87718b9bfb..daf5aef8e3a 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -3396,8 +3396,13 @@ class TestLitellmParamsEncryptedAtRest: rows = [ { "vector_store_id": "vs_rotated", + "vector_store_description": "updated in the db", "litellm_params": self._encrypted_under_another_key({"api_key": "sk-working"}, monkeypatch), }, + { + "vector_store_id": "vs_not_loaded", + "litellm_params": self._encrypted_under_another_key({"api_key": "sk-other"}, monkeypatch), + }, { "vector_store_id": "vs_readable", "litellm_params": encrypt_vector_store_litellm_params({"api_key": "sk-new"}), @@ -3423,19 +3428,29 @@ class TestLitellmParamsEncryptedAtRest: params_by_id = {vs["vector_store_id"]: vs["litellm_params"] for vs in registry.vector_stores} assert params_by_id == {"vs_rotated": {"api_key": "sk-working"}, "vs_readable": {"api_key": "sk-new"}} + rotated = registry.get_litellm_managed_vector_store_from_registry("vs_rotated") + assert rotated is not None + assert rotated["vector_store_description"] == "updated in the db" @pytest.mark.asyncio async def test_update_keeps_the_registry_copy_of_a_row_this_proxy_cannot_decrypt(self, monkeypatch): from litellm.proxy.vector_store_endpoints.management_endpoints import update_vector_store from litellm.types.vector_stores import VectorStoreUpdateRequest + from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + stored_params = self._encrypted_under_another_key({"api_key": "sk-vs-secret-9"}, monkeypatch) - row = self._row({"vector_store_id": "vs_rotated", "litellm_params": stored_params}) + row = self._row( + {"vector_store_id": "vs_rotated", "vector_store_description": "new", "litellm_params": stored_params} + ) prisma_client = MagicMock() prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=row) prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock(return_value=row) - registry = MagicMock() - registry.is_config_vector_store.return_value = False + registry = VectorStoreRegistry( + vector_stores=[ + LiteLLM_ManagedVectorStore(vector_store_id="vs_rotated", litellm_params={"api_key": "sk-vs-secret-9"}) + ] + ) with ( patch( @@ -3455,7 +3470,13 @@ class TestLitellmParamsEncryptedAtRest: user_api_key_dict=UserAPIKeyAuth(user_id="admin"), ) - registry.update_vector_store_in_registry.assert_not_called() + assert registry.vector_stores == [ + { + "vector_store_id": "vs_rotated", + "vector_store_description": "new", + "litellm_params": {"api_key": "sk-vs-secret-9"}, + } + ] @pytest.mark.asyncio async def test_store_resolved_from_the_shared_cache_has_decrypted_params(self):