refactor(vector_stores): build encrypted litellm_params without mutation

This commit is contained in:
Yucheng He 2026-09-28 17:34:33 -07:00
parent b8569cc89f
commit 4d916606ef
2 changed files with 90 additions and 47 deletions

View file

@ -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,
)

View file

@ -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.