diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 37528e7dcd5..f754e7fe411 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -1067,8 +1067,10 @@ 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 + api_key: Optional[str] = None + api_base: 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 @@ -1079,9 +1081,10 @@ class WebSearchInterceptionLogger(CustomLogger): ] if matching_tools: search_tool = matching_tools[0] - search_provider = search_tool.get("litellm_params", {}).get( - "search_provider" - ) + litellm_params = search_tool.get("litellm_params", {}) + search_provider = litellm_params.get("search_provider") + api_key = litellm_params.get("api_key") + api_base = litellm_params.get("api_base") verbose_logger.debug( f"WebSearchInterception: Found search tool '{self.search_tool_name}' " f"with provider '{search_provider}'" @@ -1095,9 +1098,10 @@ class WebSearchInterceptionLogger(CustomLogger): # 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" - ) + litellm_params = first_tool.get("litellm_params", {}) + search_provider = litellm_params.get("search_provider") + api_key = litellm_params.get("api_key") + api_base = litellm_params.get("api_base") verbose_logger.debug( f"WebSearchInterception: Using first available search tool with provider '{search_provider}'" ) @@ -1113,7 +1117,15 @@ 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: Dict[str, Any] = { + "query": query, + "search_provider": search_provider, + } + if api_key: + search_kwargs["api_key"] = api_key + if api_base: + search_kwargs["api_base"] = api_base + result = await litellm.asearch(**search_kwargs) # 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 10951265115..1c963c3fcf7 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 @@ -380,3 +380,122 @@ async def test_deployment_hook_converts_stream_and_logging_obj_syncs(): logging_obj.stream = _hook_stream assert logging_obj.stream is False + + +def _mock_proxy_server(mock_router): + """Create a mock proxy_server module with llm_router.""" + mock_module = Mock() + mock_module.llm_router = mock_router + return mock_module + + +@pytest.mark.asyncio +async def test_execute_search_loads_api_key_from_named_tool(): + """Test that _execute_search loads API key and base from router's named search tool.""" + logger = WebSearchInterceptionLogger( + enabled_providers=["bedrock"], search_tool_name="my-search" + ) + + mock_router = Mock() + mock_router.search_tools = [ + { + "search_tool_name": "my-search", + "litellm_params": { + "search_provider": "tavily", + "api_key": "tvly-secret", + "api_base": "https://custom.tavily.com", + }, + } + ] + + mock_search_result = Mock() + mock_search_result.results = [] + + with ( + patch.dict( + "sys.modules", + {"litellm.proxy.proxy_server": _mock_proxy_server(mock_router)}, + ), + patch( + "litellm.asearch", + new_callable=AsyncMock, + return_value=mock_search_result, + ) as mock_asearch, + ): + await logger._execute_search("test query") + + mock_asearch.assert_called_once_with( + query="test query", + search_provider="tavily", + api_key="tvly-secret", + api_base="https://custom.tavily.com", + ) + + +@pytest.mark.asyncio +async def test_execute_search_falls_back_to_first_tool(): + """Test that _execute_search uses first available tool when no named tool matches.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + + mock_router = Mock() + mock_router.search_tools = [ + { + "search_tool_name": "default", + "litellm_params": { + "search_provider": "google", + "api_key": "google-key", + }, + } + ] + + mock_search_result = Mock() + mock_search_result.results = [] + + with ( + patch.dict( + "sys.modules", + {"litellm.proxy.proxy_server": _mock_proxy_server(mock_router)}, + ), + patch( + "litellm.asearch", + new_callable=AsyncMock, + return_value=mock_search_result, + ) as mock_asearch, + ): + await logger._execute_search("test query") + + mock_asearch.assert_called_once_with( + query="test query", + search_provider="google", + api_key="google-key", + ) + + +@pytest.mark.asyncio +async def test_execute_search_defaults_to_perplexity(): + """Test that _execute_search falls back to perplexity when no router search tools.""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + + mock_router = Mock() + mock_router.search_tools = [] + + mock_search_result = Mock() + mock_search_result.results = [] + + with ( + patch.dict( + "sys.modules", + {"litellm.proxy.proxy_server": _mock_proxy_server(mock_router)}, + ), + patch( + "litellm.asearch", + new_callable=AsyncMock, + return_value=mock_search_result, + ) as mock_asearch, + ): + await logger._execute_search("test query") + + mock_asearch.assert_called_once_with( + query="test query", + search_provider="perplexity", + )