fix: merge websearch tool params (#32162)

* fix: pass websearch tool params

* fix: load db websearch tool params

* fix: merge search tools in proxy

* fix: satisfy websearch lint budget

* fix: enforce websearch tool auth

* fix: preserve search tools on empty sync

* chore: rerun circleci
This commit is contained in:
Krrish Dholakia 2026-07-04 19:24:35 -07:00 • committed by GitHub
parent ed07aec89f
commit 2967bc9bef
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 455 additions and 127 deletions

View file

@ -91,6 +91,7 @@ class WebSearchInterceptionLogger(CustomLogger):
messages: List[Dict],
tools: Optional[List[Dict]],
custom_llm_provider: Optional[str],
kwargs: Optional[dict[str, Any]] = None,
) -> Optional[Dict[str, Any]]:
"""
Short-circuit web-search-only requests by executing the search directly.
@ -176,7 +177,10 @@ class WebSearchInterceptionLogger(CustomLogger):
# Execute search — keep the structured SearchResponse so the native
# block can carry per-result url/title/page_age.
try:
search_result_text, structured = await self._execute_search(query)
if kwargs is None:
search_result_text, structured = await self._execute_search(query)
else:
search_result_text, structured = await self._execute_search(query, kwargs=kwargs)
except Exception as e:
verbose_logger.error(f"WebSearchInterception: Short-circuit search failed: {e}")
search_result_text, structured = f"Search failed: {e}", None
@ -936,7 +940,7 @@ class WebSearchInterceptionLogger(CustomLogger):
query = tool_call["input"].get("query")
if query:
verbose_logger.debug(f"WebSearchInterception: Queuing search for query='{query}'")
search_tasks.append(self._execute_search(query))
search_tasks.append(self._execute_search(query, kwargs=kwargs))
else:
verbose_logger.debug(f"WebSearchInterception: Tool call {tool_call['id']} has no query")
# Add empty result for tools without query
@ -1009,7 +1013,9 @@ class WebSearchInterceptionLogger(CustomLogger):
)
return patch, structured_results
async def _execute_search(self, query: str) -> Tuple[str, Optional[SearchResponse]]:
async def _execute_search(
self, query: str, kwargs: Optional[dict[str, Any]] = None
) -> Tuple[str, Optional[SearchResponse]]:
"""
Execute a single web search using router's search tools.
@ -1031,36 +1037,13 @@ class WebSearchInterceptionLogger(CustomLogger):
)
llm_router = None
# Determine search provider from router's search_tools
search_tool = self._select_search_tool_from_router(llm_router=llm_router)
search_provider: Optional[str] = None
if llm_router is not None and hasattr(llm_router, "search_tools"):
if self.search_tool_name:
# Find specific search tool by name
matching_tools = [
tool
for tool in llm_router.search_tools
if tool.get("search_tool_name") == self.search_tool_name
]
if matching_tools:
search_tool = matching_tools[0]
search_provider = search_tool.get("litellm_params", {}).get("search_provider")
verbose_logger.debug(
f"WebSearchInterception: Found search tool '{self.search_tool_name}' "
f"with provider '{search_provider}'"
)
else:
verbose_logger.debug(
f"WebSearchInterception: Search tool '{self.search_tool_name}' not found in router, "
"falling back to first available or perplexity"
)
# If no specific tool or not found, use first available
if not search_provider and llm_router.search_tools:
first_tool = llm_router.search_tools[0]
search_provider = first_tool.get("litellm_params", {}).get("search_provider")
verbose_logger.debug(
f"WebSearchInterception: Using first available search tool with provider '{search_provider}'"
)
search_litellm_params: dict[str, Any] = {}
if search_tool is not None:
await self._authorize_search_tool(search_tool=search_tool, kwargs=kwargs)
search_litellm_params = dict(search_tool.get("litellm_params", {}) or {})
search_provider = search_litellm_params.get("search_provider")
# Fallback to perplexity if no router or no search tools configured
if not search_provider:
@ -1073,7 +1056,12 @@ class WebSearchInterceptionLogger(CustomLogger):
verbose_logger.debug(
f"WebSearchInterception: Executing search for '{query}' using provider '{search_provider}'"
)
result = await litellm.asearch(query=query, search_provider=search_provider)
search_kwargs = {
key: value
for key, value in search_litellm_params.items()
if key != "search_provider" and value is not None
}
result = await litellm.asearch(query=query, search_provider=search_provider, **search_kwargs)
# Format using transformation function
search_result_text = WebSearchTransformation.format_search_response(result)
@ -1086,6 +1074,107 @@ class WebSearchInterceptionLogger(CustomLogger):
verbose_logger.error(f"WebSearchInterception: Search failed for '{query}': {str(e)}")
raise
async def _authorize_search_tool(
self,
search_tool: dict[str, Any],
kwargs: Optional[dict[str, Any]],
) -> None:
search_tool_name = search_tool.get("search_tool_name")
if not isinstance(search_tool_name, str) or not search_tool_name:
return
user_api_key_auth = self._get_user_api_key_auth_from_kwargs(kwargs)
if user_api_key_auth is None:
return
from litellm.proxy.auth.auth_checks import (
can_key_call_search_tool,
can_team_call_search_tool,
get_team_object,
)
await can_key_call_search_tool(
search_tool_name=search_tool_name,
valid_token=user_api_key_auth,
)
team_id = getattr(user_api_key_auth, "team_id", None)
if team_id:
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
team_object = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=getattr(user_api_key_auth, "parent_otel_span", None),
proxy_logging_obj=proxy_logging_obj,
)
await can_team_call_search_tool(
search_tool_name=search_tool_name,
team_object=team_object,
)
@staticmethod
def _get_user_api_key_auth_from_kwargs(kwargs: Optional[dict[str, Any]]) -> Any:
if not kwargs:
return None
for metadata_key in ("metadata", "litellm_metadata"):
metadata = kwargs.get(metadata_key)
if isinstance(metadata, dict) and metadata.get("user_api_key_auth") is not None:
return metadata["user_api_key_auth"]
litellm_params = kwargs.get("litellm_params")
if not isinstance(litellm_params, dict):
return None
for metadata_key in ("metadata", "litellm_metadata"):
metadata = litellm_params.get(metadata_key)
if isinstance(metadata, dict) and metadata.get("user_api_key_auth") is not None:
return metadata["user_api_key_auth"]
return None
def _select_search_tool_from_router(self, llm_router: Any) -> Optional[dict[str, Any]]:
if llm_router is None or not hasattr(llm_router, "search_tools"):
return None
search_tools = list(getattr(llm_router, "search_tools") or [])
return self._select_search_tool_from_list(search_tools=search_tools, source="router")
def _select_search_tool_from_list(
self,
search_tools: list[dict[str, Any]],
source: str,
) -> Optional[dict[str, Any]]:
if self.search_tool_name:
matching_tools = [tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name]
if matching_tools:
search_provider = (matching_tools[0].get("litellm_params", {}) or {}).get("search_provider")
verbose_logger.debug(
f"WebSearchInterception: Found search tool '{self.search_tool_name}' "
f"from {source} with provider '{search_provider}'"
)
return matching_tools[0]
verbose_logger.debug(
f"WebSearchInterception: Search tool '{self.search_tool_name}' not found in {source}, "
"falling back to first available or perplexity"
)
if search_tools:
first_tool = search_tools[0]
search_provider = (first_tool.get("litellm_params", {}) or {}).get("search_provider")
verbose_logger.debug(
f"WebSearchInterception: Using first available search tool from {source} "
f"with provider '{search_provider}'"
)
return first_tool
return None
async def _execute_chat_completion_agentic_loop(
self,
model: str,
@ -1145,7 +1234,7 @@ class WebSearchInterceptionLogger(CustomLogger):
if query:
verbose_logger.debug(f"WebSearchInterception: Queuing search for query='{query}'")
search_tasks.append(self._execute_search(query))
search_tasks.append(self._execute_search(query, kwargs=kwargs))
else:
verbose_logger.debug(f"WebSearchInterception: Tool call {tool_call.get('id')} has no query")
# Add empty result for tools without query

View file

@ -148,6 +148,7 @@ async def _try_websearch_short_circuit(
tools: Optional[List[Dict]],
custom_llm_provider: Optional[str],
stream: Optional[bool],
kwargs: Optional[dict] = None,
) -> Optional[Union[AnthropicMessagesResponse, AsyncIterator]]:
"""
Attempt to short-circuit a web-search-only request.
@ -177,6 +178,7 @@ async def _try_websearch_short_circuit(
messages=messages,
tools=tools,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
)
if response is not None:
anthropic_response = cast(AnthropicMessagesResponse, response)
@ -292,6 +294,7 @@ async def anthropic_messages(
tools=tools,
custom_llm_provider=custom_llm_provider,
stream=original_stream,
kwargs={**kwargs, "metadata": metadata},
)
if short_circuit_response is not None:
return short_circuit_response

View file

@ -6382,7 +6382,6 @@ class ProxyConfig:
async def _init_search_tools_in_db(self, prisma_client: PrismaClient):
"""
Initialize search tools from database into the router on startup.
Only updates router if there are tools in the database, otherwise preserves config-loaded tools.
"""
global llm_router
@ -6392,26 +6391,29 @@ class ProxyConfig:
from litellm.router_utils.search_api_router import SearchAPIRouter
try:
search_tools = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client)
db_search_tools = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client)
verbose_proxy_logger.info(f"Loading {len(search_tools)} search tool(s) from database into router")
parsed_tools = self.parse_search_tools(self.get_config_state())
config_search_tools = parsed_tools or []
# Only update router if there are tools in the database
# This prevents overwriting config-loaded tools with an empty list
if len(search_tools) > 0:
if llm_router is not None:
# Add search tools to the router
await SearchAPIRouter.update_router_search_tools(
router_instance=llm_router, search_tools=search_tools
)
verbose_proxy_logger.info(f"Successfully loaded {len(search_tools)} search tool(s) into router")
else:
verbose_proxy_logger.debug(
"Router not initialized yet, search tools will be added when router is created"
)
search_tools = self._merge_config_and_db_search_tools(
config_search_tools=config_search_tools,
db_search_tools=[dict(tool) for tool in db_search_tools],
)
verbose_proxy_logger.info(
f"Loading {len(search_tools)} search tool(s) into router "
f"({len(config_search_tools)} from config, {len(db_search_tools)} from database)"
)
if llm_router is not None and search_tools:
await SearchAPIRouter.update_router_search_tools(router_instance=llm_router, search_tools=search_tools)
verbose_proxy_logger.info(f"Successfully loaded {len(search_tools)} search tool(s) into router")
elif llm_router is not None:
verbose_proxy_logger.debug("No search tools found in config or database, skipping router update")
else:
verbose_proxy_logger.debug(
"No search tools found in database, keeping config-loaded search tools (if any)"
"Router not initialized yet, search tools will be added when router is created"
)
except Exception as e:
@ -6419,6 +6421,21 @@ class ProxyConfig:
"litellm.proxy.proxy_server.py::ProxyConfig:_init_search_tools_in_db - {}".format(str(e))
)
@staticmethod
def _merge_config_and_db_search_tools(
config_search_tools: list[SearchToolTypedDict],
db_search_tools: list[dict[str, Any]],
) -> list[dict[str, Any]]:
db_tool_names = {tool.get("search_tool_name") for tool in db_search_tools}
return [
*[
dict(config_search_tool)
for config_search_tool in config_search_tools
if config_search_tool.get("search_tool_name") not in db_tool_names
],
*db_search_tools,
]
async def _init_pass_through_endpoints_in_db(self):
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
initialize_pass_through_endpoints_in_db,

View file

@ -12,6 +12,8 @@ from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
from litellm.llms.base_llm.search.transformation import SearchResponse
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, ProxyException, UserAPIKeyAuth
from litellm.types.utils import LlmProviders
@ -56,9 +58,7 @@ def test_initialize_from_proxy_config_honors_dict_callback_specific_params():
"""A valid dict under callback_settings.websearch_interception is applied."""
logger = WebSearchInterceptionLogger.initialize_from_proxy_config(
litellm_settings={},
callback_specific_params={
"websearch_interception": {"search_tool_name": "ws-tool"}
},
callback_specific_params={"websearch_interception": {"search_tool_name": "ws-tool"}},
)
assert logger.search_tool_name == "ws-tool"
@ -119,9 +119,7 @@ async def test_async_build_agentic_loop_plan_returns_request_patch():
"response_format": "anthropic",
}
logging_obj = MagicMock()
logging_obj.model_call_details = {
"agentic_loop_params": {"model": "bedrock/invoke/claude-3-5-sonnet"}
}
logging_obj.model_call_details = {"agentic_loop_params": {"model": "bedrock/invoke/claude-3-5-sonnet"}}
kwargs = {
"temperature": 0.2,
"_websearch_interception_converted_stream": True,
@ -162,8 +160,6 @@ async def test_internal_flags_filtered_from_followup_kwargs():
to the follow-up LLM request, causing "Extra inputs are not permitted" errors
from providers like Bedrock that use strict parameter validation.
"""
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
# Simulate kwargs that would be passed during agentic loop execution
kwargs_with_internal_flags = {
"_websearch_interception_converted_stream": True,
@ -174,9 +170,7 @@ async def test_internal_flags_filtered_from_followup_kwargs():
# Apply the same filtering logic used in _execute_agentic_loop
kwargs_for_followup = {
k: v
for k, v in kwargs_with_internal_flags.items()
if not k.startswith("_websearch_interception")
k: v for k, v in kwargs_with_internal_flags.items() if not k.startswith("_websearch_interception")
}
# Verify internal flags are filtered out
@ -188,6 +182,138 @@ async def test_internal_flags_filtered_from_followup_kwargs():
assert kwargs_for_followup["max_tokens"] == 1024
@pytest.mark.asyncio
async def test_execute_search_passes_selected_search_tool_litellm_params(monkeypatch):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(
enabled_providers=["bedrock"],
search_tool_name="ui-tavily",
)
router = MagicMock()
router.search_tools = [
{
"search_tool_name": "ui-tavily",
"litellm_params": {
"search_provider": "tavily",
"api_key": "fake-ui-key",
"api_base": "https://api.tavily.com",
"timeout": 10.0,
"max_retries": 2,
"country": None,
},
}
]
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
user_api_key_auth = UserAPIKeyAuth(
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="op-allowed-search",
search_tools=["ui-tavily"],
)
)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
await logger._execute_search(
"what is litellm",
kwargs={"litellm_params": {"metadata": {"user_api_key_auth": user_api_key_auth}}},
)
mock_asearch.assert_awaited_once_with(
query="what is litellm",
search_provider="tavily",
api_key="fake-ui-key",
api_base="https://api.tavily.com",
timeout=10.0,
max_retries=2,
)
@pytest.mark.asyncio
async def test_execute_search_enforces_key_search_tool_permission(monkeypatch):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(
enabled_providers=["bedrock"],
search_tool_name="blocked-search",
)
router = MagicMock()
router.search_tools = [
{
"search_tool_name": "blocked-search",
"litellm_params": {
"search_provider": "tavily",
"api_key": "fake-ui-key",
},
}
]
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
user_api_key_auth = UserAPIKeyAuth(
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="op-key-search",
search_tools=["allowed-search"],
)
)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(ProxyException):
await logger._execute_search(
"what is litellm",
kwargs={"metadata": {"user_api_key_auth": user_api_key_auth}},
)
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_execute_search_enforces_team_search_tool_permission(monkeypatch):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(
enabled_providers=["bedrock"],
search_tool_name="blocked-search",
)
router = MagicMock()
router.search_tools = [
{
"search_tool_name": "blocked-search",
"litellm_params": {
"search_provider": "tavily",
"api_key": "fake-ui-key",
},
}
]
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
team_key_auth = UserAPIKeyAuth(team_id="team-1")
team_object = LiteLLM_TeamTable(
team_id="team-1",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="op-team-search",
search_tools=["allowed-search"],
),
)
mock_get_team_object = AsyncMock(return_value=team_object)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_object", mock_get_team_object)
with pytest.raises(ProxyException):
await logger._execute_search(
"what is litellm",
kwargs={"metadata": {"user_api_key_auth": team_key_auth}},
)
mock_get_team_object.assert_awaited_once()
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_async_pre_call_deployment_hook_provider_from_top_level_kwargs():
"""Test that async_pre_call_deployment_hook finds custom_llm_provider at top-level kwargs.
@ -216,15 +342,12 @@ async def test_async_pre_call_deployment_hook_provider_from_top_level_kwargs():
assert result is not None
# The web_search tool should be converted to litellm_web_search (OpenAI format)
assert any(
t.get("type") == "function"
and t.get("function", {}).get("name") == "litellm_web_search"
t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search"
for t in result["tools"]
)
# The non-web-search tool should be preserved
assert any(
t.get("type") == "function"
and t.get("function", {}).get("name") == "other_tool"
for t in result["tools"]
t.get("type") == "function" and t.get("function", {}).get("name") == "other_tool" for t in result["tools"]
)
@ -261,8 +384,7 @@ async def test_async_pre_call_deployment_hook_returns_full_kwargs():
assert result["custom_llm_provider"] == "openai"
# Tools should be converted
assert any(
t.get("type") == "function"
and t.get("function", {}).get("name") == "litellm_web_search"
t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search"
for t in result["tools"]
)
@ -323,8 +445,7 @@ async def test_async_pre_call_deployment_hook_nested_litellm_params_fallback():
assert result is not None
assert any(
t.get("type") == "function"
and t.get("function", {}).get("name") == "litellm_web_search"
t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search"
for t in result["tools"]
)
# Full kwargs preserved
@ -357,8 +478,7 @@ async def test_async_pre_call_deployment_hook_provider_derived_from_model_name()
# Should NOT be None — the hook should derive "openai" from "openai/gpt-4o-mini"
assert result is not None
assert any(
t.get("type") == "function"
and t.get("function", {}).get("name") == "litellm_web_search"
t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search"
for t in result["tools"]
)
# Full kwargs preserved
@ -478,9 +598,7 @@ def test_sync_forced_tool_choice_leaves_non_forced_untouched(tool_choice):
and None pass through unchanged."""
converted_tools = [{"name": LITELLM_WEB_SEARCH_TOOL_NAME}]
result = WebSearchInterceptionLogger._sync_forced_tool_choice(
tool_choice, converted_tools
)
result = WebSearchInterceptionLogger._sync_forced_tool_choice(tool_choice, converted_tools)
assert result == tool_choice

