diff --git a/tests/integration/_support/claude_code.py b/tests/integration/_support/claude_code.py index 7bb05ce941e..605ce7bd50d 100644 --- a/tests/integration/_support/claude_code.py +++ b/tests/integration/_support/claude_code.py @@ -31,6 +31,8 @@ CLAUDE_CODE_REASONING_BETAS: Final = ( "interleaved-thinking-2025-05-14", "thinking-token-count-2026-05-13", ) +CLAUDE_CODE_CACHING_BETAS: Final = ("extended-cache-ttl-2025-04-11", "prompt-caching-scope-2026-01-05") +CACHING_CLI_BETA: Final = f"{CLI_BETA},extended-cache-ttl-2025-04-11" REASONING_FIELDS: Final = ("thinking", "output_config", "reasoning_effort", "temperature") _OUTPUT_USAGE_KEYS: Final = frozenset({"output_tokens", "output_tokens_details"}) METADATA_USER_ID: Final = json.dumps( @@ -731,6 +733,114 @@ def forwarded(sent: Mapping[str, JsonValue], request: Request) -> Forwarded: ) +@dataclass(frozen=True, slots=True) +class CacheForwarded: + breakpoints: dict[str, JsonValue] + other_changes: dict[str, JsonValue] + caching_betas: tuple[str, ...] + + +def _marked(prefix: str, blocks: JsonValue) -> tuple[tuple[str, JsonValue], ...]: + if not isinstance(blocks, list): + return () + return tuple( + (f"{prefix}[{index}]", block["cache_control"]) + for index, block in enumerate(blocks) + if isinstance(block, dict) and "cache_control" in block + ) + + +def _marked_message(index: int, message: JsonValue) -> tuple[tuple[str, JsonValue], ...]: + return _marked(f"messages[{index}].content", message.get("content")) if isinstance(message, dict) else () + + +def _marked_messages(messages: JsonValue) -> tuple[tuple[str, JsonValue], ...]: + if not isinstance(messages, list): + return () + return tuple(chain.from_iterable(_marked_message(index, message) for index, message in enumerate(messages))) + + +def cache_breakpoints(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + top_level: Final = (("cache_control", body["cache_control"]),) if "cache_control" in body else () + return dict( + ( + *top_level, + *_marked("system", body.get("system")), + *_marked("tools", body.get("tools")), + *_marked_messages(body.get("messages")), + ) + ) + + +def _unmarked_block(block: JsonValue) -> JsonValue: + if not isinstance(block, dict): + return block + return {key: value for key, value in block.items() if key != "cache_control"} + + +def _unmarked_blocks(blocks: JsonValue) -> JsonValue: + return [_unmarked_block(block) for block in blocks] if isinstance(blocks, list) else blocks + + +def _unmarked_message(message: JsonValue) -> JsonValue: + if not isinstance(message, dict) or "content" not in message: + return message + return {**message, "content": _unmarked_blocks(message["content"])} + + +def _unmarked_value(key: str, value: JsonValue) -> JsonValue: + match key, value: + case ("system" | "tools", _): + return _unmarked_blocks(value) + case ("messages", list()): + return [_unmarked_message(message) for message in value] + case _: + return value + + +def with_block_breakpoint( + body: Mapping[str, JsonValue], field: str, index: int, control: Mapping[str, JsonValue] +) -> dict[str, JsonValue]: + blocks: Final = body[field] + assert isinstance(blocks, list), body + return { + **body, + field: [ + {**block, "cache_control": dict(control)} if position == index and isinstance(block, dict) else block + for position, block in enumerate(blocks) + ], + } + + +def without_cache_breakpoints(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return {key: _unmarked_value(key, value) for key, value in body.items() if key != "cache_control"} + + +def caching_betas(anthropic_beta: str) -> tuple[str, ...]: + return tuple(sorted(beta for beta in anthropic_beta.split(",") if beta in CLAUDE_CODE_CACHING_BETAS)) + + +def _without_model(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return {key: value for key, value in without_cache_breakpoints(body).items() if key != "model"} + + +def cache_forwarded(sent: Mapping[str, JsonValue], request: Request) -> CacheForwarded: + body: Final = JSON_OBJECT.validate_json(request.body) + return CacheForwarded( + breakpoints=cache_breakpoints(body), + other_changes=body_diff(_without_model(sent), _without_model(body)), + caching_betas=caching_betas(request.headers.get("anthropic-beta", "")), + ) + + +def streamed_start_usage(stream: str) -> JsonValue: + return next( + JSON_OBJECT.validate_python(data["message"]).get("usage") + for event, data in _client_events(stream) + if event == "message_start" + ) + + def _appended(block: Mapping[str, JsonValue], key: str, text: str) -> dict[str, JsonValue]: return {**block, key: f"{block.get(key) or ''}{text}"} @@ -848,7 +958,11 @@ def message_reply( def message_stream( - identity: str, model: str, content: tuple[dict[str, JsonValue], ...], usage: dict[str, JsonValue] + identity: str, + model: str, + content: tuple[dict[str, JsonValue], ...], + usage: dict[str, JsonValue], + final_usage: Mapping[str, JsonValue] | None = None, ) -> tuple[bytes, ...]: start: Final = sse_frame( "message_start", @@ -871,7 +985,7 @@ def message_stream( { "type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": None}, - "usage": _delta_usage(usage), + "usage": _delta_usage(usage) if final_usage is None else dict(final_usage), }, ) blocks: Final = chain.from_iterable(_block_frames(index, block) for index, block in enumerate(content)) diff --git a/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_breakpoint_injection_wire.py b/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_breakpoint_injection_wire.py new file mode 100644 index 00000000000..7027e624924 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_breakpoint_injection_wire.py @@ -0,0 +1,234 @@ +import uuid +from collections.abc import Callable, Mapping +from typing import Final + +import pytest +from integration._support import claude_code as cc +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue + +_MODEL: Final = "claude-sonnet-4-5" +_FIVE_MINUTES: Final = {"type": "ephemeral"} +_PONG: Final[tuple[dict[str, JsonValue], ...]] = ({"type": "text", "text": "PONG"},) +_SUBAGENT_BILLING: Final = ( + "x-anthropic-billing-header: cc_version=2.1.283.00; cc_entrypoint=sdk-cli; cc_is_subagent=true;" +) +_MAIN_AGENT_BILLING: Final = "x-anthropic-billing-header: cc_version=2.1.283.00; cc_entrypoint=sdk-cli;" + + +def _unmarked_claude_code_turn() -> dict[str, JsonValue]: + return cc.without_cache_breakpoints({**cc.claude_code_request(f"cache-bust-{uuid.uuid4().hex}"), "stream": False}) + + +def _one_shot_turn(billing: str) -> dict[str, JsonValue]: + return { + "system": [{"type": "text", "text": billing}], + "messages": [{"role": "user", "content": [{"type": "text", "text": f"Summarize {uuid.uuid4().hex}"}]}], + "max_tokens": 512, + "stream": False, + } + + +def _forward( + gateway: Gateway, + sent: Mapping[str, JsonValue], + key_fields: Mapping[str, JsonValue], + deployment: Mapping[str, JsonValue], +) -> tuple[cc.CacheForwarded, dict[str, JsonValue], dict[str, JsonValue]]: + identity: Final = f"msg_{uuid.uuid4().hex}" + + def respond(request: Request) -> Reply: + return Reply(body=cc.message_reply(identity, _MODEL, _PONG, {"input_tokens": 12, "output_tokens": 4})) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=cc.ANTHROPIC_API_KEY, **deployment + ) + key: Final = scenario.key(**key_fields) + response: Final = gateway.request( + "POST", "/v1/messages", {**sent, "model": model}, key=key, headers=cc.cli_headers(key, cc.CACHING_CLI_BETA) + ) + assert response.status_code == 200, response.text + upstream: Final = wire.drain() + rows: Final = eventually( + lambda: read_rows( + "SELECT model_id, metadata->>'litellm_gateway_injected_cache' AS injected_for " + 'FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (identity,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert len(upstream) == 1, upstream + return cc.cache_forwarded(sent, upstream[0]), cc.JSON_OBJECT.validate_json(upstream[0].body), rows[0] + + +def test_prompt_caching_key_adds_breakpoints_to_the_system_prompt_and_last_message(gateway: Gateway) -> None: + sent: Final = _unmarked_claude_code_turn() + forwarded, _, spend_row = _forward(gateway, sent, {"enable_prompt_caching": True}, {}) + assert cc.cache_breakpoints(sent) == {} + assert forwarded.breakpoints == {"system[2]": _FIVE_MINUTES, "messages[0].content[7]": _FIVE_MINUTES} + assert forwarded.other_changes == {} + assert forwarded.caching_betas == ("extended-cache-ttl-2025-04-11", "prompt-caching-scope-2026-01-05") + assert spend_row["injected_for"] == spend_row["model_id"], spend_row + + +def test_prompt_caching_key_turns_a_string_system_prompt_into_a_cached_block(gateway: Gateway) -> None: + sent: Final = {**_unmarked_claude_code_turn(), "system": "Synthetic agent identity system prompt."} + forwarded, _, _ = _forward(gateway, sent, {"enable_prompt_caching": True}, {}) + assert forwarded.breakpoints == {"system[0]": _FIVE_MINUTES, "messages[0].content[7]": _FIVE_MINUTES} + assert forwarded.other_changes == { + "system": { + "expected": "Synthetic agent identity system prompt.", + "upstream": [{"type": "text", "text": "Synthetic agent identity system prompt."}], + } + } + + +def test_key_without_prompt_caching_adds_no_breakpoints(gateway: Gateway) -> None: + sent: Final = _unmarked_claude_code_turn() + forwarded, _, spend_row = _forward(gateway, sent, {}, {}) + assert forwarded.breakpoints == {} + assert forwarded.other_changes == {} + assert spend_row["injected_for"] is None, spend_row + + +def test_prompt_caching_key_adds_nothing_when_the_client_marked_any_breakpoint(gateway: Gateway) -> None: + sent: Final = cc.with_block_breakpoint(_unmarked_claude_code_turn(), "tools", 23, _FIVE_MINUTES) + forwarded, _, spend_row = _forward(gateway, sent, {"enable_prompt_caching": True}, {}) + assert forwarded.breakpoints == {"tools[23]": _FIVE_MINUTES} + assert forwarded.other_changes == {} + assert spend_row["injected_for"] is None, spend_row + + +def test_prompt_caching_key_adds_nothing_when_extra_body_unmarks_the_client_tool_breakpoint(gateway: Gateway) -> None: + marked: Final = cc.with_block_breakpoint(_unmarked_claude_code_turn(), "tools", 23, _FIVE_MINUTES) + envelope: Final = {"tools": cc.without_cache_breakpoints(marked)["tools"]} + sent: Final = {**marked, "extra_body": envelope} + forwarded, _, spend_row = _forward(gateway, sent, {"enable_prompt_caching": True}, {}) + assert forwarded.breakpoints == {"tools[23]": _FIVE_MINUTES} + assert forwarded.other_changes == {"extra_body": {"expected": envelope, "upstream": None}} + assert spend_row["injected_for"] is None, spend_row + + +@pytest.mark.parametrize( + ("billing", "received"), + ( + pytest.param(_SUBAGENT_BILLING, {}, id="subagent-gets-no-breakpoints"), + pytest.param( + _MAIN_AGENT_BILLING, + {"system[0]": _FIVE_MINUTES, "messages[0].content[0]": _FIVE_MINUTES}, + id="main-agent-gets-breakpoints", + ), + ), +) +def test_prompt_caching_key_skips_one_shot_claude_code_subagent_requests( + gateway: Gateway, billing: str, received: dict[str, JsonValue] +) -> None: + sent: Final = _one_shot_turn(billing) + forwarded, _, _ = _forward(gateway, sent, {"enable_prompt_caching": True}, {}) + assert forwarded.breakpoints == received + assert forwarded.other_changes == {} + + +def _three_client_breakpoints() -> dict[str, JsonValue]: + system_marked: Final = cc.with_block_breakpoint(_unmarked_claude_code_turn(), "system", 1, _FIVE_MINUTES) + return {**cc.with_block_breakpoint(system_marked, "tools", 23, _FIVE_MINUTES), "cache_control": dict(_FIVE_MINUTES)} + + +def _four_client_breakpoints() -> dict[str, JsonValue]: + return cc.with_block_breakpoint(_three_client_breakpoints(), "system", 2, _FIVE_MINUTES) + + +@pytest.mark.parametrize( + ("sent", "received"), + ( + pytest.param( + _three_client_breakpoints, + { + "cache_control": _FIVE_MINUTES, + "system[1]": _FIVE_MINUTES, + "tools[23]": _FIVE_MINUTES, + "messages[0].content[7]": _FIVE_MINUTES, + }, + id="three-client-breakpoints-plus-last-message", + ), + pytest.param( + _four_client_breakpoints, + { + "cache_control": _FIVE_MINUTES, + "system[1]": _FIVE_MINUTES, + "system[2]": _FIVE_MINUTES, + "tools[23]": _FIVE_MINUTES, + }, + id="four-client-breakpoints-leave-no-room", + ), + ), +) +def test_configured_injection_points_stop_at_anthropics_four_breakpoint_limit( + gateway: Gateway, sent: Callable[[], dict[str, JsonValue]], received: dict[str, JsonValue] +) -> None: + client_body: Final = sent() + forwarded, _, _ = _forward( + gateway, + client_body, + {}, + { + "cache_control_injection_points": [ + {"location": "message", "role": "system"}, + {"location": "message", "index": -1}, + ] + }, + ) + assert forwarded.breakpoints == received + assert forwarded.other_changes == {} + + +def test_fallback_deployment_gets_its_own_breakpoints_and_the_injection_credit(gateway: Gateway) -> None: + identity: Final = f"msg_{uuid.uuid4().hex}" + sent: Final = _unmarked_claude_code_turn() + + def fail(request: Request) -> Reply: + return Reply(status=500, body=b'{"type":"error","error":{"type":"api_error","message":"primary down"}}') + + def respond(request: Request) -> Reply: + return Reply(body=cc.message_reply(identity, _MODEL, _PONG, {"input_tokens": 12, "output_tokens": 4})) + + with wire_server(fail) as primary_wire, wire_server(respond) as fallback_wire, gateway.scenario() as scenario: + primary: Final = scenario.model( + model=f"anthropic/{_MODEL}", api_base=primary_wire.url, api_key=cc.ANTHROPIC_API_KEY + ) + fallback: Final = scenario.model( + model=f"anthropic/{_MODEL}", api_base=fallback_wire.url, api_key=cc.ANTHROPIC_API_KEY + ) + key: Final = scenario.key(enable_prompt_caching=True) + response: Final = gateway.request( + "POST", + "/v1/messages", + {**sent, "model": primary, "fallbacks": [fallback]}, + key=key, + headers=cc.cli_headers(key, cc.CACHING_CLI_BETA), + ) + assert response.status_code == 200, response.text + primary_legs: Final = primary_wire.drain() + fallback_legs: Final = fallback_wire.drain() + rows: Final = eventually( + lambda: read_rows( + "SELECT spend.model_id, deployment.model_name, " + "spend.metadata->>'litellm_gateway_injected_cache' AS injected_for " + 'FROM "LiteLLM_SpendLogs" spend JOIN "LiteLLM_ProxyModelTable" deployment ' + "ON deployment.model_id = spend.model_id WHERE spend.request_id=%s", + (identity,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + injected: Final = {"system[2]": _FIVE_MINUTES, "messages[0].content[7]": _FIVE_MINUTES} + legs: Final = [cc.cache_forwarded(sent, leg) for leg in (*primary_legs, *fallback_legs)] + assert len(primary_legs) == 1, primary_legs + assert len(fallback_legs) == 1, fallback_legs + assert [(leg.breakpoints, leg.other_changes) for leg in legs] == [(injected, {}), (injected, {})], legs + assert rows[0]["model_name"] == fallback, rows + assert rows[0]["injected_for"] == rows[0]["model_id"], rows diff --git a/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_breakpoint_passthrough_wire.py b/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_breakpoint_passthrough_wire.py new file mode 100644 index 00000000000..144d209073c --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_breakpoint_passthrough_wire.py @@ -0,0 +1,86 @@ +import uuid +from collections.abc import Callable +from typing import Final + +import pytest +from integration._support import claude_code as cc +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue + +_MODEL: Final = "claude-sonnet-4-5" +_FIVE_MINUTES: Final = {"type": "ephemeral"} +_ONE_HOUR: Final = {"type": "ephemeral", "ttl": "1h"} +_ONE_HOUR_GLOBAL: Final = {"type": "ephemeral", "ttl": "1h", "scope": "global"} +_CLAUDE_CODE_BREAKPOINTS: Final = { + "system[1]": _FIVE_MINUTES, + "system[2]": _FIVE_MINUTES, + "messages[0].content[7]": _FIVE_MINUTES, +} + + +def _claude_code_default(body: dict[str, JsonValue]) -> dict[str, JsonValue]: + return body + + +def _one_hour_global_system_prompt(body: dict[str, JsonValue]) -> dict[str, JsonValue]: + return cc.with_block_breakpoint(body, "system", 2, _ONE_HOUR_GLOBAL) + + +def _one_hour_tool_breakpoint(body: dict[str, JsonValue]) -> dict[str, JsonValue]: + return cc.with_block_breakpoint(body, "tools", 23, _ONE_HOUR) + + +def _top_level_automatic_caching(body: dict[str, JsonValue]) -> dict[str, JsonValue]: + return {**body, "cache_control": dict(_FIVE_MINUTES)} + + +@pytest.mark.parametrize( + ("mark", "received"), + ( + pytest.param(_claude_code_default, _CLAUDE_CODE_BREAKPOINTS, id="claude-code-default"), + pytest.param( + _one_hour_global_system_prompt, + {**_CLAUDE_CODE_BREAKPOINTS, "system[2]": _ONE_HOUR_GLOBAL}, + id="1h-global-scope-on-system-prompt", + ), + pytest.param( + _one_hour_tool_breakpoint, + {**_CLAUDE_CODE_BREAKPOINTS, "tools[23]": _ONE_HOUR}, + id="1h-on-last-tool", + ), + pytest.param( + _top_level_automatic_caching, + {"cache_control": _FIVE_MINUTES, **_CLAUDE_CODE_BREAKPOINTS}, + id="top-level-automatic-caching", + ), + ), +) +def test_client_cache_breakpoints_reach_anthropic_unchanged( + gateway: Gateway, mark: Callable[[dict[str, JsonValue]], dict[str, JsonValue]], received: dict[str, JsonValue] +) -> None: + sent: Final = mark({**cc.claude_code_request(f"cache-bust-{uuid.uuid4().hex}"), "stream": False}) + + def respond(request: Request) -> Reply: + return Reply( + body=cc.message_reply( + f"msg_{uuid.uuid4().hex}", + _MODEL, + ({"type": "text", "text": "PONG"},), + {"input_tokens": 12, "output_tokens": 4}, + ) + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=cc.ANTHROPIC_API_KEY) + response: Final = gateway.request( + "POST", "/v1/messages", {**sent, "model": model}, headers=cc.cli_headers(gateway.key, cc.CACHING_CLI_BETA) + ) + assert response.status_code == 200, response.text + upstream: Final = wire.drain() + assert len(upstream) == 1, upstream + forwarded: Final = cc.cache_forwarded(sent, upstream[0]) + assert cc.cache_breakpoints(sent) == received + assert forwarded.breakpoints == received + assert forwarded.other_changes == {} + assert forwarded.caching_betas == ("extended-cache-ttl-2025-04-11", "prompt-caching-scope-2026-01-05") diff --git a/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_token_pricing_wire.py b/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_token_pricing_wire.py new file mode 100644 index 00000000000..e5cc0247177 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_token_pricing_wire.py @@ -0,0 +1,159 @@ +import uuid +from collections.abc import Mapping +from typing import Final + +import pytest +from integration._support import claude_code as cc +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, wire_server +from pydantic import BaseModel, ConfigDict, JsonValue + +_MODEL: Final = "claude-sonnet-4-5" +_FAST_MODE_MODEL: Final = "claude-opus-5-5" +_INPUT_RATE: Final = 1e-6 +_OUTPUT_RATE: Final = 2e-6 +_CACHE_READ_RATE: Final = 1e-7 +_CACHE_WRITE_5M_RATE: Final = 1.25e-6 +_CACHE_WRITE_1H_RATE: Final = 2e-6 +_FAST_MULTIPLIER: Final = 6.0 +_RATES: Final = { + "input_cost_per_token": _INPUT_RATE, + "output_cost_per_token": _OUTPUT_RATE, + "cache_read_input_token_cost": _CACHE_READ_RATE, + "cache_creation_input_token_cost": _CACHE_WRITE_5M_RATE, + "cache_creation_input_token_cost_above_1hr": _CACHE_WRITE_1H_RATE, +} +_CONTENT: Final[tuple[dict[str, JsonValue], ...]] = ({"type": "text", "text": "PONG"},) +_USAGE: Final[dict[str, JsonValue]] = { + "input_tokens": 100, + "output_tokens": 50, + "cache_read_input_tokens": 400, + "cache_creation_input_tokens": 300, + "cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200}, +} + + +class _BilledRow(BaseModel): + model_config = ConfigDict(frozen=True) + + spend: float + prompt_tokens: int + cache_read_cost: float + cache_creation_cost: float + + +_STREAM_IDS: Final = (pytest.param(False, id="non-streamed"), pytest.param(True, id="streamed")) + + +def _reply(identity: str, model: str, usage: dict[str, JsonValue], stream: bool) -> Reply: + if stream: + return Reply(chunks=cc.message_stream(identity, model, _CONTENT, usage), content_type="text/event-stream") + return Reply(body=cc.message_reply(identity, model, _CONTENT, usage)) + + +def _billed_row( + gateway: Gateway, + model_name: str, + usage: dict[str, JsonValue], + stream: bool, + request_fields: Mapping[str, JsonValue], + model_info: Mapping[str, JsonValue] | None, +) -> _BilledRow: + identity: Final = f"msg_{uuid.uuid4().hex}" + + def respond(request: Request) -> Reply: + return _reply(identity, model_name, usage, stream) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"anthropic/{model_name}", + api_base=wire.url, + api_key=cc.ANTHROPIC_API_KEY, + model_info=model_info, + **_RATES, + ) + body: Final = { + **cc.claude_code_request(f"cache-bust-{uuid.uuid4().hex}"), + **request_fields, + "stream": stream, + "model": model, + } + response: Final = gateway.request("POST", "/v1/messages", body) + assert response.status_code == 200, response.text + assert len(wire.drain()) == 1 + rows: Final = eventually( + lambda: read_rows( + "SELECT spend, prompt_tokens, " + "(metadata->'cost_breakdown'->>'cache_read_cost')::float AS cache_read_cost, " + "(metadata->'cost_breakdown'->>'cache_creation_cost')::float AS cache_creation_cost " + 'FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (identity,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + return _BilledRow.model_validate(rows[0]) + + +@pytest.mark.parametrize("stream", _STREAM_IDS) +def test_cache_reads_and_5m_and_1h_cache_writes_are_each_billed_at_their_own_rate( + gateway: Gateway, stream: bool +) -> None: + row: Final = _billed_row(gateway, _MODEL, _USAGE, stream, {}, None) + input_tokens, output_tokens, cache_read, cache_write_5m, cache_write_1h = 100, 50, 400, 100, 200 + assert row.prompt_tokens == input_tokens + cache_read + cache_write_5m + cache_write_1h, row + assert row.cache_read_cost == pytest.approx(cache_read * _CACHE_READ_RATE), row + assert row.cache_creation_cost == pytest.approx( + cache_write_5m * _CACHE_WRITE_5M_RATE + cache_write_1h * _CACHE_WRITE_1H_RATE + ), row + assert row.spend == pytest.approx( + input_tokens * _INPUT_RATE + + output_tokens * _OUTPUT_RATE + + cache_read * _CACHE_READ_RATE + + cache_write_5m * _CACHE_WRITE_5M_RATE + + cache_write_1h * _CACHE_WRITE_1H_RATE + ), row + + +@pytest.mark.parametrize("stream", _STREAM_IDS) +def test_fast_mode_multiplies_cache_reads_and_writes_along_with_input_and_output( + gateway: Gateway, stream: bool +) -> None: + row: Final = _billed_row( + gateway, + _FAST_MODE_MODEL, + {**_USAGE, "speed": "fast"}, + stream, + {"speed": "fast"}, + {"provider_specific_entry": {"fast": _FAST_MULTIPLIER}}, + ) + input_tokens, output_tokens, cache_read, cache_write_5m, cache_write_1h = 100, 50, 400, 100, 200 + assert row.spend == pytest.approx( + _FAST_MULTIPLIER + * ( + input_tokens * _INPUT_RATE + + output_tokens * _OUTPUT_RATE + + cache_read * _CACHE_READ_RATE + + cache_write_5m * _CACHE_WRITE_5M_RATE + + cache_write_1h * _CACHE_WRITE_1H_RATE + ) + ), row + + +@pytest.mark.parametrize("stream", _STREAM_IDS) +def test_fast_mode_cost_breakdown_reports_the_multiplied_cache_costs(gateway: Gateway, stream: bool) -> None: + pytest.skip("BUG: fast mode bills 6x cache costs but cost_breakdown.cache_read_cost/cache_creation_cost stay 1x") + row: Final = _billed_row( + gateway, + _FAST_MODE_MODEL, + {**_USAGE, "speed": "fast"}, + stream, + {"speed": "fast"}, + {"provider_specific_entry": {"fast": _FAST_MULTIPLIER}}, + ) + cache_read, cache_write_5m, cache_write_1h = 400, 100, 200 + assert row.cache_read_cost == pytest.approx(_FAST_MULTIPLIER * cache_read * _CACHE_READ_RATE), row + assert row.cache_creation_cost == pytest.approx( + _FAST_MULTIPLIER * (cache_write_5m * _CACHE_WRITE_5M_RATE + cache_write_1h * _CACHE_WRITE_1H_RATE) + ), row diff --git a/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_usage_response_wire.py b/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_usage_response_wire.py new file mode 100644 index 00000000000..fcb6d1af3cf --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/caching/test_anthropic_cache_usage_response_wire.py @@ -0,0 +1,75 @@ +import uuid +from typing import Final + +from integration._support import claude_code as cc +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue + +_MODEL: Final = "claude-sonnet-4-5" +_ANTHROPIC_USAGE: Final = { + "input_tokens": 100, + "output_tokens": 50, + "cache_read_input_tokens": 400, + "cache_creation_input_tokens": 300, + "cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200}, +} + + +def _client_body(stream: bool) -> dict[str, JsonValue]: + return {**cc.claude_code_request(f"cache-bust-{uuid.uuid4().hex}"), "stream": stream} + + +def test_streamed_cache_read_and_write_tokens_reach_the_client_unchanged(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + return Reply( + chunks=cc.message_stream( + f"msg_{uuid.uuid4().hex}", + _MODEL, + ({"type": "text", "text": "PONG"},), + _ANTHROPIC_USAGE, + final_usage=_ANTHROPIC_USAGE, + ), + content_type="text/event-stream", + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=cc.ANTHROPIC_API_KEY) + response: Final = gateway.request("POST", "/v1/messages", {**_client_body(stream=True), "model": model}) + assert response.status_code == 200, response.text + assert len(wire.drain()) == 1 + assert cc.streamed_start_usage(response.text) == { + "input_tokens": 100, + "cache_read_input_tokens": 400, + "cache_creation_input_tokens": 300, + "cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200}, + }, response.text + assert cc.streamed_usage(response.text) == { + "input_tokens": 100, + "output_tokens": 50, + "cache_read_input_tokens": 400, + "cache_creation_input_tokens": 300, + "cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200}, + }, response.text + + +def test_non_streamed_cache_read_and_write_tokens_reach_the_client_unchanged(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + return Reply( + body=cc.message_reply( + f"msg_{uuid.uuid4().hex}", _MODEL, ({"type": "text", "text": "PONG"},), _ANTHROPIC_USAGE + ) + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=cc.ANTHROPIC_API_KEY) + response: Final = gateway.request("POST", "/v1/messages", {**_client_body(stream=False), "model": model}) + assert response.status_code == 200, response.text + assert len(wire.drain()) == 1 + assert response.json()["usage"] == { + "input_tokens": 100, + "output_tokens": 50, + "cache_read_input_tokens": 400, + "cache_creation_input_tokens": 300, + "cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200}, + }, response.text