mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(vector_stores): build encrypted litellm_params without mutation
This commit is contained in:
parent
b8569cc89f
commit
4d916606ef
2 changed files with 90 additions and 47 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue