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.
This commit is contained in:
Yucheng He 2026-09-28 14:27:42 -07:00
parent 6d4ccf7e97
commit fc801cd870
8 changed files with 331 additions and 3 deletions

View file

@ -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),

View file

@ -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",), ()),
]

View file

@ -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()

View file

@ -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

View file

@ -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())

View file

@ -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

View file

@ -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"))

View file

@ -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,