mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
test(anthropic): fix duplicate WebSearch tool, drop mutation in stream builders, ignore pings
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b93c26a789
commit
337d5cc482
9 changed files with 134 additions and 138 deletions
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from itertools import chain
|
||||
from typing import Final
|
||||
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
|
@ -640,10 +641,12 @@ def sse_events(text: str) -> tuple[tuple[str, dict[str, object]], ...]:
|
|||
frames: Final = tuple(frame for frame in text.split("\n\n") if frame.strip())
|
||||
return tuple(
|
||||
(
|
||||
next(line.removeprefix("event: ") for line in frame.splitlines() if line.startswith("event: ")),
|
||||
event,
|
||||
json.loads(next(line.removeprefix("data: ") for line in frame.splitlines() if line.startswith("data: "))),
|
||||
)
|
||||
for frame in frames
|
||||
if (event := next(line.removeprefix("event: ") for line in frame.splitlines() if line.startswith("event: ")))
|
||||
!= "ping"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -690,6 +693,37 @@ def text_stream(identity: str, model: str, text: str, usage: dict[str, int]) ->
|
|||
)
|
||||
|
||||
|
||||
def _tool_use_frames(index: int, tool_id: str, name: str, tool_input: JsonValue) -> tuple[bytes, ...]:
|
||||
arguments: Final = json.dumps(tool_input)
|
||||
return (
|
||||
sse_frame(
|
||||
"content_block_start",
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": index,
|
||||
"content_block": {"type": "tool_use", "id": tool_id, "name": name, "input": {}},
|
||||
},
|
||||
),
|
||||
sse_frame(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": index,
|
||||
"delta": {"type": "input_json_delta", "partial_json": arguments[: len(arguments) // 2]},
|
||||
},
|
||||
),
|
||||
sse_frame(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": index,
|
||||
"delta": {"type": "input_json_delta", "partial_json": arguments[len(arguments) // 2 :]},
|
||||
},
|
||||
),
|
||||
sse_frame("content_block_stop", {"type": "content_block_stop", "index": index}),
|
||||
)
|
||||
|
||||
|
||||
def tool_use_stream(
|
||||
identity: str,
|
||||
model: str,
|
||||
|
|
@ -698,7 +732,7 @@ def tool_use_stream(
|
|||
tool_calls: tuple[tuple[str, str, JsonValue], ...],
|
||||
usage: dict[str, int],
|
||||
) -> tuple[bytes, ...]:
|
||||
frames: list[bytes] = [
|
||||
head: Final = (
|
||||
sse_frame(
|
||||
"message_start",
|
||||
{
|
||||
|
|
@ -728,37 +762,8 @@ def tool_use_stream(
|
|||
{"type": "content_block_delta", "index": 0, "delta": {"type": "signature_delta", "signature": signature}},
|
||||
),
|
||||
sse_frame("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
]
|
||||
for index, (tool_id, name, tool_input) in enumerate(tool_calls, start=1):
|
||||
arguments: Final = json.dumps(tool_input)
|
||||
frames += [
|
||||
sse_frame(
|
||||
"content_block_start",
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": index,
|
||||
"content_block": {"type": "tool_use", "id": tool_id, "name": name, "input": {}},
|
||||
},
|
||||
),
|
||||
sse_frame(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": index,
|
||||
"delta": {"type": "input_json_delta", "partial_json": arguments[: len(arguments) // 2]},
|
||||
},
|
||||
),
|
||||
sse_frame(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": index,
|
||||
"delta": {"type": "input_json_delta", "partial_json": arguments[len(arguments) // 2 :]},
|
||||
},
|
||||
),
|
||||
sse_frame("content_block_stop", {"type": "content_block_stop", "index": index}),
|
||||
]
|
||||
frames += [
|
||||
)
|
||||
tail: Final = (
|
||||
sse_frame(
|
||||
"message_delta",
|
||||
{
|
||||
|
|
@ -768,5 +773,13 @@ def tool_use_stream(
|
|||
},
|
||||
),
|
||||
sse_frame("message_stop", {"type": "message_stop"}),
|
||||
]
|
||||
return tuple(frames)
|
||||
)
|
||||
frames: Final = (
|
||||
*head,
|
||||
*chain.from_iterable(
|
||||
_tool_use_frames(index, tool_id, name, tool_input)
|
||||
for index, (tool_id, name, tool_input) in enumerate(tool_calls, start=1)
|
||||
),
|
||||
*tail,
|
||||
)
|
||||
return frames
|
||||
|
|
|
|||
|
|
@ -67,22 +67,22 @@ def test_compaction_edit_and_applied_edit_block_round_trip_through_anthropic(gat
|
|||
),
|
||||
(),
|
||||
)
|
||||
seen: list[dict[str, JsonValue]] = []
|
||||
first_expected: Final = {**request_body, "model": cc.FABLE}
|
||||
second_expected: Final = {**turn3, "model": cc.FABLE}
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/v1/messages", request.target
|
||||
body: Final = cc.JSON_OBJECT.validate_json(request.body)
|
||||
seen.append(body)
|
||||
expected: Final = {**(turn3 if len(seen) == 2 else request_body), "model": cc.FABLE}
|
||||
assert body == expected, {
|
||||
key: (expected.get(key), body.get(key))
|
||||
for key in expected.keys() | body.keys()
|
||||
if expected.get(key) != body.get(key)
|
||||
if body == first_expected:
|
||||
assert body["context_management"] == _CONTEXT_MANAGEMENT
|
||||
return Reply(content_type="text/event-stream", chunks=_compaction_stream(identity))
|
||||
assert body == second_expected, {
|
||||
key: (second_expected.get(key), body.get(key))
|
||||
for key in second_expected.keys() | body.keys()
|
||||
if second_expected.get(key) != body.get(key)
|
||||
}
|
||||
assert body["context_management"] == _CONTEXT_MANAGEMENT
|
||||
if len(seen) == 1:
|
||||
return Reply(content_type="text/event-stream", chunks=_compaction_stream(identity))
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=cc.text_stream("msg_cm_next", cc.FABLE, "OK", {"input_tokens": 20, "output_tokens": 2}),
|
||||
|
|
|
|||
|
|
@ -79,16 +79,18 @@ def _cited_stream(identity: str) -> tuple[bytes, ...]:
|
|||
|
||||
|
||||
def test_base64_pdf_document_with_citations_reaches_anthropic_identical(gateway: Gateway) -> None:
|
||||
request_body: Final = cc.frontier_request(f"cache-bust-{uuid.uuid4().hex}", "high", 64000)
|
||||
request_body["messages"] = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
dict(_DOC_BLOCK),
|
||||
{"type": "text", "text": f"What is on page one? {uuid.uuid4().hex}"},
|
||||
],
|
||||
}
|
||||
]
|
||||
request_body: Final = {
|
||||
**cc.frontier_request(f"cache-bust-{uuid.uuid4().hex}", "high", 64000),
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
dict(_DOC_BLOCK),
|
||||
{"type": "text", "text": f"What is on page one? {uuid.uuid4().hex}"},
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
|
|
|
|||
|
|
@ -33,31 +33,33 @@ def test_anthropic_overloaded_primary_falls_back_to_second_deployment(gateway: G
|
|||
wire_server(lambda request: _error_529()) as primary,
|
||||
wire_server(respond_fallback) as fallback,
|
||||
):
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
{
|
||||
"model_name": "cc-primary",
|
||||
"litellm_params": {
|
||||
"model": f"anthropic/{cc.SONNET}",
|
||||
"api_key": cc.ANTHROPIC_API_KEY,
|
||||
"api_base": primary.url,
|
||||
"model_info": {"id": "primary-cc"},
|
||||
config: Final = {
|
||||
**yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()),
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "cc-primary",
|
||||
"litellm_params": {
|
||||
"model": f"anthropic/{cc.SONNET}",
|
||||
"api_key": cc.ANTHROPIC_API_KEY,
|
||||
"api_base": primary.url,
|
||||
"model_info": {"id": "primary-cc"},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "cc-fallback-group",
|
||||
"litellm_params": {
|
||||
"model": f"anthropic/{cc.SONNET}",
|
||||
"api_key": cc.ANTHROPIC_API_KEY,
|
||||
"api_base": fallback.url,
|
||||
"model_info": {"id": "fallback-cc"},
|
||||
{
|
||||
"model_name": "cc-fallback-group",
|
||||
"litellm_params": {
|
||||
"model": f"anthropic/{cc.SONNET}",
|
||||
"api_key": cc.ANTHROPIC_API_KEY,
|
||||
"api_base": fallback.url,
|
||||
"model_info": {"id": "fallback-cc"},
|
||||
},
|
||||
},
|
||||
],
|
||||
"router_settings": {
|
||||
"num_retries": 0,
|
||||
"disable_cooldowns": True,
|
||||
"fallbacks": [{"cc-primary": ["cc-fallback-group"]}],
|
||||
},
|
||||
]
|
||||
config["router_settings"] = {
|
||||
"num_retries": 0,
|
||||
"disable_cooldowns": True,
|
||||
"fallbacks": [{"cc-primary": ["cc-fallback-group"]}],
|
||||
}
|
||||
path: Final = tmp_path / "fallbacks.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
|
|
|
|||
|
|
@ -25,29 +25,27 @@ def test_tool_result_image_block_and_pasted_image_reach_anthropic_identical(gate
|
|||
({"type": "tool_use", "id": "toolu_img", "name": "Read", "input": {"file_path": "/tmp/cc_probe/dot.png"}},),
|
||||
(("toolu_img", [dict(_IMAGE_BLOCK)]),),
|
||||
)
|
||||
pasted: Final = cc.frontier_request(f"cache-bust-{uuid.uuid4().hex}", "high", 64000)
|
||||
pasted["messages"] = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
dict(_IMAGE_BLOCK),
|
||||
{"type": "text", "text": f"What colour is this? {uuid.uuid4().hex}"},
|
||||
],
|
||||
}
|
||||
]
|
||||
seen: list[dict[str, JsonValue]] = []
|
||||
pasted: Final = {
|
||||
**cc.frontier_request(f"cache-bust-{uuid.uuid4().hex}", "high", 64000),
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
dict(_IMAGE_BLOCK),
|
||||
{"type": "text", "text": f"What colour is this? {uuid.uuid4().hex}"},
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
first_expected: Final = {**with_image_result, "model": cc.FABLE}
|
||||
second_expected: Final = {**pasted, "model": cc.FABLE}
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/v1/messages", request.target
|
||||
body: Final = cc.JSON_OBJECT.validate_json(request.body)
|
||||
seen.append(body)
|
||||
expected: Final = {**(pasted if len(seen) == 2 else with_image_result), "model": cc.FABLE}
|
||||
assert body == expected, {
|
||||
key: (expected.get(key), body.get(key))
|
||||
for key in expected.keys() | body.keys()
|
||||
if expected.get(key) != body.get(key)
|
||||
}
|
||||
if body != first_expected:
|
||||
assert body == second_expected, body
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=cc.text_stream(
|
||||
|
|
|
|||
|
|
@ -129,7 +129,8 @@ def test_interleaved_thinking_text_and_tool_use_history_reaches_anthropic_identi
|
|||
)
|
||||
turn2: Final = _turn2(turn1)
|
||||
turn3: Final = _turn3(turn2)
|
||||
seen: list[dict[str, JsonValue]] = []
|
||||
turn2_expected: Final = {**turn2, "model": cc.FABLE}
|
||||
turn3_expected: Final = {**turn3, "model": cc.FABLE}
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
|
|
@ -137,14 +138,12 @@ def test_interleaved_thinking_text_and_tool_use_history_reaches_anthropic_identi
|
|||
upstream_beta: Final = request.headers.get("anthropic-beta", "")
|
||||
assert upstream_beta.split(",").count("interleaved-thinking-2025-05-14") == 1, upstream_beta
|
||||
body: Final = cc.JSON_OBJECT.validate_json(request.body)
|
||||
seen.append(body)
|
||||
expected: Final = {**turn3, "model": cc.FABLE} if len(seen) == 2 else {**turn2, "model": cc.FABLE}
|
||||
assert body == expected, _diff(expected, body)
|
||||
if len(seen) == 1:
|
||||
if body == turn2_expected:
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=cc.text_stream("msg_il_turn2", cc.FABLE, "got A", {"input_tokens": 20, "output_tokens": 4}),
|
||||
)
|
||||
assert body == turn3_expected, _diff(turn3_expected, body)
|
||||
return Reply(content_type="text/event-stream", chunks=_interleaved_stream(identity))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
|
|
|
|||
|
|
@ -30,15 +30,14 @@ def test_claude_code_mid_loop_model_switch_replays_history_byte_identical(gatewa
|
|||
),
|
||||
(("toolu_read_1", "1\tPROBE\n2\t"),),
|
||||
)
|
||||
seen: list[dict[str, JsonValue]] = []
|
||||
first_expected: Final = {**turn1, "model": cc.FABLE}
|
||||
second_expected: Final = {**turn2, "model": cc.OPUS}
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/v1/messages", request.target
|
||||
body: Final = cc.JSON_OBJECT.validate_json(request.body)
|
||||
seen.append(body)
|
||||
if len(seen) == 1:
|
||||
assert body == {**turn1, "model": cc.FABLE}, body.get("model")
|
||||
if body == first_expected:
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=cc.tool_use_stream(
|
||||
|
|
@ -50,11 +49,10 @@ def test_claude_code_mid_loop_model_switch_replays_history_byte_identical(gatewa
|
|||
{"input_tokens": 20, "output_tokens": 10},
|
||||
),
|
||||
)
|
||||
expected: Final = {**turn2, "model": cc.OPUS}
|
||||
assert body == expected, {
|
||||
key: (expected.get(key), body.get(key))
|
||||
for key in expected.keys() | body.keys()
|
||||
if expected.get(key) != body.get(key)
|
||||
assert body == second_expected, {
|
||||
key: (second_expected.get(key), body.get(key))
|
||||
for key in second_expected.keys() | body.keys()
|
||||
if second_expected.get(key) != body.get(key)
|
||||
}
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
|
|
@ -74,7 +72,6 @@ def test_claude_code_mid_loop_model_switch_replays_history_byte_identical(gatewa
|
|||
)
|
||||
assert response2.status_code == 200, response2.text
|
||||
assert len(wire.drain()) == 2
|
||||
assert seen[1]["model"] == cc.OPUS
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT model FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
|
||||
|
|
|
|||
|
|
@ -43,20 +43,19 @@ def test_claude_code_tool_loop_round_trips_thinking_tool_use_and_tool_result(gat
|
|||
),
|
||||
(("toolu_read_1", "1\tPROBE\n2\t"),),
|
||||
)
|
||||
seen: list[dict[str, JsonValue]] = []
|
||||
first_expected: Final = {**turn1, "model": cc.FABLE}
|
||||
second_expected: Final = {**turn2, "model": cc.FABLE}
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/v1/messages", request.target
|
||||
body: Final = cc.JSON_OBJECT.validate_json(request.body)
|
||||
seen.append(body)
|
||||
expected: Final = {**turn2, "model": cc.FABLE} if len(seen) == 2 else {**turn1, "model": cc.FABLE}
|
||||
assert body == expected, _diff(expected, body)
|
||||
if len(seen) == 1:
|
||||
if body == first_expected:
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=cc.tool_use_stream(identity1, cc.FABLE, _THINKING, _SIGNATURE, calls, _USAGE),
|
||||
)
|
||||
assert body == second_expected, _diff(second_expected, body)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=cc.text_stream(identity2, cc.FABLE, "PROBE", {"input_tokens": 30, "output_tokens": 3}),
|
||||
|
|
|
|||
|
|
@ -5,24 +5,8 @@ from integration._support.client import Gateway, eventually
|
|||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from integration.messages_endpoint import _claude_code as cc
|
||||
from pydantic import JsonValue
|
||||
|
||||
WEB_SEARCH_TOOL: Final = {
|
||||
"name": "WebSearch",
|
||||
"description": "Search the web. Returns result blocks with titles and URLs.",
|
||||
"input_schema": cc.schema(
|
||||
{
|
||||
"query": cc.field("The search query to use", type="string", minLength=2),
|
||||
"allowed_domains": cc.field(
|
||||
"Only include search results from these domains", type="array", items={"type": "string"}
|
||||
),
|
||||
"blocked_domains": cc.field(
|
||||
"Never include search results from these domains", type="array", items={"type": "string"}
|
||||
),
|
||||
},
|
||||
("query",),
|
||||
),
|
||||
}
|
||||
WEB_SEARCH_TOOL: Final = {"type": "web_search_20250305", "name": "web_search", "max_uses": 8}
|
||||
|
||||
|
||||
def _web_search_stream(identity: str) -> tuple[bytes, ...]:
|
||||
|
|
@ -121,15 +105,16 @@ def _web_search_stream(identity: str) -> tuple[bytes, ...]:
|
|||
|
||||
def test_claude_code_web_search_tool_passthrough_and_cited_response(gateway: Gateway) -> None:
|
||||
identity: Final = f"msg_ws_{uuid.uuid4().hex}"
|
||||
base: Final = cc.frontier_request(
|
||||
f"cache-bust-{uuid.uuid4().hex}",
|
||||
"high",
|
||||
64000,
|
||||
prompt_text="Use web search to find the current LiteLLM version and answer in one word",
|
||||
)
|
||||
request_body: Final = {
|
||||
**cc.frontier_request(
|
||||
f"cache-bust-{uuid.uuid4().hex}",
|
||||
"high",
|
||||
64000,
|
||||
prompt_text="Use web search to find the current LiteLLM version and answer in one word",
|
||||
),
|
||||
**base,
|
||||
"tools": [*base["tools"], WEB_SEARCH_TOOL],
|
||||
}
|
||||
request_body["tools"] = [*request_body["tools"], WEB_SEARCH_TOOL]
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
|
|
@ -142,6 +127,7 @@ def test_claude_code_web_search_tool_passthrough_and_cited_response(gateway: Gat
|
|||
if expected.get(key) != body.get(key)
|
||||
}
|
||||
assert body["tools"][-1] == WEB_SEARCH_TOOL
|
||||
assert len({tool["name"] for tool in body["tools"]}) == len(body["tools"]), body["tools"]
|
||||
return Reply(content_type="text/event-stream", chunks=_web_search_stream(identity))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue