This commit is contained in:
Devon Krisman 2026-08-26 21:06:15 -04:00 • committed by GitHub
commit 2406a00b11
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 208 additions and 4 deletions

View file

@ -244,6 +244,10 @@ class WebSearchInterceptionLogger(CustomLogger):
(t for t in tools if is_anthropic_native_web_search_tool(t)),
None,
)
# The flag covers native tools the pre-request/deployment hooks already converted
emit_native_blocks: Final = native_tool is not None or (
kwargs is not None and bool(kwargs.get(WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY))
)
# Execute search — keep the structured SearchResponse so the native
# block can carry per-result url/title/page_age.
@ -257,9 +261,9 @@ class WebSearchInterceptionLogger(CustomLogger):
search_result_text, structured = f"Search failed: {e}", None
content: Final[list[dict[str, object]]] = []
if native_tool is not None:
if emit_native_blocks:
tool_use_id: Final = f"srvtoolu_{uuid.uuid4().hex}"
tool_name: Final = native_tool.get("name") or "web_search"
tool_name: Final = (native_tool.get("name") if native_tool is not None else None) or "web_search"
content.append(
{
"type": "server_tool_use",
@ -292,7 +296,7 @@ class WebSearchInterceptionLogger(CustomLogger):
verbose_logger.debug(
"WebSearchInterception: Short-circuit search completed, returning synthetic response (%s chars, native_blocks=%s)",
len(search_result_text),
native_tool is not None,
emit_native_blocks,
)
return response

View file

@ -113,8 +113,31 @@ class FakeAnthropicMessagesStreamIterator:
}
chunks.append(f"event: content_block_delta\ndata: {json.dumps(content_block_delta)}\n\n".encode())
elif block_type == "server_tool_use":
server_tool_start: Final = { # mutable-ok: SSE event dict
"type": "content_block_start",
"index": index,
"content_block": { # mutable-ok: SSE event dict
"type": "server_tool_use",
"id": block_dict.get("id"),
"name": block_dict.get("name"),
"input": {}, # mutable-ok: protocol sends empty input in content_block_start
},
}
chunks.append(f"event: content_block_start\ndata: {json.dumps(server_tool_start)}\n\n".encode())
server_tool_delta: Final = { # mutable-ok: SSE event dict
"type": "content_block_delta",
"index": index,
"delta": { # mutable-ok: SSE event dict
"type": "input_json_delta",
"partial_json": json.dumps(block_dict.get("input", {})), # mutable-ok: json.dumps input default
},
}
chunks.append(f"event: content_block_delta\ndata: {json.dumps(server_tool_delta)}\n\n".encode())
else:
passthrough_start: Final = {
# Anthropic does not delta server-generated result blocks either
passthrough_start: Final = { # mutable-ok: SSE event dict
"type": "content_block_start",
"index": index,
"content_block": block_dict,

View file

@ -60,6 +60,53 @@ class TestTryShortCircuitSearch:
assert "Result" in text_block["text"]
mock_search.assert_called_once_with("Search for Claude Code releases")
@pytest.mark.asyncio
async def test_emits_native_blocks_when_tools_already_converted(self):
"""Pre-request/deployment hooks convert the native tool to the LiteLLM
standard shape before the short-circuit runs and record that in
WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY; the flag alone must be enough for
native blocks to be emitted."""
from litellm.integrations.websearch_interception.handler import (
WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY,
)
logger = WebSearchInterceptionLogger(enabled_providers=["hosted_vllm"])
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,
)
result = await logger.try_short_circuit_search(
model="hosted_vllm/some-model",
messages=[{"role": "user", "content": "Search for something"}],
tools=[
{
"name": "litellm_web_search",
"description": "Search the web for information.",
"input_schema": {
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"],
},
}
],
custom_llm_provider="hosted_vllm",
kwargs={WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: True},
)
assert result is not None
block_types = [b["type"] for b in result["content"]]
assert "server_tool_use" in block_types
assert "web_search_tool_result" in block_types
server_tool_use = next(
b for b in result["content"] if b["type"] == "server_tool_use"
)
assert server_tool_use["name"] == "web_search"
@pytest.mark.asyncio
async def test_does_not_short_circuit_mixed_tools(self):
"""Mix of web_search and other tools → NOT short-circuited"""

View file

@ -0,0 +1,130 @@
"""
Tests for FakeAnthropicMessagesStreamIterator
Source: litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py
"""
import json
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
FakeAnthropicMessagesStreamIterator,
)
WEB_SEARCH_RESULT_BLOCK = {
"type": "web_search_tool_result",
"tool_use_id": "srvtoolu_01",
"content": [
{
"type": "web_search_result",
"url": "https://example.com/a",
"title": "Example A",
"encrypted_content": "abc123",
}
],
}
def _response_with_content(content: list) -> dict:
return {
"id": "msg_01",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
"content": content,
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 5},
}
def _events_by_type(iterator: FakeAnthropicMessagesStreamIterator) -> list[dict]:
events = []
for chunk in iterator.chunks:
for line in chunk.decode().splitlines():
if line.startswith("data: "):
events.append(json.loads(line[6:]))
return events
def test_server_tool_use_block_streams_start_and_input_delta():
server_tool_use = {
"type": "server_tool_use",
"id": "srvtoolu_01",
"name": "web_search",
"input": {"query": "example search"},
}
iterator = FakeAnthropicMessagesStreamIterator(
response=_response_with_content([server_tool_use])
)
events = _events_by_type(iterator)
starts = [e for e in events if e.get("type") == "content_block_start"]
assert starts[0]["content_block"]["type"] == "server_tool_use"
assert starts[0]["content_block"]["id"] == "srvtoolu_01"
assert starts[0]["content_block"]["name"] == "web_search"
deltas = [e for e in events if e.get("type") == "content_block_delta"]
assert deltas[0]["delta"]["type"] == "input_json_delta"
assert json.loads(deltas[0]["delta"]["partial_json"]) == {
"query": "example search"
}
stops = [e for e in events if e.get("type") == "content_block_stop"]
assert len(stops) == 1
def test_web_search_tool_result_block_rides_in_content_block_start():
iterator = FakeAnthropicMessagesStreamIterator(
response=_response_with_content([WEB_SEARCH_RESULT_BLOCK])
)
events = _events_by_type(iterator)
starts = [e for e in events if e.get("type") == "content_block_start"]
assert len(starts) == 1
assert starts[0]["content_block"] == WEB_SEARCH_RESULT_BLOCK
stops = [e for e in events if e.get("type") == "content_block_stop"]
assert len(stops) == 1
def test_mixed_content_keeps_block_indices_aligned():
content = [
{"type": "text", "text": "Here is what I found."},
{
"type": "server_tool_use",
"id": "srvtoolu_02",
"name": "web_search",
"input": {"query": "q"},
},
WEB_SEARCH_RESULT_BLOCK,
]
iterator = FakeAnthropicMessagesStreamIterator(
response=_response_with_content(content)
)
events = _events_by_type(iterator)
starts = [e for e in events if e.get("type") == "content_block_start"]
assert [s["index"] for s in starts] == [0, 1, 2]
assert [s["content_block"]["type"] for s in starts] == [
"text",
"server_tool_use",
"web_search_tool_result",
]
stops = [e for e in events if e.get("type") == "content_block_stop"]
assert [s["index"] for s in stops] == [0, 1, 2]
def test_text_only_block_behavior_unchanged():
iterator = FakeAnthropicMessagesStreamIterator(
response=_response_with_content([{"type": "text", "text": "hello"}])
)
events = _events_by_type(iterator)
starts = [e for e in events if e.get("type") == "content_block_start"]
assert starts[0]["content_block"] == {"type": "text", "text": ""}
deltas = [e for e in events if e.get("type") == "content_block_delta"]
assert deltas[0]["delta"] == {"type": "text_delta", "text": "hello"}