diff --git a/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py b/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py index 64a35b1995e..4f54ca291b9 100644 --- a/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py +++ b/litellm/proxy/vector_store_endpoints/litellm_params_encryption.py @@ -1,8 +1,10 @@ """Encryption at rest for the secret values in a managed vector store's ``litellm_params``.""" import os -from collections.abc import Callable, Mapping +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 @@ -41,20 +43,54 @@ def _map_litellm_param_strings( ``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 _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 + 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]: @@ -89,17 +125,10 @@ def decrypt_vector_store_litellm_params(vector_store: LiteLLM_ManagedVectorStore 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.""" - undecrypted: Final[list[str]] = [] - - def find(key: str, value: str, is_secret: bool) -> str: - if is_secret and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): - undecrypted.append(key) - return value - litellm_params: Final = vector_store.get("litellm_params") - if isinstance(litellm_params, dict): - _map_litellm_param_strings(litellm_params, find) - return bool(undecrypted) + 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: @@ -123,35 +152,28 @@ async def reencrypt_vector_store_litellm_params(prisma_client: "PrismaClient", n """ salt_key: Final = os.getenv("LITELLM_SALT_KEY") new_encryption_key: Final = new_master_key if salt_key is None else salt_key - undecryptable_keys: Final[list[str]] = [] 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: - if is_secret and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): - undecryptable_keys.append(key) return value return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(plaintext, new_encryption_key=new_encryption_key) - table: Final = ManagedVectorStoresRepository(prisma_client).table - 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") - if not isinstance(litellm_params, dict): - continue - undecryptable_keys.clear() - reencrypted = _map_litellm_param_strings(litellm_params, reencrypt) - 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", - stored.get("vector_store_id"), - sorted(undecryptable_keys), - ) - vector_store_id = stored.get("vector_store_id") - if reencrypted == litellm_params or vector_store_id is None: - continue - rewrites[vector_store_id] = reencrypted + 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(): @@ -160,3 +182,22 @@ async def reencrypt_vector_store_litellm_params(prisma_client: "PrismaClient", n 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/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 585ac836c95..f8ec07690b3 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -57,6 +57,8 @@ IGNORE_FUNCTIONS = [ "_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.