test(websearch): replace mock with pure unit tests and cover empty queries

Signed-off-by: agustin18 <agustinsaiz02@gmail.com>
This commit is contained in:
agustin18 2026-09-27 19:33:24 +00:00
parent c65292e699
commit e542343994
2 changed files with 101 additions and 109 deletions

View file

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

View file

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