diff --git a/tests/integration/messages_endpoint/_claude_code.py b/tests/integration/messages_endpoint/_claude_code.py index 323312186c1..53c439fbc82 100644 --- a/tests/integration/messages_endpoint/_claude_code.py +++ b/tests/integration/messages_endpoint/_claude_code.py @@ -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 diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_compaction_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_compaction_wire.py index a1e903ab5c6..0132cf7a1e6 100644 --- a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_compaction_wire.py +++ b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_compaction_wire.py @@ -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}), diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_document_input_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_document_input_wire.py index 7fa67549372..d4a827d7ece 100644 --- a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_document_input_wire.py +++ b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_document_input_wire.py @@ -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" diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_fallback_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_fallback_wire.py index 0075acb6da8..830f84c858c 100644 --- a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_fallback_wire.py +++ b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_fallback_wire.py @@ -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)) diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_image_input_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_image_input_wire.py index 0deed7db56a..6d2f4d06dec 100644 --- a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_image_input_wire.py +++ b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_image_input_wire.py @@ -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( diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_interleaved_thinking_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_interleaved_thinking_wire.py index c6a46ca2806..f5da9b070f5 100644 --- a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_interleaved_thinking_wire.py +++ b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_interleaved_thinking_wire.py @@ -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: diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_model_switch_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_model_switch_wire.py index ccdd4b5bd06..c04b974da01 100644 --- a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_model_switch_wire.py +++ b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_model_switch_wire.py @@ -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', diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_tool_loop_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_tool_loop_wire.py index f386a206451..0fd9828186d 100644 --- a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_tool_loop_wire.py +++ b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_tool_loop_wire.py @@ -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}), diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_web_search_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_web_search_wire.py index a275b2e3305..aed0a5abd40 100644 --- a/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_web_search_wire.py +++ b/tests/integration/messages_endpoint/providers/anthropic/test_claude_code_web_search_wire.py @@ -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: