diff --git a/litellm/proxy/db/master_key_migration.py b/litellm/proxy/db/master_key_migration.py index d100554201a..4d8defd51f3 100644 --- a/litellm/proxy/db/master_key_migration.py +++ b/litellm/proxy/db/master_key_migration.py @@ -43,6 +43,9 @@ _SECRET_COLUMNS: Final = ( _SecretColumn("LiteLLM_UserTable", "user_id", "metadata", only_rows_with_marked_ciphertexts=True), _SecretColumn("LiteLLM_DeletedTeamTable", "id", "metadata", only_rows_with_marked_ciphertexts=True), _SecretColumn("LiteLLM_DeletedVerificationToken", "id", "metadata", only_rows_with_marked_ciphertexts=True), + _SecretColumn( + "LiteLLM_ManagedVectorStoresTable", "vector_store_id", "litellm_params", only_rows_with_marked_ciphertexts=True + ), ) _STORED_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index d37dfe87ad5..03c79eac593 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -133,6 +133,7 @@ from litellm.proxy.utils import ( handle_exception_on_proxy, is_valid_api_key, ) +from litellm.proxy.vector_store_endpoints.litellm_params_encryption import reencrypt_vector_store_litellm_params from litellm.repositories.base_repository import BaseRepository from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.config_repository import ConfigParam, ConfigRepository @@ -5340,6 +5341,12 @@ async def _rotate_master_key( except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation verbose_proxy_logger.warning("Failed to rotate SSO identity assertions: %s", str(e)) + # 4e. process managed vector store table + try: + await reencrypt_vector_store_litellm_params(prisma_client=prisma_client, new_master_key=new_master_key) + except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation + verbose_proxy_logger.warning("Failed to rotate vector store credentials: %s", str(e)) + # 5. process credentials table try: credentials = await _credentials_table(prisma_client).find_many() diff --git a/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py b/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py new file mode 100644 index 00000000000..4f54ca291b9 --- /dev/null +++ b/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py @@ -0,0 +1,203 @@ +"""Encryption at rest for the secret values in a managed vector store's ``litellm_params``.""" + +import os +from collections.abc import Callable, Iterator, Mapping +from datetime import timedelta +from itertools import chain +from types import MappingProxyType +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 +from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker +from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.repositories.table_repositories import ManagedVectorStoresRepository +from litellm.types.vector_stores import LiteLLM_ManagedVectorStore + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +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], + parent_is_secret: bool = False, + depth: int = 0, +) -> 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. A list item is passed + with the key of the list that holds it. + """ + return { + key: _map_param_value( + key, + value, + transform, + parent_is_secret or VECTOR_STORE_SECRET_PARAM_MASKER.is_sensitive_key(key), + depth, + ) + for key, value in litellm_params.items() + } + + +def _map_param_value( + key: str, value: object, transform: Callable[[str, str, bool], str], is_secret: bool, depth: int +) -> object: + if isinstance(value, str): + return transform(key, value, is_secret) + if depth >= DEFAULT_MAX_RECURSE_DEPTH: + return value + if _is_json_object(value): + return _map_litellm_param_strings(value, transform, is_secret, depth + 1) + if _is_json_array(value): + return [_map_param_value(key, item, transform, is_secret, depth + 1) for item in value] + return value + + +def _secret_strings( + litellm_params: Mapping[str, object], parent_is_secret: bool = False, depth: int = 0 +) -> Iterator[tuple[str, str]]: + """Every ``(key, value)`` string that ``_map_litellm_param_strings`` would pass with ``is_secret`` true.""" + return chain.from_iterable( + _secret_strings_in( + key, value, parent_is_secret or VECTOR_STORE_SECRET_PARAM_MASKER.is_sensitive_key(key), depth + ) + for key, value in litellm_params.items() + ) + + +def _secret_strings_in(key: str, value: object, is_secret: bool, depth: int) -> Iterator[tuple[str, str]]: + if isinstance(value, str): + return iter(((key, value),) if is_secret else ()) + if depth >= DEFAULT_MAX_RECURSE_DEPTH: + return iter(()) + if _is_json_object(value): + return _secret_strings(value, is_secret, depth + 1) + if _is_json_array(value): + return chain.from_iterable(_secret_strings_in(key, item, is_secret, depth + 1) for item in value) + return iter(()) + + +def encrypt_vector_store_litellm_params(litellm_params: Mapping[str, object]) -> dict[str, object]: + """Encrypt the secret values of a vector store's ``litellm_params`` for storage in the database.""" + + def encrypt(key: str, value: str, is_secret: bool) -> str: + if not is_secret or not value: + return value + try: + return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(value) + except Exception as e: # noqa: BLE001 # no master key or salt key configured, so the value is stored as given + verbose_proxy_logger.warning( + "Vector store litellm_params[%r] stored unencrypted: %s", key, type(e).__name__ + ) + return value + + return _map_litellm_param_strings(litellm_params, encrypt) + + +def decrypt_vector_store_litellm_params(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_ManagedVectorStore: + """Return ``vector_store`` with the encrypted values of its ``litellm_params`` decrypted; plaintext is kept.""" + + def decrypt(key: str, value: str, is_secret: bool) -> str: + decrypted: Final = _decrypt_secret_value(key, value) if is_secret else None + return value if decrypted is None else decrypted + + litellm_params: Final = vector_store.get("litellm_params") + if not isinstance(litellm_params, dict): + return vector_store + return vector_store | LiteLLM_ManagedVectorStore(litellm_params=_map_litellm_param_strings(litellm_params, decrypt)) + + +def holds_undecrypted_secret(vector_store: LiteLLM_ManagedVectorStore) -> bool: + """Whether a decrypted ``vector_store`` still has an encrypted secret value, one this proxy's key cannot read.""" + litellm_params: Final = vector_store.get("litellm_params") + return isinstance(litellm_params, dict) and any( + value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) for _, value in _secret_strings(litellm_params) + ) + + +def _decrypt_secret_value(key: str, value: str) -> str | None: + """The plaintext of an encrypted ``litellm_params`` value, or None when ``value`` is not one this proxy can read.""" + if not value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): + return None + return decrypt_value_helper( + value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), + key=key, + exception_type="debug", + return_original_value=False, + ) + + +async def reencrypt_vector_store_litellm_params(prisma_client: "PrismaClient", new_master_key: str) -> int: + """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. 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 + + def reencrypt(key: str, value: str, is_secret: bool) -> str: + plaintext: Final = _decrypt_secret_value(key, value) if is_secret else None + if plaintext is None: + return value + return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(plaintext, new_encryption_key=new_encryption_key) + + stored_rows: Final = tuple( + LiteLLM_ManagedVectorStore(**dict(row)) + for row in await ManagedVectorStoresRepository(prisma_client).table.find_many() + ) + for stored in stored_rows: + _warn_about_undecryptable_values(stored) + rewrites: Final = MappingProxyType( + { + vector_store_id: reencrypted + for stored in stored_rows + if isinstance(litellm_params := stored.get("litellm_params"), dict) + and (vector_store_id := stored.get("vector_store_id")) is not None + and (reencrypted := _map_litellm_param_strings(litellm_params, reencrypt)) != litellm_params + } + ) + 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) + + +def _warn_about_undecryptable_values(vector_store: LiteLLM_ManagedVectorStore) -> None: + litellm_params: Final = vector_store.get("litellm_params") + undecryptable_keys: Final = ( + sorted( + key + for key, value in _secret_strings(litellm_params) + if value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) and _decrypt_secret_value(key, value) is None + ) + if isinstance(litellm_params, dict) + else [] + ) + if undecryptable_keys: + verbose_proxy_logger.warning( + "Vector store %s: litellm_params %s do not decrypt with the current key and were not re-encrypted", + vector_store.get("vector_store_id"), + undecryptable_keys, + ) diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index cae144bb266..d116add8f4f 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -23,7 +23,6 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.constants import REDACTED_BY_LITELM_STRING from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._types import ( LiteLLM_ManagedVectorStoresTable, ResponseLiteLLM_ManagedVectorStore, @@ -31,6 +30,12 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user +from litellm.proxy.vector_store_endpoints.litellm_params_encryption import ( + VECTOR_STORE_SECRET_PARAM_MASKER, + decrypt_vector_store_litellm_params, + encrypt_vector_store_litellm_params, + holds_undecrypted_secret, +) from litellm.proxy.vector_store_endpoints.utils import ( can_user_access_vector_store, filter_listable_vector_stores, @@ -54,7 +59,27 @@ def _vector_store_table(prisma_client: "PrismaClient") -> "TableActions[_VectorS def _row_to_vector_store(row: "_VectorStoreRow") -> LiteLLM_ManagedVectorStore: - return LiteLLM_ManagedVectorStore(**row.model_dump()) + 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): @@ -83,9 +108,6 @@ def _with_ownership(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_Managed return vector_store | ownership -_LITELLM_PARAMS_MASKER: Final = SensitiveDataMasker(extra_sensitive_patterns=frozenset(("connection",))) - - _REDACT_LITELLM_PARAMS_MAX_DEPTH: Final = 10 @@ -123,7 +145,7 @@ def _redact_sensitive_litellm_params(litellm_params: object, _depth: int = 0) -> return litellm_params out: Final[dict[str, object]] = {} for k, v in litellm_params.items(): - if _LITELLM_PARAMS_MASKER.is_sensitive_key(k): + if VECTOR_STORE_SECRET_PARAM_MASKER.is_sensitive_key(k): out[k] = REDACTED_BY_LITELM_STRING elif isinstance(v, dict): out[k] = _redact_sensitive_litellm_params(v, _depth + 1) @@ -247,7 +269,7 @@ async def create_vector_store_in_db( # on the deployment and never reach the database. if litellm_params: litellm_params_dict: Final = GenericLiteLLMParams(**litellm_params).model_dump(exclude_none=True) - data_to_create["litellm_params"] = safe_dumps(litellm_params_dict) + data_to_create["litellm_params"] = safe_dumps(encrypt_vector_store_litellm_params(litellm_params_dict)) else: # Provide empty dict if no litellm_params provided data_to_create["litellm_params"] = safe_dumps({}) @@ -420,9 +442,10 @@ async def list_vector_stores( # 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: + 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 @@ -650,7 +673,7 @@ async def update_vector_store( if "litellm_params" in update_data: _input_litellm_params: Final[dict] = update_data.get("litellm_params", {}) or {} litellm_params_dict: Final = GenericLiteLLMParams(**_input_litellm_params).model_dump(exclude_none=True) - update_data["litellm_params"] = safe_dumps(litellm_params_dict) + update_data["litellm_params"] = safe_dumps(encrypt_vector_store_litellm_params(litellm_params_dict)) # Update in database updated: Final = await _vector_store_table(prisma_client).update( @@ -667,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: + 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/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 8363aaee99a..ea434802704 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -17,6 +17,7 @@ from litellm.proxy._types import ( LitellmUserRoles, UserAPIKeyAuth, ) +from litellm.proxy.vector_store_endpoints.litellm_params_encryption import decrypt_vector_store_litellm_params from litellm.types.utils import LlmProviders from litellm.types.vector_stores import LiteLLM_ManagedVectorStore from litellm.utils import ProviderConfigManager @@ -308,7 +309,9 @@ async def get_litellm_managed_vector_store( ) if not rows: return None - return _normalize_litellm_params(LiteLLM_ManagedVectorStore(**rows[0].model_dump())) + return decrypt_vector_store_litellm_params( + _normalize_litellm_params(LiteLLM_ManagedVectorStore(**rows[0].model_dump())) + ) except Exception as e: verbose_proxy_logger.warning( "Failed to resolve vector store id=%s from shared cache: %s", diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index e78d2aa5f6a..2a088d38484 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -501,13 +501,17 @@ class VectorStoreRegistry: """ vector_stores_from_db: Final[list[LiteLLM_ManagedVectorStore]] = [] if prisma_client is not None: + from litellm.proxy.vector_store_endpoints.litellm_params_encryption import ( + decrypt_vector_store_litellm_params, + ) + _vector_stores_from_db: Final = await ManagedVectorStoresRepository(prisma_client).table.find_many( order={"created_at": "desc"}, ) for vector_store in _vector_stores_from_db: _dict_vector_store = dict(vector_store) _litellm_managed_vector_store = LiteLLM_ManagedVectorStore(**_dict_vector_store) - vector_stores_from_db.append(_litellm_managed_vector_store) + vector_stores_from_db.append(decrypt_vector_store_litellm_params(_litellm_managed_vector_store)) return vector_stores_from_db def get_credentials_for_vector_store(self, vector_store_id: str) -> dict[str, object]: diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 659dc438f2d..f8ec07690b3 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -56,6 +56,9 @@ IGNORE_FUNCTIONS = [ "_filter_argument_value", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by blocking the tool call at the cap. "_redact_scanned_content", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by returning "[REDACTED]" at the cap. "replace_ciphertexts", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); walks stored JSON, which has no cycles, and leaves values below the cap untouched. + "_map_litellm_param_strings", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); walks vector store litellm_params JSON, which has no cycles, and leaves values below the cap untouched. + "_map_param_value", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); list items of vector store litellm_params, same walk as _map_litellm_param_strings. + "_secret_strings_in", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); read-only twin of _map_param_value over vector store litellm_params. "_iter_fallback_targets", # max depth set (2 * ROUTER_MAX_FALLBACKS); fails closed by raising ValueError at the cap. "_mergeable_branch", # max depth set (_MAX_SCHEMA_FLATTEN_DEPTH=32) plus a seen_refs cycle guard; passes the schema through untouched at the cap. "json_string_leaves", # max depth set (MAX_STRUCTURED_CONTENT_SCAN_DEPTH); fails closed by raising at the cap so nothing goes unscanned. diff --git a/tests/test_litellm/proxy/db/test_master_key_migration.py b/tests/test_litellm/proxy/db/test_master_key_migration.py index 9c0fc163b9f..b420d9e3a4e 100644 --- a/tests/test_litellm/proxy/db/test_master_key_migration.py +++ b/tests/test_litellm/proxy/db/test_master_key_migration.py @@ -175,6 +175,32 @@ async def test_reencryption_moves_every_stored_shape_to_the_new_key_and_nothing_ ) +@pytest.mark.asyncio +async def test_reencryption_moves_marked_vector_store_credentials_and_skips_plaintext_rows(): + tables: Tables = { + "LiteLLM_ManagedVectorStoresTable": [ + { + "vector_store_id": "vs-new", + "litellm_params": { + "api_key": "litellm_enc::" + _encrypted("vector-store-key"), + "api_base": "https://vector.example/v1", + }, + }, + {"vector_store_id": "vs-legacy", "litellm_params": {"api_key": "sk-legacy-plaintext"}}, + ] + } + database = _FakeDatabase(tables) + + migrated = await reencrypt_stored_values(database, from_key=PREVIOUS_KEY, to_key=NEW_KEY) + + assert migrated == 1 + assert database.writes == [("LiteLLM_ManagedVectorStoresTable", "litellm_params", "vs-new")] + rotated = tables["LiteLLM_ManagedVectorStoresTable"][0]["litellm_params"] + assert decrypt_if_encrypted_with(rotated["api_key"].removeprefix("litellm_enc::"), NEW_KEY) == "vector-store-key" + assert rotated["api_base"] == "https://vector.example/v1" + assert tables["LiteLLM_ManagedVectorStoresTable"][1]["litellm_params"] == {"api_key": "sk-legacy-plaintext"} + + @pytest.mark.asyncio async def test_count_follows_the_values_from_the_previous_key_to_the_new_one(): database = _FakeDatabase(_seeded_tables()) 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 a5d2828dd9c..30429746ad4 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 @@ -18567,6 +18567,56 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( ) +@pytest.mark.asyncio +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()), + patch("litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_server_credentials_master_key"), + 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"), + ): + await _rotate_master_key( + prisma_client=mock_prisma_client, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"), + current_master_key="sk-old-master-key", + 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 async def test_check_encryption_endpoint_rejects_proxy_admin_viewer(): """The residual scan walks and decrypt-classifies every credential-bearing table, 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 new file mode 100644 index 00000000000..aaeba1000d2 --- /dev/null +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_litellm_params_encryption.py @@ -0,0 +1,246 @@ +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm._logging import verbose_proxy_logger +from litellm.proxy import proxy_server +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with +from litellm.proxy.vector_store_endpoints.litellm_params_encryption import ( + decrypt_vector_store_litellm_params, + encrypt_vector_store_litellm_params, + holds_undecrypted_secret, + reencrypt_vector_store_litellm_params, +) +from litellm.types.vector_stores import LiteLLM_ManagedVectorStore +from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + +_SALT_KEY = "sk-vector-store-test-salt" +_ENCRYPTED_PREFIX = "litellm_enc::" + + +@pytest.fixture +def salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", _SALT_KEY) + monkeypatch.setattr(proxy_server, "general_settings", {}) + + +def _decrypted_under(value: str, key: str): + assert value.startswith(_ENCRYPTED_PREFIX) + return decrypt_if_encrypted_with(value.removeprefix(_ENCRYPTED_PREFIX), key) + + +def test_encrypt_litellm_params_encrypts_only_secret_values(salt_key): + params = { + "api_key": "sk-vs-secret", + "api_base": "https://vector.example/v1", + "aws_secret_access_key": "aws-secret", + "aws_region_name": "us-east-1", + "valkey_password": "valkey-secret", + "max_retries": 2, + "litellm_embedding_config": {"api_key": "sk-embed-secret", "api_base": "https://embed.example"}, + "vertex_credentials": {"client_email": "svc@example.iam", "private_key": "pem-secret"}, + } + + encrypted = encrypt_vector_store_litellm_params(params) + + for key in ("api_key", "aws_secret_access_key", "valkey_password"): + assert _decrypted_under(encrypted[key], _SALT_KEY) == params[key] + assert _decrypted_under(encrypted["litellm_embedding_config"]["api_key"], _SALT_KEY) == "sk-embed-secret" + assert _decrypted_under(encrypted["vertex_credentials"]["client_email"], _SALT_KEY) == "svc@example.iam" + assert _decrypted_under(encrypted["vertex_credentials"]["private_key"], _SALT_KEY) == "pem-secret" + assert encrypted["api_base"] == "https://vector.example/v1" + assert encrypted["aws_region_name"] == "us-east-1" + assert encrypted["max_retries"] == 2 + assert encrypted["litellm_embedding_config"]["api_base"] == "https://embed.example" + for secret in ("sk-vs-secret", "aws-secret", "valkey-secret", "sk-embed-secret", "pem-secret", "svc@example"): + assert secret not in json.dumps(encrypted) + + +def test_ciphertext_supplied_by_a_caller_is_never_decrypted_back(salt_key): + ciphertext = encrypt_vector_store_litellm_params({"api_key": "someone-elses-secret"})["api_key"] + supplied = {"api_key": ciphertext, "api_base": ciphertext, "nested": {"url": ciphertext}} + + stored = encrypt_vector_store_litellm_params(supplied) + read_back = decrypt_vector_store_litellm_params( + LiteLLM_ManagedVectorStore(vector_store_id="vs", litellm_params=stored) + ) + + assert read_back["litellm_params"] == supplied + assert "someone-elses-secret" not in json.dumps(read_back) + + +def test_encrypt_litellm_params_without_a_master_or_salt_key_stores_values_as_given(monkeypatch): + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr(proxy_server, "master_key", None) + + with patch.object(verbose_proxy_logger, "warning") as warning: + assert encrypt_vector_store_litellm_params({"api_key": "sk-vs-secret"}) == {"api_key": "sk-vs-secret"} + + warning.assert_called_once() + assert "sk-vs-secret" not in str(warning.call_args) + + +def test_decrypt_litellm_params_round_trips_and_keeps_legacy_plaintext(salt_key): + params = { + "api_key": "sk-vs-secret", + "api_base": "https://vector.example/v1", + "litellm_embedding_config": {"api_key": "sk-embed-secret"}, + "vertex_credentials": {"client_email": "svc@example.iam", "private_key": "pem-secret"}, + } + encrypted_store = LiteLLM_ManagedVectorStore( + vector_store_id="vs_new", + custom_llm_provider="openai", + litellm_params=encrypt_vector_store_litellm_params(params), + ) + legacy_store = LiteLLM_ManagedVectorStore( + vector_store_id="vs_legacy", custom_llm_provider="openai", litellm_params={"api_key": "sk-legacy-plaintext"} + ) + + assert decrypt_vector_store_litellm_params(encrypted_store)["litellm_params"] == params + assert decrypt_vector_store_litellm_params(legacy_store) == legacy_store + assert decrypt_vector_store_litellm_params(LiteLLM_ManagedVectorStore(vector_store_id="vs_none")) == { + "vector_store_id": "vs_none" + } + + +def test_decrypt_litellm_params_keeps_a_value_it_cannot_decrypt(salt_key, monkeypatch): + encrypted = encrypt_vector_store_litellm_params({"api_key": "sk-vs-secret"}) + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-some-other-salt") + + store = LiteLLM_ManagedVectorStore(vector_store_id="vs", litellm_params=encrypted) + assert decrypt_vector_store_litellm_params(store)["litellm_params"] == encrypted + + +@pytest.mark.asyncio +async def test_vector_stores_loaded_from_db_have_decrypted_litellm_params(salt_key): + new_row = { + "vector_store_id": "vs_new", + "custom_llm_provider": "openai", + "litellm_params": encrypt_vector_store_litellm_params({"api_key": "sk-new-secret", "api_base": "https://b"}), + } + legacy_row = { + "vector_store_id": "vs_legacy", + "custom_llm_provider": "openai", + "litellm_params": {"api_key": "sk-legacy-plaintext"}, + } + prisma_client = MagicMock() + prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[new_row, legacy_row]) + + loaded = await VectorStoreRegistry._get_vector_stores_from_db(prisma_client=prisma_client) + + assert [vs["litellm_params"] for vs in loaded] == [ + {"api_key": "sk-new-secret", "api_base": "https://b"}, + {"api_key": "sk-legacy-plaintext"}, + ] + + +def _rotation_rows(encrypted_params): + rows = [ + {"vector_store_id": "vs_new", "litellm_params": encrypted_params}, + {"vector_store_id": "vs_legacy", "litellm_params": {"api_key": "sk-legacy-plaintext"}}, + ] + prisma_client = MagicMock() + prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=rows) + prisma_client.db.litellm_managedvectorstorestable.update_many = AsyncMock() + prisma_client.tx.return_value.__aenter__.return_value = prisma_client.db + return prisma_client + + +@pytest.mark.asyncio +async def test_reencrypt_moves_encrypted_rows_to_the_new_master_key_and_leaves_plaintext_rows(monkeypatch): + old_key, new_key = "sk-current-master-key", "sk-rotated-master-key" + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(proxy_server, "master_key", old_key) + prisma_client = _rotation_rows( + encrypt_vector_store_litellm_params({"api_key": "sk-new-secret", "api_base": "https://b"}) + ) + + 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_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" + assert _decrypted_under(stored["api_key"], old_key) is None + assert stored["api_base"] == "https://b" + + +@pytest.mark.asyncio +async def test_reencrypt_keeps_values_under_the_salt_key_when_one_is_set(salt_key): + prisma_client = _rotation_rows(encrypt_vector_store_litellm_params({"api_key": "sk-new-secret"})) + + 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_many.await_args.kwargs["data"]["litellm_params"] + ) + assert _decrypted_under(stored["api_key"], _SALT_KEY) == "sk-new-secret" + + +@pytest.mark.asyncio +async def test_reencrypt_warns_about_values_it_cannot_decrypt(salt_key, monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-removed-salt") + unreadable = encrypt_vector_store_litellm_params({"api_key": "sk-lost-secret"}) + monkeypatch.setenv("LITELLM_SALT_KEY", _SALT_KEY) + prisma_client = _rotation_rows(unreadable) + + with patch.object(verbose_proxy_logger, "warning") as warning: + 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_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) + + +@pytest.mark.asyncio +async def test_reencrypt_treats_an_empty_salt_key_as_set(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "") + monkeypatch.setattr(proxy_server, "general_settings", {}) + prisma_client = _rotation_rows(encrypt_vector_store_litellm_params({"api_key": "sk-new-secret"})) + + 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_many.await_args.kwargs["data"]["litellm_params"] + ) + assert decrypt_vector_store_litellm_params(LiteLLM_ManagedVectorStore(vector_store_id="vs", litellm_params=stored))[ + "litellm_params" + ] == {"api_key": "sk-new-secret"} + + +def test_holds_undecrypted_secret(salt_key, monkeypatch): + readable = encrypt_vector_store_litellm_params({"api_key": "sk-a", "nested": {"api_key": "sk-b"}}) + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-other-salt") + unreadable_nested = {"api_key": "sk-plain", "nested": encrypt_vector_store_litellm_params({"api_key": "sk-b"})} + monkeypatch.setenv("LITELLM_SALT_KEY", _SALT_KEY) + + def holds(params): + return holds_undecrypted_secret( + decrypt_vector_store_litellm_params(LiteLLM_ManagedVectorStore(vector_store_id="vs", litellm_params=params)) + ) + + 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 2f7d4b350be..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 @@ -3256,3 +3256,250 @@ class TestConfigOwnedVectorStores: assert response["status"] == "success", response prisma.db.litellm_managedvectorstorestable.delete.assert_awaited_once_with(where={"vector_store_id": self.DB_ID}) assert registry.get_litellm_managed_vector_store_from_registry(self.DB_ID) is None + + +class TestLitellmParamsEncryptedAtRest: + SALT_KEY = "sk-vector-store-endpoint-salt" + + @pytest.fixture(autouse=True) + def _salt_key(self, monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setenv("LITELLM_SALT_KEY", self.SALT_KEY) + monkeypatch.setattr(proxy_server, "general_settings", {}) + + @staticmethod + def _row(data): + row = MagicMock() + row.model_dump.return_value = { + **data, + "litellm_params": json.loads(data["litellm_params"]) + if isinstance(data.get("litellm_params"), str) + else data.get("litellm_params"), + } + return row + + @pytest.mark.asyncio + async def test_create_writes_encrypted_secrets_and_registers_decrypted_params(self): + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with + + prisma_client = MagicMock() + table = prisma_client.db.litellm_managedvectorstorestable + table.find_unique = AsyncMock(return_value=None) + table.create = AsyncMock(side_effect=lambda data: self._row(data)) + registry = MagicMock() + + with patch.object(litellm, "vector_store_registry", registry): + created = await create_vector_store_in_db( + vector_store_id="vs_encrypted", + custom_llm_provider="openai", + prisma_client=prisma_client, + litellm_params={"api_key": "sk-vs-secret-9", "api_base": "https://vector.example/v1"}, + ) + + stored = json.loads(table.create.await_args.kwargs["data"]["litellm_params"]) + assert "sk-vs-secret-9" not in json.dumps(stored) + assert stored["api_key"].startswith("litellm_enc::") + assert decrypt_if_encrypted_with(stored["api_key"].removeprefix("litellm_enc::"), self.SALT_KEY) == ( + "sk-vs-secret-9" + ) + assert stored["api_base"] == "https://vector.example/v1" + assert created["litellm_params"]["api_key"] == "sk-vs-secret-9" + registered = registry.add_vector_store_to_registry.call_args.kwargs["vector_store"] + assert registered["litellm_params"]["api_key"] == "sk-vs-secret-9" + + @pytest.mark.asyncio + async def test_new_endpoint_response_redacts_the_decrypted_key(self): + from litellm.constants import REDACTED_BY_LITELM_STRING + + prisma_client = MagicMock() + table = prisma_client.db.litellm_managedvectorstorestable + table.find_unique = AsyncMock(return_value=None) + table.create = AsyncMock(side_effect=lambda data: self._row(data)) + + with ( + patch( + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new_callable=AsyncMock, + ), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch.object(litellm, "vector_store_registry", None), + ): + response = await new_vector_store( + vector_store=LiteLLM_ManagedVectorStore( + vector_store_id="vs_encrypted", + custom_llm_provider="openai", + litellm_params={"api_key": "sk-vs-secret-9", "api_base": "https://vector.example/v1"}, + ), + user_api_key_dict=UserAPIKeyAuth(user_id="admin"), + ) + + assert response["vector_store"]["litellm_params"]["api_key"] == REDACTED_BY_LITELM_STRING + assert "litellm_enc::" not in json.dumps(response, default=str) + + @pytest.mark.asyncio + async def test_update_registers_decrypted_params_for_an_encrypted_row(self): + from litellm.proxy.vector_store_endpoints.management_endpoints import update_vector_store + from litellm.types.vector_stores import VectorStoreUpdateRequest + from litellm.proxy.vector_store_endpoints.litellm_params_encryption import encrypt_vector_store_litellm_params + + stored_params = encrypt_vector_store_litellm_params({"api_key": "sk-vs-secret-9", "api_base": "https://b"}) + row = self._row({"vector_store_id": "vs_encrypted", "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 + + with ( + patch( + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.vector_store_endpoints.management_endpoints._check_vector_store_access", + new_callable=AsyncMock, + return_value=True, + ), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch.object(litellm, "vector_store_registry", registry), + ): + await update_vector_store( + data=VectorStoreUpdateRequest(vector_store_id="vs_encrypted", vector_store_description="new"), + user_api_key_dict=UserAPIKeyAuth(user_id="admin"), + ) + + registered = registry.update_vector_store_in_registry.call_args.kwargs["updated_data"] + assert registered["litellm_params"] == {"api_key": "sk-vs-secret-9", "api_base": "https://b"} + + @staticmethod + def _encrypted_under_another_key(params, monkeypatch): + from litellm.proxy.vector_store_endpoints.litellm_params_encryption import encrypt_vector_store_litellm_params + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-rotated-away-salt") + encrypted = encrypt_vector_store_litellm_params(params) + monkeypatch.setenv("LITELLM_SALT_KEY", TestLitellmParamsEncryptedAtRest.SALT_KEY) + return encrypted + + @pytest.mark.asyncio + async def test_list_keeps_the_registry_copy_of_a_row_this_proxy_cannot_decrypt(self, monkeypatch): + from litellm.proxy.vector_store_endpoints.management_endpoints import list_vector_stores + from litellm.proxy.vector_store_endpoints.litellm_params_encryption import encrypt_vector_store_litellm_params + from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + + registry = VectorStoreRegistry( + vector_stores=[ + LiteLLM_ManagedVectorStore(vector_store_id="vs_rotated", litellm_params={"api_key": "sk-working"}), + LiteLLM_ManagedVectorStore(vector_store_id="vs_readable", litellm_params={"api_key": "sk-old"}), + ] + ) + 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"}), + }, + ] + prisma_client = MagicMock() + prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=rows) + + with ( + patch( + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.vector_store_endpoints.management_endpoints.filter_listable_vector_stores", + new_callable=AsyncMock, + side_effect=lambda stores, _: list(stores), + ), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch.object(litellm, "vector_store_registry", registry), + ): + await list_vector_stores(user_api_key_dict=UserAPIKeyAuth(user_id="admin")) + + 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", "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 = VectorStoreRegistry( + vector_stores=[ + LiteLLM_ManagedVectorStore(vector_store_id="vs_rotated", litellm_params={"api_key": "sk-vs-secret-9"}) + ] + ) + + with ( + patch( + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.vector_store_endpoints.management_endpoints._check_vector_store_access", + new_callable=AsyncMock, + return_value=True, + ), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch.object(litellm, "vector_store_registry", registry), + ): + await update_vector_store( + data=VectorStoreUpdateRequest(vector_store_id="vs_rotated", vector_store_description="new"), + user_api_key_dict=UserAPIKeyAuth(user_id="admin"), + ) + + 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): + from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable + from litellm.proxy.vector_store_endpoints.utils import get_litellm_managed_vector_store + from litellm.proxy.vector_store_endpoints.litellm_params_encryption import encrypt_vector_store_litellm_params + + cached_row = LiteLLM_ManagedVectorStoresTable( + vector_store_id="vs_encrypted", + custom_llm_provider="openai", + litellm_params=encrypt_vector_store_litellm_params({"api_key": "sk-vs-secret-9"}), + ) + + with ( + patch.object(litellm, "vector_store_registry", None), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", + new_callable=AsyncMock, + return_value=[cached_row], + ), + ): + resolved = await get_litellm_managed_vector_store("vs_encrypted") + + assert resolved is not None + assert resolved["litellm_params"] == {"api_key": "sk-vs-secret-9"}