mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 10fbe70c6a into b781d157d7
This commit is contained in:
commit
d12802a83b
2 changed files with 165 additions and 43 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,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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue