mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(search_tools): retry rotation writes until the row stops changing
Rotate a search tool row again for as long as it keeps being edited instead of giving up after five attempts, and stop with a warning only when the conditional write fails on an unchanged row. Build the decrypted read result without mutating it in place.
This commit is contained in:
parent
2b3a0f0bc2
commit
74fac7c09a
2 changed files with 50 additions and 30 deletions
|
|
@ -76,35 +76,32 @@ 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)
|
||||
|
||||
|
||||
_ROTATE_ROW_ATTEMPTS: Final = 5
|
||||
|
||||
|
||||
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 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
|
||||
rows_updated: Final = await table.update_many(
|
||||
where={"search_tool_id": search_tool_id, "litellm_params": {"equals": safe_dumps(stored_litellm_params)}},
|
||||
data={
|
||||
"litellm_params": safe_dumps(
|
||||
{
|
||||
key: _reencrypt_search_tool_value(value, new_master_key)
|
||||
for key, value in stored_litellm_params.items()
|
||||
}
|
||||
)
|
||||
},
|
||||
)
|
||||
if rows_updated:
|
||||
return
|
||||
reread: Final = await table.find_unique(where={"search_tool_id": search_tool_id})
|
||||
reread_litellm_params: Final = None if reread is None else dict(reread).get("litellm_params")
|
||||
if not isinstance(reread_litellm_params, Mapping):
|
||||
return
|
||||
if reread_litellm_params == stored_litellm_params:
|
||||
verbose_proxy_logger.warning(
|
||||
"Search tool %s was not re-encrypted: its stored litellm_params did not match on write", search_tool_id
|
||||
)
|
||||
return
|
||||
await _rotate_search_tool_row(table, search_tool_id, reread_litellm_params, new_master_key)
|
||||
|
||||
|
||||
async def rotate_search_tools_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
|
||||
|
|
@ -112,7 +109,7 @@ async def rotate_search_tools_master_key(prisma_client: PrismaClient, new_master
|
|||
|
||||
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.
|
||||
litellm_params that were read, and is re-read and rotated again while it keeps being edited in between.
|
||||
"""
|
||||
table: Final = _search_tools_table(prisma_client)
|
||||
for row in await table.find_many():
|
||||
|
|
@ -140,10 +137,15 @@ class SearchToolRegistry:
|
|||
Returns:
|
||||
Dict with datetime fields converted to ISO strings
|
||||
"""
|
||||
result: Final = dict(prisma_obj)
|
||||
stored_litellm_params: Final = result.get("litellm_params")
|
||||
if isinstance(stored_litellm_params, Mapping):
|
||||
result["litellm_params"] = decrypt_search_tool_litellm_params(stored_litellm_params)
|
||||
stored_litellm_params: Final = dict(prisma_obj).get("litellm_params")
|
||||
result: Final = {
|
||||
**dict(prisma_obj),
|
||||
**(
|
||||
{"litellm_params": decrypt_search_tool_litellm_params(stored_litellm_params)}
|
||||
if isinstance(stored_litellm_params, Mapping)
|
||||
else {}
|
||||
),
|
||||
}
|
||||
# Convert datetime objects to ISO format strings
|
||||
if "created_at" in result and result["created_at"]:
|
||||
result["created_at"] = prisma_obj.created_at.isoformat()
|
||||
|
|
|
|||
|
|
@ -1359,3 +1359,21 @@ async def test_master_key_rotation_keeps_an_edit_made_while_it_runs(salt_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
|
||||
|
||||
|
||||
class _TableWhoseConditionalWritesNeverMatch(_InMemorySearchToolsTable):
|
||||
async def update_many(self, where, data):
|
||||
return 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_master_key_rotation_leaves_a_row_that_never_matches_and_finishes(salt_key):
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import rotate_search_tools_master_key
|
||||
|
||||
stored = {"api_key": encrypt_value_helper("tvly-unmatched")}
|
||||
table = _TableWhoseConditionalWritesNeverMatch([_stored_row("unmatched-id", "unmatched", dict(stored))])
|
||||
|
||||
await rotate_search_tools_master_key(prisma_client=_prisma_client_over(table), new_master_key="sk-new-master-key")
|
||||
|
||||
assert table.rows["unmatched-id"].litellm_params == stored
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue