mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
refactor(websearch): simplify explicit search tool validation
This commit is contained in:
parent
ce7fbf91b2
commit
56ec8f7af4
2 changed files with 49 additions and 245 deletions
|
|
@ -95,10 +95,6 @@ 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."""
|
||||
|
||||
|
|
@ -256,8 +252,6 @@ 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
|
||||
|
|
@ -1164,8 +1158,6 @@ 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}"
|
||||
|
|
@ -1322,8 +1314,6 @@ 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}")
|
||||
|
|
@ -1559,9 +1549,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
return None
|
||||
|
||||
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 self._select_search_tool_from_list(search_tools=[], source="router")
|
||||
search_tools: Final = list(getattr(llm_router, "search_tools") or [])
|
||||
search_tools: Final = 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(
|
||||
|
|
@ -1570,13 +1558,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
source: str,
|
||||
) -> "_SearchToolConfig | None":
|
||||
if self.search_tool_name:
|
||||
matching_tools = [
|
||||
tool
|
||||
for tool in search_tools
|
||||
if isinstance(tool, Mapping) and tool.get("search_tool_name") == self.search_tool_name
|
||||
]
|
||||
matching_tools = [tool for tool in search_tools if 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")
|
||||
raise ValueError(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")
|
||||
|
|
@ -1584,7 +1568,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
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(
|
||||
raise ValueError(
|
||||
f"Configured search tool '{self.search_tool_name}' does not define a valid search provider"
|
||||
)
|
||||
|
||||
|
|
@ -1681,8 +1665,6 @@ 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}")
|
||||
|
|
|
|||
|
|
@ -232,17 +232,37 @@ async def test_execute_search_passes_selected_search_tool_litellm_params(monkeyp
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"search_tools",
|
||||
("search_tools", "error"),
|
||||
[
|
||||
pytest.param(None, id="router-not-configured"),
|
||||
pytest.param([], id="router-has-no-tools"),
|
||||
pytest.param(None, "was not found", id="router-not-configured"),
|
||||
pytest.param(
|
||||
[{"search_tool_name": "other-search", "litellm_params": {"search_provider": "tavily"}}],
|
||||
"was not found",
|
||||
id="requested-tool-not-configured",
|
||||
),
|
||||
pytest.param(
|
||||
[{"search_tool_name": "parallel-search", "litellm_params": "not-a-mapping"}],
|
||||
"does not define a valid search provider",
|
||||
id="invalid-parameters",
|
||||
),
|
||||
pytest.param(
|
||||
[{"search_tool_name": "parallel-search", "litellm_params": {}}],
|
||||
"does not define a valid search provider",
|
||||
id="missing-provider",
|
||||
),
|
||||
pytest.param(
|
||||
[{"search_tool_name": "parallel-search", "litellm_params": {"search_provider": " "}}],
|
||||
"does not define a valid search provider",
|
||||
id="whitespace-provider",
|
||||
),
|
||||
pytest.param(
|
||||
[{"search_tool_name": "parallel-search", "litellm_params": {"search_provider": 123}}],
|
||||
"does not define a valid search provider",
|
||||
id="invalid-provider",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_execute_search_rejects_missing_explicit_search_tool(monkeypatch, search_tools):
|
||||
async def test_execute_search_rejects_invalid_explicit_search_tool(monkeypatch, search_tools, error):
|
||||
import litellm
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
|
|
@ -253,41 +273,7 @@ async def test_execute_search_rejects_missing_explicit_search_tool(monkeypatch,
|
|||
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"
|
||||
):
|
||||
with pytest.raises(ValueError, match=f"Configured search tool 'parallel-search' {error}"):
|
||||
await logger._execute_search("what is litellm")
|
||||
|
||||
mock_asearch.assert_not_awaited()
|
||||
|
|
@ -307,11 +293,7 @@ async def test_execute_search_honors_explicit_parallel_search_tool(monkeypatch):
|
|||
},
|
||||
{
|
||||
"search_tool_name": "parallel-search",
|
||||
"litellm_params": {
|
||||
"search_provider": "parallel_ai",
|
||||
"api_key": "parallel-key",
|
||||
"api_base": "https://api.parallel.ai",
|
||||
},
|
||||
"litellm_params": {"search_provider": "parallel_ai", "api_key": "parallel-key"},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
|
@ -326,7 +308,6 @@ async def test_execute_search_honors_explicit_parallel_search_tool(monkeypatch):
|
|||
query="what is litellm",
|
||||
search_provider="parallel_ai",
|
||||
api_key="parallel-key",
|
||||
api_base="https://api.parallel.ai",
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -335,7 +316,6 @@ async def test_execute_search_honors_explicit_parallel_search_tool(monkeypatch):
|
|||
("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(
|
||||
[
|
||||
{
|
||||
|
|
@ -368,70 +348,6 @@ async def test_execute_search_preserves_implicit_provider_selection(monkeypatch,
|
|||
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.
|
||||
|
|
@ -608,7 +524,8 @@ async def test_execute_search_enforces_team_search_tool_permission(monkeypatch):
|
|||
{"type": "web_search_20250305", "name": "web_search"},
|
||||
id="chat-completion",
|
||||
),
|
||||
pytest.param(CallTypes.aresponses, {"type": "web_search"}, id="responses"),
|
||||
pytest.param(CallTypes.responses, {"type": "web_search"}, id="responses"),
|
||||
pytest.param(CallTypes.aresponses, {"type": "web_search"}, id="async-responses"),
|
||||
pytest.param(
|
||||
CallTypes.anthropic_messages,
|
||||
{"type": "web_search_20250305", "name": "web_search"},
|
||||
|
|
@ -616,147 +533,52 @@ async def test_execute_search_enforces_team_search_tool_permission(monkeypatch):
|
|||
),
|
||||
],
|
||||
)
|
||||
async def test_deployment_hook_rejects_missing_explicit_search_tool_before_conversion(
|
||||
async def test_deployment_hook_dispatcher_propagates_missing_explicit_search_tool(
|
||||
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()
|
||||
kwargs = {
|
||||
"model": "bedrock/claude-sonnet-4",
|
||||
"tools": [web_search_tool],
|
||||
"custom_llm_provider": "bedrock",
|
||||
}
|
||||
|
||||
monkeypatch.setattr(proxy_server, "llm_router", MagicMock(search_tools=[]))
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
MagicMock(search_tools=[{"search_tool_name": "other-search", "litellm_params": {"search_provider": "tavily"}}]),
|
||||
)
|
||||
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,
|
||||
)
|
||||
await async_pre_call_deployment_hook(kwargs=kwargs, call_type=call_type.value)
|
||||
|
||||
assert kwargs["tools"] == [web_search_tool]
|
||||
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
|
||||
async def test_deployment_hook_skips_explicit_tool_validation_for_non_search_responses(monkeypatch):
|
||||
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,
|
||||
"tools": [{"type": "function", "name": "calculator"}],
|
||||
"custom_llm_provider": "bedrock",
|
||||
},
|
||||
call_type=call_type,
|
||||
call_type=CallTypes.aresponses,
|
||||
)
|
||||
|
||||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue