mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
fc801cd870
commit
2b3a0f0bc2
3 changed files with 86 additions and 16 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue