mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
74fac7c09a
commit
f3ca85874c
4 changed files with 35 additions and 36 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue