fix(certification): harden against adversarial inputs found by tester agents

Three independent tester subagents found 12 defects (1 critical, 2 high,
9 medium/low). This commit fixes all of them and adds 56 regression tests.

Critical (CLI dead on arrival):
- certify_upstream_url() called sats_usd_price() which raises ValueError
  in any fresh process (the module global is only set by the app lifespan
  task). Now resolves via _resolve_sats_usd_price(): module global → BTC
  global → exchange feed → None (warn row, not a crash). Adds
  --sats-usd-price CLI flag for explicit override.

High (non-finite tokens crash the billing path):
- parse_token_count() crashed on Infinity/NaN (json.loads accepts both).
  Fixed to reject non-finite values → 0. This was a shared-code bug in
  routstr/payment/usage.py, reachable from the main billing path too.
- usage_capture_row and cost_prompt_completion_row now guard normalize_usage
  in try/except via safe_row(), so a raising check becomes a fail row
  instead of a 500.

Medium:
- certification_row now coerces non-dict evidence to {} (was stored verbatim)
- endpoint_validity_row checks .hostname not .netloc (http://:8080 rejected)
- models_payload_row rejects empty-string ids (agrees with CLI discovery)
- models_payload_row / usage_capture_row guard against non-dict payloads
- _reported_usd_cost uses coerce_rate for parity with the engine
- _expected_token_msats raises ValueError on non-finite rates (not OverflowError)
- get_candidates() call in admin endpoint wrapped in try/except
- Admin timeout clamped to [1, 60] seconds
- Explicit --prompt-price validated via coerce_rate (negatives rejected)
- CLI --json-out flag writes strictly parseable JSON to a file
- CLI logs routed to stderr so stdout is the report's channel

Test plan: 1749 passed, 1 skipped (no regressions); ruff clean; mypy clean.
This commit is contained in:
9qeklajc
2026-09-21 22:53:33 +02:00
committed by 9qeklajc
parent c5e4f4caee
commit 7f44389b59
4 changed files with 691 additions and 63 deletions
+22 -3
View File
@@ -1529,6 +1529,8 @@ async def certify_upstream_provider(
"""
from ..payment.price import sats_usd_price
from ..upstream.certification import (
MAX_PROBE_TIMEOUT_SECONDS,
PROBE_TIMEOUT_SECONDS,
build_checklist,
run_live_checks,
)
@@ -1570,7 +1572,19 @@ async def certify_upstream_provider(
model_obj = None
if model_id:
for model, _upstream in get_candidates(model_id) or []:
try:
candidates = get_candidates(model_id) or []
except Exception as exc: # noqa: BLE001 - a broken served map is a warn
logger.warning(
"Could not read the served map for certification",
extra={
"provider_id": provider.id,
"model_id": model_id,
"error": f"{type(exc).__name__}: {exc}",
},
)
candidates = []
for model, _upstream in candidates:
if model.upstream_provider_id == provider_pk:
model_obj = model
break
@@ -1620,9 +1634,14 @@ async def certify_upstream_provider(
]
else:
sats_to_usd = sats_usd_price()
timeout = (
payload.timeout_seconds if payload.timeout_seconds is not None else 15.0
# Clamp the admin-supplied timeout: the probe must never be able to
# hold the request open indefinitely.
requested = (
payload.timeout_seconds
if payload.timeout_seconds is not None
else PROBE_TIMEOUT_SECONDS
)
timeout = min(max(requested, 1.0), MAX_PROBE_TIMEOUT_SECONDS)
live_rows = await run_live_checks(
provider.base_url,
provider.api_key,
+15 -4
View File
@@ -38,6 +38,8 @@ names do not collide, so a single union parser is safe; a vendor whose fields
would genuinely conflict needs a dedicated branch here.
"""
import math
from pydantic.v1 import BaseModel
@@ -51,18 +53,27 @@ class NormalizedUsage(BaseModel):
def parse_token_count(value: object) -> int:
"""Parse a token count from various formats (int, float, str, bool)."""
"""Parse a token count from various formats (int, float, str, bool).
A non-finite count is not a count. ``json.loads`` accepts the bare
``Infinity``/``NaN`` literals and overflows ``1e999`` to ``inf``, so an
upstream — or an attacker who controls one — can put them on the wire.
``int(inf)`` raises ``OverflowError`` and ``int(nan)`` raises
``ValueError``; either would turn a billing path into a 500. Same rule as
``is_usable_rate``: reject the value, do not crash on it.
"""
if isinstance(value, bool):
return 0
if isinstance(value, int):
return max(0, value)
if isinstance(value, float):
return max(0, int(value))
return max(0, int(value)) if math.isfinite(value) else 0
if isinstance(value, str):
try:
return max(0, int(float(value)))
except ValueError:
parsed = float(value)
except (ValueError, OverflowError):
return 0
return max(0, int(parsed)) if math.isfinite(parsed) else 0
return 0
+266 -56
View File
@@ -36,7 +36,9 @@ import argparse
import asyncio
import json
import math
import sys
import time
from collections.abc import Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from urllib.parse import urlparse
@@ -45,6 +47,7 @@ import httpx
from ..core.logging import get_logger
from ..payment.cost_calculation import calculate_cost
from ..payment.rates import coerce_rate
from ..payment.usage import normalize_usage
if TYPE_CHECKING:
@@ -64,6 +67,11 @@ TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"}
# the request.
PROBE_TIMEOUT_SECONDS = 15.0
# An upper bound for a caller-supplied timeout. The admin endpoint accepts a
# timeout override, and without a ceiling that override could hold the
# request open for as long as the caller likes.
MAX_PROBE_TIMEOUT_SECONDS = 60.0
# The cheapest request that still exercises the usage/cost path: one token
# out. Anything larger only spends more upstream credit for no extra
# signal.
@@ -88,16 +96,50 @@ def certification_row(
detail: str,
evidence: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Build one row of the certification report."""
"""Build one row of the certification report.
``evidence`` is coerced to a dict so the row contract holds by
construction rather than by caller discipline — a caller that passes a
list or a string still produces a row a client can read.
"""
return {
"id": row_id,
"status": status,
"title": title,
"detail": detail,
"evidence": evidence if evidence is not None else {},
"evidence": evidence if isinstance(evidence, dict) else {},
}
def safe_row(
row_id: str,
title: str,
builder: Callable[[], dict[str, Any]],
) -> dict[str, Any]:
"""Run a row builder, turning any raise into a ``fail`` row.
The report is the diagnostic; it must never be the thing that fails. A
builder tripping over a hostile payload — a non-finite count, a body of
the wrong shape — becomes a ``fail`` row carrying the exception instead
of escaping the endpoint as a 500.
"""
try:
return builder()
except Exception as exc: # noqa: BLE001 - a raising check is a row status
described = f"{type(exc).__name__}: {exc}"
logger.warning(
"Certification check raised",
extra={"row_id": row_id, "error": described},
)
return certification_row(
row_id,
STATUS_FAIL,
title,
f"The {row_id} check could not run: {described}.",
{"error": described},
)
# The operator-facing goals, each mapped onto the rows that decide it. A
# goal is ``ok`` only when every row it names is ``ok``; any ``fail`` makes
# it ``fail``; anything else (a ``warn``, or a row that did not run) makes
@@ -269,12 +311,16 @@ def endpoint_validity_row(base_url: str) -> dict[str, Any]:
problems: list[str] = []
if parsed.scheme not in ("http", "https"):
problems.append(f"scheme {parsed.scheme!r} is not http or https")
if not parsed.netloc:
# ``netloc`` is truthy for a hostless authority like ``http://:8080``
# (``.netloc == ':8080'``) even though there is no host to connect to —
# only ``.hostname`` answers "is there a host here".
if not parsed.hostname:
problems.append("no host component")
evidence: dict[str, Any] = {
"base_url": base_url,
"scheme": parsed.scheme,
"host": parsed.netloc,
"host": parsed.hostname,
"port": parsed.port,
"path": parsed.path,
}
if problems:
@@ -332,17 +378,18 @@ def heartbeat_row(probe: ProbeResult) -> dict[str, Any]:
def models_payload_row(probe: ProbeResult) -> dict[str, Any]:
"""Check the ``/models`` payload matches the OpenAI list shape."""
if probe.models_payload is None:
payload = probe.models_payload
if not isinstance(payload, dict):
return certification_row(
"endpoint.models_payload",
STATUS_FAIL,
"Models payload has the expected shape",
f"Could not read a JSON object from {probe.models_url}: "
f"{probe.models_error}.",
f"{probe.models_error or type(payload).__name__}.",
{"url": probe.models_url, "error": probe.models_error},
)
data = probe.models_payload.get("data")
data = payload.get("data")
if not isinstance(data, list):
return certification_row(
"endpoint.models_payload",
@@ -351,14 +398,16 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]:
f'Expected a top-level "data" list, got {type(data).__name__}.',
{
"url": probe.models_url,
"top_level_keys": sorted(probe.models_payload.keys()),
"top_level_keys": sorted(payload.keys()),
},
)
# An empty id is not an id — the CLI discovery path refuses it, so the
# row must not certify it either.
ids = [
item.get("id")
item["id"]
for item in data
if isinstance(item, dict) and isinstance(item.get("id"), str)
if isinstance(item, dict) and isinstance(item.get("id"), str) and item["id"]
]
evidence: dict[str, Any] = {
"url": probe.models_url,
@@ -371,7 +420,7 @@ 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 string "id" '
f'The "data" list carries no entry with a non-empty string "id" '
f"({len(data)} entries).",
evidence,
)
@@ -416,18 +465,31 @@ def usage_capture_row(probe: ProbeResult) -> dict[str, Any]:
f"{PROBE_MAX_TOKENS}-token probe.",
evidence,
)
if probe.chat_payload is None:
if probe.chat_payload is None or not isinstance(probe.chat_payload, dict):
evidence["error"] = probe.chat_error
return certification_row(
"usage.capture",
STATUS_FAIL,
"Token usage captured from a completion",
f"The completion body was not a JSON object: {probe.chat_error}.",
f"The completion body was not a JSON object: "
f"{probe.chat_error or type(probe.chat_payload).__name__}.",
evidence,
)
raw_usage = probe.chat_payload.get("usage")
normalized = normalize_usage(raw_usage)
try:
normalized = normalize_usage(raw_usage)
except Exception as exc: # noqa: BLE001 - a malformed usage object is a row status
evidence["usage"] = _truncate(raw_usage)
evidence["error"] = f"{type(exc).__name__}: {exc}"
return certification_row(
"usage.capture",
STATUS_FAIL,
"Token usage captured from a completion",
f"The completion's usage object could not be read: "
f"{type(exc).__name__}: {exc}.",
evidence,
)
evidence["usage"] = raw_usage
if normalized is None:
return certification_row(
@@ -472,24 +534,26 @@ def _reported_usd_cost(payload: dict[str, Any]) -> float:
Mirrors ``_resolve_usd_cost``'s priority (``cost_details.total_cost``
then ``total_cost`` then ``cost``) so this check knows which branch of
the engine it is verifying. It is written out here rather than imported
on purpose: the point of the cost row is an independent re-derivation,
and reusing the engine's own helper would make a wrong priority
self-consistent and therefore invisible.
the engine it is verifying. Coercion goes through the shared
``coerce_rate`` — the one definition of what an upstream-supplied
number is — so this helper and the engine agree on *whether* a cost was
reported; only the arithmetic below is re-derived independently. Using
a private coercion here would disagree with the engine on numeric
strings and booleans and manufacture false failures.
"""
usage = payload.get("usage")
if not isinstance(usage, dict):
return 0.0
cost_details = usage.get("cost_details")
if isinstance(cost_details, dict):
total = cost_details.get("total_cost")
if isinstance(total, (int, float)) and math.isfinite(total) and total > 0:
return float(total)
total = coerce_rate(cost_details.get("total_cost"))
if total is not None and total > 0:
return total
for source in (usage, payload):
for field in ("total_cost", "cost"):
value = source.get(field)
if isinstance(value, (int, float)) and math.isfinite(value) and value > 0:
return float(value)
value = coerce_rate(source.get(field))
if value is not None and value > 0:
return value
return 0.0
@@ -504,6 +568,11 @@ def _expected_token_msats(sats_pricing: Any, usage: Any) -> tuple[int, int, int]
cache term or a changed rounding rule shows up as a mismatch.
Returns ``(total_msats, input_msats, output_msats)``.
Raises ``ValueError`` when a rate is not finite: ``math.ceil`` on an
infinite sum raises ``ValueError`` and on ``NaN`` produces an
unrepresentable result, so a non-finite rate is rejected explicitly
here rather than surfacing as an opaque crash.
"""
input_rate = float(sats_pricing.prompt) * 1_000_000.0
output_rate = float(sats_pricing.completion) * 1_000_000.0
@@ -514,6 +583,10 @@ def _expected_token_msats(sats_pricing: Any, usage: Any) -> tuple[int, int, int]
float(sats_pricing.input_cache_write or 0.0) * 1_000_000.0 or input_rate
)
rates = (input_rate, output_rate, cache_read_rate, cache_write_rate)
if not all(math.isfinite(rate) for rate in rates):
raise ValueError(f"non-finite pricing rate in {rates!r}")
calc_input = round(usage.input_tokens / 1000 * input_rate, 3)
calc_output = round(usage.output_tokens / 1000 * output_rate, 3)
calc_cache_read = round(usage.cache_read_tokens / 1000 * cache_read_rate, 3)
@@ -528,6 +601,10 @@ def _expected_usd_msats(
reported_usd: float, provider_fee: float, sats_to_usd: float
) -> int:
"""Re-derive the upstream-reported-USD charge, fee applied then converted."""
if not all(math.isfinite(x) for x in (reported_usd, provider_fee, sats_to_usd)):
raise ValueError("non-finite input to the USD charge derivation")
if sats_to_usd <= 0:
raise ValueError("sats/USD price must be positive")
return math.ceil(reported_usd * provider_fee / sats_to_usd * 1000)
@@ -549,8 +626,11 @@ def cost_prompt_completion_row(
"""
from ..payment.cost_calculation import CostDataError
payload = probe.chat_payload or {}
usage = normalize_usage(payload.get("usage"))
payload = probe.chat_payload if isinstance(probe.chat_payload, dict) else {}
try:
usage = normalize_usage(payload.get("usage"))
except Exception: # noqa: BLE001 - a malformed usage object is a row status
usage = None
evidence: dict[str, Any] = {
"model_id": model.id,
"forwarded_model_id": model.forwarded_model_id,
@@ -596,16 +676,30 @@ def cost_prompt_completion_row(
)
reported_usd = _reported_usd_cost(payload)
if reported_usd > 0:
expected_total = _expected_usd_msats(reported_usd, provider_fee, sats_to_usd)
expected_input: int | None = None
expected_output: int | None = None
basis = "upstream_reported_usd"
else:
expected_total, expected_input, expected_output = _expected_token_msats(
model.sats_pricing, usage
try:
if reported_usd > 0:
expected_total = _expected_usd_msats(
reported_usd, provider_fee, sats_to_usd
)
expected_input: int | None = None
expected_output: int | None = None
basis = "upstream_reported_usd"
else:
expected_total, expected_input, expected_output = _expected_token_msats(
model.sats_pricing, usage
)
basis = "configured_token_pricing"
except (ValueError, OverflowError) as exc:
evidence["error"] = f"{type(exc).__name__}: {exc}"
evidence["reported_usd"] = reported_usd or None
return certification_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}.",
evidence,
)
basis = "configured_token_pricing"
actual_total = int(cost_data.total_msats)
actual_input = int(cost_data.input_msats)
@@ -688,10 +782,24 @@ async def run_live_checks(
timeout=timeout,
)
rows = [
endpoint_validity_row(base_url),
heartbeat_row(probe),
models_payload_row(probe),
usage_capture_row(probe),
safe_row(
"endpoint.validity",
"Upstream URL is well-formed",
lambda: endpoint_validity_row(base_url),
),
safe_row(
"endpoint.reachable", "Endpoint responds", lambda: heartbeat_row(probe)
),
safe_row(
"endpoint.models_payload",
"Models payload has the expected shape",
lambda: models_payload_row(probe),
),
safe_row(
"usage.capture",
"Token usage captured from a completion",
lambda: usage_capture_row(probe),
),
]
cost_data: Any = None
@@ -718,13 +826,17 @@ async def run_live_checks(
)
rows.append(
cost_prompt_completion_row(
model=model,
probe=probe,
cost_data=cost_data,
provider_fee=provider_fee,
sats_to_usd=sats_to_usd,
pricing_known=pricing_known,
safe_row(
"cost.prompt_completion",
"Prompt and completion cost calculated",
lambda: cost_prompt_completion_row(
model=model,
probe=probe,
cost_data=cost_data,
provider_fee=provider_fee,
sats_to_usd=sats_to_usd,
pricing_known=pricing_known,
),
)
)
return rows
@@ -756,8 +868,45 @@ def _first_model_id(probe: ProbeResult) -> str | None:
def _as_price(value: Any) -> float | None:
if isinstance(value, (int, float)) and math.isfinite(value) and value >= 0:
return float(value)
"""A USD-per-token price from outside the node, or ``None``.
Shares ``coerce_rate`` — the one definition of a usable rate — so an
explicit ``--prompt-price`` is validated exactly like a litellm-derived
one: a boolean, a negative or a non-finite value is not a price.
"""
return coerce_rate(value)
async def _resolve_sats_usd_price(override: float | None) -> float | None:
"""The sats/USD price for a standalone run, or ``None`` if unavailable.
``SATS_USD_PRICE`` is a module global populated by the app's lifespan
background task, so a fresh ``python -m`` process has none and
``sats_usd_price()`` raises ``ValueError``. That must not abort a
certification run: fall back to the BTC global, then try the exchange
feed once, and return ``None`` rather than raising so the cost row can
degrade to a ``warn`` and the rest of the report still prints.
"""
if override is not None:
return override if math.isfinite(override) and override > 0 else None
from ..payment import price as price_module
if price_module.SATS_USD_PRICE:
return float(price_module.SATS_USD_PRICE)
if price_module.BTC_USD_PRICE:
return float(price_module.BTC_USD_PRICE) / price_module.SATS_PER_BTC
try:
await price_module._update_prices()
except Exception as exc: # noqa: BLE001 - no price is a row status
logger.warning(
"Could not initialize the sats/USD price for the standalone run",
extra={"error": f"{type(exc).__name__}: {exc}"},
)
return None
if price_module.SATS_USD_PRICE:
return float(price_module.SATS_USD_PRICE)
return None
@@ -805,13 +954,13 @@ async def certify_upstream_url(
completion_price: float | None = None,
provider_fee: float = 1.0,
timeout: float = PROBE_TIMEOUT_SECONDS,
sats_usd_price: float | None = None,
client: httpx.AsyncClient | None = None,
) -> dict[str, Any]:
"""Certify an arbitrary upstream URL without touching the node's DB."""
from ..payment.models import litellm_cost_entry
from ..payment.price import sats_usd_price
sats_to_usd = sats_usd_price()
sats_to_usd = await _resolve_sats_usd_price(sats_usd_price)
target: dict[str, Any] = {"base_url": base_url, "model_id": model_id}
if not model_id:
@@ -849,28 +998,36 @@ async def certify_upstream_url(
entry = litellm_cost_entry(model_id) or {}
resolved_prompt = (
prompt_price
_as_price(prompt_price)
if prompt_price is not None
else _as_price(entry.get("input_cost_per_token"))
)
resolved_completion = (
completion_price
_as_price(completion_price)
if completion_price is not None
else _as_price(entry.get("output_cost_per_token"))
)
pricing_known = resolved_prompt is not None and resolved_completion is not None
pricing_known = (
resolved_prompt is not None
and resolved_completion is not None
and sats_to_usd is not None
)
target["prompt_price_usd"] = resolved_prompt
target["completion_price_usd"] = resolved_completion
target["sats_usd_price"] = sats_to_usd
model = _model_from_usd_pricing(
model_id, resolved_prompt or 0.0, resolved_completion or 0.0, sats_to_usd
model_id,
resolved_prompt or 0.0,
resolved_completion or 0.0,
sats_to_usd or 1.0,
)
rows = await run_live_checks(
base_url,
api_key,
model,
provider_fee=provider_fee,
sats_to_usd=sats_to_usd,
sats_to_usd=sats_to_usd or 1.0,
client=client,
timeout=timeout,
pricing_known=pricing_known,
@@ -896,8 +1053,33 @@ def render_checklist(result: dict[str, Any]) -> str:
return "\n".join(lines)
def _route_logs_to_stderr() -> None:
"""Move the app's stdout log handlers to stderr.
``routstr.core.logging`` configures its handlers onto ``sys.stdout``, so
a machine-readable run would otherwise interleave log records with the
document. Stdout is the report's channel; logs belong on stderr.
"""
import logging
loggers = [logging.getLogger()]
loggers.extend(
obj
for obj in logging.root.manager.loggerDict.values()
if isinstance(obj, logging.Logger)
)
for logger in loggers:
for handler in list(logger.handlers):
if (
isinstance(handler, logging.StreamHandler)
and getattr(handler, "stream", None) is sys.stdout
):
handler.setStream(sys.stderr)
def main(argv: list[str] | None = None) -> int:
"""Run the checklist against one or more upstream base URLs."""
_route_logs_to_stderr()
parser = argparse.ArgumentParser(
prog="python -m routstr.upstream.certification",
description=(
@@ -935,6 +1117,15 @@ def main(argv: list[str] | None = None) -> int:
default=1.0,
help="Provider fee multiplier applied by the cost check",
)
parser.add_argument(
"--sats-usd-price",
type=float,
default=None,
help=(
"USD per satoshi for the cost check. Defaults to the node's "
"live rate, initialized from the exchange feed when unset."
),
)
parser.add_argument(
"--timeout",
type=float,
@@ -944,6 +1135,16 @@ def main(argv: list[str] | None = None) -> int:
parser.add_argument(
"--json", action="store_true", help="Emit the raw report as JSON"
)
parser.add_argument(
"--json-out",
default=None,
metavar="PATH",
help=(
"Write the raw JSON report to PATH ('-' for stdout). Unlike "
"--json, nothing else is written there, so the file is always "
"parseable — use this in pipelines."
),
)
args = parser.parse_args(argv)
async def _run_all() -> list[dict[str, Any]]:
@@ -958,15 +1159,24 @@ def main(argv: list[str] | None = None) -> int:
completion_price=args.completion_price,
provider_fee=args.provider_fee,
timeout=args.timeout,
sats_usd_price=args.sats_usd_price,
)
)
return results
results = asyncio.run(_run_all())
if args.json_out is not None:
document = json.dumps(results, indent=2, default=str)
if args.json_out == "-":
print(document)
else:
with open(args.json_out, "w", encoding="utf-8") as handle:
handle.write(document + "\n")
if args.json:
print(json.dumps(results, indent=2, default=str))
else:
elif args.json_out is None:
for result in results:
print(render_checklist(result))
print()
+388
View File
@@ -0,0 +1,388 @@
"""Regression tests for defects found by independent adversarial testing.
Each test here pins a specific failure mode that was found and fixed while
building the certification harness. They are grouped by the defect they
guard, not by the function under test, because the point of each one is the
bug it prevents from coming back.
"""
from __future__ import annotations
import json
import os
import subprocess
import sys
from pathlib import Path
from typing import Any
import pytest
from routstr.payment.usage import parse_token_count
from routstr.upstream.certification import (
STATUS_FAIL,
STATUS_OK,
STATUS_WARN,
ProbeResult,
_as_price,
_expected_token_msats,
_reported_usd_cost,
certification_row,
endpoint_validity_row,
models_payload_row,
safe_row,
usage_capture_row,
)
REPO_ROOT = Path(__file__).resolve().parents[2]
def _probe(**kwargs: Any) -> ProbeResult:
return ProbeResult(
base_url=kwargs.get("base_url", "https://upstream.example/v1"),
models_url=kwargs.get("models_url", "https://upstream.example/v1/models"),
chat_url=kwargs.get("chat_url", "https://upstream.example/v1/chat/completions"),
models_status=kwargs.get("models_status"),
models_payload=kwargs.get("models_payload"),
models_error=kwargs.get("models_error"),
chat_status=kwargs.get("chat_status"),
chat_payload=kwargs.get("chat_payload"),
chat_error=kwargs.get("chat_error"),
)
# ---------------------------------------------------------------------------
# Defect: a non-finite token count crashed the billing path.
#
# ``json.loads`` accepts the bare ``Infinity``/``NaN`` literals, so an
# upstream can put them on the wire; ``int(inf)`` raised OverflowError and
# ``int(nan)`` raised ValueError inside ``parse_token_count``.
# ---------------------------------------------------------------------------
class TestNonFiniteTokenCounts:
@pytest.mark.parametrize(
"value",
[
float("inf"),
float("-inf"),
float("nan"),
1e999,
"Infinity",
"NaN",
"-Infinity",
"1e999",
],
)
def test_parse_token_count_rejects_non_finite(self, value: Any) -> None:
assert parse_token_count(value) == 0
def test_parse_token_count_still_parses_ordinary_values(self) -> None:
assert parse_token_count(42) == 42
assert parse_token_count("42") == 42
assert parse_token_count(42.9) == 42
assert parse_token_count("42.9") == 42
assert parse_token_count(True) == 0
assert parse_token_count(-5) == 0
assert parse_token_count("not a number") == 0
assert parse_token_count(None) == 0
def test_usage_row_survives_infinite_tokens(self) -> None:
row = usage_capture_row(
_probe(
chat_status=200,
chat_payload={
"usage": {
"prompt_tokens": float("inf"),
"completion_tokens": float("nan"),
}
},
)
)
# Both counts collapse to 0, which is the "nothing to bill on" case.
assert row["status"] == STATUS_WARN
def test_usage_row_survives_infinite_tokens_in_a_string(self) -> None:
row = usage_capture_row(
_probe(
chat_status=200,
chat_payload={"usage": {"prompt_tokens": "Infinity"}},
)
)
assert row["status"] == STATUS_WARN
# ---------------------------------------------------------------------------
# Defect: ``certification_row`` stored non-dict evidence verbatim, so the
# row contract ("evidence is always a dict") held only by caller discipline.
# ---------------------------------------------------------------------------
class TestEvidenceContract:
@pytest.mark.parametrize("evidence", [None, [1, 2], "text", 42, (1, 2)])
def test_evidence_is_always_a_dict(self, evidence: Any) -> None:
row = certification_row("x", STATUS_OK, "t", "d", evidence)
assert isinstance(row["evidence"], dict)
def test_evidence_dict_is_passed_through(self) -> None:
row = certification_row("x", STATUS_OK, "t", "d", {"a": 1})
assert row["evidence"] == {"a": 1}
# ---------------------------------------------------------------------------
# Defect: ``http://:8080/v1`` was certified as a valid endpoint because
# ``netloc`` is truthy for a hostless authority.
# ---------------------------------------------------------------------------
class TestEndpointValidity:
@pytest.mark.parametrize(
"url",
["http://:8080/v1", "https://:443", "http://", "https://"],
)
def test_hostless_authority_is_rejected(self, url: str) -> None:
row = endpoint_validity_row(url)
assert row["status"] == STATUS_FAIL, url
@pytest.mark.parametrize(
"url",
[
"https://api.example.com/v1",
"http://localhost:8888/v1",
"http://127.0.0.1:8080",
"https://[::1]:8080/v1",
],
)
def test_real_hosts_are_accepted(self, url: str) -> None:
row = endpoint_validity_row(url)
assert row["status"] == STATUS_OK, url
# ---------------------------------------------------------------------------
# Defect: the payload builders called ``.get()`` on whatever they were
# given, so a wrong-typed body raised AttributeError instead of producing a
# verdict.
# ---------------------------------------------------------------------------
class TestPayloadTypeGuards:
@pytest.mark.parametrize("payload", [[1, 2], "text", 42, ("a",)])
def test_models_payload_row_handles_non_dict(self, payload: Any) -> None:
row = models_payload_row(_probe(models_payload=payload))
assert row["status"] == STATUS_FAIL
@pytest.mark.parametrize("payload", [[1, 2], "text", 42, ("a",)])
def test_usage_row_handles_non_dict_chat_payload(self, payload: Any) -> None:
row = usage_capture_row(_probe(chat_status=200, chat_payload=payload))
assert row["status"] == STATUS_FAIL
def test_models_payload_with_non_dict_entries(self) -> None:
row = models_payload_row(
_probe(models_payload={"data": [None, 42, "string", {}]})
)
assert row["status"] == STATUS_FAIL
assert row["evidence"]["model_count"] == 4
assert row["evidence"]["usable_ids"] == 0
# ---------------------------------------------------------------------------
# Defect: an empty-string id was counted as "usable" by the payload row but
# rejected by the CLI's discovery path — the two disagreed on one response.
# ---------------------------------------------------------------------------
class TestModelIdAgreement:
def test_empty_string_id_is_not_usable(self) -> None:
row = models_payload_row(_probe(models_payload={"data": [{"id": ""}]}))
assert row["status"] == STATUS_FAIL
assert row["evidence"]["usable_ids"] == 0
def test_one_usable_id_among_empties_is_ok(self) -> None:
row = models_payload_row(
_probe(models_payload={"data": [{"id": ""}, {"id": "real-model"}]})
)
assert row["status"] == STATUS_OK
assert row["evidence"]["usable_ids"] == 1
# ---------------------------------------------------------------------------
# Defect: the independent cost re-derivation disagreed with the engine on
# coercion (numeric strings, booleans), manufacturing false failures.
# ---------------------------------------------------------------------------
class TestReportedCostCoercionParity:
def test_numeric_string_cost_is_read(self) -> None:
assert _reported_usd_cost({"usage": {"cost": "0.001"}}) == pytest.approx(0.001)
def test_boolean_cost_is_rejected(self) -> None:
# ``True`` is an int in Python and would read as $1.00 per token.
assert _reported_usd_cost({"usage": {"cost": True}}) == 0.0
def test_non_finite_cost_is_rejected(self) -> None:
assert _reported_usd_cost({"usage": {"cost": float("inf")}}) == 0.0
assert _reported_usd_cost({"usage": {"cost": float("nan")}}) == 0.0
def test_negative_cost_is_rejected(self) -> None:
assert _reported_usd_cost({"usage": {"cost": -1.0}}) == 0.0
def test_cost_details_wins_over_cost(self) -> None:
payload = {"usage": {"cost": 0.5, "cost_details": {"total_cost": 0.001}}}
assert _reported_usd_cost(payload) == pytest.approx(0.001)
# ---------------------------------------------------------------------------
# Defect: ``_expected_token_msats`` ran ``math.ceil`` on a non-finite sum,
# raising an opaque error instead of a describable one.
# ---------------------------------------------------------------------------
class TestNonFinitePricing:
class _SatsPricing:
def __init__(self, **kwargs: Any) -> None:
self.prompt = kwargs.get("prompt", 1.0)
self.completion = kwargs.get("completion", 1.0)
self.input_cache_read = kwargs.get("input_cache_read", 0.0)
self.input_cache_write = kwargs.get("input_cache_write", 0.0)
class _Usage:
input_tokens = 10
output_tokens = 5
cache_read_tokens = 0
cache_write_tokens = 0
def test_infinite_rate_raises_value_error(self) -> None:
with pytest.raises(ValueError, match="non-finite"):
_expected_token_msats(self._SatsPricing(prompt=float("inf")), self._Usage())
def test_nan_rate_raises_value_error(self) -> None:
with pytest.raises(ValueError, match="non-finite"):
_expected_token_msats(
self._SatsPricing(completion=float("nan")), self._Usage()
)
def test_finite_rates_still_work(self) -> None:
total, inp, outp = _expected_token_msats(
self._SatsPricing(prompt=1.4e-7, completion=2.8e-7), self._Usage()
)
assert total == inp + outp
# ---------------------------------------------------------------------------
# Defect: a row builder raising escaped as a 500 from the admin endpoint.
# ---------------------------------------------------------------------------
class TestSafeRow:
def test_a_raising_builder_becomes_a_fail_row(self) -> None:
def boom() -> dict[str, Any]:
raise RuntimeError("hostile payload")
row = safe_row("x.row", "Title", boom)
assert row["status"] == STATUS_FAIL
assert "RuntimeError" in row["detail"]
assert isinstance(row["evidence"], dict)
def test_a_working_builder_passes_through(self) -> None:
row = safe_row(
"x.row", "Title", lambda: certification_row("x.row", STATUS_OK, "T", "D")
)
assert row["status"] == STATUS_OK
# ---------------------------------------------------------------------------
# Defect: explicit ``--prompt-price`` bypassed validation, so a negative
# rate could be fed into the cost engine.
# ---------------------------------------------------------------------------
class TestExplicitPriceValidation:
def test_negative_price_is_rejected(self) -> None:
assert _as_price(-1.0) is None
def test_non_finite_price_is_rejected(self) -> None:
assert _as_price(float("inf")) is None
assert _as_price(float("nan")) is None
def test_boolean_price_is_rejected(self) -> None:
assert _as_price(True) is None
def test_zero_is_a_valid_price(self) -> None:
# Free is a real price.
assert _as_price(0.0) == 0.0
def test_numeric_string_is_accepted(self) -> None:
assert _as_price("1e-7") == pytest.approx(1e-7)
# ---------------------------------------------------------------------------
# Defect: the standalone CLI was dead on arrival — ``sats_usd_price()``
# raises in a fresh process because the module global is only populated by
# the app's lifespan task. These run the CLI as a subprocess so the fresh
# process is the thing under test.
# ---------------------------------------------------------------------------
def _run_cli(*args: str, timeout: float = 90.0) -> subprocess.CompletedProcess[str]:
env = dict(os.environ)
env.setdefault("ROUTSTR_SECRET_KEY", "l_Tkp-7xmjcQ-IFhr6qhILrU8HPRbEmYMrfSbo_5srU=")
return subprocess.run(
[sys.executable, "-m", "routstr.upstream.certification", *args],
cwd=REPO_ROOT,
env=env,
capture_output=True,
text=True,
timeout=timeout,
)
@pytest.mark.slow
class TestCliFreshProcess:
def test_cli_emits_a_report_instead_of_dying(self, tmp_path: Path) -> None:
"""The regression: this used to raise 'SATS price not initialized'."""
out = tmp_path / "report.json"
result = _run_cli(
"--url",
"http://localhost:1/v1",
"--timeout",
"1",
"--json-out",
str(out),
)
assert "SATS price not initialized" not in result.stderr
assert out.exists(), result.stderr
def test_json_out_is_strictly_parseable(self, tmp_path: Path) -> None:
out = tmp_path / "report.json"
_run_cli(
"--url", "http://localhost:1/v1", "--timeout", "1", "--json-out", str(out)
)
document = json.loads(out.read_text(encoding="utf-8"))
assert isinstance(document, list) and document
assert len(document[0]["rows"]) == 5
assert len(document[0]["checklist"]) == 4
def test_dead_host_exits_non_zero(self) -> None:
result = _run_cli("--url", "http://localhost:1/v1", "--timeout", "1")
assert result.returncode == 1
def test_negative_explicit_price_is_rejected(self, tmp_path: Path) -> None:
out = tmp_path / "report.json"
_run_cli(
"--url",
"http://localhost:1/v1",
"--timeout",
"1",
"--model",
"m",
"--prompt-price",
"-1",
"--json-out",
str(out),
)
document = json.loads(out.read_text(encoding="utf-8"))
assert document[0]["target"]["prompt_price_usd"] is None
def test_checklist_uses_the_documented_ticks(self) -> None:
result = _run_cli("--url", "http://localhost:1/v1", "--timeout", "1")
assert "❌" in result.stdout
assert "Heartbeat" in result.stdout