diff --git a/litellm/proxy/search_endpoints/search_tool_registry.py b/litellm/proxy/search_endpoints/search_tool_registry.py index 99366ed7d8f..19f62511bef 100644 --- a/litellm/proxy/search_endpoints/search_tool_registry.py +++ b/litellm/proxy/search_endpoints/search_tool_registry.py @@ -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() 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 e5e99cf7720..ce792246769 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 @@ -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