diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 6ebf485d717..53c7c32f0f9 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -8,6 +8,7 @@ server-side using litellm router's search tools. import asyncio import math +import re import uuid from collections.abc import AsyncIterator, Mapping, Sequence from dataclasses import dataclass @@ -92,6 +93,11 @@ WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks" _RESPONSE_CONTENT_FIELD: Final = "content" +_CLAUDE_CODE_SEARCH_PREFIX_RE: Final = re.compile( + r"^\s*(?:Perform\s+a\s+)?(?:web\s+)?search\s+for\s+(?:the\s+)?query:\s*", + re.IGNORECASE, +) + _ResponseT: Final = TypeVar("_ResponseT") @@ -288,6 +294,27 @@ class WebSearchInterceptionLogger(CustomLogger): """ return validated_max_agentic_loops(max_agentic_loops, field="websearch_interception_params.max_agentic_loops") + @classmethod + def _extract_short_circuit_query(cls, raw_message: str) -> str: + """Extract the intended search terms from a short-circuit user message. + + Clients such as Claude Code format standalone WebSearch requests with an + instructional wrapper (e.g. 'Perform a web search for the query: '). + Strip known instructional prefixes so the search backend receives only + the intended search terms rather than query wrapper keywords. + """ + prefix_match: Final[re.Match[str] | None] = _CLAUDE_CODE_SEARCH_PREFIX_RE.match(raw_message) + stripped_prefix: Final[str] = ( + raw_message[prefix_match.end() :].strip() if prefix_match is not None else raw_message.strip() + ) + if prefix_match is not None and ( + (stripped_prefix.startswith('"') and stripped_prefix.endswith('"')) + or (stripped_prefix.startswith("'") and stripped_prefix.endswith("'")) + ): + return stripped_prefix[1:-1].strip() + + return stripped_prefix if (stripped_prefix or prefix_match is not None) else raw_message.strip() + async def try_short_circuit_search( self, model: str, @@ -358,7 +385,11 @@ class WebSearchInterceptionLogger(CustomLogger): get_last_user_message, ) - query: Final = get_last_user_message(cast(list[AllMessageValues], messages)) + raw_query: Final = get_last_user_message(cast(list[AllMessageValues], messages)) + if not raw_query: + return None + + query: Final = self._extract_short_circuit_query(raw_query) if not query: return None diff --git a/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py index 8294add60c7..4ae815186c3 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py @@ -26,9 +26,7 @@ class TestTryShortCircuitSearch: """Single web_search_20250305 tool → short-circuit fires""" logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) - with patch.object( - logger, "_execute_search", new_callable=AsyncMock - ) as mock_search: + with patch.object(logger, "_execute_search", new_callable=AsyncMock) as mock_search: mock_search.return_value = ( "Title: Result\nURL: https://example.com\nSnippet: test", None, @@ -36,12 +34,8 @@ class TestTryShortCircuitSearch: result = await logger.try_short_circuit_search( model="github_copilot/claude-sonnet-4", - messages=[ - {"role": "user", "content": "Search for Claude Code releases"} - ], - tools=[ - {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} - ], + messages=[{"role": "user", "content": "Search for Claude Code releases"}], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], custom_llm_provider="github_copilot", ) @@ -113,9 +107,7 @@ class TestTryShortCircuitSearch: result = await logger.try_short_circuit_search( model="github_copilot/claude-sonnet-4", messages=[{"role": "user", "content": "Search for something"}], - tools=[ - {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} - ], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], custom_llm_provider="github_copilot", ) @@ -129,16 +121,12 @@ class TestTryShortCircuitSearch: use the agentic loop which includes a follow-up LLM synthesis step. The short-circuit must not fire for them. """ - logger = WebSearchInterceptionLogger( - enabled_providers=["bedrock", "github_copilot"] - ) + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock", "github_copilot"]) result = await logger.try_short_circuit_search( model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", messages=[{"role": "user", "content": "Search for something"}], - tools=[ - {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} - ], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], custom_llm_provider="bedrock", ) @@ -152,9 +140,29 @@ class TestTryShortCircuitSearch: result = await logger.try_short_circuit_search( model="github_copilot/claude-sonnet-4", messages=[], - tools=[ - {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} - ], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + custom_llm_provider="github_copilot", + ) + + assert result is None + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "empty_prompt", + [ + " ", + "Perform a web search for the query: ", + 'Perform a web search for the query: ""', + ], + ) + async def test_does_not_short_circuit_empty_or_whitespace_user_message(self, empty_prompt: str): + """User message that is whitespace or empty query → returns None without executing search.""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": empty_prompt}], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], custom_llm_provider="github_copilot", ) @@ -165,17 +173,13 @@ class TestTryShortCircuitSearch: """Search failure → response with error message, not exception""" logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) - with patch.object( - logger, "_execute_search", new_callable=AsyncMock - ) as mock_search: + with patch.object(logger, "_execute_search", new_callable=AsyncMock) as mock_search: mock_search.side_effect = RuntimeError("Tavily API error") result = await logger.try_short_circuit_search( model="github_copilot/claude-sonnet-4", messages=[{"role": "user", "content": "Search for something"}], - tools=[ - {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} - ], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], custom_llm_provider="github_copilot", ) @@ -188,9 +192,7 @@ class TestTryShortCircuitSearch: """Synthetic response has all required AnthropicMessagesResponse fields""" logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) - with patch.object( - logger, "_execute_search", new_callable=AsyncMock - ) as mock_search: + with patch.object(logger, "_execute_search", new_callable=AsyncMock) as mock_search: mock_search.return_value = ("search results here", None) result = await logger.try_short_circuit_search( @@ -218,6 +220,66 @@ class TestTryShortCircuitSearch: # --------------------------------------------------------------------------- +class TestQueryExtraction: + """Direct unit tests for query extraction without application test doubles.""" + + @pytest.mark.parametrize( + "raw_prompt,expected_query", + [ + ( + "Perform a web search for the query: LiteLLM latest version release", + "LiteLLM latest version release", + ), + ( + 'Perform a web search for the query: "who is the maintainer of litellm"', + "who is the maintainer of litellm", + ), + ( + " perform a web search for the query: fastapi SSE streaming ", + "fastapi SSE streaming", + ), + ( + "Search for the query: python 3.13 changelog", + "python 3.13 changelog", + ), + ( + "web search for query: litellm", + "litellm", + ), + ( + "Search for Claude Code releases", + "Search for Claude Code releases", + ), + ( + '"LiteLLM latest version release"', + '"LiteLLM latest version release"', + ), + ( + "'fastapi SSE streaming'", + "'fastapi SSE streaming'", + ), + ( + "Perform a web search for the query: ", + "", + ), + ( + 'Perform a web search for the query: ""', + "", + ), + ( + "", + "", + ), + ], + ) + def test_extract_short_circuit_query_strips_claude_code_prefix(self, raw_prompt: str, expected_query: str): + """Claude Code wraps standalone searches in 'Perform a web search for the query: ...'. + _extract_short_circuit_query must extract only the actual query terms and preserve + exact-phrase quotes when no wrapper was present, without application doubles. + """ + assert WebSearchInterceptionLogger._extract_short_circuit_query(raw_prompt) == expected_query + + # --------------------------------------------------------------------------- # Integration with entry point # --------------------------------------------------------------------------- @@ -251,9 +313,7 @@ class TestShortCircuitEntryPoint: ) logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) - with patch.object( - logger, "_execute_search", new_callable=AsyncMock - ) as mock_search: + with patch.object(logger, "_execute_search", new_callable=AsyncMock) as mock_search: mock_search.return_value = ("results", None) with patch("litellm.callbacks", [logger]): result = await _try_websearch_short_circuit( @@ -279,9 +339,7 @@ class TestShortCircuitEntryPoint: ) logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) - with patch.object( - logger, "_execute_search", new_callable=AsyncMock - ) as mock_search: + with patch.object(logger, "_execute_search", new_callable=AsyncMock) as mock_search: mock_search.return_value = ("streaming results", None) with patch("litellm.callbacks", [logger]): result = await _try_websearch_short_circuit( @@ -344,9 +402,7 @@ class TestShortCircuitEntryPoint: ) logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) - with patch.object( - logger, "_execute_search", new_callable=AsyncMock - ) as mock_search: + with patch.object(logger, "_execute_search", new_callable=AsyncMock) as mock_search: mock_search.return_value = ("streaming results", None) with patch("litellm.callbacks", [logger]): # Simulate what anthropic_messages() does: original_stream=True @@ -374,9 +430,7 @@ class TestShortCircuitEntryPoint: ) logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) - with patch.object( - logger, "_execute_search", new_callable=AsyncMock - ) as mock_search: + with patch.object(logger, "_execute_search", new_callable=AsyncMock) as mock_search: mock_search.return_value = ("results", None) with patch("litellm.callbacks", [logger]): # Simulate the caller having derived custom_llm_provider from @@ -392,3 +446,40 @@ class TestShortCircuitEntryPoint: assert result is not None text_block = next(b for b in result["content"] if b["type"] == "text") assert text_block["text"] == "results" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "raw_prompt,expected_query", + [ + ( + "Perform a web search for the query: LiteLLM latest version release", + "LiteLLM latest version release", + ), + ( + 'Perform a web search for the query: "who is the maintainer of litellm"', + "who is the maintainer of litellm", + ), + ( + '"LiteLLM latest version release"', + '"LiteLLM latest version release"', + ), + ], + ) + async def test_short_circuits_strips_claude_code_instructional_prefix(self, raw_prompt: str, expected_query: str): + """Verify end-to-end that try_short_circuit_search forwards the cleaned query to search.""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + with patch.object(logger, "_execute_search", new_callable=AsyncMock) as mock_search: + mock_search.return_value = ("Results", None) + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": raw_prompt}], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + custom_llm_provider="github_copilot", + ) + + assert result is not None + mock_search.assert_called_once_with(expected_query) + tool_use_block = next(b for b in result["content"] if b["type"] == "server_tool_use") + assert tool_use_block["input"]["query"] == expected_query