mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
ed07aec89f
commit
2967bc9bef
5 changed files with 455 additions and 127 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue