mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
4f53bbf4d2
commit
36ed5b9a40
2 changed files with 140 additions and 9 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue