mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vector_stores): keep working credentials across rotation, walk lists, rotate in one transaction
This commit is contained in:
parent
3e497771df
commit
e3a6679eb4
4 changed files with 107 additions and 28 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue