From 56ec8f7af42c29eb0ac08f0bff3bb1b6565183ac Mon Sep 17 00:00:00 2001 From: George Pickett Date: Mon, 24 Aug 2026 12:09:11 -0700 Subject: [PATCH] refactor(websearch): simplify explicit search tool validation --- .../websearch_interception/handler.py | 26 +- .../test_websearch_interception_handler.py | 268 +++--------------- 2 files changed, 49 insertions(+), 245 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 26766cfb13f..d025590def9 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -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}") diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py index aa4dd929ff7..ec4bc1f49eb 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py @@ -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