refactor(search_tools): rotate edited rows in a loop, drop the step comment

Retry the conditional rotation write in a loop instead of recursion so sustained
edits cannot deepen the call stack, drop the step comment on the rotation call,
and stop mutating local state in the rotation tests.
This commit is contained in:
Yucheng He 2026-09-28 16:45:30 -07:00
parent 74fac7c09a
commit f3ca85874c
4 changed files with 35 additions and 36 deletions

View file

@ -5339,7 +5339,6 @@ async def _rotate_master_key(
except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation
verbose_proxy_logger.warning("Failed to rotate SSO identity assertions: %s", str(e))
# 4e. process search tools table
try:
await rotate_search_tools_master_key(prisma_client=prisma_client, new_master_key=new_master_key)
except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation

View file

@ -79,29 +79,29 @@ def _reencrypt_search_tool_value(value: object, new_master_key: str) -> object:
async def _rotate_search_tool_row(
table: SearchToolTableClient, search_tool_id: str, stored_litellm_params: Mapping[str, object], new_master_key: str
) -> None:
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
expected_litellm_params: Mapping[str, object] | None = stored_litellm_params
while expected_litellm_params is not None:
rows_updated = await table.update_many(
where={"search_tool_id": search_tool_id, "litellm_params": {"equals": safe_dumps(expected_litellm_params)}},
data={
"litellm_params": safe_dumps(
{
key: _reencrypt_search_tool_value(value, new_master_key)
for key, value in expected_litellm_params.items()
}
)
},
)
return
await _rotate_search_tool_row(table, search_tool_id, reread_litellm_params, new_master_key)
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 isinstance(reread_litellm_params, Mapping) and reread_litellm_params == expected_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
expected_litellm_params = reread_litellm_params if isinstance(reread_litellm_params, Mapping) else None
async def rotate_search_tools_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:

View file

@ -1,5 +1,6 @@
import contextlib
import json
from types import SimpleNamespace
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
@ -1151,10 +1152,7 @@ async def test_create_search_tool_survives_a_failing_router_refresh():
assert response.json()["search_tool_name"] == "tavily-search"
class _StoredSearchToolRow:
def __init__(self, **fields):
self.__dict__.update(fields)
class _StoredSearchToolRow(SimpleNamespace):
def __iter__(self):
return iter(self.__dict__.items())

View file

@ -18583,17 +18583,21 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch):
)
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-old-master-key")
stored = {"search_provider": "tavily", "api_key": encrypt_value_helper("tvly-secret")}
class _Row(SimpleNamespace):
def __iter__(self):
return iter(vars(self).items())
row = _Row(
search_tool_id="search-tool-1",
litellm_params={"search_provider": "tavily", "api_key": encrypt_value_helper("tvly-secret")},
)
async def _update_many(where, data):
assert where["search_tool_id"] == "search-tool-1"
if json.loads(where["litellm_params"]["equals"]) != stored:
expected_litellm_params = json.loads(where["litellm_params"]["equals"])
if where["search_tool_id"] != row.search_tool_id or expected_litellm_params != row.litellm_params:
return 0
stored.update(json.loads(data["litellm_params"]))
row.litellm_params = json.loads(data["litellm_params"])
return 1
mock_prisma_client = AsyncMock()
@ -18601,9 +18605,7 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch):
mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[])
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.find_many = AsyncMock(return_value=[row])
mock_prisma_client.db.litellm_searchtoolstable.update_many = AsyncMock(side_effect=_update_many)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
@ -18618,8 +18620,8 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch):
new_master_key="sk-new-master-key",
)
assert decrypt_if_encrypted_with(stored["api_key"], "sk-new-master-key") == "tvly-secret"
assert stored["search_provider"] == "tavily"
assert decrypt_if_encrypted_with(row.litellm_params["api_key"], "sk-new-master-key") == "tvly-secret"
assert row.litellm_params["search_provider"] == "tavily"
@pytest.mark.asyncio