mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(websearch): allow providers to skip short-circuit
This commit is contained in:
parent
e9417603a3
commit
c99acf1bd9
4 changed files with 73 additions and 16 deletions
|
|
@ -49,6 +49,27 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY = "_websearch_interception_emit_native_blocks"
|
|||
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY = "websearch_native_blocks"
|
||||
|
||||
|
||||
def _normalize_provider_names(
|
||||
providers: List[Union[LlmProviders, str]],
|
||||
) -> List[str]:
|
||||
return [p.value if isinstance(p, LlmProviders) else p for p in providers]
|
||||
|
||||
|
||||
def _provider_names_from_config(
|
||||
providers: Optional[List[str]],
|
||||
) -> Optional[List[Union[LlmProviders, str]]]:
|
||||
if providers is None:
|
||||
return None
|
||||
|
||||
normalized: List[Union[LlmProviders, str]] = []
|
||||
for provider in providers:
|
||||
try:
|
||||
normalized.append(LlmProviders(provider))
|
||||
except ValueError:
|
||||
normalized.append(provider)
|
||||
return normalized
|
||||
|
||||
|
||||
class WebSearchInterceptionLogger(CustomLogger):
|
||||
"""
|
||||
CustomLogger that intercepts WebSearch tool calls for models that don't
|
||||
|
|
@ -65,6 +86,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
self,
|
||||
enabled_providers: Optional[List[Union[LlmProviders, str]]] = None,
|
||||
search_tool_name: Optional[str] = None,
|
||||
disable_short_circuit_providers: Optional[
|
||||
List[Union[LlmProviders, str]]
|
||||
] = None,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
|
|
@ -74,15 +98,18 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
Default: None (all providers enabled)
|
||||
search_tool_name: Name of search tool configured in router's search_tools.
|
||||
If None, will attempt to use first available search tool.
|
||||
disable_short_circuit_providers: Providers where web-search-only
|
||||
requests should use normal dispatch.
|
||||
"""
|
||||
super().__init__()
|
||||
# Convert enum values to strings for comparison
|
||||
if enabled_providers is None:
|
||||
self.enabled_providers = [LlmProviders.BEDROCK.value]
|
||||
else:
|
||||
self.enabled_providers = [
|
||||
p.value if isinstance(p, LlmProviders) else p for p in enabled_providers
|
||||
]
|
||||
self.enabled_providers = _normalize_provider_names(enabled_providers)
|
||||
self.disable_short_circuit_providers = _normalize_provider_names(
|
||||
disable_short_circuit_providers or []
|
||||
)
|
||||
self.search_tool_name = search_tool_name
|
||||
self._request_has_websearch = False # Track if current request has web search
|
||||
|
||||
|
|
@ -123,6 +150,8 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
and provider_str not in self.enabled_providers
|
||||
):
|
||||
return None
|
||||
if provider_str in self.disable_short_circuit_providers:
|
||||
return None
|
||||
|
||||
# Only short-circuit for providers without native Anthropic Messages
|
||||
# support. Providers that have a BaseAnthropicMessagesConfig (bedrock,
|
||||
|
|
@ -326,23 +355,20 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
# Extract parameters from config
|
||||
enabled_providers_str = config.get("enabled_providers", None)
|
||||
search_tool_name = config.get("search_tool_name", None)
|
||||
disable_short_circuit_providers_str = config.get(
|
||||
"disable_short_circuit_providers", None
|
||||
)
|
||||
|
||||
# Convert string provider names to LlmProviders enum values
|
||||
enabled_providers: Optional[List[Union[LlmProviders, str]]] = None
|
||||
if enabled_providers_str is not None:
|
||||
enabled_providers = []
|
||||
for provider in enabled_providers_str:
|
||||
try:
|
||||
# Try to convert string to LlmProviders enum
|
||||
provider_enum = LlmProviders(provider)
|
||||
enabled_providers.append(provider_enum)
|
||||
except ValueError:
|
||||
# If conversion fails, keep as string
|
||||
enabled_providers.append(provider)
|
||||
enabled_providers = _provider_names_from_config(enabled_providers_str)
|
||||
disable_short_circuit_providers = _provider_names_from_config(
|
||||
disable_short_circuit_providers_str
|
||||
)
|
||||
|
||||
return cls(
|
||||
enabled_providers=enabled_providers,
|
||||
search_tool_name=search_tool_name,
|
||||
disable_short_circuit_providers=disable_short_circuit_providers,
|
||||
)
|
||||
|
||||
async def async_pre_request_hook(
|
||||
|
|
|
|||
|
|
@ -13,11 +13,15 @@ class WebSearchInterceptionConfig(TypedDict, total=False):
|
|||
litellm_settings:
|
||||
websearch_interception_params:
|
||||
enabled_providers: ["bedrock"]
|
||||
disable_short_circuit_providers: ["hosted_vllm"]
|
||||
search_tool_name: "my-perplexity-search"
|
||||
"""
|
||||
|
||||
enabled_providers: List[str]
|
||||
"""List of LLM provider names to enable interception for (e.g., ['bedrock', 'vertex_ai'])"""
|
||||
|
||||
disable_short_circuit_providers: List[str]
|
||||
"""List of LLM provider names where web-search-only requests should use normal dispatch."""
|
||||
|
||||
search_tool_name: Optional[str]
|
||||
"""Name of search tool configured in router's search_tools. If None, uses first available."""
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ def test_initialize_from_proxy_config():
|
|||
litellm_settings = {
|
||||
"websearch_interception_params": {
|
||||
"enabled_providers": ["bedrock", "vertex_ai"],
|
||||
"disable_short_circuit_providers": ["hosted_vllm"],
|
||||
"search_tool_name": "my-search",
|
||||
}
|
||||
}
|
||||
|
|
@ -31,6 +32,7 @@ def test_initialize_from_proxy_config():
|
|||
|
||||
assert LlmProviders.BEDROCK.value in logger.enabled_providers
|
||||
assert LlmProviders.VERTEX_AI.value in logger.enabled_providers
|
||||
assert "hosted_vllm" in logger.disable_short_circuit_providers
|
||||
assert logger.search_tool_name == "my-search"
|
||||
|
||||
|
||||
|
|
@ -132,8 +134,6 @@ async def test_internal_flags_filtered_from_followup_kwargs():
|
|||
to the follow-up LLM request, causing "Extra inputs are not permitted" errors
|
||||
from providers like Bedrock that use strict parameter validation.
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
|
||||
# Simulate kwargs that would be passed during agentic loop execution
|
||||
kwargs_with_internal_flags = {
|
||||
"_websearch_interception_converted_stream": True,
|
||||
|
|
|
|||
|
|
@ -121,6 +121,33 @@ class TestTryShortCircuitSearch:
|
|||
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_does_not_short_circuit_disabled_provider(self):
|
||||
"""Provider in disable_short_circuit_providers -> NOT short-circuited"""
|
||||
logger = WebSearchInterceptionLogger(
|
||||
enabled_providers=["hosted_vllm"],
|
||||
disable_short_circuit_providers=["hosted_vllm"],
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
result = await logger.try_short_circuit_search(
|
||||
model="hosted_vllm/Qwen/Qwen2.5-Coder-32B-Instruct",
|
||||
messages=[{"role": "user", "content": "Search for something"}],
|
||||
tools=[
|
||||
{
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
"max_uses": 8,
|
||||
}
|
||||
],
|
||||
custom_llm_provider="hosted_vllm",
|
||||
)
|
||||
|
||||
assert result is None
|
||||
mock_search.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_does_not_short_circuit_bedrock(self):
|
||||
"""Bedrock has native agentic loop support → NOT short-circuited.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue