From 25a63544657b67a6d2c882a35566e5f98effbe62 Mon Sep 17 00:00:00 2001 From: agustin18 Date: Sun, 27 Sep 2026 19:57:45 +0000 Subject: [PATCH] refactor(websearch): enforce Final annotations, prevent rebinding, and retain e2e coverage Signed-off-by: agustin18 --- .../websearch_interception/handler.py | 16 ++++--- .../test_websearch_short_circuit.py | 43 +++++++++++++++++-- 2 files changed, 49 insertions(+), 10 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 266dea59b8a..53c7c32f0f9 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -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, diff --git a/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py index 3f17d1b8aad..a3702bbd29b 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py @@ -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