View file

@ -189,9 +189,7 @@ def test_ProxyConfig__load_yaml_file_raises_on_missing_file():
@pytest.mark.asyncio
async def test_ProxyConfig__get_config_from_file_loads_yaml(tmp_path):
f = tmp_path / "c.yaml"
f.write_text(
"model_list: []\ngeneral_settings: {}\nlitellm_settings:\n drop_params: true\n"
)
f.write_text("model_list: []\ngeneral_settings: {}\nlitellm_settings:\n drop_params: true\n")
pc = ProxyConfig()
result = await pc._get_config_from_file(config_file_path=str(f))
assert result == {
@ -534,6 +532,139 @@ def test_ProxyConfig_parse_search_tools_missing_returns_none():
assert pc.parse_search_tools({}) is None
def test_ProxyConfig_merge_config_and_db_search_tools_returns_superset():
config_tools = [
{
"search_tool_name": "config-search",
"litellm_params": {"search_provider": "tavily"},
}
]
db_tools = [
{
"search_tool_name": "db-search",
"litellm_params": {
"search_provider": "exa_ai",
"api_key": "fake-db-key",
},
}
]
merged = ProxyConfig._merge_config_and_db_search_tools(
config_search_tools=config_tools,
db_search_tools=db_tools,
)
assert [tool["search_tool_name"] for tool in merged] == ["config-search", "db-search"]
assert merged[1]["litellm_params"]["api_key"] == "fake-db-key"
def test_ProxyConfig_merge_config_and_db_search_tools_prefers_db_duplicate():
config_tools = [
{
"search_tool_name": "shared-search",
"litellm_params": {"search_provider": "tavily"},
},
{
"search_tool_name": "config-only",
"litellm_params": {"search_provider": "perplexity"},
},
]
db_tools = [
{
"search_tool_name": "shared-search",
"litellm_params": {
"search_provider": "exa_ai",
"api_key": "fake-db-key",
},
}
]
merged = ProxyConfig._merge_config_and_db_search_tools(
config_search_tools=config_tools,
db_search_tools=db_tools,
)
assert [tool["search_tool_name"] for tool in merged] == ["config-only", "shared-search"]
assert merged[1]["litellm_params"]["search_provider"] == "exa_ai"
assert merged[1]["litellm_params"]["api_key"] == "fake-db-key"
@pytest.mark.asyncio
async def test_ProxyConfig__init_search_tools_in_db_loads_merged_tools(monkeypatch):
from litellm.proxy import proxy_server
from litellm.router_utils.search_api_router import SearchAPIRouter
pc = ProxyConfig()
pc.update_config_state(
{
"search_tools": [
{
"search_tool_name": "shared-search",
"litellm_params": {"search_provider": "tavily"},
},
{
"search_tool_name": "config-only",
"litellm_params": {"search_provider": "perplexity"},
},
]
}
)
db_tools = [
{
"search_tool_name": "shared-search",
"litellm_params": {
"search_provider": "exa_ai",
"api_key": "fake-db-key",
},
}
]
fake_router = MagicMock()
mock_get_db_tools = AsyncMock(return_value=db_tools)
mock_update_router = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", fake_router)
monkeypatch.setattr(
"litellm.proxy.search_endpoints.search_tool_registry.SearchToolRegistry.get_all_search_tools_from_db",
mock_get_db_tools,
)
monkeypatch.setattr(SearchAPIRouter, "update_router_search_tools", mock_update_router)
await pc._init_search_tools_in_db(prisma_client=MagicMock())
mock_get_db_tools.assert_awaited_once()
mock_update_router.assert_awaited_once()
update_kwargs = mock_update_router.await_args.kwargs
assert update_kwargs["router_instance"] is fake_router
assert [tool["search_tool_name"] for tool in update_kwargs["search_tools"]] == [
"config-only",
"shared-search",
]
assert update_kwargs["search_tools"][1]["litellm_params"]["api_key"] == "fake-db-key"
@pytest.mark.asyncio
async def test_ProxyConfig__init_search_tools_in_db_skips_empty_router_update(monkeypatch):
from litellm.proxy import proxy_server
from litellm.router_utils.search_api_router import SearchAPIRouter
pc = ProxyConfig()
pc.update_config_state({})
mock_get_db_tools = AsyncMock(return_value=[])
mock_update_router = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", MagicMock())
monkeypatch.setattr(
"litellm.proxy.search_endpoints.search_tool_registry.SearchToolRegistry.get_all_search_tools_from_db",
mock_get_db_tools,
)
monkeypatch.setattr(SearchAPIRouter, "update_router_search_tools", mock_update_router)
await pc._init_search_tools_in_db(prisma_client=MagicMock())
mock_get_db_tools.assert_awaited_once()
mock_update_router.assert_not_awaited()
# ---------------------------------------------------------------------------
# ProxyConfig._load_environment_variables
# ---------------------------------------------------------------------------
@ -542,9 +673,7 @@ def test_ProxyConfig_parse_search_tools_missing_returns_none():
def test_ProxyConfig__load_environment_variables_sets_env(monkeypatch):
monkeypatch.delenv("TEST_LOAD_ENV_X", raising=False)
pc = ProxyConfig()
pc._load_environment_variables(
{"environment_variables": {"TEST_LOAD_ENV_X": "hello"}}
)
pc._load_environment_variables({"environment_variables": {"TEST_LOAD_ENV_X": "hello"}})
result = {
"TEST_LOAD_ENV_X": os.environ.get("TEST_LOAD_ENV_X"),
"set": True,
@ -602,9 +731,7 @@ async def test_ProxyConfig_load_config_missing_file_raises(monkeypatch):
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_forwards_callback_specific_params(
tmp_path, monkeypatch
):
async def test_ProxyConfig_load_config_forwards_callback_specific_params(tmp_path, monkeypatch):
"""Regression: callback_settings from config must be forwarded to
initialize_callbacks_on_proxy as callback_specific_params.
@ -645,16 +772,12 @@ async def test_ProxyConfig_load_config_forwards_callback_specific_params(
# The callbacks branch must forward the loaded callback_settings.
assert captured.get("callback_specific_params") == {
"datadog_cost_management": {
"cost_tag_keys": ["capability", "platform", "ai_product"]
}
"datadog_cost_management": {"cost_tag_keys": ["capability", "platform", "ai_product"]}
}
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_blank_callback_settings_does_not_crash(
tmp_path, monkeypatch
):
async def test_ProxyConfig_load_config_blank_callback_settings_does_not_crash(tmp_path, monkeypatch):
"""Regression: `callback_settings:` with no body loads as None because
dict.get() only falls back to the default when the key is absent. The None
was forwarded verbatim to initialize_callbacks_on_proxy, where the first
@ -678,17 +801,13 @@ async def test_ProxyConfig_load_config_blank_callback_settings_does_not_crash(
CompressionInterceptionLogger,
)
original_callbacks = (
list(litellm.callbacks) if isinstance(litellm.callbacks, list) else []
)
original_callbacks = list(litellm.callbacks) if isinstance(litellm.callbacks, list) else []
litellm.callbacks = []
try:
pc = ProxyConfig()
await pc.load_config(router=None, config_file_path=str(f))
assert any(
isinstance(c, CompressionInterceptionLogger) for c in litellm.callbacks
)
assert any(isinstance(c, CompressionInterceptionLogger) for c in litellm.callbacks)
finally:
litellm.callbacks = original_callbacks
@ -1071,7 +1190,9 @@ def test_ProxyConfig_decrypt_model_list_from_db_resolves_env_refs_after_db_decry
lambda value, key, return_original_value: (
"os.environ/LITELLM_DB_MODEL_API_KEY"
if key == "api_key"
else "os.environ/LITELLM_MASTER_KEY" if key == "api_base" else value
else "os.environ/LITELLM_MASTER_KEY"
if key == "api_base"
else value
),
)
pc = ProxyConfig()
@ -1102,9 +1223,7 @@ def test_ProxyConfig_decrypt_model_list_from_db_keeps_team_env_refs_literal_afte
monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret")
monkeypatch.setattr(
"litellm.proxy.proxy_server.decrypt_value_helper",
lambda value, key, return_original_value: (
"os.environ/LITELLM_MASTER_KEY" if key == "api_key" else value
),
lambda value, key, return_original_value: "os.environ/LITELLM_MASTER_KEY" if key == "api_key" else value,
)
monkeypatch.setattr("litellm.proxy.proxy_server.get_secret", fail_on_call)
pc = ProxyConfig()
@ -1128,9 +1247,7 @@ def test_ProxyConfig_decrypt_model_list_from_db_keeps_team_env_refs_literal_afte
def test_ProxyConfig_decrypt_model_list_from_db_invalid_params_skips():
pc = ProxyConfig()
bad = SimpleNamespace(
model_id="m-1", model_name="x", model_info={}, litellm_params="not-a-dict"
)
bad = SimpleNamespace(model_id="m-1", model_name="x", model_info={}, litellm_params="not-a-dict")
out = pc.decrypt_model_list_from_db(new_models=[bad])
# Invalid entries skipped — empty list returned.
assert out == []
@ -1180,9 +1297,7 @@ async def test_ProxyConfig__update_llm_router_bad_proxy_logging_raises(monkeypat
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-x")
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings", {"alerting": ["email"]}
)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"alerting": ["email"]})
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", pc)
# Passing None for proxy_logging_obj triggers AttributeError in _add_general_settings_from_db_config
# when it calls proxy_logging_obj.update_values.
@ -1433,9 +1548,7 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router():
fake_router.update_settings = MagicMock()
fake_prisma = MagicMock()
fake_prisma.db.litellm_config.find_first = AsyncMock(
return_value=SimpleNamespace(
param_value={"timeout": 30, "retries": 2, "fallbacks": []}
)
return_value=SimpleNamespace(param_value={"timeout": 30, "retries": 2, "fallbacks": []})
)
config_data = {"router_settings": {"timeout": 10}}
await pc._add_router_settings_from_db_config(
@ -1446,9 +1559,7 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router():
snapshot = {
"called": fake_router.update_settings.called,
"call_count": fake_router.update_settings.call_count,
"kwargs_keys": sorted(
list(fake_router.update_settings.call_args.kwargs.keys())
),
"kwargs_keys": sorted(list(fake_router.update_settings.call_args.kwargs.keys())),
}
assert snapshot == {
"called": True,
@ -1461,9 +1572,7 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router():
async def test_ProxyConfig__add_router_settings_from_db_config_none_router_noop():
pc = ProxyConfig()
# No router and no prisma — should silently return.
await pc._add_router_settings_from_db_config(
config_data={}, llm_router=None, prisma_client=None
)
await pc._add_router_settings_from_db_config(config_data={}, llm_router=None, prisma_client=None)
# Error-style: bad call signature raises.
with pytest.raises(TypeError):
await pc._add_router_settings_from_db_config() # type: ignore[call-arg]
@ -1568,9 +1677,7 @@ async def test_ProxyConfig__update_general_settings_updates_max_parallel(monkeyp
snapshot = {
"max_parallel_requests": ps.general_settings.get("max_parallel_requests"),
"global_max_parallel_requests": ps.general_settings.get(
"global_max_parallel_requests"
),
"global_max_parallel_requests": ps.general_settings.get("global_max_parallel_requests"),
"ui_access_mode": ps.general_settings.get("ui_access_mode"),
}
assert snapshot == {
@ -1712,9 +1819,7 @@ async def test_ProxyConfig__update_config_from_db_does_not_log_general_settings_
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_redacts_secret_litellm_setting_keeps_plain(
tmp_path, monkeypatch
):
async def test_ProxyConfig_load_config_redacts_secret_litellm_setting_keeps_plain(tmp_path, monkeypatch):
"""Regression for LIT-4152 on the ``litellm_settings`` apply loop.
``load_config`` logged ``setting litellm.<key>=<value>`` verbatim at DEBUG,
@ -1735,11 +1840,7 @@ async def test_ProxyConfig_load_config_redacts_secret_litellm_setting_keeps_plai
api_key_secret = "sk-lit4152-litellm-settings-secret-abcdef1234567890"
f = tmp_path / "c.yaml"
f.write_text(
"model_list: []\n"
"general_settings: {}\n"
"litellm_settings:\n"
f" api_key: {api_key_secret}\n"
" num_retries: 7\n"
f"model_list: []\ngeneral_settings: {{}}\nlitellm_settings:\n api_key: {api_key_secret}\n num_retries: 7\n"
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)