mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
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:
+22
-3
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user