mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 4b1937eba8 into f285229b51
This commit is contained in:
commit
847f6b6c56
8 changed files with 420 additions and 4 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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",), ()),
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue