mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
test(websearch): replace mock with pure unit tests and cover empty queries
Signed-off-by: agustin18 <agustinsaiz02@gmail.com>
This commit is contained in:
parent
c65292e699
commit
e542343994
2 changed files with 101 additions and 109 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue