fix: harden provider certification and native pricing

This commit is contained in:
9qeklajc
2026-10-03 20:00:26 +02:00
parent df89c0a5b6
commit c4e708bb32
18 changed files with 1285 additions and 261 deletions
+8 -2
View File
@@ -239,7 +239,10 @@ def create_model_mappings(
override_row, provider_fee = overrides_by_key[model_key] override_row, provider_fee = overrides_by_key[model_key]
try: try:
model_to_use = _row_to_model( model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee override_row,
apply_provider_fee=True,
provider_fee=provider_fee,
provider_type=upstream.provider_type,
) )
except Exception as exc: except Exception as exc:
# Stored pricing is JSON from whatever wrote the row, so # Stored pricing is JSON from whatever wrote the row, so
@@ -315,7 +318,10 @@ def create_model_mappings(
try: try:
model_to_use = _row_to_model( model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee override_row,
apply_provider_fee=True,
provider_fee=provider_fee,
provider_type=upstream_for_override.provider_type,
) )
except Exception as exc: except Exception as exc:
logger.warning( logger.warning(
+15 -4
View File
@@ -617,7 +617,10 @@ async def upsert_provider_model(
await refresh_model_maps() await refresh_model_maps()
await _refresh_provider_model_paths(provider_pk) await _refresh_provider_model_paths(provider_pk)
return _row_to_model( return _row_to_model(
row, apply_provider_fee=True, provider_fee=provider.provider_fee row,
apply_provider_fee=True,
provider_fee=provider.provider_fee,
provider_type=provider.provider_type,
).dict() # type: ignore ).dict() # type: ignore
@@ -653,7 +656,10 @@ async def get_provider_model(provider_id: str, model_id: str) -> dict[str, objec
# is not a usable number must be shown as it is, not encoded as `null`. # is not a usable number must be shown as it is, not encoded as `null`.
return json_compliant( # type: ignore[return-value] return json_compliant( # type: ignore[return-value]
_row_to_model( _row_to_model(
row, apply_provider_fee=False, provider_fee=provider.provider_fee row,
apply_provider_fee=False,
provider_fee=provider.provider_fee,
provider_type=provider.provider_type,
).dict() ).dict()
) )
@@ -1311,7 +1317,10 @@ def _evaluate_model_row(
) -> _ModelEvaluation: ) -> _ModelEvaluation:
try: try:
configured: Model | None = _build_model_from_row( configured: Model | None = _build_model_from_row(
row, apply_provider_fee=True, provider_fee=provider.provider_fee row,
apply_provider_fee=True,
provider_fee=provider.provider_fee,
provider_type=provider.provider_type,
) )
build_error = None build_error = None
except Exception as exc: except Exception as exc:
@@ -1352,7 +1361,9 @@ def _aggregate_row(
""" """
evidence: dict[str, object] = {"checked": checked, "flagged": list(flagged)} evidence: dict[str, object] = {"checked": checked, "flagged": list(flagged)}
if checked == 0: if checked == 0:
return _report_row(row_id, "ok", title, empty_detail, evidence) return _report_row(
row_id, "warn", title, f"Not evaluated — {empty_detail}", evidence
)
if flagged: if flagged:
return _report_row(row_id, fail_status, title, flagged_detail, evidence) return _report_row(row_id, fail_status, title, flagged_detail, evidence)
return _report_row(row_id, "ok", title, ok_detail, evidence) return _report_row(row_id, "ok", title, ok_detail, evidence)
+22 -4
View File
@@ -316,8 +316,16 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis
return [] return []
def allows_cache_pricing_backfill(provider_type: str | None) -> bool:
return provider_type not in {"ppqai", "venice"}
def _build_model_from_row( def _build_model_from_row(
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 row: ModelRow,
apply_provider_fee: bool = False,
provider_fee: float = 1.01,
*,
provider_type: str | None = None,
) -> Model: ) -> Model:
"""The deterministic USD view of a stored model row, before the sats conversion.""" """The deterministic USD view of a stored model row, before the sats conversion."""
architecture = json.loads(row.architecture) architecture = json.loads(row.architecture)
@@ -346,7 +354,8 @@ def _build_model_from_row(
# forwarded_model_id="deepseek-v4-flash") would otherwise look up the alias # forwarded_model_id="deepseek-v4-flash") would otherwise look up the alias
# and miss the cache rate. # and miss the cache rate.
pricing_model_id = getattr(row, "forwarded_model_id", None) or row.id pricing_model_id = getattr(row, "forwarded_model_id", None) or row.id
parsed_pricing = backfill_cache_pricing(pricing_model_id, parsed_pricing) if allows_cache_pricing_backfill(provider_type):
parsed_pricing = backfill_cache_pricing(pricing_model_id, parsed_pricing)
if apply_provider_fee: if apply_provider_fee:
parsed_pricing = Pricing.parse_obj( parsed_pricing = Pricing.parse_obj(
@@ -383,9 +392,15 @@ def _build_model_from_row(
def _row_to_model( def _row_to_model(
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 row: ModelRow,
apply_provider_fee: bool = False,
provider_fee: float = 1.01,
*,
provider_type: str | None = None,
) -> Model: ) -> Model:
model = _build_model_from_row(row, apply_provider_fee, provider_fee) model = _build_model_from_row(
row, apply_provider_fee, provider_fee, provider_type=provider_type
)
try: try:
sats_to_usd = sats_usd_price() sats_to_usd = sats_usd_price()
@@ -430,6 +445,9 @@ async def list_models(
provider_fee=providers_by_id[r.upstream_provider_id].provider_fee provider_fee=providers_by_id[r.upstream_provider_id].provider_fee
if r.upstream_provider_id in providers_by_id if r.upstream_provider_id in providers_by_id
else 1.01, else 1.01,
provider_type=providers_by_id[r.upstream_provider_id].provider_type
if r.upstream_provider_id in providers_by_id
else None,
) )
except Exception as e: except Exception as e:
# Stored pricing/architecture is JSON from whatever wrote the row, so # Stored pricing/architecture is JSON from whatever wrote the row, so
+13 -4
View File
@@ -5699,12 +5699,17 @@ class BaseUpstreamProvider:
code=error_code, code=error_code,
) )
@property
def allow_cache_pricing_backfill(self) -> bool:
from ..payment.models import allows_cache_pricing_backfill
return allows_cache_pricing_backfill(self.provider_type)
def _apply_provider_fee_to_model(self, model: Model) -> Model: def _apply_provider_fee_to_model(self, model: Model) -> Model:
"""Apply provider fee to model's USD pricing and calculate max costs. """Apply provider fee to model's USD pricing and calculate max costs.
Cache rates missing from the upstream pricing feed are backfilled from Providers with native price catalogs can disable generic cache-rate
litellm's cost map first, so they carry the provider fee like every backfill to avoid treating another provider's rates as their own.
other price component.
Args: Args:
model: Model object to update model: Model object to update
@@ -5712,7 +5717,11 @@ class BaseUpstreamProvider:
Returns: Returns:
Model with provider fee applied to pricing and max costs calculated Model with provider fee applied to pricing and max costs calculated
""" """
base_pricing = backfill_cache_pricing(model.id, model.pricing) base_pricing = (
backfill_cache_pricing(model.id, model.pricing)
if self.allow_cache_pricing_backfill
else model.pricing
)
adjusted_pricing = Pricing.parse_obj( adjusted_pricing = Pricing.parse_obj(
{k: v * self.provider_fee for k, v in base_pricing.dict().items()} {k: v * self.provider_fee for k, v in base_pricing.dict().items()}
) )
+113 -95
View File
@@ -1,12 +1,12 @@
"""Live certification checks for an upstream provider endpoint. """Live certification checks for an upstream provider endpoint.
Extends the read-only pricing rows, which never touch the network, with the Extends the read-only pricing rows, which never touch the network, with the
ones that must: a ``/models`` heartbeat and a one-token completion. ones that must: a ``/models`` heartbeat and a bounded completion.
Probes call the upstream directly with ``httpx``, never through the node's Probes call the upstream directly with ``httpx``, never through the node's
billing path — no reservation, no Cashu. Upstream spend is one one-token billing path — no reservation, no Cashu. Output budgets start at 32 tokens
completion, plus two or three one-token completions on a ~4.4k-token prompt and rise only on recognized limit rejections, up to 2048. Cache checks add
when the cache checks are enabled (per certified model path). repeated completions on a ~4.4k-token prompt (per certified model path).
They sit behind ``POST …/certify`` rather than the read-only ``GET …/report`` They sit behind ``POST …/certify`` rather than the read-only ``GET …/report``
because they can block for the length of the timeout. because they can block for the length of the timeout.
""" """
@@ -21,7 +21,7 @@ import os
import sys import sys
import time import time
from collections.abc import Callable from collections.abc import Callable
from dataclasses import dataclass from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from urllib.parse import urlparse from urllib.parse import urlparse
@@ -31,6 +31,14 @@ from ..core.logging import get_logger
from ..payment.cost_calculation import _resolve_usd_cost, calculate_cost from ..payment.cost_calculation import _resolve_usd_cost, calculate_cost
from ..payment.rates import coerce_rate from ..payment.rates import coerce_rate
from ..payment.usage import normalize_usage from ..payment.usage import normalize_usage
from .certification_probe import (
PROBE_MAX_TOKENS,
PROBE_TIMEOUT_SECONDS,
send_completion,
)
from .certification_probe import (
wants_max_completion_tokens as wants_max_completion_tokens,
)
from .model_paths import is_openrouter_base_url from .model_paths import is_openrouter_base_url
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -45,12 +53,7 @@ STATUS_FAIL = "fail"
TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"} TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"}
# Bounded so a dead upstream fails the row rather than wedging the request. PROBE_PROMPT = "Reply with the single word: ok"
PROBE_TIMEOUT_SECONDS = 15.0
# The cheapest request that still exercises the usage/cost path.
PROBE_MAX_TOKENS = 1
PROBE_PROMPT = "ping"
# ``calculate_cost`` demands a reservation ceiling; any value at or above the # ``calculate_cost`` demands a reservation ceiling; any value at or above the
# real charge behaves identically. # real charge behaves identically.
@@ -128,12 +131,12 @@ CHECKLIST_GOALS: tuple[tuple[str, str, tuple[str, ...]], ...] = (
), ),
( (
"caching", "caching",
"Prompt caching — cache hits reported and billed at the cache rate", "Prompt caching — hits reported and cached cost calculation verified",
("cache.reported", "cache.billing"), ("cache.reported", "cache.billing"),
), ),
( (
"margin", "margin",
"Margin — node charge covers the upstream's cost", "Pricing target — token estimate covers the fee-adjusted target",
("cost.margin",), ("cost.margin",),
), ),
) )
@@ -166,7 +169,7 @@ def build_checklist(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
@dataclass @dataclass
class ProbeResult: class ProbeResult:
"""Raw outcome of the two live HTTP calls a probe makes.""" """Outcome of model discovery and the bounded completion attempts."""
base_url: str base_url: str
models_url: str models_url: str
@@ -180,19 +183,11 @@ class ProbeResult:
chat_payload: dict[str, Any] | None = None chat_payload: dict[str, Any] | None = None
chat_error: str | None = None chat_error: str | None = None
chat_latency_ms: float | None = None chat_latency_ms: float | None = None
# ``max_completion_tokens`` once the upstream rejected ``max_tokens``.
token_limit_field: str = "max_tokens" token_limit_field: str = "max_tokens"
token_limit: int = PROBE_MAX_TOKENS
chat_attempts: list[dict[str, Any]] = field(default_factory=list)
def wants_max_completion_tokens(status: int | None, payload: Any) -> bool: provider_type: str | None = None
"""Whether a 400 names ``max_completion_tokens`` as the field to use. completion_skip_reason: str | None = None
OpenAI's o-series and gpt-5 reject ``max_tokens`` on chat completions with
"Unsupported parameter: 'max_tokens' ... Use 'max_completion_tokens'".
"""
if status != 400 or payload is None:
return False
return "max_completion_tokens" in json.dumps(payload, default=str)
@dataclass @dataclass
@@ -269,12 +264,11 @@ async def probe_upstream(
upstream: "BaseUpstreamProvider | None" = None, upstream: "BaseUpstreamProvider | None" = None,
model: "Model | None" = None, model: "Model | None" = None,
) -> ProbeResult: ) -> ProbeResult:
"""Call the upstream's ``/models`` and a one-token completion. """Discover models and probe supported chat APIs with bounded correction.
A completion refused with a 400 naming ``max_completion_tokens`` is Each HTTP call, including its body read, has an elapsed-time deadline.
retried once with that field (OpenAI o-series, gpt-5). Each HTTP call, Transport failures are rows, not failed admin requests. Native System One
including its body read, has an elapsed-time deadline. A transport inference is deliberately unverified until a supported fixture exists.
failure is a ``fail`` row, not a failed admin request.
""" """
shape = probe_shape(base_url, api_key, upstream, model) shape = probe_shape(base_url, api_key, upstream, model)
result = ProbeResult( result = ProbeResult(
@@ -282,7 +276,14 @@ async def probe_upstream(
models_url=shape.models_url, models_url=shape.models_url,
chat_url=shape.chat_url, chat_url=shape.chat_url,
endpoint_tag=endpoint_tag, endpoint_tag=endpoint_tag,
provider_type=upstream.provider_type if upstream is not None else None,
) )
if result.provider_type == "typesafe":
result.completion_skip_reason = (
"Not evaluated — TypeSafe uses native /v1/systemone, not chat "
"completions. No supported native inference fixture is configured; "
"usage, billing and cache behavior remain unverified."
)
headers = shape.headers headers = shape.headers
owns_client = client is None owns_client = client is None
@@ -313,22 +314,11 @@ async def probe_upstream(
result.models_error = f"{type(exc).__name__}: {exc}" result.models_error = f"{type(exc).__name__}: {exc}"
result.models_latency_ms = round((time.monotonic() - started) * 1000, 2) result.models_latency_ms = round((time.monotonic() - started) * 1000, 2)
if not model_id: if not model_id or result.completion_skip_reason:
return result return result
await _probe_chat( await _probe_chat(
client, result, model_id, shape, timeout, upstream, model, "max_tokens" client, result, model_id, shape, timeout, upstream, model, "max_tokens"
) )
if wants_max_completion_tokens(result.chat_status, result.chat_payload):
await _probe_chat(
client,
result,
model_id,
shape,
timeout,
upstream,
model,
"max_completion_tokens",
)
finally: finally:
if owns_client: if owns_client:
await client.aclose() await client.aclose()
@@ -346,7 +336,7 @@ async def _probe_chat(
model: "Model | None", model: "Model | None",
token_field: str, token_field: str,
) -> None: ) -> None:
"""Send the one-token completion and record its outcome on ``result``.""" """Record the effective provider-shaped field, budget and every attempt."""
request_body: dict[str, Any] = { request_body: dict[str, Any] = {
"model": model_id, "model": model_id,
"messages": [{"role": "user", "content": PROBE_PROMPT}], "messages": [{"role": "user", "content": PROBE_PROMPT}],
@@ -358,35 +348,21 @@ async def _probe_chat(
"order": [result.endpoint_tag], "order": [result.endpoint_tag],
"allow_fallbacks": False, "allow_fallbacks": False,
} }
result.token_limit_field = token_field outcome = await send_completion(
result.chat_status = None client,
result.chat_payload = None result.chat_url,
result.chat_error = None shape_body(request_body, upstream, model),
started = time.monotonic() shape.headers,
try: shape.chat_params,
async with asyncio.timeout(timeout): timeout,
response = await client.post( )
result.chat_url, result.token_limit_field = outcome.token_limit_field
json=shape_body(request_body, upstream, model), result.token_limit = outcome.token_limit
headers=shape.headers, result.chat_status = outcome.status
params=shape.chat_params, result.chat_payload = outcome.payload
) result.chat_error = outcome.error
result.chat_status = response.status_code result.chat_latency_ms = outcome.latency_ms
result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) result.chat_attempts = outcome.attempts
try:
payload = response.json()
except Exception as exc: # noqa: BLE001 - any decode failure is the signal
result.chat_error = f"{type(exc).__name__}: {exc}"
else:
if isinstance(payload, dict):
result.chat_payload = payload
else:
result.chat_error = (
f"expected a JSON object, got {type(payload).__name__}"
)
except Exception as exc: # noqa: BLE001 - transport failure is a row status
result.chat_error = f"{type(exc).__name__}: {exc}"
result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2)
# Row builders are pure: the network lives only in ``probe_upstream`` and # Row builders are pure: the network lives only in ``probe_upstream`` and
@@ -473,13 +449,16 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]:
{"url": probe.models_url, "error": probe.models_error}, {"url": probe.models_url, "error": probe.models_error},
) )
data = payload.get("data") data_key, id_key = (
("models", "name") if probe.provider_type == "typesafe" else ("data", "id")
)
data = payload.get(data_key)
if not isinstance(data, list): if not isinstance(data, list):
return certification_row( return certification_row(
"endpoint.models_payload", "endpoint.models_payload",
STATUS_FAIL, STATUS_FAIL,
"Models payload has the expected shape", "Models payload has the expected shape",
f'Expected a top-level "data" list, got {type(data).__name__}.', f'Expected a top-level "{data_key}" list, got {type(data).__name__}.',
{ {
"url": probe.models_url, "url": probe.models_url,
"top_level_keys": sorted(payload.keys()), "top_level_keys": sorted(payload.keys()),
@@ -487,9 +466,9 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]:
) )
ids = [ ids = [
item["id"] item[id_key]
for item in data for item in data
if isinstance(item, dict) and isinstance(item.get("id"), str) and item["id"] if isinstance(item, dict) and isinstance(item.get(id_key), str) and item[id_key]
] ]
evidence: dict[str, Any] = { evidence: dict[str, Any] = {
"url": probe.models_url, "url": probe.models_url,
@@ -502,15 +481,15 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]:
"endpoint.models_payload", "endpoint.models_payload",
STATUS_FAIL, STATUS_FAIL,
"Models payload has the expected shape", "Models payload has the expected shape",
f'The "data" list carries no entry with a non-empty string "id" ' f'The "{data_key}" list carries no entry with a non-empty string '
f"({len(data)} entries).", f'"{id_key}" ({len(data)} entries).',
evidence, evidence,
) )
return certification_row( return certification_row(
"endpoint.models_payload", "endpoint.models_payload",
STATUS_OK, STATUS_OK,
"Models payload has the expected shape", "Models payload has the expected shape",
f"{len(ids)} of {len(data)} entries carry a string id.", f"{len(ids)} of {len(data)} entries carry a string {id_key}.",
evidence, evidence,
) )
@@ -518,14 +497,25 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]:
def usage_capture_row(probe: ProbeResult) -> dict[str, Any]: def usage_capture_row(probe: ProbeResult) -> dict[str, Any]:
"""Check a completion comes back with token usage the node can bill on. """Check a completion comes back with token usage the node can bill on.
A missing ``usage`` object means the node has nothing to price and the Missing upstream usage is a coverage gap. The proxy can estimate usage
request settles for free. Broken, but still usable, so ``warn``. before settlement; this direct probe does not exercise that fallback.
""" """
evidence: dict[str, Any] = { evidence: dict[str, Any] = {
"url": probe.chat_url, "url": probe.chat_url,
"status_code": probe.chat_status, "status_code": probe.chat_status,
"latency_ms": probe.chat_latency_ms, "latency_ms": probe.chat_latency_ms,
"token_limit_field": probe.token_limit_field,
"token_limit": probe.token_limit,
"attempts": probe.chat_attempts,
} }
if probe.completion_skip_reason:
return certification_row(
"usage.capture",
STATUS_WARN,
"Token usage captured from a completion",
probe.completion_skip_reason,
evidence,
)
if probe.chat_status is None: if probe.chat_status is None:
evidence["error"] = probe.chat_error evidence["error"] = probe.chat_error
return certification_row( return certification_row(
@@ -536,13 +526,17 @@ def usage_capture_row(probe: ProbeResult) -> dict[str, Any]:
evidence, evidence,
) )
if not 200 <= probe.chat_status < 300: if not 200 <= probe.chat_status < 300:
evidence["body"] = _truncate(probe.chat_payload) evidence["body"] = (
probe.chat_attempts[-1]["body"]
if probe.chat_attempts
else _truncate(probe.chat_payload)
)
return certification_row( return certification_row(
"usage.capture", "usage.capture",
STATUS_FAIL, STATUS_FAIL,
"Token usage captured from a completion", "Token usage captured from a completion",
f"{probe.chat_url} answered {probe.chat_status} for a " f"{probe.chat_url} answered {probe.chat_status} for a "
f"{PROBE_MAX_TOKENS}-token probe.", f"{probe.token_limit}-token probe ({probe.token_limit_field}).",
evidence, evidence,
) )
if probe.chat_payload is None or not isinstance(probe.chat_payload, dict): if probe.chat_payload is None or not isinstance(probe.chat_payload, dict):
@@ -576,8 +570,9 @@ def usage_capture_row(probe: ProbeResult) -> dict[str, Any]:
"usage.capture", "usage.capture",
STATUS_WARN, STATUS_WARN,
"Token usage captured from a completion", "Token usage captured from a completion",
'The completion carried no "usage" object, so the node has no ' 'The completion carried no "usage" object. The proxy may estimate '
"token counts to bill on and the request would settle as (0+0).", "missing usage before settlement; this probe does not verify that "
"fallback or the resulting client debit.",
evidence, evidence,
) )
evidence["input_tokens"] = normalized.input_tokens evidence["input_tokens"] = normalized.input_tokens
@@ -782,8 +777,8 @@ def cost_prompt_completion_row(
"cost.prompt_completion", "cost.prompt_completion",
STATUS_FAIL, STATUS_FAIL,
"Prompt and completion cost calculated", "Prompt and completion cost calculated",
f"The expected charge could not be derived from the configured " f"The expected charge could not be derived from the selected "
f"pricing: {type(exc).__name__}: {exc}.", f"billing basis: {type(exc).__name__}: {exc}.",
evidence, evidence,
) )
@@ -826,12 +821,17 @@ def cost_prompt_completion_row(
): ):
mismatches.append(f"input {actual_input} != {expected_input}") mismatches.append(f"input {actual_input} != {expected_input}")
basis_label = (
"upstream-reported USD converted with the provider fee and exchange rate"
if basis == "upstream_reported_usd"
else "configured token pricing"
)
if mismatches: if mismatches:
return certification_row( return certification_row(
"cost.prompt_completion", "cost.prompt_completion",
STATUS_FAIL, STATUS_FAIL,
"Prompt and completion cost calculated", "Prompt and completion cost calculated",
"The computed charge disagrees with the configured pricing: " f"The computed charge disagrees with {basis_label}: "
+ "; ".join(mismatches) + "; ".join(mismatches)
+ ".", + ".",
evidence, evidence,
@@ -840,10 +840,10 @@ def cost_prompt_completion_row(
"cost.prompt_completion", "cost.prompt_completion",
STATUS_OK, STATUS_OK,
"Prompt and completion cost calculated", "Prompt and completion cost calculated",
f"Charged {actual_total} msats ({actual_input} input + " f"Calculated {actual_total} msats ({actual_input} input + "
f"{actual_output} output) for {usage.input_tokens} prompt and " f"{actual_output} output) for {usage.input_tokens} prompt and "
f"{usage.output_tokens} completion tokens, matching the configured " f"{usage.output_tokens} completion tokens, matching {basis_label}. "
f"pricing.", "This is a calculation, not a measured client debit.",
evidence, evidence,
) )
@@ -901,8 +901,27 @@ async def run_live_checks(
), ),
] ]
from .certification_cache import run_cache_checks, skipped_cache_rows
if probe.completion_skip_reason:
rows.append(
certification_row(
"cost.prompt_completion",
STATUS_WARN,
"Prompt and completion cost calculated",
probe.completion_skip_reason,
{},
)
)
rows.extend(skipped_cache_rows(probe.completion_skip_reason))
return rows
cost_data: Any = None cost_data: Any = None
if probe.chat_payload is not None and probe.chat_status is not None: if (
probe.chat_payload is not None
and probe.chat_status is not None
and 200 <= probe.chat_status < 300
):
try: try:
cost_data = await calculate_cost( cost_data = await calculate_cost(
probe.chat_payload, probe.chat_payload,
@@ -939,8 +958,6 @@ async def run_live_checks(
) )
) )
from .certification_cache import run_cache_checks, skipped_cache_rows
if not check_cache: if not check_cache:
rows.extend(skipped_cache_rows("Skipped — cache checks disabled.")) rows.extend(skipped_cache_rows("Skipped — cache checks disabled."))
elif ( elif (
@@ -973,6 +990,7 @@ async def run_live_checks(
endpoint_tag=endpoint_tag, endpoint_tag=endpoint_tag,
upstream=upstream, upstream=upstream,
token_limit_field=probe.token_limit_field, token_limit_field=probe.token_limit_field,
token_limit=probe.token_limit,
) )
) )
return rows return rows
+171 -88
View File
@@ -1,10 +1,10 @@
"""Prompt-cache and margin certification for an upstream provider. """Prompt-cache and margin certification for an upstream provider.
Three questions the one-token probe cannot answer: Three questions the short completion probe cannot answer:
* does the upstream *report* prompt-cache hits in a dialect the node parses, * does the upstream *report* prompt-cache hits in a dialect the node parses,
* does the node bill cached reads at the discounted rate (client side), and * does the cost engine calculate cached-completion charges correctly, and
* does the node's charge cover what the upstream charged (node side). * does the token estimate cover the reported upstream cost target.
The cache probe sends the same long system prompt twice; the second call is The cache probe sends the same long system prompt twice; the second call is
the one expected to report cached reads. Calls go straight to the upstream the one expected to report cached reads. Calls go straight to the upstream
@@ -13,8 +13,6 @@ with ``httpx`` and never enter the billing path.
from __future__ import annotations from __future__ import annotations
import asyncio
import time
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
@@ -36,11 +34,13 @@ from .certification import (
_fixed_token_pricing_active, _fixed_token_pricing_active,
_reported_usd_cost, _reported_usd_cost,
_token_rates, _token_rates,
_truncate,
certification_row, certification_row,
probe_shape, probe_shape,
safe_row, safe_row,
shape_body, shape_body,
) )
from .certification_probe import CompletionProbe, send_completion
if TYPE_CHECKING: if TYPE_CHECKING:
from ..payment.models import Model from ..payment.models import Model
@@ -58,8 +58,8 @@ ROW_BILLING = "cache.billing"
ROW_MARGIN = "cost.margin" ROW_MARGIN = "cost.margin"
TITLE_REPORTED = "Upstream reports prompt-cache hits" TITLE_REPORTED = "Upstream reports prompt-cache hits"
TITLE_BILLING = "Cached tokens billed at the cache-read rate" TITLE_BILLING = "Cached completion cost calculated"
TITLE_MARGIN = "Node charge covers upstream cost" TITLE_MARGIN = "Token estimate covers fee-adjusted target"
def cache_probe_prefix() -> str: def cache_probe_prefix() -> str:
@@ -86,6 +86,9 @@ class CacheProbeResult:
payloads: list[dict[str, Any] | None] = field(default_factory=list) payloads: list[dict[str, Any] | None] = field(default_factory=list)
errors: list[str | None] = field(default_factory=list) errors: list[str | None] = field(default_factory=list)
latencies_ms: list[float | None] = field(default_factory=list) latencies_ms: list[float | None] = field(default_factory=list)
attempts: list[dict[str, Any]] = field(default_factory=list)
token_limit_field: str = "max_tokens"
token_limit: int = PROBE_MAX_TOKENS
@property @property
def second_payload(self) -> dict[str, Any] | None: def second_payload(self) -> dict[str, Any] | None:
@@ -104,6 +107,7 @@ def _request_body(
fmt: str, fmt: str,
endpoint_tag: str | None, endpoint_tag: str | None,
token_field: str = "max_tokens", token_field: str = "max_tokens",
token_limit: int = PROBE_MAX_TOKENS,
) -> dict[str, Any]: ) -> dict[str, Any]:
system: Any system: Any
if fmt == "cache_control": if fmt == "cache_control":
@@ -122,7 +126,7 @@ def _request_body(
{"role": "system", "content": system}, {"role": "system", "content": system},
{"role": "user", "content": CACHE_PROBE_QUESTION}, {"role": "user", "content": CACHE_PROBE_QUESTION},
], ],
token_field: PROBE_MAX_TOKENS, token_field: token_limit,
"stream": False, "stream": False,
} }
if endpoint_tag: if endpoint_tag:
@@ -133,45 +137,11 @@ def _request_body(
return body return body
async def _post_completion( def _record(result: CacheProbeResult, outcome: CompletionProbe) -> None:
client: httpx.AsyncClient, result.statuses.append(outcome.status)
url: str, result.payloads.append(outcome.payload)
body: Any, result.errors.append(outcome.error)
headers: dict[str, str], result.latencies_ms.append(outcome.latency_ms)
params: dict[str, str],
timeout: float,
) -> tuple[int | None, dict[str, Any] | None, str | None, float]:
started = time.monotonic()
try:
async with asyncio.timeout(timeout):
response = await client.post(url, json=body, headers=headers, params=params)
except Exception as exc: # noqa: BLE001 - transport failure is a row status
latency = round((time.monotonic() - started) * 1000, 2)
return None, None, f"{type(exc).__name__}: {exc}", latency
latency = round((time.monotonic() - started) * 1000, 2)
try:
payload = response.json()
except Exception as exc: # noqa: BLE001 - any decode failure is the signal
return response.status_code, None, f"{type(exc).__name__}: {exc}", latency
if not isinstance(payload, dict):
return (
response.status_code,
None,
f"expected a JSON object, got {type(payload).__name__}",
latency,
)
return response.status_code, payload, None, latency
def _record(
result: CacheProbeResult,
outcome: tuple[int | None, dict[str, Any] | None, str | None, float],
) -> None:
status, payload, error, latency = outcome
result.statuses.append(status)
result.payloads.append(payload)
result.errors.append(error)
result.latencies_ms.append(latency)
def _is_2xx(status: int | None) -> bool: def _is_2xx(status: int | None) -> bool:
@@ -189,6 +159,7 @@ async def probe_cache(
upstream: "BaseUpstreamProvider | None" = None, upstream: "BaseUpstreamProvider | None" = None,
model: "Model | None" = None, model: "Model | None" = None,
token_limit_field: str = "max_tokens", token_limit_field: str = "max_tokens",
token_limit: int = PROBE_MAX_TOKENS,
) -> CacheProbeResult: ) -> CacheProbeResult:
"""Send the same long prompt twice. """Send the same long prompt twice.
@@ -198,17 +169,27 @@ async def probe_cache(
succeeded. Each call's elapsed deadline includes the response body. succeeded. Each call's elapsed deadline includes the response body.
""" """
shape = probe_shape(base_url, api_key, upstream, model) shape = probe_shape(base_url, api_key, upstream, model)
result = CacheProbeResult(chat_url=shape.chat_url, endpoint_tag=endpoint_tag) result = CacheProbeResult(
chat_url=shape.chat_url,
endpoint_tag=endpoint_tag,
token_limit_field=token_limit_field,
token_limit=token_limit,
)
prefix = cache_probe_prefix() prefix = cache_probe_prefix()
owns_client = client is None owns_client = client is None
http = client if client is not None else httpx.AsyncClient(timeout=timeout) http = client if client is not None else httpx.AsyncClient(timeout=timeout)
async def post( async def post(fmt: str, phase: str) -> CompletionProbe:
fmt: str, body = _request_body(
) -> tuple[int | None, dict[str, Any] | None, str | None, float]: model_id,
body = _request_body(model_id, prefix, fmt, endpoint_tag, token_limit_field) prefix,
return await _post_completion( fmt,
endpoint_tag,
result.token_limit_field,
result.token_limit,
)
outcome = await send_completion(
http, http,
result.chat_url, result.chat_url,
shape_body(body, upstream, model), shape_body(body, upstream, model),
@@ -216,16 +197,23 @@ async def probe_cache(
shape.chat_params, shape.chat_params,
timeout, timeout,
) )
result.token_limit_field = outcome.token_limit_field
result.token_limit = outcome.token_limit
result.attempts.extend(
{**attempt, "request_format": fmt, "phase": phase}
for attempt in outcome.attempts
)
return outcome
try: try:
first = await post("cache_control") first = await post("cache_control", "warm_up")
if first[0] in (400, 422): if first.status in (400, 422):
result.request_format = "plain" result.request_format = "plain"
first = await post("plain") first = await post("plain", "warm_up")
_record(result, first) _record(result, first)
if not _is_2xx(first[0]): if not _is_2xx(first.status):
return result return result
second = await post(result.request_format) second = await post(result.request_format, "repeated")
_record(result, second) _record(result, second)
finally: finally:
if owns_client: if owns_client:
@@ -262,16 +250,29 @@ def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]:
"endpoint_tag": probe.endpoint_tag, "endpoint_tag": probe.endpoint_tag,
"statuses": probe.statuses, "statuses": probe.statuses,
"latencies_ms": probe.latencies_ms, "latencies_ms": probe.latencies_ms,
"attempts": probe.attempts,
"errors": probe.errors,
"token_limit_field": probe.token_limit_field,
"token_limit": probe.token_limit,
} }
payload = probe.second_payload payload = probe.second_payload
if payload is None or not _is_2xx(probe.statuses[-1] if probe.statuses else None): if payload is None or not _is_2xx(probe.statuses[-1] if probe.statuses else None):
evidence["error"] = probe.second_error status = probe.statuses[-1] if probe.statuses else None
if status is not None and not _is_2xx(status):
body = (
probe.attempts[-1].get("body")
if probe.attempts
else _truncate(probe.payloads[-1] if probe.payloads else None)
)
reason = f"HTTP {status}" + (f": {body}" if body else "")
else:
reason = probe.second_error or "no decodable completion body"
evidence["error"] = reason
return certification_row( return certification_row(
ROW_REPORTED, ROW_REPORTED,
STATUS_FAIL, STATUS_FAIL,
TITLE_REPORTED, TITLE_REPORTED,
f"The cache probe did not get two successful completions: " f"The cache probe did not get two successful completions: {reason}.",
f"{probe.second_error or 'no response body'}.",
evidence, evidence,
) )
@@ -308,17 +309,17 @@ def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]:
STATUS_FAIL, STATUS_FAIL,
TITLE_REPORTED, TITLE_REPORTED,
"The upstream reported cache tokens under fields the node does not " "The upstream reported cache tokens under fields the node does not "
f"parse ({', '.join(raw_keys)}); cached reads would be billed at " f"parse ({', '.join(raw_keys)}). Token-based pricing cannot apply "
"the full input rate.", "a cache-read rate to those fields; reported-USD billing may take "
"precedence.",
evidence, evidence,
) )
return certification_row( return certification_row(
ROW_REPORTED, ROW_REPORTED,
STATUS_WARN, STATUS_WARN,
TITLE_REPORTED, TITLE_REPORTED,
"Two identical prompts produced no cache hit. Either the model does " "Two identical prompts produced no reported cache hit. Caching may "
"not support prompt caching or the upstream hides it; clients pay the " "be unsupported or unreported; no cache discount is verified.",
"full input rate on repeated prompts.",
evidence, evidence,
) )
@@ -329,6 +330,8 @@ def cache_billing_row(
probe: CacheProbeResult, probe: CacheProbeResult,
cost_data: Any, cost_data: Any,
pricing_known: bool = True, pricing_known: bool = True,
provider_fee: float = 1.0,
sats_to_usd: float | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
usage = _usage_of(probe.second_payload) usage = _usage_of(probe.second_payload)
evidence: dict[str, Any] = {"model_id": model.id} evidence: dict[str, Any] = {"model_id": model.id}
@@ -345,7 +348,7 @@ def cache_billing_row(
ROW_BILLING, ROW_BILLING,
STATUS_WARN, STATUS_WARN,
TITLE_BILLING, TITLE_BILLING,
"No pricing is known for this model, so the cache discount cannot " "No pricing is known for this model, so the cache charge cannot "
"be verified.", "be verified.",
evidence, evidence,
) )
@@ -395,14 +398,59 @@ def cache_billing_row(
} }
) )
evidence["basis"] = (
"upstream_reported_usd" if reported_usd > 0 else "configured_token_pricing"
)
if reported_usd > 0: if reported_usd > 0:
evidence.update(
{
"token_estimate_msats": expected_total,
"provider_fee": provider_fee,
"sats_usd_price": sats_to_usd,
"expected_total_msats": None,
"upstream_discount_verified": False,
}
)
if sats_to_usd is None:
return certification_row(
ROW_BILLING,
STATUS_WARN,
TITLE_BILLING,
"No exchange rate was supplied, so the reported-USD calculation "
"cannot be verified.",
evidence,
)
try:
expected_total = _expected_usd_msats(
reported_usd, provider_fee, sats_to_usd
)
except (ValueError, OverflowError) as exc:
evidence["error"] = f"{type(exc).__name__}: {exc}"
return certification_row(
ROW_BILLING,
STATUS_FAIL,
TITLE_BILLING,
f"The expected reported-USD charge could not be derived: {exc}.",
evidence,
)
evidence["expected_total_msats"] = expected_total
if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS:
return certification_row(
ROW_BILLING,
STATUS_FAIL,
TITLE_BILLING,
f"The engine calculated {actual_total} msats but reported USD, "
f"the provider fee and exchange rate imply {expected_total} msats.",
evidence,
)
return certification_row( return certification_row(
ROW_BILLING, ROW_BILLING,
STATUS_OK, STATUS_OK,
TITLE_BILLING, TITLE_BILLING,
f"Billed {actual_total} msats from the upstream-reported cost, " f"Calculated {actual_total} msats, matching the independently "
f"which already carries the cache discount " "converted upstream-reported USD with the provider fee. "
f"(full token price would be {full_total} msats).", f"The configured full-input comparator is {full_total} msats; "
"an upstream cache discount and actual client debit are not verified.",
evidence, evidence,
) )
if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS: if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS:
@@ -410,7 +458,7 @@ def cache_billing_row(
ROW_BILLING, ROW_BILLING,
STATUS_FAIL, STATUS_FAIL,
TITLE_BILLING, TITLE_BILLING,
f"The engine charged {actual_total} msats but the configured cache " f"The engine calculated {actual_total} msats but the configured cache "
f"rate implies {expected_total} msats.", f"rate implies {expected_total} msats.",
evidence, evidence,
) )
@@ -424,16 +472,16 @@ def cache_billing_row(
ROW_BILLING, ROW_BILLING,
STATUS_WARN, STATUS_WARN,
TITLE_BILLING, TITLE_BILLING,
f"Cached reads are billed at the full input rate ({actual_total} " f"Cached reads use the full input rate ({actual_total} calculated "
f"msats) because {reason}; clients pay more than the upstream " f"msats) because {reason}. Upstream cost, any upstream discount "
"charges.", "and actual client debit are unverified.",
evidence, evidence,
) )
return certification_row( return certification_row(
ROW_BILLING, ROW_BILLING,
STATUS_OK, STATUS_OK,
TITLE_BILLING, TITLE_BILLING,
f"Charged {actual_total} msats for {usage.cache_read_tokens} cached " f"Calculated {actual_total} msats for {usage.cache_read_tokens} cached "
f"tokens, {full_total - actual_total} msats below the full input price.", f"tokens, {full_total - actual_total} msats below the full input price.",
evidence, evidence,
) )
@@ -449,12 +497,9 @@ def cost_margin_row(
) -> dict[str, Any]: ) -> dict[str, Any]:
"""Configured token pricing must cover what the upstream reports charging. """Configured token pricing must cover what the upstream reports charging.
Responses that carry a USD cost are billed from it, so they cannot lose Reported USD takes precedence in the cost engine; the token-price estimate
money themselves; they are used here as a price sample. The configured is a separate fallback comparison, not a measured debit. A fee-target gap
token pricing is what every other path bills from (streams, upstreams need not be a loss against raw upstream cost. Missing cost stays unverified.
that omit cost, the served ``/v1/models`` list), so a sample where it
falls below the fee-adjusted upstream cost means those paths underprice.
Upstreams that report no cost give no sample and the row stays a warn.
``model`` carries the pricing the proxy reserves and token-bills with, ``model`` carries the pricing the proxy reserves and token-bills with,
which on a pinned endpoint is that endpoint's own rates. which on a pinned endpoint is that endpoint's own rates.
@@ -477,6 +522,7 @@ def cost_margin_row(
samples: list[dict[str, Any]] = [] samples: list[dict[str, Any]] = []
short: list[str] = [] short: list[str] = []
below_raw: list[str] = []
for payload in payloads: for payload in payloads:
if not isinstance(payload, dict): if not isinstance(payload, dict):
continue continue
@@ -489,6 +535,7 @@ def cost_margin_row(
upstream_total = _expected_usd_msats( upstream_total = _expected_usd_msats(
reported_usd, provider_fee, sats_to_usd reported_usd, provider_fee, sats_to_usd
) )
raw_upstream_total = _expected_usd_msats(reported_usd, 1.0, sats_to_usd)
except (ValueError, OverflowError) as exc: except (ValueError, OverflowError) as exc:
evidence["error"] = f"{type(exc).__name__}: {exc}" evidence["error"] = f"{type(exc).__name__}: {exc}"
return certification_row( return certification_row(
@@ -502,11 +549,17 @@ def cost_margin_row(
"usage": usage.dict(), "usage": usage.dict(),
"reported_usd": reported_usd, "reported_usd": reported_usd,
"upstream_msats_with_fee": upstream_total, "upstream_msats_with_fee": upstream_total,
"upstream_msats": raw_upstream_total,
"configured_msats": configured_total, "configured_msats": configured_total,
"token_estimate_minus_upstream_msats": configured_total
- raw_upstream_total,
"token_estimate_minus_fee_target_msats": configured_total - upstream_total,
} }
samples.append(sample) samples.append(sample)
if configured_total + COST_TOLERANCE_MSATS < upstream_total: if configured_total + COST_TOLERANCE_MSATS < upstream_total:
short.append(f"{configured_total} < {upstream_total}") short.append(f"{configured_total} < {upstream_total}")
if configured_total + COST_TOLERANCE_MSATS < raw_upstream_total:
below_raw.append(f"{configured_total} < {raw_upstream_total}")
evidence["samples"] = samples evidence["samples"] = samples
if not samples: if not samples:
@@ -524,17 +577,30 @@ def cost_margin_row(
ROW_MARGIN, ROW_MARGIN,
STATUS_FAIL, STATUS_FAIL,
TITLE_MARGIN, TITLE_MARGIN,
"Configured pricing is below the upstream's reported cost " "The token-pricing estimate is below the fee-adjusted cost target "
f"(configured < upstream msats: {'; '.join(short)}); token-billed " f"(estimate < target msats: {'; '.join(short)}). "
"requests lose money.", + (
"It is also below raw upstream cost "
f"(estimate < raw msats: {'; '.join(below_raw)}). "
if below_raw
else "Raw upstream cost is covered; the gap is in the fee target. "
)
+ "Responses reporting USD use reported-cost billing instead; "
"this check does not measure client debits.",
evidence, evidence,
) )
return certification_row( return certification_row(
ROW_MARGIN, ROW_MARGIN,
STATUS_OK, STATUS_OK,
TITLE_MARGIN, TITLE_MARGIN,
f"Configured pricing covers the upstream's reported cost on " f"Token pricing covers the fee-adjusted cost target on "
f"{len(samples)} sampled completion(s).", f"{len(samples)} sampled completion(s)."
+ (
" The fee multiplier is below break-even: raw upstream cost exceeds "
f"the token estimate on samples ({'; '.join(below_raw)})."
if below_raw
else ""
),
evidence, evidence,
) )
@@ -570,6 +636,7 @@ async def run_cache_checks(
endpoint_tag: str | None = None, endpoint_tag: str | None = None,
upstream: "BaseUpstreamProvider | None" = None, upstream: "BaseUpstreamProvider | None" = None,
token_limit_field: str = "max_tokens", token_limit_field: str = "max_tokens",
token_limit: int = PROBE_MAX_TOKENS,
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
"""Run the cache probe and build the three cache/margin rows.""" """Run the cache probe and build the three cache/margin rows."""
probe = await probe_cache( probe = await probe_cache(
@@ -582,8 +649,15 @@ async def run_cache_checks(
upstream=upstream, upstream=upstream,
model=model, model=model,
token_limit_field=token_limit_field, token_limit_field=token_limit_field,
token_limit=token_limit,
)
cost_data = await _price_payload(
probe.second_payload
if probe.statuses and _is_2xx(probe.statuses[-1])
else None,
model,
provider_fee,
) )
cost_data = await _price_payload(probe.second_payload, model, provider_fee)
return [ return [
safe_row(ROW_REPORTED, TITLE_REPORTED, lambda: cache_reported_row(probe)), safe_row(ROW_REPORTED, TITLE_REPORTED, lambda: cache_reported_row(probe)),
safe_row( safe_row(
@@ -594,6 +668,8 @@ async def run_cache_checks(
probe=probe, probe=probe,
cost_data=cost_data, cost_data=cost_data,
pricing_known=pricing_known, pricing_known=pricing_known,
provider_fee=provider_fee,
sats_to_usd=sats_to_usd,
), ),
), ),
safe_row( safe_row(
@@ -601,7 +677,14 @@ async def run_cache_checks(
TITLE_MARGIN, TITLE_MARGIN,
lambda: cost_margin_row( lambda: cost_margin_row(
model=model, model=model,
payloads=[probe_payload, *probe.payloads], payloads=[
probe_payload,
*[
payload
for status, payload in zip(probe.statuses, probe.payloads)
if _is_2xx(status)
],
],
provider_fee=provider_fee, provider_fee=provider_fee,
sats_to_usd=sats_to_usd, sats_to_usd=sats_to_usd,
pricing_known=pricing_known, pricing_known=pricing_known,
+156
View File
@@ -0,0 +1,156 @@
"""Bounded completion probes shared by ordinary and cache certification."""
from __future__ import annotations
import asyncio
import json
import re
import time
from dataclasses import dataclass, field
from typing import Any
import httpx
PROBE_TIMEOUT_SECONDS = 15.0
PROBE_TOKEN_BUDGETS = (32, 128, 512, 2048)
PROBE_MAX_TOKENS = PROBE_TOKEN_BUDGETS[0]
# One field correction plus the four bounded budgets, never general retries.
PROBE_MAX_ATTEMPTS = len(PROBE_TOKEN_BUDGETS) + 1
def wants_max_completion_tokens(status: int | None, payload: Any) -> bool:
if status != 400 or payload is None:
return False
text = json.dumps(payload, default=str).lower()
return "max_completion_tokens" in text and (
"unsupported" in text
and "max_tokens" in text
or "use 'max_completion_tokens'" in text
)
def next_probe_budget(status: int | None, payload: Any, current: int) -> int | None:
if status not in (400, 422) or payload is None:
return None
text = json.dumps(payload, default=str).lower()
minimum = re.search(
r"(?:max_tokens|max_completion_tokens|max_output_tokens)"
r".{0,80}?(?:at least|minimum(?: of)?|>=|greater than or equal to)"
r"\D{0,12}(\d+)",
text,
)
required = int(minimum[1]) if minimum else None
if required is not None:
if required <= current:
return None
elif not (
"max_tokens or model output limit was reached" in text
or (
any(name in text for name in ("max_tokens", "max_completion_tokens"))
and any(
hint in text for hint in ("higher max", "exhausted", "increase max")
)
)
):
return None
return next(
(n for n in PROBE_TOKEN_BUDGETS if n > current and n >= (required or 0)),
None,
)
def diagnostic_text(value: Any, headers: dict[str, str], limit: int = 1000) -> str:
"""Redact known request credentials before bounding diagnostic text."""
text = value if isinstance(value, str) else json.dumps(value, default=str)
for name, secret in headers.items():
if name.lower() in {"authorization", "api-key", "x-api-key"} and secret:
text = text.replace(secret, "[redacted]")
if name.lower() == "authorization" and " " in secret:
text = text.replace(secret.split(" ", 1)[1], "[redacted]")
return text if len(text) <= limit else text[:limit] + "…"
@dataclass
class CompletionProbe:
status: int | None = None
payload: dict[str, Any] | None = None
error: str | None = None
latency_ms: float = 0.0
token_limit_field: str = "max_tokens"
token_limit: int = PROBE_MAX_TOKENS
attempts: list[dict[str, Any]] = field(default_factory=list)
async def send_completion(
client: httpx.AsyncClient,
url: str,
body: dict[str, Any],
headers: dict[str, str],
params: dict[str, str],
timeout: float,
) -> CompletionProbe:
"""Adapt only recognized token-limit rejections, preserving every call.
The body is already provider-shaped. This policy is exclusive to synthetic
certification requests; user request budgets are never changed here.
"""
body = dict(body)
result = CompletionProbe()
for _ in range(PROBE_MAX_ATTEMPTS):
token_field = (
"max_completion_tokens" if "max_completion_tokens" in body else "max_tokens"
)
budget = int(body[token_field])
if not 1 <= budget <= PROBE_TOKEN_BUDGETS[-1]:
raise ValueError("certification token budget is outside the probe ceiling")
result.token_limit_field = token_field
result.token_limit = budget
result.status, result.payload, result.error = None, None, None
diagnostic: Any = None
started = time.monotonic()
try:
async with asyncio.timeout(timeout):
response = await client.post(
url, json=body, headers=headers, params=params
)
result.status = response.status_code
diagnostic = response.text
try:
payload = response.json()
except Exception as exc: # noqa: BLE001 - decoding failure is evidence
result.error = f"{type(exc).__name__}: {exc}"
else:
if isinstance(payload, dict):
result.payload = payload
else:
result.error = (
f"expected a JSON object, got {type(payload).__name__}"
)
except Exception as exc: # noqa: BLE001 - cancellation still propagates
result.error = f"{type(exc).__name__}: {exc}"
result.latency_ms = round((time.monotonic() - started) * 1000, 2)
if result.error:
result.error = diagnostic_text(result.error, headers)
result.attempts.append(
{
"status_code": result.status,
"latency_ms": result.latency_ms,
"token_limit_field": token_field,
"token_limit": budget,
"error": result.error,
"body": diagnostic_text(diagnostic, headers)
if diagnostic is not None
and (result.error or not 200 <= (result.status or 0) < 300)
else None,
}
)
if token_field == "max_tokens" and wants_max_completion_tokens(
result.status, result.payload
):
body["max_completion_tokens"] = body.pop("max_tokens")
continue
next_budget = next_probe_budget(result.status, result.payload, budget)
if next_budget is None:
break
body[token_field] = next_budget
return result
+4 -1
View File
@@ -123,7 +123,10 @@ async def get_all_models_with_overrides(
if model_key is not None and model_key in overrides_by_key: if model_key is not None and model_key in overrides_by_key:
override_row, provider_fee = overrides_by_key[model_key] override_row, provider_fee = overrides_by_key[model_key]
all_models[(model.id.lower(), provider_key)] = _row_to_model( all_models[(model.id.lower(), provider_key)] = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee override_row,
apply_provider_fee=True,
provider_fee=provider_fee,
provider_type=upstream.provider_type,
) )
elif model.enabled: elif model.enabled:
all_models[(model.id.lower(), provider_key)] = model all_models[(model.id.lower(), provider_key)] = model
+15 -7
View File
@@ -824,7 +824,9 @@ async def refresh_model_paths_periodically(
break break
def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None: def _price_in_sats(
model: dict[str, Any], provider_fee: float, provider_type: str | None = None
) -> None:
"""Run a path's USD rates through the ``/v1/models`` pricing pipeline. """Run a path's USD rates through the ``/v1/models`` pricing pipeline.
Metadata copied from the provider model cache is already priced. OpenRouter Metadata copied from the provider model cache is already priced. OpenRouter
@@ -842,13 +844,16 @@ def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None:
TopProvider, TopProvider,
_calculate_usd_max_costs, _calculate_usd_max_costs,
_update_model_sats_pricing, _update_model_sats_pricing,
allows_cache_pricing_backfill,
backfill_cache_pricing, backfill_cache_pricing,
) )
from ..payment.price import sats_usd_price from ..payment.price import sats_usd_price
try: try:
model_id = model.get("forwarded_model_id") or model["id"] model_id = model.get("forwarded_model_id") or model["id"]
usd = backfill_cache_pricing(model_id, Pricing.parse_obj(pricing)) usd = Pricing.parse_obj(pricing)
if allows_cache_pricing_backfill(provider_type):
usd = backfill_cache_pricing(model_id, usd)
usd = Pricing.parse_obj({k: v * provider_fee for k, v in usd.dict().items()}) usd = Pricing.parse_obj({k: v * provider_fee for k, v in usd.dict().items()})
priced = Model( priced = Model(
id=model_id, id=model_id,
@@ -921,10 +926,13 @@ def apply_model_path_pricing(
metadata.get("pricing"), dict metadata.get("pricing"), dict
): ):
return model return model
pricing = backfill_cache_pricing( from ..payment.models import allows_cache_pricing_backfill
model.forwarded_model_id or row.model_id,
Pricing.parse_obj(metadata["pricing"]), pricing = Pricing.parse_obj(metadata["pricing"])
) if allows_cache_pricing_backfill(row.provider_type):
pricing = backfill_cache_pricing(
model.forwarded_model_id or row.model_id, pricing
)
pricing = Pricing.parse_obj( pricing = Pricing.parse_obj(
{key: float(value) * provider_fee for key, value in pricing.dict().items()} {key: float(value) * provider_fee for key, value in pricing.dict().items()}
) )
@@ -963,7 +971,7 @@ def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]:
if not isinstance(model, dict): if not isinstance(model, dict):
model = {} model = {}
model.setdefault("id", row.model_id) model.setdefault("id", row.model_id)
_price_in_sats(model, provider_fee) _price_in_sats(model, provider_fee, row.provider_type)
return { return {
"path": row.path, "path": row.path,
"provider": { "provider": {
+36 -45
View File
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import math
import random import random
import time import time
from dataclasses import dataclass, field from dataclasses import dataclass, field
@@ -226,6 +227,27 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
if ppqai_model.id in self.IGNORED_MODEL_IDS: if ppqai_model.id in self.IGNORED_MODEL_IDS:
continue continue
api_pricing = ppqai_model.pricing.api or {}
input_price = api_pricing.get("input_per_1M")
if input_price is None:
input_price = ppqai_model.pricing.input_per_1M_tokens
output_price = api_pricing.get("output_per_1M")
if output_price is None:
output_price = ppqai_model.pricing.output_per_1M_tokens
if any(
rate is None or not math.isfinite(rate) or rate < 0
for rate in (input_price, output_price)
):
logger.warning(
"Skipping PPQ model without valid native token prices",
extra={"model_id": ppqai_model.id},
)
continue
assert input_price is not None and output_price is not None
pricing = Pricing(
prompt=input_price / 1_000_000,
completion=output_price / 1_000_000,
)
or_model = next( or_model = next(
( (
model model
@@ -238,44 +260,20 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
) )
if or_model: if or_model:
input_price = None # OpenRouter supplies metadata, not PPQ billing rates.
if ppqai_model.pricing.api: models.append(
input_price = ppqai_model.pricing.api.get("input_per_1M") or_model.copy(
elif ppqai_model.pricing.input_per_1M_tokens: deep=True,
input_price = ppqai_model.pricing.input_per_1M_tokens update={
"id": ppqai_model.id,
if input_price is not None: "pricing": pricing,
or_model.pricing.prompt = input_price / 1_000_000 "sats_pricing": None,
"context_length": ppqai_model.context_length
output_price = None or or_model.context_length,
if ppqai_model.pricing.api: },
output_price = ppqai_model.pricing.api.get("output_per_1M") )
elif ppqai_model.pricing.output_per_1M_tokens: )
output_price = ppqai_model.pricing.output_per_1M_tokens
if output_price is not None:
or_model.pricing.completion = output_price / 1_000_000
if cl := ppqai_model.context_length:
or_model.context_length = cl
models.append(or_model)
else: else:
input_price = 0.0
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M", 0.0
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
output_price = 0.0
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M", 0.0
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
models.append( models.append(
Model( Model(
id=ppqai_model.id, id=ppqai_model.id,
@@ -290,14 +288,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
tokenizer="Unknown", tokenizer="Unknown",
instruct_type=None, instruct_type=None,
), ),
pricing=Pricing( pricing=pricing,
prompt=input_price / 1_000_000,
completion=output_price / 1_000_000,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
),
) )
) )
except Exception as e: except Exception as e:
@@ -448,7 +448,7 @@ async def test_cache_rate_ignores_an_enabled_model_that_is_not_served(
assert resp.status_code == 200, resp.text assert resp.status_code == 200, resp.text
row = _find_row(resp.json()["rows"], "pricing.cache_rate") row = _find_row(resp.json()["rows"], "pricing.cache_rate")
assert row["status"] == "ok", row assert row["status"] == "warn", row
assert row["evidence"]["checked"] == 0 assert row["evidence"]["checked"] == 0
assert "unserved-model" not in json.dumps(row["evidence"]) assert "unserved-model" not in json.dumps(row["evidence"])
+9 -3
View File
@@ -343,9 +343,7 @@ async def test_certify_margin_bills_pinned_path_pricing(
}, },
base_url=base_url, base_url=base_url,
) )
model_path = encode_model_path( model_path = encode_model_path(base_url, "cert-test-model", "deepinfra/fp8")
base_url, "cert-test-model", "deepinfra/fp8"
)
integration_session.add( integration_session.add(
ModelPathRow( ModelPathRow(
model_id="cert-test-model", model_id="cert-test-model",
@@ -834,6 +832,14 @@ async def test_certify_explicit_discovered_model_without_override(
assert _find_row(rows, "endpoint.reachable")["status"] == "ok" assert _find_row(rows, "endpoint.reachable")["status"] == "ok"
assert _find_row(rows, "usage.capture")["status"] == "ok" assert _find_row(rows, "usage.capture")["status"] == "ok"
assert _find_row(rows, "cost.prompt_completion")["status"] == "ok" assert _find_row(rows, "cost.prompt_completion")["status"] == "ok"
pricing_rows = [row for row in rows if row["id"].startswith("pricing.")]
assert len(pricing_rows) == 4
assert all(row["evidence"]["checked"] == 0 for row in pricing_rows)
assert all(row["status"] == "warn" for row in pricing_rows)
pricing_goal = next(
goal for goal in resp.json()["checklist"] if goal["goal"] == "pricing_v1_models"
)
assert pricing_goal["status"] == "warn"
@pytest.mark.integration @pytest.mark.integration
+8 -2
View File
@@ -272,7 +272,10 @@ def test_create_model_mappings_applies_custom_provider_fees_before_advertising(
2: SimpleNamespace(id="shared-model", upstream_provider_id=2, enabled=True), 2: SimpleNamespace(id="shared-model", upstream_provider_id=2, enabled=True),
} }
def fake_row_to_model(row, *, apply_provider_fee, provider_fee) -> Model: # type: ignore[no-untyped-def] def fake_row_to_model(
row, *, apply_provider_fee, provider_fee, provider_type
) -> Model: # type: ignore[no-untyped-def]
assert provider_type == providers[row.upstream_provider_id - 1].provider_type
assert apply_provider_fee is True assert apply_provider_fee is True
return create_test_model( return create_test_model(
row.id, row.id,
@@ -943,7 +946,10 @@ def test_create_model_mappings_uppercase_prefixed_base_keeps_top_tier() -> None:
"prefixed", "https://prefixed.example/v1", db_id=1, models=[prefixed_cheap] "prefixed", "https://prefixed.example/v1", db_id=1, models=[prefixed_cheap]
) )
forwarded_provider = create_test_provider( forwarded_provider = create_test_provider(
"forwarded", "https://forwarded.example/v1", db_id=2, models=[forwarded_expensive] "forwarded",
"https://forwarded.example/v1",
db_id=2,
models=[forwarded_expensive],
) )
_, provider_map, unique_models = create_model_mappings( _, provider_map, unique_models = create_model_mappings(
+22 -4
View File
@@ -11,7 +11,12 @@ import pytest
import respx import respx
from routstr.payment.cost_calculation import CostData, CostDataError from routstr.payment.cost_calculation import CostData, CostDataError
from routstr.upstream.certification import STATUS_FAIL, STATUS_OK, STATUS_WARN from routstr.upstream.certification import (
PROBE_MAX_TOKENS,
STATUS_FAIL,
STATUS_OK,
STATUS_WARN,
)
from routstr.upstream.certification_cache import ( from routstr.upstream.certification_cache import (
CacheProbeResult, CacheProbeResult,
_raw_cache_keys, _raw_cache_keys,
@@ -210,6 +215,8 @@ class TestCacheBillingRow:
model=_model(), model=_model(),
probe=_probe([_payload(UNCACHED), cached]), probe=_probe([_payload(UNCACHED), cached]),
cost_data=_cost(100), cost_data=_cost(100),
provider_fee=1.0,
sats_to_usd=SATS_USD,
) )
assert row["status"] == STATUS_OK assert row["status"] == STATUS_OK
assert row["evidence"]["reported_usd"] == 5e-5 assert row["evidence"]["reported_usd"] == 5e-5
@@ -242,12 +249,21 @@ class TestCostMarginRow:
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 1e-3}) payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 1e-3})
row = self._row([payload]) row = self._row([payload])
assert row["status"] == STATUS_FAIL assert row["status"] == STATUS_FAIL
assert "lose money" in row["detail"] assert "below raw upstream cost" in row["detail"]
assert "does not measure client debits" in row["detail"]
def test_fee_scales_upstream_cost(self) -> None: def test_fee_scales_upstream_cost(self) -> None:
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7})
assert self._row([payload], fee=1.0)["status"] == STATUS_OK assert self._row([payload], fee=1.0)["status"] == STATUS_OK
assert self._row([payload], fee=2.0)["status"] == STATUS_FAIL marked_up = self._row([payload], fee=2.0)
assert marked_up["status"] == STATUS_FAIL
assert "Raw upstream cost is covered" in marked_up["detail"]
assert "lose money" not in marked_up["detail"]
sample = marked_up["evidence"]["samples"][0]
assert sample["upstream_msats"] == 2
assert sample["upstream_msats_with_fee"] == 4
assert sample["token_estimate_minus_upstream_msats"] == 0
assert sample["token_estimate_minus_fee_target_msats"] == -2
def test_deepseek_endpoint_price_exceeds_configured_model_price(self) -> None: def test_deepseek_endpoint_price_exceeds_configured_model_price(self) -> None:
sats_usd = 0.0008616302499999999 sats_usd = 0.0008616302499999999
@@ -332,6 +348,8 @@ class TestCostMarginRow:
) )
assert row["status"] == STATUS_OK assert row["status"] == STATUS_OK
assert row["title"] == "Token estimate covers fee-adjusted target"
assert "raw upstream cost exceeds" in row["detail"]
def test_warn_when_pricing_unknown(self) -> None: def test_warn_when_pricing_unknown(self) -> None:
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7})
@@ -373,7 +391,7 @@ class TestProbeCache:
assert first == second assert first == second
system = first["messages"][0]["content"] system = first["messages"][0]["content"]
assert system[0]["cache_control"] == {"type": "ephemeral"} assert system[0]["cache_control"] == {"type": "ephemeral"}
assert first["max_tokens"] == 1 assert first["max_tokens"] == PROBE_MAX_TOKENS
assert route.calls[0].request.headers["Authorization"] == "Bearer k" assert route.calls[0].request.headers["Authorization"] == "Bearer k"
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -0,0 +1,419 @@
"""Regressions from the independent three-pass certification audit."""
import json
from typing import Any
import httpx
import pytest
from routstr.payment import price
from routstr.payment.cost_calculation import CostData
from routstr.upstream.certification import (
ProbeResult,
_model_from_usd_pricing,
build_checklist,
cost_prompt_completion_row,
run_live_checks,
usage_capture_row,
)
from routstr.upstream.certification_cache import (
CacheProbeResult,
cache_billing_row,
cache_reported_row,
probe_cache,
)
from routstr.upstream.certification_probe import (
PROBE_TOKEN_BUDGETS,
next_probe_budget,
)
SATS_USD = 0.0005
URL = "https://mock.example/v1"
CACHED = {
"prompt_tokens": 50,
"completion_tokens": 0,
"prompt_tokens_details": {"cached_tokens": 45},
"cost": 5e-5,
}
EXHAUSTED = {
"error": {
"message": "Could not finish the message because max_tokens "
"or model output limit was reached. Please try again with higher max_tokens."
}
}
def model(name: str = "test-model") -> Any:
return _model_from_usd_pricing(name, 1e-6, 2e-6, SATS_USD)
def cost(total: int) -> CostData:
return CostData(base_msats=0, input_msats=total, output_msats=0, total_msats=total)
def cached_probe(usage: dict[str, Any]) -> CacheProbeResult:
return CacheProbeResult(
chat_url=f"{URL}/chat/completions",
statuses=[200, 200],
payloads=[{"usage": usage}, {"usage": usage}],
)
def rows_by_id(rows: list[dict[str, Any]]) -> dict[str, Any]:
return {row["id"]: row for row in rows}
@pytest.fixture(autouse=True)
def local_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
from routstr.core.settings import settings
monkeypatch.setattr(price, "SATS_USD_PRICE", SATS_USD)
monkeypatch.setattr(settings, "fixed_pricing", False)
@pytest.mark.parametrize("wrong", [0, 1, 999999])
def test_cached_usd_rejects_wrong_calculation(wrong: int) -> None:
row = cache_billing_row(
model=model(),
probe=cached_probe(CACHED),
cost_data=cost(wrong),
provider_fee=1.25,
sats_to_usd=SATS_USD,
)
assert row["status"] == "fail"
assert row["evidence"]["basis"] == "upstream_reported_usd"
assert row["evidence"]["expected_total_msats"] == 125
def test_cached_usd_verifies_fee_without_claiming_discount() -> None:
row = cache_billing_row(
model=model(),
probe=cached_probe(CACHED),
cost_data=cost(125),
provider_fee=1.25,
sats_to_usd=SATS_USD,
)
assert row["status"] == "ok"
assert row["evidence"]["expected_total_msats"] == 125
assert row["evidence"]["upstream_discount_verified"] is False
assert "discount" not in row["title"].lower()
assert "not verified" in row["detail"]
def test_equal_full_price_does_not_certify_discount() -> None:
row = cache_billing_row(
model=model(),
probe=cached_probe(CACHED),
cost_data=cost(100),
sats_to_usd=SATS_USD,
)
assert row["status"] == "ok"
assert (
row["evidence"]["actual_total_msats"]
== row["evidence"]["full_price_total_msats"]
)
assert row["evidence"]["upstream_discount_verified"] is False
assert "already carries" not in row["detail"]
assert "cache rate" not in build_checklist([row])[4]["label"]
@pytest.mark.parametrize(
"rate,status", [(None, "warn"), (0.0, "fail"), (float("nan"), "fail")]
)
def test_cached_usd_requires_valid_exchange_rate(
rate: float | None, status: str
) -> None:
row = cache_billing_row(
model=model(), probe=cached_probe(CACHED), cost_data=cost(100), sats_to_usd=rate
)
assert row["status"] == status
@pytest.mark.parametrize("free", [False, True])
def test_unknown_cost_does_not_claim_upstream_overpayment(free: bool) -> None:
usage = {key: value for key, value in CACHED.items() if key != "cost"}
priced = _model_from_usd_pricing("test-model", 0 if free else 1e-6, 0, SATS_USD)
row = cache_billing_row(
model=priced, probe=cached_probe(usage), cost_data=cost(0 if free else 100)
)
assert row["status"] == "warn"
assert "pay more" not in row["detail"]
assert "unverified" in row["detail"]
def test_missing_usage_does_not_claim_free_settlement() -> None:
probe = ProbeResult(
base_url=URL,
models_url=f"{URL}/models",
chat_url=f"{URL}/chat/completions",
chat_status=200,
chat_payload={"choices": [{"message": {"content": "ok"}}]},
)
row = usage_capture_row(probe)
assert row["status"] == "warn"
assert "estimate" in row["detail"]
assert "(0+0)" not in row["detail"]
def test_short_usd_basis_is_not_called_configured_pricing() -> None:
probe = ProbeResult(
base_url=URL,
models_url=f"{URL}/models",
chat_url=f"{URL}/chat/completions",
chat_status=200,
chat_payload={"usage": CACHED},
)
row = cost_prompt_completion_row(
model=model(),
probe=probe,
cost_data=cost(100),
provider_fee=1,
sats_to_usd=SATS_USD,
)
assert row["status"] == "ok"
assert "upstream-reported USD" in row["detail"]
assert "matching the configured" not in row["detail"]
def test_no_evaluated_models_warns_all_pricing_rows() -> None:
from routstr.core.admin import (
_report_row_cache_rate,
_report_row_enabled_models_served,
_report_row_sats_pricing_present,
_report_row_served_matches_configured,
)
rows = [
builder([])
for builder in (
_report_row_cache_rate,
_report_row_enabled_models_served,
_report_row_sats_pricing_present,
_report_row_served_matches_configured,
)
]
assert all(
row["status"] == "warn" and row["evidence"]["checked"] == 0 for row in rows
)
assert all("Not evaluated" in row["detail"] for row in rows)
assert (
next(c for c in build_checklist(rows) if c["goal"] == "pricing_v1_models")[
"status"
]
== "warn"
)
@pytest.mark.asyncio
async def test_typesafe_native_listing_without_unsupported_chat() -> None:
from routstr.upstream.typesafe import TypeSafeUpstreamProvider
calls = []
def handle(request: httpx.Request) -> httpx.Response:
calls.append(request)
assert request.method == "GET"
return httpx.Response(200, json={"models": [{"name": "jev-1.13"}]})
upstream = TypeSafeUpstreamProvider(api_key="test-key")
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
rows = rows_by_id(
await run_live_checks(
upstream.base_url,
"test-key",
model("jev-1.13.0"),
provider_fee=1,
sats_to_usd=SATS_USD,
upstream=upstream,
client=client,
)
)
assert len(calls) == 1
assert rows["endpoint.models_payload"]["status"] == "ok"
assert rows["endpoint.models_payload"]["evidence"]["sample_ids"] == ["jev-1.13"]
for key in (
"usage.capture",
"cost.prompt_completion",
"cache.reported",
"cache.billing",
"cost.margin",
):
assert rows[key]["status"] == "warn"
assert "unverified" in rows[key]["detail"]
@pytest.mark.asyncio
async def test_corrected_field_and_budget_propagate_to_cache() -> None:
bodies = []
def handle(request: httpx.Request) -> httpx.Response:
if request.method == "GET":
return httpx.Response(200, json={"data": [{"id": "test-model"}]})
body = json.loads(request.content)
bodies.append(body)
if "max_tokens" in body:
return httpx.Response(
400,
json={
"error": {
"message": "Unsupported parameter: 'max_tokens'. Use 'max_completion_tokens' instead."
}
},
)
if body["max_completion_tokens"] < 128:
return httpx.Response(
400,
json={
"error": {
"message": "max_completion_tokens must be at least 128 for this deployment."
}
},
)
return httpx.Response(200, json={"usage": CACHED})
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
rows = rows_by_id(
await run_live_checks(
URL, "", model(), provider_fee=1, sats_to_usd=SATS_USD, client=client
)
)
assert rows["usage.capture"]["status"] == "ok"
assert len(bodies) == 5
assert [body.get("max_completion_tokens") for body in bodies] == [
None,
32,
128,
128,
128,
]
evidence = rows["usage.capture"]["evidence"]
assert evidence["token_limit_field"] == "max_completion_tokens"
assert evidence["token_limit"] == 128
assert [a["status_code"] for a in evidence["attempts"]] == [400, 400, 200]
assert rows["cache.reported"]["evidence"]["token_limit"] == 128
@pytest.mark.asyncio
async def test_openai_provider_shaped_field_is_recorded(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from routstr.upstream.openai import OpenAIUpstreamProvider
monkeypatch.setattr("routstr.upstream.openai._rejects_max_tokens", lambda _: True)
def handle(request: httpx.Request) -> httpx.Response:
if request.method == "GET":
return httpx.Response(200, json={"data": [{"id": "gpt-5"}]})
body = json.loads(request.content)
assert body["max_completion_tokens"] == 32
assert "max_tokens" not in body
return httpx.Response(200, json={"usage": CACHED})
upstream = OpenAIUpstreamProvider(api_key="test-key")
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
rows = rows_by_id(
await run_live_checks(
upstream.base_url,
"test-key",
model("gpt-5"),
provider_fee=1,
sats_to_usd=SATS_USD,
upstream=upstream,
client=client,
)
)
assert (
rows["usage.capture"]["evidence"]["token_limit_field"]
== "max_completion_tokens"
)
@pytest.mark.asyncio
async def test_repeated_exhaustion_stops_at_ceiling_and_preserves_failure() -> None:
budgets = []
def handle(request: httpx.Request) -> httpx.Response:
if request.method == "GET":
return httpx.Response(200, json={"data": [{"id": "test-model"}]})
budgets.append(json.loads(request.content)["max_tokens"])
return httpx.Response(400, json=EXHAUSTED)
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
rows = rows_by_id(
await run_live_checks(
URL, "", model(), provider_fee=1, sats_to_usd=SATS_USD, client=client
)
)
assert budgets == list(PROBE_TOKEN_BUDGETS)
assert rows["usage.capture"]["status"] == "fail"
assert len(rows["usage.capture"]["evidence"]["attempts"]) == len(
PROBE_TOKEN_BUDGETS
)
assert rows["cache.reported"]["status"] == "warn"
@pytest.mark.parametrize(
"status,message,current,expected",
[
(400, "max_tokens must be at least 16 for this deployment", 1, 32),
(400, "max_tokens must be at least 2 for this deployment", 1, 32),
(400, "max_tokens must be at least 9000", 32, None),
(400, "context window exceeded", 32, None),
(400, "unknown model", 32, None),
(429, "increase max_tokens", 32, None),
],
)
def test_budget_corrections_are_narrow_and_bounded(
status: int, message: str, current: int, expected: int | None
) -> None:
assert (
next_probe_budget(
status,
{"error": {"metadata": {"raw": json.dumps({"message": message})}}},
current,
)
== expected
)
@pytest.mark.asyncio
async def test_cache_retains_fallback_and_http_error_without_credentials() -> None:
replies = iter(
[
httpx.Response(
400, json={"error": "cache_control not allowed; key=private-test-key"}
),
httpx.Response(200, json={"usage": CACHED}),
httpx.Response(429, json={"error": "rate limited; key=private-test-key"}),
]
)
async with httpx.AsyncClient(
transport=httpx.MockTransport(lambda _: next(replies))
) as client:
probe = await probe_cache(URL, "private-test-key", "test-model", client=client)
row = cache_reported_row(probe)
assert probe.statuses == [200, 429]
assert [a["status_code"] for a in row["evidence"]["attempts"]] == [400, 200, 429]
assert [a["request_format"] for a in probe.attempts] == [
"cache_control",
"plain",
"plain",
]
assert "cache_control not allowed" in probe.attempts[0]["body"]
assert "HTTP 429" in row["detail"] and "rate limited" in row["detail"]
assert "no response body" not in row["detail"]
assert "private-test-key" not in json.dumps(row)
@pytest.mark.asyncio
async def test_cache_decode_failure_has_bounded_body() -> None:
async with httpx.AsyncClient(
transport=httpx.MockTransport(
lambda _: httpx.Response(200, text="not JSON " + "x" * 3000)
)
) as client:
probe = await probe_cache(URL, "", "test-model", client=client)
row = cache_reported_row(probe)
assert row["status"] == "fail"
assert "JSONDecodeError" in row["detail"]
assert all(len(a["body"]) <= 1001 for a in probe.attempts)
+4 -1
View File
@@ -13,6 +13,7 @@ import httpx
import pytest import pytest
from routstr.upstream.certification import ( from routstr.upstream.certification import (
PROBE_MAX_TOKENS,
STATUS_OK, STATUS_OK,
certify_upstream_url, certify_upstream_url,
wants_max_completion_tokens, wants_max_completion_tokens,
@@ -79,7 +80,9 @@ async def test_probe_retries_with_max_completion_tokens() -> None:
# Rejected probe, retried probe, then both cache-probe calls reuse the # Rejected probe, retried probe, then both cache-probe calls reuse the
# accepted field instead of being rejected again. # accepted field instead of being rejected again.
assert ["max_tokens" in body for body in bodies] == [True, False, False, False] assert ["max_tokens" in body for body in bodies] == [True, False, False, False]
assert all(body.get("max_completion_tokens") == 1 for body in bodies[1:]) assert all(
body.get("max_completion_tokens") == PROBE_MAX_TOKENS for body in bodies[1:]
)
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -0,0 +1,131 @@
import json
from contextlib import asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from routstr.core.db import ModelRow
from routstr.payment.models import Pricing, _row_to_model, list_models
from routstr.upstream.helpers import get_all_models_with_overrides
from routstr.upstream.model_paths import _price_in_sats
def _row(cache_rate: float = 0) -> ModelRow:
return ModelRow(
id="vendor/model",
name="Model",
created=0,
description="",
context_length=8192,
architecture=json.dumps(
{
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "unknown",
}
),
pricing=json.dumps(
{"prompt": 4e-6, "completion": 8e-6, "input_cache_read": cache_rate}
),
enabled=True,
upstream_provider_id=1,
)
def _session(row: ModelRow, provider_type: str):
provider = SimpleNamespace(
id=1, enabled=True, provider_fee=1.1, provider_type=provider_type
)
return SimpleNamespace(
exec=AsyncMock(
side_effect=[
SimpleNamespace(all=lambda: [row]),
SimpleNamespace(all=lambda: [provider]),
]
)
)
@pytest.mark.parametrize("provider_type", ["ppqai", "venice"])
@pytest.mark.parametrize("cache_rate", [0, 7e-7])
def test_db_conversion_preserves_native_or_explicit_cache_prices(
provider_type, cache_rate
) -> None:
row = _row(cache_rate)
stored = row.pricing
with (
patch(
"routstr.payment.models.backfill_cache_pricing",
return_value=Pricing(prompt=4e-6, completion=8e-6, input_cache_read=9e-7),
) as backfill,
patch("routstr.payment.models.sats_usd_price", return_value=0.001),
):
model = _row_to_model(row, True, 1.1, provider_type=provider_type)
backfill.assert_not_called()
assert model.pricing.input_cache_read == pytest.approx(cache_rate * 1.1)
assert row.pricing == stored
@pytest.mark.asyncio
@pytest.mark.parametrize("provider_type", ["ppqai", "venice"])
async def test_admin_listing_passes_provider_policy_to_database_conversion(
provider_type,
) -> None:
with (
patch("routstr.payment.models.backfill_cache_pricing") as backfill,
patch("routstr.payment.models.sats_usd_price", return_value=0.001),
):
models = await list_models(_session(_row(), provider_type), upstream_id=1)
backfill.assert_not_called()
assert len(models) == 1
assert models[0].pricing.input_cache_read == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("provider_type", ["ppqai", "venice"])
async def test_runtime_overrides_do_not_reintroduce_generic_cache_prices(
provider_type,
) -> None:
row = _row()
session = _session(row, provider_type)
@asynccontextmanager
async def create_session():
yield session
upstream = SimpleNamespace(
db_id=1,
provider_type=provider_type,
base_url="https://example.invalid",
get_cached_models=MagicMock(
return_value=[SimpleNamespace(id=row.id, enabled=True)]
),
)
with (
patch("routstr.upstream.helpers.create_session", create_session),
patch("routstr.payment.models.backfill_cache_pricing") as backfill,
patch("routstr.payment.models.sats_usd_price", return_value=0.001),
):
models = await get_all_models_with_overrides([upstream])
backfill.assert_not_called()
assert len(models) == 1
assert models[0].pricing.input_cache_read == 0
@pytest.mark.parametrize("provider_type", ["ppqai", "venice"])
def test_path_sats_conversion_respects_native_cache_policy(provider_type) -> None:
model = {
"id": "vendor/model",
"pricing": {"prompt": 4e-6, "completion": 8e-6},
"context_length": 8192,
}
with (
patch("routstr.payment.models.backfill_cache_pricing") as backfill,
patch("routstr.payment.price.sats_usd_price", return_value=0.001),
):
_price_in_sats(model, 1.1, provider_type)
backfill.assert_not_called()
assert model["pricing"]["input_cache_read"] == 0
assert model["sats_pricing"]["input_cache_read"] == 0
+138
View File
@@ -0,0 +1,138 @@
import json
from unittest.mock import AsyncMock, patch
import httpx
import pytest
from routstr.payment.models import Architecture, Model, Pricing
from routstr.upstream.ppqai import PPQAIUpstreamProvider
from routstr.upstream.venice import VeniceUpstreamProvider
def _model() -> Model:
return Model(
id="vendor/model",
name="Model",
created=0,
description="Metadata",
context_length=8192,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Unknown",
),
pricing=Pricing(
prompt=1e-6,
completion=2e-6,
input_cache_read=1e-8,
input_cache_write=3e-6,
request=0.1,
),
)
def _entry(model_id: str, pricing: dict) -> dict:
return {
"id": model_id,
"name": model_id,
"created_at": 0,
"context_length": 8192,
"pricing": pricing,
}
async def _fetch(
entries: list[dict], metadata: list[dict] | None = None
) -> list[Model]:
response = httpx.Response(
200,
request=httpx.Request("GET", "https://api.ppq.ai/models"),
text=json.dumps({"data": entries}),
)
with (
patch(
"routstr.upstream.ppqai._safe_read_request",
AsyncMock(return_value=response),
),
patch(
"routstr.upstream.ppqai.async_fetch_openrouter_models",
AsyncMock(
return_value=metadata if metadata is not None else [_model().dict()]
),
),
):
return await PPQAIUpstreamProvider("test-only").fetch_models()
@pytest.mark.asyncio
async def test_ppq_does_not_inherit_other_provider_cache_or_request_rates() -> None:
(model,) = await _fetch(
[_entry("vendor/model", {"api": {"input_per_1M": 4, "output_per_1M": 8}})]
)
assert model.pricing == Pricing(prompt=4e-6, completion=8e-6)
@pytest.mark.asyncio
async def test_ppq_alias_matches_do_not_share_mutated_prices() -> None:
models = await _fetch(
[
_entry("vendor/model", {"api": {"input_per_1M": 4, "output_per_1M": 8}}),
_entry("model", {"api": {"input_per_1M": 6, "output_per_1M": 9}}),
]
)
assert [m.id for m in models] == ["vendor/model", "model"]
assert [m.pricing.prompt for m in models] == [4e-6, 6e-6]
assert models[0] is not models[1]
@pytest.mark.asyncio
@pytest.mark.parametrize("metadata", [None, []])
async def test_ppq_partial_api_prices_fall_back_per_field_preserving_zero(
metadata,
) -> None:
(model,) = await _fetch(
[
_entry(
"vendor/model",
{
"api": {"input_per_1M": 0},
"input_per_1M_tokens": 5,
"output_per_1M_tokens": 8,
},
)
],
metadata,
)
assert model.pricing.prompt == 0
assert model.pricing.completion == 8e-6
@pytest.mark.asyncio
@pytest.mark.parametrize("rate", [None, -1, float("inf"), float("nan")])
async def test_ppq_unpriced_or_invalid_native_rate_is_not_replaced_by_openrouter(
rate,
) -> None:
assert (
await _fetch(
[
_entry(
"vendor/model", {"api": {"input_per_1M": rate, "output_per_1M": 8}}
)
]
)
== []
)
@pytest.mark.parametrize("provider", [PPQAIUpstreamProvider, VeniceUpstreamProvider])
def test_native_catalog_providers_do_not_backfill_generic_cache_rates(provider) -> None:
model = _model()
model.pricing = Pricing(prompt=4e-6, completion=8e-6)
with patch("routstr.upstream.base.backfill_cache_pricing") as backfill:
adjusted = provider("test-only", provider_fee=1.1)._apply_provider_fee_to_model(
model
)
backfill.assert_not_called()
assert adjusted.pricing.input_cache_read == 0
assert adjusted.pricing.prompt == pytest.approx(4.4e-6)