This commit is contained in:
agustin18 2026-09-30 10:30:34 -04:00 • committed by GitHub
commit d12802a83b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 165 additions and 43 deletions

View file

@ -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: <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

View file

@ -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