mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 573b02a35f into e58a561caa
This commit is contained in:
commit
96a115f0f3
6 changed files with 269 additions and 148 deletions
|
|
@ -522,7 +522,7 @@ ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES = {
|
|||
|
||||
# LiteLLM standard web search tool name
|
||||
# Used for web search interception across providers
|
||||
LITELLM_WEB_SEARCH_TOOL_NAME = "litellm_web_search"
|
||||
LITELLM_WEB_SEARCH_TOOL_NAME = "WebSearch"
|
||||
|
||||
DEFAULT_IMAGE_ENDPOINT_MODEL = "dall-e-2"
|
||||
DEFAULT_VIDEO_ENDPOINT_MODEL = "sora-2"
|
||||
|
|
|
|||
|
|
@ -86,10 +86,14 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
|
||||
Claude Code sends web search as a separate, standalone /v1/messages
|
||||
request with a simple prompt and only web_search tool(s). For providers
|
||||
that don't natively support web search (e.g. github_copilot), there is
|
||||
no need to route this through the backend LLM — we can detect the
|
||||
pattern, execute the search via Tavily/Perplexity, and return a
|
||||
synthetic Anthropic response immediately.
|
||||
that don't natively support web search, we execute the search via the
|
||||
configured provider (SearXNG/Tavily/Perplexity) and return a synthetic
|
||||
response in native Anthropic format (server_tool_use +
|
||||
web_search_tool_result) so Claude Code's WebSearchTool parser works.
|
||||
|
||||
Providers with native Anthropic Messages support (anthropic, bedrock,
|
||||
vertex_ai, azure_ai) are skipped — their API handles web search
|
||||
natively and returns the correct format already.
|
||||
|
||||
Args:
|
||||
model: Model name from the request
|
||||
|
|
@ -112,12 +116,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
):
|
||||
return None
|
||||
|
||||
# Only short-circuit for providers without native Anthropic Messages
|
||||
# support. Providers that have a BaseAnthropicMessagesConfig (bedrock,
|
||||
# vertex_ai, azure_ai, anthropic) already use the agentic loop, which
|
||||
# includes a follow-up LLM call to synthesize the answer from search
|
||||
# results. Short-circuiting those would skip that synthesis step and
|
||||
# return raw search text — a regression for existing users.
|
||||
# Skip providers with native Anthropic Messages support — their API
|
||||
# handles web_search_20250305 natively, returning server_tool_use +
|
||||
# web_search_tool_result in the correct format already.
|
||||
try:
|
||||
provider_enum = LlmProviders(provider_str)
|
||||
anthropic_config = (
|
||||
|
|
@ -128,7 +129,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
if anthropic_config is not None:
|
||||
verbose_logger.debug(
|
||||
f"WebSearchInterception: Skipping short-circuit for {provider_str} "
|
||||
"(provider has native Anthropic Messages support, using agentic loop)"
|
||||
"(provider has native web search support)"
|
||||
)
|
||||
return None
|
||||
except (ValueError, Exception):
|
||||
|
|
@ -161,21 +162,69 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
)
|
||||
search_result_text = f"Search failed: {e}"
|
||||
|
||||
# Build synthetic Anthropic response
|
||||
# Parse search results into structured hits for web_search_tool_result
|
||||
search_hits = []
|
||||
for block in search_result_text.split("\n\n"):
|
||||
title, url, snippet = "", "", ""
|
||||
for line in block.strip().splitlines():
|
||||
if line.startswith("Title: "):
|
||||
title = line[7:]
|
||||
elif line.startswith("URL: "):
|
||||
url = line[5:]
|
||||
elif line.startswith("Snippet: "):
|
||||
snippet = line[9:]
|
||||
if url:
|
||||
hit: Dict[str, Any] = {
|
||||
"type": "web_search_result",
|
||||
"url": url,
|
||||
"title": title or url,
|
||||
"encrypted_content": snippet or "",
|
||||
"page_age": None,
|
||||
}
|
||||
search_hits.append(hit)
|
||||
|
||||
tool_use_id = f"srvtoolu_{str(uuid.uuid4()).replace('-', '')[:24]}"
|
||||
|
||||
# Build response in native Anthropic format so Claude Code's
|
||||
# WebSearchTool parser sees server_tool_use + web_search_tool_result.
|
||||
content: List[Dict[str, Any]] = [
|
||||
{
|
||||
"type": "server_tool_use",
|
||||
"id": tool_use_id,
|
||||
"name": "web_search",
|
||||
"input": {"query": query},
|
||||
},
|
||||
{
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": tool_use_id,
|
||||
"content": search_hits,
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": search_result_text,
|
||||
},
|
||||
]
|
||||
|
||||
response: Dict[str, Any] = {
|
||||
"id": f"msg_{str(uuid.uuid4())}",
|
||||
"id": f"msg_{str(uuid.uuid4()).replace('-', '')[:20]}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": model,
|
||||
"content": [{"type": "text", "text": search_result_text}],
|
||||
"content": content,
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 0, "output_tokens": 0},
|
||||
"usage": {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
"server_tool_use": {
|
||||
"web_search_requests": 1,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
verbose_logger.debug(
|
||||
"WebSearchInterception: Short-circuit search completed, "
|
||||
f"returning synthetic response ({len(search_result_text)} chars)"
|
||||
f"returning native format ({len(search_hits)} hits)"
|
||||
)
|
||||
return response
|
||||
|
||||
|
|
@ -880,8 +929,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
)
|
||||
llm_router = None
|
||||
|
||||
# Determine search provider from router's search_tools
|
||||
# Determine search provider and api_base from router's search_tools
|
||||
search_provider: Optional[str] = None
|
||||
api_base: Optional[str] = None
|
||||
if llm_router is not None and hasattr(llm_router, "search_tools"):
|
||||
if self.search_tool_name:
|
||||
# Find specific search tool by name
|
||||
|
|
@ -892,9 +942,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
]
|
||||
if matching_tools:
|
||||
search_tool = matching_tools[0]
|
||||
search_provider = search_tool.get("litellm_params", {}).get(
|
||||
"search_provider"
|
||||
)
|
||||
litellm_params = search_tool.get("litellm_params", {})
|
||||
search_provider = litellm_params.get("search_provider")
|
||||
api_base = litellm_params.get("api_base")
|
||||
verbose_logger.debug(
|
||||
f"WebSearchInterception: Found search tool '{self.search_tool_name}' "
|
||||
f"with provider '{search_provider}'"
|
||||
|
|
@ -908,9 +958,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
# If no specific tool or not found, use first available
|
||||
if not search_provider and llm_router.search_tools:
|
||||
first_tool = llm_router.search_tools[0]
|
||||
search_provider = first_tool.get("litellm_params", {}).get(
|
||||
"search_provider"
|
||||
)
|
||||
litellm_params = first_tool.get("litellm_params", {})
|
||||
search_provider = litellm_params.get("search_provider")
|
||||
api_base = api_base or litellm_params.get("api_base")
|
||||
verbose_logger.debug(
|
||||
f"WebSearchInterception: Using first available search tool with provider '{search_provider}'"
|
||||
)
|
||||
|
|
@ -926,7 +976,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
verbose_logger.debug(
|
||||
f"WebSearchInterception: Executing search for '{query}' using provider '{search_provider}'"
|
||||
)
|
||||
result = await litellm.asearch(query=query, search_provider=search_provider)
|
||||
result = await litellm.asearch(query=query, search_provider=search_provider, api_base=api_base)
|
||||
|
||||
# Format using transformation function
|
||||
search_result_text = WebSearchTransformation.format_search_response(result)
|
||||
|
|
|
|||
|
|
@ -13,36 +13,28 @@ from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME
|
|||
|
||||
def get_litellm_web_search_tool() -> Dict[str, Any]:
|
||||
"""
|
||||
Get the standard LiteLLM web search tool definition.
|
||||
Get the web search tool definition in Anthropic format.
|
||||
|
||||
This is the canonical tool definition that all native web search tools
|
||||
(like Anthropic's web_search_20250305, Claude Code's web_search, etc.)
|
||||
are converted to for interception.
|
||||
Uses the same name and schema that Claude Code expects so it appears
|
||||
as the native WebSearch tool in the client.
|
||||
|
||||
Returns:
|
||||
Dict containing the Anthropic-style tool definition with:
|
||||
- name: Tool name
|
||||
- description: What the tool does
|
||||
- input_schema: JSON schema for tool parameters
|
||||
|
||||
Example:
|
||||
>>> tool = get_litellm_web_search_tool()
|
||||
>>> tool['name']
|
||||
'litellm_web_search'
|
||||
Dict containing the Anthropic-style tool definition.
|
||||
"""
|
||||
return {
|
||||
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"description": (
|
||||
"Search the web for information. Use this when you need current "
|
||||
"information or answers to questions that require up-to-date data."
|
||||
"Search the web for current information. Returns search results "
|
||||
"with titles, URLs, and snippets. Use this tool when you need "
|
||||
"up-to-date information beyond your knowledge cutoff."
|
||||
),
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "The search query to execute",
|
||||
}
|
||||
"description": "The search query to use",
|
||||
},
|
||||
},
|
||||
"required": ["query"],
|
||||
},
|
||||
|
|
@ -51,7 +43,7 @@ def get_litellm_web_search_tool() -> Dict[str, Any]:
|
|||
|
||||
def get_litellm_web_search_tool_openai() -> Dict[str, Any]:
|
||||
"""
|
||||
Get the standard LiteLLM web search tool definition in OpenAI format.
|
||||
Get the web search tool definition in OpenAI format.
|
||||
|
||||
Used by async_pre_call_deployment_hook which runs in the chat completions
|
||||
path where tools must be in OpenAI format (type: "function" with
|
||||
|
|
@ -65,16 +57,17 @@ def get_litellm_web_search_tool_openai() -> Dict[str, Any]:
|
|||
"function": {
|
||||
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"description": (
|
||||
"Search the web for information. Use this when you need current "
|
||||
"information or answers to questions that require up-to-date data."
|
||||
"Search the web for current information. Returns search results "
|
||||
"with titles, URLs, and snippets. Use this tool when you need "
|
||||
"up-to-date information beyond your knowledge cutoff."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "The search query to execute",
|
||||
}
|
||||
"description": "The search query to use",
|
||||
},
|
||||
},
|
||||
"required": ["query"],
|
||||
},
|
||||
|
|
@ -82,45 +75,35 @@ def get_litellm_web_search_tool_openai() -> Dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
_WEB_SEARCH_NAMES = {LITELLM_WEB_SEARCH_TOOL_NAME, "WebSearch", "web_search", "litellm_web_search"}
|
||||
|
||||
|
||||
def is_web_search_tool_chat_completion(tool: Dict[str, Any]) -> bool:
|
||||
"""
|
||||
Check if a tool is a web search tool for Chat Completions API (strict check).
|
||||
Check if a tool is a web search tool for Chat Completions API.
|
||||
|
||||
This is a stricter version that ONLY checks for the exact LiteLLM web search tool name.
|
||||
Use this for Chat Completions API to avoid false positives with user-defined tools.
|
||||
|
||||
Detects ONLY:
|
||||
- LiteLLM standard: name == "litellm_web_search" (Anthropic format)
|
||||
- OpenAI format: type == "function" with function.name == "litellm_web_search"
|
||||
Detects:
|
||||
- Anthropic format: name in {litellm_web_search, WebSearch, web_search}
|
||||
- OpenAI format: type == "function" with function.name in same set
|
||||
|
||||
Args:
|
||||
tool: Tool dictionary to check
|
||||
|
||||
Returns:
|
||||
True if tool is exactly the LiteLLM web search tool
|
||||
|
||||
Example:
|
||||
>>> is_web_search_tool_chat_completion({"name": "litellm_web_search"})
|
||||
True
|
||||
>>> is_web_search_tool_chat_completion({"type": "function", "function": {"name": "litellm_web_search"}})
|
||||
True
|
||||
>>> is_web_search_tool_chat_completion({"name": "web_search"})
|
||||
False
|
||||
>>> is_web_search_tool_chat_completion({"name": "WebSearch"})
|
||||
False
|
||||
True if tool is a web search tool
|
||||
"""
|
||||
tool_name = tool.get("name", "")
|
||||
tool_type = tool.get("type", "")
|
||||
|
||||
# Check for OpenAI format: {"type": "function", "function": {"name": "litellm_web_search"}}
|
||||
# Check for OpenAI format
|
||||
if tool_type == "function" and "function" in tool:
|
||||
function_def = tool.get("function", {})
|
||||
function_name = function_def.get("name", "")
|
||||
if function_name == LITELLM_WEB_SEARCH_TOOL_NAME:
|
||||
if function_name in _WEB_SEARCH_NAMES:
|
||||
return True
|
||||
|
||||
# Check for LiteLLM standard tool (Anthropic format)
|
||||
if tool_name == LITELLM_WEB_SEARCH_TOOL_NAME:
|
||||
# Check for Anthropic format
|
||||
if tool_name in _WEB_SEARCH_NAMES:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
|
@ -175,8 +158,8 @@ def is_web_search_tool(tool: Dict[str, Any]) -> bool:
|
|||
if tool_name == "web_search" and tool_type:
|
||||
return True
|
||||
|
||||
# Check for legacy WebSearch format
|
||||
if tool_name == "WebSearch":
|
||||
# Check for legacy names
|
||||
if tool_name in ("WebSearch", "litellm_web_search"):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -101,11 +101,11 @@ class WebSearchTransformation:
|
|||
block_input = getattr(block, "input", {})
|
||||
|
||||
# Check for LiteLLM standard or legacy web search tools
|
||||
# Handles: litellm_web_search, WebSearch, web_search
|
||||
if block_type == "tool_use" and block_name in (
|
||||
LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"WebSearch",
|
||||
"web_search",
|
||||
"litellm_web_search",
|
||||
):
|
||||
# Convert to dict for easier handling
|
||||
tool_call = {
|
||||
|
|
@ -195,6 +195,7 @@ class WebSearchTransformation:
|
|||
LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"WebSearch",
|
||||
"web_search",
|
||||
"litellm_web_search",
|
||||
):
|
||||
# Parse arguments (might be JSON string)
|
||||
if isinstance(function_arguments, str):
|
||||
|
|
|
|||
|
|
@ -131,6 +131,18 @@ class FakeAnthropicMessagesStreamIterator:
|
|||
f"event: content_block_delta\ndata: {json.dumps(content_block_delta)}\n\n".encode()
|
||||
)
|
||||
|
||||
elif block_type in ("server_tool_use", "web_search_tool_result"):
|
||||
# Emit the full block as content_block_start — same as
|
||||
# Anthropic's native streaming format for server-side tools.
|
||||
content_block_start = {
|
||||
"type": "content_block_start",
|
||||
"index": index,
|
||||
"content_block": block_dict,
|
||||
}
|
||||
chunks.append(
|
||||
f"event: content_block_start\ndata: {json.dumps(content_block_start)}\n\n".encode()
|
||||
)
|
||||
|
||||
content_block_stop = {"type": "content_block_stop", "index": index}
|
||||
chunks.append(
|
||||
f"event: content_block_stop\ndata: {json.dumps(content_block_stop)}\n\n".encode()
|
||||
|
|
|
|||
|
|
@ -3,6 +3,9 @@ Unit tests for WebSearch Short-Circuit
|
|||
|
||||
Tests the short-circuit path that detects web-search-only /v1/messages requests
|
||||
and executes the search directly without routing through the backend LLM.
|
||||
|
||||
The response uses native Anthropic format (server_tool_use + web_search_tool_result)
|
||||
so Claude Code's WebSearchTool parser works correctly.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
|
@ -48,11 +51,103 @@ class TestTryShortCircuitSearch:
|
|||
assert result["type"] == "message"
|
||||
assert result["role"] == "assistant"
|
||||
assert result["stop_reason"] == "end_turn"
|
||||
assert len(result["content"]) == 1
|
||||
assert result["content"][0]["type"] == "text"
|
||||
assert "Result" in result["content"][0]["text"]
|
||||
# Native format: server_tool_use + web_search_tool_result + text
|
||||
assert result["content"][0]["type"] == "server_tool_use"
|
||||
assert result["content"][0]["name"] == "web_search"
|
||||
assert result["content"][1]["type"] == "web_search_tool_result"
|
||||
assert result["content"][2]["type"] == "text"
|
||||
assert "Result" in result["content"][2]["text"]
|
||||
mock_search.assert_called_once_with("Search for Claude Code releases")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_format_search_hits(self):
|
||||
"""Search results are structured as web_search_result hits"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
|
||||
|
||||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = (
|
||||
"Title: First Result\nURL: https://example.com/1\nSnippet: first\n\n"
|
||||
"Title: Second Result\nURL: https://example.com/2\nSnippet: second"
|
||||
)
|
||||
|
||||
result = await logger.try_short_circuit_search(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
messages=[{"role": "user", "content": "Search query"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search"}],
|
||||
custom_llm_provider="github_copilot",
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
hits = result["content"][1]["content"]
|
||||
assert len(hits) == 2
|
||||
assert hits[0]["type"] == "web_search_result"
|
||||
assert hits[0]["url"] == "https://example.com/1"
|
||||
assert hits[0]["title"] == "First Result"
|
||||
assert hits[1]["url"] == "https://example.com/2"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_server_tool_use_has_query(self):
|
||||
"""server_tool_use block contains the original search query"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
|
||||
|
||||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "Title: R\nURL: https://x.com\nSnippet: s"
|
||||
|
||||
result = await logger.try_short_circuit_search(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
messages=[{"role": "user", "content": "trending AI topics"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search"}],
|
||||
custom_llm_provider="github_copilot",
|
||||
)
|
||||
|
||||
stu = result["content"][0]
|
||||
assert stu["type"] == "server_tool_use"
|
||||
assert stu["name"] == "web_search"
|
||||
assert stu["input"]["query"] == "trending AI topics"
|
||||
assert stu["id"].startswith("srvtoolu_")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_use_id_links_blocks(self):
|
||||
"""server_tool_use.id matches web_search_tool_result.tool_use_id"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
|
||||
|
||||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "Title: R\nURL: https://x.com\nSnippet: s"
|
||||
|
||||
result = await logger.try_short_circuit_search(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
messages=[{"role": "user", "content": "query"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search"}],
|
||||
custom_llm_provider="github_copilot",
|
||||
)
|
||||
|
||||
assert result["content"][0]["id"] == result["content"][1]["tool_use_id"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_includes_web_search_requests(self):
|
||||
"""Usage includes server_tool_use.web_search_requests count"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
|
||||
|
||||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "Title: R\nURL: https://x.com\nSnippet: s"
|
||||
|
||||
result = await logger.try_short_circuit_search(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
messages=[{"role": "user", "content": "query"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search"}],
|
||||
custom_llm_provider="github_copilot",
|
||||
)
|
||||
|
||||
assert result["usage"]["server_tool_use"]["web_search_requests"] == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_does_not_short_circuit_mixed_tools(self):
|
||||
"""Mix of web_search and other tools → NOT short-circuited"""
|
||||
|
|
@ -115,27 +210,42 @@ class TestTryShortCircuitSearch:
|
|||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_does_not_short_circuit_bedrock(self):
|
||||
"""Bedrock has native agentic loop support → NOT short-circuited.
|
||||
async def test_does_not_short_circuit_native_providers(self):
|
||||
"""Providers with native Anthropic Messages support (anthropic, bedrock,
|
||||
vertex_ai) are skipped — their API handles web search natively."""
|
||||
for provider in ["anthropic", "bedrock"]:
|
||||
logger = WebSearchInterceptionLogger(
|
||||
enabled_providers=[provider, "github_copilot"]
|
||||
)
|
||||
|
||||
Providers with a BaseAnthropicMessagesConfig (bedrock, vertex_ai, etc.)
|
||||
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"]
|
||||
)
|
||||
result = await logger.try_short_circuit_search(
|
||||
model=f"{provider}/claude-sonnet-4",
|
||||
messages=[{"role": "user", "content": "search query"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search"}],
|
||||
custom_llm_provider=provider,
|
||||
)
|
||||
|
||||
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}
|
||||
],
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
assert result is None, f"Short-circuit should NOT fire for native provider {provider}"
|
||||
|
||||
assert result is None
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_circuits_non_native_providers(self):
|
||||
"""Non-native providers (github_copilot, etc.) get short-circuited."""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
|
||||
|
||||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "Title: R\nURL: https://x.com\nSnippet: s"
|
||||
|
||||
result = await logger.try_short_circuit_search(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
messages=[{"role": "user", "content": "search query"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search"}],
|
||||
custom_llm_provider="github_copilot",
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["content"][0]["type"] == "server_tool_use"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_does_not_short_circuit_no_messages(self):
|
||||
|
|
@ -173,7 +283,10 @@ class TestTryShortCircuitSearch:
|
|||
)
|
||||
|
||||
assert result is not None
|
||||
assert "Search failed" in result["content"][0]["text"]
|
||||
# Error text is in the last content block (text)
|
||||
text_block = result["content"][-1]
|
||||
assert text_block["type"] == "text"
|
||||
assert "Search failed" in text_block["text"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_has_valid_structure(self):
|
||||
|
|
@ -183,7 +296,7 @@ class TestTryShortCircuitSearch:
|
|||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "search results here"
|
||||
mock_search.return_value = "Title: R\nURL: https://x.com\nSnippet: test"
|
||||
|
||||
result = await logger.try_short_circuit_search(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
|
|
@ -203,11 +316,8 @@ class TestTryShortCircuitSearch:
|
|||
assert result["stop_sequence"] is None
|
||||
assert "usage" in result
|
||||
assert "content" in result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Query extraction tests
|
||||
# ---------------------------------------------------------------------------
|
||||
# Content has 3 blocks: server_tool_use, web_search_tool_result, text
|
||||
assert len(result["content"]) == 3
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -237,7 +347,7 @@ class TestShortCircuitEntryPoint:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_dict_when_not_streaming(self):
|
||||
"""Non-streaming short-circuit → returns dict"""
|
||||
"""Non-streaming short-circuit → returns dict with native format"""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
|
||||
_try_websearch_short_circuit,
|
||||
)
|
||||
|
|
@ -246,7 +356,7 @@ class TestShortCircuitEntryPoint:
|
|||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "results"
|
||||
mock_search.return_value = "Title: R\nURL: https://x.com\nSnippet: results"
|
||||
with patch("litellm.callbacks", [logger]):
|
||||
result = await _try_websearch_short_circuit(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
|
|
@ -257,7 +367,8 @@ class TestShortCircuitEntryPoint:
|
|||
)
|
||||
|
||||
assert isinstance(result, dict)
|
||||
assert result["content"][0]["text"] == "results"
|
||||
assert result["content"][0]["type"] == "server_tool_use"
|
||||
assert result["content"][1]["type"] == "web_search_tool_result"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_stream_iterator_when_streaming(self):
|
||||
|
|
@ -273,7 +384,9 @@ class TestShortCircuitEntryPoint:
|
|||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "streaming results"
|
||||
mock_search.return_value = (
|
||||
"Title: Result\nURL: https://example.com\nSnippet: streaming results"
|
||||
)
|
||||
with patch("litellm.callbacks", [logger]):
|
||||
result = await _try_websearch_short_circuit(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
|
|
@ -291,12 +404,12 @@ class TestShortCircuitEntryPoint:
|
|||
chunks.append(chunk)
|
||||
|
||||
assert len(chunks) > 0
|
||||
# First chunk should be message_start
|
||||
assert b"event: message_start" in chunks[0]
|
||||
# Last chunk should be message_stop
|
||||
assert b"event: message_stop" in chunks[-1]
|
||||
# Should contain the search results text
|
||||
# Should contain server_tool_use and web_search_tool_result blocks
|
||||
all_data = b"".join(chunks)
|
||||
assert b"server_tool_use" in all_data
|
||||
assert b"web_search_tool_result" in all_data
|
||||
assert b"streaming results" in all_data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -321,12 +434,7 @@ class TestShortCircuitEntryPoint:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_uses_original_stream_not_hook_converted(self):
|
||||
"""Verify that the entry point passes original_stream to the short-circuit.
|
||||
|
||||
The pre-request hook converts stream=True → stream=False for the agentic
|
||||
loop. The short-circuit must use the ORIGINAL stream value so streaming
|
||||
callers get SSE events instead of a plain dict.
|
||||
"""
|
||||
"""Verify that the entry point passes original_stream to the short-circuit."""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
|
||||
FakeAnthropicMessagesStreamIterator,
|
||||
)
|
||||
|
|
@ -338,47 +446,14 @@ class TestShortCircuitEntryPoint:
|
|||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "streaming results"
|
||||
mock_search.return_value = "Title: R\nURL: https://x.com\nSnippet: s"
|
||||
with patch("litellm.callbacks", [logger]):
|
||||
# Simulate what anthropic_messages() does: original_stream=True
|
||||
# is passed to the short-circuit, even though the hook would have
|
||||
# already converted stream to False in request_kwargs.
|
||||
result = await _try_websearch_short_circuit(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
messages=[{"role": "user", "content": "search query"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search"}],
|
||||
custom_llm_provider="github_copilot",
|
||||
stream=True, # original_stream, NOT the hook-converted value
|
||||
stream=True,
|
||||
)
|
||||
|
||||
# Must return a stream iterator, not a plain dict
|
||||
assert isinstance(result, FakeAnthropicMessagesStreamIterator)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_circuits_with_provider_from_model_string(self):
|
||||
"""Provider embedded in model string (custom_llm_provider=None) should
|
||||
still fire the short-circuit when the caller propagates the derived
|
||||
provider.
|
||||
"""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
|
||||
_try_websearch_short_circuit,
|
||||
)
|
||||
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
|
||||
with patch.object(
|
||||
logger, "_execute_search", new_callable=AsyncMock
|
||||
) as mock_search:
|
||||
mock_search.return_value = "results"
|
||||
with patch("litellm.callbacks", [logger]):
|
||||
# Simulate the caller having derived custom_llm_provider from
|
||||
# the model string before calling _try_websearch_short_circuit
|
||||
result = await _try_websearch_short_circuit(
|
||||
model="github_copilot/claude-sonnet-4",
|
||||
messages=[{"role": "user", "content": "search query"}],
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search"}],
|
||||
custom_llm_provider="github_copilot",
|
||||
stream=False,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["content"][0]["text"] == "results"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue