mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(search_tools): encrypt search tool litellm_params at rest (#43631)
* 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. * 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. * 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. * 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. * test(search_tools): drop the rotation test docstring * Store search tool params as written when no encryption key is configured * Rotate search tools under the salt key, keep non-ciphertext values and loaded tools that do not decrypt * Treat a search tool as undecryptable only when its provider is ciphertext-length * Drop suppressions the type discipline gate on main now reports as unused * Show the loaded search tool in the admin list and info views when its DB params do not decrypt * Keep the DB row's other fields when the admin views substitute loaded params
This commit is contained in:
parent
5d42cb7cfa
commit
12d4b75b7a
11 changed files with 696 additions and 10 deletions
|
|
@ -39,6 +39,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 (
|
||||
|
|
@ -5339,6 +5340,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()
|
||||
|
|
|
|||
|
|
@ -9011,11 +9011,15 @@ class ProxyConfig:
|
|||
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import (
|
||||
SearchToolRegistry,
|
||||
keep_loaded_search_tools_that_do_not_decrypt,
|
||||
)
|
||||
from litellm.router_utils.search_api_router import SearchAPIRouter
|
||||
|
||||
try:
|
||||
db_search_tools: Final = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client)
|
||||
db_search_tools: Final = keep_loaded_search_tools_that_do_not_decrypt(
|
||||
await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client),
|
||||
loaded_search_tools=llm_router.search_tools if llm_router is not None else (),
|
||||
)
|
||||
|
||||
parsed_tools: Final = self.parse_search_tools(self.get_config_state())
|
||||
config_search_tools: Final = parsed_tools or []
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
CRUD ENDPOINTS FOR SEARCH TOOLS
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any, Final, TypeAlias
|
||||
|
||||
|
|
@ -17,7 +17,10 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import SearchToolRegistry
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import (
|
||||
SearchToolRegistry,
|
||||
keep_loaded_search_tools_that_do_not_decrypt,
|
||||
)
|
||||
from litellm.types.search import (
|
||||
ListSearchToolsResponse,
|
||||
SearchTool,
|
||||
|
|
@ -65,6 +68,18 @@ async def _refresh_router_search_tools() -> None:
|
|||
verbose_proxy_logger.exception("Search tool router refresh failed after a management write: %s", e)
|
||||
|
||||
|
||||
def _with_loaded_tools_where_undecryptable(db_search_tools: Sequence[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
kept_search_tools: Final = keep_loaded_search_tools_that_do_not_decrypt(
|
||||
db_search_tools, loaded_search_tools=llm_router.search_tools if llm_router is not None else ()
|
||||
)
|
||||
return [
|
||||
{**db_tool, "litellm_params": kept_tool.get("litellm_params")}
|
||||
for db_tool, kept_tool in zip(db_search_tools, kept_search_tools, strict=True)
|
||||
]
|
||||
|
||||
|
||||
async def _team_object_from_db(team_id: str, user_api_key_dict: UserAPIKeyAuth) -> LiteLLM_TeamTable:
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
|
|
@ -187,7 +202,9 @@ async def list_search_tools(
|
|||
raise HTTPException(status_code=500, detail="Prisma client not initialized")
|
||||
|
||||
try:
|
||||
search_tools_from_db = await SEARCH_TOOL_REGISTRY.get_all_search_tools_from_db(prisma_client=prisma_client)
|
||||
search_tools_from_db = _with_loaded_tools_where_undecryptable(
|
||||
await SEARCH_TOOL_REGISTRY.get_all_search_tools_from_db(prisma_client=prisma_client)
|
||||
)
|
||||
|
||||
db_tool_names: Final = {tool.get("search_tool_name") for tool in search_tools_from_db}
|
||||
|
||||
|
|
@ -514,15 +531,16 @@ async def get_search_tool_info(search_tool_id: str):
|
|||
raise HTTPException(status_code=500, detail="Prisma client not initialized")
|
||||
|
||||
try:
|
||||
result: Final = await SEARCH_TOOL_REGISTRY.get_search_tool_by_id_from_db(
|
||||
db_result: Final = await SEARCH_TOOL_REGISTRY.get_search_tool_by_id_from_db(
|
||||
search_tool_id=search_tool_id, prisma_client=prisma_client
|
||||
)
|
||||
|
||||
if result is None:
|
||||
if db_result is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Search tool with ID {search_tool_id} not found",
|
||||
)
|
||||
result: Final = _with_loaded_tools_where_undecryptable((db_result,))[0]
|
||||
|
||||
# Mask sensitive data
|
||||
litellm_params_dict: Final = dict(result.get("litellm_params", {}))
|
||||
|
|
|
|||
|
|
@ -2,16 +2,26 @@
|
|||
Search Tool Registry for managing search tool configurations.
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final, Protocol
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy.auth.master_key_boot_check import SALT_KEY_ENV_VAR
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
_get_salt_key,
|
||||
decrypt_if_encrypted_with,
|
||||
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
|
||||
from litellm.types.search import SearchTool
|
||||
from litellm.types.utils import SearchProviders
|
||||
|
||||
|
||||
class SearchToolRecord(Protocol):
|
||||
|
|
@ -32,6 +42,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 +60,136 @@ def _search_tools_table(prisma_client: PrismaClient) -> SearchToolTableClient:
|
|||
return _search_tools_table_of(SearchToolsRepository(prisma_client))
|
||||
|
||||
|
||||
_STORED_LITELLM_PARAMS: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _stored_litellm_params(row: SearchToolRecord) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _STORED_LITELLM_PARAMS.validate_python(dict(row).get("litellm_params"))
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _encrypted_search_tool_value(value: object) -> object:
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
try:
|
||||
return encrypt_value_helper(value=value)
|
||||
except Exception: # noqa: BLE001 # no salt key or master key configured: store the value as written
|
||||
return value
|
||||
|
||||
|
||||
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: _encrypted_search_tool_value(value) for key, value in litellm_params.items()}
|
||||
|
||||
|
||||
def _search_tool_plaintext(value: str) -> str | None:
|
||||
signing_key: Final = _get_salt_key()
|
||||
return None if signing_key is None else decrypt_if_encrypted_with(value, signing_key)
|
||||
|
||||
|
||||
def _decrypted_search_tool_value(value: object) -> object:
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
plaintext: Final = _search_tool_plaintext(value)
|
||||
return value if plaintext is None else plaintext
|
||||
|
||||
|
||||
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: _decrypted_search_tool_value(value) for key, value in litellm_params.items()}
|
||||
|
||||
|
||||
def _reencrypt_search_tool_value(value: object, encryption_key: str) -> object:
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
plaintext: Final = _search_tool_plaintext(value)
|
||||
return value if plaintext is None else encrypt_value_helper(value=plaintext, new_encryption_key=encryption_key)
|
||||
|
||||
|
||||
async def _rotate_search_tool_row(
|
||||
table: SearchToolTableClient, search_tool_id: str, stored_litellm_params: Mapping[str, object], encryption_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, encryption_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 _stored_litellm_params(reread)
|
||||
if 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
|
||||
|
||||
|
||||
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 the key in force after
|
||||
rotation (LITELLM_SALT_KEY when set, otherwise 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.
|
||||
"""
|
||||
salt_key: Final = os.environ.get(SALT_KEY_ENV_VAR)
|
||||
encryption_key: Final = new_master_key if salt_key is None else salt_key
|
||||
table: Final = _search_tools_table(prisma_client)
|
||||
for row in await table.find_many():
|
||||
stored_litellm_params = _stored_litellm_params(row)
|
||||
if stored_litellm_params is not None:
|
||||
await _rotate_search_tool_row(table, row.search_tool_id, stored_litellm_params, encryption_key)
|
||||
|
||||
|
||||
_KNOWN_SEARCH_PROVIDERS: Final = frozenset(provider.value for provider in SearchProviders)
|
||||
# An empty string encrypted with aes-256-gcm, the shortest ciphertext either algorithm produces
|
||||
_SHORTEST_CIPHERTEXT_LENGTH: Final = 47
|
||||
|
||||
|
||||
def _did_not_decrypt(search_tool: Mapping[str, object]) -> bool:
|
||||
litellm_params: Final = search_tool.get("litellm_params")
|
||||
search_provider: Final = litellm_params.get("search_provider") if isinstance(litellm_params, Mapping) else None
|
||||
return (
|
||||
isinstance(search_provider, str)
|
||||
and search_provider not in _KNOWN_SEARCH_PROVIDERS
|
||||
and len(search_provider) >= _SHORTEST_CIPHERTEXT_LENGTH
|
||||
)
|
||||
|
||||
|
||||
def keep_loaded_search_tools_that_do_not_decrypt(
|
||||
db_search_tools: Sequence[Mapping[str, object]], loaded_search_tools: Sequence[Mapping[str, object]]
|
||||
) -> Sequence[Mapping[str, object]]:
|
||||
"""Replace each DB search tool whose params do not decrypt with the current key by its loaded version."""
|
||||
loaded_by_id: Final = {tool.get("search_tool_id"): tool for tool in loaded_search_tools}
|
||||
kept: Final = tuple(
|
||||
loaded_by_id.get(tool.get("search_tool_id"), tool) if _did_not_decrypt(tool) else tool
|
||||
for tool in db_search_tools
|
||||
)
|
||||
for db_tool, kept_tool in zip(db_search_tools, kept):
|
||||
if kept_tool is not db_tool:
|
||||
verbose_proxy_logger.warning(
|
||||
"Search tool %s has litellm_params that do not decrypt with the current key; keeping the loaded "
|
||||
"version. Restart the proxy if the master key was rotated.",
|
||||
db_tool.get("search_tool_id"),
|
||||
)
|
||||
return kept
|
||||
|
||||
|
||||
class SearchToolRegistry:
|
||||
"""
|
||||
Handles adding, removing, and getting search tools in DB + in memory.
|
||||
|
|
@ -59,7 +201,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 +209,15 @@ class SearchToolRegistry:
|
|||
Returns:
|
||||
Dict with datetime fields converted to ISO strings
|
||||
"""
|
||||
result: Final = dict(prisma_obj)
|
||||
stored_litellm_params: Final = _stored_litellm_params(prisma_obj)
|
||||
result: Final = {
|
||||
**dict(prisma_obj),
|
||||
**(
|
||||
{"litellm_params": decrypt_search_tool_litellm_params(stored_litellm_params)}
|
||||
if stored_litellm_params is not None
|
||||
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 +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", {}))
|
||||
|
||||
# Create search tool in DB
|
||||
|
|
@ -162,7 +314,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,348 @@ 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.fixture
|
||||
def master_key_only(monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
|
||||
monkeypatch.setattr(ps, "master_key", "sk-old-master-key")
|
||||
monkeypatch.setattr(ps, "general_settings", {})
|
||||
return "sk-old-master-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_search_tool_is_stored_as_written_when_no_encryption_key_is_configured(monkeypatch):
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import SearchToolRegistry
|
||||
|
||||
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
|
||||
monkeypatch.setattr(ps, "master_key", None)
|
||||
monkeypatch.setattr(ps, "general_settings", {})
|
||||
table = _InMemorySearchToolsTable()
|
||||
|
||||
created = await SearchToolRegistry().add_search_tool_to_db(
|
||||
search_tool={"search_tool_name": "agentcore-search", "litellm_params": SECRET_PARAMS},
|
||||
prisma_client=_prisma_client_over(table),
|
||||
)
|
||||
|
||||
assert table.rows[created["search_tool_id"]].litellm_params == SECRET_PARAMS
|
||||
|
||||
|
||||
@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(master_key_only):
|
||||
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(master_key_only):
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_master_key_rotation_with_a_salt_key_keeps_search_tools_readable(salt_key, monkeypatch):
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import (
|
||||
SearchToolRegistry,
|
||||
rotate_search_tools_master_key,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(ps, "master_key", "sk-old-master-key")
|
||||
table = _InMemorySearchToolsTable(
|
||||
[
|
||||
_stored_row(
|
||||
"salted-id",
|
||||
"salted",
|
||||
{"search_provider": encrypt_value_helper("tavily"), "api_key": encrypt_value_helper("tvly-salted")},
|
||||
)
|
||||
]
|
||||
)
|
||||
prisma_client = _prisma_client_over(table)
|
||||
|
||||
await rotate_search_tools_master_key(prisma_client=prisma_client, new_master_key="sk-new-master-key")
|
||||
monkeypatch.setattr(ps, "master_key", "sk-new-master-key")
|
||||
|
||||
loaded = await SearchToolRegistry().get_search_tool_by_id_from_db("salted-id", prisma_client=prisma_client)
|
||||
assert loaded["litellm_params"] == {"search_provider": "tavily", "api_key": "tvly-salted"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("legacy_value", ["****", ".", "--", "*"])
|
||||
async def test_plaintext_values_that_are_not_base64_load_and_rotate_unchanged(salt_key, legacy_value):
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import (
|
||||
SearchToolRegistry,
|
||||
rotate_search_tools_master_key,
|
||||
)
|
||||
|
||||
legacy_params = {"search_provider": "perplexity", "api_key": legacy_value, "api_base": "https://api.perplexity.ai"}
|
||||
table = _InMemorySearchToolsTable([_stored_row("legacy-id", "legacy", dict(legacy_params))])
|
||||
prisma_client = _prisma_client_over(table)
|
||||
|
||||
loaded = await SearchToolRegistry().get_search_tool_by_id_from_db("legacy-id", prisma_client=prisma_client)
|
||||
await rotate_search_tools_master_key(prisma_client=prisma_client, new_master_key="sk-new-master-key")
|
||||
|
||||
assert loaded["litellm_params"] == legacy_params
|
||||
assert table.rows["legacy-id"].litellm_params == legacy_params
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_and_info_show_the_loaded_tool_when_db_params_do_not_decrypt(master_key_only):
|
||||
"""After /key/regenerate rewrites the rows and before a restart, the admin views read the loaded tool."""
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import SearchToolRegistry
|
||||
|
||||
rewritten_params = {
|
||||
"search_provider": encrypt_value_helper("perplexity", new_encryption_key="sk-new-master-key"),
|
||||
"api_key": encrypt_value_helper("pplx-loaded-key", new_encryption_key="sk-new-master-key"),
|
||||
"api_base": encrypt_value_helper("https://api.perplexity.ai", new_encryption_key="sk-new-master-key"),
|
||||
}
|
||||
table = _InMemorySearchToolsTable([_stored_row("rotated-id", "rotated", rewritten_params)])
|
||||
loaded_tool = {
|
||||
"search_tool_id": "rotated-id",
|
||||
"search_tool_name": "rotated",
|
||||
"litellm_params": {
|
||||
"search_provider": "perplexity",
|
||||
"api_key": "pplx-loaded-key",
|
||||
"api_base": "https://api.perplexity.ai",
|
||||
},
|
||||
}
|
||||
fake_router = MagicMock()
|
||||
fake_router.search_tools = [loaded_tool]
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", _prisma_client_over(table)
|
||||
), # test-quality-ok: proxy globals are the only seam; see the module note above
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.llm_router", fake_router
|
||||
), # test-quality-ok: proxy globals are the only seam; see the module note above
|
||||
patch( # test-quality-ok: proxy globals are the only seam; see the module note above
|
||||
"litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", SearchToolRegistry()
|
||||
),
|
||||
_override_auth(UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user")),
|
||||
):
|
||||
listed = TestClient(app).get("/search_tools/list")
|
||||
info = TestClient(app).get("/search_tools/rotated-id")
|
||||
|
||||
assert listed.status_code == 200
|
||||
assert info.status_code == 200
|
||||
listed_params = [tool["litellm_params"] for tool in listed.json()["search_tools"]]
|
||||
assert [params["search_provider"] for params in listed_params] == ["perplexity"]
|
||||
assert info.json()["litellm_params"]["search_provider"] == "perplexity"
|
||||
assert info.json()["litellm_params"]["api_base"] == listed_params[0]["api_base"] != rewritten_params["api_base"]
|
||||
assert "pplx-loaded-key" not in listed.text + info.text
|
||||
assert info.json()["created_at"] == listed.json()["search_tools"][0]["created_at"] == "2026-09-01T00:00:00"
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -18562,6 +18562,63 @@ 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.delenv("LITELLM_SALT_KEY", raising=False)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_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,
|
||||
|
|
|
|||
|
|
@ -2029,6 +2029,61 @@ async def test_ProxyConfig__init_search_tools_in_db_clears_router_when_last_tool
|
|||
assert fake_router.search_tools == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig__init_search_tools_in_db_keeps_loaded_tools_whose_params_do_not_decrypt(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
pc = ProxyConfig()
|
||||
pc.update_config_state({})
|
||||
loaded_tool = {
|
||||
"search_tool_id": "rotated-id",
|
||||
"search_tool_name": "rotated-search",
|
||||
"litellm_params": {"search_provider": "perplexity", "api_key": "pplx-loaded"},
|
||||
}
|
||||
fake_router = MagicMock()
|
||||
fake_router.search_tools = [
|
||||
loaded_tool,
|
||||
{
|
||||
"search_tool_id": "typo-id",
|
||||
"search_tool_name": "typo-search",
|
||||
"litellm_params": {"search_provider": "tavily"},
|
||||
},
|
||||
]
|
||||
db_tools = [
|
||||
{
|
||||
"search_tool_id": "rotated-id",
|
||||
"search_tool_name": "rotated-search",
|
||||
"litellm_params": {
|
||||
"search_provider": "zM9FVihBfZj0LRkl6_J4TeIEO8ijpxKov0QnfZa1uM9J1lO7Txy9IQ==",
|
||||
"api_key": "c2VhbGVkLWtleQ",
|
||||
},
|
||||
},
|
||||
{
|
||||
"search_tool_id": "fresh-id",
|
||||
"search_tool_name": "fresh-search",
|
||||
"litellm_params": {"search_provider": "tavily", "api_key": "tvly-fresh"},
|
||||
},
|
||||
{
|
||||
"search_tool_id": "typo-id",
|
||||
"search_tool_name": "typo-search",
|
||||
"litellm_params": {"search_provider": "Tavily", "api_key": "tvly-edited"},
|
||||
},
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", fake_router)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.search_endpoints.search_tool_registry.SearchToolRegistry.get_all_search_tools_from_db",
|
||||
AsyncMock(return_value=db_tools),
|
||||
)
|
||||
|
||||
await pc._init_search_tools_in_db(prisma_client=MagicMock())
|
||||
|
||||
assert [tool["litellm_params"] for tool in fake_router.search_tools] == [
|
||||
{"search_provider": "perplexity", "api_key": "pplx-loaded"},
|
||||
{"search_provider": "tavily", "api_key": "tvly-fresh"},
|
||||
{"search_provider": "Tavily", "api_key": "tvly-edited"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_reload_search_tools_from_db_refreshes_router(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue