test(e2e): apply review nits to cost calculation suite

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-15 23:28:44 +00:00
parent 61ac4f5739
commit 269afbe382
6 changed files with 620 additions and 457 deletions

View file

@ -13,7 +13,7 @@ from __future__ import annotations
import importlib.util
import sys
from collections.abc import Callable
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from pathlib import Path
from types import ModuleType
@ -34,16 +34,16 @@ def _load_cost_rows() -> ModuleType:
"""Load quota_management/spend_tracking/cost_rows.py by path (the e2e tree
has no package layout), the same trick the mcp suite uses for
logging/datadog_reader.py."""
path = (
path: Final = (
Path(__file__).resolve().parent.parent
/ "quota_management"
/ "spend_tracking"
/ "cost_rows.py"
)
name = "e2e_spend_tracking_cost_rows"
spec = importlib.util.spec_from_file_location(name, path)
name: Final = "e2e_spend_tracking_cost_rows"
spec: Final = importlib.util.spec_from_file_location(name, path)
assert spec is not None and spec.loader is not None
module = importlib.util.module_from_spec(spec)
module: Final = importlib.util.module_from_spec(spec)
sys.modules[name] = module
spec.loader.exec_module(module)
return module
@ -59,7 +59,7 @@ class SpendCostBreakdown(Protocol):
total_cost: float | None
service_tier: str | None
def model_dump(self) -> dict[str, object]: ...
def model_dump(self) -> Mapping[str, object]: ...
class SpendRowMetadata(Protocol):
@ -89,7 +89,9 @@ class CostRowsModule(Protocol):
]
cost_rows: Final[CostRowsModule] = cast(CostRowsModule, _load_cost_rows())
cost_rows: Final[CostRowsModule] = cast( # cast-ok: cost_rows.py is loaded by path, so basedpyright has no importable name for it; its surface is declared in CostRowsModule
CostRowsModule, _load_cost_rows()
)
@dataclass(frozen=True, slots=True)
@ -101,7 +103,7 @@ class CostCalcClient:
@pytest.fixture(scope="session")
def client() -> CostCalcClient:
proxy = build_proxy_client(
proxy: Final = build_proxy_client(
base_url=COST_MAP_PROXY_URL,
control_plane_base_url=COST_MAP_PROXY_URL,
replica_urls=(COST_MAP_PROXY_URL,),
@ -118,18 +120,18 @@ def register_scenario_deployment(
) -> tuple[str, ScenarioHandle]:
"""Register the case's scenario on the sidecar plus a deployment pointed at
it; both are torn down by ``resources``. Returns the callable model_name."""
scenario: Scenario = case.scenario(
scenario: Final[Scenario] = case.scenario(
scenario_id=f"sc-{marker}", model=model, text=f"scripted answer {marker}"
)
handle = register_scenario(scenario)
handle: Final = register_scenario(scenario)
resources.defer(lambda: delete_scenario(handle))
model_name = f"{model.model_name}-{marker}"
model_id = client.proxy.register_model(
model_name: Final = f"{model.model_name}-{marker}"
model_id: Final = client.proxy.register_model(
ModelNewBody(
model_name=model_name,
litellm_params=LiteLLMParamsBody(
model=model.litellm_model,
api_key="sk-scripted-provider",
api_key=model.api_key,
api_base=handle.api_base(),
),
model_info=ModelInfoBody(),

View file

@ -18,9 +18,11 @@ creation), the case is absent from the matrix rather than silently zero.
from __future__ import annotations
import json
from collections.abc import Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import Final, Literal
from types import MappingProxyType
from typing import Final, Literal, TypeAlias
from pydantic import BaseModel, ConfigDict, TypeAdapter
@ -64,8 +66,8 @@ class CostMapEntry(BaseModel):
_COST_MAP_ADAPTER: Final = TypeAdapter(dict[str, CostMapEntry])
_COST_MAP: Final[dict[str, CostMapEntry]] = _COST_MAP_ADAPTER.validate_python(
json.loads(COST_MAP_PATH.read_text())
_COST_MAP: Final[Mapping[str, CostMapEntry]] = MappingProxyType(
_COST_MAP_ADAPTER.validate_python(json.loads(COST_MAP_PATH.read_text()))
)
TIER_THRESHOLD_TOKENS: Final = 200_000
@ -109,7 +111,7 @@ class FrontierModel:
# Response-model override targets: emit a sibling's bare provider-facing name so
# the biller's provider-prefixed lookup lands on that sibling's map key.
_OVERRIDE_MODELS: Final[dict[str, str]] = {
_OVERRIDE_MODELS: Final[Mapping[str, str]] = MappingProxyType({
"gpt-5.6": "gpt-5.4-mini",
"gpt-5.5-pro": "gpt-5.3-codex",
"gpt-5.3-codex": "gpt-5.5-pro",
@ -124,9 +126,9 @@ _OVERRIDE_MODELS: Final[dict[str, str]] = {
"fireworks_ai/kimi-k3": "qwen3p8-max",
"fireworks_ai/qwen3p8-max": "kimi-k3",
"fireworks_ai/deepseek-v4p1-flash": "kimi-k3",
}
})
_OVERRIDE_MAP_KEYS: Final[dict[str, str]] = {
_OVERRIDE_MAP_KEYS: Final[Mapping[str, str]] = MappingProxyType({
"gpt-5.4-mini": "gpt-5.4-mini",
"gpt-5.6": "gpt-5.6",
"gpt-5.3-codex": "gpt-5.3-codex",
@ -139,7 +141,7 @@ _OVERRIDE_MAP_KEYS: Final[dict[str, str]] = {
"moonshotai/Kimi-K3": "together_ai/moonshotai/Kimi-K3",
"qwen3p8-max": "fireworks_ai/qwen3p8-max",
"kimi-k3": "fireworks_ai/kimi-k3",
}
})
_FRONTIER_SPECS: Final[tuple[tuple[str, str, Wire], ...]] = (
@ -176,7 +178,7 @@ def _frontier() -> tuple[FrontierModel, ...]:
FRONTIER_MODELS: Final[tuple[FrontierModel, ...]] = _frontier()
# Token kinds each wire can report, gating which pricing cases apply.
_WIRE_CAPS: Final[dict[str, frozenset[str]]] = {
_WIRE_CAPS: Final[Mapping[str, frozenset[str]]] = MappingProxyType({
"openai_chat": frozenset(
{
"cache_read", "cache_write_5m", "cache_write_1h", "reasoning", "audio",
@ -205,9 +207,9 @@ _WIRE_CAPS: Final[dict[str, frozenset[str]]] = {
"web_search", "response_model", "absent_usage",
}
),
}
})
CaseName = Literal[
CaseName: TypeAlias = Literal[
"basic",
"cache_read",
"cache_write_5m",
@ -260,7 +262,7 @@ _BASIC_USAGE: Final = ScriptedUsage(fresh_input_tokens=120, output_tokens=40)
def _web_search_case(model: FrontierModel) -> Case:
counts_exactly = model.wire in ("openai_responses", "anthropic_messages", "gemini_generate")
counts_exactly: Final = model.wire in ("openai_responses", "anthropic_messages", "gemini_generate")
return Case(
name="web_search",
usage=ScriptedUsage(fresh_input_tokens=100, output_tokens=30, web_search_calls=3),
@ -269,26 +271,24 @@ def _web_search_case(model: FrontierModel) -> Case:
def cases_for(model: FrontierModel) -> tuple[Case, ...]:
rates = model.rates
caps = _WIRE_CAPS[model.wire]
cases: list[Case] = [Case(name="basic", usage=_BASIC_USAGE)]
if rates.cache_read_input_token_cost is not None and "cache_read" in caps:
cases.append(
rates: Final = model.rates
caps: Final = _WIRE_CAPS[model.wire]
candidates: Final[tuple[Case | None, ...]] = (
Case(name="basic", usage=_BASIC_USAGE),
(
Case(name="cache_read", usage=ScriptedUsage(fresh_input_tokens=100, cache_read_tokens=50, output_tokens=30))
)
if rates.cache_creation_input_token_cost is not None and "cache_write_5m" in caps:
cases.append(
if rates.cache_read_input_token_cost is not None and "cache_read" in caps
else None
),
(
Case(
name="cache_write_5m",
usage=ScriptedUsage(fresh_input_tokens=90, cache_write_5m_tokens=60, output_tokens=30),
)
)
if (
rates.cache_creation_input_token_cost_above_1hr is not None
and rates.cache_creation_input_token_cost is not None
and "cache_write_1h" in caps
):
cases.append(
if rates.cache_creation_input_token_cost is not None and "cache_write_5m" in caps
else None
),
(
Case(
name="cache_write_1h",
usage=ScriptedUsage(
@ -298,52 +298,61 @@ def cases_for(model: FrontierModel) -> tuple[Case, ...]:
output_tokens=30,
),
)
)
if rates.output_cost_per_reasoning_token is not None and "reasoning" in caps:
cases.append(
if (
rates.cache_creation_input_token_cost_above_1hr is not None
and rates.cache_creation_input_token_cost is not None
and "cache_write_1h" in caps
)
else None
),
(
Case(
name="reasoning",
usage=ScriptedUsage(fresh_input_tokens=100, output_tokens=30, reasoning_tokens=70),
)
)
if (
rates.input_cost_per_audio_token is not None
and rates.output_cost_per_audio_token is not None
and "audio" in caps
):
cases.append(
if rates.output_cost_per_reasoning_token is not None and "reasoning" in caps
else None
),
(
Case(
name="audio",
usage=ScriptedUsage(
fresh_input_tokens=100, audio_input_tokens=25, output_tokens=30, audio_output_tokens=15
),
)
)
if (
rates.input_cost_per_token_above_200k_tokens is not None
and rates.output_cost_per_token_above_200k_tokens is not None
):
cases.append(
if (
rates.input_cost_per_audio_token is not None
and rates.output_cost_per_audio_token is not None
and "audio" in caps
)
else None
),
(
Case(
name="tiered",
usage=ScriptedUsage(
fresh_input_tokens=TIER_THRESHOLD_TOKENS + 1, output_tokens=30
),
)
)
if rates.input_cost_per_token_flex is not None and rates.output_cost_per_token_flex is not None:
cases.append(
if (
rates.input_cost_per_token_above_200k_tokens is not None
and rates.output_cost_per_token_above_200k_tokens is not None
)
else None
),
(
Case(name="service_tier_flex", usage=_BASIC_USAGE, service_tier="flex")
)
if rates.input_cost_per_token_priority is not None and rates.output_cost_per_token_priority is not None:
cases.append(
if rates.input_cost_per_token_flex is not None and rates.output_cost_per_token_flex is not None
else None
),
(
Case(name="service_tier_priority", usage=_BASIC_USAGE, service_tier="priority")
)
if rates.search_context_cost_per_query is not None and "web_search" in caps:
cases.append(_web_search_case(model))
cases.append(Case(name="stream", usage=_BASIC_USAGE, stream=True))
if "absent_usage" in caps:
cases.append(
if rates.input_cost_per_token_priority is not None and rates.output_cost_per_token_priority is not None
else None
),
_web_search_case(model) if rates.search_context_cost_per_query is not None and "web_search" in caps else None,
Case(name="stream", usage=_BASIC_USAGE, stream=True),
(
Case(
name="stream_no_usage",
usage=_BASIC_USAGE,
@ -355,10 +364,16 @@ def cases_for(model: FrontierModel) -> tuple[Case, ...]:
# wires recount tokens proxy-side and bill a nonzero amount.
expect_zero_bill=model.wire == "openai_responses",
)
)
if "response_model" in caps:
cases.append(Case(name="response_model_override", usage=_BASIC_USAGE, response_model_override=True))
return tuple(cases)
if "absent_usage" in caps
else None
),
(
Case(name="response_model_override", usage=_BASIC_USAGE, response_model_override=True)
if "response_model" in caps
else None
),
)
return tuple(case for case in candidates if case is not None)
@dataclass(frozen=True, slots=True)
@ -387,38 +402,41 @@ def expected_breakdown(model: FrontierModel, case: Case) -> ExpectedCost:
to the tier's variants, falling back to the base rate when a variant is
unset -- mirroring _get_token_base_cost in litellm's cost calculator.
"""
rates = model.override_rates if case.response_model_override else model.rates
u = case.usage
prompt_tokens = (
rates: Final = model.override_rates if case.response_model_override else model.rates
u: Final = case.usage
prompt_tokens: Final = (
u.fresh_input_tokens + u.cache_read_tokens + u.cache_write_5m_tokens
+ u.cache_write_1h_tokens + u.audio_input_tokens
)
tiered = prompt_tokens > TIER_THRESHOLD_TOKENS
in_rate = rates.input_cost_per_token or 0.0
out_rate = rates.output_cost_per_token or 0.0
if case.service_tier == "flex":
in_rate = rates.input_cost_per_token_flex or in_rate
out_rate = rates.output_cost_per_token_flex or out_rate
if case.service_tier == "priority":
in_rate = rates.input_cost_per_token_priority or in_rate
out_rate = rates.output_cost_per_token_priority or out_rate
if tiered:
in_rate = rates.input_cost_per_token_above_200k_tokens or in_rate
out_rate = rates.output_cost_per_token_above_200k_tokens or out_rate
input_cost = (
tiered: Final = prompt_tokens > TIER_THRESHOLD_TOKENS
in_rate: Final = (
(rates.input_cost_per_token_above_200k_tokens if tiered else None)
or (rates.input_cost_per_token_priority if case.service_tier == "priority" else None)
or (rates.input_cost_per_token_flex if case.service_tier == "flex" else None)
or rates.input_cost_per_token
or 0.0
)
out_rate: Final = (
(rates.output_cost_per_token_above_200k_tokens if tiered else None)
or (rates.output_cost_per_token_priority if case.service_tier == "priority" else None)
or (rates.output_cost_per_token_flex if case.service_tier == "flex" else None)
or rates.output_cost_per_token
or 0.0
)
input_cost: Final = (
u.fresh_input_tokens * in_rate
+ u.cache_read_tokens * (rates.cache_read_input_token_cost or 0.0)
+ u.cache_write_5m_tokens * (rates.cache_creation_input_token_cost or 0.0)
+ u.cache_write_1h_tokens * (rates.cache_creation_input_token_cost_above_1hr or 0.0)
+ u.audio_input_tokens * (rates.input_cost_per_audio_token or 0.0)
)
output_cost = (
output_cost: Final = (
u.output_tokens * out_rate
+ u.reasoning_tokens * (rates.output_cost_per_reasoning_token or out_rate)
+ u.audio_output_tokens * (rates.output_cost_per_audio_token or out_rate)
)
search = rates.search_context_cost_per_query
tool_cost = case.billed_web_search_calls * (
search: Final = rates.search_context_cost_per_query
tool_cost: Final = case.billed_web_search_calls * (
search.search_context_size_medium if search and search.search_context_size_medium else 0.0
)
return ExpectedCost(input_cost=input_cost, output_cost=output_cost, tool_cost=tool_cost)
@ -432,7 +450,7 @@ 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,
everyone else reports the totals the wire emitted."""
u = case.usage
u: Final = case.usage
if model.wire == "anthropic_messages":
return (
u.fresh_input_tokens + u.cache_read_tokens + u.cache_write_5m_tokens + u.cache_write_1h_tokens,

View file

@ -12,6 +12,7 @@ from e2e_config import SCRIPTED_PROVIDER_CONTROL_URL, SCRIPTED_PROVIDER_PROXY_BA
from e2e_http import URL, NoBody, unwrap, post
from e2e_http import delete as http_delete
from scripted_provider import (
WIRE_MOUNTS,
Scenario,
ScenarioDeleted,
ScenarioRegistered,
@ -29,19 +30,12 @@ class ScenarioHandle:
return f"{self.proxy_base}/{self.scenario_id}/{self._mount()}"
def _mount(self) -> str:
return {
"openai_chat": "openai",
"openai_responses": "openai",
"anthropic_messages": "anthropic",
"gemini_generate": "gemini",
"together_chat": "together",
"fireworks_chat": "fireworks",
}[self.wire]
return WIRE_MOUNTS[self.wire]
def register_scenario(scenario: Scenario) -> ScenarioHandle:
"""POST the scenario to the sidecar's control API and return its handle."""
result = unwrap(
result: Final = unwrap(
post(
URL(f"{SCRIPTED_PROVIDER_CONTROL_URL}/_scenarios"),
headers=NoBody(),

View file

@ -31,14 +31,16 @@ import json
import sys
import threading
import time
from collections.abc import Mapping
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Final, Literal
from types import MappingProxyType
from typing import Final, Literal, TypeAlias
from urllib.parse import urlsplit
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
Wire = Literal[
Wire: TypeAlias = Literal[
"openai_chat",
"openai_responses",
"anthropic_messages",
@ -47,17 +49,19 @@ Wire = Literal[
"fireworks_chat",
]
_WIRE_MOUNTS: Final[dict[str, str]] = {
"openai_chat": "openai",
"openai_responses": "openai",
"anthropic_messages": "anthropic",
"gemini_generate": "gemini",
"together_chat": "together",
"fireworks_chat": "fireworks",
}
WIRE_MOUNTS: Final[Mapping[str, str]] = MappingProxyType(
{
"openai_chat": "openai",
"openai_responses": "openai",
"anthropic_messages": "anthropic",
"gemini_generate": "gemini",
"together_chat": "together",
"fireworks_chat": "fireworks",
}
)
StreamUsage = Literal["final_chunk", "absent"]
ServiceTier = Literal["flex", "priority"]
StreamUsage: TypeAlias = Literal["final_chunk", "absent"]
ServiceTier: TypeAlias = Literal["flex", "priority"]
class ScriptedUsage(BaseModel):
@ -106,7 +110,7 @@ class Scenario(BaseModel):
@property
def mount(self) -> str:
return _WIRE_MOUNTS[self.wire]
return WIRE_MOUNTS[self.wire]
class ScenarioRegistered(BaseModel):
@ -128,369 +132,494 @@ class RenderedResponse:
body: bytes
def _json_bytes(payload: dict[str, object]) -> bytes:
return json.dumps(payload).encode("utf-8")
def _jobj(*pairs: tuple[str, object]) -> Mapping[str, object]:
"""A JSON object payload built in one shot and frozen."""
return MappingProxyType(dict(pairs))
def _sse(events: tuple[tuple[str | None, dict[str, object] | str], ...]) -> bytes:
frames: list[str] = []
for event_name, data in events:
head = f"event: {event_name}\n" if event_name is not None else ""
payload = data if isinstance(data, str) else json.dumps(data)
frames.append(f"{head}data: {payload}\n\n")
return "".join(frames).encode("utf-8")
def _jobj_opt(*pairs: tuple[str, object] | None) -> Mapping[str, object]:
"""``_jobj`` where a ``None`` pair means the field is absent."""
return MappingProxyType(dict(pair for pair in pairs if pair is not None))
def _json_bytes(payload: Mapping[str, object]) -> bytes:
return json.dumps(payload, default=dict).encode("utf-8")
def _sse_frame(event_name: str | None, data: Mapping[str, object] | str) -> str:
head: Final = f"event: {event_name}\n" if event_name is not None else ""
payload: Final = data if isinstance(data, str) else json.dumps(data, default=dict)
return f"{head}data: {payload}\n\n"
def _sse(events: tuple[tuple[str | None, Mapping[str, object] | str], ...]) -> bytes:
return "".join(_sse_frame(event_name, data) for event_name, data in events).encode("utf-8")
# ---------- per-wire usage shapes ----------
def _openai_usage(u: ScriptedUsage) -> dict[str, object]:
prompt_tokens = (
def _openai_usage(u: ScriptedUsage) -> Mapping[str, object]:
prompt_tokens: Final = (
u.fresh_input_tokens
+ u.cache_read_tokens
+ u.cache_write_5m_tokens
+ u.cache_write_1h_tokens
+ u.audio_input_tokens
)
completion_tokens = u.output_tokens + u.reasoning_tokens + u.audio_output_tokens
prompt_details: dict[str, object] = {}
if u.cache_read_tokens:
prompt_details["cached_tokens"] = u.cache_read_tokens
if u.cache_write_5m_tokens or u.cache_write_1h_tokens:
prompt_details["cache_write_tokens"] = u.cache_write_5m_tokens + u.cache_write_1h_tokens
prompt_details["cache_creation_token_details"] = {
"ephemeral_5m_input_tokens": u.cache_write_5m_tokens,
"ephemeral_1h_input_tokens": u.cache_write_1h_tokens,
}
if u.audio_input_tokens:
prompt_details["audio_tokens"] = u.audio_input_tokens
completion_details: dict[str, object] = {}
if u.reasoning_tokens:
completion_details["reasoning_tokens"] = u.reasoning_tokens
if u.audio_output_tokens:
completion_details["audio_tokens"] = u.audio_output_tokens
usage: dict[str, object] = {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
}
if prompt_details:
usage["prompt_tokens_details"] = prompt_details
if completion_details:
usage["completion_tokens_details"] = completion_details
return usage
completion_tokens: Final = u.output_tokens + u.reasoning_tokens + u.audio_output_tokens
prompt_details: Final = _jobj_opt(
("cached_tokens", u.cache_read_tokens) if u.cache_read_tokens else None,
(
("cache_write_tokens", u.cache_write_5m_tokens + u.cache_write_1h_tokens)
if u.cache_write_5m_tokens or u.cache_write_1h_tokens
else None
),
(
(
"cache_creation_token_details",
_jobj(
("ephemeral_5m_input_tokens", u.cache_write_5m_tokens),
("ephemeral_1h_input_tokens", u.cache_write_1h_tokens),
),
)
if u.cache_write_5m_tokens or u.cache_write_1h_tokens
else None
),
("audio_tokens", u.audio_input_tokens) if u.audio_input_tokens else None,
)
completion_details: Final = _jobj_opt(
("reasoning_tokens", u.reasoning_tokens) if u.reasoning_tokens else None,
("audio_tokens", u.audio_output_tokens) if u.audio_output_tokens else None,
)
return _jobj_opt(
("prompt_tokens", prompt_tokens),
("completion_tokens", completion_tokens),
("total_tokens", prompt_tokens + completion_tokens),
("prompt_tokens_details", prompt_details) if prompt_details else None,
("completion_tokens_details", completion_details) if completion_details else None,
)
def _anthropic_usage(u: ScriptedUsage) -> dict[str, object]:
def _anthropic_usage(u: ScriptedUsage) -> Mapping[str, object]:
# Anthropic reports uncached-only input_tokens; cache reads and writes ride
# top-level fields, with the 5m/1h write split under cache_creation.
usage: dict[str, object] = {
"input_tokens": u.fresh_input_tokens,
"output_tokens": u.output_tokens,
}
if u.cache_read_tokens:
usage["cache_read_input_tokens"] = u.cache_read_tokens
if u.cache_write_5m_tokens or u.cache_write_1h_tokens:
usage["cache_creation_input_tokens"] = u.cache_write_5m_tokens + u.cache_write_1h_tokens
usage["cache_creation"] = {
"ephemeral_5m_input_tokens": u.cache_write_5m_tokens,
"ephemeral_1h_input_tokens": u.cache_write_1h_tokens,
}
if u.web_search_calls:
usage["server_tool_use"] = {"web_search_requests": u.web_search_calls}
return usage
return _jobj_opt(
("input_tokens", u.fresh_input_tokens),
("output_tokens", u.output_tokens),
("cache_read_input_tokens", u.cache_read_tokens) if u.cache_read_tokens else None,
(
("cache_creation_input_tokens", u.cache_write_5m_tokens + u.cache_write_1h_tokens)
if u.cache_write_5m_tokens or u.cache_write_1h_tokens
else None
),
(
(
"cache_creation",
_jobj(
("ephemeral_5m_input_tokens", u.cache_write_5m_tokens),
("ephemeral_1h_input_tokens", u.cache_write_1h_tokens),
),
)
if u.cache_write_5m_tokens or u.cache_write_1h_tokens
else None
),
(
("server_tool_use", _jobj(("web_search_requests", u.web_search_calls)))
if u.web_search_calls
else None
),
)
def _gemini_usage(u: ScriptedUsage) -> dict[str, object]:
def _gemini_usage(u: ScriptedUsage) -> Mapping[str, object]:
# promptTokenCount carries the cached count inside it; TEXT modality is the
# cached-inclusive text count so litellm's implicit-caching subtraction lands
# on the fresh figure. candidatesTokenCount includes reasoning + audio.
prompt_tokens = u.fresh_input_tokens + u.cache_read_tokens + u.audio_input_tokens
candidates = u.output_tokens + u.reasoning_tokens + u.audio_output_tokens
usage: dict[str, object] = {
"promptTokenCount": prompt_tokens,
"candidatesTokenCount": candidates,
"totalTokenCount": prompt_tokens + candidates,
}
if u.cache_read_tokens:
usage["cachedContentTokenCount"] = u.cache_read_tokens
if u.reasoning_tokens:
usage["thoughtsTokenCount"] = u.reasoning_tokens
prompt_details = [{"modality": "TEXT", "tokenCount": u.fresh_input_tokens + u.cache_read_tokens}]
if u.audio_input_tokens:
prompt_details.append({"modality": "AUDIO", "tokenCount": u.audio_input_tokens})
usage["promptTokensDetails"] = prompt_details
if u.audio_output_tokens:
usage["candidatesTokensDetails"] = [
{"modality": "TEXT", "tokenCount": u.output_tokens + u.reasoning_tokens},
{"modality": "AUDIO", "tokenCount": u.audio_output_tokens},
]
return usage
prompt_tokens: Final = u.fresh_input_tokens + u.cache_read_tokens + u.audio_input_tokens
candidates: Final = u.output_tokens + u.reasoning_tokens + u.audio_output_tokens
return _jobj_opt(
("promptTokenCount", prompt_tokens),
("candidatesTokenCount", candidates),
("totalTokenCount", prompt_tokens + candidates),
("cachedContentTokenCount", u.cache_read_tokens) if u.cache_read_tokens else None,
("thoughtsTokenCount", u.reasoning_tokens) if u.reasoning_tokens else None,
(
"promptTokensDetails",
(
_jobj(("modality", "TEXT"), ("tokenCount", u.fresh_input_tokens + u.cache_read_tokens)),
*(
(_jobj(("modality", "AUDIO"), ("tokenCount", u.audio_input_tokens)),)
if u.audio_input_tokens
else ()
),
),
),
(
(
"candidatesTokensDetails",
(
_jobj(("modality", "TEXT"), ("tokenCount", u.output_tokens + u.reasoning_tokens)),
_jobj(("modality", "AUDIO"), ("tokenCount", u.audio_output_tokens)),
),
)
if u.audio_output_tokens
else None
),
)
def _responses_usage(u: ScriptedUsage) -> dict[str, object]:
input_tokens = u.fresh_input_tokens + u.cache_read_tokens + u.audio_input_tokens
output_tokens = u.output_tokens + u.reasoning_tokens + u.audio_output_tokens
usage: dict[str, object] = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
}
input_details: dict[str, object] = {}
if u.cache_read_tokens:
input_details["cached_tokens"] = u.cache_read_tokens
if input_details:
usage["input_tokens_details"] = input_details
if u.reasoning_tokens:
usage["output_tokens_details"] = {"reasoning_tokens": u.reasoning_tokens}
return usage
def _responses_usage(u: ScriptedUsage) -> Mapping[str, object]:
input_tokens: Final = u.fresh_input_tokens + u.cache_read_tokens + u.audio_input_tokens
output_tokens: Final = u.output_tokens + u.reasoning_tokens + u.audio_output_tokens
input_details: Final = _jobj_opt(
("cached_tokens", u.cache_read_tokens) if u.cache_read_tokens else None,
)
return _jobj_opt(
("input_tokens", input_tokens),
("output_tokens", output_tokens),
("total_tokens", input_tokens + output_tokens),
("input_tokens_details", input_details) if input_details else None,
(
("output_tokens_details", _jobj(("reasoning_tokens", u.reasoning_tokens)))
if u.reasoning_tokens
else None
),
)
# ---------- per-wire responses ----------
def _openai_message(scenario: Scenario) -> dict[str, object]:
message: dict[str, object] = {"role": "assistant", "content": scenario.output.text}
if scenario.usage.web_search_calls:
message["annotations"] = [
{
"type": "url_citation",
"url_citation": {
"url": "https://scripted.example/source",
"title": "scripted source",
"start_index": 0,
"end_index": 1,
},
}
for _ in range(scenario.usage.web_search_calls)
]
return message
def _openai_message(scenario: Scenario) -> Mapping[str, object]:
return _jobj_opt(
("role", "assistant"),
("content", scenario.output.text),
(
(
"annotations",
tuple(
_jobj(
("type", "url_citation"),
(
"url_citation",
_jobj(
("url", "https://scripted.example/source"),
("title", "scripted source"),
("start_index", 0),
("end_index", 1),
),
),
)
for _ in range(scenario.usage.web_search_calls)
),
)
if scenario.usage.web_search_calls
else None
),
)
def _openai_chat_body(scenario: Scenario, requested_model: str) -> dict[str, object]:
body: dict[str, object] = {
"id": f"chatcmpl-{scenario.scenario_id}",
"object": "chat.completion",
"created": int(time.time()),
"model": scenario.output.response_model or requested_model,
"choices": [
{
"index": 0,
"message": _openai_message(scenario),
"finish_reason": scenario.output.finish_reason,
}
],
"usage": _openai_usage(scenario.usage),
}
if scenario.service_tier is not None:
body["service_tier"] = scenario.service_tier
if scenario.output.provider_cost is not None:
body["cost"] = scenario.output.provider_cost
return body
def _openai_chat_body(scenario: Scenario, requested_model: str) -> Mapping[str, object]:
return _jobj_opt(
("id", f"chatcmpl-{scenario.scenario_id}"),
("object", "chat.completion"),
("created", int(time.time())),
("model", scenario.output.response_model or requested_model),
(
"choices",
(
_jobj(
("index", 0),
("message", _openai_message(scenario)),
("finish_reason", scenario.output.finish_reason),
),
),
),
("usage", _openai_usage(scenario.usage)),
("service_tier", scenario.service_tier) if scenario.service_tier is not None else None,
("cost", scenario.output.provider_cost) if scenario.output.provider_cost is not None else None,
)
def _openai_chunk(scenario: Scenario, requested_model: str, **kw: object) -> dict[str, object]:
chunk: dict[str, object] = {
"id": f"chatcmpl-{scenario.scenario_id}",
"object": "chat.completion.chunk",
"created": int(time.time()),
"model": scenario.output.response_model or requested_model,
}
chunk.update(kw)
return chunk
def _openai_chunk(
scenario: Scenario,
requested_model: str,
choices: tuple[Mapping[str, object], ...] = (),
usage: Mapping[str, object] | None = None,
) -> Mapping[str, object]:
return _jobj_opt(
("id", f"chatcmpl-{scenario.scenario_id}"),
("object", "chat.completion.chunk"),
("created", int(time.time())),
("model", scenario.output.response_model or requested_model),
("choices", choices),
("usage", usage),
)
def _openai_chat_sse(scenario: Scenario, requested_model: str) -> bytes:
_EMPTY_DELTA: Final[dict[str, object]] = {}
delta: dict[str, object] = {"role": "assistant", "content": scenario.output.text}
if scenario.usage.web_search_calls:
delta["annotations"] = _openai_message(scenario)["annotations"]
events: list[tuple[str | None, dict[str, object] | str]] = [
delta: Final = _jobj_opt(
("role", "assistant"),
("content", scenario.output.text),
(
None,
_openai_chunk(
scenario,
requested_model,
choices=[{"index": 0, "delta": {"role": "assistant"}, "finish_reason": None}],
),
("annotations", _openai_message(scenario)["annotations"])
if scenario.usage.web_search_calls
else None
),
)
return _sse(
(
None,
_openai_chunk(
scenario,
requested_model,
choices=[{"index": 0, "delta": delta, "finish_reason": None}],
(
None,
_openai_chunk(
scenario,
requested_model,
choices=(_jobj(("index", 0), ("delta", _jobj(("role", "assistant"))), ("finish_reason", None)),),
),
),
),
(
None,
_openai_chunk(
scenario,
requested_model,
choices=[
{
"index": 0,
"delta": _EMPTY_DELTA,
"finish_reason": scenario.output.finish_reason,
}
],
(
None,
_openai_chunk(
scenario,
requested_model,
choices=(_jobj(("index", 0), ("delta", delta), ("finish_reason", None)),),
),
),
),
]
if scenario.stream_usage == "final_chunk":
events.append(
(None, _openai_chunk(scenario, requested_model, choices=(), usage=_openai_usage(scenario.usage)))
(
None,
_openai_chunk(
scenario,
requested_model,
choices=(
_jobj(
("index", 0),
("delta", _jobj()),
("finish_reason", scenario.output.finish_reason),
),
),
),
),
*(
((None, _openai_chunk(scenario, requested_model, usage=_openai_usage(scenario.usage))),)
if scenario.stream_usage == "final_chunk"
else ()
),
(None, "[DONE]"),
)
events.append((None, "[DONE]"))
return _sse(tuple(events))
)
def _anthropic_body(scenario: Scenario, requested_model: str) -> dict[str, object]:
return {
"id": f"msg_{scenario.scenario_id}",
"type": "message",
"role": "assistant",
"model": scenario.output.response_model or requested_model,
"content": [{"type": "text", "text": scenario.output.text}],
"stop_reason": "end_turn" if scenario.output.finish_reason == "stop" else scenario.output.finish_reason,
"usage": _anthropic_usage(scenario.usage),
}
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,
),
("usage", _anthropic_usage(scenario.usage)),
)
def _anthropic_sse(scenario: Scenario, requested_model: str) -> bytes:
emit_usage = scenario.stream_usage == "final_chunk"
input_usage = {k: v for k, v in _anthropic_usage(scenario.usage).items() if k != "output_tokens"}
message_start: dict[str, object] = {
"type": "message_start",
"message": {
"id": f"msg_{scenario.scenario_id}",
"type": "message",
"role": "assistant",
"model": scenario.output.response_model or requested_model,
"content": [],
"stop_reason": None,
**({"usage": input_usage} if emit_usage else {}),
},
}
message_delta: dict[str, object] = {
"type": "message_delta",
"delta": {
"stop_reason": "end_turn" if scenario.output.finish_reason == "stop" else scenario.output.finish_reason
},
**({"usage": {"output_tokens": scenario.usage.output_tokens}} if emit_usage else {}),
}
emit_usage: Final = scenario.stream_usage == "final_chunk"
input_usage: Final = _jobj(
*(
(key, value)
for key, value in _anthropic_usage(scenario.usage).items()
if key != "output_tokens"
)
)
message_start: Final = _jobj(
("type", "message_start"),
(
"message",
_jobj_opt(
("id", f"msg_{scenario.scenario_id}"),
("type", "message"),
("role", "assistant"),
("model", scenario.output.response_model or requested_model),
("content", ()),
("stop_reason", None),
("usage", input_usage) if emit_usage else None,
),
),
)
message_delta: Final = _jobj_opt(
("type", "message_delta"),
(
"delta",
_jobj(
(
"stop_reason",
"end_turn" if scenario.output.finish_reason == "stop" else scenario.output.finish_reason,
)
),
),
(
("usage", _jobj(("output_tokens", scenario.usage.output_tokens)))
if emit_usage
else None
),
)
return _sse(
(
("message_start", message_start),
(
"content_block_start",
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
},
_jobj(
("type", "content_block_start"),
("index", 0),
("content_block", _jobj(("type", "text"), ("text", ""))),
),
),
(
"content_block_delta",
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": scenario.output.text},
},
_jobj(
("type", "content_block_delta"),
("index", 0),
("delta", _jobj(("type", "text_delta"), ("text", scenario.output.text))),
),
),
("content_block_stop", {"type": "content_block_stop", "index": 0}),
("content_block_stop", _jobj(("type", "content_block_stop"), ("index", 0))),
("message_delta", message_delta),
("message_stop", {"type": "message_stop"}),
("message_stop", _jobj(("type", "message_stop"))),
)
)
def _gemini_body(scenario: Scenario, requested_model: str) -> dict[str, object]:
candidate: dict[str, object] = {
"content": {"parts": [{"text": scenario.output.text}], "role": "model"},
"finishReason": "STOP" if scenario.output.finish_reason == "stop" else scenario.output.finish_reason.upper(),
"index": 0,
}
if scenario.usage.web_search_calls:
candidate["groundingMetadata"] = {
"webSearchQueries": [f"query {i}" for i in range(scenario.usage.web_search_calls)]
}
return {
"candidates": [candidate],
"usageMetadata": _gemini_usage(scenario.usage),
"modelVersion": scenario.output.response_model or requested_model,
}
def _gemini_body(scenario: Scenario, requested_model: str) -> Mapping[str, object]:
return _jobj(
(
"candidates",
(
_jobj_opt(
(
"content",
_jobj(
("parts", (_jobj(("text", scenario.output.text)),)),
("role", "model"),
),
),
(
"finishReason",
"STOP" if scenario.output.finish_reason == "stop" else scenario.output.finish_reason.upper(),
),
("index", 0),
(
(
"groundingMetadata",
_jobj(
(
"webSearchQueries",
tuple(f"query {i}" for i in range(scenario.usage.web_search_calls)),
)
),
)
if scenario.usage.web_search_calls
else None
),
),
),
),
("usageMetadata", _gemini_usage(scenario.usage)),
("modelVersion", scenario.output.response_model or requested_model),
)
def _gemini_sse(scenario: Scenario, requested_model: str) -> bytes:
first = _gemini_body(scenario, requested_model)
if scenario.stream_usage == "absent":
first = {k: v for k, v in first.items() if k != "usageMetadata"}
events: list[tuple[str | None, dict[str, object] | str]] = [(None, first)]
if scenario.stream_usage == "final_chunk":
events.append(
(
None,
{
"candidates": [],
"usageMetadata": _gemini_usage(scenario.usage),
"modelVersion": scenario.output.response_model or requested_model,
},
)
)
return _sse(tuple(events))
def _responses_body(scenario: Scenario, requested_model: str) -> dict[str, object]:
output: list[dict[str, object]] = [
{"type": "web_search_call", "id": f"ws_{i}", "status": "completed"}
for i in range(scenario.usage.web_search_calls)
]
output.append(
{
"type": "message",
"id": f"msg_{scenario.scenario_id}",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": scenario.output.text,
"annotations": [],
}
],
}
emit_usage: Final = scenario.stream_usage == "final_chunk"
first: Final = (
_jobj(*((key, value) for key, value in _gemini_body(scenario, requested_model).items() if key != "usageMetadata"))
if scenario.stream_usage == "absent"
else _gemini_body(scenario, requested_model)
)
return _sse(
(
(None, first),
*(
(
(
None,
_jobj(
("candidates", ()),
("usageMetadata", _gemini_usage(scenario.usage)),
("modelVersion", scenario.output.response_model or requested_model),
),
),
)
if emit_usage
else ()
),
)
)
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",
(
*(
_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", ()),
),
),
),
),
),
),
("usage", _responses_usage(scenario.usage)),
)
return {
"id": f"resp_{scenario.scenario_id}",
"object": "response",
"created_at": int(time.time()),
"status": "completed",
"model": scenario.output.response_model or requested_model,
"output": output,
"usage": _responses_usage(scenario.usage),
}
def _responses_sse(scenario: Scenario, requested_model: str) -> bytes:
completed = _responses_body(scenario, requested_model)
if scenario.stream_usage == "absent":
completed = {k: v for k, v in completed.items() if k != "usage"}
created = {**completed, "status": "in_progress", "usage": None}
completed: 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")),
("status", "in_progress"),
("usage", None),
)
return _sse(
(
("response.created", {"type": "response.created", "response": created}),
("response.created", _jobj(("type", "response.created"), ("response", created))),
(
"response.output_text.delta",
{
"type": "response.output_text.delta",
"item_id": f"msg_{scenario.scenario_id}",
"output_index": scenario.usage.web_search_calls,
"content_index": 0,
"delta": scenario.output.text,
},
_jobj(
("type", "response.output_text.delta"),
("item_id", f"msg_{scenario.scenario_id}"),
("output_index", scenario.usage.web_search_calls),
("content_index", 0),
("delta", scenario.output.text),
),
),
("response.completed", {"type": "response.completed", "response": completed}),
("response.completed", _jobj(("type", "response.completed"), ("response", completed))),
)
)
@ -538,11 +667,11 @@ class _ScenarioStore:
_REQUEST_BODY: Final = TypeAdapter(dict[str, object])
def _request_body(body: bytes) -> dict[str, object]:
def _request_body(body: bytes) -> Mapping[str, object]:
try:
return _REQUEST_BODY.validate_json(body)
except ValueError:
return {}
return MappingProxyType({})
def _request_wants_stream(path_tail: str, body: bytes) -> bool:
@ -554,52 +683,66 @@ def _request_wants_stream(path_tail: str, body: bytes) -> bool:
def _request_model(body: bytes) -> str:
model = _request_body(body).get("model")
model: Final = _request_body(body).get("model")
return model if isinstance(model, str) else "unknown"
def handle_request(store: _ScenarioStore, method: str, raw_path: str, body: bytes) -> RenderedResponse:
path = urlsplit(raw_path).path
segments = [segment for segment in path.split("/") if segment]
if method == "GET" and segments == ["health"]:
return RenderedResponse(200, "application/json", _json_bytes({"status": "ok"}))
path: Final = urlsplit(raw_path).path
segments: Final = tuple(segment for segment in path.split("/") if segment)
if method == "GET" and segments == ("health",):
return RenderedResponse(200, "application/json", _json_bytes(_jobj(("status", "ok"))))
if segments and segments[0] == "_scenarios":
if method == "POST" and len(segments) == 1:
try:
scenario = Scenario.model_validate_json(body)
scenario: Final = Scenario.model_validate_json(body)
except ValidationError as exc:
return RenderedResponse(400, "application/json", _json_bytes({"error": str(exc)}))
return RenderedResponse(
400, "application/json", _json_bytes(_jobj(("error", str(exc))))
)
store.put(scenario)
return RenderedResponse(200, "application/json", _json_bytes({"scenario_id": scenario.scenario_id}))
if method == "DELETE" and len(segments) == 2:
deleted = store.drop(segments[1])
return RenderedResponse(
200 if deleted else 404, "application/json", _json_bytes({"deleted": deleted})
200, "application/json", _json_bytes(_jobj(("scenario_id", scenario.scenario_id)))
)
return RenderedResponse(404, "application/json", _json_bytes({"error": "unknown control route"}))
if method == "DELETE" and len(segments) == 2:
deleted: Final = store.drop(segments[1])
return RenderedResponse(
200 if deleted else 404,
"application/json",
_json_bytes(_jobj(("deleted", deleted))),
)
return RenderedResponse(
404, "application/json", _json_bytes(_jobj(("error", "unknown control route")))
)
if len(segments) < 2 or method != "POST":
return RenderedResponse(404, "application/json", _json_bytes({"error": f"no route for {method} {path}"}))
return RenderedResponse(
404, "application/json", _json_bytes(_jobj(("error", f"no route for {method} {path}")))
)
scenario_id, mount = segments[0], segments[1]
scenario = store.get(scenario_id)
if scenario is None:
return RenderedResponse(404, "application/json", _json_bytes({"error": f"unknown scenario {scenario_id}"}))
if scenario.mount != mount:
found: Final = store.get(scenario_id)
if found is None:
return RenderedResponse(
404, "application/json", _json_bytes(_jobj(("error", f"unknown scenario {scenario_id}")))
)
if found.mount != mount:
return RenderedResponse(
400,
"application/json",
_json_bytes({"error": f"scenario {scenario_id} is wire {scenario.wire}, not mount {mount}"}),
_json_bytes(
_jobj(("error", f"scenario {scenario_id} is wire {found.wire}, not mount {mount}"))
),
)
tail = "/".join(segments[2:])
return _render(scenario, stream=_request_wants_stream(tail, body), requested_model=_request_model(body))
tail: Final = "/".join(segments[2:])
return _render(found, stream=_request_wants_stream(tail, body), requested_model=_request_model(body))
class _ScriptedHandler(BaseHTTPRequestHandler):
store: Final[_ScenarioStore] = _ScenarioStore()
def _dispatch(self, method: str) -> None:
length = int(self.headers.get("content-length") or 0)
body = self.rfile.read(length) if length else b""
rendered = handle_request(self.store, method, self.path, body)
length: Final = int(self.headers.get("content-length") or 0)
body: Final = self.rfile.read(length) if length else b""
rendered: Final = handle_request(self.store, method, self.path, body)
self.send_response(rendered.status_code)
self.send_header("content-type", rendered.content_type)
self.send_header("content-length", str(len(rendered.body)))
@ -621,11 +764,11 @@ DEFAULT_PORT: Final = 9100
def serve(port: int = DEFAULT_PORT, bind_host: str = "127.0.0.1") -> None:
server = ThreadingHTTPServer((bind_host, port), _ScriptedHandler)
server: Final = ThreadingHTTPServer((bind_host, port), _ScriptedHandler)
sys.stderr.write(f"scripted-provider listening on http://{bind_host}:{port}\n")
server.serve_forever()
if __name__ == "__main__":
port_arg = int(sys.argv[1]) if len(sys.argv) > 1 else DEFAULT_PORT
port_arg: Final = int(sys.argv[1]) if len(sys.argv) > 1 else DEFAULT_PORT
serve(port=port_arg)

View file

@ -11,6 +11,7 @@ tests/e2e/cost_map.json.
from __future__ import annotations
import pytest
from typing import Final
from conftest import CostCalcClient, cost_rows, register_scenario_deployment
from cost_matrix import (
@ -25,11 +26,11 @@ from e2e_config import unique_marker
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, ChatStreamOptions
pytestmark = [pytest.mark.e2e, pytest.mark.cost_map_stack]
pytestmark: Final = [pytest.mark.e2e, pytest.mark.cost_map_stack] # mutable-ok: pytest only accepts a list for pytestmark
_MATRIX: list[tuple[FrontierModel, Case]] = [
_MATRIX: Final[tuple[tuple[FrontierModel, Case], ...]] = tuple(
(model, case) for model in FRONTIER_MODELS for case in cases_for(model)
]
)
def _case_id(param: tuple[FrontierModel, Case]) -> str:
@ -40,7 +41,7 @@ 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=f"{marker} scripted pricing call"),),
stream=case.stream,
stream_options=ChatStreamOptions(include_usage=True) if case.stream else None,
service_tier=case.service_tier,
@ -58,9 +59,9 @@ class TestTokenPricing:
model_case: tuple[FrontierModel, Case],
) -> None:
model, case = model_case
marker = unique_marker()
marker: Final = unique_marker()
model_name, _handle = register_scenario_deployment(client, resources, model, case, marker)
response = client.proxy.transport.send(
response: Final = client.proxy.transport.send(
"/chat/completions",
headers=client.proxy.transport.bearer(scoped_key),
json=_chat_body(model_name, marker, case),
@ -71,7 +72,7 @@ class TestTokenPricing:
)
assert response.stream_error is None, f"stream carried an error event: {response.stream_error}"
expected = expected_cost(model, case)
expected: Final = expected_cost(model, case)
if case.exact_spend and not case.stream:
# Streamed responses commit headers before the bill is computed, so
# the x-litellm-response-cost header is asserted only on non-stream
@ -82,7 +83,7 @@ class TestTokenPricing:
f"x-litellm-response-cost {response.response_cost} != expected {expected}"
)
row = cost_rows.poll_cost_row_where(
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,

View file

@ -12,6 +12,9 @@ the proxy to POST /responses) and a streamed Anthropic-messages case.
from __future__ import annotations
import pytest
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from conftest import CostCalcClient, cost_rows, register_scenario_deployment
from cost_matrix import (
@ -26,12 +29,14 @@ from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, ChatStreamOptions
from scripted_provider import ScriptedUsage
pytestmark = [pytest.mark.e2e, pytest.mark.cost_map_stack]
pytestmark: Final = [pytest.mark.e2e, pytest.mark.cost_map_stack] # mutable-ok: pytest only accepts a list for pytestmark
_MODELS: dict[str, FrontierModel] = {model.map_key: model for model in FRONTIER_MODELS}
_MODELS: Final[Mapping[str, FrontierModel]] = MappingProxyType(
{model.map_key: model for model in FRONTIER_MODELS}
)
# One scripted usage per wire, every reportable token kind nonzero.
_WIRE_USAGE: dict[str, tuple[str, ScriptedUsage]] = {
_WIRE_USAGE: Final[Mapping[str, tuple[str, ScriptedUsage]]] = MappingProxyType({
"openai_chat": (
"gpt-5.6",
ScriptedUsage(
@ -89,7 +94,7 @@ _WIRE_USAGE: dict[str, tuple[str, ScriptedUsage]] = {
"fireworks_ai/kimi-k3",
ScriptedUsage(fresh_input_tokens=80, cache_read_tokens=40, output_tokens=25),
),
}
})
class TestWireFormats:
@ -103,22 +108,22 @@ class TestWireFormats:
wire: str,
) -> None:
map_key, usage = _WIRE_USAGE[wire]
model = _MODELS[map_key]
case = Case(name="basic", usage=usage)
marker = unique_marker()
model: Final = _MODELS[map_key]
case: Final = Case(name="basic", usage=usage)
marker: Final = unique_marker()
model_name, _handle = register_scenario_deployment(client, resources, model, case, marker)
response = client.proxy.transport.send(
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 wire call")],
messages=(ChatMessage(role="user", content=f"{marker} scripted wire call"),),
),
)
assert response.ok, f"{wire}: proxy returned {response.status_code}: {response.body[:400]}"
expected = expected_breakdown(model, case)
row = cost_rows.poll_cost_row_where(
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,
@ -128,7 +133,7 @@ class TestWireFormats:
f"{wire}: spend {row.spend} != expected {expected.total} "
f"(breakdown {row.breakdown.model_dump()})"
)
breakdown = row.breakdown
breakdown: Final = row.breakdown
assert breakdown.input_cost is not None and cost_rows.approx_equal(
breakdown.input_cost, expected.input_cost
), (
@ -153,16 +158,16 @@ class TestWireFormats:
self, client: CostCalcClient, resources: ResourceManager, scoped_key: str
) -> None:
map_key, usage = _WIRE_USAGE["anthropic_messages"]
model = _MODELS[map_key]
case = Case(name="stream", usage=usage, stream=True)
marker = unique_marker()
model: Final = _MODELS[map_key]
case: Final = Case(name="stream", usage=usage, stream=True)
marker: Final = unique_marker()
model_name, _handle = register_scenario_deployment(client, resources, model, case, marker)
response = client.proxy.transport.send(
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 anthropic stream")],
messages=(ChatMessage(role="user", content=f"{marker} scripted anthropic stream"),),
stream=True,
stream_options=ChatStreamOptions(include_usage=True),
),
@ -172,8 +177,8 @@ class TestWireFormats:
assert response.stream_done, "anthropic stream did not reach its terminal event"
assert response.stream_error is None, f"stream carried an error event: {response.stream_error}"
expected = expected_breakdown(model, case)
row = cost_rows.poll_cost_row_where(
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,