diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 8d95d936..cb4c3af6 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -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, diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index 08d675ad..7bee961b 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -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 diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index 6952a698..b36e9e46 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -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() diff --git a/tests/unit/test_certification_hardening.py b/tests/unit/test_certification_hardening.py new file mode 100644 index 00000000..ae707448 --- /dev/null +++ b/tests/unit/test_certification_hardening.py @@ -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