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:
Yucheng He 2026-09-28 16:09:00 -07:00
parent 2b3a0f0bc2
commit 74fac7c09a
2 changed files with 50 additions and 30 deletions

View file

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

View file

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