From 2b3a0f0bc2f5cef00fcca3663bd2219330c06688 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 28 Sep 2026 15:56:23 -0700 Subject: [PATCH] fix(search_tools): keep edits made while the master key rotates Write each rotated search tool row only if it still holds the litellm_params that were read, and re-read and rotate it again if it was edited in between, so a PUT that lands during /key/regenerate is not overwritten. --- .../search_endpoints/search_tool_registry.py | 49 ++++++++++++++----- .../test_search_tool_management.py | 44 +++++++++++++++++ .../test_key_management_endpoints.py | 9 ++-- 3 files changed, 86 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/search_endpoints/search_tool_registry.py b/litellm/proxy/search_endpoints/search_tool_registry.py index 822e6b5cdc2..99366ed7d8f 100644 --- a/litellm/proxy/search_endpoints/search_tool_registry.py +++ b/litellm/proxy/search_endpoints/search_tool_registry.py @@ -33,6 +33,8 @@ class SearchToolTableClient(Protocol): async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> SearchToolRecord: ... + async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... + async def delete(self, where: Mapping[str, object]) -> SearchToolRecord: ... @@ -74,28 +76,49 @@ 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) -async def rotate_search_tools_master_key(prisma_client: PrismaClient, new_master_key: str) -> None: - """Re-encrypt the litellm_params values that decrypt under the current key with new_master_key. +_ROTATE_ROW_ATTEMPTS: Final = 5 - Values that do not decrypt under the current key (plaintext rows written before encryption, or - ciphertext under another key) are kept as stored. - """ - table: Final = _search_tools_table(prisma_client) - for row in await table.find_many(): - stored_litellm_params = dict(row).get("litellm_params") - if not isinstance(stored_litellm_params, Mapping): - continue - await table.update( - where={"search_tool_id": row.search_tool_id}, + +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 stored_litellm_params.items() + 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 + ) + + +async def rotate_search_tools_master_key(prisma_client: PrismaClient, new_master_key: str) -> None: + """Re-encrypt the litellm_params values that decrypt under the current key with new_master_key. + + 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. + """ + table: Final = _search_tools_table(prisma_client) + for row in await table.find_many(): + stored_litellm_params = dict(row).get("litellm_params") + if isinstance(stored_litellm_params, Mapping): + await _rotate_search_tool_row(table, row.search_tool_id, stored_litellm_params, new_master_key) class SearchToolRegistry: 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 6e0bf9de4fe..e5e99cf7720 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 @@ -1192,6 +1192,28 @@ class _InMemorySearchToolsTable: setattr(row, column, json.loads(value) if column in ("litellm_params", "search_tool_info") else value) return row + async def update_many(self, where, data): + row = self.rows.get(where["search_tool_id"]) + if row is None or row.litellm_params != json.loads(where["litellm_params"]["equals"]): + return 0 + await self.update(where={"search_tool_id": row.search_tool_id}, data=data) + return 1 + + +class _TableWithEditDuringRotation(_InMemorySearchToolsTable): + """Applies an admin edit to a row right before the rotation's first conditional write to it.""" + + def __init__(self, rows, edited_id, edited_params): + super().__init__(rows) + self.pending_edit = (edited_id, edited_params) + + async def update_many(self, where, data): + if self.pending_edit and self.pending_edit[0] == where["search_tool_id"]: + edited_id, edited_params = self.pending_edit + self.pending_edit = None + self.rows[edited_id].litellm_params = edited_params + return await super().update_many(where, data) + def _stored_row(search_tool_id: str, name: str, litellm_params: dict) -> _StoredSearchToolRow: return _StoredSearchToolRow( @@ -1315,3 +1337,25 @@ async def test_master_key_rotation_reencrypts_only_values_the_current_key_decryp assert table.rows["legacy-id"].litellm_params == legacy_params assert table.rows["foreign-id"].litellm_params == {"api_key": foreign_ciphertext} assert json.dumps({row_id: row.litellm_params for row_id, row in table.rows.items()}) == after_first_rotation + + +@pytest.mark.asyncio +async def test_master_key_rotation_keeps_an_edit_made_while_it_runs(salt_key): + from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_if_encrypted_with, + encrypt_value_helper, + ) + from litellm.proxy.search_endpoints.search_tool_registry import rotate_search_tools_master_key + + new_key = "sk-new-master-key" + table = _TableWithEditDuringRotation( + [_stored_row("edited-id", "edited", {"api_key": encrypt_value_helper("tvly-before-edit")})], + edited_id="edited-id", + edited_params={"api_key": encrypt_value_helper("tvly-after-edit"), "max_results": 3}, + ) + + await rotate_search_tools_master_key(prisma_client=_prisma_client_over(table), new_master_key=new_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 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 9739f7cb6bb..c13465191e3 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 @@ -18589,9 +18589,12 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch): def __iter__(self): return iter(vars(self).items()) - async def _update(where, data): - assert where == {"search_tool_id": "search-tool-1"} + async def _update_many(where, data): + assert where["search_tool_id"] == "search-tool-1" + if json.loads(where["litellm_params"]["equals"]) != stored: + return 0 stored.update(json.loads(data["litellm_params"])) + return 1 mock_prisma_client = AsyncMock() mock_prisma_client.db = MagicMock() @@ -18601,7 +18604,7 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch): 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.update = AsyncMock(side_effect=_update) + mock_prisma_client.db.litellm_searchtoolstable.update_many = AsyncMock(side_effect=_update_many) user_api_key_dict = UserAPIKeyAuth( user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234",