fix(websearch_interception): load API keys from router configuration

Load search provider, API key, and API base from the router's
search_tools config instead of relying on environment variables.
This commit is contained in:
Quentin Machu 2026-05-03 23:14:21 -04:00 • committed by Quentin Machu
parent 4f53bbf4d2
commit 36ed5b9a40
No known key found for this signature in database
2 changed files with 140 additions and 9 deletions

View file

@ -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)

View file

@ -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",
)