test(e2e): add tool-call, terminal, and image-input shapes to the cost matrix

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-16 16:38:15 +00:00
parent 415b06f5ff
commit feb69c5f78
4 changed files with 688 additions and 76 deletions

View file

@ -17,7 +17,11 @@ creation), the case is absent from the matrix rather than silently zero.
from __future__ import annotations
import base64
import json
import random
import struct
import zlib
from collections.abc import Mapping
from dataclasses import dataclass
from pathlib import Path
@ -26,7 +30,7 @@ from typing import Final, Literal, TypeAlias
from pydantic import BaseModel, ConfigDict, TypeAdapter
from scripted_provider import Scenario, ScriptedOutput, ScriptedUsage, Wire
from scripted_provider import Scenario, ScriptedOutput, ScriptedToolCall, ScriptedUsage, Wire
COST_MAP_PATH: Final = Path(__file__).resolve().parent.parent / "cost_map.json"
@ -182,26 +186,37 @@ _WIRE_CAPS: Final[Mapping[str, frozenset[str]]] = MappingProxyType({
"openai_chat": frozenset(
{
"cache_read", "cache_write_5m", "cache_write_1h", "reasoning", "audio",
"web_search", "response_model", "absent_usage",
"web_search", "response_model", "absent_usage", "tool_call", "image_input",
}
),
"openai_responses": frozenset(
{
"cache_read", "reasoning", "web_search", "response_model", "absent_usage",
"tool_call", "image_input", "responses_terminal",
}
),
"openai_responses": frozenset({"cache_read", "reasoning", "web_search", "response_model", "absent_usage"}),
"anthropic_messages": frozenset(
{"cache_read", "cache_write_5m", "cache_write_1h", "web_search", "response_model", "absent_usage"}
{
"cache_read", "cache_write_5m", "cache_write_1h", "web_search",
"response_model", "absent_usage", "tool_call", "image_input",
}
),
"gemini_generate": frozenset(
{"cache_read", "reasoning", "audio", "web_search", "response_model", "absent_usage"}
{
"cache_read", "reasoning", "audio", "web_search", "response_model",
"absent_usage", "tool_call", "image_input", "prompt_blocked",
}
),
"together_chat": frozenset(
{
"cache_read", "cache_write_5m", "cache_write_1h", "reasoning", "audio",
"web_search", "response_model", "absent_usage",
"web_search", "response_model", "absent_usage", "tool_call", "image_input",
}
),
"fireworks_chat": frozenset(
{
"cache_read", "cache_write_5m", "cache_write_1h", "reasoning", "audio",
"web_search", "response_model", "absent_usage",
"web_search", "response_model", "absent_usage", "tool_call", "image_input",
}
),
})
@ -220,6 +235,15 @@ CaseName: TypeAlias = Literal[
"stream",
"stream_no_usage",
"response_model_override",
"stream_response_model_override",
"tool_call",
"stream_no_usage_tool_call",
"stream_no_usage_image_input",
"stream_no_usage_incomplete",
"stream_unvalidated",
"stream_no_usage_unvalidated",
"prompt_blocked",
"stream_prompt_blocked",
]
@ -237,6 +261,9 @@ class Case:
billed_web_search_calls: int = 0
response_model_override: bool = False
exact_spend: bool = True
tool_call: bool = False
image_input: bool = False
terminal: Literal["completed", "incomplete", "unvalidated", "prompt_blocked"] = "completed"
def scenario(self, scenario_id: str, model: FrontierModel, text: str) -> Scenario:
return Scenario(
@ -246,6 +273,10 @@ class Case:
output=ScriptedOutput(
text=text,
response_model=model.override_model if self.response_model_override else None,
tool_call=ScriptedToolCall(name="get_weather", arguments=TOOL_CALL_ARGUMENTS)
if self.tool_call
else None,
terminal=self.terminal,
),
stream_usage=self.stream_usage,
service_tier=self.service_tier,
@ -254,6 +285,15 @@ class Case:
_BASIC_USAGE: Final = ScriptedUsage(fresh_input_tokens=120, output_tokens=40)
TOOL_CALL_ARGUMENTS: Final = json.dumps({
"city": "Berlin",
"days": 7,
"units": "metric",
"notes": "filler " * 30,
})
_PROMPT_BLOCKED_USAGE: Final = ScriptedUsage(fresh_input_tokens=1000, output_tokens=0)
def _web_search_case(model: FrontierModel) -> Case:
counts_exactly: Final = model.wire in ("openai_responses", "anthropic_messages", "gemini_generate")
@ -362,6 +402,100 @@ def cases_for(model: FrontierModel) -> tuple[Case, ...]:
if "response_model" in caps
else None
),
(
Case(
name="stream_response_model_override",
usage=_BASIC_USAGE,
stream=True,
response_model_override=True,
)
if "response_model" in caps
else None
),
(
Case(name="tool_call", usage=_BASIC_USAGE, tool_call=True)
if "tool_call" in caps
else None
),
(
Case(
name="stream_no_usage_tool_call",
usage=_BASIC_USAGE,
stream=True,
stream_usage="absent",
tool_call=True,
exact_spend=False,
)
if "absent_usage" in caps and "tool_call" in caps
else None
),
(
Case(
name="stream_no_usage_image_input",
usage=_BASIC_USAGE,
stream=True,
stream_usage="absent",
image_input=True,
exact_spend=False,
)
if "absent_usage" in caps and "image_input" in caps
else None
),
(
Case(
name="stream_no_usage_incomplete",
usage=_BASIC_USAGE,
stream=True,
stream_usage="absent",
terminal="incomplete",
exact_spend=False,
)
if "responses_terminal" in caps
else None
),
(
Case(
name="stream_unvalidated",
usage=_BASIC_USAGE,
stream=True,
terminal="unvalidated",
)
if "responses_terminal" in caps
else None
),
(
Case(
name="stream_no_usage_unvalidated",
usage=_BASIC_USAGE,
stream=True,
stream_usage="absent",
terminal="unvalidated",
exact_spend=False,
)
if "responses_terminal" in caps
else None
),
(
Case(
name="prompt_blocked",
usage=_PROMPT_BLOCKED_USAGE,
terminal="prompt_blocked",
response_model_override=True,
)
if "prompt_blocked" in caps
else None
),
(
Case(
name="stream_prompt_blocked",
usage=_PROMPT_BLOCKED_USAGE,
stream=True,
terminal="prompt_blocked",
response_model_override=True,
)
if "prompt_blocked" in caps
else None
),
)
return tuple(case for case in candidates if case is not None)
@ -436,6 +570,42 @@ def expected_cost(model: FrontierModel, case: Case) -> float:
return expected_breakdown(model, case).total
def recount_cost(
model: FrontierModel, case: Case, prompt_tokens: int, completion_tokens: int
) -> float:
"""What the proxy's own token recount should cost at the case's rates,
without pinning the tokenizer's exact counts."""
rates: Final = model.override_rates if case.response_model_override else model.rates
return prompt_tokens * (rates.input_cost_per_token or 0.0) + completion_tokens * (
rates.output_cost_per_token or 0.0
)
def _png_chunk(tag: bytes, payload: bytes) -> bytes:
return struct.pack(">I", len(payload)) + tag + payload + struct.pack(">I", zlib.crc32(tag + payload))
def image_input_data_url() -> str:
"""A deterministic 256x256 RGB noise PNG as a data URL; noise compresses
poorly on purpose so the base64 payload stays well above 100 KB and would
blow up the prompt recount if the URL were ever tokenized as text."""
rng: Final = random.Random(0)
side: Final = 256
raw: Final = b"".join(
b"\x00" + rng.randbytes(side * 3) for _ in range(side)
)
png: Final = (
b"\x89PNG\r\n\x1a\n"
+ _png_chunk(b"IHDR", struct.pack(">IIBBBBB", side, side, 8, 2, 0, 0, 0))
+ _png_chunk(b"IDAT", zlib.compress(raw))
+ _png_chunk(b"IEND", b"")
)
return "data:image/png;base64," + base64.b64encode(png).decode()
IMAGE_INPUT_DATA_URL: Final = image_input_data_url()
def expected_token_columns(model: FrontierModel, case: Case) -> tuple[int, int]:
"""(prompt_tokens, completion_tokens) the spend row should carry, per the
wire's normalization: Anthropic folds cache read/write into prompt_tokens,

View file

@ -38,7 +38,7 @@ from types import MappingProxyType
from typing import Final, Literal, TypeAlias
from urllib.parse import urlsplit
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError, model_validator
Wire: TypeAlias = Literal[
"openai_chat",
@ -62,6 +62,26 @@ WIRE_MOUNTS: Final[Mapping[str, str]] = MappingProxyType(
StreamUsage: TypeAlias = Literal["final_chunk", "absent"]
ServiceTier: TypeAlias = Literal["flex", "priority"]
TerminalKind: TypeAlias = Literal["completed", "incomplete", "unvalidated", "prompt_blocked"]
# Which terminal variant each wire can represent.
_TERMINAL_CAPS: Final[Mapping[str, frozenset[str]]] = MappingProxyType(
{
"openai_responses": frozenset({"incomplete", "unvalidated"}),
"gemini_generate": frozenset({"prompt_blocked"}),
}
)
class ScriptedToolCall(BaseModel):
"""A single function call the scripted output emits instead of text.
``arguments`` is the wire's JSON string (~250 chars), sliced into deltas
for streams."""
model_config = ConfigDict(frozen=True)
name: str
arguments: str
class ScriptedUsage(BaseModel):
@ -96,6 +116,12 @@ class ScriptedOutput(BaseModel):
# OpenAI-compatible providers can report a provider-computed cost; emitted as
# the top-level "cost" field on the together/fireworks wire.
provider_cost: float | None = None
# When set, the response is a tool call only: no text content on any wire.
tool_call: ScriptedToolCall | None = None
# Terminal shape: "unvalidated" makes the Responses terminal response fail
# pydantic validation so the proxy takes its model_construct dict path;
# "prompt_blocked" is a Gemini promptFeedback-only body.
terminal: TerminalKind = "completed"
class Scenario(BaseModel):
@ -108,6 +134,17 @@ class Scenario(BaseModel):
stream_usage: StreamUsage = "final_chunk"
service_tier: ServiceTier | None = None
@model_validator(mode="after")
def _check_terminal_supported(self) -> Scenario:
if (
self.output.terminal != "completed"
and self.output.terminal not in _TERMINAL_CAPS.get(self.wire, frozenset())
):
raise ValueError(
f"wire {self.wire} cannot emit terminal={self.output.terminal}"
)
return self
@property
def mount(self) -> str:
return WIRE_MOUNTS[self.wire]
@ -291,10 +328,38 @@ def _responses_usage(u: ScriptedUsage) -> Mapping[str, object]:
# ---------- per-wire responses ----------
def _split_arguments(arguments: str) -> tuple[str, ...]:
"""Slice a tool-call arguments JSON string into 2-3 streamed deltas."""
third: Final = max(1, len(arguments) // 3)
return tuple(
slice_
for slice_ in (arguments[:third], arguments[third : 2 * third], arguments[2 * third :])
if slice_
)
def _openai_message(scenario: Scenario) -> Mapping[str, object]:
tool_call: Final = scenario.output.tool_call
return _jobj_opt(
("role", "assistant"),
("content", scenario.output.text),
("content", None if tool_call is not None else scenario.output.text),
(
(
"tool_calls",
(
_jobj(
("id", f"call_{scenario.scenario_id}"),
("type", "function"),
(
"function",
_jobj(("name", tool_call.name), ("arguments", tool_call.arguments)),
),
),
),
)
if tool_call is not None
else None
),
(
(
"annotations",
@ -332,7 +397,12 @@ def _openai_chat_body(scenario: Scenario, requested_model: str) -> Mapping[str,
_jobj(
("index", 0),
("message", _openai_message(scenario)),
("finish_reason", scenario.output.finish_reason),
(
"finish_reason",
"tool_calls"
if scenario.output.tool_call is not None
else scenario.output.finish_reason,
),
),
),
),
@ -359,6 +429,7 @@ def _openai_chunk(
def _openai_chat_sse(scenario: Scenario, requested_model: str) -> bytes:
tool_call: Final = scenario.output.tool_call
delta: Final = _jobj_opt(
("role", "assistant"),
("content", scenario.output.text),
@ -368,6 +439,43 @@ def _openai_chat_sse(scenario: Scenario, requested_model: str) -> bytes:
else None
),
)
body_deltas: Final[tuple[Mapping[str, object], ...]] = (
(
_jobj(
("role", "assistant"),
(
"tool_calls",
(
_jobj(
("index", 0),
("id", f"call_{scenario.scenario_id}"),
("type", "function"),
(
"function",
_jobj(("name", tool_call.name), ("arguments", "")),
),
),
),
),
),
*(
_jobj(
(
"tool_calls",
(
_jobj(
("index", 0),
("function", _jobj(("arguments", arguments_slice))),
),
),
)
)
for arguments_slice in _split_arguments(tool_call.arguments)
),
)
if tool_call is not None
else (delta,)
)
return _sse(
(
(
@ -378,13 +486,16 @@ def _openai_chat_sse(scenario: Scenario, requested_model: str) -> bytes:
choices=(_jobj(("index", 0), ("delta", _jobj(("role", "assistant"))), ("finish_reason", None)),),
),
),
(
None,
_openai_chunk(
scenario,
requested_model,
choices=(_jobj(("index", 0), ("delta", delta), ("finish_reason", None)),),
),
*(
(
None,
_openai_chunk(
scenario,
requested_model,
choices=(_jobj(("index", 0), ("delta", body_delta), ("finish_reason", None)),),
),
)
for body_delta in body_deltas
),
(
None,
@ -395,7 +506,12 @@ def _openai_chat_sse(scenario: Scenario, requested_model: str) -> bytes:
_jobj(
("index", 0),
("delta", _jobj()),
("finish_reason", scenario.output.finish_reason),
(
"finish_reason",
"tool_calls"
if tool_call is not None
else scenario.output.finish_reason,
),
),
),
),
@ -410,17 +526,34 @@ def _openai_chat_sse(scenario: Scenario, requested_model: str) -> bytes:
)
def _anthropic_content(scenario: Scenario) -> tuple[Mapping[str, object], ...]:
tool_call: Final = scenario.output.tool_call
if tool_call is not None:
return (
_jobj(
("type", "tool_use"),
("id", f"toolu_{scenario.scenario_id}"),
("name", tool_call.name),
("input", json.loads(tool_call.arguments)),
),
)
return (_jobj(("type", "text"), ("text", scenario.output.text)),)
def _anthropic_stop_reason(scenario: Scenario) -> str:
if scenario.output.tool_call is not None:
return "tool_use"
return "end_turn" if scenario.output.finish_reason == "stop" else scenario.output.finish_reason
def _anthropic_body(scenario: Scenario, requested_model: str) -> Mapping[str, object]:
return _jobj(
("id", f"msg_{scenario.scenario_id}"),
("type", "message"),
("role", "assistant"),
("model", scenario.output.response_model or requested_model),
("content", (_jobj(("type", "text"), ("text", scenario.output.text)),)),
(
"stop_reason",
"end_turn" if scenario.output.finish_reason == "stop" else scenario.output.finish_reason,
),
("content", _anthropic_content(scenario)),
("stop_reason", _anthropic_stop_reason(scenario)),
("usage", _anthropic_usage(scenario.usage)),
)
@ -453,12 +586,7 @@ def _anthropic_sse(scenario: Scenario, requested_model: str) -> bytes:
("type", "message_delta"),
(
"delta",
_jobj(
(
"stop_reason",
"end_turn" if scenario.output.finish_reason == "stop" else scenario.output.finish_reason,
)
),
_jobj(("stop_reason", _anthropic_stop_reason(scenario))),
),
(
("usage", _jobj(("output_tokens", scenario.usage.output_tokens)))
@ -474,16 +602,45 @@ def _anthropic_sse(scenario: Scenario, requested_model: str) -> bytes:
_jobj(
("type", "content_block_start"),
("index", 0),
("content_block", _jobj(("type", "text"), ("text", ""))),
(
"content_block",
_jobj(
("type", "tool_use"),
("id", f"toolu_{scenario.scenario_id}"),
("name", scenario.output.tool_call.name),
("input", _jobj()),
)
if scenario.output.tool_call is not None
else _jobj(("type", "text"), ("text", "")),
),
),
),
(
"content_block_delta",
_jobj(
("type", "content_block_delta"),
("index", 0),
("delta", _jobj(("type", "text_delta"), ("text", scenario.output.text))),
),
*(
tuple(
(
"content_block_delta",
_jobj(
("type", "content_block_delta"),
("index", 0),
(
"delta",
_jobj(("type", "input_json_delta"), ("partial_json", arguments_slice)),
),
),
)
for arguments_slice in _split_arguments(scenario.output.tool_call.arguments)
)
if scenario.output.tool_call is not None
else (
(
"content_block_delta",
_jobj(
("type", "content_block_delta"),
("index", 0),
("delta", _jobj(("type", "text_delta"), ("text", scenario.output.text))),
),
),
)
),
("content_block_stop", _jobj(("type", "content_block_stop"), ("index", 0))),
("message_delta", message_delta),
@ -492,7 +649,49 @@ def _anthropic_sse(scenario: Scenario, requested_model: str) -> bytes:
)
def _gemini_prompt_blocked_body(scenario: Scenario, requested_model: str) -> Mapping[str, object]:
return _jobj(
(
"promptFeedback",
_jobj(
("blockReason", "SAFETY"),
(
"safetyRatings",
(
_jobj(
("category", "HARM_CATEGORY_HARASSMENT"),
("probability", "HIGH"),
("blocked", True),
),
),
),
),
),
("usageMetadata", _gemini_usage(scenario.usage)),
("modelVersion", scenario.output.response_model or requested_model),
)
def _gemini_parts(scenario: Scenario) -> tuple[Mapping[str, object], ...]:
tool_call: Final = scenario.output.tool_call
if tool_call is not None:
return (
_jobj(
(
"functionCall",
_jobj(
("name", tool_call.name),
("args", json.loads(tool_call.arguments)),
),
)
),
)
return (_jobj(("text", scenario.output.text)),)
def _gemini_body(scenario: Scenario, requested_model: str) -> Mapping[str, object]:
if scenario.output.terminal == "prompt_blocked":
return _gemini_prompt_blocked_body(scenario, requested_model)
return _jobj(
(
"candidates",
@ -501,7 +700,7 @@ def _gemini_body(scenario: Scenario, requested_model: str) -> Mapping[str, objec
(
"content",
_jobj(
("parts", (_jobj(("text", scenario.output.text)),)),
("parts", _gemini_parts(scenario)),
("role", "model"),
),
),
@ -559,67 +758,148 @@ def _gemini_sse(scenario: Scenario, requested_model: str) -> bytes:
)
def _responses_body(scenario: Scenario, requested_model: str) -> Mapping[str, object]:
return _jobj(
("id", f"resp_{scenario.scenario_id}"),
("object", "response"),
("created_at", int(time.time())),
("status", "completed"),
("model", scenario.output.response_model or requested_model),
(
"output",
def _responses_output(scenario: Scenario) -> tuple[Mapping[str, object], ...]:
tool_call: Final = scenario.output.tool_call
return (
*(
(
*(
_jobj(("type", "web_search_call"), ("id", f"ws_{i}"), ("status", "completed"))
for i in range(scenario.usage.web_search_calls)
),
_jobj(
("type", "message"),
("id", f"msg_{scenario.scenario_id}"),
("status", "completed"),
("role", "assistant"),
(
"content",
(
_jobj(
("type", "output_text"),
("text", scenario.output.text),
("annotations", ()),
),
),
_jobj(("type", "scripted_future_item"), ("id", f"fut_{scenario.scenario_id}"), ("status", "completed")),
)
if scenario.output.terminal == "unvalidated"
else ()
),
*(
_jobj(("type", "web_search_call"), ("id", f"ws_{i}"), ("status", "completed"))
for i in range(scenario.usage.web_search_calls)
),
_jobj(
("type", "function_call"),
("id", f"fc_{scenario.scenario_id}"),
("call_id", f"call_{scenario.scenario_id}"),
("name", tool_call.name),
("arguments", tool_call.arguments),
("status", "completed"),
)
if tool_call is not None
else _jobj(
("type", "message"),
("id", f"msg_{scenario.scenario_id}"),
("status", "completed"),
("role", "assistant"),
(
"content",
(
_jobj(
("type", "output_text"),
("text", scenario.output.text),
("annotations", ()),
),
),
),
),
)
def _responses_body(scenario: Scenario, requested_model: str) -> Mapping[str, object]:
incomplete: Final = scenario.output.terminal == "incomplete"
return _jobj_opt(
("id", f"resp_{scenario.scenario_id}"),
("object", "response"),
(
"created_at",
"not-a-number" if scenario.output.terminal == "unvalidated" else int(time.time()),
),
("status", "incomplete" if incomplete else "completed"),
(
("incomplete_details", _jobj(("reason", "max_output_tokens")))
if incomplete
else None
),
("model", scenario.output.response_model or requested_model),
("output", _responses_output(scenario)),
("usage", _responses_usage(scenario.usage)),
)
def _responses_sse(scenario: Scenario, requested_model: str) -> bytes:
completed: Final = (
tool_call: Final = scenario.output.tool_call
terminal: Final = (
_jobj(*((key, value) for key, value in _responses_body(scenario, requested_model).items() if key != "usage"))
if scenario.stream_usage == "absent"
else _responses_body(scenario, requested_model)
)
created: Final = _jobj(
*((key, value) for key, value in completed.items() if key not in ("status", "usage")),
*((key, value) for key, value in terminal.items() if key not in ("status", "usage")),
("status", "in_progress"),
("usage", None),
)
return _sse(
terminal_event: Final = (
"response.incomplete" if scenario.output.terminal == "incomplete" else "response.completed"
)
output_index: Final = (
scenario.usage.web_search_calls + (1 if scenario.output.terminal == "unvalidated" else 0)
)
middle_events: Final[tuple[tuple[str, Mapping[str, object]], ...]] = (
(
("response.created", _jobj(("type", "response.created"), ("response", created))),
(
"response.output_item.added",
_jobj(
("type", "response.output_item.added"),
("output_index", output_index),
(
"item",
_jobj(
("type", "function_call"),
("id", f"fc_{scenario.scenario_id}"),
("call_id", f"call_{scenario.scenario_id}"),
("name", tool_call.name),
("arguments", ""),
("status", "in_progress"),
),
),
),
),
*(
(
"response.function_call_arguments.delta",
_jobj(
("type", "response.function_call_arguments.delta"),
("item_id", f"fc_{scenario.scenario_id}"),
("output_index", output_index),
("delta", arguments_slice),
),
)
for arguments_slice in _split_arguments(tool_call.arguments)
),
(
"response.function_call_arguments.done",
_jobj(
("type", "response.function_call_arguments.done"),
("item_id", f"fc_{scenario.scenario_id}"),
("output_index", output_index),
("arguments", tool_call.arguments),
),
),
)
if tool_call is not None
else (
(
"response.output_text.delta",
_jobj(
("type", "response.output_text.delta"),
("item_id", f"msg_{scenario.scenario_id}"),
("output_index", scenario.usage.web_search_calls),
("output_index", output_index),
("content_index", 0),
("delta", scenario.output.text),
),
),
("response.completed", _jobj(("type", "response.completed"), ("response", completed))),
)
)
return _sse(
(
("response.created", _jobj(("type", "response.created"), ("response", created))),
*middle_events,
(terminal_event, _jobj(("type", terminal_event), ("response", terminal))),
)
)

View file

@ -16,15 +16,26 @@ from typing import Final
from conftest import CostCalcClient, cost_rows, register_scenario_deployment
from cost_matrix import (
FRONTIER_MODELS,
IMAGE_INPUT_DATA_URL,
Case,
FrontierModel,
cases_for,
expected_cost,
expected_token_columns,
recount_cost,
)
from e2e_config import unique_marker
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, ChatStreamOptions
from models import (
ChatBody,
ChatMessage,
ChatStreamOptions,
ChatTool,
ChatToolFunction,
ImageContentPart,
ImageUrl,
TextContentPart,
)
pytestmark: Final = [pytest.mark.e2e, pytest.mark.cost_map_stack] # mutable-ok: pytest only accepts a list for pytestmark
@ -41,10 +52,37 @@ def _case_id(param: tuple[FrontierModel, Case]) -> str:
def _chat_body(model_name: str, marker: str, case: Case) -> ChatBody:
return ChatBody(
model=model_name,
messages=(ChatMessage(role="user", content=f"{marker} scripted pricing call"),),
messages=(
ChatMessage(
role="user",
content=(
[
TextContentPart(text=f"{marker} scripted pricing call"),
ImageContentPart(image_url=ImageUrl(url=IMAGE_INPUT_DATA_URL)),
]
if case.image_input
else f"{marker} scripted pricing call"
),
),
),
stream=case.stream,
stream_options=ChatStreamOptions(include_usage=True) if case.stream else None,
service_tier=case.service_tier,
tools=(
(
ChatTool(
function=ChatToolFunction(
name="get_weather",
parameters={
"type": "object",
"properties": {"city": {"type": "string"}},
},
)
),
)
if case.tool_call
else None
),
)
@ -92,8 +130,23 @@ class TestTokenPricing:
if not case.exact_spend:
# stream_usage=absent: the provider reported no usage, so the row's
# token counts are the proxy's own recount; only assert a bill landed.
assert row.spend is not None and row.spend > 0, f"no-usage stream billed nothing: {row}"
# token counts are the proxy's own recount; assert the recount
# billed both directions at the case's rates.
assert row.prompt_tokens is not None and row.prompt_tokens > 0, (
f"no-usage stream counted no input tokens: {row}"
)
assert row.completion_tokens is not None and row.completion_tokens > 0, (
f"no-usage stream counted no output tokens: {row}"
)
if case.image_input:
assert row.prompt_tokens < 4000, (
f"image data URL looks tokenized as text: prompt_tokens={row.prompt_tokens}"
)
assert row.spend is not None and cost_rows.approx_equal(
row.spend,
recount_cost(model, case, row.prompt_tokens, row.completion_tokens),
), f"no-usage stream spend {row.spend} != recount at map rates: {row}"
cost_rows.assert_total_is_sum_of_components(row)
return
assert row.spend is not None and cost_rows.approx_equal(row.spend, expected), (

View file

@ -26,7 +26,7 @@ from cost_matrix import (
)
from e2e_config import unique_marker
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, ChatStreamOptions
from models import ChatBody, ChatMessage, ChatStreamOptions, ChatTool, ChatToolFunction
from scripted_provider import ScriptedUsage
pytestmark: Final = [pytest.mark.e2e, pytest.mark.cost_map_stack] # mutable-ok: pytest only accepts a list for pytestmark
@ -96,6 +96,53 @@ _WIRE_USAGE: Final[Mapping[str, tuple[str, ScriptedUsage]]] = MappingProxyType({
),
})
_SHAPE_USAGE: Final = ScriptedUsage(fresh_input_tokens=80, output_tokens=25)
# Renderer-level shapes the pricing matrix gates per cap, pinned here once per
# wire so the sidecar emits prove they survive the proxy end to end.
_SHAPES: Final[tuple[tuple[str, str, Case], ...]] = (
*(
(
f"tool_call_{'stream' if stream else 'sync'}",
wire,
Case(name="tool_call", usage=_SHAPE_USAGE, stream=stream, tool_call=True),
)
for wire in _WIRE_USAGE
for stream in (False, True)
),
(
"responses_incomplete",
"openai_responses",
Case(name="stream_no_usage_incomplete", usage=_SHAPE_USAGE, stream=True, terminal="incomplete"),
),
(
"responses_unvalidated",
"openai_responses",
Case(name="stream_unvalidated", usage=_SHAPE_USAGE, stream=True, terminal="unvalidated"),
),
(
"gemini_prompt_blocked",
"gemini_generate",
Case(
name="prompt_blocked",
usage=ScriptedUsage(fresh_input_tokens=1000, output_tokens=0),
terminal="prompt_blocked",
response_model_override=True,
),
),
(
"gemini_prompt_blocked_stream",
"gemini_generate",
Case(
name="stream_prompt_blocked",
usage=ScriptedUsage(fresh_input_tokens=1000, output_tokens=0),
stream=True,
terminal="prompt_blocked",
response_model_override=True,
),
),
)
class TestWireFormats:
@pytest.mark.parametrize("wire", tuple(_WIRE_USAGE))
@ -189,3 +236,65 @@ class TestWireFormats:
f"(breakdown {row.breakdown.model_dump()})"
)
cost_rows.assert_total_is_sum_of_components(row)
@pytest.mark.parametrize("shape_wire_case", _SHAPES, ids=lambda entry: entry[0])
@pytest.mark.covers("quota_management.spend_tracking.scripted_wire.logs_cost")
def test_response_shape_bills_reported_usage(
self,
client: CostCalcClient,
resources: ResourceManager,
scoped_key: str,
shape_wire_case: tuple[str, str, Case],
) -> None:
shape, wire, case = shape_wire_case
map_key, _usage = _WIRE_USAGE[wire]
model: Final = _MODELS[map_key]
marker: Final = unique_marker()
model_name, _handle = register_scenario_deployment(client, resources, model, case, marker)
response: Final = client.proxy.transport.send(
"/chat/completions",
headers=client.proxy.transport.bearer(scoped_key),
json=ChatBody(
model=model_name,
messages=(ChatMessage(role="user", content=f"{marker} scripted {shape}"),),
stream=case.stream,
stream_options=ChatStreamOptions(include_usage=True) if case.stream else None,
tools=(
(
ChatTool(
function=ChatToolFunction(
name="get_weather",
parameters={"type": "object", "properties": {"city": {"type": "string"}}},
)
),
)
if case.tool_call
else None
),
),
stream=case.stream,
)
assert response.ok, f"{shape}: proxy returned {response.status_code}: {response.body[:400]}"
if case.stream:
assert response.stream_done, f"{shape}: stream did not reach its terminal event"
assert response.stream_error is None, f"{shape}: stream error: {response.stream_error}"
expected: Final = expected_breakdown(model, case)
row: Final = cost_rows.poll_cost_row_where(
client.proxy,
scoped_key,
lambda r: r.metadata is not None and r.metadata.cost_breakdown is not None,
)
assert row is not None, f"{shape}: no spend row landed"
assert row.spend is not None and cost_rows.approx_equal(row.spend, expected.total), (
f"{shape}: spend {row.spend} != expected {expected.total} "
f"(breakdown {row.breakdown.model_dump()})"
)
prompt_tokens, completion_tokens = expected_token_columns(model, case)
assert row.prompt_tokens == prompt_tokens, (
f"{shape}: prompt_tokens {row.prompt_tokens} != {prompt_tokens}"
)
assert row.completion_tokens == completion_tokens, (
f"{shape}: completion_tokens {row.completion_tokens} != {completion_tokens}"
)
cost_rows.assert_total_is_sum_of_components(row)