From e542343994ecffadf2b50cdf598d7367f7bfcd82 Mon Sep 17 00:00:00 2001 From: agustin18 Date: Sun, 27 Sep 2026 19:33:24 +0000 Subject: [PATCH] test(websearch): replace mock with pure unit tests and cover empty queries Signed-off-by: agustin18 --- .../websearch_interception/handler.py | 3 +- .../test_websearch_short_circuit.py | 207 +++++++++--------- 2 files changed, 101 insertions(+), 109 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index cf46f5827f0..c534115181a 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -309,7 +309,8 @@ class WebSearchInterceptionLogger(CustomLogger): (cleaned.startswith('"') and cleaned.endswith('"')) or (cleaned.startswith("'") and cleaned.endswith("'")) ): cleaned = cleaned[1:-1].strip() - return cleaned if cleaned else raw_message.strip() + return cleaned if (cleaned or prefix) 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 81d3aa0b26f..3f17d1b8aad 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py @@ -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,30 +140,49 @@ 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", + ) + + assert result is None + + @pytest.mark.asyncio async def test_search_failure_returns_error_text(self): """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 +195,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 +223,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 +316,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 +342,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 +405,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 +433,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,69 +449,3 @@ 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", - ), - ( - " 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", - ), - ( - "Search for Claude Code releases", - "Search for Claude Code releases", - ), - ( - '"LiteLLM latest version release"', - '"LiteLLM latest version release"', - ), - ( - "'fastapi SSE streaming'", - "'fastapi SSE streaming'", - ), - ], - ) - async def test_short_circuits_strips_claude_code_instructional_prefix( - self, raw_prompt, expected_query - ): - """Claude Code wraps standalone searches in 'Perform a web search for the query: ...'. - The short-circuit path must extract only the actual query terms so search backends - receive the intended query rather than instructional prefix words. - """ - 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 -