mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge 784e9eb3c9 into 77765fd302
This commit is contained in:
commit
2406a00b11
4 changed files with 208 additions and 4 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
Loading…
Add table
Reference in a new issue