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,