mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
parent
415b06f5ff
commit
feb69c5f78
4 changed files with 688 additions and 76 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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))),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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), (
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue