diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index b77cd17ccc0..f9d2fa3065b 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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 diff --git a/litellm/proxy/search_endpoints/search_tool_registry.py b/litellm/proxy/search_endpoints/search_tool_registry.py index 19f62511bef..f6e6bea645a 100644 --- a/litellm/proxy/search_endpoints/search_tool_registry.py +++ b/litellm/proxy/search_endpoints/search_tool_registry.py @@ -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: diff --git a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py index ce792246769..7e8932baa0f 100644 --- a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py +++ b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py @@ -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()) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index c13465191e3..f182995c5f8 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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