From fc801cd87077706672ff1a694c4e11b2d25692a7 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 28 Sep 2026 14:27:42 -0700 Subject: [PATCH 1/5] fix(search_tools): encrypt search tool litellm_params at rest Encrypt every string value of a search tool's litellm_params on create and update and decrypt on every DB read, so legacy plaintext rows load unchanged. Include the table in master key rotation, LITELLM_MIGRATE_FROM_MASTER_KEY and the migrate-encryption scan. --- litellm/proxy/db/master_key_migration.py | 1 + .../credential_migration.py | 1 + .../key_management_endpoints.py | 7 + .../search_endpoints/search_tool_registry.py | 63 ++++++- .../proxy/db/test_master_key_migration.py | 21 +++ .../test_search_tool_management.py | 167 ++++++++++++++++++ .../test_credential_migration.py | 22 +++ .../test_key_management_endpoints.py | 52 ++++++ 8 files changed, 331 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/db/master_key_migration.py b/litellm/proxy/db/master_key_migration.py index d100554201a..490003be35a 100644 --- a/litellm/proxy/db/master_key_migration.py +++ b/litellm/proxy/db/master_key_migration.py @@ -38,6 +38,7 @@ _SECRET_COLUMNS: Final = ( _SecretColumn("LiteLLM_MCPUserCredentials", "id", "credential_b64", is_json=False), _SecretColumn("LiteLLM_MCPUserEnvVars", "id", "values_b64", is_json=False), _SecretColumn("LiteLLM_SSOIdentityAssertion", "user_id", "assertion_b64", is_json=False), + _SecretColumn("LiteLLM_SearchToolsTable", "search_tool_id", "litellm_params"), _SecretColumn("LiteLLM_TeamTable", "team_id", "metadata", only_rows_with_marked_ciphertexts=True), _SecretColumn("LiteLLM_VerificationToken", "token", "metadata", only_rows_with_marked_ciphertexts=True), _SecretColumn("LiteLLM_UserTable", "user_id", "metadata", only_rows_with_marked_ciphertexts=True), diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index 915cce87dbd..73e26aa827a 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -460,6 +460,7 @@ _COVERED_TABLE_SPECS: Final = [ ("mcp_server", "litellm_mcpservertable", ("credentials", "env_vars", "static_headers", "env"), ()), ("mcp_user_credentials", "litellm_mcpusercredentials", (), ("credential_b64",)), ("mcp_user_env_vars", "litellm_mcpuserenvvars", (), ("values_b64",)), + ("search_tools", "litellm_searchtoolstable", ("litellm_params",), ()), ] diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 7e159ec90e7..b77cd17ccc0 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -124,6 +124,7 @@ from litellm.proxy.management_helpers.team_member_permission_checks import ( TeamMemberPermissionChecks, ) from litellm.proxy.management_helpers.utils import management_endpoint_wrapper +from litellm.proxy.search_endpoints.search_tool_registry import rotate_search_tools_master_key from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start from litellm.proxy.spend_tracking.spend_tracking_utils import _is_master_key from litellm.proxy.utils import ( @@ -5338,6 +5339,12 @@ 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 + verbose_proxy_logger.warning("Failed to rotate search tool credentials: %s", str(e)) + # 5. process credentials table try: credentials = await _credentials_table(prisma_client).find_many() diff --git a/litellm/proxy/search_endpoints/search_tool_registry.py b/litellm/proxy/search_endpoints/search_tool_registry.py index b25263e4c64..822e6b5cdc2 100644 --- a/litellm/proxy/search_endpoints/search_tool_registry.py +++ b/litellm/proxy/search_endpoints/search_tool_registry.py @@ -8,6 +8,7 @@ from typing import Final, Protocol from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.proxy.utils import PrismaClient from litellm.repositories.table_repositories import SearchToolsRepository @@ -48,6 +49,55 @@ def _search_tools_table(prisma_client: PrismaClient) -> SearchToolTableClient: return _search_tools_table_of(SearchToolsRepository(prisma_client)) +def encrypt_search_tool_litellm_params(litellm_params: Mapping[str, object]) -> Mapping[str, object]: + """Encrypt every string value of a search tool's litellm_params for storage.""" + return { + key: encrypt_value_helper(value=value) if isinstance(value, str) else value + for key, value in litellm_params.items() + } + + +def decrypt_search_tool_litellm_params(litellm_params: Mapping[str, object]) -> Mapping[str, object]: + """Decrypt stored litellm_params values; values that are not ciphertext are returned unchanged.""" + return { + key: decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=True) + if isinstance(value, str) + else value + for key, value in litellm_params.items() + } + + +def _reencrypt_search_tool_value(value: object, new_master_key: str) -> object: + if not isinstance(value, str): + return value + plaintext: Final = decrypt_value_helper(value=value, key="search_tool", exception_type="debug") + 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. + + 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}, + data={ + "litellm_params": safe_dumps( + { + key: _reencrypt_search_tool_value(value, new_master_key) + for key, value in stored_litellm_params.items() + } + ) + }, + ) + + class SearchToolRegistry: """ Handles adding, removing, and getting search tools in DB + in memory. @@ -59,7 +109,7 @@ class SearchToolRegistry: @staticmethod def _convert_prisma_to_dict(prisma_obj: SearchToolRecord) -> dict: """ - Convert Prisma result to dict with datetime objects as ISO format strings. + Convert Prisma result to dict with decrypted litellm_params and datetime objects as ISO format strings. Args: prisma_obj: Prisma model instance @@ -68,6 +118,9 @@ class SearchToolRegistry: 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) # Convert datetime objects to ISO format strings if "created_at" in result and result["created_at"]: result["created_at"] = prisma_obj.created_at.isoformat() @@ -92,7 +145,9 @@ class SearchToolRegistry: """ try: search_tool_name: Final = search_tool.get("search_tool_name") - litellm_params: Final[str] = safe_dumps(dict(search_tool.get("litellm_params", {}))) + litellm_params: Final[str] = safe_dumps( + encrypt_search_tool_litellm_params(search_tool.get("litellm_params", {})) + ) search_tool_info: Final[str] = safe_dumps(search_tool.get("search_tool_info", {})) # Create search tool in DB @@ -162,7 +217,9 @@ class SearchToolRegistry: """ try: search_tool_name: Final = search_tool.get("search_tool_name") - litellm_params: Final[str] = safe_dumps(dict(search_tool.get("litellm_params", {}))) + litellm_params: Final[str] = safe_dumps( + encrypt_search_tool_litellm_params(search_tool.get("litellm_params", {})) + ) search_tool_info: Final[str] = safe_dumps(search_tool.get("search_tool_info", {})) # Update in DB diff --git a/tests/test_litellm/proxy/db/test_master_key_migration.py b/tests/test_litellm/proxy/db/test_master_key_migration.py index 9c0fc163b9f..ff7881a192b 100644 --- a/tests/test_litellm/proxy/db/test_master_key_migration.py +++ b/tests/test_litellm/proxy/db/test_master_key_migration.py @@ -175,6 +175,27 @@ async def test_reencryption_moves_every_stored_shape_to_the_new_key_and_nothing_ ) +@pytest.mark.asyncio +async def test_search_tool_litellm_params_are_moved_to_the_new_key(): + tables: Tables = { + "LiteLLM_SearchToolsTable": [ + { + "search_tool_id": "search-tool-1", + "litellm_params": {"search_provider": _encrypted("tavily"), "api_key": _encrypted("tvly-secret")}, + }, + {"search_tool_id": "legacy-search-tool", "litellm_params": {"api_key": "tvly-plaintext"}}, + ] + } + + migrated = await reencrypt_stored_values(_FakeDatabase(tables), from_key=PREVIOUS_KEY, to_key=NEW_KEY) + + assert migrated == 2 + search_tool_params = tables["LiteLLM_SearchToolsTable"][0]["litellm_params"] + assert decrypt_if_encrypted_with(search_tool_params["api_key"], NEW_KEY) == "tvly-secret" + assert decrypt_if_encrypted_with(search_tool_params["search_provider"], NEW_KEY) == "tavily" + assert tables["LiteLLM_SearchToolsTable"][1]["litellm_params"] == {"api_key": "tvly-plaintext"} + + @pytest.mark.asyncio async def test_count_follows_the_values_from_the_previous_key_to_the_new_one(): database = _FakeDatabase(_seeded_tables()) 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 70e9a96b316..6e0bf9de4fe 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,4 +1,5 @@ import contextlib +import json from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch @@ -1148,3 +1149,169 @@ async def test_create_search_tool_survives_a_failing_router_refresh(): assert response.status_code == 200 assert response.json()["search_tool_name"] == "tavily-search" + + +class _StoredSearchToolRow: + def __init__(self, **fields): + self.__dict__.update(fields) + + def __iter__(self): + return iter(self.__dict__.items()) + + +class _InMemorySearchToolsTable: + """Stands in for prisma's litellm_searchtoolstable: JSON columns are stored parsed, as prisma returns them.""" + + def __init__(self, rows=()): + self.rows = {row.search_tool_id: row for row in rows} + + async def create(self, data): + row = _StoredSearchToolRow( + search_tool_id=f"id-{len(self.rows)}", + search_tool_name=data["search_tool_name"], + litellm_params=json.loads(data["litellm_params"]), + search_tool_info=json.loads(data["search_tool_info"]), + created_at=data["created_at"], + updated_at=data["updated_at"], + ) + self.rows[row.search_tool_id] = row + return row + + async def find_unique(self, where): + return self.rows.get(where.get("search_tool_id")) or next( + (row for row in self.rows.values() if row.search_tool_name == where.get("search_tool_name")), + None, + ) + + async def find_many(self, order=None): + return list(self.rows.values()) + + async def update(self, where, data): + row = self.rows[where["search_tool_id"]] + for column, value in data.items(): + setattr(row, column, json.loads(value) if column in ("litellm_params", "search_tool_info") else value) + return row + + +def _stored_row(search_tool_id: str, name: str, litellm_params: dict) -> _StoredSearchToolRow: + return _StoredSearchToolRow( + search_tool_id=search_tool_id, + search_tool_name=name, + litellm_params=litellm_params, + search_tool_info={}, + created_at=datetime(2026, 9, 1), + updated_at=datetime(2026, 9, 1), + ) + + +def _prisma_client_over(table: _InMemorySearchToolsTable) -> MagicMock: + prisma_client = MagicMock() + prisma_client.db.litellm_searchtoolstable = table + return prisma_client + + +SALT_KEY = "sk-search-tool-salt" +SECRET_PARAMS = { + "search_provider": "bedrock_agentcore", + "api_key": "tvly-secret-api-key-0001", + "aws_secret_access_key": "aws-secret-0002", + "timeout": 30, +} + + +@pytest.fixture +def salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY) + monkeypatch.setattr(ps, "general_settings", {}) + return SALT_KEY + + +@pytest.mark.asyncio +async def test_search_tool_litellm_params_are_encrypted_at_rest_and_decrypted_on_read(salt_key): + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with + from litellm.proxy.search_endpoints.search_tool_registry import SearchToolRegistry + + table = _InMemorySearchToolsTable() + prisma_client = _prisma_client_over(table) + registry = SearchToolRegistry() + + created = await registry.add_search_tool_to_db( + search_tool={"search_tool_name": "agentcore-search", "litellm_params": SECRET_PARAMS}, + prisma_client=prisma_client, + ) + await registry.update_search_tool_in_db( + search_tool_id=created["search_tool_id"], + search_tool={ + "search_tool_name": "agentcore-search", + "litellm_params": {**SECRET_PARAMS, "api_key": "tvly-rotated-api-key-0003"}, + }, + prisma_client=prisma_client, + ) + + stored = table.rows[created["search_tool_id"]].litellm_params + assert "tvly-" not in json.dumps(stored) + assert "aws-secret-0002" not in json.dumps(stored) + assert decrypt_if_encrypted_with(stored["api_key"], salt_key) == "tvly-rotated-api-key-0003" + assert decrypt_if_encrypted_with(stored["aws_secret_access_key"], salt_key) == "aws-secret-0002" + assert stored["timeout"] == 30 + + expected = {**SECRET_PARAMS, "api_key": "tvly-rotated-api-key-0003"} + loaded = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client) + assert [tool["litellm_params"] for tool in loaded] == [expected] + by_id = await registry.get_search_tool_by_id_from_db(created["search_tool_id"], prisma_client=prisma_client) + by_name = await registry.get_search_tool_by_name_from_db("agentcore-search", prisma_client=prisma_client) + assert by_id["litellm_params"] == by_name["litellm_params"] == expected + + +@pytest.mark.asyncio +async def test_plaintext_search_tool_rows_written_before_encryption_still_load(salt_key): + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.proxy.search_endpoints.search_tool_registry import SearchToolRegistry + + encrypted_row = _stored_row( + "encrypted-id", + "encrypted", + {"search_provider": encrypt_value_helper("tavily"), "api_key": encrypt_value_helper("tvly-new")}, + ) + legacy_row = _stored_row( + "legacy-id", "legacy", {"search_provider": "perplexity", "api_key": "pplx-legacy", "max_results": 5} + ) + prisma_client = _prisma_client_over(_InMemorySearchToolsTable([encrypted_row, legacy_row])) + + loaded = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client) + + assert [tool["litellm_params"] for tool in loaded] == [ + {"search_provider": "tavily", "api_key": "tvly-new"}, + {"search_provider": "perplexity", "api_key": "pplx-legacy", "max_results": 5}, + ] + + +@pytest.mark.asyncio +async def test_master_key_rotation_reencrypts_only_values_the_current_key_decrypts(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" + foreign_ciphertext = encrypt_value_helper("tvly-foreign", new_encryption_key="sk-some-other-key") + legacy_params = {"search_provider": "perplexity", "api_key": "pplx-legacy"} + table = _InMemorySearchToolsTable( + [ + _stored_row("encrypted-id", "encrypted", {"api_key": encrypt_value_helper("tvly-new"), "timeout": 30}), + _stored_row("legacy-id", "legacy", dict(legacy_params)), + _stored_row("foreign-id", "foreign", {"api_key": foreign_ciphertext}), + ] + ) + + await rotate_search_tools_master_key(prisma_client=_prisma_client_over(table), new_master_key=new_key) + after_first_rotation = json.dumps({row_id: row.litellm_params for row_id, row in table.rows.items()}) + await rotate_search_tools_master_key(prisma_client=_prisma_client_over(table), new_master_key=new_key) + + encrypted_params = table.rows["encrypted-id"].litellm_params + assert decrypt_if_encrypted_with(encrypted_params["api_key"], new_key) == "tvly-new" + assert encrypted_params["timeout"] == 30 + 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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py b/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py index 0ecc4f8d7cb..c638f1b30b2 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py +++ b/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py @@ -458,6 +458,28 @@ async def test_scan_covered_tables_classifies_legacy_and_v2(salt_key, monkeypatc assert by_loc["credentials"].legacy == 0 +@pytest.mark.asyncio +async def test_scan_covered_tables_classifies_search_tool_params(salt_key, monkeypatch): + legacy = _legacy_ct("tvly-legacy", monkeypatch) + _enable_aes(monkeypatch) + v2 = encrypt_value_helper("tvly-migrated") + + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_searchtoolstable.find_many = AsyncMock( + return_value=[ + SimpleNamespace(litellm_params={"api_key": legacy, "timeout": 30}), + SimpleNamespace(litellm_params={"api_key": v2, "search_provider": "tavily"}), + ] + ) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + + by_loc = {r.location: r for r in await cm._scan_covered_tables(client)} + + assert (by_loc["search_tools"].legacy, by_loc["search_tools"].already_v2) == (1, 1) + assert by_loc["search_tools"].plaintext == 1 + + @pytest.mark.asyncio @pytest.mark.parametrize("column", ("static_headers", "env")) @pytest.mark.parametrize("algorithm", ("xsalsa20-poly1305", "aes-256-gcm")) 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 aa6be328f4a..9739f7cb6bb 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 @@ -18567,6 +18567,58 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( ) +@pytest.mark.asyncio +async def test_rotate_master_key_rotates_search_tools(monkeypatch): + """Master-key rotation re-encrypts the search tools table (step 4e).""" + from types import SimpleNamespace + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_if_encrypted_with, + encrypt_value_helper, + ) + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _rotate_master_key, + ) + + 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()) + + async def _update(where, data): + assert where == {"search_tool_id": "search-tool-1"} + stored.update(json.loads(data["litellm_params"])) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db = MagicMock() + 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.update = AsyncMock(side_effect=_update) + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="test-user", + ) + + await _rotate_master_key( + prisma_client=mock_prisma_client, + user_api_key_dict=user_api_key_dict, + current_master_key="sk-old-master-key", + 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" + + @pytest.mark.asyncio async def test_check_encryption_endpoint_rejects_proxy_admin_viewer(): """The residual scan walks and decrypt-classifies every credential-bearing table, From 2b3a0f0bc2f5cef00fcca3663bd2219330c06688 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 28 Sep 2026 15:56:23 -0700 Subject: [PATCH 2/5] 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", From 74fac7c09a591e788d3b8084516d3ad0188c990f Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 28 Sep 2026 16:09:00 -0700 Subject: [PATCH 3/5] fix(search_tools): retry rotation writes until the row stops changing Rotate a search tool row again for as long as it keeps being edited instead of giving up after five attempts, and stop with a warning only when the conditional write fails on an unchanged row. Build the decrypted read result without mutating it in place. --- .../search_endpoints/search_tool_registry.py | 62 ++++++++++--------- .../test_search_tool_management.py | 18 ++++++ 2 files changed, 50 insertions(+), 30 deletions(-) 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 From f3ca85874c42a32df328851ade8c479654b39116 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 28 Sep 2026 16:45:30 -0700 Subject: [PATCH 4/5] 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. --- .../key_management_endpoints.py | 1 - .../search_endpoints/search_tool_registry.py | 44 +++++++++---------- .../test_search_tool_management.py | 6 +-- .../test_key_management_endpoints.py | 20 +++++---- 4 files changed, 35 insertions(+), 36 deletions(-) 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 From 4b1937eba8422c7b0466c34f333e22ad854918c1 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 11:20:22 -0700 Subject: [PATCH 5/5] test(search_tools): drop the rotation test docstring --- .../proxy/management_endpoints/test_key_management_endpoints.py | 1 - 1 file changed, 1 deletion(-) 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 f182995c5f8..1b0903fff25 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 @@ -18569,7 +18569,6 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( @pytest.mark.asyncio async def test_rotate_master_key_rotates_search_tools(monkeypatch): - """Master-key rotation re-encrypts the search tools table (step 4e).""" from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock