mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(websearch): strip Claude Code instructional prefix in short-circuit search
Signed-off-by: agustin18 <agustinsaiz02@gmail.com>
This commit is contained in:
parent
22b36cbcf6
commit
b94efc3399
2 changed files with 83 additions and 1 deletions
|
|
@ -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: <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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue