refactor(websearch): enforce Final annotations, prevent rebinding, and retain e2e coverage

Signed-off-by: agustin18 <agustinsaiz02@gmail.com>
This commit is contained in:
agustin18 2026-09-27 19:57:45 +00:00
parent 5fd3e07583
commit 25a6354465
2 changed files with 49 additions and 10 deletions

View file

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

View file

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