mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(websearch): enforce Final annotations, prevent rebinding, and retain e2e coverage
Signed-off-by: agustin18 <agustinsaiz02@gmail.com>
This commit is contained in:
parent
5fd3e07583
commit
25a6354465
2 changed files with 49 additions and 10 deletions
|
|
@ -303,13 +303,17 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
Strip known instructional prefixes so the search backend receives only
|
||||
the intended search terms rather than query wrapper keywords.
|
||||
"""
|
||||
prefix = _CLAUDE_CODE_SEARCH_PREFIX_RE.match(raw_message)
|
||||
cleaned = (raw_message[prefix.end() :] if prefix else raw_message).strip()
|
||||
if prefix and (
|
||||
(cleaned.startswith('"') and cleaned.endswith('"')) or (cleaned.startswith("'") and cleaned.endswith("'"))
|
||||
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("'"))
|
||||
):
|
||||
cleaned = cleaned[1:-1].strip()
|
||||
return cleaned if (cleaned or prefix) else raw_message.strip()
|
||||
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,
|
||||
|
|
|
|||
|
|
@ -155,9 +155,7 @@ class TestTryShortCircuitSearch:
|
|||
'Perform a web search for the query: ""',
|
||||
],
|
||||
)
|
||||
async def test_does_not_short_circuit_empty_or_whitespace_user_message(
|
||||
self, empty_prompt: str
|
||||
):
|
||||
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"])
|
||||
|
||||
|
|
@ -170,7 +168,6 @@ class TestTryShortCircuitSearch:
|
|||
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_failure_returns_error_text(self):
|
||||
"""Search failure → response with error message, not exception"""
|
||||
|
|
@ -449,3 +446,41 @@ 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"), None)
|
||||
if tool_use_block is not None:
|
||||
assert tool_use_block["input"]["query"] == expected_query
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue