mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix: harden provider certification and native pricing
This commit is contained in:
@@ -239,7 +239,10 @@ def create_model_mappings(
|
||||
override_row, provider_fee = overrides_by_key[model_key]
|
||||
try:
|
||||
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:
|
||||
# Stored pricing is JSON from whatever wrote the row, so
|
||||
@@ -315,7 +318,10 @@ def create_model_mappings(
|
||||
|
||||
try:
|
||||
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:
|
||||
logger.warning(
|
||||
|
||||
+15
-4
@@ -617,7 +617,10 @@ async def upsert_provider_model(
|
||||
await refresh_model_maps()
|
||||
await _refresh_provider_model_paths(provider_pk)
|
||||
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
|
||||
|
||||
|
||||
@@ -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`.
|
||||
return json_compliant( # type: ignore[return-value]
|
||||
_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()
|
||||
)
|
||||
|
||||
@@ -1311,7 +1317,10 @@ def _evaluate_model_row(
|
||||
) -> _ModelEvaluation:
|
||||
try:
|
||||
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
|
||||
except Exception as exc:
|
||||
@@ -1352,7 +1361,9 @@ def _aggregate_row(
|
||||
"""
|
||||
evidence: dict[str, object] = {"checked": checked, "flagged": list(flagged)}
|
||||
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:
|
||||
return _report_row(row_id, fail_status, title, flagged_detail, evidence)
|
||||
return _report_row(row_id, "ok", title, ok_detail, evidence)
|
||||
|
||||
@@ -316,8 +316,16 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis
|
||||
return []
|
||||
|
||||
|
||||
def allows_cache_pricing_backfill(provider_type: str | None) -> bool:
|
||||
return provider_type not in {"ppqai", "venice"}
|
||||
|
||||
|
||||
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:
|
||||
"""The deterministic USD view of a stored model row, before the sats conversion."""
|
||||
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
|
||||
# and miss the cache rate.
|
||||
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:
|
||||
parsed_pricing = Pricing.parse_obj(
|
||||
@@ -383,9 +392,15 @@ def _build_model_from_row(
|
||||
|
||||
|
||||
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 = _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:
|
||||
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
|
||||
if r.upstream_provider_id in providers_by_id
|
||||
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:
|
||||
# Stored pricing/architecture is JSON from whatever wrote the row, so
|
||||
|
||||
@@ -5699,12 +5699,17 @@ class BaseUpstreamProvider:
|
||||
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:
|
||||
"""Apply provider fee to model's USD pricing and calculate max costs.
|
||||
|
||||
Cache rates missing from the upstream pricing feed are backfilled from
|
||||
litellm's cost map first, so they carry the provider fee like every
|
||||
other price component.
|
||||
Providers with native price catalogs can disable generic cache-rate
|
||||
backfill to avoid treating another provider's rates as their own.
|
||||
|
||||
Args:
|
||||
model: Model object to update
|
||||
@@ -5712,7 +5717,11 @@ class BaseUpstreamProvider:
|
||||
Returns:
|
||||
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(
|
||||
{k: v * self.provider_fee for k, v in base_pricing.dict().items()}
|
||||
)
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
"""Live certification checks for an upstream provider endpoint.
|
||||
|
||||
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
|
||||
billing path — no reservation, no Cashu. Upstream spend is one one-token
|
||||
completion, plus two or three one-token completions on a ~4.4k-token prompt
|
||||
when the cache checks are enabled (per certified model path).
|
||||
billing path — no reservation, no Cashu. Output budgets start at 32 tokens
|
||||
and rise only on recognized limit rejections, up to 2048. Cache checks add
|
||||
repeated completions on a ~4.4k-token prompt (per certified model path).
|
||||
They sit behind ``POST …/certify`` rather than the read-only ``GET …/report``
|
||||
because they can block for the length of the timeout.
|
||||
"""
|
||||
@@ -21,7 +21,7 @@ import os
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any
|
||||
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.rates import coerce_rate
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -45,12 +53,7 @@ STATUS_FAIL = "fail"
|
||||
|
||||
TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"}
|
||||
|
||||
# Bounded so a dead upstream fails the row rather than wedging the request.
|
||||
PROBE_TIMEOUT_SECONDS = 15.0
|
||||
|
||||
# The cheapest request that still exercises the usage/cost path.
|
||||
PROBE_MAX_TOKENS = 1
|
||||
PROBE_PROMPT = "ping"
|
||||
PROBE_PROMPT = "Reply with the single word: ok"
|
||||
|
||||
# ``calculate_cost`` demands a reservation ceiling; any value at or above the
|
||||
# real charge behaves identically.
|
||||
@@ -128,12 +131,12 @@ CHECKLIST_GOALS: tuple[tuple[str, str, tuple[str, ...]], ...] = (
|
||||
),
|
||||
(
|
||||
"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"),
|
||||
),
|
||||
(
|
||||
"margin",
|
||||
"Margin — node charge covers the upstream's cost",
|
||||
"Pricing target — token estimate covers the fee-adjusted target",
|
||||
("cost.margin",),
|
||||
),
|
||||
)
|
||||
@@ -166,7 +169,7 @@ def build_checklist(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
|
||||
@dataclass
|
||||
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
|
||||
models_url: str
|
||||
@@ -180,19 +183,11 @@ class ProbeResult:
|
||||
chat_payload: dict[str, Any] | None = None
|
||||
chat_error: str | None = None
|
||||
chat_latency_ms: float | None = None
|
||||
# ``max_completion_tokens`` once the upstream rejected ``max_tokens``.
|
||||
token_limit_field: str = "max_tokens"
|
||||
|
||||
|
||||
def wants_max_completion_tokens(status: int | None, payload: Any) -> bool:
|
||||
"""Whether a 400 names ``max_completion_tokens`` as the field to use.
|
||||
|
||||
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)
|
||||
token_limit: int = PROBE_MAX_TOKENS
|
||||
chat_attempts: list[dict[str, Any]] = field(default_factory=list)
|
||||
provider_type: str | None = None
|
||||
completion_skip_reason: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -269,12 +264,11 @@ async def probe_upstream(
|
||||
upstream: "BaseUpstreamProvider | None" = None,
|
||||
model: "Model | None" = None,
|
||||
) -> 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
|
||||
retried once with that field (OpenAI o-series, gpt-5). Each HTTP call,
|
||||
including its body read, has an elapsed-time deadline. A transport
|
||||
failure is a ``fail`` row, not a failed admin request.
|
||||
Each HTTP call, including its body read, has an elapsed-time deadline.
|
||||
Transport failures are rows, not failed admin requests. Native System One
|
||||
inference is deliberately unverified until a supported fixture exists.
|
||||
"""
|
||||
shape = probe_shape(base_url, api_key, upstream, model)
|
||||
result = ProbeResult(
|
||||
@@ -282,7 +276,14 @@ async def probe_upstream(
|
||||
models_url=shape.models_url,
|
||||
chat_url=shape.chat_url,
|
||||
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
|
||||
|
||||
owns_client = client is None
|
||||
@@ -313,22 +314,11 @@ async def probe_upstream(
|
||||
result.models_error = f"{type(exc).__name__}: {exc}"
|
||||
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
|
||||
await _probe_chat(
|
||||
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:
|
||||
if owns_client:
|
||||
await client.aclose()
|
||||
@@ -346,7 +336,7 @@ async def _probe_chat(
|
||||
model: "Model | None",
|
||||
token_field: str,
|
||||
) -> 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] = {
|
||||
"model": model_id,
|
||||
"messages": [{"role": "user", "content": PROBE_PROMPT}],
|
||||
@@ -358,35 +348,21 @@ async def _probe_chat(
|
||||
"order": [result.endpoint_tag],
|
||||
"allow_fallbacks": False,
|
||||
}
|
||||
result.token_limit_field = token_field
|
||||
result.chat_status = None
|
||||
result.chat_payload = None
|
||||
result.chat_error = None
|
||||
started = time.monotonic()
|
||||
try:
|
||||
async with asyncio.timeout(timeout):
|
||||
response = await client.post(
|
||||
result.chat_url,
|
||||
json=shape_body(request_body, upstream, model),
|
||||
headers=shape.headers,
|
||||
params=shape.chat_params,
|
||||
)
|
||||
result.chat_status = response.status_code
|
||||
result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2)
|
||||
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)
|
||||
outcome = await send_completion(
|
||||
client,
|
||||
result.chat_url,
|
||||
shape_body(request_body, upstream, model),
|
||||
shape.headers,
|
||||
shape.chat_params,
|
||||
timeout,
|
||||
)
|
||||
result.token_limit_field = outcome.token_limit_field
|
||||
result.token_limit = outcome.token_limit
|
||||
result.chat_status = outcome.status
|
||||
result.chat_payload = outcome.payload
|
||||
result.chat_error = outcome.error
|
||||
result.chat_latency_ms = outcome.latency_ms
|
||||
result.chat_attempts = outcome.attempts
|
||||
|
||||
|
||||
# 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},
|
||||
)
|
||||
|
||||
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):
|
||||
return certification_row(
|
||||
"endpoint.models_payload",
|
||||
STATUS_FAIL,
|
||||
"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,
|
||||
"top_level_keys": sorted(payload.keys()),
|
||||
@@ -487,9 +466,9 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]:
|
||||
)
|
||||
|
||||
ids = [
|
||||
item["id"]
|
||||
item[id_key]
|
||||
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] = {
|
||||
"url": probe.models_url,
|
||||
@@ -502,15 +481,15 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]:
|
||||
"endpoint.models_payload",
|
||||
STATUS_FAIL,
|
||||
"Models payload has the expected shape",
|
||||
f'The "data" list carries no entry with a non-empty string "id" '
|
||||
f"({len(data)} entries).",
|
||||
f'The "{data_key}" list carries no entry with a non-empty string '
|
||||
f'"{id_key}" ({len(data)} entries).',
|
||||
evidence,
|
||||
)
|
||||
return certification_row(
|
||||
"endpoint.models_payload",
|
||||
STATUS_OK,
|
||||
"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,
|
||||
)
|
||||
|
||||
@@ -518,14 +497,25 @@ def models_payload_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.
|
||||
|
||||
A missing ``usage`` object means the node has nothing to price and the
|
||||
request settles for free. Broken, but still usable, so ``warn``.
|
||||
Missing upstream usage is a coverage gap. The proxy can estimate usage
|
||||
before settlement; this direct probe does not exercise that fallback.
|
||||
"""
|
||||
evidence: dict[str, Any] = {
|
||||
"url": probe.chat_url,
|
||||
"status_code": probe.chat_status,
|
||||
"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:
|
||||
evidence["error"] = probe.chat_error
|
||||
return certification_row(
|
||||
@@ -536,13 +526,17 @@ def usage_capture_row(probe: ProbeResult) -> dict[str, Any]:
|
||||
evidence,
|
||||
)
|
||||
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(
|
||||
"usage.capture",
|
||||
STATUS_FAIL,
|
||||
"Token usage captured from a completion",
|
||||
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,
|
||||
)
|
||||
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",
|
||||
STATUS_WARN,
|
||||
"Token usage captured from a completion",
|
||||
'The completion carried no "usage" object, so the node has no '
|
||||
"token counts to bill on and the request would settle as (0+0).",
|
||||
'The completion carried no "usage" object. The proxy may estimate '
|
||||
"missing usage before settlement; this probe does not verify that "
|
||||
"fallback or the resulting client debit.",
|
||||
evidence,
|
||||
)
|
||||
evidence["input_tokens"] = normalized.input_tokens
|
||||
@@ -782,8 +777,8 @@ def cost_prompt_completion_row(
|
||||
"cost.prompt_completion",
|
||||
STATUS_FAIL,
|
||||
"Prompt and completion cost calculated",
|
||||
f"The expected charge could not be derived from the configured "
|
||||
f"pricing: {type(exc).__name__}: {exc}.",
|
||||
f"The expected charge could not be derived from the selected "
|
||||
f"billing basis: {type(exc).__name__}: {exc}.",
|
||||
evidence,
|
||||
)
|
||||
|
||||
@@ -826,12 +821,17 @@ def cost_prompt_completion_row(
|
||||
):
|
||||
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:
|
||||
return certification_row(
|
||||
"cost.prompt_completion",
|
||||
STATUS_FAIL,
|
||||
"Prompt and completion cost calculated",
|
||||
"The computed charge disagrees with the configured pricing: "
|
||||
f"The computed charge disagrees with {basis_label}: "
|
||||
+ "; ".join(mismatches)
|
||||
+ ".",
|
||||
evidence,
|
||||
@@ -840,10 +840,10 @@ def cost_prompt_completion_row(
|
||||
"cost.prompt_completion",
|
||||
STATUS_OK,
|
||||
"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"{usage.output_tokens} completion tokens, matching the configured "
|
||||
f"pricing.",
|
||||
f"{usage.output_tokens} completion tokens, matching {basis_label}. "
|
||||
"This is a calculation, not a measured client debit.",
|
||||
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
|
||||
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:
|
||||
cost_data = await calculate_cost(
|
||||
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:
|
||||
rows.extend(skipped_cache_rows("Skipped — cache checks disabled."))
|
||||
elif (
|
||||
@@ -973,6 +990,7 @@ async def run_live_checks(
|
||||
endpoint_tag=endpoint_tag,
|
||||
upstream=upstream,
|
||||
token_limit_field=probe.token_limit_field,
|
||||
token_limit=probe.token_limit,
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
"""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 node bill cached reads at the discounted rate (client side), and
|
||||
* does the node's charge cover what the upstream charged (node side).
|
||||
* does the cost engine calculate cached-completion charges correctly, and
|
||||
* 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 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
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
@@ -36,11 +34,13 @@ from .certification import (
|
||||
_fixed_token_pricing_active,
|
||||
_reported_usd_cost,
|
||||
_token_rates,
|
||||
_truncate,
|
||||
certification_row,
|
||||
probe_shape,
|
||||
safe_row,
|
||||
shape_body,
|
||||
)
|
||||
from .certification_probe import CompletionProbe, send_completion
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..payment.models import Model
|
||||
@@ -58,8 +58,8 @@ ROW_BILLING = "cache.billing"
|
||||
ROW_MARGIN = "cost.margin"
|
||||
|
||||
TITLE_REPORTED = "Upstream reports prompt-cache hits"
|
||||
TITLE_BILLING = "Cached tokens billed at the cache-read rate"
|
||||
TITLE_MARGIN = "Node charge covers upstream cost"
|
||||
TITLE_BILLING = "Cached completion cost calculated"
|
||||
TITLE_MARGIN = "Token estimate covers fee-adjusted target"
|
||||
|
||||
|
||||
def cache_probe_prefix() -> str:
|
||||
@@ -86,6 +86,9 @@ class CacheProbeResult:
|
||||
payloads: list[dict[str, Any] | None] = field(default_factory=list)
|
||||
errors: list[str | 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
|
||||
def second_payload(self) -> dict[str, Any] | None:
|
||||
@@ -104,6 +107,7 @@ def _request_body(
|
||||
fmt: str,
|
||||
endpoint_tag: str | None,
|
||||
token_field: str = "max_tokens",
|
||||
token_limit: int = PROBE_MAX_TOKENS,
|
||||
) -> dict[str, Any]:
|
||||
system: Any
|
||||
if fmt == "cache_control":
|
||||
@@ -122,7 +126,7 @@ def _request_body(
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": CACHE_PROBE_QUESTION},
|
||||
],
|
||||
token_field: PROBE_MAX_TOKENS,
|
||||
token_field: token_limit,
|
||||
"stream": False,
|
||||
}
|
||||
if endpoint_tag:
|
||||
@@ -133,45 +137,11 @@ def _request_body(
|
||||
return body
|
||||
|
||||
|
||||
async def _post_completion(
|
||||
client: httpx.AsyncClient,
|
||||
url: str,
|
||||
body: Any,
|
||||
headers: dict[str, str],
|
||||
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 _record(result: CacheProbeResult, outcome: CompletionProbe) -> None:
|
||||
result.statuses.append(outcome.status)
|
||||
result.payloads.append(outcome.payload)
|
||||
result.errors.append(outcome.error)
|
||||
result.latencies_ms.append(outcome.latency_ms)
|
||||
|
||||
|
||||
def _is_2xx(status: int | None) -> bool:
|
||||
@@ -189,6 +159,7 @@ async def probe_cache(
|
||||
upstream: "BaseUpstreamProvider | None" = None,
|
||||
model: "Model | None" = None,
|
||||
token_limit_field: str = "max_tokens",
|
||||
token_limit: int = PROBE_MAX_TOKENS,
|
||||
) -> CacheProbeResult:
|
||||
"""Send the same long prompt twice.
|
||||
|
||||
@@ -198,17 +169,27 @@ async def probe_cache(
|
||||
succeeded. Each call's elapsed deadline includes the response body.
|
||||
"""
|
||||
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()
|
||||
|
||||
owns_client = client is None
|
||||
http = client if client is not None else httpx.AsyncClient(timeout=timeout)
|
||||
|
||||
async def post(
|
||||
fmt: str,
|
||||
) -> tuple[int | None, dict[str, Any] | None, str | None, float]:
|
||||
body = _request_body(model_id, prefix, fmt, endpoint_tag, token_limit_field)
|
||||
return await _post_completion(
|
||||
async def post(fmt: str, phase: str) -> CompletionProbe:
|
||||
body = _request_body(
|
||||
model_id,
|
||||
prefix,
|
||||
fmt,
|
||||
endpoint_tag,
|
||||
result.token_limit_field,
|
||||
result.token_limit,
|
||||
)
|
||||
outcome = await send_completion(
|
||||
http,
|
||||
result.chat_url,
|
||||
shape_body(body, upstream, model),
|
||||
@@ -216,16 +197,23 @@ async def probe_cache(
|
||||
shape.chat_params,
|
||||
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:
|
||||
first = await post("cache_control")
|
||||
if first[0] in (400, 422):
|
||||
first = await post("cache_control", "warm_up")
|
||||
if first.status in (400, 422):
|
||||
result.request_format = "plain"
|
||||
first = await post("plain")
|
||||
first = await post("plain", "warm_up")
|
||||
_record(result, first)
|
||||
if not _is_2xx(first[0]):
|
||||
if not _is_2xx(first.status):
|
||||
return result
|
||||
second = await post(result.request_format)
|
||||
second = await post(result.request_format, "repeated")
|
||||
_record(result, second)
|
||||
finally:
|
||||
if owns_client:
|
||||
@@ -262,16 +250,29 @@ def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]:
|
||||
"endpoint_tag": probe.endpoint_tag,
|
||||
"statuses": probe.statuses,
|
||||
"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
|
||||
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(
|
||||
ROW_REPORTED,
|
||||
STATUS_FAIL,
|
||||
TITLE_REPORTED,
|
||||
f"The cache probe did not get two successful completions: "
|
||||
f"{probe.second_error or 'no response body'}.",
|
||||
f"The cache probe did not get two successful completions: {reason}.",
|
||||
evidence,
|
||||
)
|
||||
|
||||
@@ -308,17 +309,17 @@ def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]:
|
||||
STATUS_FAIL,
|
||||
TITLE_REPORTED,
|
||||
"The upstream reported cache tokens under fields the node does not "
|
||||
f"parse ({', '.join(raw_keys)}); cached reads would be billed at "
|
||||
"the full input rate.",
|
||||
f"parse ({', '.join(raw_keys)}). Token-based pricing cannot apply "
|
||||
"a cache-read rate to those fields; reported-USD billing may take "
|
||||
"precedence.",
|
||||
evidence,
|
||||
)
|
||||
return certification_row(
|
||||
ROW_REPORTED,
|
||||
STATUS_WARN,
|
||||
TITLE_REPORTED,
|
||||
"Two identical prompts produced no cache hit. Either the model does "
|
||||
"not support prompt caching or the upstream hides it; clients pay the "
|
||||
"full input rate on repeated prompts.",
|
||||
"Two identical prompts produced no reported cache hit. Caching may "
|
||||
"be unsupported or unreported; no cache discount is verified.",
|
||||
evidence,
|
||||
)
|
||||
|
||||
@@ -329,6 +330,8 @@ def cache_billing_row(
|
||||
probe: CacheProbeResult,
|
||||
cost_data: Any,
|
||||
pricing_known: bool = True,
|
||||
provider_fee: float = 1.0,
|
||||
sats_to_usd: float | None = None,
|
||||
) -> dict[str, Any]:
|
||||
usage = _usage_of(probe.second_payload)
|
||||
evidence: dict[str, Any] = {"model_id": model.id}
|
||||
@@ -345,7 +348,7 @@ def cache_billing_row(
|
||||
ROW_BILLING,
|
||||
STATUS_WARN,
|
||||
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.",
|
||||
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:
|
||||
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(
|
||||
ROW_BILLING,
|
||||
STATUS_OK,
|
||||
TITLE_BILLING,
|
||||
f"Billed {actual_total} msats from the upstream-reported cost, "
|
||||
f"which already carries the cache discount "
|
||||
f"(full token price would be {full_total} msats).",
|
||||
f"Calculated {actual_total} msats, matching the independently "
|
||||
"converted upstream-reported USD with the provider fee. "
|
||||
f"The configured full-input comparator is {full_total} msats; "
|
||||
"an upstream cache discount and actual client debit are not verified.",
|
||||
evidence,
|
||||
)
|
||||
if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS:
|
||||
@@ -410,7 +458,7 @@ def cache_billing_row(
|
||||
ROW_BILLING,
|
||||
STATUS_FAIL,
|
||||
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.",
|
||||
evidence,
|
||||
)
|
||||
@@ -424,16 +472,16 @@ def cache_billing_row(
|
||||
ROW_BILLING,
|
||||
STATUS_WARN,
|
||||
TITLE_BILLING,
|
||||
f"Cached reads are billed at the full input rate ({actual_total} "
|
||||
f"msats) because {reason}; clients pay more than the upstream "
|
||||
"charges.",
|
||||
f"Cached reads use the full input rate ({actual_total} calculated "
|
||||
f"msats) because {reason}. Upstream cost, any upstream discount "
|
||||
"and actual client debit are unverified.",
|
||||
evidence,
|
||||
)
|
||||
return certification_row(
|
||||
ROW_BILLING,
|
||||
STATUS_OK,
|
||||
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.",
|
||||
evidence,
|
||||
)
|
||||
@@ -449,12 +497,9 @@ def cost_margin_row(
|
||||
) -> dict[str, Any]:
|
||||
"""Configured token pricing must cover what the upstream reports charging.
|
||||
|
||||
Responses that carry a USD cost are billed from it, so they cannot lose
|
||||
money themselves; they are used here as a price sample. The configured
|
||||
token pricing is what every other path bills from (streams, upstreams
|
||||
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.
|
||||
Reported USD takes precedence in the cost engine; the token-price estimate
|
||||
is a separate fallback comparison, not a measured debit. A fee-target gap
|
||||
need not be a loss against raw upstream cost. Missing cost stays unverified.
|
||||
|
||||
``model`` carries the pricing the proxy reserves and token-bills with,
|
||||
which on a pinned endpoint is that endpoint's own rates.
|
||||
@@ -477,6 +522,7 @@ def cost_margin_row(
|
||||
|
||||
samples: list[dict[str, Any]] = []
|
||||
short: list[str] = []
|
||||
below_raw: list[str] = []
|
||||
for payload in payloads:
|
||||
if not isinstance(payload, dict):
|
||||
continue
|
||||
@@ -489,6 +535,7 @@ def cost_margin_row(
|
||||
upstream_total = _expected_usd_msats(
|
||||
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:
|
||||
evidence["error"] = f"{type(exc).__name__}: {exc}"
|
||||
return certification_row(
|
||||
@@ -502,11 +549,17 @@ def cost_margin_row(
|
||||
"usage": usage.dict(),
|
||||
"reported_usd": reported_usd,
|
||||
"upstream_msats_with_fee": upstream_total,
|
||||
"upstream_msats": raw_upstream_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)
|
||||
if configured_total + COST_TOLERANCE_MSATS < 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
|
||||
|
||||
if not samples:
|
||||
@@ -524,17 +577,30 @@ def cost_margin_row(
|
||||
ROW_MARGIN,
|
||||
STATUS_FAIL,
|
||||
TITLE_MARGIN,
|
||||
"Configured pricing is below the upstream's reported cost "
|
||||
f"(configured < upstream msats: {'; '.join(short)}); token-billed "
|
||||
"requests lose money.",
|
||||
"The token-pricing estimate is below the fee-adjusted cost target "
|
||||
f"(estimate < target msats: {'; '.join(short)}). "
|
||||
+ (
|
||||
"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,
|
||||
)
|
||||
return certification_row(
|
||||
ROW_MARGIN,
|
||||
STATUS_OK,
|
||||
TITLE_MARGIN,
|
||||
f"Configured pricing covers the upstream's reported cost on "
|
||||
f"{len(samples)} sampled completion(s).",
|
||||
f"Token pricing covers the fee-adjusted cost target on "
|
||||
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,
|
||||
)
|
||||
|
||||
@@ -570,6 +636,7 @@ async def run_cache_checks(
|
||||
endpoint_tag: str | None = None,
|
||||
upstream: "BaseUpstreamProvider | None" = None,
|
||||
token_limit_field: str = "max_tokens",
|
||||
token_limit: int = PROBE_MAX_TOKENS,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Run the cache probe and build the three cache/margin rows."""
|
||||
probe = await probe_cache(
|
||||
@@ -582,8 +649,15 @@ async def run_cache_checks(
|
||||
upstream=upstream,
|
||||
model=model,
|
||||
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 [
|
||||
safe_row(ROW_REPORTED, TITLE_REPORTED, lambda: cache_reported_row(probe)),
|
||||
safe_row(
|
||||
@@ -594,6 +668,8 @@ async def run_cache_checks(
|
||||
probe=probe,
|
||||
cost_data=cost_data,
|
||||
pricing_known=pricing_known,
|
||||
provider_fee=provider_fee,
|
||||
sats_to_usd=sats_to_usd,
|
||||
),
|
||||
),
|
||||
safe_row(
|
||||
@@ -601,7 +677,14 @@ async def run_cache_checks(
|
||||
TITLE_MARGIN,
|
||||
lambda: cost_margin_row(
|
||||
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,
|
||||
sats_to_usd=sats_to_usd,
|
||||
pricing_known=pricing_known,
|
||||
|
||||
@@ -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
|
||||
@@ -123,7 +123,10 @@ async def get_all_models_with_overrides(
|
||||
if model_key is not None and model_key in overrides_by_key:
|
||||
override_row, provider_fee = overrides_by_key[model_key]
|
||||
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:
|
||||
all_models[(model.id.lower(), provider_key)] = model
|
||||
|
||||
@@ -824,7 +824,9 @@ async def refresh_model_paths_periodically(
|
||||
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.
|
||||
|
||||
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,
|
||||
_calculate_usd_max_costs,
|
||||
_update_model_sats_pricing,
|
||||
allows_cache_pricing_backfill,
|
||||
backfill_cache_pricing,
|
||||
)
|
||||
from ..payment.price import sats_usd_price
|
||||
|
||||
try:
|
||||
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()})
|
||||
priced = Model(
|
||||
id=model_id,
|
||||
@@ -921,10 +926,13 @@ def apply_model_path_pricing(
|
||||
metadata.get("pricing"), dict
|
||||
):
|
||||
return model
|
||||
pricing = backfill_cache_pricing(
|
||||
model.forwarded_model_id or row.model_id,
|
||||
Pricing.parse_obj(metadata["pricing"]),
|
||||
)
|
||||
from ..payment.models import allows_cache_pricing_backfill
|
||||
|
||||
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(
|
||||
{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):
|
||||
model = {}
|
||||
model.setdefault("id", row.model_id)
|
||||
_price_in_sats(model, provider_fee)
|
||||
_price_in_sats(model, provider_fee, row.provider_type)
|
||||
return {
|
||||
"path": row.path,
|
||||
"provider": {
|
||||
|
||||
+36
-45
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
import random
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
@@ -226,6 +227,27 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
if ppqai_model.id in self.IGNORED_MODEL_IDS:
|
||||
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(
|
||||
(
|
||||
model
|
||||
@@ -238,44 +260,20 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
)
|
||||
|
||||
if or_model:
|
||||
input_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
input_price = ppqai_model.pricing.api.get("input_per_1M")
|
||||
elif ppqai_model.pricing.input_per_1M_tokens:
|
||||
input_price = ppqai_model.pricing.input_per_1M_tokens
|
||||
|
||||
if input_price is not None:
|
||||
or_model.pricing.prompt = input_price / 1_000_000
|
||||
|
||||
output_price = None
|
||||
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)
|
||||
# OpenRouter supplies metadata, not PPQ billing rates.
|
||||
models.append(
|
||||
or_model.copy(
|
||||
deep=True,
|
||||
update={
|
||||
"id": ppqai_model.id,
|
||||
"pricing": pricing,
|
||||
"sats_pricing": None,
|
||||
"context_length": ppqai_model.context_length
|
||||
or or_model.context_length,
|
||||
},
|
||||
)
|
||||
)
|
||||
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(
|
||||
Model(
|
||||
id=ppqai_model.id,
|
||||
@@ -290,14 +288,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
tokenizer="Unknown",
|
||||
instruct_type=None,
|
||||
),
|
||||
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,
|
||||
),
|
||||
pricing=pricing,
|
||||
)
|
||||
)
|
||||
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
|
||||
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 "unserved-model" not in json.dumps(row["evidence"])
|
||||
|
||||
|
||||
@@ -343,9 +343,7 @@ async def test_certify_margin_bills_pinned_path_pricing(
|
||||
},
|
||||
base_url=base_url,
|
||||
)
|
||||
model_path = encode_model_path(
|
||||
base_url, "cert-test-model", "deepinfra/fp8"
|
||||
)
|
||||
model_path = encode_model_path(base_url, "cert-test-model", "deepinfra/fp8")
|
||||
integration_session.add(
|
||||
ModelPathRow(
|
||||
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, "usage.capture")["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
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
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
|
||||
return create_test_model(
|
||||
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]
|
||||
)
|
||||
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(
|
||||
|
||||
@@ -11,7 +11,12 @@ import pytest
|
||||
import respx
|
||||
|
||||
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 (
|
||||
CacheProbeResult,
|
||||
_raw_cache_keys,
|
||||
@@ -210,6 +215,8 @@ class TestCacheBillingRow:
|
||||
model=_model(),
|
||||
probe=_probe([_payload(UNCACHED), cached]),
|
||||
cost_data=_cost(100),
|
||||
provider_fee=1.0,
|
||||
sats_to_usd=SATS_USD,
|
||||
)
|
||||
assert row["status"] == STATUS_OK
|
||||
assert row["evidence"]["reported_usd"] == 5e-5
|
||||
@@ -242,12 +249,21 @@ class TestCostMarginRow:
|
||||
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 1e-3})
|
||||
row = self._row([payload])
|
||||
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:
|
||||
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=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:
|
||||
sats_usd = 0.0008616302499999999
|
||||
@@ -332,6 +348,8 @@ class TestCostMarginRow:
|
||||
)
|
||||
|
||||
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:
|
||||
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7})
|
||||
@@ -373,7 +391,7 @@ class TestProbeCache:
|
||||
assert first == second
|
||||
system = first["messages"][0]["content"]
|
||||
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"
|
||||
|
||||
@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)
|
||||
@@ -13,6 +13,7 @@ import httpx
|
||||
import pytest
|
||||
|
||||
from routstr.upstream.certification import (
|
||||
PROBE_MAX_TOKENS,
|
||||
STATUS_OK,
|
||||
certify_upstream_url,
|
||||
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
|
||||
# accepted field instead of being rejected again.
|
||||
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
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user