From 01f94ac87682ea7bb5cb65acdd094cbfd3e5d001 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 30 Jun 2026 12:07:31 +0000 Subject: [PATCH] fix: pass search tool credentials to litellm.asearch() in websearch interception _execute_search was extracting only search_provider from the matched search tool's litellm_params but silently dropping api_key, api_base, and timeout. This meant user-configured credentials (e.g. EXA_API_KEY via os.environ/) were never forwarded to the actual search call, causing auth failures for providers that require explicit API keys. Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- .../websearch_interception/handler.py | 51 +++--- .../test_websearch_interception_handler.py | 161 +++++++++++++++--- 2 files changed, 165 insertions(+), 47 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 2e11405af3f..fab3a9ee3f5 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -1009,38 +1009,37 @@ class WebSearchInterceptionLogger(CustomLogger): ) llm_router = None - # Determine search provider from router's search_tools + # Determine search provider and credentials from router's search_tools search_provider: Optional[str] = None + search_api_key: Optional[str] = None + search_api_base: Optional[str] = None + search_timeout: Optional[float] = None + resolved_tool: Optional[dict] = 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}'" - ) + resolved_tool = next( + (t for t in llm_router.search_tools if t.get("search_tool_name") == self.search_tool_name), + None, + ) + if resolved_tool: + verbose_logger.debug(f"WebSearchInterception: Found search tool '{self.search_tool_name}'") 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}'" - ) + if resolved_tool is None and llm_router.search_tools: + resolved_tool = llm_router.search_tools[0] + verbose_logger.debug("WebSearchInterception: Using first available search tool") + + if resolved_tool is not None: + litellm_params = resolved_tool.get("litellm_params", {}) + search_provider = litellm_params.get("search_provider") + search_api_key = litellm_params.get("api_key") + search_api_base = litellm_params.get("api_base") + search_timeout = litellm_params.get("timeout") - # Fallback to perplexity if no router or no search tools configured if not search_provider: search_provider = "perplexity" verbose_logger.debug( @@ -1051,7 +1050,13 @@ 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) + result = await litellm.asearch( + query=query, + search_provider=search_provider, + api_key=search_api_key, + api_base=search_api_base, + timeout=search_timeout, + ) # Format using transformation function search_result_text = WebSearchTransformation.format_search_response(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 ea117380edf..e6d08bae24b 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 @@ -4,7 +4,7 @@ Unit tests for WebSearch Interception Handler Tests the WebSearchInterceptionLogger class and helper functions. """ -from unittest.mock import AsyncMock, MagicMock, Mock +from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -56,9 +56,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 +117,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, @@ -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 @@ -216,15 +210,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 +252,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 +313,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 +346,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 +466,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 @@ -509,3 +495,130 @@ async def test_pre_request_hook_syncs_forced_tool_choice(): "type": "tool", "name": LITELLM_WEB_SEARCH_TOOL_NAME, } + + +@pytest.mark.asyncio +async def test_execute_search_passes_api_key_from_search_tool_config(): + """Regression: _execute_search must forward api_key, api_base, and timeout + from the matched search tool's litellm_params to litellm.asearch(). + + Before the fix, only search_provider was extracted; api_key/api_base/timeout + were silently dropped, so user-configured credentials (e.g. EXA_API_KEY via + os.environ/) were never sent to the search provider. + """ + logger = WebSearchInterceptionLogger( + enabled_providers=["bedrock"], + search_tool_name="exa_search", + ) + + fake_router = MagicMock() + fake_router.search_tools = [ + { + "search_tool_name": "exa_search", + "litellm_params": { + "search_provider": "exa_ai", + "api_key": "sk-exa-test-key-123", + "api_base": "https://custom.exa.ai", + "timeout": 30.0, + }, + }, + ] + + mock_search_response = MagicMock() + mock_search_response.results = [] + + with ( + patch( + "litellm.integrations.websearch_interception.handler.litellm.asearch", + new_callable=AsyncMock, + return_value=mock_search_response, + ) as mock_asearch, + patch("litellm.proxy.proxy_server.llm_router", fake_router), + ): + await logger._execute_search("what is litellm") + + mock_asearch.assert_called_once() + call_kwargs = mock_asearch.call_args + assert call_kwargs.kwargs["query"] == "what is litellm" + assert call_kwargs.kwargs["search_provider"] == "exa_ai" + assert call_kwargs.kwargs["api_key"] == "sk-exa-test-key-123" + assert call_kwargs.kwargs["api_base"] == "https://custom.exa.ai" + assert call_kwargs.kwargs["timeout"] == 30.0 + + +@pytest.mark.asyncio +async def test_execute_search_falls_back_to_first_tool_with_credentials(): + """When search_tool_name is not set, _execute_search uses the first + available search tool and still forwards its credentials.""" + logger = WebSearchInterceptionLogger( + enabled_providers=["bedrock"], + search_tool_name=None, + ) + + fake_router = MagicMock() + fake_router.search_tools = [ + { + "search_tool_name": "default-tavily", + "litellm_params": { + "search_provider": "tavily", + "api_key": "tvly-fallback-key", + }, + }, + ] + + mock_search_response = MagicMock() + mock_search_response.results = [] + + with ( + patch( + "litellm.integrations.websearch_interception.handler.litellm.asearch", + new_callable=AsyncMock, + return_value=mock_search_response, + ) as mock_asearch, + patch("litellm.proxy.proxy_server.llm_router", fake_router), + ): + await logger._execute_search("latest AI news") + + mock_asearch.assert_called_once() + call_kwargs = mock_asearch.call_args + assert call_kwargs.kwargs["search_provider"] == "tavily" + assert call_kwargs.kwargs["api_key"] == "tvly-fallback-key" + + +@pytest.mark.asyncio +async def test_execute_search_named_tool_not_found_falls_back(): + """When the named search tool doesn't exist in the router, _execute_search + falls back to the first available tool and uses its credentials.""" + logger = WebSearchInterceptionLogger( + enabled_providers=["bedrock"], + search_tool_name="nonexistent_tool", + ) + + fake_router = MagicMock() + fake_router.search_tools = [ + { + "search_tool_name": "google_pse", + "litellm_params": { + "search_provider": "google_pse", + "api_key": "google-pse-key", + }, + }, + ] + + mock_search_response = MagicMock() + mock_search_response.results = [] + + with ( + patch( + "litellm.integrations.websearch_interception.handler.litellm.asearch", + new_callable=AsyncMock, + return_value=mock_search_response, + ) as mock_asearch, + patch("litellm.proxy.proxy_server.llm_router", fake_router), + ): + await logger._execute_search("test query") + + mock_asearch.assert_called_once() + call_kwargs = mock_asearch.call_args + assert call_kwargs.kwargs["search_provider"] == "google_pse" + assert call_kwargs.kwargs["api_key"] == "google-pse-key"