fix(vector_stores): keep working credentials across rotation, walk lists, rotate in one transaction

This commit is contained in:
Yucheng He 2026-09-28 16:55:21 -07:00
parent 3e497771df
commit e3a6679eb4
4 changed files with 107 additions and 28 deletions

View file

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

View file

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

View file

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

View file

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