From c99acf1bd946b56161a0a30fc3f24ff5b1d9a066 Mon Sep 17 00:00:00 2001 From: Deepak kudi Date: Thu, 4 Jun 2026 12:54:28 +0530 Subject: [PATCH] fix(websearch): allow providers to skip short-circuit --- .../websearch_interception/handler.py | 54 ++++++++++++++----- .../integrations/websearch_interception.py | 4 ++ .../test_websearch_interception_handler.py | 4 +- .../test_websearch_short_circuit.py | 27 ++++++++++ 4 files changed, 73 insertions(+), 16 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 37528e7dcd5..e9173d86fa0 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -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( diff --git a/litellm/types/integrations/websearch_interception.py b/litellm/types/integrations/websearch_interception.py index d8a36169b88..ab1680e8a85 100644 --- a/litellm/types/integrations/websearch_interception.py +++ b/litellm/types/integrations/websearch_interception.py @@ -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.""" 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..553e347bb1d 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 @@ -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, diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py index 7de8892b8fc..e2185fe96fb 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py @@ -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.