diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 6ebf485d717..6a4929a6b22 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,20 @@ 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. + """ + cleaned = _CLAUDE_CODE_SEARCH_PREFIX_RE.sub("", raw_message).strip() + if (cleaned.startswith('"') and cleaned.endswith('"')) or (cleaned.startswith("'") and cleaned.endswith("'")): + cleaned = cleaned[1:-1].strip() + return cleaned if cleaned else raw_message.strip() + async def try_short_circuit_search( self, model: str, @@ -358,7 +378,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..2616a81ce96 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py @@ -392,3 +392,61 @@ 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", + ), + ( + " 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", + ), + ( + "Search for Claude Code releases", + "Search for Claude Code releases", + ), + ], + ) + async def test_short_circuits_strips_claude_code_instructional_prefix( + self, raw_prompt, expected_query + ): + """Claude Code wraps standalone searches in 'Perform a web search for the query: ...'. + The short-circuit path must extract only the actual query terms so search backends + receive the intended query rather than instructional prefix words. + """ + 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"), None + ) + if tool_use_block is not None: + assert tool_use_block["input"]["query"] == expected_query +