This commit is contained in:
yucheng-berri 2026-09-30 16:55:20 -04:00 • committed by GitHub
commit 847f6b6c56
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 420 additions and 4 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 (
@ -5340,6 +5341,11 @@ 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))
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
@ -32,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: ...
@ -48,6 +51,73 @@ 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_tool_row(
table: SearchToolTableClient, search_tool_id: str, stored_litellm_params: Mapping[str, object], new_master_key: str
) -> None:
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()
}
)
},
)
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:
"""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 while it keeps being 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:
"""
Handles adding, removing, and getting search tools in DB + in memory.
@ -59,7 +129,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
@ -67,7 +137,15 @@ class SearchToolRegistry:
Returns:
Dict with datetime fields converted to ISO strings
"""
result: Final = dict(prisma_obj)
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()
@ -92,7 +170,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 +242,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,6 @@
import contextlib
import json
from types import SimpleNamespace
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
@ -1148,3 +1150,228 @@ 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(SimpleNamespace):
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
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(
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
@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
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

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,62 @@ async def test_rotate_master_key_rotates_sso_identity_assertions(
)
@pytest.mark.asyncio
async def test_rotate_master_key_rotates_search_tools(monkeypatch):
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")
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):
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
row.litellm_params = json.loads(data["litellm_params"])
return 1
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])
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",
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(row.litellm_params["api_key"], "sk-new-master-key") == "tvly-secret"
assert row.litellm_params["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,