fix(websearch): reject invalid explicit search tool selections

This commit is contained in:
George Pickett 2026-08-24 11:24:05 -07:00
parent 9bcc00b1f1
commit ce7fbf91b2
2 changed files with 412 additions and 18 deletions

View file

@ -95,6 +95,10 @@ class _SearchToolConfig(TypedDict, total=False):
litellm_params: Mapping[str, object] | None
class _SearchToolConfigurationError(ValueError):
"""An explicitly configured search tool cannot be safely executed."""
class _DeploymentKwargsView(TypedDict):
"""Typed reads of the untyped request kwargs seen by the deployment hook."""
@ -252,6 +256,8 @@ class WebSearchInterceptionLogger(CustomLogger):
search_result_text, structured = await self._execute_search(query)
else:
search_result_text, structured = await self._execute_search(query, kwargs=kwargs)
except _SearchToolConfigurationError:
raise
except Exception as e:
verbose_logger.error("WebSearchInterception: Short-circuit search failed: %s", e)
search_result_text, structured = f"Search failed: {e}", None
@ -329,15 +335,25 @@ class WebSearchInterceptionLogger(CustomLogger):
if not tools:
return None
if call_type in (CallTypes.responses, CallTypes.aresponses):
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
# Check if any tool is a web search tool (native or already LiteLLM standard)
has_websearch: Final = any(is_web_search_tool(t) for t in tools)
is_responses_call: Final = call_type in (CallTypes.responses, CallTypes.aresponses)
has_websearch: Final = (
any(is_web_search_tool_responses(tool) for tool in tools)
if is_responses_call
else any(is_web_search_tool(tool) for tool in tools)
)
if not has_websearch:
return None
if self.search_tool_name:
try:
from litellm.proxy.proxy_server import llm_router
except ImportError:
llm_router = None
self._select_search_tool_from_router(llm_router=llm_router)
if is_responses_call:
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
verbose_logger.debug("WebSearchInterception: Converting native web_search tools to LiteLLM standard")
# If the client sent an Anthropic-native web_search_* tool, mark the
@ -1148,6 +1164,8 @@ class WebSearchInterceptionLogger(CustomLogger):
@staticmethod
def _extract_search_text(result: object) -> str:
if isinstance(result, _SearchToolConfigurationError):
raise result
if isinstance(result, Exception):
verbose_logger.error("WebSearchInterception: Responses search failed with error: %s", result)
return f"Search failed: {result}"
@ -1304,6 +1322,8 @@ class WebSearchInterceptionLogger(CustomLogger):
final_search_results: Final[list[str]] = []
structured_results: Final[list[SearchResponse | None]] = []
for i, result in enumerate(search_results):
if isinstance(result, _SearchToolConfigurationError):
raise result
if isinstance(result, Exception):
verbose_logger.error("WebSearchInterception: Search %s failed with error: %s", i, result)
final_search_results.append(f"Search failed: {result}")
@ -1540,7 +1560,7 @@ class WebSearchInterceptionLogger(CustomLogger):
def _select_search_tool_from_router(self, llm_router: object) -> "_SearchToolConfig | None":
if llm_router is None or not hasattr(llm_router, "search_tools"):
return None
return self._select_search_tool_from_list(search_tools=[], source="router")
search_tools: Final = list(getattr(llm_router, "search_tools") or [])
return self._select_search_tool_from_list(search_tools=search_tools, source="router")
@ -1550,21 +1570,31 @@ class WebSearchInterceptionLogger(CustomLogger):
source: str,
) -> "_SearchToolConfig | None":
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(
"WebSearchInterception: Found search tool '%s' from %s with provider '%s'",
self.search_tool_name,
source,
search_provider,
matching_tools = [
tool
for tool in search_tools
if isinstance(tool, Mapping) and tool.get("search_tool_name") == self.search_tool_name
]
if not matching_tools:
raise _SearchToolConfigurationError(f"Configured search tool '{self.search_tool_name}' was not found")
selected_tool: Final = matching_tools[0]
litellm_params: Final = selected_tool.get("litellm_params")
selected_search_provider: Final = (
litellm_params.get("search_provider") if isinstance(litellm_params, Mapping) else None
)
if not isinstance(selected_search_provider, str) or not selected_search_provider.strip():
raise _SearchToolConfigurationError(
f"Configured search tool '{self.search_tool_name}' does not define a valid search provider"
)
return matching_tools[0]
verbose_logger.debug(
"WebSearchInterception: Search tool '%s' not found in %s, falling back to first available or perplexity",
"WebSearchInterception: Found search tool '%s' from %s with provider '%s'",
self.search_tool_name,
source,
selected_search_provider,
)
return selected_tool
if search_tools:
first_tool: Final = search_tools[0]
@ -1651,6 +1681,8 @@ class WebSearchInterceptionLogger(CustomLogger):
# has no equivalent of Anthropic's web_search_tool_result block.
final_search_results: Final[list[str]] = []
for i, result in enumerate(search_results):
if isinstance(result, _SearchToolConfigurationError):
raise result
if isinstance(result, Exception):
verbose_logger.error("WebSearchInterception: Search %s failed with error: %s", i, result)
final_search_results.append(f"Search failed: {result}")

View file

@ -14,7 +14,7 @@ from litellm.integrations.websearch_interception.handler import (
)
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
from litellm.types.utils import CallTypes, LlmProviders
def test_initialize_from_proxy_config():
@ -230,6 +230,208 @@ async def test_execute_search_passes_selected_search_tool_litellm_params(monkeyp
assert forwarded_kwargs["max_retries"] == 2
@pytest.mark.asyncio
@pytest.mark.parametrize(
"search_tools",
[
pytest.param(None, id="router-not-configured"),
pytest.param([], id="router-has-no-tools"),
pytest.param(
[{"search_tool_name": "other-search", "litellm_params": {"search_provider": "tavily"}}],
id="requested-tool-not-configured",
),
],
)
async def test_execute_search_rejects_missing_explicit_search_tool(monkeypatch, search_tools):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(search_tool_name="parallel-search")
router = None if search_tools is None else MagicMock(search_tools=search_tools)
mock_asearch = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(ValueError, match="Configured search tool 'parallel-search' was not found"):
await logger._execute_search("what is litellm")
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"litellm_params",
[
pytest.param(None, id="missing-parameters"),
pytest.param("not-a-mapping", id="invalid-parameters"),
pytest.param({}, id="missing-provider"),
pytest.param({"search_provider": None}, id="null-provider"),
pytest.param({"search_provider": ""}, id="empty-provider"),
pytest.param({"search_provider": " "}, id="whitespace-provider"),
pytest.param({"search_provider": 123}, id="invalid-provider"),
],
)
async def test_execute_search_rejects_malformed_explicit_search_tool(monkeypatch, litellm_params):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(search_tool_name="parallel-search")
router = MagicMock(
search_tools=[{"search_tool_name": "parallel-search", "litellm_params": litellm_params}],
)
mock_asearch = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(
ValueError, match="Configured search tool 'parallel-search' does not define a valid search provider"
):
await logger._execute_search("what is litellm")
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_execute_search_honors_explicit_parallel_search_tool(monkeypatch):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(search_tool_name="parallel-search")
router = MagicMock(
search_tools=[
{
"search_tool_name": "other-search",
"litellm_params": {"search_provider": "tavily", "api_key": "other-key"},
},
{
"search_tool_name": "parallel-search",
"litellm_params": {
"search_provider": "parallel_ai",
"api_key": "parallel-key",
"api_base": "https://api.parallel.ai",
},
},
],
)
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
await logger._execute_search("what is litellm")
mock_asearch.assert_awaited_once_with(
query="what is litellm",
search_provider="parallel_ai",
api_key="parallel-key",
api_base="https://api.parallel.ai",
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("search_tools", "expected_search_kwargs"),
[
pytest.param(None, {"search_provider": "perplexity"}, id="router-not-configured"),
pytest.param([], {"search_provider": "perplexity"}, id="router-has-no-tools"),
pytest.param(
[
{
"search_tool_name": "first-search",
"litellm_params": {"search_provider": "tavily", "api_key": "first-key"},
},
{
"search_tool_name": "parallel-search",
"litellm_params": {"search_provider": "parallel_ai", "api_key": "parallel-key"},
},
],
{"search_provider": "tavily", "api_key": "first-key"},
id="first-configured-tool",
),
],
)
async def test_execute_search_preserves_implicit_provider_selection(monkeypatch, search_tools, expected_search_kwargs):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger()
router = None if search_tools is None else MagicMock(search_tools=search_tools)
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
await logger._execute_search("what is litellm")
mock_asearch.assert_awaited_once_with(query="what is litellm", **expected_search_kwargs)
@pytest.mark.asyncio
@pytest.mark.parametrize("interception_surface", ["short-circuit", "responses", "anthropic", "chat-completion"])
async def test_missing_explicit_search_tool_propagates_through_interception_surfaces(monkeypatch, interception_surface):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(
enabled_providers=["github_copilot"],
search_tool_name="parallel-search",
)
router = MagicMock(
search_tools=[
{
"search_tool_name": "other-search",
"litellm_params": {"search_provider": "tavily", "api_key": "other-key"},
}
],
)
mock_asearch = AsyncMock()
tool_calls = [{"id": "toolu_123", "call_id": "toolu_123", "input": {"query": "what is litellm"}}]
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
if interception_surface == "short-circuit":
operation = logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "what is litellm"}],
tools=[{"type": "web_search_20250305", "name": "web_search"}],
custom_llm_provider="github_copilot",
)
elif interception_surface == "responses":
operation = logger._build_responses_request_patch(
model="test-model",
messages=[],
tool_calls=tool_calls,
optional_params={},
kwargs={},
)
elif interception_surface == "anthropic":
operation = logger._build_anthropic_request_patch(
model="test-model",
messages=[],
tool_calls=tool_calls,
thinking_blocks=[],
anthropic_messages_optional_request_params={},
logging_obj=None,
kwargs={},
)
else:
operation = logger._build_chat_completion_request_patch(
model="test-model",
messages=[],
tool_calls=tool_calls,
optional_params={},
kwargs={},
)
with pytest.raises(ValueError, match="Configured search tool 'parallel-search' was not found"):
await operation
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_execute_search_attributes_spend_to_the_calling_key(monkeypatch):
"""An intercepted search is billed and logged against the key that made the LLM request.
@ -397,6 +599,166 @@ async def test_execute_search_enforces_team_search_tool_permission(monkeypatch):
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("call_type", "web_search_tool"),
[
pytest.param(
CallTypes.acompletion,
{"type": "web_search_20250305", "name": "web_search"},
id="chat-completion",
),
pytest.param(CallTypes.aresponses, {"type": "web_search"}, id="responses"),
pytest.param(
CallTypes.anthropic_messages,
{"type": "web_search_20250305", "name": "web_search"},
id="anthropic-messages",
),
],
)
async def test_deployment_hook_rejects_missing_explicit_search_tool_before_conversion(
monkeypatch, call_type, web_search_tool
):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="parallel-search")
router = MagicMock(
search_tools=[
{
"search_tool_name": "other-search",
"litellm_params": {"search_provider": "tavily", "api_key": "other-key"},
}
]
)
mock_asearch = AsyncMock()
kwargs = {
"model": "bedrock/claude-sonnet-4",
"messages": [{"role": "user", "content": "what is litellm"}],
"tools": [web_search_tool],
"custom_llm_provider": "bedrock",
}
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(ValueError, match="Configured search tool 'parallel-search' was not found"):
await logger.async_pre_call_deployment_hook(kwargs=kwargs, call_type=call_type)
assert kwargs["tools"] == [web_search_tool]
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_deployment_hook_dispatcher_propagates_missing_explicit_search_tool(monkeypatch):
import litellm
from litellm.proxy import proxy_server
from litellm.utils import async_pre_call_deployment_hook
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="parallel-search")
mock_asearch = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", MagicMock(search_tools=[]))
monkeypatch.setattr(litellm, "callbacks", [logger])
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(ValueError, match="Configured search tool 'parallel-search' was not found"):
await async_pre_call_deployment_hook(
kwargs={
"model": "bedrock/claude-sonnet-4",
"messages": [{"role": "user", "content": "what is litellm"}],
"tools": [{"type": "web_search_20250305", "name": "web_search"}],
"custom_llm_provider": "bedrock",
},
call_type=CallTypes.acompletion.value,
)
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("call_type", "custom_llm_provider", "tools"),
[
pytest.param(
CallTypes.acompletion,
"bedrock",
[{"type": "function", "function": {"name": "calculator", "parameters": {}}}],
id="chat-without-web-search",
),
pytest.param(
CallTypes.aresponses,
"bedrock",
[{"type": "function", "name": "calculator", "parameters": {}}],
id="responses-without-web-search",
),
pytest.param(
CallTypes.acompletion,
"openai",
[{"type": "web_search_20250305", "name": "web_search"}],
id="disabled-chat-provider",
),
pytest.param(CallTypes.aresponses, "openai", [{"type": "web_search"}], id="disabled-responses-provider"),
],
)
async def test_deployment_hook_skips_search_tool_validation_for_unaffected_requests(
monkeypatch, call_type, custom_llm_provider, tools
):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="parallel-search")
mock_asearch = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", MagicMock(search_tools=[]))
monkeypatch.setattr(litellm, "asearch", mock_asearch)
result = await logger.async_pre_call_deployment_hook(
kwargs={
"model": "test-model",
"messages": [{"role": "user", "content": "what is litellm"}],
"tools": tools,
"custom_llm_provider": custom_llm_provider,
},
call_type=call_type,
)
assert result is None
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_deployment_hook_converts_web_search_with_valid_explicit_parallel_tool(monkeypatch):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="parallel-search")
router = MagicMock(
search_tools=[
{"search_tool_name": "other-search", "litellm_params": {"search_provider": "tavily"}},
{"search_tool_name": "parallel-search", "litellm_params": {"search_provider": "parallel_ai"}},
]
)
mock_asearch = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
result = await logger.async_pre_call_deployment_hook(
kwargs={
"model": "bedrock/claude-sonnet-4",
"messages": [{"role": "user", "content": "what is litellm"}],
"tools": [{"type": "web_search_20250305", "name": "web_search"}],
"custom_llm_provider": "bedrock",
},
call_type=CallTypes.acompletion,
)
assert result is not None
assert result["tools"][0]["function"]["name"] == LITELLM_WEB_SEARCH_TOOL_NAME
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.