fix(search_tools): keep edits made while the master key rotates

Write each rotated search tool row only if it still holds the litellm_params
that were read, and re-read and rotate it again if it was edited in between,
so a PUT that lands during /key/regenerate is not overwritten.
This commit is contained in:
Yucheng He 2026-09-28 15:56:23 -07:00
parent fc801cd870
commit 2b3a0f0bc2
3 changed files with 86 additions and 16 deletions

View file

@ -33,6 +33,8 @@ class SearchToolTableClient(Protocol):
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> SearchToolRecord: ...
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
async def delete(self, where: Mapping[str, object]) -> SearchToolRecord: ...
@ -74,28 +76,49 @@ def _reencrypt_search_tool_value(value: object, new_master_key: str) -> object:
return value if plaintext is None else encrypt_value_helper(value=plaintext, new_encryption_key=new_master_key)
async def rotate_search_tools_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
"""Re-encrypt the litellm_params values that decrypt under the current key with new_master_key.
_ROTATE_ROW_ATTEMPTS: Final = 5
Values that do not decrypt under the current key (plaintext rows written before encryption, or
ciphertext under another key) are kept as stored.
"""
table: Final = _search_tools_table(prisma_client)
for row in await table.find_many():
stored_litellm_params = dict(row).get("litellm_params")
if not isinstance(stored_litellm_params, Mapping):
continue
await table.update(
where={"search_tool_id": row.search_tool_id},
async def _rotate_search_tool_row(
table: SearchToolTableClient, search_tool_id: str, stored_litellm_params: Mapping[str, object], new_master_key: str
) -> None:
current_litellm_params: Mapping[str, object] = stored_litellm_params
for _ in range(_ROTATE_ROW_ATTEMPTS):
rows_updated = await table.update_many(
where={"search_tool_id": search_tool_id, "litellm_params": {"equals": safe_dumps(current_litellm_params)}},
data={
"litellm_params": safe_dumps(
{
key: _reencrypt_search_tool_value(value, new_master_key)
for key, value in stored_litellm_params.items()
for key, value in current_litellm_params.items()
}
)
},
)
if rows_updated:
return
reread = await table.find_unique(where={"search_tool_id": search_tool_id})
reread_litellm_params = None if reread is None else dict(reread).get("litellm_params")
if not isinstance(reread_litellm_params, Mapping):
return
current_litellm_params = reread_litellm_params
verbose_proxy_logger.warning(
"Search tool %s kept changing during master key rotation and was not re-encrypted", search_tool_id
)
async def rotate_search_tools_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
"""Re-encrypt the litellm_params values that decrypt under the current key with new_master_key.
Values that do not decrypt under the current key (plaintext rows written before encryption, or
ciphertext under another key) are kept as stored. Each row is written only if it still holds the
litellm_params that were read, and is re-read and rotated again if it was edited in between.
"""
table: Final = _search_tools_table(prisma_client)
for row in await table.find_many():
stored_litellm_params = dict(row).get("litellm_params")
if isinstance(stored_litellm_params, Mapping):
await _rotate_search_tool_row(table, row.search_tool_id, stored_litellm_params, new_master_key)
class SearchToolRegistry:

View file

@ -1192,6 +1192,28 @@ class _InMemorySearchToolsTable:
setattr(row, column, json.loads(value) if column in ("litellm_params", "search_tool_info") else value)
return row
async def update_many(self, where, data):
row = self.rows.get(where["search_tool_id"])
if row is None or row.litellm_params != json.loads(where["litellm_params"]["equals"]):
return 0
await self.update(where={"search_tool_id": row.search_tool_id}, data=data)
return 1
class _TableWithEditDuringRotation(_InMemorySearchToolsTable):
"""Applies an admin edit to a row right before the rotation's first conditional write to it."""
def __init__(self, rows, edited_id, edited_params):
super().__init__(rows)
self.pending_edit = (edited_id, edited_params)
async def update_many(self, where, data):
if self.pending_edit and self.pending_edit[0] == where["search_tool_id"]:
edited_id, edited_params = self.pending_edit
self.pending_edit = None
self.rows[edited_id].litellm_params = edited_params
return await super().update_many(where, data)
def _stored_row(search_tool_id: str, name: str, litellm_params: dict) -> _StoredSearchToolRow:
return _StoredSearchToolRow(
@ -1315,3 +1337,25 @@ async def test_master_key_rotation_reencrypts_only_values_the_current_key_decryp
assert table.rows["legacy-id"].litellm_params == legacy_params
assert table.rows["foreign-id"].litellm_params == {"api_key": foreign_ciphertext}
assert json.dumps({row_id: row.litellm_params for row_id, row in table.rows.items()}) == after_first_rotation
@pytest.mark.asyncio
async def test_master_key_rotation_keeps_an_edit_made_while_it_runs(salt_key):
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_if_encrypted_with,
encrypt_value_helper,
)
from litellm.proxy.search_endpoints.search_tool_registry import rotate_search_tools_master_key
new_key = "sk-new-master-key"
table = _TableWithEditDuringRotation(
[_stored_row("edited-id", "edited", {"api_key": encrypt_value_helper("tvly-before-edit")})],
edited_id="edited-id",
edited_params={"api_key": encrypt_value_helper("tvly-after-edit"), "max_results": 3},
)
await rotate_search_tools_master_key(prisma_client=_prisma_client_over(table), new_master_key=new_key)
rotated = table.rows["edited-id"].litellm_params
assert decrypt_if_encrypted_with(rotated["api_key"], new_key) == "tvly-after-edit"
assert rotated["max_results"] == 3

View file

@ -18589,9 +18589,12 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch):
def __iter__(self):
return iter(vars(self).items())
async def _update(where, data):
assert where == {"search_tool_id": "search-tool-1"}
async def _update_many(where, data):
assert where["search_tool_id"] == "search-tool-1"
if json.loads(where["litellm_params"]["equals"]) != stored:
return 0
stored.update(json.loads(data["litellm_params"]))
return 1
mock_prisma_client = AsyncMock()
mock_prisma_client.db = MagicMock()
@ -18601,7 +18604,7 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch):
mock_prisma_client.db.litellm_searchtoolstable.find_many = AsyncMock(
return_value=[_Row(search_tool_id="search-tool-1", litellm_params=dict(stored))]
)
mock_prisma_client.db.litellm_searchtoolstable.update = AsyncMock(side_effect=_update)
mock_prisma_client.db.litellm_searchtoolstable.update_many = AsyncMock(side_effect=_update_many)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",